{"metadata":{"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"}],"dockerImageVersionId":30646,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"pytorch_gpu","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.7.16"},"papermill":{"default_parameters":{},"duration":34779.689211,"end_time":"2024-01-15T06:22:09.132326","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-01-14T20:42:29.443115","version":"2.4.0"},"widgets":{"application/vnd.jupyter.widget-state+json":{"state":{"0fe72cc2bc1f4bf18dcd2cb619b0793f":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ProgressView","bar_style":"success","description":"","description_tooltip":null,"layout":"IPY_MODEL_5587c3ca80394c29a768b17bb1bec50e","max":110547680,"min":0,"orientation":"horizontal","style":"IPY_MODEL_d1d8e16f213540a3aae74b868d8db955","value":110547680}},"1c3436fb934544dba3efc117a0e37463":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"42db75fbc1f445d0ba18d57ab83e6c02":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_973200bb56c142e8aec7db76c99b1e82","placeholder":"​","style":"IPY_MODEL_6d712cb50fb04747ba6f6a73a23fbf47","value":" 111M/111M [00:00&lt;00:00, 243MB/s]"}},"46a63b6a748b4499bddfc2ee33e2deeb":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"5587c3ca80394c29a768b17bb1bec50e":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"6d712cb50fb04747ba6f6a73a23fbf47":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"973200bb56c142e8aec7db76c99b1e82":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"b64b10ba7e134b74a106b9b41c81c704":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"c96336b5d74849e78f312d497a83c145":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_ee51f9cfd0c54209bb3ee15dfca54124","IPY_MODEL_0fe72cc2bc1f4bf18dcd2cb619b0793f","IPY_MODEL_42db75fbc1f445d0ba18d57ab83e6c02"],"layout":"IPY_MODEL_b64b10ba7e134b74a106b9b41c81c704"}},"d1d8e16f213540a3aae74b868d8db955":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"ee51f9cfd0c54209bb3ee15dfca54124":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_46a63b6a748b4499bddfc2ee33e2deeb","placeholder":"​","style":"IPY_MODEL_1c3436fb934544dba3efc117a0e37463","value":"model.safetensors: 100%"}}},"version_major":2,"version_minor":0}}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport math\nimport glob\nimport gc\nimport tqdm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision\nimport torchvision.transforms as T\nimport torchvision.transforms.functional as TF\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom torch.utils.data import Dataset, DataLoader\nfrom fastai.vision.all import *\nfrom typing import Optional\nfrom torch.nn.functional import one_hot\nfrom sklearn.model_selection import KFold\nimport random\nimport skimage\n!pip install segmentation_models_pytorch\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.execute_input":"2024-01-14T20:42:32.902639Z","iopub.status.busy":"2024-01-14T20:42:32.901810Z","iopub.status.idle":"2024-01-14T20:43:02.190494Z","shell.execute_reply":"2024-01-14T20:43:02.189464Z"},"papermill":{"duration":29.297804,"end_time":"2024-01-14T20:43:02.192930","exception":false,"start_time":"2024-01-14T20:42:32.895126","status":"completed"},"tags":[],"collapsed":true,"jupyter":{"outputs_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ENCODER = \"tu-seresnext50_32x4d\"\nWEIGHTS = \"imagenet\"\nBATCH_SIZE = 8\nDEPTH = 1\nPATCH_SIZE = 512\nEPOCHS = 50\nLR = 5e-4\nTRAIN = ['C:/Users/Angel/VESSELS/kidney_1_dense',\n         'C:/Users/Angel/VESSELS/kidney_3_dense']\n#   '/kaggle/input/blood-vessel-segmentation/train/kidney_1_voi'\n#   '/kaggle/input/blood-vessel-segmentation/train/kidney_2'\n#   '/kaggle/input/blood-vessel-segmentation/train/kidney_3_dense'\nMODEL_NAME = 'SMP_2D_MIXED_' + ENCODER + '_' + WEIGHTS + '_' + str(PATCH_SIZE)\nSEED = 143\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.execute_input":"2024-01-14T20:43:02.210608Z","iopub.status.busy":"2024-01-14T20:43:02.209892Z","iopub.status.idle":"2024-01-14T20:43:02.267228Z","shell.execute_reply":"2024-01-14T20:43:02.266272Z"},"papermill":{"duration":0.068329,"end_time":"2024-01-14T20:43:02.269197","exception":false,"start_time":"2024-01-14T20:43:02.200868","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"execution":{"iopub.execute_input":"2024-01-14T20:43:03.727238Z","iopub.status.busy":"2024-01-14T20:43:03.726914Z","iopub.status.idle":"2024-01-14T20:43:03.732568Z","shell.execute_reply":"2024-01-14T20:43:03.731651Z"},"papermill":{"duration":0.016658,"end_time":"2024-01-14T20:43:03.734553","exception":false,"start_time":"2024-01-14T20:43:03.717895","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_volume(dataset, labeled=True, slice_range=None):\n    ''' Load slices into a volume. Keeps the memory requirement\n        as low as possible by using uint8 and uint16 in CPU memory.\n    '''\n    if labeled:\n        path = os.path.join(dataset, \"labels\", \"*.tif\")\n    else:\n        path = os.path.join(dataset, \"images\", \"*.tif\")\n        \n    dataset = sorted(glob.glob(path))\n    volume = None\n    target = None\n    keys = []\n    offset = 0 if slice_range is None else slice_range[0]\n    depth = len(dataset) if slice_range is None else slice_range[1]-slice_range[0]\n    \n    for z, path in enumerate(tqdm.tqdm(dataset)):\n        if slice_range is not None:\n            if z < slice_range[0]: continue\n            if z >= slice_range[1]: continue\n        \n        parts = path.split(os.path.sep)\n        key = parts[-3] + \"_\" + parts[-1].split(\".\")[0]\n        keys.append(key)\n                \n        if labeled:\n            label = cv2.imread(path, cv2.IMREAD_ANYDEPTH)#IMREAD_GRAYSCALE)#IMREAD_ANYDEPTH)\n            label = np.array(label,dtype=np.uint8)\n            if target is None:\n                target = np.zeros((1,depth, *label.shape[-2:]), dtype=np.uint8)\n            target[:,z-offset] = label\n        \n        path = path.replace(\"labels\",\"images\")\n        path = path.replace(\"kidney_3_dense\",\"kidney_3_sparse\")\n        image = cv2.imread(path, cv2.IMREAD_ANYDEPTH)#IMREAD_GRAYSCALE)#IMREAD_ANYDEPTH)\n        image = np.array(image,dtype=np.uint16)\n        \n        if volume is None:\n            volume = np.zeros((1,depth, *image.shape[-2:]), dtype=np.uint16)\n        volume[:,z-offset] = image\n    \n    return volume, target, keys","metadata":{"execution":{"iopub.execute_input":"2024-01-14T20:43:03.752043Z","iopub.status.busy":"2024-01-14T20:43:03.751704Z","iopub.status.idle":"2024-01-14T20:43:03.763956Z","shell.execute_reply":"2024-01-14T20:43:03.762992Z"},"papermill":{"duration":0.023474,"end_time":"2024-01-14T20:43:03.766073","exception":false,"start_time":"2024-01-14T20:43:03.742599","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\n# Multiple Datasets version\nclass Crops3D_Dataset(torch.utils.data.Dataset):\n    '''   \n    Dataset of 3D crops DEPTHxPATCH_SIZExPATCH_SIZE.\n    The crops will be centered at labels centroid.\n    The crops will be augmented on train mode.\n    The crops will be from each dataset successively.\n    The crops will be also augmented on less populated datasets\n    in order to match the more populated ones.\n    The augmentations will be different at every invocation.\n    The augmentations will be applied successively.\n    DEPTH_ZOOM_OUT: The amount of times that effective DEPTH \n    is applied to original crop and downscaled on returned crop.\n    PATCH_ZOOM_OUT: The amount of times that effective PATCH_SIZE \n    is applied to original crop and downscaled on returned crop.\n    data_augmentation: The portion of data augmentation to the largest dataset\n    sigma: The sigma applied to the normal distribution \n    deallocation of centroids durind data augmentation.\n    pfullrot: The probability to apply a random full rotation (any angle) during data augmentation.\n    '''  \n    def __init__(self, Datasets, DEPTH=64, PATCH_SIZE=256, DEPTH_ZOOM_OUT=1, PATCH_ZOOM_OUT=1,\n                       mode ='train', data_augmentation=0.2, sigma=1, pfullrot=.2):\n        if DEPTH_ZOOM_OUT > 1: DEPTH *= DEPTH_ZOOM_OUT\n        if PATCH_ZOOM_OUT > 1: PATCH_SIZE *= PATCH_ZOOM_OUT\n        self.pfullrot = pfullrot\n        self.DEPTH_ZOOM_OUT = DEPTH_ZOOM_OUT\n        self.PATCH_ZOOM_OUT = PATCH_ZOOM_OUT\n        self.PATCH = PATCH_SIZE\n        self.DEPTH = DEPTH\n        self.sigma = sigma\n        self.min = []\n        self.max = []\n        self.crops = []\n        self.volume = []\n        self.target = []\n        self.LEN = []\n        self.aug_crops = []\n        self.IDX = []\n        max_iter = 50\n        # Closer to patch limits convolutions have less true information\n        # In order to get a well centered crop the labels on those planes\n        # should count less for centroid calculation\n        alpha = -10/(DEPTH*DEPTH)\n        weights = np.arange(DEPTH) - DEPTH/2 + .5\n        weights = np.exp(alpha*weights*weights)\n        weights = weights/np.max(weights)\n        for Dataset in Datasets:\n            print(\"loading volume \" + Dataset.split('/')[-1] + \"...\")\n            volume,target,_ = load_volume(Dataset)\n            pmin,pmax = np.percentile(volume, (1, 99))\n            self.min.append(pmin)\n            self.max.append(pmax)\n            self.volume.append(volume)\n            self.target.append(target)\n#           axis 0 label, orthogonal to Z\n#           axis 1 label, orthogonal to X\n#           axis 2 label, orthogonal to Y\n            crops = []\n            SHAPE = list(target.shape[-3:])\n            for axis in [0,1,2]:\n                D = SHAPE[axis]\n                print(\"processing\",Dataset.split('/')[-1],\"axis\",axis,\"subvolumes...\")\n#               let's find the range without empty masks \n                d_min = 0\n                while len(np.argwhere(target[0].take(d_min,axis=axis) > 0)) == 0: d_min += 1\n                d_max = D - 1\n                while len(np.argwhere(target[0].take(d_max,axis=axis) > 0)) == 0: d_max -= 1\n#               let's loop over the inner axis values that will lead to a \n#               volume d-DEPTH/2:d+DEPTH/2 without empty masks\n                for d in tqdm.tqdm(range(d_min + DEPTH//2,d_max - DEPTH + DEPTH//2 + 1)):\n                    FOUND = False\n#                   let's condense all the depth labels on a plane\n#                   that plane will contain the count the total \n#                   of labels present on the depth of each pixel\n#                   PLANE = (target[0].take(range(d-DEPTH//2,d+DEPTH-DEPTH//2),axis=axis) > 0).swapaxes(0,axis).sum(0)\n                    d -= DEPTH//2\n                    SV = (target[0].take(range(d,d+DEPTH),axis=axis) > 0).swapaxes(0,axis)\n                    PLANE = np.average(SV,0,weights)\n                    nonzero = PLANE > 0\n                    total = np.sum(PLANE[nonzero])\n                    if total > 0:\n#                       let's rotate 90 degrees with respect prevoius subvolume\n#                       this way the consecutive crops will be more different\n#                       let's find the corresponding rot to d, 0 1 2 3 0 1 2 3 ...\n                        rot = int(4*(d/4-d//4))\n                        PLANE = np.rot90(PLANE,rot,[-2,-1])\n                        nonzero = np.rot90(nonzero,rot,[-2,-1])\n#                       let's find the centroid of those labels in that plane\n                        H,W = PLANE.shape\n                        indices = np.indices((H,W))\n                        h = int(np.sum(PLANE[nonzero]*indices[0,nonzero])/total)\n                        w = int(np.sum(PLANE[nonzero]*indices[1,nonzero])/total)\n#                       Let's define the starting patch as the one with his\n#                       bottom right corner closer to the centroid\n                        h_ = max([0,h - PATCH_SIZE])\n                        w_ = max([0,w - PATCH_SIZE])\n#                       let's find a nice picture\n                        for i in range(max_iter):\n                            PATCH = PLANE[h_:h_+PATCH_SIZE,w_:w_+PATCH_SIZE]\n                            nonzero = PATCH > 0\n                            total = np.sum(PATCH[nonzero])\n                            if total > 0:\n#                               let's find the centroid of that patch\n                                indices = np.indices(PATCH.shape)\n                                h = h_ + int(np.sum(PATCH[nonzero]*indices[0,nonzero])/total)\n                                w = w_ + int(np.sum(PATCH[nonzero]*indices[1,nonzero])/total)\n                                CM = [h,w]\n                                FOUND = True\n                            else:\n                                break                           \n#                           let's define the new patch tryng to center\n#                           the centroid on it\n                            if H > PATCH_SIZE:\n#                               if it actually fits\n                                h_n = min([max([0,h - PATCH_SIZE//2]),H - PATCH_SIZE])\n                            else:\n                                h_n = 0\n                            if W > PATCH_SIZE:\n#                               if it actually fits\n                                w_n = min([max([0,w - PATCH_SIZE//2]),W - PATCH_SIZE])\n                            else:\n                                w_n = 0\n#                           Stop if you've already been here\n                            if (h_n - h_)*(h_n - h_) + (w_n - w_)*(w_n - w_) == 0: break\n                            h_,w_ = h_n,w_n\n                            \n                        if FOUND: crops.append([axis,d,CM])\n        \n            random.shuffle(crops)\n            self.crops.append(crops)\n            self.LEN.append(len(crops))\n           \n        if mode == 'train':\n            self.MAX_LEN = int(np.max(self.LEN)*(1 + data_augmentation))\n        else:\n            self.MAX_LEN = np.max(self.LEN)\n            \n        for kidney in range(len(Datasets)):\n            self.IDX.append(list(range(self.MAX_LEN)))\n            random.shuffle(self.IDX[kidney])\n\n    def __MAXLEN__(self):\n        return self.MAX_LEN        \n   \n    def __len__(self):\n        return len(self.LEN)*self.MAX_LEN\n        \n    def get_crop(self,kidney,axis,d,CM,P,PAD):\n        DEPTH = self.DEPTH\n#       axis 0 label, orthogonal to Z\n#       axis 1 label, orthogonal to X\n#       axis 2 label, orthogonal to Y\n        crop = self.volume[kidney][0].take(range(d,d+DEPTH),axis).swapaxes(0,axis)\n        target = self.target[kidney][0].take(range(d,d+DEPTH),axis).swapaxes(0,axis)\n        rot = int(4*(d/4-d//4))\n        flip = int(2*(d/2-d//2))\n#       Simple transformations to make more different neighbor subvolumes\n        crop = np.rot90(crop,rot,(-2,-1))\n        target = np.rot90(target,rot,(-2,-1))\n        if flip:\n            crop = np.flip(crop,-3)\n            target = np.flip(target,-3)\n#       Ending paddings to center the center of masses\n        X,Y = crop.shape[-2:]\n        XPR = 0\n        XR = CM[0] + P - P//2 + PAD\n        if XR > X: XPR = XR - X\n        YPR = 0\n        YR = CM[1] + P - P//2 + PAD\n        if YR > Y: YPR = YR - Y\n#       Leading paddings to center the center of masses\n        XPL = 0\n        X_START = CM[0] - P//2 - PAD\n        if X_START < 0:\n            XPL = - X_START\n            X_START = 0\n        YPL = 0\n        Y_START = CM[1] - P//2 - PAD\n        if Y_START < 0: \n            YPL = - Y_START\n            Y_START = 0\n            \n        crop = np.pad(crop,((0,0),(XPL,XPR),(YPL,YPR)),'reflect')\n        target = np.pad(target,((0,0),(XPL,XPR),(YPL,YPR)),'reflect')\n        \n        crop = crop[...,X_START:X_START+P+2*PAD,Y_START:Y_START+P+2*PAD]\n        target = target[...,X_START:X_START+P+2*PAD,Y_START:Y_START+P+2*PAD]\n        if self.DEPTH_ZOOM_OUT > 1 or self.PATCH_ZOOM_OUT > 1:\n            blocks = np.ones(len(crop.shape),dtype=np.int32)\n            blocks[-3] = self.DEPTH_ZOOM_OUT\n            blocks[-2:] = self.PATCH_ZOOM_OUT\n            blocks = tuple(blocks)\n            crop = skimage.measure.block_reduce(crop,blocks, np.mean)\n            target = skimage.measure.block_reduce(target,blocks, np.min)\n\n        crop = torch.from_numpy(crop.astype(np.float32)).float().to(device)\n        crop = (crop - self.min[kidney])/(self.max[kidney] - self.min[kidney])\n        \n        target = torch.from_numpy(target > 0).float().to(device)\n        if crop.shape[-3] > 1:\n            return crop.unsqueeze(0),target.unsqueeze(0)\n        else:\n            return crop,target\n    \n    def get_augmented(self, kidney, aug_idx):\n        P = self.PATCH\n        axis,d,CM = self.crops[kidney][aug_idx]\n#       Deallocating CM aside from the center\n        CM = np.random.normal(CM, self.sigma).astype(np.int32)\n#-----------------------------------------------------------------------------\n#       Full rotations leads to weird looking masks, I decided to rot90 only\n#       The centroid deallocation together with random flips and rot90\n#       should be enough\n#_____________________________________________________________________________\n        fullrot = False\n        if torch.rand(1) < self.pfullrot: fullrot = True\n        if fullrot:\n#           The minimum pad that needs to be added in order\n#           to ensure no empty values after any rotation\n            PAD = int(P*(np.sqrt(2) - 1)//2 + 1)\n#_____________________________________________________________________________\n        else:\n            PAD = 0\n        crop,target = Crops3D_Dataset.get_crop(self,kidney,axis,d,CM,P,PAD)\n        rot = torch.randint(0,3,(1,)).item() + 1\n        rng = torch.get_rng_state()\n        if fullrot: crop = T.RandomRotation(180,interpolation=TF.InterpolationMode.BILINEAR)(crop)\n        crop = torch.rot90(crop, rot, dims=[-2, -1])\n        crop = T.RandomHorizontalFlip()(crop)\n        crop = T.RandomVerticalFlip()(crop)\n\n        torch.set_rng_state(rng)\n        if fullrot: target = T.RandomRotation(180)(target)\n        target = torch.rot90(target, rot, dims=[-2, -1])\n        target = T.RandomHorizontalFlip()(target)\n        target = T.RandomVerticalFlip()(target)\n        if self.PATCH_ZOOM_OUT > 1:\n            PAD = PAD//self.PATCH_ZOOM_OUT\n            P = P//self.PATCH_ZOOM_OUT\n        crop = crop[...,PAD:PAD+P,PAD:PAD+P]\n        target = target[...,PAD:PAD+P,PAD:PAD+P]\n        if torch.rand(1)>.5:\n            crop = crop.flip((-3,))\n            target = target.flip((-3,))\n        return crop,target            \n\n    def __getitem__(self, idx):\n#       we will point at each kidney consecutively\n        N = len(self.LEN)\n        kidney = int(N*(idx/N-idx//N))\n        idx = self.IDX[kidney][idx//N]\n        P = self.PATCH\n        if idx < len(self.crops[kidney]):\n            axis,d,CM = self.crops[kidney][idx]\n            crop,target = Crops3D_Dataset.get_crop(self,kidney,axis,d,CM,P,0)\n        else:\n#           we will point to each kidney crop consecutively\n            idx = idx - len(self.crops[kidney])*(idx//len(self.crops[kidney])) - 1\n            crop,target = Crops3D_Dataset.get_augmented(self,kidney,idx)\n\n        return crop,target","metadata":{"execution":{"iopub.execute_input":"2024-01-14T20:43:03.784487Z","iopub.status.busy":"2024-01-14T20:43:03.784135Z","iopub.status.idle":"2024-01-14T20:43:03.836031Z","shell.execute_reply":"2024-01-14T20:43:03.835170Z"},"papermill":{"duration":0.063914,"end_time":"2024-01-14T20:43:03.838079","exception":false,"start_time":"2024-01-14T20:43:03.774165","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Fake_Val_Dataset(torch.utils.data.Dataset):\n    '''   \n    Dataset with a single dummy sample\n    '''  \n    def __init__(self,DEPTH,PATCH_SIZE):\n        self.sample = torch.zeros((DEPTH,PATCH_SIZE,PATCH_SIZE),dtype=torch.float32)\n        if DEPTH > 1:\n            self.sample = self.sample.unsqueeze(0)\n\n    def __len__(self):\n        return 1\n    \n    def __getitem__(self, idx):\n        return self.sample,self.sample","metadata":{"execution":{"iopub.execute_input":"2024-01-14T20:43:03.856780Z","iopub.status.busy":"2024-01-14T20:43:03.855965Z","iopub.status.idle":"2024-01-14T20:43:03.862427Z","shell.execute_reply":"2024-01-14T20:43:03.861516Z"},"papermill":{"duration":0.018245,"end_time":"2024-01-14T20:43:03.864621","exception":false,"start_time":"2024-01-14T20:43:03.846376","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def mixedidx(K,MAXLEN):\n#   Allows random stratified indexes\n    idx = torch.zeros(K*MAXLEN).long()\n    for k in range(K):\n        idx[np.arange(0,K*MAXLEN,K).astype(np.int32)+k]  = (torch.arange(0,K*MAXLEN,K)+k)[torch.randperm(MAXLEN)]\n\n    return idx","metadata":{"execution":{"iopub.execute_input":"2024-01-14T20:43:03.882514Z","iopub.status.busy":"2024-01-14T20:43:03.882147Z","iopub.status.idle":"2024-01-14T20:43:03.888053Z","shell.execute_reply":"2024-01-14T20:43:03.887015Z"},"papermill":{"duration":0.017161,"end_time":"2024-01-14T20:43:03.890006","exception":false,"start_time":"2024-01-14T20:43:03.872845","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# https://github.com/pytorch/pytorch/blob/main/torch/utils/data/sampler.py\nimport torch\nfrom torch import Tensor\n\nfrom typing import Iterator, Iterable, Optional, Sequence, List, TypeVar, Generic, Sized, Union\n\nclass MyRandomSampler(Sampler[int]):\n#   Allows random stratified indexes\n    r\"\"\"Samples elements randomly. If without replacement, then sample from a shuffled dataset.\n    I need for my Datset that sampler gives at each epoch\n\n    If with replacement, then user can specify :attr:`num_samples` to draw.\n\n    Args:\n        data_source (Dataset): dataset to sample from\n        num_samples (int): number of samples to draw, default=`len(dataset)`.\n        generator (Generator): Generator used in sampling.\n    \"\"\"\n\n    data_source: Sized\n    replacement: bool\n\n    def __init__(self, data_source,\n                 MAXLEN: [int],\n                 num_kidneys: [int],\n                 num_samples: Optional[int] = None, generator=None) -> None:\n        self.data_source = data_source\n        self.MAXLEN = MAXLEN\n        self.num_kidneys = num_kidneys\n        self._num_samples = num_samples\n        self.generator = generator\n\n        if not isinstance(self.num_samples, int) or self.num_samples <= 0:\n            raise ValueError(f\"num_samples should be a positive integer value, but got num_samples={self.num_samples}\")\n\n    @property\n    def num_samples(self) -> int:\n        # dataset size might change at runtime\n        if self._num_samples is None:\n            return len(self.data_source)\n        return self._num_samples\n\n    def __iter__(self) -> Iterator[int]:\n        n = len(self.data_source)\n        if self.generator is None:\n            seed = int(torch.empty((), dtype=torch.int64).random_().item())\n            generator = torch.Generator()\n            generator.manual_seed(seed)\n        else:\n            generator = self.generator\n\n        for _ in range(self.num_samples // n):\n            yield from mixedidx(self.num_kidneys,self.MAXLEN).tolist()\n        yield from mixedidx(self.num_kidneys,self.MAXLEN).tolist()[:self.num_samples % n]\n\n    def __len__(self) -> int:\n        return self.num_samples","metadata":{"execution":{"iopub.execute_input":"2024-01-14T20:43:03.908140Z","iopub.status.busy":"2024-01-14T20:43:03.907784Z","iopub.status.idle":"2024-01-14T20:43:03.919482Z","shell.execute_reply":"2024-01-14T20:43:03.918451Z"},"papermill":{"duration":0.023292,"end_time":"2024-01-14T20:43:03.921532","exception":false,"start_time":"2024-01-14T20:43:03.898240","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# based on:\n# https://github.com/kevinzakka/pytorch-goodies/blob/master/losses.py\n\nclass myDiceLoss(nn.Module):\n    r\"\"\"Criterion that computes Sørensen-Dice Coefficient loss.\n\n    According to [1], we compute the Sørensen-Dice Coefficient as follows:\n\n    .. math::\n\n        \\text{Dice}(x, class) = \\frac{2 |X| \\cap |Y|}{|X| + |Y|}\n\n    where:\n       - :math:`X` expects to be the scores of each class.\n       - :math:`Y` expects to be the one-hot tensor with the class labels.\n\n    the loss, is finally computed as:\n\n    .. math::\n\n        \\text{loss}(x, class) = 1 - \\text{Dice}(x, class)\n\n    [1] https://en.wikipedia.org/wiki/S%C3%B8rensen%E2%80%93Dice_coefficient\n\n    Shape:\n        - Input: :math:`(N, C, H, W)` where C = number of classes.\n        - Target: :math:`(N, H, W)` where each value is\n          :math:`0 ≤ targets[i] ≤ C−1`.\n\n    Examples:\n        >>> N = 5  # num_classes\n        >>> loss = tgm.losses.DiceLoss()\n        >>> input = torch.randn(1, N, 3, 5, requires_grad=True)\n        >>> target = torch.empty(1, 3, 5, dtype=torch.long).random_(N)\n        >>> output = loss(input, target)\n        >>> output.backward()\n    \"\"\"\n\n    def __init__(self) -> None:\n        super(myDiceLoss, self).__init__()\n        self.eps: float = 1e-6\n\n    def forward(\n            self,\n            input: torch.Tensor,\n            target: torch.Tensor) -> torch.Tensor:\n        if not torch.is_tensor(input):\n            raise TypeError(\"Input type is not a torch.Tensor. Got {}\"\n                            .format(type(input)))\n        if not len(input.shape) == 4:\n            raise ValueError(\"Invalid input shape, we expect BxNxDxHxW. Got: {}\"\n                             .format(input.shape))\n        if not input.shape[-2:] == target.shape[-2:]:\n            raise ValueError(\"input and target shapes must be the same. Got: {}\"\n                             .format(input.shape, input.shape))\n        if not input.device == target.device:\n            raise ValueError(\n                \"input and target must be in the same device. Got: {}\" .format(\n                    input.device, target.device))\n        # compute softmax over the classes axis\n        input_soft = F.softmax(input, dim=-3)\n\n        # create the labels one hot tensor\n        target_one_hot = F.one_hot(target.long(), num_classes=input.shape[1])\n        target_one_hot = torch.swapaxes(target_one_hot,1,-1).squeeze(-1)\n\n        # compute the actual dice score\n        dims = (1, 2, 3)\n        intersection = torch.sum(input_soft * target_one_hot, dims)\n        cardinality = torch.sum(input_soft + target_one_hot, dims)\n\n        dice_score = 2. * intersection / (cardinality + self.eps)\n        return torch.mean(1. - dice_score)\n\n\n\n######################\n# functional interface\n######################\n\n\ndef dice_loss(\n        input: torch.Tensor,\n        target: torch.Tensor) -> torch.Tensor:\n    r\"\"\"Function that computes Sørensen-Dice Coefficient loss.\n\n    See :class:`~torchgeometry.losses.DiceLoss` for details.\n    \"\"\"\n    return myDiceLoss()(input, target)","metadata":{"execution":{"iopub.execute_input":"2024-01-14T20:43:03.941748Z","iopub.status.busy":"2024-01-14T20:43:03.941356Z","iopub.status.idle":"2024-01-14T20:43:03.953872Z","shell.execute_reply":"2024-01-14T20:43:03.952923Z"},"papermill":{"duration":0.024868,"end_time":"2024-01-14T20:43:03.955885","exception":false,"start_time":"2024-01-14T20:43:03.931017","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#https://forums.fast.ai/t/plotting-metrics-after-learning/69937/3    \n@patch\n@delegates(subplots)\ndef plot_metrics(self: Recorder, nrows=None, ncols=None, figsize=None, **kwargs):\n    metrics = np.stack(self.values)\n    names = self.metric_names[1:-1]\n    n = len(names) - 1\n    if nrows is None and ncols is None:\n        nrows = int(math.sqrt(n))\n        ncols = int(np.ceil(n / nrows))\n    elif nrows is None: nrows = int(np.ceil(n / ncols))\n    elif ncols is None: ncols = int(np.ceil(n / nrows))\n    figsize = figsize or (ncols * 6, nrows * 4)\n    fig, axs = subplots(nrows, ncols, figsize=figsize, **kwargs)\n    axs = [ax if i < n else ax.set_axis_off() for i, ax in enumerate(axs.flatten())][:n]\n    for i, (name, ax) in enumerate(zip(names, [axs[0]] + axs)):\n        ax.plot(metrics[:, i], color='#1f77b4' if i == 0 else '#ff7f0e', label='valid' if i > 0 else 'train')\n        ax.set_title(name if i > 1 else 'losses')\n        ax.legend(loc='best')\n    plt.show()","metadata":{"execution":{"iopub.execute_input":"2024-01-14T20:43:03.976334Z","iopub.status.busy":"2024-01-14T20:43:03.975967Z","iopub.status.idle":"2024-01-14T20:43:03.986608Z","shell.execute_reply":"2024-01-14T20:43:03.985630Z"},"papermill":{"duration":0.023733,"end_time":"2024-01-14T20:43:03.988756","exception":false,"start_time":"2024-01-14T20:43:03.965023","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.cuda import amp","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(SEED)\n\nmodel = smp.Unet(\n    encoder_name=ENCODER,\n    encoder_weights=WEIGHTS,\n    in_channels=1,\n    classes=2,\n    decoder_attention_type='scse')\n\nds_train = Crops3D_Dataset(TRAIN,DEPTH=DEPTH,PATCH_SIZE=PATCH_SIZE)\nds_valid = Fake_Val_Dataset(DEPTH=DEPTH,PATCH_SIZE=PATCH_SIZE)\n\ndl_train = torch.utils.data.DataLoader(ds_train,\n                                       batch_size=BATCH_SIZE,\n                                       num_workers=0,\n                                       sampler=MyRandomSampler(np.arange(len(ds_train)),len(TRAIN),ds_train.__MAXLEN__()),\n                                       drop_last=True)\ndl_valid = torch.utils.data.DataLoader(ds_valid)\ndls = DataLoaders(dl_train,dl_valid)\n    \nlearn = Learner(dls, \n                model,\n                lr=LR,\n                loss_func=dice_loss,\n                cbs=[\n                    GradientClip(3.0),\n                    SaveModelCallback(every_epoch=True),\n                    ShowGraphCallback()]\n)\nlearn.fit_one_cycle(EPOCHS)\ntorch.save(model,MODEL_NAME)\nlearn.recorder.plot_metrics()\ndel model,ds_train,dl_train,ds_valid,dl_valid,dls\ngc.collect()","metadata":{"execution":{"iopub.execute_input":"2024-01-14T20:43:04.007204Z","iopub.status.busy":"2024-01-14T20:43:04.006865Z","iopub.status.idle":"2024-01-15T06:22:02.095840Z","shell.execute_reply":"2024-01-15T06:22:02.094879Z"},"papermill":{"duration":34738.100875,"end_time":"2024-01-15T06:22:02.098081","exception":false,"start_time":"2024-01-14T20:43:03.997206","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}