{"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 PyTorch ⚡ Train","metadata":{}},{"cell_type":"markdown","source":"# Installs","metadata":{}},{"cell_type":"code","source":"!pip install segmentation_models_pytorch\n\n","metadata":{"execution":{"iopub.status.busy":"2022-08-05T10:04:37.487335Z","iopub.execute_input":"2022-08-05T10:04:37.488051Z","iopub.status.idle":"2022-08-05T10:04:56.947511Z","shell.execute_reply.started":"2022-08-05T10:04:37.488006Z","shell.execute_reply":"2022-08-05T10:04:56.945643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"###############################################################\n##### @Title:  HuBMAP baseline\n##### @Time:  2022/07/29\n##### @Author: frank\n##### @Struct: \n        #  part0: data preprocess\n        #  part1: build_transforme() & build_dataset() & build_dataloader()\n        #  part2: build_model()\n        #  part3: build_loss()\n        #  part4: build_metric()\n        #  part5: train_one_epoch() & valid_one_epoch() & test_one_epoch()\n##### @Describe: \n        # The Devastator - \"hubmap-2022-256x256\" 作者\n        # \n##### @Reference:\n        # [Training] - FastAI Baseline: https://www.kaggle.com/code/thedevastator/training-fastai-baseline\n        # [Inference] - FastAI Baselin: https://www.kaggle.com/code/thedevastator/inference-fastai-baseline\n###############################################################\nimport os\nimport pdb\nimport cv2\nimport time\nimport glob\nimport random\n\nimport rasterio\nfrom rasterio.windows import Window\nimport tifffile\n\nfrom cv2 import transform\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\n\nfrom fastai.vision.all import *\n\nimport torch # PyTorch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda import amp # https://pytorch.org/docs/stable/notes/amp_examples.html\n\nfrom sklearn.model_selection import StratifiedGroupKFold, KFold # Sklearn\nimport albumentations as A # Augmentations\n\nimport segmentation_models_pytorch as smp # smp\nsys.path.append('../resnet_unet/')\nfrom resnext_unet import *\n\ndef set_seed(seed=42):\n    ##### why 42? The Answer to the Ultimate Question of Life, the Universe, and Everything is 42.\n    random.seed(seed) # python\n    np.random.seed(seed) # numpy\n    torch.manual_seed(seed) # pytorch\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n###############################################################\n##### part0: data preprocess\n###############################################################\n# functions to convert encoding to mask and mask to encoding\ndef enc2mask(encs, shape):\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for m,enc in enumerate(encs):\n        if isinstance(enc,np.float) and np.isnan(enc): continue\n        s = enc.split()\n        for i in range(len(s)//2):\n            start = int(s[2*i]) - 1\n            length = int(s[2*i+1])\n            img[start:start+length] = 1 + m\n    return img.reshape(shape).T\n\ndef mask2enc(mask, n=1):\n    pixels = mask.T.flatten()\n    encs = []\n    for i in range(1,n+1):\n        p = (pixels == i).astype(np.int8)\n        if p.sum() == 0: encs.append(np.nan)\n        else:\n            p = np.concatenate([[0], p, [0]])\n            runs = np.where(p[1:] != p[:-1])[0] + 1\n            runs[1::2] -= runs[::2]\n            encs.append(' '.join(str(x) for x in runs))\n    return encs\n\n#https://www.kaggle.com/bguberfain/memory-aware-rle-encoding\n#with transposed mask\ndef rle_encode_less_memory(img):\n    #the image should be transposed\n    pixels = img.T.flatten()\n    \n    # This simplified method requires first and last pixel to be zero\n    pixels[0] = 0\n    pixels[-1] = 0\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    runs[1::2] -= runs[::2]\n    \n    return ' '.join(str(x) for x in runs)\n\n\ndef plot_visual(image, mask, pred, fold, idx, cmap):\n    \n    image = ((image.transpose(1,2,0)*std + mean)*255.0).astype(np.uint8)\n    mask = mask.squeeze()\n    pred = pred.squeeze()\n    \n    plt.figure(figsize=(16, 10))\n\n    plt.subplot(1, 3, 1)\n    plt.imshow(image, vmin=0, vmax=255)\n    plt.title(\"image\", fontsize=10)\n    plt.axis(\"off\")\n\n    plt.subplot(1, 3, 2)\n    plt.imshow(image, vmin=0, vmax=255)\n    plt.imshow(mask, cmap=cmap, alpha=0.5)\n    plt.title(f\"mask\", fontsize=10)    \n    plt.axis(\"off\")\n    \n    plt.subplot(1, 3, 3)\n    plt.imshow(image, vmin=0, vmax=255)\n    plt.imshow(pred, cmap=cmap, alpha=0.5)\n    plt.title(f\"pred\", fontsize=10)    \n    plt.axis(\"off\")\n\n    plt.savefig(f\"./visual/{fold}/{idx}.png\")\n    plt.close()\n    \n###############################################################\n##### part1: build_transforms & build_dataset & build_dataloader\n###############################################################\ndef build_transforms(CFG):\n    data_transforms = {\n        \"train\": A.Compose([\n            A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, border_mode=cv2.BORDER_REFLECT),\n            A.OneOf([\n                A.OpticalDistortion(p=0.3),\n                A.GridDistortion(p=.1),\n                # IAAPiecewiseAffine(p=0.3),\n            ], p=0.3),\n            \n            A.OneOf([\n                A.HueSaturationValue(10,15,10),\n                A.CLAHE(clip_limit=2),\n                A.RandomBrightnessContrast(),            \n            ], p=0.3),\n            ], p=1.0),\n        \n        \"valid_test\": A.Compose([\n            A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n            ], p=1.0)\n        }\n    return data_transforms\n\n\n# https://www.kaggle.com/datasets/thedevastator/hubmap-2022-256x256\nmean = 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 build_dataset(Dataset):\n    def __init__(self, df=None, label=True, transforms=None, CFG=None, idx = None, sz=None, reduce=reduce):\n        self.label = label\n        if self.label:\n            ###########################################\n            ##### >>>>>>> Use \"hubmap-2022-256x256\" Dataset <<<<<<\n            ############################################\n            self.df = df\n            ids = df.id.astype(str).values\n            self.third_data_train = os.path.join(CFG.third_data_path, \"train\")\n            self.third_data_mask = os.path.join(CFG.third_data_path, \"masks\")\n            self.file_names = [file_name for file_name in os.listdir(self.third_data_train) if file_name.split('_')[0] in ids]\n            self.organ_to_label = {'kidney' : 0, 'prostate' : 1, 'largeintestine' : 2, 'spleen' : 3, 'lung' : 4}        \n            self.label = label\n            self.transforms = transforms\n        else:\n            ###########################################\n            ##### >>>>>>> Use Original Dataset <<<<<<\n            ############################################            \n            self.original_data_train = os.path.join(CFG.data_path, \"test_images\")\n            self.data = rasterio.open(os.path.join(self.original_data_train, idx+'.tiff'), transform = identity,\n                                 num_threads='all_cpus')\n            # some images have issues with their format \n            # and must be saved correctly before reading with rasterio\n            if self.data.count != 3:\n                subdatasets = self.data.subdatasets\n                self.layers = []\n                if len(subdatasets) > 0:\n                    for i, subdataset in enumerate(subdatasets, 0):\n                        self.layers.append(rasterio.open(subdataset))\n            self.shape = self.data.shape\n            self.reduce = reduce\n            self.sz = reduce*sz\n            self.pad0 = (self.sz - self.shape[0]%self.sz)%self.sz\n            self.pad1 = (self.sz - self.shape[1]%self.sz)%self.sz\n            self.n0max = (self.shape[0] + self.pad0)//self.sz\n            self.n1max = (self.shape[1] + self.pad1)//self.sz\n        \n            \n    def __len__(self):\n        if self.label:\n            return len(self.df)\n        else:\n            return self.n0max*self.n1max\n    \n    def __getitem__(self, idx):\n        \n        if self.label:\n            file_name = self.file_names[idx]\n            img = cv2.cvtColor(cv2.imread(os.path.join(self.third_data_train, file_name)), cv2.COLOR_BGR2RGB)\n            mask = cv2.imread(os.path.join(self.third_data_mask, file_name),cv2.IMREAD_GRAYSCALE)\n            \n            if self.transforms:\n                data = self.transforms(image=img, mask=mask)\n                img  = data['image']\n                mask  = data['mask']\n            \n            return img2tensor((img/255.0 - mean)/std), img2tensor(mask)\n        else:\n            # the code below may be a little bit difficult to understand,\n            # but the thing it does is mapping the original image to\n            # tiles created with adding padding, as done in\n            # https://www.kaggle.com/iafoss/256x256-images ,\n            # and then the tiles are loaded with rasterio\n            # n0,n1 - are the x and y index of the tile (idx = n0*self.n1max + n1)\n            n0,n1 = idx//self.n1max, idx%self.n1max\n            # x0,y0 - are the coordinates of the lower left corner of the tile in the image\n            # negative numbers correspond to padding (which must not be loaded)\n            x0,y0 = -self.pad0//2 + n0*self.sz, -self.pad1//2 + n1*self.sz\n            # make sure that the region to read is within the image\n            p00,p01 = max(0,x0), min(x0+self.sz,self.shape[0])\n            p10,p11 = max(0,y0), min(y0+self.sz,self.shape[1])\n            img = np.zeros((self.sz,self.sz,3),np.uint8)\n            # mapping the loade region to the tile\n            if self.data.count == 3:\n                img[(p00-x0):(p01-x0),(p10-y0):(p11-y0)] = np.moveaxis(self.data.read([1,2,3],\n                    window=Window.from_slices((p00,p01),(p10,p11))), 0, -1)\n            else:\n                for i,layer in enumerate(self.layers):\n                    img[(p00-x0):(p01-x0),(p10-y0):(p11-y0),i] =\\\n                    layer.read(1,window=Window.from_slices((p00,p01),(p10,p11)))\n            \n            if self.reduce != 1:\n                img = cv2.resize(img,(self.sz//reduce,self.sz//reduce),\n                                interpolation = cv2.INTER_AREA)\n            #check for empty imges\n            hsv = cv2.cvtColor(img, cv2.COLOR_BGR2HSV)\n            h,s,v = cv2.split(hsv)\n            if (s>s_th).sum() <= p_th or img.sum() <= p_th:\n                #images with -1 will be skipped\n                return img2tensor((img/255.0 - mean)/std), -1\n            else: return img2tensor((img/255.0 - mean)/std), idx\n        \n        \n\ndef build_dataset_dataloader(df, fold, data_transforms, CFG):\n    train_df = df[df.fold != fold].reset_index(drop=True)\n    valid_df = df[df.fold == fold].reset_index(drop=True)\n    train_dataset = build_dataset(df = train_df, label=True, transforms=data_transforms['train'], CFG=CFG)\n    valid_dataset = build_dataset(df = valid_df, label=True, transforms=data_transforms['valid_test'], CFG=CFG)\n    train_loader = DataLoader(train_dataset, batch_size=CFG.train_bs, num_workers=CFG.num_worker, \n                              shuffle=True, pin_memory=True, drop_last=False)\n    valid_loader = DataLoader(valid_dataset, batch_size=CFG.valid_bs, num_workers=CFG.num_worker, \n                              shuffle=False, pin_memory=True)\n    return train_loader, valid_loader\n\n###############################################################\n##### >>>>>>> part2: build_model <<<<<<\n###############################################################\n# document: https://smp.readthedocs.io/en/latest/encoders_timm.html\ndef build_model(CFG, test_flag=False):\n    if test_flag:\n        pretrain_weights = None\n    else:\n        pretrain_weights = \"imagenet\"\n    # model = smp.Unet(\n    #         encoder_name=CFG.backbone,\n    #         encoder_weights=pretrain_weights, \n    #         in_channels=3,             \n    #         classes=CFG.num_classes,   \n    #         activation=None,\n    #     )\n    model = UneXt50()\n    model.to(CFG.device)\n    return model\n\n###############################################################\n##### >>>>>>> part3: build_loss <<<<<<\n###############################################################\ndef build_loss():\n    # BCELoss     = smp.losses.SoftBCEWithLogitsLoss()\n    BCELoss     = nn.BCEWithLogitsLoss()\n    DiceLoss    = smp.losses.DiceLoss(mode='binary')\n    return {\"BCELoss\":BCELoss, \"DiceLoss\":DiceLoss}\n\n###############################################################\n##### >>>>>>> part4: build_metric <<<<<<\n###############################################################\n# def dice_coef(y_true, y_pred, thr=0.5, dim=(2,3), epsilon=0.001):\n#     y_true = y_true.to(torch.float32)\n#     y_pred = (y_pred>thr).to(torch.float32)\n#     inter = (y_true*y_pred).sum(dim=dim)\n#     den = y_true.sum(dim=dim) + y_pred.sum(dim=dim)\n#     dice = ((2*inter+epsilon)/(den+epsilon)).mean(dim=(1,0))\n#     return dice\n    \nclass 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, pred, targ):\n        pred = torch.sigmoid(pred)\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###############################################################\n##### >>>>>>> part5: train & validation & test <<<<<<\n###############################################################\ndef train_one_epoch(model, train_loader, optimizer, losses_dict, CFG):\n    model.train()\n    scaler = amp.GradScaler() \n    losses_all, bce_all, dice_all = 0, 0, 0\n    \n    pbar = tqdm(enumerate(train_loader), total=len(train_loader), desc='Train ')\n    for _, (images, masks) in pbar:\n        optimizer.zero_grad()\n\n        images = images.to(CFG.device, dtype=torch.float) # [b, c, w, h]\n        masks  = masks.to(CFG.device, dtype=torch.float)  # [b, c, w, h]\n        \n        with amp.autocast(enabled=True):\n            y_preds = model(images) # [b, c, w, h]\n            \n            bce_loss = losses_dict[\"BCELoss\"](y_preds, masks)\n            # dice_loss = 0.3 * losses_dict[\"DiceLoss\"](y_preds, masks)\n            losses = bce_loss # + dice_loss\n        \n        scaler.scale(losses).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        losses_all += losses.item() \n        bce_all += bce_loss.item()\n        # dice_all += dice_loss.item() \n    \n    current_lr = optimizer.param_groups[0]['lr']\n    print(\"lr: {:.5f}\".format(current_lr), flush=True)\n    print(\"loss: {:.3f}, bce_all: {:.3f}, dice_all: {:.3f}\".format(losses_all/len(train_loader), bce_all/len(train_loader), dice_all), flush=True)\n        \n@torch.no_grad()\ndef valid_one_epoch(model, valid_loader, metric, CFG):\n    model.eval()\n    metric.reset()\n    \n    val_images = []\n    val_masks = []\n    val_preds = []\n    \n    pbar = tqdm(enumerate(valid_loader), total=len(valid_loader), desc='Valid ')\n    for _, (images, masks) in pbar:\n        images  = images.to(CFG.device, dtype=torch.float) # [b, c, w, h]\n        masks   = masks.to(CFG.device, dtype=torch.float)  # [b, c, w, h]\n        \n        y_preds = model(images) \n        y_preds   = torch.nn.Sigmoid()(y_preds) # [b, c, w, h]\n        \n        val_masks.append(masks)\n        val_preds.append(y_preds)\n        val_images.append(images)\n        metric.accumulate(y_preds.detach(), masks)\n    \n    val_images = torch.cat(val_images)\n    val_masks = torch.cat(val_masks)\n    val_preds = torch.cat(val_preds)\n    val_images = val_images.cpu().numpy()\n    val_masks = val_masks.cpu().numpy()\n    val_preds = val_preds.cpu().numpy()\n    \n    val_dice = metric.value\n    print(\"val_dice: {:.4f}\".format(val_dice), flush=True)\n\n    return val_dice, val_images, val_masks, val_preds\n\n\n#iterator like wrapper that returns predicted masks\nclass test_one_epoch:\n    def __init__(self, models, dl, tta:bool=True, half:bool=False, CFG=None):\n        self.models = models\n        self.dl = dl\n        self.tta = tta\n        self.half = half\n        \n    def __iter__(self):\n        count=0\n        with torch.no_grad():\n            for x,y in iter(self.dl):\n                if ((y>=0).sum() > 0): #exclude empty images\n                    x = x[y>=0].to(CFG.device)\n                    y = y[y>=0]\n                    if self.half: x = x.half()\n                    py = None\n                    for model in self.models:\n                        p = model(x)\n                        p = torch.sigmoid(p).detach()\n                        if py is None: py = p\n                        else: py += p\n                    if self.tta:\n                        #x,y,xy flips as TTA\n                        flips = [[-1],[-2],[-2,-1]]\n                        for f in flips:\n                            xf = torch.flip(x,f)\n                            for model in self.models:\n                                p = model(xf)\n                                p = torch.flip(p,f)\n                                py += torch.sigmoid(p).detach()\n                        py /= (1+len(flips))        \n                    py /= len(self.models) # [bs, 1, 256, 256]\n                    py = F.upsample(py, scale_factor=CFG.reduce, mode=\"bilinear\")\n                    py = py.permute(0,2,3,1).float().cpu()\n                    \n                    batch_size = len(py)\n                    for i in range(batch_size):\n                        yield py[i],y[i]\n                        count += 1\n                    \n    def __len__(self):\n        return len(self.dl.dataset)\n    \n\nif __name__ == '__main__':\n    ###############################################################\n    ##### >>>>>>> config <<<<<<\n    ###############################################################\n    class CFG:\n        # step1: hyper-parameter\n        seed = 42 \n        device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n        num_worker = 16 # if debug\n        data_path = \"../input/hubmap-organ-segmentation\"\n        third_data_path = \"../input/hubmap-2022-256x256\"\n        ckpt_path = \"../input/ckpt-frank/resnet18_img512_bs8_fold4\" # for submit\n        # step2: data\n        n_fold = 4\n        img_size = [512, 512]\n        train_bs = 8\n        valid_bs = train_bs * 2\n        # step3: model\n        backbone = 'resnet18'\n        num_classes = 1\n        # step4: optimizer\n        epoch = 50\n        lr = 1e-4\n        wd = 1e-5\n        lr_drop = 30\n        # step5: infer\n        thr = 0.3\n        reduce = 4\n        sz = 256\n        s_th = 40  #saturation blancking threshold\n        p_th = 1000*(sz//256)**2 #threshold for the minimum number of pixels\n    \n    set_seed(CFG.seed)\n    if not os.path.exists(CFG.ckpt_path):\n        os.makedirs(CFG.ckpt_path)\n\n    train_val_flag = True\n    if train_val_flag:\n        ###############################################################\n        ##### part0: data preprocess\n        ###############################################################\n        df = pd.read_csv(os.path.join(CFG.data_path, \"train.csv\"))\n\n        ###############################################################\n        ##### >>>>>>> trick1: cross validation train <<<<<<\n        ###############################################################\n        kf = KFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\n        df.loc[:,'fold'] = -1\n        for fold, (train_idx, val_idx) in enumerate(kf.split(X=df['id'], y=df['organ'])):\n            df.iloc[val_idx, -1] = fold\n        \n        for fold in range(CFG.n_fold):\n            print(f'#'*40, flush=True)\n            print(f'###### Fold: {fold}', flush=True)\n            print(f'#'*40, flush=True)\n\n            ###############################################################\n            ##### >>>>>>> step2: combination <<<<<<\n            ###############################################################\n            data_transforms = build_transforms(CFG)  \n            train_loader, valid_loader = build_dataset_dataloader(df, fold, data_transforms, CFG) # dataset & dtaloader\n            \n            model = build_model(CFG) # model\n            optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.wd) # optimizer\n            lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, CFG.lr_drop) \n            losses_dict = build_loss() # loss\n            metric=Dice_soft()\n            \n            best_val_dice = 0\n            best_epoch = 0\n            \n            for epoch in range(1, CFG.epoch+1):\n                start_time = time.time()\n                ###############################################################\n                ##### >>>>>>> step3: train & val <<<<<<d\n                ###############################################################\n                train_one_epoch(model, train_loader, optimizer, losses_dict, CFG)\n                lr_scheduler.step()\n                val_dice, _, _, _ = valid_one_epoch(model, valid_loader, metric, CFG)\n                \n                ###############################################################\n                ##### >>>>>>> step4: save best model <<<<<<\n                ###############################################################\n                is_best = (val_dice > best_val_dice)\n                best_val_dice = max(best_val_dice, val_dice)\n                if is_best:\n                    save_path = f\"{CFG.ckpt_path}/best_fold{fold}_epoch{epoch}_dice{best_val_dice:.4f}.pth\"\n                    if os.path.isfile(save_path):\n                        os.remove(save_path) \n                    torch.save(model.state_dict(), save_path)\n                \n                epoch_time = time.time() - start_time\n                print(\"epoch:{}, time:{:.2f}s, best:{:.2f}\\n\".format(epoch, epoch_time, best_val_dice), flush=True)\n    \n    \n    visual_flag = False\n    if visual_flag:\n        ###############################################################\n        ##### part0: data preprocess\n        ###############################################################\n        df = pd.read_csv(os.path.join(CFG.data_path, \"train.csv\"))\n\n        kf = KFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\n        df.loc[:,'fold'] = -1\n        for fold, (train_idx, val_idx) in enumerate(kf.split(X=df['id'], y=df['organ'])):\n            df.iloc[val_idx, -1] = fold\n        \n        ###############################################################\n        ##### >>>>>>> step2: infer & visual <<<<<<\n        ###############################################################\n        # only save best checkpoint\n        ##### cd ../input/ckpt-frank/resnet18_img512_bs8_fold2\n        ##### cp best_fold0_epoch62_dice0.6791.pth best_flag_fold0_epoch62_dice0.6791.pth\n        ##### cp best_fold1_epoch80_dice0.6604.pth best_flag_fold1_epoch80_dice0.6604.pth\n        ##### rm -rf best_fold*\n        for fold in range(CFG.n_fold):\n            data_transforms = build_transforms(CFG)  \n            train_loader, valid_loader = build_dataset_dataloader(df, fold, data_transforms, CFG) # dataset & dtaloader & transform\n            model = build_model(CFG, test_flag=True) # model\n            \n            sub_ckpt_path = \"../input/ckpt-frank/resnet18_img512_bs8_fold2/best_flag_fold1_epoch80_dice0.6604.pth\"\n            model.load_state_dict(torch.load(sub_ckpt_path))\n            model.eval()\n            val_dice, val_images, val_masks, val_preds = valid_one_epoch(model, valid_loader, CFG)\n            \n            for idx in range(val_images.shape[0]):\n                plot_visual(val_images[idx], val_masks[idx], val_preds[idx], fold, idx, \"bwr\")\n                \n        pdb.set_trace()\n        \n        \n    test_flag = False\n    if test_flag:\n        set_seed(CFG.seed)\n        \n        ###############################################################\n        ##### part0: load model\n        ###############################################################\n        # attention: change the corresponding upload path to kaggle!!!!!\n        ckpt_paths  = glob.glob(f'{CFG.ckpt_path}/best_flag_*')\n        assert len(ckpt_paths) == CFG.n_fold, \"ckpt path error!\"\n        \n        models = []\n        for sub_ckpt_path in ckpt_paths:\n            state_dict = torch.load(sub_ckpt_path, map_location=torch.device('cpu'))\n            model = smp.Unet(\n                encoder_name='resnet18',\n                encoder_weights=None, \n                in_channels=3,             \n                classes=1,   \n                activation=None,\n            )\n            model.load_state_dict(state_dict)\n            model.float()\n            model.eval()\n            model.to(CFG.device)\n            models.append(model)\n        \n        ###############################################################\n        ##### part0: data preprocess\n        ###############################################################\n        df_sample = pd.read_csv(os.path.join(CFG.data_path, \"sample_submission.csv\"))\n        \n        \n        sz = 256    # the size of tiles\n        reduce = 4  # reduce the original images by 4 times\n        TH = 0.225  # threshold for positive predictions\n        s_th = 40  #saturation blancking threshold\n        p_th = 1000*(sz//256)**2 #threshold for the minimum number of pixels\n        identity = rasterio.Affine(1, 0, 0, 0, 1, 0)\n\n        names,preds = [],[]\n        for idx,row in tqdm(df_sample.iterrows(),total=len(df_sample)):\n            idx = str(row['id'])\n            ds = build_dataset(idx=idx, label=False, sz=sz, reduce=reduce, CFG=CFG)\n            #rasterio cannot be used with multiple workers\n            dl = DataLoader(ds, CFG.valid_bs, num_workers=0, shuffle=False, pin_memory=True)\n            mp = test_one_epoch(models,dl)\n            #generate masks\n            mask = torch.zeros(len(ds),ds.sz,ds.sz,dtype=torch.int8)\n            for p,i in iter(mp): mask[i.item()] = p.squeeze(-1) > TH\n            \n            #reshape tiled masks into a single mask and crop padding\n            mask = mask.view(ds.n0max,ds.n1max,ds.sz,ds.sz).\\\n                permute(0,2,1,3).reshape(ds.n0max*ds.sz,ds.n1max*ds.sz)\n            mask = mask[ds.pad0//2:-(ds.pad0-ds.pad0//2) if ds.pad0 > 0 else ds.n0max*ds.sz,\n                ds.pad1//2:-(ds.pad1-ds.pad1//2) if ds.pad1 > 0 else ds.n1max*ds.sz]\n            \n            #convert to rle\n            #https://www.kaggle.com/bguberfain/memory-aware-rle-encoding\n            rle = rle_encode_less_memory(mask.numpy())\n            names.append(idx)\n            preds.append(rle)\n      \n\n        df = pd.DataFrame({'id':names,'rle':preds})\n        df.to_csv('submission.csv',index=False)\n\n\n       ","metadata":{"execution":{"iopub.status.busy":"2022-08-05T10:04:56.955245Z","iopub.execute_input":"2022-08-05T10:04:56.958217Z","iopub.status.idle":"2022-08-05T10:04:57.18968Z","shell.execute_reply.started":"2022-08-05T10:04:56.958168Z","shell.execute_reply":"2022-08-05T10:04:57.186666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}