{"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":"code","source":"import argparse\nimport gc\nimport importlib\nimport sys\nimport os\nimport zipfile\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport torch\nimport numpy as np\nimport random\nimport warnings\nimport cv2\nfrom fastai.callback.progress import CSVLogger\nfrom fastai.optimizer import Adam\nfrom torch.utils.data import DataLoader\nfrom fastai.learner import Learner\nfrom fastai.vision.data import ImageDataLoaders\nfrom fastprogress import progress_bar\nimport glob","metadata":{"execution":{"iopub.status.busy":"2022-08-05T14:03:06.996936Z","iopub.execute_input":"2022-08-05T14:03:06.997293Z","iopub.status.idle":"2022-08-05T14:03:10.702734Z","shell.execute_reply.started":"2022-08-05T14:03:06.997212Z","shell.execute_reply":"2022-08-05T14:03:10.693200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from functools import partial\nfrom types import SimpleNamespace\nimport torch\nfrom albumentations import *\nfrom lovasz import lovasz_hinge\nfrom fastai.optimizer import SGD, Adam, QHAdam, OptimWrapper\n\n\ndef symmetric_lovasz(outputs, targets):\n    return 0.5 * (lovasz_hinge(outputs, targets) + lovasz_hinge(-outputs, 1.0 - targets))\n\n\ndef dice_loss(pred, target, smooth=1.):\n    pred = pred.contiguous()\n    target = target.contiguous()\n    intersection = (pred * target).sum(dim=2).sum(dim=2)\n    loss = (1 - ((2. * intersection + smooth) / (pred.sum(dim=2).sum(dim=2) + target.sum(dim=2).sum(dim=2) + smooth)))\n    return loss.mean()\n\n\nconfig = SimpleNamespace(**{})\n\nconfig.batch_size = 16\nconfig.nfolds = 5\nconfig.fold = 0\nconfig.SEED = 2020\n\nconfig.NUM_WORKERS = 4\n\nconfig.device = \"cuda:0\" if torch.cuda.is_available() else \"cpu\"\n# config.device = \"cpu\"\nif config.device != \"cpu\":\n    torch.cuda.set_device(config.device)\n\nconfig.pretrained_root = '../input/my-efficientnet-pytorch/'\n\nconfig.efficient_net_encoders = {\n    \"efficientnet-b0\": {\n        \"out_channels\": (3, 32, 24, 40, 112, 320),\n        \"stage_idxs\": (3, 5, 9, 16),\n        \"weight_path\": config.pretrained_root + \"efficientnet-b0-08094119.pth\"\n    },\n    \"efficientnet-b1\": {\n        \"out_channels\": (3, 32, 24, 40, 112, 320),\n        \"stage_idxs\": (5, 8, 16, 23),\n        \"weight_path\": config.pretrained_root + \"efficientnet-b1-dbc7070a.pth\"\n    },\n    \"efficientnet-b2\": {\n        \"out_channels\": (3, 32, 24, 48, 120, 352),\n        \"stage_idxs\": (5, 8, 16, 23),\n        \"weight_path\": config.pretrained_root + \"efficientnet-b2-27687264.pth\"\n    },\n    \"efficientnet-b3\": {\n        \"out_channels\": (3, 40, 32, 48, 136, 384),\n        \"stage_idxs\": (5, 8, 18, 26),\n        \"weight_path\": config.pretrained_root + \"efficientnet-b3-c8376fa2.pth\"\n    },\n    \"efficientnet-b4\": {\n        \"out_channels\": (3, 48, 32, 56, 160, 448),\n        \"stage_idxs\": (6, 10, 22, 32),\n        \"weight_path\": config.pretrained_root + \"efficientnet-b4-e116e8b3.pth\"\n    },\n    \"efficientnet-b5\": {\n        \"out_channels\": (3, 48, 40, 64, 176, 512),\n        \"stage_idxs\": (8, 13, 27, 39),\n        \"weight_path\": config.pretrained_root + \"efficientnet-b5-586e6cc6.pth\"\n    },\n    \"efficientnet-b6\": {\n        \"out_channels\": (3, 56, 40, 72, 200, 576),\n        \"stage_idxs\": (9, 15, 31, 45),\n        \"weight_path\": config.pretrained_root + \"efficientnet-b6-c76e70fd.pth\"\n    },\n    \"efficientnet-b7\": {\n        \"out_channels\": (3, 64, 48, 80, 224, 640),\n        \"stage_idxs\": (11, 18, 38, 55),\n        \"weight_path\": config.pretrained_root + \"efficientnet-b7-dcc49843.pth\"\n    }\n}\n\nconfig.train_p = 1.0\nconfig.train_transform = Compose([\n    HorizontalFlip(p=0.5),\n    VerticalFlip(),\n    RandomRotate90(p=1),\n    # Morphology\n    ShiftScaleRotate(shift_limit=0, scale_limit=(-0.2, 0.2), rotate_limit=(-30, 30),\n                     interpolation=1, border_mode=0, value=(0, 0, 0), p=0.5),\n    GaussNoise(var_limit=(0, 50.0), mean=0, p=0.5),\n    GaussianBlur(blur_limit=(3, 7), p=0.5),\n    # Color\n    RandomBrightnessContrast(brightness_limit=0.35, contrast_limit=0.5,\n                             brightness_by_max=True, p=0.5),\n    HueSaturationValue(hue_shift_limit=30, sat_shift_limit=30,\n                       val_shift_limit=0, p=0.5),\n    OneOf([\n        OpticalDistortion(p=0.3),\n        GridDistortion(p=.1),\n        IAAPiecewiseAffine(p=0.3),\n    ], p=0.3),\n], p=config.train_p)\n\nconfig.val_p = 0.4\nconfig.val_transform = Compose([\n    RandomRotate90(p=0.2),\n\n    RandomBrightnessContrast(brightness_limit=0.35, contrast_limit=0.5,\n                             brightness_by_max=True, p=0.2),\n    HueSaturationValue(hue_shift_limit=30, sat_shift_limit=30,\n                       val_shift_limit=0, p=0.3),\n], p=config.val_p)\n\n# config.train_transform_list = [\n#     HorizontalFlip(p=1),\n#     VerticalFlip(),\n#     RandomRotate90(p=1),\n#     # Morphology\n#     ShiftScaleRotate(shift_limit=0, scale_limit=(-0.2, 0.2), rotate_limit=(-30, 30),\n#                      interpolation=1, border_mode=0, value=(0, 0, 0), p=1),\n#     GaussNoise(var_limit=(0, 50.0), mean=0, p=1),\n#     GaussianBlur(blur_limit=(3, 7), p=1),\n#     # Color\n#     RandomBrightnessContrast(brightness_limit=0.35, contrast_limit=0.5,\n#                              brightness_by_max=True, p=1),\n#     HueSaturationValue(hue_shift_limit=30, sat_shift_limit=30,\n#                        val_shift_limit=0, p=1),\n#     OneOf([\n#         OpticalDistortion(p=1),\n#         GridDistortion(p=1),\n#         IAAPiecewiseAffine(p=1),\n#     ], p=1),\n# ]\n#\n#\n# config.val_transform_list = [\n#     RandomRotate90(p=1),\n#\n#     RandomBrightnessContrast(brightness_limit=0.35, contrast_limit=0.5,\n#                                brightness_by_max=True, p=1),\n#\n#     HueSaturationValue(hue_shift_limit=30, sat_shift_limit=30,\n#                        val_shift_limit=0, p=1),\n# ]\n\nconfig.weight = None\n\nconfig.load_best_weight = False\n\nconfig.head_epoch = 6\nconfig.head_lr_max = 0.5e-4\n\nconfig.full_epoch = 2000\nconfig.full_lr_max = slice(2e-5, 2e-4)\n\nconfig.model = 'efficientnet-b7'\n\nconfig.loss = symmetric_lovasz\n\n# config.optimizer = partial(OptimWrapper, opt=torch.optim.AdamW)\nconfig.optimizer = Adam\n\nconfig.freeze_layer = -1\n\nconfig.image_size = 256\n\nconfig.train_dataset = \"hap\" # only all, hap, hubmap\n\nconfig.dice_dataset = \"hap\" # only all, hap, hubmap\n\nconfig.only_dice = 0\n","metadata":{"execution":{"iopub.status.busy":"2022-08-05T14:05:50.724895Z","iopub.execute_input":"2022-08-05T14:05:50.725416Z","iopub.status.idle":"2022-08-05T14:05:50.764069Z","shell.execute_reply.started":"2022-08-05T14:05:50.725371Z","shell.execute_reply":"2022-08-05T14:05:50.762903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nfrom fastai.callback.core import Callback\nfrom fastai.callback.tracker import TrackerCallback\nfrom fastcore.basics import store_attr, range_of\nimport os\n\nfrom matplotlib import pyplot as plt\n\n\nclass MySaveModelCallback(TrackerCallback):\n    \"A `TrackerCallback` that saves the model's best during training and loads it at the end.\"\n    order = TrackerCallback.order + 1\n\n    def __init__(self,\n                 monitor='valid_loss',  # value (usually loss or metric) being monitored.\n                 comp=None,  # numpy comparison operator; np.less if monitor is loss, np.greater if monitor is metric.\n                 min_delta=0.,  # minimum delta between the last monitor value and the best monitor value.\n                 fname='model',  # model name to be used when saving model.\n                 every_epoch=False,\n                 # if true, save model after every epoch; else save only when model is better than existing best.\n                 at_end=False,\n                 # if true, save model when training ends; else load best model if there is only one saved model.\n                 with_opt=False,  # if true, save optimizer state (if any available) when saving model.\n                 reset_on_fit=True\n                 # before model fitting, reset value being monitored to -infinity (if monitor is metric) or +infinity (if monitor is loss).\n                 ):\n        super().__init__(monitor=monitor, comp=comp, min_delta=min_delta, reset_on_fit=reset_on_fit)\n        assert not (every_epoch and at_end), \"every_epoch and at_end cannot both be set to True\"\n        # keep track of file path for loggers\n        self.last_saved_path = None\n        store_attr('fname,every_epoch,at_end,with_opt')\n\n    def _save(self, name):\n        self.last_saved_path = self.learn.save(name, with_opt=self.with_opt)\n\n    def after_epoch(self):\n        \"Compare the value monitored to its best score and save if best.\"\n        if self.every_epoch:\n            if (self.epoch % self.every_epoch) == 0: self._save(f'{self.fname}_{self.epoch}')\n        else:  # every improvement\n            super().after_epoch()\n            if self.new_best and self.best != 0:\n                print(f'Better model found at epoch {self.epoch} with {self.monitor} value: {self.best}.')\n                os.system(f\"rm models/{self.fname}*.pth\")\n                self._save(f'{self.fname}_{self.best:.4}')\n\n\n","metadata":{"execution":{"iopub.status.busy":"2022-08-05T14:03:12.387314Z","iopub.execute_input":"2022-08-05T14:03:12.387977Z","iopub.status.idle":"2022-08-05T14:03:12.403136Z","shell.execute_reply.started":"2022-08-05T14:03:12.387939Z","shell.execute_reply":"2022-08-05T14:03:12.402161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from albumentations import Compose\nfrom torch.utils.data import Dataset\nimport os\nimport cv2\nimport pandas as pd\nimport numpy as np\nfrom sklearn.model_selection import KFold\nimport torch\nimport glob\n\n\ndef img2tensor(img, dtype: np.dtype = np.float32):\n    if img.ndim == 2: img = np.expand_dims(img, 2)\n    img = np.transpose(img, (2, 0, 1))\n    return torch.from_numpy(img.astype(dtype, copy=False))\n\n\nclass TrainDataset(Dataset):\n    def __init__(self, train=True, config=None):\n        if config.train_dataset == \"hap\":\n            ids = glob.glob(\"../input/hubmap-2022-256x256/train/*.png\")\n        elif config.train_dataset == \"hubmap\":\n            ids = glob.glob(\"../hubmap-256x256/train/*.png\")\n        else:\n            ids = glob.glob(\"../all_256/train/*.png\")\n        kf = KFold(n_splits=config.nfolds, random_state=config.SEED, shuffle=True)\n        self.fnames = []\n        if train:\n            for fold, i in enumerate(kf.split(ids)):\n                if fold == config.fold:\n                    self.fnames.extend([ids[id] for id in i[0]])\n        else:\n            for fold, i in enumerate(kf.split(ids)):\n                if fold == config.fold:\n                    self.fnames.extend([ids[id] for id in i[1]])\n        self.train = train\n        if train:\n            self.tfms = config.train_transform\n        else:\n            self.tfms = config.val_transform\n        self.config = config\n\n    def __len__(self):\n        return len(self.fnames)\n\n    def __getitem__(self, idx):\n        image_fname = self.fnames[idx]\n        img = cv2.cvtColor(cv2.imread(image_fname), cv2.COLOR_BGR2RGB)\n        mask_fname = (image_fname.replace(\"train\", \"masks\", 1)).replace(\"train\", \"mask\", 1)\n        mask = cv2.imread(mask_fname, cv2.IMREAD_GRAYSCALE)\n        if self.tfms is not None:\n            augmented = self.tfms(image=img, mask=mask)\n            img, mask = augmented['image'], augmented['mask']\n        return img2tensor(img / 255.0), img2tensor(mask)\n\n\nclass DiceDataset(Dataset):\n    def __init__(self, config=None):\n        if config.dice_dataset == \"hap\":\n            ids = glob.glob(\"../input/hubmap-2022-256x256/train/*.png\")\n        elif config.dice_dataset == \"hubmap\":\n            ids = glob.glob(\"../hubmap-256x256/train/*.png\")\n        else:\n            ids = glob.glob(\"../all_256/train/*.png\")\n        self.fnames = ids\n        self.config = config\n\n    def __len__(self):\n        return len(self.fnames)\n\n    def __getitem__(self, idx):\n        image_fname = self.fnames[idx]\n        img = cv2.cvtColor(cv2.imread(image_fname), cv2.COLOR_BGR2RGB)\n        mask_fname = (image_fname.replace(\"train\", \"masks\", 1)).replace(\"train\", \"mask\", 1)\n        mask = cv2.imread(mask_fname, cv2.IMREAD_GRAYSCALE)\n        return img2tensor(img / 255.0), img2tensor(mask)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-05T14:03:12.407247Z","iopub.execute_input":"2022-08-05T14:03:12.407655Z","iopub.status.idle":"2022-08-05T14:03:12.426349Z","shell.execute_reply.started":"2022-08-05T14:03:12.407628Z","shell.execute_reply":"2022-08-05T14:03:12.425014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nfrom fastai.vision.all import *\nimport sys\nimport torch.nn.functional as F\n\nfrom fastai.layers import ConvLayer, SelfAttention, PixelShuffle_ICNR\n\nsys.path.insert(0, '../input/my-efficientnet-pytorch/EfficientNet-PyTorch/EfficientNet-PyTorch/EfficientNet-PyTorch-master')\n\nfrom efficientnet_pytorch import EfficientNet\nfrom efficientnet_pytorch.utils import url_map, url_map_advprop, get_model_params\n\nclass FPN(nn.Module):\n    def __init__(self, input_channels: list, output_channels: list):\n        super().__init__()\n        self.convs = nn.ModuleList(\n            [nn.Sequential(nn.Conv2d(in_ch, out_ch * 2, kernel_size=3, padding=1),\n                           nn.ReLU(inplace=True), nn.BatchNorm2d(out_ch * 2),\n                           nn.Conv2d(out_ch * 2, out_ch, kernel_size=3, padding=1))\n             for in_ch, out_ch in zip(input_channels, output_channels)])\n\n    def forward(self, xs: list, last_layer):\n        hcs = [F.interpolate(c(x), scale_factor=2 ** (len(self.convs) - i), mode='bilinear')\n               for i, (c, x) in enumerate(zip(self.convs, xs))]\n        hcs.append(last_layer)\n        return torch.cat(hcs, dim=1)\n\n\nclass UnetBlock(nn.Module):\n    def __init__(self, up_in_c: int, x_in_c: int, nf: int = None, blur: bool = False,\n                 self_attention: bool = False, **kwargs):\n        super().__init__()\n        self.shuf = PixelShuffle_ICNR(up_in_c, up_in_c // 2, blur=blur, **kwargs)\n        self.bn = nn.BatchNorm2d(x_in_c)\n        ni = up_in_c // 2 + x_in_c\n        nf = nf if nf is not None else max(up_in_c // 2, 32)\n        self.conv1 = ConvLayer(ni, nf, norm_type=None, **kwargs)\n        self.conv2 = ConvLayer(nf, nf, norm_type=None,\n                               xtra=SelfAttention(nf) if self_attention else None, **kwargs)\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, up_in: torch.Tensor, left_in: torch.Tensor) -> torch.Tensor:\n        s = left_in\n        up_out = self.shuf(up_in)\n        cat_x = self.relu(torch.cat([up_out, self.bn(s)], dim=1))\n        return self.conv2(self.conv1(cat_x))\n\n\nclass _ASPPModule(nn.Module):\n    def __init__(self, inplanes, planes, kernel_size, padding, dilation, groups=1):\n        super().__init__()\n        self.atrous_conv = nn.Conv2d(inplanes, planes, kernel_size=kernel_size,\n                                     stride=1, padding=padding, dilation=dilation, bias=False, groups=groups)\n        self.bn = nn.BatchNorm2d(planes)\n        self.relu = nn.ReLU()\n\n        self._init_weight()\n\n    def forward(self, x):\n        x = self.atrous_conv(x)\n        x = self.bn(x)\n\n        return self.relu(x)\n\n    def _init_weight(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                torch.nn.init.kaiming_normal_(m.weight)\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n\n\nclass ASPP(nn.Module):\n    def __init__(self, inplanes=512, mid_c=256, dilations=[6, 12, 18, 24], out_c=None):\n        super().__init__()\n        self.aspps = [_ASPPModule(inplanes, mid_c, 1, padding=0, dilation=1)] + \\\n                     [_ASPPModule(inplanes, mid_c, 3, padding=d, dilation=d, groups=4) for d in dilations]\n        self.aspps = nn.ModuleList(self.aspps)\n        self.global_pool = nn.Sequential(nn.AdaptiveMaxPool2d((1, 1)),\n                                         nn.Conv2d(inplanes, mid_c, 1, stride=1, bias=False),\n                                         nn.BatchNorm2d(mid_c), nn.ReLU())\n        out_c = out_c if out_c is not None else mid_c\n        self.out_conv = nn.Sequential(nn.Conv2d(mid_c * (2 + len(dilations)), out_c, 1, bias=False),\n                                      nn.BatchNorm2d(out_c), nn.ReLU(inplace=True))\n        self.conv1 = nn.Conv2d(mid_c * (2 + len(dilations)), out_c, 1, bias=False)\n        self._init_weight()\n\n    def forward(self, x):\n        x0 = self.global_pool(x)\n        xs = [aspp(x) for aspp in self.aspps]\n        x0 = F.interpolate(x0, size=xs[0].size()[2:], mode='bilinear', align_corners=True)\n        x = torch.cat([x0] + xs, dim=1)\n        return self.out_conv(x)\n\n    def _init_weight(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                torch.nn.init.kaiming_normal_(m.weight)\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n\n\nclass EfficientNetEncoder(EfficientNet):\n    def __init__(self, stage_idxs, out_channels, depth=5, config=None):\n\n        blocks_args, global_params = get_model_params(config.model, override_params=None, image_size = config.image_size)\n        super().__init__(blocks_args, global_params)\n\n        super().__init__(blocks_args, global_params)\n\n        cfg = config.efficient_net_encoders[config.model]\n\n        self._stage_idxs = stage_idxs\n        self._out_channels = out_channels\n        self._depth = depth\n        self._in_channels = 3\n\n        del self._fc\n        self.load_state_dict(torch.load(cfg['weight_path']))\n\n    def get_stages(self):\n        return [\n            nn.Identity(),\n            nn.Sequential(self._conv_stem, self._bn0, self._swish),\n            self._blocks[:self._stage_idxs[0]],\n            self._blocks[self._stage_idxs[0]:self._stage_idxs[1]],\n            self._blocks[self._stage_idxs[1]:self._stage_idxs[2]],\n            self._blocks[self._stage_idxs[2]:],\n        ]\n\n    def forward(self, x):\n        stages = self.get_stages()\n\n        block_number = 0.\n        drop_connect_rate = self._global_params.drop_connect_rate\n\n        features = []\n        for i in range(self._depth + 1):\n\n            # Identity and Sequential stages\n            if i < 2:\n                x = stages[i](x)\n\n            # Block stages need drop_connect rate\n            else:\n                for module in stages[i]:\n                    drop_connect = drop_connect_rate * block_number / len(self._blocks)\n                    block_number += 1.\n                    x = module(x, drop_connect)\n\n            features.append(x)\n\n        return features\n\n    def load_state_dict(self, state_dict, **kwargs):\n        state_dict.pop(\"_fc.bias\")\n        state_dict.pop(\"_fc.weight\")\n        super().load_state_dict(state_dict, **kwargs)\n\n\nclass EffUnet(nn.Module):\n    def __init__(self, stride=1,config=None):\n        super().__init__()\n\n        cfg = config.efficient_net_encoders[config.model]\n        stage_idxs = cfg['stage_idxs']\n        out_channels = cfg['out_channels']\n\n        self.encoder = EfficientNetEncoder(stage_idxs, out_channels, config=config)\n\n        # aspp with customized dilatations\n        self.aspp = ASPP(out_channels[-1], 256, out_c=384,\n                         dilations=[stride * 1, stride * 2, stride * 3, stride * 4])\n        self.drop_aspp = nn.Dropout2d(0.5)\n        # decoder\n        self.dec4 = UnetBlock(384, out_channels[-2], 256)\n        self.dec3 = UnetBlock(256, out_channels[-3], 128)\n        self.dec2 = UnetBlock(128, out_channels[-4], 64)\n        self.dec1 = UnetBlock(64, out_channels[-5], 32)\n        self.fpn = FPN([384, 256, 128, 64], [16] * 4)\n        self.drop = nn.Dropout2d(0.1)\n        self.final_conv = ConvLayer(32 + 16 * 4, 1, ks=1, norm_type=None, act_cls=None)\n\n    def forward(self, x):\n        enc0, enc1, enc2, enc3, enc4 = self.encoder(x)[-5:]\n        enc5 = self.aspp(enc4)\n        dec3 = self.dec4(self.drop_aspp(enc5), enc3)\n        dec2 = self.dec3(dec3, enc2)\n        dec1 = self.dec2(dec2, enc1)\n        dec0 = self.dec1(dec1, enc0)\n        x = self.fpn([enc5, dec3, dec2, dec1], dec0)\n        x = self.final_conv(self.drop(x))\n        x = F.interpolate(x, scale_factor=2, mode='bilinear')\n        return x\n","metadata":{"execution":{"iopub.status.busy":"2022-08-05T14:03:42.480246Z","iopub.execute_input":"2022-08-05T14:03:42.480666Z","iopub.status.idle":"2022-08-05T14:03:42.545847Z","shell.execute_reply.started":"2022-08-05T14:03:42.480633Z","shell.execute_reply":"2022-08-05T14:03:42.544932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from fastai.learner import Metric\nfrom fastai.torch_core import flatten_check\nimport torch\nimport torch.nn.functional as F\nimport numpy as np\n\n\nclass Dice_soft(Metric):\n    def __init__(self, axis=1):\n        self.axis = axis\n\n    def reset(self): self.inter, self.union = 0, 0\n\n    def accumulate(self, learn):\n        pred, targ = flatten_check(torch.sigmoid(learn.pred), learn.y)\n        self.inter += (pred * targ).float().sum().item()\n        self.union += (pred + targ).float().sum().item()\n\n    @property\n    def value(self): return 2.0 * self.inter / self.union if self.union > 0 else None\n\n\n# dice with automatic threshold selection\nclass Dice_th(Metric):\n    def __init__(self, ths=np.arange(0.1, 0.9, 0.05), axis=1):\n        self.axis = axis\n        self.ths = ths\n\n    def reset(self):\n        self.inter = torch.zeros(len(self.ths))\n        self.union = torch.zeros(len(self.ths))\n\n    def accumulate(self, learn):\n        pred, targ = flatten_check(torch.sigmoid(learn.pred), learn.y)\n        for i, th in enumerate(self.ths):\n            p = (pred > th).float()\n            self.inter[i] += (p * targ).float().sum().item()\n            self.union[i] += (p + targ).float().sum().item()\n\n    @property\n    def value(self):\n        dices = torch.where(self.union > 0.0,\n                            2.0 * self.inter / self.union, torch.zeros_like(self.union))\n        return dices.max()\n\n\nclass Model_pred:\n    def __init__(self, model, dl, tta: bool = True, half: bool = False, config=None):\n        self.model = model\n        self.dl = dl\n        self.tta = tta\n        self.half = half\n        self.config = config\n\n    def __iter__(self):\n        self.model.eval()\n        name_list = [i[i.rindex(\"/\"):] for i in self.dl.dataset.fnames]\n        count = 0\n        with torch.no_grad():\n            for x, y in iter(self.dl):\n                if self.config.device != \"cpu\":\n                    x = x.to(self.config.device)\n                if self.half:\n                    x = x.half()\n                p = self.model(x)\n                py = torch.sigmoid(p).detach()\n                if self.tta:\n                    # x,y,xy flips as TTA\n                    flips = [[-1], [-2], [-2, -1]]\n                    for f in flips:\n                        p = self.model(torch.flip(x, f))\n                        p = torch.flip(p, f)\n                        py += torch.sigmoid(p).detach()\n                    py /= (1 + len(flips))\n                if y is not None and len(y.shape) == 4 and py.shape != y.shape:\n                    py = F.upsample(py, size=(y.shape[-2], y.shape[-1]), mode=\"bilinear\")\n                py = py.permute(0, 2, 3, 1).float().cpu()\n                batch_size = len(py)\n                for i in range(batch_size):\n                    taget = y[i].detach().cpu() if y is not None else None\n                    yield py[i], taget, name_list[count]\n                    count += 1\n\n    def __len__(self):\n        return len(self.dl.dataset)\n\n\nclass Dice_th_pred(Metric):\n    def __init__(self, ths=np.arange(0.1, 0.9, 0.01), axis=1):\n        self.axis = axis\n        self.ths = ths\n        self.reset()\n\n    def reset(self):\n        self.inter = torch.zeros(len(self.ths))\n        self.union = torch.zeros(len(self.ths))\n\n    def accumulate(self, p, t):\n        pred, targ = flatten_check(p, t)\n        for i, th in enumerate(self.ths):\n            p = (pred > th).float()\n            self.inter[i] += (p * targ).float().sum().item()\n            self.union[i] += (p + targ).float().sum().item()\n\n    @property\n    def value(self):\n        dices = torch.where(self.union > 0.0, 2.0 * self.inter / self.union,\n                            torch.zeros_like(self.union))\n        return dices\n","metadata":{"execution":{"iopub.status.busy":"2022-08-05T14:03:42.682625Z","iopub.execute_input":"2022-08-05T14:03:42.683009Z","iopub.status.idle":"2022-08-05T14:03:42.707840Z","shell.execute_reply.started":"2022-08-05T14:03:42.682975Z","shell.execute_reply":"2022-08-05T14:03:42.706877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(config):\n    random.seed(config.SEED)\n    os.environ['PYTHONHASHSEED'] = str(config.SEED)\n    np.random.seed(config.SEED)\n    torch.manual_seed(config.SEED)\n    if config.device != \"cpu\":\n        torch.cuda.manual_seed(config.SEED)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\n\ndef save_img(data, name, out):\n    data = data.float().cpu().numpy()\n    img = cv2.imencode('.png', (data * 255).astype(np.uint8))[1]\n    out.writestr(name, img)\n\n\ndef score(weight_path):\n    score_lindex = weight_path.rindex(\"_\") + 1\n    score_rindex = weight_path.rindex(\".\")\n    return float(weight_path[score_lindex:score_rindex])\n\n\ndef main(config):\n    fold = config.fold\n\n    if not os.path.exists(f\"models/fold_{config.fold}\"):\n        os.mkdir(f\"models/fold_{config.fold}\")\n\n    split_layers = lambda m: [\n        list(m.encoder.parameters()),\n        list(m.aspp.parameters()) + list(m.dec4.parameters()) +\n        list(m.dec3.parameters()) + list(m.dec2.parameters()) +\n        list(m.dec1.parameters()) + list(m.fpn.parameters()) +\n        list(m.final_conv.parameters())\n    ]\n    dice = Dice_th_pred(np.arange(0.2, 0.7, 0.01))\n\n    ds_t = TrainDataset(train=True, config=config)\n    ds_v = TrainDataset(train=False, config=config)\n    data = ImageDataLoaders.from_dsets(ds_t, ds_v, bs=config.batch_size,\n                                       num_workers=config.NUM_WORKERS, pin_memory=True).to(config.device)\n    model = EffUnet(config=config).to(config.device)\n\n    if config.load_best_weight:\n        models = glob.glob(f\"models/fold_{config.fold}/{config.model}*_fold_{config.fold}*.pth\")\n        models = sorted(models, key=lambda i: score(i), reverse=True)\n        state_dict = torch.load(models[0], map_location=torch.device(config.device))\n        print(\"Load Pretrained Model: \" + models[0])\n        model.load_state_dict(state_dict)\n\n    if config.weight is not None:\n        state_dict = torch.load(config.weight, map_location=torch.device(config.device))\n        print(\"Load Pretrained Model: \" + config.weight)\n        model.load_state_dict(state_dict)\n\n    learn = Learner(data, model, loss_func=config.loss, opt_func=config.optimizer,\n                    metrics=[Dice_soft(), Dice_th()],\n                    splitter=split_layers).to_fp16()\n\n\n    if config.only_dice != 1:\n        # start with training the head\n        learn.freeze_to(config.freeze_layer)\n        for param in learn.opt.param_groups[0]['params']:\n            param.requires_grad = False\n\n        learn.fit_one_cycle(config.head_epoch, lr_max=config.head_lr_max)\n\n        # continue training full model\n        learn.unfreeze()\n        learn.fit_one_cycle(config.full_epoch, lr_max=config.full_lr_max, reset_opt=True,\n                            cbs=[MySaveModelCallback(monitor='dice_th', comp=np.greater,\n                                                     fname=f'fold_{fold}/{config.model}_fold_{fold}',\n                                                     every_epoch=False,\n                                                     at_end=True),\n                                 CSVLogger(fname=f\"models/fold_{fold}/history_fold_{fold}.csv\")])\n\n        dataframe_metrix = pd.read_csv(f\"models/fold_{fold}/history_fold_{fold}.csv\")\n\n        plt.subplot(1, 2, 1, frameon=False)\n        plt.title(f'fold_{fold}_train_loss')\n        plt.xlabel('Epoch')\n        plt.plot(dataframe_metrix[\"epoch\"], dataframe_metrix[\"train_loss\"], \"r\")\n\n        plt.subplot(1, 2, 2, frameon=False)\n        plt.title(f'fold_{fold}_test_dice')\n        plt.xlabel('Epoch')\n        plt.plot(dataframe_metrix[\"epoch\"], dataframe_metrix[\"dice_th\"], \"b\")\n\n        plt.savefig(f\"models/fold_{fold}/metric_fold_{fold}.jpg\")\n        plt.close()\n        best_weight_path = glob.glob(f\"models/fold_{fold}/{config.model}*.pth\")[0]\n        model.load_state_dict(torch.load(best_weight_path, map_location=torch.device(config.device)))\n\n    dice_dataset = DiceDataset(config=config)\n    dice_dataloader = DataLoader(dice_dataset, shuffle=True, batch_size=1, num_workers=config.NUM_WORKERS, pin_memory=True)\n\n    # model evaluation on val and saving the masks\n    mp = Model_pred(model, dice_dataloader, config=config)\n    # with zipfile.ZipFile('val_masks_tta.zip', 'a') as out:\n    for p in progress_bar(mp):\n        dice.accumulate(p[0], p[1])\n    # save_img(p[0], p[2], out)\n    gc.collect()\n    dices = dice.value\n    noise_ths = dice.ths\n    best_dice = dices.max()\n    best_thr = noise_ths[dices.argmax()]\n    plt.figure(figsize=(8, 4))\n    plt.plot(noise_ths, dices, color='blue')\n    plt.vlines(x=best_thr, ymin=dices.min(), ymax=dices.max(), colors='black')\n    d = dices.max() - dices.min()\n    plt.text(noise_ths[-1] - 0.1, best_dice - 0.1 * d, f'DICE = {best_dice:.3f}', fontsize=12)\n    plt.text(noise_ths[-1] - 0.1, best_dice - 0.2 * d, f'TH = {best_thr:.3f}', fontsize=12)\n    plt.savefig(f'models/fold_{fold}/save.jpg')\n    plt.close()\n\n    best_weight_path = glob.glob(f\"models/fold_{fold}/{config.model}*.pth\")[0]\n    down_index = best_weight_path.index(\"_fold_\")\n    new_weight_path = best_weight_path[:down_index] + f\"_{best_thr:.3f}\" + best_weight_path[down_index:]\n    os.rename(best_weight_path, new_weight_path)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T14:12:57.256592Z","iopub.execute_input":"2022-08-05T14:12:57.258078Z","iopub.status.idle":"2022-08-05T14:12:57.290494Z","shell.execute_reply.started":"2022-08-05T14:12:57.258034Z","shell.execute_reply":"2022-08-05T14:12:57.289173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! mkdir models","metadata":{"execution":{"iopub.status.busy":"2022-08-05T14:04:05.163907Z","iopub.execute_input":"2022-08-05T14:04:05.164497Z","iopub.status.idle":"2022-08-05T14:04:06.362974Z","shell.execute_reply.started":"2022-08-05T14:04:05.164452Z","shell.execute_reply":"2022-08-05T14:04:06.361552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"main(config)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T14:13:00.896737Z","iopub.execute_input":"2022-08-05T14:13:00.897226Z","iopub.status.idle":"2022-08-05T14:28:56.281198Z","shell.execute_reply.started":"2022-08-05T14:13:00.897183Z","shell.execute_reply":"2022-08-05T14:28:56.280118Z"},"trusted":true},"execution_count":null,"outputs":[]}]}