{"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":"# HuBMAP - Efficient Sampling Baseline (deepflash2, pytorch, fastai) [train]\n\n> Kernel for model training with efficient region based sampling.\n\nRequires deepflash2 (git version), zarr, and segmentation-models-pytorch\n","metadata":{"id":"OcsetTMwKXqC"}},{"cell_type":"markdown","source":"### Installation and package loading","metadata":{}},{"cell_type":"code","source":"!pip install deepflash2\n!pip install segmentation_models_pytorch\n!pip install fastdownload\n!pip install fastai --upgrade","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"scrolled":true,"execution":{"iopub.status.busy":"2022-03-13T15:26:42.638576Z","iopub.execute_input":"2022-03-13T15:26:42.639014Z","iopub.status.idle":"2022-03-13T15:28:32.585699Z","shell.execute_reply.started":"2022-03-13T15:26:42.638921Z","shell.execute_reply":"2022-03-13T15:28:32.584837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import zarr, cv2\nimport numpy as np, pandas as pd, segmentation_models_pytorch as smp\nfrom fastai.vision.all import *\nfrom deepflash2.all import *\nimport albumentations as alb","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:32.587618Z","iopub.execute_input":"2022-03-13T15:28:32.587900Z","iopub.status.idle":"2022-03-13T15:28:39.399679Z","shell.execute_reply.started":"2022-03-13T15:28:32.587872Z","shell.execute_reply":"2022-03-13T15:28:39.398693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Helper functions and patches","metadata":{}},{"cell_type":"code","source":"@patch\ndef read_img(self:BaseDataset, file, *args, **kwargs):\n    return zarr.open(str(file), mode='r')\n\n@patch\ndef _name_fn(self:BaseDataset, g):\n    \"Name of preprocessed and compressed data.\"\n    return f'{g}'\n\n@patch\ndef apply(self:DeformationField, data, offset=(0, 0), pad=(0, 0), order=1):\n    \"Apply deformation field to image using interpolation\"\n    outshape = tuple(int(s - p) for (s, p) in zip(self.shape, pad))\n    coords = [np.squeeze(d).astype('float32').reshape(*outshape) for d in self.get(offset, pad)]\n    # Get slices to avoid loading all data (.zarr files)\n    sl = []\n    for i in range(len(coords)):\n        cmin, cmax = int(coords[i].min()), int(coords[i].max())\n        dmax = data.shape[i]\n        if cmin<0: \n            cmax = max(-cmin, cmax)\n            cmin = 0 \n        elif cmax>dmax:\n            cmin = min(cmin, 2*dmax-cmax)\n            cmax = dmax\n            coords[i] -= cmin\n        else: coords[i] -= cmin\n        sl.append(slice(cmin, cmax))    \n    if len(data.shape) == len(self.shape) + 1:\n        tile = np.empty((*outshape, data.shape[-1]))\n        for c in range(data.shape[-1]):\n            # Adding divide\n            tile[..., c] = cv2.remap(data[sl[0],sl[1], c]/255, coords[1],coords[0], interpolation=order, borderMode=cv2.BORDER_REFLECT)\n    else:\n        tile = cv2.remap(data[sl[0], sl[1]], coords[1], coords[0], interpolation=order, borderMode=cv2.BORDER_REFLECT)\n    return tile\n\ndef dice(im1, im2):\n    \"\"\"\n    Computes the Dice coefficient, a measure of set similarity.\n    Parameters\n    ----------\n    im1 : array-like, bool\n        Any array of arbitrary size. If not boolean, will be converted.\n    im2 : array-like, bool\n        Any other array of identical size. If not boolean, will be converted.\n    Returns\n    -------\n    dice : float\n        Dice coefficient as a float on range [0,1].\n        Maximum similarity = 1\n        No similarity = 0\n        \n    Notes\n    -----\n    The order of inputs for `dice` is irrelevant. The result will be\n    identical if `im1` and `im2` are switched.\n    \"\"\"\n    im1 = np.asarray(im1).astype(np.bool)\n    im2 = np.asarray(im2).astype(np.bool)\n\n    if im1.shape != im2.shape:\n        raise ValueError(\"Shape mismatch: im1 and im2 must have the same shape.\")\n\n    # Compute Dice coefficient\n    intersection = np.logical_and(im1, im2)\n\n    return 2. * intersection.sum() / (im1.sum() + im2.sum())","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-03-13T15:28:39.401070Z","iopub.execute_input":"2022-03-13T15:28:39.401560Z","iopub.status.idle":"2022-03-13T15:28:39.425264Z","shell.execute_reply.started":"2022-03-13T15:28:39.401522Z","shell.execute_reply":"2022-03-13T15:28:39.424352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Config","metadata":{}},{"cell_type":"code","source":"class CONFIG():\n    \n    # data paths\n    data_path = Path('../input/hubmap-kidney-segmentation')\n    data_path_zarr = Path('../input/k/felipesens/hubmap-zarr/train_scale2/')\n    mask_preproc_dir = '/kaggle/input/hubmap-labels-pdf-0-5-0-25-0-01/masks_scale2'\n    \n    # deepflash2 dataset\n    scale = 1.5 # data is already downscaled to 2, so absulute downscale is 3\n    tile_shape = (256, 256)\n    padding = (0,0) # Border overlap for prediction\n    n_jobs = 1\n    sample_mult = 100 # Sample 100 tiles from each image, per epoch\n    val_length = 500 # Randomly sample 500 validation tiles\n    stats = np.array([0.61561477, 0.5179343 , 0.64067212]), np.array([0.2915353 , 0.31549066, 0.28647661])\n    \n    # deepflash2 augmentation options\n    flip = False\n    rotation_range_deg = (0, 0)\n\n    # pytorch model (segmentation_models_pytorch)\n    encoder_name = \"efficientnet-b0\"\n    encoder_weights = 'imagenet'\n    in_channels = 3\n    classes = 2\n    \n    # fastai Learner \n    mixed_precision_training = True\n    batch_size = 16\n    weight_decay = 0.01\n    loss_func = CrossEntropyLossFlat(axis=1)\n    metrics = [Dice()]\n    optimizer = ranger\n    max_learning_rate = 1e-3\n    epochs = 12\n    \ncfg = CONFIG()","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:39.426606Z","iopub.execute_input":"2022-03-13T15:28:39.427153Z","iopub.status.idle":"2022-03-13T15:28:39.438026Z","shell.execute_reply.started":"2022-03-13T15:28:39.427101Z","shell.execute_reply":"2022-03-13T15:28:39.436905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tfms = []#alb.Compose([])","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:39.441348Z","iopub.execute_input":"2022-03-13T15:28:39.441884Z","iopub.status.idle":"2022-03-13T15:28:39.447970Z","shell.execute_reply.started":"2022-03-13T15:28:39.441839Z","shell.execute_reply":"2022-03-13T15:28:39.446861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(cfg.data_path/'train.csv')\ndf_info = pd.read_csv(cfg.data_path/'HuBMAP-20-dataset_information.csv')\n\nfiles = [x for x in cfg.data_path_zarr.iterdir() if x.is_dir() if not x.name.startswith('.')]\nprint(len(files))\ntrain_files = files[0:10]\nvalid_files = files[10:]\nprint(train_files, valid_files, len(train_files), len(valid_files))\nlabel_fn = lambda o: o","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:39.449492Z","iopub.execute_input":"2022-03-13T15:28:39.450090Z","iopub.status.idle":"2022-03-13T15:28:39.812320Z","shell.execute_reply.started":"2022-03-13T15:28:39.450045Z","shell.execute_reply":"2022-03-13T15:28:39.811386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"# Model\nmodel = smp.Unet(encoder_name=cfg.encoder_name, \n                 encoder_weights=cfg.encoder_weights, \n                 in_channels=cfg.in_channels, \n                 classes=cfg.classes,\n                 #decoder_attention_type= \"scse\"\n                )","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Datasets\nds_kwargs = {\n    #'tile_shape':cfg.tile_shape,\n    'padding':cfg.padding,\n    'scale': cfg.scale,\n    'n_jobs': cfg.n_jobs, \n    'preproc_dir': cfg.mask_preproc_dir, \n    'val_length':cfg.val_length, \n    'sample_mult':cfg.sample_mult,\n    'loss_weights':False,\n    'flip' : cfg.flip,\n    'rotation_range_deg': cfg.rotation_range_deg,\n    #'albumentations_tfms': tfms\n}\n\ntrain_ds = RandomTileDataset(train_files, label_fn=label_fn, normalize=False, use_preprocessed_labels=True, **ds_kwargs)\nvalid_ds = TileDataset(valid_files, label_fn=label_fn, normalize=False, use_preprocessed_labels=True, **ds_kwargs, is_zarr=True)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:30:45.470652Z","iopub.execute_input":"2022-03-13T15:30:45.471006Z","iopub.status.idle":"2022-03-13T15:30:46.096132Z","shell.execute_reply.started":"2022-03-13T15:30:45.470973Z","shell.execute_reply":"2022-03-13T15:30:46.095435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile = train_ds[0]\nshow(tile[0], tile[1])\ntile[0].type()","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:42.454763Z","iopub.execute_input":"2022-03-13T15:28:42.455333Z","iopub.status.idle":"2022-03-13T15:28:43.092887Z","shell.execute_reply.started":"2022-03-13T15:28:42.455287Z","shell.execute_reply":"2022-03-13T15:28:43.091966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tile = valid_ds[489]\nshow(tile[0], tile[1])\ntile[0].type()","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:43.093888Z","iopub.execute_input":"2022-03-13T15:28:43.094190Z","iopub.status.idle":"2022-03-13T15:28:43.415778Z","shell.execute_reply.started":"2022-03-13T15:28:43.094157Z","shell.execute_reply":"2022-03-13T15:28:43.414912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class _T(Transform):  \n    def encodes(self, x):  \n        if(x.shape[0] == 3): return x.float()\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:43.417098Z","iopub.execute_input":"2022-03-13T15:28:43.417612Z","iopub.status.idle":"2022-03-13T15:28:43.424525Z","shell.execute_reply.started":"2022-03-13T15:28:43.417576Z","shell.execute_reply":"2022-03-13T15:28:43.423524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls = DataLoaders.from_dsets(train_ds, valid_ds, bs=cfg.batch_size, after_item=_T, after_batch=Normalize.from_stats(*cfg.stats))\nif torch.cuda.is_available(): dls.cuda()\ncbs = [SaveModelCallback(monitor='dice')]","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:43.426115Z","iopub.execute_input":"2022-03-13T15:28:43.426921Z","iopub.status.idle":"2022-03-13T15:28:46.454982Z","shell.execute_reply.started":"2022-03-13T15:28:43.426857Z","shell.execute_reply":"2022-03-13T15:28:46.454062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Dataloader and learner\ndls = DataLoaders.from_dsets(train_ds, valid_ds, bs=cfg.batch_size, after_item=_T, after_batch=Normalize.from_stats(*cfg.stats))\nif torch.cuda.is_available(): dls.cuda(), model.cuda()\ncbs = [SaveModelCallback(monitor='dice')]\nlearn = Learner(dls, model, metrics=cfg.metrics, wd=cfg.weight_decay, loss_func=cfg.loss_func, opt_func=ranger, cbs=cbs)\nif cfg.mixed_precision_training: learn.to_fp16()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Fit\nlearn.fit_one_cycle(cfg.epochs, lr_max=cfg.max_learning_rate)\nlearn.recorder.plot_metrics()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def find_best_and_worst():\n    min_dice = 1\n    max_dice = 0\n    best = None\n    worst = None\n    for i, (x, y) in enumerate(dls.valid):\n        with torch.no_grad():\n            preds = learn.model(x)\n            preds, y = preds.float().cpu(), y.cpu()\n            for j, (p, t, img) in enumerate(zip(preds,y,x)):\n                d = dice(torch.sigmoid(p[1]) > 0.4, t)\n                if d > max_dice: \n                    best = (p, t, img.cpu())\n                    max_dice = d\n                    print(\"best found\", d)\n                if d < min_dice and d != 0 and t.sum() > 1000: \n                    worst = (p, t, img.cpu())\n                    min_dice = d\n                    print(\"worst found\", d)\n    return best, worst","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:46.456343Z","iopub.execute_input":"2022-03-13T15:28:46.456692Z","iopub.status.idle":"2022-03-13T15:28:46.467830Z","shell.execute_reply.started":"2022-03-13T15:28:46.456657Z","shell.execute_reply":"2022-03-13T15:28:46.466676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best, worst = find_best_and_worst()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_pedictions(find_best_and_worst())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.rcParams['figure.figsize'] = [20, 5]\n\ndef show_prediction(pred, mask, image):\n    plt.subplot(1,4,1)\n    plt.imshow(pred)\n    plt.subplot(1,4,2)\n    plt.imshow(mask)\n    plt.subplot(1,4,3)\n    plt.imshow(mask != pred)\n    plt.subplot(1,4,4)\n    plt.imshow(image.permute(1,2,0))\n    plt.imshow(mask, 'RdYlBu', alpha=0.4)\n    plt.imshow(pred, alpha=0.4)\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:46.469342Z","iopub.execute_input":"2022-03-13T15:28:46.469782Z","iopub.status.idle":"2022-03-13T15:28:46.482338Z","shell.execute_reply.started":"2022-03-13T15:28:46.469718Z","shell.execute_reply":"2022-03-13T15:28:46.481545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_pedictions(preds):\n    for p in preds:\n        pred = torch.sigmoid(-p[0][0]) > 0.4\n        mask = p[1]\n        img = p[2] \n        show_prediction(pred, mask, img)","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:46.483918Z","iopub.execute_input":"2022-03-13T15:28:46.484411Z","iopub.status.idle":"2022-03-13T15:28:46.493581Z","shell.execute_reply.started":"2022-03-13T15:28:46.484377Z","shell.execute_reply":"2022-03-13T15:28:46.492708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred = torch.sigmoid(-best[0][0]) > 0.4\nmask = best[1]\nimg = best[2]\nshow_prediction(pred, mask, img)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred = torch.sigmoid(-worst[0][0]) > 0.4\nmask = worst[1]\nimg = worst[2]\nshow_prediction(pred, mask, img)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch import optim","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:28:46.495050Z","iopub.execute_input":"2022-03-13T15:28:46.495462Z","iopub.status.idle":"2022-03-13T15:28:46.502544Z","shell.execute_reply.started":"2022-03-13T15:28:46.495425Z","shell.execute_reply":"2022-03-13T15:28:46.501435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls.train","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"OPTIMIZERS = [ranger, RMSProp, Adam, SGD]\nEPOCHS = 5\nfor opt in OPTIMIZERS:\n    print(f\"Testing {opt}\")\n    \n    # Model\n    model = smp.Unet(encoder_name=cfg.encoder_name, \n                 encoder_weights=cfg.encoder_weights, \n                 in_channels=cfg.in_channels, \n                 classes=cfg.classes,\n                 #decoder_attention_type= \"scse\"\n                )\n    \n    # Dataloader and learner\n    if torch.cuda.is_available(): model.cuda()\n    learn = Learner(dls, model, metrics=cfg.metrics, wd=cfg.weight_decay, loss_func=cfg.loss_func, opt_func=opt, cbs=cbs)\n    if cfg.mixed_precision_training: learn.to_fp16()\n    \n    # Fit\n    learn.fit_one_cycle(EPOCHS, lr_max=cfg.max_learning_rate)\n    learn.recorder.plot_metrics()\n    \n    plot_pedictions(find_best_and_worst())\n    ","metadata":{"execution":{"iopub.status.busy":"2022-03-12T23:05:06.679767Z","iopub.execute_input":"2022-03-12T23:05:06.680088Z","iopub.status.idle":"2022-03-12T23:14:45.537233Z","shell.execute_reply.started":"2022-03-12T23:05:06.680056Z","shell.execute_reply":"2022-03-12T23:14:45.536190Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"opt = RMSProp\nBACKBONES = [\n    \"efficientnet-b0\",\n    \"efficientnet-b1\",\n    \"efficientnet-b2\",\n    \"efficientnet-b3\",\n    \"efficientnet-b4\",\n    \"efficientnet-b5\",\n    \"efficientnet-b6\",\n    \"resnet18\",\n    \"resnet34\",\n    \"resnet50\",\n    \"resnet101\"\n] \nEPOCHS = 5\nfor bkb in BACKBONES:\n    print(f\"Testing {bkb}\")\n    \n    # Model\n    model = smp.Unet(encoder_name=bkb, \n                 encoder_weights=cfg.encoder_weights, \n                 in_channels=cfg.in_channels, \n                 classes=cfg.classes,\n                 #decoder_attention_type= \"scse\"\n                )\n    \n    # Dataloader and learner\n    if torch.cuda.is_available(): model.cuda()\n    learn = Learner(dls, model, metrics=cfg.metrics, wd=cfg.weight_decay, loss_func=cfg.loss_func, opt_func=opt, cbs=cbs)\n    if cfg.mixed_precision_training: learn.to_fp16()\n    \n    # Fit\n    learn.fit_one_cycle(EPOCHS, lr_max=cfg.max_learning_rate)\n    learn.recorder.plot_metrics()\n    \n    plot_pedictions(find_best_and_worst())\n    ","metadata":{"execution":{"iopub.status.busy":"2022-03-12T23:39:13.089620Z","iopub.execute_input":"2022-03-12T23:39:13.090010Z","iopub.status.idle":"2022-03-13T00:10:57.055019Z","shell.execute_reply.started":"2022-03-12T23:39:13.089977Z","shell.execute_reply":"2022-03-13T00:10:57.053963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"opt = RMSProp\nbkb = \"efficientnet-b6\"\nTRANSFORMS = [\n    [\n        alb.HorizontalFlip(),\n        alb.GridDistortion(p=.1),\n        alb.ToFloat(1),\n        alb.RandomBrightnessContrast(),\n    ],\n    [\n        alb.HorizontalFlip(),\n        alb.GridDistortion(p=.1),\n        alb.ToFloat(1),\n        alb.RandomBrightnessContrast(),\n        alb.Cutout(),\n    ],\n    [\n        alb.HorizontalFlip(),\n        alb.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, border_mode=cv2.BORDER_REFLECT),\n        alb.GridDistortion(p=.1),\n        alb.ToFloat(1),\n        alb.RandomBrightnessContrast(),\n    ],\n    [\n        alb.HorizontalFlip(),\n        alb.VerticalFlip(),\n        alb.RandomRotate90(),\n        alb.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, border_mode=cv2.BORDER_REFLECT)\n    ],\n    [\n        alb.HorizontalFlip(),\n        alb.GridDistortion(p=.1),\n    ],\n    \n]\nEPOCHS = 5\nfor tfms in TRANSFORMS:\n    print(f\"Testing {tfms}\")\n    \n    # Model\n    model = smp.Unet(encoder_name=bkb, \n                 encoder_weights=cfg.encoder_weights, \n                 in_channels=cfg.in_channels, \n                 classes=cfg.classes,\n                 #decoder_attention_type= \"scse\"\n                )\n    \n    # Dataloader and learner\n    train_ds = RandomTileDataset(train_files, label_fn=label_fn, normalize=False, use_preprocessed_labels=True, **ds_kwargs, albumentations_tfms=tfms)\n    valid_ds = TileDataset(valid_files, label_fn=label_fn, normalize=False, use_preprocessed_labels=True, **ds_kwargs, is_zarr=True)\n    dls = DataLoaders.from_dsets(train_ds, valid_ds, bs=cfg.batch_size, after_item=_T, after_batch=Normalize.from_stats(*cfg.stats))\n    if torch.cuda.is_available(): dls.cuda(), model.cuda()\n    cbs = [SaveModelCallback(monitor='dice')]\n    learn = Learner(dls, model, metrics=cfg.metrics, wd=cfg.weight_decay, loss_func=cfg.loss_func, opt_func=ranger, cbs=cbs)\n    if cfg.mixed_precision_training: learn.to_fp16()\n    \n    # Fit\n    learn.fit_one_cycle(EPOCHS, lr_max=cfg.max_learning_rate)\n    learn.recorder.plot_metrics()\n    \n    plot_pedictions(find_best_and_worst())\n    ","metadata":{"execution":{"iopub.status.busy":"2022-03-13T00:38:37.356641Z","iopub.execute_input":"2022-03-13T00:38:37.357013Z","iopub.status.idle":"2022-03-13T01:00:13.463671Z","shell.execute_reply.started":"2022-03-13T00:38:37.356976Z","shell.execute_reply":"2022-03-13T01:00:13.462619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"opt = RMSProp\nbkb = \"efficientnet-b6\"\ntfms = [alb.HorizontalFlip(),\n        alb.GridDistortion(p=.1),\n        alb.ToFloat(1),\n        alb.RandomBrightnessContrast(),]\nEPOCHS = 12\nRES_BATCH = [\n    (16, (384, 384)),\n    (8,  (512, 512)),\n    (16, (256, 256)),\n]\nfor batch_size, tile_shape in RES_BATCH:\n    print(f\"Testing {batch_size, tile_shape}\")\n    \n    # Model\n    model = smp.Unet(encoder_name=bkb, \n                 encoder_weights=cfg.encoder_weights, \n                 in_channels=cfg.in_channels, \n                 classes=cfg.classes,\n                 #decoder_attention_type= \"scse\"\n                )\n    \n    # Dataloader and learner\n    train_ds = RandomTileDataset(train_files, label_fn=label_fn, normalize=False, use_preprocessed_labels=True, **ds_kwargs, albumentations_tfms=tfms, tile_shape=tile_shape)\n    valid_ds = TileDataset(valid_files, label_fn=label_fn, normalize=False, use_preprocessed_labels=True, **ds_kwargs, is_zarr=True, tile_shape=tile_shape)\n    dls = DataLoaders.from_dsets(train_ds, valid_ds, bs=batch_size, after_item=_T, after_batch=Normalize.from_stats(*cfg.stats))\n    if torch.cuda.is_available(): dls.cuda(), model.cuda()\n    cbs = [SaveModelCallback(monitor='dice')]\n    learn = Learner(dls, model, metrics=cfg.metrics, wd=cfg.weight_decay, loss_func=cfg.loss_func, opt_func=ranger, cbs=cbs)\n    if cfg.mixed_precision_training: learn.to_fp16()\n    \n    # Fit\n    learn.fit_one_cycle(EPOCHS, lr_max=cfg.max_learning_rate)\n    learn.recorder.plot_metrics()\n    \n    plot_pedictions(find_best_and_worst())\n    ","metadata":{"execution":{"iopub.status.busy":"2022-03-13T15:32:46.346832Z","iopub.execute_input":"2022-03-13T15:32:46.347157Z","iopub.status.idle":"2022-03-13T16:33:19.314642Z","shell.execute_reply.started":"2022-03-13T15:32:46.347128Z","shell.execute_reply":"2022-03-13T16:33:19.313664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"opt = RMSProp\nbkb = \"efficientnet-b6\"\ntfms = [alb.HorizontalFlip(),\n        alb.GridDistortion(p=.1),\n        alb.ToFloat(1),\n        alb.RandomBrightnessContrast(),]\nEPOCHS = 12\nbatch_size, tile_shape = (8,  (512, 512))\nATTENTION = [\"scse\", None]\n\nfor atte in ATTENTION:\n    print(f\"Testing {atte}\")\n    \n    # Model\n    model = smp.Unet(encoder_name=bkb, \n                 encoder_weights=cfg.encoder_weights, \n                 in_channels=cfg.in_channels, \n                 classes=cfg.classes,\n                 decoder_attention_type= atte\n                )\n    \n    # Dataloader and learner\n    train_ds = RandomTileDataset(train_files, label_fn=label_fn, normalize=False, use_preprocessed_labels=True, **ds_kwargs, albumentations_tfms=tfms, tile_shape=tile_shape)\n    valid_ds = TileDataset(valid_files, label_fn=label_fn, normalize=False, use_preprocessed_labels=True, **ds_kwargs, is_zarr=True, tile_shape=tile_shape)\n    dls = DataLoaders.from_dsets(train_ds, valid_ds, bs=batch_size, after_item=_T, after_batch=Normalize.from_stats(*cfg.stats))\n    if torch.cuda.is_available(): dls.cuda(), model.cuda()\n    cbs = [SaveModelCallback(monitor='dice')]\n    learn = Learner(dls, model, metrics=cfg.metrics, wd=cfg.weight_decay, loss_func=cfg.loss_func, opt_func=ranger, cbs=cbs)\n    if cfg.mixed_precision_training: learn.to_fp16()\n    \n    # Fit\n    learn.fit_one_cycle(EPOCHS, lr_max=cfg.max_learning_rate)\n    learn.recorder.plot_metrics()\n    \n    plot_pedictions(find_best_and_worst())\n    ","metadata":{"execution":{"iopub.status.busy":"2022-03-13T16:41:49.811340Z","iopub.execute_input":"2022-03-13T16:41:49.811719Z","iopub.status.idle":"2022-03-13T17:46:48.991418Z","shell.execute_reply.started":"2022-03-13T16:41:49.811679Z","shell.execute_reply":"2022-03-13T17:46:48.990435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"opt = RMSProp\nbkb = \"efficientnet-b6\"\ntfms = [alb.HorizontalFlip(),\n        alb.GridDistortion(p=.1),\n        alb.ToFloat(1),\n        alb.RandomBrightnessContrast(),]\nEPOCHS = 12\nbatch_size, tile_shape = (8,  (512, 512))\nATTENTION = [\"scse\", None]\n\nfor atte in ATTENTION:\n    print(f\"Testing {atte}\")\n    \n    # Model\n    model = smp.UnetPlusPlus(encoder_name=bkb, \n                 encoder_weights=cfg.encoder_weights, \n                 in_channels=cfg.in_channels, \n                 classes=cfg.classes,\n                 decoder_attention_type= atte\n                )\n    \n    # Dataloader and learner\n    train_ds = RandomTileDataset(train_files, label_fn=label_fn, normalize=False, use_preprocessed_labels=True, **ds_kwargs, albumentations_tfms=tfms, tile_shape=tile_shape)\n    valid_ds = TileDataset(valid_files, label_fn=label_fn, normalize=False, use_preprocessed_labels=True, **ds_kwargs, is_zarr=True, tile_shape=tile_shape)\n    dls = DataLoaders.from_dsets(train_ds, valid_ds, bs=batch_size, after_item=_T, after_batch=Normalize.from_stats(*cfg.stats))\n    if torch.cuda.is_available(): dls.cuda(), model.cuda()\n    cbs = [SaveModelCallback(monitor='dice')]\n    learn = Learner(dls, model, metrics=cfg.metrics, wd=cfg.weight_decay, loss_func=cfg.loss_func, opt_func=ranger, cbs=cbs)\n    if cfg.mixed_precision_training: learn.to_fp16()\n    \n    # Fit\n    learn.fit_one_cycle(EPOCHS, lr_max=cfg.max_learning_rate)\n    learn.recorder.plot_metrics()\n    \n    plot_pedictions(find_best_and_worst())\n    ","metadata":{"execution":{"iopub.status.busy":"2022-03-13T17:46:48.995424Z","iopub.execute_input":"2022-03-13T17:46:48.995829Z","iopub.status.idle":"2022-03-13T19:07:14.709342Z","shell.execute_reply.started":"2022-03-13T17:46:48.995782Z","shell.execute_reply":"2022-03-13T19:07:14.708261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Save Model\nstate = {'model': learn.model.state_dict(), 'stats':cfg.stats}\ntorch.save(state, f'unet_{cfg.encoder_name}.pth', pickle_protocol=2, _use_new_zipfile_serialization=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}