{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Introduction\nThis notebook implements sliding window inference. If you find this notebook useful, feel free to copy and customize it to suit your needs. And don't forget to upvote 🥰🥰","metadata":{}},{"cell_type":"markdown","source":"## Import Libraries","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport tifffile as tiff","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-17T03:17:41.944140Z","iopub.execute_input":"2022-07-17T03:17:41.945129Z","iopub.status.idle":"2022-07-17T03:17:42.110403Z","shell.execute_reply.started":"2022-07-17T03:17:41.945003Z","shell.execute_reply":"2022-07-17T03:17:42.109413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Configuration\nWe will be using window size of 512x512 with half overlap. You can change importance map to say gaussian center weighted.","metadata":{}},{"cell_type":"code","source":"'''\n    Configuration\n'''\nWINDOW_SIZE    = (512, 512)\nSTRIDE         = (256, 256) # Half overlap\nNUM_CLASSES    = 1\nIMPORTANCE_MAP = np.ones((*WINDOW_SIZE, NUM_CLASSES), dtype='float32') # Treat every pixel equally. You can change this to, for example, gaussian center weighted map","metadata":{"execution":{"iopub.status.busy":"2022-07-17T03:31:12.821116Z","iopub.execute_input":"2022-07-17T03:31:12.821455Z","iopub.status.idle":"2022-07-17T03:31:12.827322Z","shell.execute_reply.started":"2022-07-17T03:31:12.821425Z","shell.execute_reply":"2022-07-17T03:31:12.826204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create dummy model\nOur sliding window inference can work with multiple models at once","metadata":{}},{"cell_type":"code","source":"class DummyModel():\n    def __init__(self, num_classes):\n        self.num_classes = num_classes\n    \n    # Make random predictions\n    def predict(self, X):\n        return np.random.rand(*X.shape[:-1], self.num_classes)","metadata":{"execution":{"iopub.status.busy":"2022-07-17T03:32:07.348910Z","iopub.execute_input":"2022-07-17T03:32:07.349413Z","iopub.status.idle":"2022-07-17T03:32:07.361178Z","shell.execute_reply.started":"2022-07-17T03:32:07.349372Z","shell.execute_reply":"2022-07-17T03:32:07.359850Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create 5 models\nmodels = [DummyModel(NUM_CLASSES)] * 5\nmodels[0].predict(np.random.rand(4, 3000, 3000, 3)).shape","metadata":{"execution":{"iopub.status.busy":"2022-07-17T03:32:16.435059Z","iopub.execute_input":"2022-07-17T03:32:16.435390Z","iopub.status.idle":"2022-07-17T03:32:17.644919Z","shell.execute_reply.started":"2022-07-17T03:32:16.435361Z","shell.execute_reply":"2022-07-17T03:32:17.643837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Sliding Window Inference","metadata":{}},{"cell_type":"code","source":"def sliding_window_inference(X, models, num_class, window_size, stride, importance_map):\n        '''\n            Sliding window inference\n            --------\n            X : numpy.ndarray\n                Input 2D volume with shape = (Batch, Height, Width, Channel)\n            models: List \n                List of models with predict function\n            num_class: Integer\n                Number of output classes\n            window_size: Tuple\n                Window size\n            stride: Tuple\n                Stride\n            importance_map: numpy.ndarray\n                Patch pixel importance map\n            --------\n            return : numpy.ndarray\n                Output segmentations with shape = (Batch, Height Width, num_class) \n        '''\n        h, w = X.shape[1:3]\n        w_h, w_w = window_size\n        s_h, s_w = stride\n        \n        result = np.zeros((*X.shape[:-1], num_class), dtype='float32')\n        overlap = np.zeros((*X.shape[:-1], num_class), dtype='float32')\n        \n        # Generate a list of starting points a.k.a top left of window\n        starting_points = [(x, y)  for x in set( list(range(0, h - w_h, s_h)) + [h - w_h] ) \n                                   for y in set( list(range(0, w - w_w, s_w)) + [w - w_w] )]\n\n        # Get list of patches\n        patches = np.empty((len(starting_points), *WINDOW_SIZE, X.shape[-1]), dtype='float32')\n        for i, (x, y) in enumerate(starting_points):\n            patches[i] = X[:, x:x + w_h, y:y + w_w, :]\n\n        # Inference with each model\n        for model in models:\n            y_pred = model.predict(patches)\n            for i in range(len(y_pred)):\n                x, y = starting_points[i]\n                result[:, x:x + w_h, y:y + w_w, :] += y_pred[i] * importance_map # Multiply by importance map\n                overlap[:, x:x + w_h, y:y + w_w, :] += importance_map\n                \n        assert np.sum(overlap == 0.) == 0, \"Sliding window does not cover all volume\" # Something went wrong!\n\n        return result / overlap # Normalize output segmentation map","metadata":{"execution":{"iopub.status.busy":"2022-07-17T03:30:40.531899Z","iopub.execute_input":"2022-07-17T03:30:40.532236Z","iopub.status.idle":"2022-07-17T03:30:40.543824Z","shell.execute_reply.started":"2022-07-17T03:30:40.532207Z","shell.execute_reply":"2022-07-17T03:30:40.542652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load example image\nimage = tiff.imread(\"../input/hubmap-organ-segmentation/train_images/10044.tiff\")\nimage.shape","metadata":{"execution":{"iopub.status.busy":"2022-07-17T03:26:44.038349Z","iopub.execute_input":"2022-07-17T03:26:44.040782Z","iopub.status.idle":"2022-07-17T03:26:44.063759Z","shell.execute_reply.started":"2022-07-17T03:26:44.040739Z","shell.execute_reply":"2022-07-17T03:26:44.062898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create batch\nbatch = np.expand_dims(image, axis=0)\nprint(batch.shape)","metadata":{"execution":{"iopub.status.busy":"2022-07-17T03:28:00.754911Z","iopub.execute_input":"2022-07-17T03:28:00.755473Z","iopub.status.idle":"2022-07-17T03:28:00.761303Z","shell.execute_reply.started":"2022-07-17T03:28:00.755435Z","shell.execute_reply":"2022-07-17T03:28:00.760056Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Inference\nmask = sliding_window_inference(batch, models, num_class=NUM_CLASSES, window_size=WINDOW_SIZE, stride=STRIDE, importance_map=IMPORTANCE_MAP)\nprint(mask.shape)","metadata":{"execution":{"iopub.status.busy":"2022-07-17T03:32:21.269741Z","iopub.execute_input":"2022-07-17T03:32:21.270090Z","iopub.status.idle":"2022-07-17T03:32:23.389346Z","shell.execute_reply.started":"2022-07-17T03:32:21.270060Z","shell.execute_reply":"2022-07-17T03:32:23.388346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Hooray! Our sliding window works as expected!","metadata":{}},{"cell_type":"code","source":"print(\"Done!\")","metadata":{},"execution_count":null,"outputs":[]}]}