{"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":"# EfficientNet Baseline on GPU Training\n\n## Reference\n\n- [HuBMAP fast.ai starter (EfficientNet) by wangkui](https://www.kaggle.com/code/befunny/hubmap-fast-ai-starter-efficientnet)\n- [Converting to 256x256 by The Devastator](https://www.kaggle.com/code/thedevastator/converting-to-256x256)\n- [[Training] - FastAI Baseline by The Devastator](https://www.kaggle.com/code/thedevastator/training-fastai-baseline)\n- [[Inference] - FastAI Baseline by The Devastator](https://www.kaggle.com/code/thedevastator/inference-fastai-baseline)","metadata":{}},{"cell_type":"code","source":"%reload_ext autoreload\n%autoreload 2\n%matplotlib inline\n\nfrom fastai.vision.all import *\nfrom torch.utils.data import Dataset, DataLoader\nimport pandas as pd\nimport numpy as np\nimport os\nimport cv2\nimport gc\nimport random\nfrom albumentations import *\nfrom sklearn.model_selection import KFold\nimport matplotlib.pyplot as plt\nfrom lovasz import lovasz_hinge\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class config:\n    bs = 32\n    nfolds = 4\n    SEED = 2020\n    TRAIN = '../input/hubmap-2022-256x256/train/'\n    MASKS = '../input/hubmap-2022-256x256/masks/'\n    LABELS = '../input/hubmap-organ-segmentation/train.csv'\n    NUM_WORKERS = 4\n    device = \"cuda:0\" if torch.cuda.is_available() else \"cpu\"\n    pretrained_root = '../input/efficientnet-pytorch/'\n    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\": 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\": 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\": 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\": 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\": 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\": 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\": 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\": pretrained_root + \"efficientnet-b7-dcc49843.pth\"\n        }\n    }\n    p = 1.0\n    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=p)\n    # device = \"cuda:0\" if torch.cuda.is_available() else \"cpu\"\n    device = \"cpu\" # do not use cuda for out of memory exception in Kaggle\n    head_epoch = 1\n    head_lr_max = 0.5e-2\n\n    full_epoch = 1\n    full_lr_max = slice(2e-4,2e-3)\n\n    model = 'efficientnet-b5'","metadata":{"trusted":true},"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 = True\n    \nseed_everything(config.SEED)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mean = np.array([0.7720342, 0.74582646, 0.76392896])\nstd = np.array([0.24745085, 0.26182273, 0.25782376])\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\nclass HuBMAPDataset(Dataset):\n    def __init__(self, fold=0, train=True, tfms=None):\n        ids = pd.read_csv(config.LABELS).id.astype(str).values\n        kf = KFold(n_splits=config.nfolds,random_state=config.SEED,shuffle=True)\n        ids = set(ids[list(kf.split(ids))[fold][0 if train else 1]])\n        self.fnames = [fname for fname in os.listdir(config.TRAIN) if fname.split('_')[0] in ids]\n        self.train = train\n        self.tfms = tfms\n        \n    def __len__(self):\n        return len(self.fnames)\n    \n    def __getitem__(self, idx):\n        fname = self.fnames[idx]\n        img = cv2.cvtColor(cv2.imread(os.path.join(config.TRAIN,fname)), cv2.COLOR_BGR2RGB)\n        mask = cv2.imread(os.path.join(config.MASKS,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 - mean)/std),img2tensor(mask)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ds = HuBMAPDataset(tfms=config.train_transform)\n# dl = DataLoader(ds,batch_size=64,shuffle=False,num_workers=config.NUM_WORKERS)\n# imgs,masks = next(iter(dl))\n\n# plt.figure(figsize=(16,16))\n# for i,(img,mask) in enumerate(zip(imgs,masks)):\n#     img = ((img.permute(1,2,0)*std + mean)*255.0).numpy().astype(np.uint8)\n#     plt.subplot(8,8,i+1)\n#     plt.imshow(img,vmin=0,vmax=255)\n#     plt.imshow(mask.squeeze().numpy(), alpha=0.2)\n#     plt.axis('off')\n#     plt.subplots_adjust(wspace=None, hspace=None)\n    \n# del ds,dl,imgs,masks","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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\nclass UnetBlock(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:Tensor, left_in:Tensor) -> 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        \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\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_()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('../input/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\n\nclass EfficientNetEncoder(EfficientNet):\n    def __init__(self, stage_idxs, out_channels, model_name, depth=5):\n\n        blocks_args, global_params = get_model_params(model_name, override_params=None)\n        super().__init__(blocks_args, global_params)\n        \n        cfg = config.efficient_net_encoders[model_name]\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, model_name, stride=1):\n        super().__init__()\n        \n        cfg = config.efficient_net_encoders[model_name]\n        stage_idxs = cfg['stage_idxs']\n        out_channels = cfg['out_channels']\n        \n        self.encoder = EfficientNetEncoder(stage_idxs, out_channels, model_name)\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    \n#split the model to encoder and decoder for fast.ai\nsplit_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            ]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def symmetric_lovasz(outputs, targets):\n    return 0.5*(lovasz_hinge(outputs, targets) + lovasz_hinge(-outputs, 1.0 - targets))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Dice_soft(Metric):\n    def __init__(self, axis=1): \n        self.axis = axis \n    def reset(self): self.inter,self.union = 0,0\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    @property\n    def value(self): return 2.0 * self.inter/self.union if self.union > 0 else None\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()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#iterator like wrapper that returns predicted and gt masks\nclass Model_pred:\n    def __init__(self, model, dl, tta:bool=True, half:bool=False):\n        self.model = model\n        self.dl = dl\n        self.tta = tta\n        self.half = half\n        \n    def __iter__(self):\n        self.model.eval()\n        name_list = self.dl.dataset.fnames\n        count=0\n        with torch.no_grad():\n            for x,y in iter(self.dl):\n                if config.device != \"cpu\":\n                    x = x.cuda()\n                if self.half: 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    \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    \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)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from fastai.callback.tracker import TrackerCallback\nfrom fastcore.basics import store_attr\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    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, # if true, save model after every epoch; else save only when model is better than existing best.\n        at_end=False, # 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 # 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:\n                print(f'Better model found at epoch {self.epoch} with {self.monitor} value: {self.best}.')\n                self._save(f'{self.fname}_{self.best:.4}')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dice = Dice_th_pred(np.arange(0.2,0.7,0.01))\nif not os.path.exists(\"models\"):\n        os.mkdir(\"models\")\n\nfor fold in range(config.nfolds):\n    if not os.path.exists(f\"models/fold_{fold}\"):\n        os.mkdir(f\"models/fold_{fold}\")\n    ds_t = HuBMAPDataset(fold=fold, train=True, tfms=config.train_transform)\n    ds_v = HuBMAPDataset(fold=fold, train=False)\n    data = ImageDataLoaders.from_dsets(ds_t,ds_v,bs=config.bs,\n                num_workers=config.NUM_WORKERS,pin_memory=True)\n    model = EffUnet(config.model)\n    \n    if config.device != \"cpu\":\n        data.to(config.device)\n        model.to(config.device)\n        \n    learn = Learner(data, model, loss_func=symmetric_lovasz,\n                metrics=[Dice_soft(),Dice_th()], \n                splitter=split_layers)\n    \n    learn.freeze_to(-1) #doesn't work\n    for param in learn.opt.param_groups[0]['params']:\n        param.requires_grad = False\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,\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    mp = Model_pred(learn.model,learn.dls.loaders[1])\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()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dices = dice.value\nnoise_ths = dice.ths\nbest_dice = dices.max()\nbest_thr = noise_ths[dices.argmax()]\nplt.figure(figsize=(8,4))\nplt.plot(noise_ths, dices, color='blue')\nplt.vlines(x=best_thr, ymin=dices.min(), ymax=dices.max(), colors='black')\nd = dices.max() - dices.min()\nplt.text(noise_ths[-1]-0.1, best_dice-0.1*d, f'DICE = {best_dice:.3f}', fontsize=12);\nplt.text(noise_ths[-1]-0.1, best_dice-0.2*d, f'TH = {best_thr:.3f}', fontsize=12);\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}