{"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":"# 🛠 Install Libraries","metadata":{}},{"cell_type":"code","source":"!pip install -q ../input/pytorch-segmentation-models-lib/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4\n!pip install -q ../input/pytorch-segmentation-models-lib/efficientnet_pytorch-0.6.3/efficientnet_pytorch-0.6.3\n!pip install -q ../input/pytorch-segmentation-models-lib/timm-0.4.12-py3-none-any.whl\n!pip install -q ../input/pytorch-segmentation-models-lib/segmentation_models_pytorch-0.2.0-py3-none-any.whl","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-12T04:51:29.434481Z","iopub.execute_input":"2022-07-12T04:51:29.435147Z","iopub.status.idle":"2022-07-12T04:53:23.583350Z","shell.execute_reply.started":"2022-07-12T04:51:29.435056Z","shell.execute_reply":"2022-07-12T04:53:23.582551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!pip uninstall -y scikit-learn\n#!pip install --pre --extra-index https://pypi.anaconda.org/scipy-wheels-nightly/simple scikit-learn","metadata":{"execution":{"iopub.status.busy":"2022-07-12T04:53:23.585435Z","iopub.execute_input":"2022-07-12T04:53:23.585714Z","iopub.status.idle":"2022-07-12T04:53:23.590784Z","shell.execute_reply.started":"2022-07-12T04:53:23.585676Z","shell.execute_reply":"2022-07-12T04:53:23.589672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pdb\nfrom tkinter.messagebox import NO\nimport cv2\nimport time\nimport glob\nimport random\n\nfrom cv2 import transform\nimport cupy as cp # https://cupy.dev/ => pip install cupy-cuda102\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom tqdm import tqdm\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\n#from sklearn.model_selection import StratifiedGroupKFold # Sklearn\nimport albumentations as A # Augmentations\nimport segmentation_models_pytorch as smp # smp\n","metadata":{"execution":{"iopub.status.busy":"2022-07-12T04:53:23.592346Z","iopub.execute_input":"2022-07-12T04:53:23.592720Z","iopub.status.idle":"2022-07-12T04:53:34.039011Z","shell.execute_reply.started":"2022-07-12T04:53:23.592686Z","shell.execute_reply":"2022-07-12T04:53:34.038137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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","metadata":{"execution":{"iopub.status.busy":"2022-07-12T04:53:34.041404Z","iopub.execute_input":"2022-07-12T04:53:34.041701Z","iopub.status.idle":"2022-07-12T04:53:34.047863Z","shell.execute_reply.started":"2022-07-12T04:53:34.041660Z","shell.execute_reply":"2022-07-12T04:53:34.046089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"###############################################################\n##### ***part0: data preprocess***\n###############################################################","metadata":{}},{"cell_type":"code","source":"def get_metadata(row):\n    data = row['id'].split('_')\n    case = int(data[0].replace('case',''))\n    day = int(data[1].replace('day',''))\n    slice_ = int(data[-1])\n    row['case'] = case\n    row['day'] = day\n    row['slice'] = slice_\n    return row\n\ndef path2info(row):\n    path = row['image_path']\n    data = path.split('/')\n    slice_ = int(data[-1].split('_')[1])\n    case = int(data[-3].split('_')[0].replace('case',''))\n    day = int(data[-3].split('_')[1].replace('day',''))\n    width = int(data[-1].split('_')[2])\n    height = int(data[-1].split('_')[3])\n    row['height'] = height\n    row['width'] = width\n    row['case'] = case\n    row['day'] = day\n    row['slice'] = slice_\n    # row['id'] = f'case{case}_day{day}_slice_{slice_}'\n    return row","metadata":{"execution":{"iopub.status.busy":"2022-07-12T04:53:34.048948Z","iopub.execute_input":"2022-07-12T04:53:34.049505Z","iopub.status.idle":"2022-07-12T04:53:34.060104Z","shell.execute_reply.started":"2022-07-12T04:53:34.049467Z","shell.execute_reply":"2022-07-12T04:53:34.059227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def mask2rle(msk, thr=0.5):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    msk    = cp.array(msk)\n    pixels = msk.flatten()\n    pad    = cp.array([0])\n    pixels = cp.concatenate([pad, pixels, pad])\n    runs   = cp.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\ndef masks2rles(msks, ids, heights, widths):\n    pred_strings = []; pred_ids = []; pred_classes = [];\n    for idx in range(msks.shape[0]):\n        height = heights[idx].item()\n        width = widths[idx].item()\n        msk = cv2.resize(msks[idx], \n                        dsize=(width, height), \n                        interpolation=cv2.INTER_NEAREST) # back to original shape\n        rle = [None]*3\n        for midx in [0, 1, 2]:\n            rle[midx] = mask2rle(msk[...,midx])\n        pred_strings.extend(rle)\n        pred_ids.extend([ids[idx]]*len(rle))\n        pred_classes.extend(['large_bowel', 'small_bowel', 'stomach'])\n    return pred_strings, pred_ids, pred_classes","metadata":{"execution":{"iopub.status.busy":"2022-07-12T04:53:34.061106Z","iopub.execute_input":"2022-07-12T04:53:34.061368Z","iopub.status.idle":"2022-07-12T04:53:34.073948Z","shell.execute_reply.started":"2022-07-12T04:53:34.061334Z","shell.execute_reply":"2022-07-12T04:53:34.073193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"###############################################################\n##### part1: build_transforms & build_dataset & build_dataloader\n###############################################################","metadata":{}},{"cell_type":"code","source":"def build_transforms(CFG):\n    data_transforms = {\n        \"train\": A.Compose([\n            A.OneOf([\n                A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST, p=1.0),\n            ], p=1),\n\n            A.HorizontalFlip(p=0.5),\n            # A.VerticalFlip(p=0.5),\n            A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.05, rotate_limit=10, p=0.5),\n            A.OneOf([\n                A.GridDistortion(num_steps=5, distort_limit=0.05, p=1.0),\n                # A.OpticalDistortion(distort_limit=0.05, shift_limit=0.05, p=1.0),\n                A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=1.0)\n            ], p=0.25),\n            A.CoarseDropout(max_holes=8, max_height=CFG.img_size[0]//20, max_width=CFG.img_size[1]//20,\n                            min_holes=5, fill_value=0, mask_fill_value=0, p=0.5),\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\nclass build_dataset(Dataset):\n    def __init__(self, df, label=True, transforms=None, cfg=None):\n        self.df = df\n        self.label = label\n        self.img_paths = df['image_path'].tolist() # image\n        self.ids = df['id'].tolist()\n\n        if 'mask_path' in df.columns:\n            self.mask_paths  = df['mask_path'].tolist() # mask\n        else:\n            self.mask_paths = None\n\n        self.transforms = transforms\n        self.n_25d_shift = cfg.n_25d_shift\n\n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        #### load id\n        id       = self.ids[index]\n        #### load image\n        img_path  = self.img_paths[index]\n        img = self.load_2_5d_slice(img_path) # [h, w, c]\n        h, w = img.shape[:2]\n        \n        if self.label: # train\n            #### load mask\n            mask_path = self.mask_paths[index]\n            mask = np.load(mask_path).astype('float32')\n            mask/=255.0 # scale mask to [0, 1]\n\n            ### augmentations\n            data = self.transforms(image=img, mask=mask)\n            img  = data['image']\n            mask  = data['mask']\n            img = np.transpose(img, (2, 0, 1)) # [h, w, c] => [c, h, w]\n            mask = np.transpose(mask, (2, 0, 1)) # [h, w, c] => [c, h, w]\n            return torch.tensor(img), torch.tensor(mask)\n        \n        else:  # test\n            ### augmentations\n            data = self.transforms(image=img)\n            img  = data['image']\n            img = np.transpose(img, (2, 0, 1)) # [h, w, c] => [c, h, w]\n            return torch.tensor(img), id, h, w\n\n    def load_2_5d_slice(self, middle_img_path):\n        #### 步骤1: 获取中间图片的基本信息\n        #### eg: middle_img_path: 'slice_0005_266_266_1.50_1.50.png' \n        middle_slice_num = os.path.basename(middle_img_path).split('_')[1] # eg: 0005\n        middle_str = 'slice_'+middle_slice_num\n\n        new_25d_imgs = []\n\n        ##### 步骤2：按照左右n_25d_shift数量进行填充，如果没有相应图片填充为Nan.\n        ##### 注：经过EDA发现同一天的所有患者图片的shape是一致的\n        for i in range(-self.n_25d_shift, self.n_25d_shift+1): # eg: i = {-2, -1, 0, 1, 2}\n            shift_slice_num = int(middle_slice_num) + i\n            shift_str = 'slice_'+str(shift_slice_num).zfill(4)\n            shift_img_path = middle_img_path.replace(middle_str, shift_str)\n            \n            if os.path.exists(shift_img_path):\n                shift_img = cv2.imread(shift_img_path, cv2.IMREAD_UNCHANGED) # [w, h]\n                new_25d_imgs.append(shift_img)\n            else:\n                new_25d_imgs.append(None)\n        \n        ##### 步骤3：从中心开始往外循环，依次填补None的值\n        ##### eg: n_25d_shift = 2, 那么形成5个channel, idx为[0, 1, 2, 3, 4], 所以依次处理的idx为[1, 3, 0, 4]\n        shift_left_idxs = []\n        shift_right_idxs = []\n        for related_idx in range(self.n_25d_shift):\n            shift_left_idxs.append(self.n_25d_shift - related_idx - 1)\n            shift_right_idxs.append(self.n_25d_shift + related_idx + 1)\n\n        for left_idx, right_idx in zip(shift_left_idxs, shift_right_idxs):\n            if new_25d_imgs[left_idx] is None:\n                new_25d_imgs[left_idx] = new_25d_imgs[left_idx+1]\n            if new_25d_imgs[right_idx] is None:\n                new_25d_imgs[right_idx] = new_25d_imgs[right_idx-1]\n\n        new_25d_imgs = np.stack(new_25d_imgs, axis=2).astype('float32') # [w, h, c]\n        mx_pixel = new_25d_imgs.max()\n        if mx_pixel != 0:\n            new_25d_imgs /= mx_pixel\n        return new_25d_imgs","metadata":{"execution":{"iopub.status.busy":"2022-07-12T04:53:34.075551Z","iopub.execute_input":"2022-07-12T04:53:34.076052Z","iopub.status.idle":"2022-07-12T04:53:34.098740Z","shell.execute_reply.started":"2022-07-12T04:53:34.076016Z","shell.execute_reply":"2022-07-12T04:53:34.098093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"###############################################################\n##### >>>>>>> trick: construct 2.5d slice images <<<<<<\n###############################################################","metadata":{}},{"cell_type":"code","source":"\ndef build_dataloader(df, fold, data_transforms, CFG):\n    train_df = df.query(\"fold!=@fold\").reset_index(drop=True)\n    valid_df = df.query(\"fold==@fold\").reset_index(drop=True)\n    train_dataset = build_dataset(train_df, label=True, transforms=data_transforms['train'], cfg=CFG)\n    valid_dataset = build_dataset(valid_df, label=True, transforms=data_transforms['valid_test'], cfg=CFG)\n\n    train_loader = DataLoader(train_dataset, batch_size=CFG.train_bs, num_workers=CFG.num_worker, shuffle=True, pin_memory=True, drop_last=False)\n    valid_loader = DataLoader(valid_dataset, batch_size=CFG.valid_bs, num_workers=CFG.num_worker, shuffle=False, pin_memory=True)\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2022-07-12T04:53:34.100485Z","iopub.execute_input":"2022-07-12T04:53:34.100722Z","iopub.status.idle":"2022-07-12T04:53:34.112318Z","shell.execute_reply.started":"2022-07-12T04:53:34.100678Z","shell.execute_reply":"2022-07-12T04:53:34.111647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"###############################################################\n##### >>>>>>> part2: build_model <<<<<<\n###############################################################","metadata":{}},{"cell_type":"code","source":"def 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=2*CFG.n_25d_shift+1,             \n            classes=CFG.num_classes,   \n            activation=None,\n        )\n    model.to(CFG.device)\n    return model\n","metadata":{"execution":{"iopub.status.busy":"2022-07-12T04:53:34.113479Z","iopub.execute_input":"2022-07-12T04:53:34.113888Z","iopub.status.idle":"2022-07-12T04:53:34.122570Z","shell.execute_reply.started":"2022-07-12T04:53:34.113851Z","shell.execute_reply":"2022-07-12T04:53:34.121650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"###############################################################\n##### >>>>>>> part3: build_loss <<<<<<\n###############################################################","metadata":{}},{"cell_type":"code","source":"def build_loss():\n    BCELoss     = smp.losses.SoftBCEWithLogitsLoss()\n    TverskyLoss = smp.losses.TverskyLoss(mode='multilabel', log_loss=False)\n    return {\"BCELoss\":BCELoss, \"TverskyLoss\":TverskyLoss}","metadata":{"execution":{"iopub.status.busy":"2022-07-12T04:53:34.126220Z","iopub.execute_input":"2022-07-12T04:53:34.126432Z","iopub.status.idle":"2022-07-12T04:53:34.132814Z","shell.execute_reply.started":"2022-07-12T04:53:34.126409Z","shell.execute_reply":"2022-07-12T04:53:34.131989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"###############################################################\n##### >>>>>>> part4: build_metric <<<<<<\n###############################################################","metadata":{}},{"cell_type":"code","source":"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\ndef iou_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    union = (y_true + y_pred - y_true*y_pred).sum(dim=dim)\n    iou = ((inter+epsilon)/(union+epsilon)).mean(dim=(1,0))\n    return iou","metadata":{"execution":{"iopub.status.busy":"2022-07-12T04:53:34.134008Z","iopub.execute_input":"2022-07-12T04:53:34.134381Z","iopub.status.idle":"2022-07-12T04:53:34.143988Z","shell.execute_reply.started":"2022-07-12T04:53:34.134346Z","shell.execute_reply":"2022-07-12T04:53:34.143209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"###############################################################\n##### >>>>>>> part5: train & validation & test <<<<<<\n###############################################################","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, train_loader, optimizer, losses_dict, CFG):\n    model.train()\n    scaler = amp.GradScaler() \n    losses_all, bce_all, tverskly_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            tverskly_loss = losses_dict[\"TverskyLoss\"](y_preds, masks)\n            losses = bce_loss + tverskly_loss\n        \n        scaler.scale(losses).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        losses_all += losses.item() / images.shape[0]\n        bce_all += bce_loss.item() / images.shape[0]\n        tverskly_all += tverskly_loss.item() / images.shape[0]\n    \n    current_lr = optimizer.param_groups[0]['lr']\n    print(\"lr: {:.4f}\".format(current_lr), flush=True)\n    print(\"loss: {:.3f}, bce_all: {:.3f}, tverskly_all: {:.3f}\".format(losses_all, bce_all, tverskly_all), flush=True)\n        \n@torch.no_grad()\ndef valid_one_epoch(model, valid_loader, CFG):\n    model.eval()\n    val_scores = []\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_dice = dice_coef(masks, y_preds).cpu().detach().numpy()\n        val_jaccard = iou_coef(masks, y_preds).cpu().detach().numpy()\n        val_scores.append([val_dice, val_jaccard])\n        \n    val_scores  = np.mean(val_scores, axis=0)\n    val_dice, val_jaccard = val_scores\n    print(\"val_dice: {:.4f}, val_jaccard: {:.4f}\".format(val_dice, val_jaccard), flush=True)\n    \n    return val_dice, val_jaccard\n\n@torch.no_grad()\ndef test_one_epoch(ckpt_paths, test_loader, CFG):\n    pred_strings = []\n    pred_ids = []\n    pred_classes = []\n    \n    pbar = tqdm(enumerate(test_loader), total=len(test_loader), desc='Test: ')\n    for _, (images, ids, h, w) in pbar:\n\n        images  = images.to(CFG.device, dtype=torch.float) # [b, c, w, h]\n        size = images.size()\n        masks = torch.zeros((size[0], 3, size[2], size[3]), device=CFG.device, dtype=torch.float32) # [b, c, w, h]\n        \n        ############################################\n        ##### >>>>>>> cross validation infer <<<<<<\n        ############################################\n        for sub_ckpt_path in ckpt_paths:\n            model = build_model(CFG, test_flag=True)\n            model.load_state_dict(torch.load(sub_ckpt_path))\n            model.eval()\n            y_preds = model(images) # [b, c, w, h]\n            y_preds   = torch.nn.Sigmoid()(y_preds)\n            masks += y_preds/len(ckpt_paths)\n        \n        masks = (masks.permute((0, 2, 3, 1))>CFG.thr).to(torch.uint8).cpu().detach().numpy() # [n, h, w, c]\n        result = masks2rles(masks, ids, h, w)\n        pred_strings.extend(result[0])\n        pred_ids.extend(result[1])\n        pred_classes.extend(result[2])\n    return pred_strings, pred_ids, pred_classes\n\n\nif __name__ == '__main__':\n    ###############################################################\n    ##### >>>>>>> config <<<<<<\n    ###############################################################\n    class CFG:\n        # step1: hyper-parameter\n        seed = 970301  # birthday\n        num_worker = 0 # debug => 0\n        device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n        ckpt_fold = \"cpktkinkinb7\"\n        ckpt_name = \"efficientnetb7_img224224_bs128_fold4\"\n        \n        # step2: data\n        n_25d_shift = 2\n        n_fold = 4\n        img_size = [224, 224]\n        train_bs = 16\n        valid_bs = train_bs * 2\n\n        # step3: model\n        backbone = 'efficientnet-b7'\n        num_classes = 3\n\n        # step4: optimizer\n        epoch = 12\n        lr = 1e-3\n        wd = 1e-5\n        lr_drop = 8\n\n        # step5: infer\n        thr = 0.45\n    \n    set_seed(CFG.seed)\n    ckpt_path = f\"../input/{CFG.ckpt_fold}/{CFG.ckpt_name}\"\n    if not os.path.exists(ckpt_path):\n        os.makedirs(ckpt_path)\n\n    train_val_flag = False\n    if train_val_flag:\n        ###############################################################\n        ##### part0: data preprocess\n        ###############################################################\n        # document: https://pandas.pydata.org/docs/reference/frame.html\n        df = pd.read_csv('kaggle/input/uwmgi-mask-dataset/train.csv')\n        df['segmentation'] = df.segmentation.fillna('') # .fillna(): 填充NaN的值为空\n        # rle mask length\n        df['rle_len'] = df.segmentation.map(len) # .map(): 特定列中的每一个元素应用一个函数len\n        # image/mask path\n        df['image_path'] = df.image_path.str.replace('/kaggle/', 'kaggle/') # .str: 特定列应用python字符串处理方法\n        df['mask_path'] = df.mask_path.str.replace('/kaggle/', 'kaggle/')\n        df['mask_path'] = df.mask_path.str.replace('/png/','/np').str.replace('.png','.npy')\n        # rle list of each id\n        df2 = df.groupby(['id'])['segmentation'].agg(list).to_frame().reset_index() # .grouby(): 特定列划分group.\n        # total length of all rles of each id\n        df2 = df2.merge(df.groupby(['id'])['rle_len'].agg(sum).to_frame().reset_index()) # .agg(): 特定列应用operations\n        df = df.drop(columns=['segmentation', 'class', 'rle_len']) # .drop(): 特定列的删除\n        df = df.groupby(['id']).head(1).reset_index(drop=True)\n        # empty mask\n        df = df.merge(df2, on=['id']) # .merge(): 特定列的合并\n        df['empty'] = (df.rle_len==0) \n\n        ###############################################################\n        ##### >>>>>>> trick1: cross validation train <<<<<<\n        ###############################################################\n        # document: http://scikit-learn.org/stable/modules/generated/sklearn.model_selection.StratifiedGroupKFold.html\n        skf = StratifiedGroupKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\n        for fold, (train_idx, val_idx) in enumerate(skf.split(df, df['empty'], groups = df[\"case\"])):\n            df.loc[val_idx, 'fold'] = fold\n        \n        for fold in range(CFG.n_fold):\n            print(f'#'*80, flush=True)\n            print(f'###### Fold: {fold}', flush=True)\n            print(f'#'*80, flush=True)\n\n            ###############################################################\n            ##### >>>>>>> step2: combination <<<<<<\n            ##### build_transforme() & build_dataset() & build_dataloader()\n            ##### build_model() & build_loss()\n            ###############################################################\n            data_transforms = build_transforms(CFG)  \n            train_loader, valid_loader = build_dataloader(df, fold, data_transforms, CFG) # dataset & dtaloader\n            model = build_model(CFG) # model\n            optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.wd)\n            lr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, CFG.lr_drop) \n            losses_dict = build_loss() # loss\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 <<<<<<\n                ###############################################################\n                train_one_epoch(model, train_loader, optimizer, losses_dict, CFG)\n                lr_scheduler.step()\n                val_dice, val_jaccard = valid_one_epoch(model, valid_loader, 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\"{ckpt_path}/best_fold{fold}.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    test_flag = True\n    if test_flag:\n        set_seed(CFG.seed)\n        ###############################################################\n        ##### part0: data preprocess\n        ###############################################################\n        sub_df = pd.read_csv('../input/uw-madison-gi-tract-image-segmentation/sample_submission.csv')\n        if not len(sub_df):\n            sub_firset = True\n            sub_df = pd.read_csv('../input/uw-madison-gi-tract-image-segmentation/train.csv')[:1000*3]\n            sub_df = sub_df.drop(columns=['class','segmentation']).drop_duplicates()\n            paths = glob(f'../input/uw-madison-gi-tract-image-segmentation/train/**/*png',recursive=True)\n        else:\n            sub_firset = False\n            sub_df = sub_df.drop(columns=['class','predicted']).drop_duplicates()\n            paths = glob(f'../input/uw-madison-gi-tract-image-segmentation/test/**/*png',recursive=True)\n        sub_df = sub_df.apply(get_metadata,axis=1)\n        path_df = pd.DataFrame(paths, columns=['image_path'])\n        path_df = path_df.apply(path2info, axis=1)\n        test_df = sub_df.merge(path_df, on=['case','day','slice'], how='left')\n\n        data_transforms = build_transforms(CFG)\n        test_dataset = build_dataset(test_df, label=False, transforms=data_transforms['valid_test'], cfg=CFG)\n        test_loader  = DataLoader(test_dataset, batch_size=CFG.valid_bs, num_workers=2, shuffle=False, pin_memory=False)\n\n        ###############################################################\n        ##### >>>>>>> step2: infer <<<<<<\n        ###############################################################\n        # attention: change the corresponding upload path to kaggle.\n        ckpt_paths  = glob(f'{ckpt_path}/*.pth')\n        assert len(ckpt_paths) == CFG.n_fold, \"ckpt path error!\"\n\n        pred_strings, pred_ids, pred_classes = test_one_epoch(ckpt_paths, test_loader, CFG)\n\n        ###############################################################\n        ##### step3: submit\n        ###############################################################\n        pred_df = pd.DataFrame({\n            \"id\":pred_ids,\n            \"class\":pred_classes,\n            \"predicted\":pred_strings\n        })\n        if not sub_firset:\n            sub_df = pd.read_csv('../input/uw-madison-gi-tract-image-segmentation/sample_submission.csv')\n            del sub_df['predicted']\n        else:\n            sub_df = pd.read_csv('../input/uw-madison-gi-tract-image-segmentation/train.csv')[:1000*3]\n            del sub_df['segmentation']\n            \n        sub_df = sub_df.merge(pred_df, on=['id','class'])\n        sub_df.to_csv('submission.csv',index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-12T04:54:31.371966Z","iopub.execute_input":"2022-07-12T04:54:31.372460Z","iopub.status.idle":"2022-07-12T04:59:27.286293Z","shell.execute_reply.started":"2022-07-12T04:54:31.372422Z","shell.execute_reply":"2022-07-12T04:59:27.285522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}