{"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":"<div style=\"height:200px;width:100%;margin: 0;\">\n    <img src=\"https://storage.googleapis.com/kaggle-competitions/kaggle/34547/logos/header.png?t=2022-02-15-22-37-27\" style=\"width:100%;\" />\n</div>","metadata":{}},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"additionals\"><center>Installs</center></h3>","metadata":{}},{"cell_type":"code","source":"! pip install segmentation_models_pytorch","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:35.179623Z","iopub.execute_input":"2022-08-25T10:33:35.180093Z","iopub.status.idle":"2022-08-25T10:33:45.912933Z","shell.execute_reply.started":"2022-08-25T10:33:35.179986Z","shell.execute_reply":"2022-08-25T10:33:45.911685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"additionals\"><center>Imports</center></h3>","metadata":{}},{"cell_type":"code","source":"import os\nimport pdb\nimport cv2\nimport time\nimport glob\nimport random\nimport tifffile\n\nfrom cv2 import transform\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\nfrom sklearn.model_selection import StratifiedGroupKFold, KFold # Sklearn\nimport albumentations as A # Augmentations\nimport segmentation_models_pytorch as smp # smp","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:45.915474Z","iopub.execute_input":"2022-08-25T10:33:45.916117Z","iopub.status.idle":"2022-08-25T10:33:50.425460Z","shell.execute_reply.started":"2022-08-25T10:33:45.916075Z","shell.execute_reply":"2022-08-25T10:33:50.424238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"additionals\"><center>Configurations</center></h3>","metadata":{}},{"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-08-25T10:33:50.427744Z","iopub.execute_input":"2022-08-25T10:33:50.429089Z","iopub.status.idle":"2022-08-25T10:33:50.437053Z","shell.execute_reply.started":"2022-08-25T10:33:50.429031Z","shell.execute_reply":"2022-08-25T10:33:50.434766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_tiff(path, scale=None, verbose=0): \n    \"\"\"\n    Input: tiff file\n    Output: ndarray of shape (3000 * 3000 * 3)\n    \n    \"\"\"\n    image = tifffile.imread(path)\n    if len(image.shape) == 5:\n        image = image.squeeze().transpose(1, 2, 0)\n    \n    if verbose:\n        print(f\"[{path}] Image shape: {image.shape}\")\n    \n    if scale:\n        new_size = (image.shape[1] // scale, image.shape[0] // scale)\n        image = cv2.resize(image, new_size)\n        if verbose:\n            print(f\"[{path}] Resized Image shape: {image.shape}\")\n        \n    mx = np.max(image)\n    image = image.astype(np.float32)\n    if mx:\n        image /= mx # scale image to [0, 1]\n    return image\n\n# tiff_test_file = \"../input/hubmap-organ-segmentation/train_images/10044.tiff\"\ndef rle_decode(mask_rle, shape):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height,width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n    '''\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape).T  # Needed to align to RLE direction\n\ndef rle_encode(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.T.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:50.443869Z","iopub.execute_input":"2022-08-25T10:33:50.444383Z","iopub.status.idle":"2022-08-25T10:33:50.457816Z","shell.execute_reply.started":"2022-08-25T10:33:50.444335Z","shell.execute_reply":"2022-08-25T10:33:50.456175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"additionals\"><center>Preprocessing</center></h3>","metadata":{}},{"cell_type":"code","source":"def do_random_flip(image, mask):\n    if np.random.rand()>0.5:\n        image = cv2.flip(image,0)\n        mask = cv2.flip(mask,0)\n    if np.random.rand()>0.5:\n        image = cv2.flip(image,1)\n        mask = cv2.flip(mask,1)\n    if np.random.rand()>0.5:\n        image = image.transpose(1,0,2)\n        mask = mask.transpose(1,0)\n    \n    image = np.ascontiguousarray(image)\n    mask = np.ascontiguousarray(mask)\n    return image, mask\n\ndef do_random_rot90(image, mask):\n    r = np.random.choice([\n        0,\n        cv2.ROTATE_90_CLOCKWISE,\n        cv2.ROTATE_90_COUNTERCLOCKWISE,\n        cv2.ROTATE_180,\n    ])\n    if r==0:\n        return image, mask\n    else:\n        image = cv2.rotate(image, r)\n        mask = cv2.rotate(mask, r)\n        return image, mask\n    \ndef do_random_contast(image, mask, mag=0.3):\n    alpha = 1 + random.uniform(-1,1)*mag\n    image = image * alpha\n    image = np.clip(image,0,1)\n    return image, mask\n\ndef do_random_hsv(image, mask, mag=[0.15,0.25,0.25]):\n    image = (image*255).astype(np.uint8)\n    hsv = cv2.cvtColor(image, cv2.COLOR_BGR2HSV)\n\n    h = hsv[:, :, 0].astype(np.float32)  # hue\n    s = hsv[:, :, 1].astype(np.float32)  # saturation\n    v = hsv[:, :, 2].astype(np.float32)  # value\n    h = (h*(1 + random.uniform(-1,1)*mag[0]))%180\n    s =  s*(1 + random.uniform(-1,1)*mag[1])\n    v =  v*(1 + random.uniform(-1,1)*mag[2])\n\n    hsv[:, :, 0] = np.clip(h,0,180).astype(np.uint8)\n    hsv[:, :, 1] = np.clip(s,0,255).astype(np.uint8)\n    hsv[:, :, 2] = np.clip(v,0,255).astype(np.uint8)\n    image = cv2.cvtColor(hsv, cv2.COLOR_HSV2BGR)\n    image = image.astype(np.float32)/255\n    return image, mask\n\ndef do_random_noise(image, mask, mag=0.1):\n    height, width = image.shape[:2]\n    noise = np.random.uniform(-1,1, (height, width,1))*mag\n    image = image + noise\n    image = np.clip(image,0,1)\n    return image, mask\n\ndef do_random_rotate_scale(image, mask, angle=30, scale=[0.8,1.2] ):\n    angle = np.random.uniform(-angle, angle)\n    scale = np.random.uniform(*scale) if scale is not None else 1\n    \n    height, width = image.shape[:2]\n    center = (height // 2, width // 2)\n    \n    transform = cv2.getRotationMatrix2D(center, angle, scale)\n    image = cv2.warpAffine( image, transform, (width, height), flags=cv2.INTER_LINEAR,\n                            borderMode=cv2.BORDER_CONSTANT, borderValue=(0,0,0))\n    mask  = cv2.warpAffine( mask, transform, (width, height), flags=cv2.INTER_LINEAR,\n                            borderMode=cv2.BORDER_CONSTANT, borderValue=0)\n    return image, mask\n\ndef build_transforms(CFG):\n    data_transforms = {\n        \"train\": A.Compose([\n            A.OneOf([\n                A.Resize(*(CFG.img_size, CFG.img_size), interpolation=cv2.INTER_NEAREST, p=1.0),   \n                # 最邻近插值算法\n            ], p=1),\n            A.HorizontalFlip(p=0.5),\n            ], p=1.0),\n        \n        \"valid_test\": A.Compose([\n            A.Resize(*(CFG.img_size, CFG.img_size), interpolation=cv2.INTER_NEAREST),\n            ], p=1.0)\n        }\n    return data_transforms","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:50.459738Z","iopub.execute_input":"2022-08-25T10:33:50.460593Z","iopub.status.idle":"2022-08-25T10:33:50.482866Z","shell.execute_reply.started":"2022-08-25T10:33:50.460541Z","shell.execute_reply":"2022-08-25T10:33:50.481789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_augment(image, mask):\n    image, mask = do_random_flip(image, mask)\n    image, mask = do_random_rot90(image, mask)\n\n    for fn in np.random.choice([\n        lambda image, mask: (image, mask),\n        lambda image, mask: do_random_noise(image, mask, mag=0.1),\n        lambda image, mask: do_random_contast(image, mask, mag=0.40),\n        lambda image, mask: do_random_hsv(image, mask, mag=[0.40, 0.40, 0])\n    ], 2): image, mask = fn(image, mask)\n\n    for fn in np.random.choice([\n        lambda image, mask: (image, mask),\n        lambda image, mask: do_random_rotate_scale(image, mask, angle=45, scale=[0.50, 2.0]),\n    ], 1): image, mask = fn(image, mask)\n\n    return image, mask\n\n\ndef valid_augment(image, mask):\n    #image, mask  = do_crop(image, mask, image_size, xy=(None,None))\n    return image, mask","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:50.484279Z","iopub.execute_input":"2022-08-25T10:33:50.484845Z","iopub.status.idle":"2022-08-25T10:33:50.496456Z","shell.execute_reply.started":"2022-08-25T10:33:50.484806Z","shell.execute_reply":"2022-08-25T10:33:50.495371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"additionals\"><center>Dataset & Dataloader</center></h3>","metadata":{}},{"cell_type":"code","source":"class build_dataset(Dataset):\n    def __init__(self, df, label=True, transforms=None, CFG=None, flag='train'):\n        ###########################################\n        ##### >>>>>>> Use \"hubmap-2022-256x256\" Dataset <<<<<<\n        ############################################\n        self.df = df\n        self.flag = flag\n        ids = df.id.astype(str).values\n        if CFG is not None:\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        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        if self.flag == 'test':\n            img_path = os.path.join(CFG.data_path, \"test_images/\" + str(self.df.loc[index, 'id'])+\".tiff\")\n        elif self.flag == 'train':\n            img_path = os.path.join(CFG.data_path, \"train_images/\" + str(self.df.loc[index, 'id'])+\".tiff\")\n        else:\n            raise ValueError\n            \n        img_height = self.df.loc[index, 'img_height']\n        img_width = self.df.loc[index, 'img_width']\n        id_ = self.df.loc[index, 'id']\n        img = read_tiff(img_path)\n        \n        if self.label:\n            file_name = self.file_names[index]\n            sub_df = self.df.iloc[index]\n            organ = self.organ_to_label[sub_df.organ]\n            \n            image = cv2.cvtColor(cv2.imread(os.path.join(self.third_data_train, file_name)), cv2.COLOR_BGR2RGB).astype(np.float32)/255\n            mask = cv2.imread(os.path.join(self.third_data_mask, file_name), cv2.IMREAD_GRAYSCALE).astype(np.float32)/255\n            image = cv2.resize(image,dsize=(CFG.img_size, CFG.img_size),interpolation=cv2.INTER_LINEAR)\n            mask  = cv2.resize(mask, dsize=(CFG.img_size, CFG.img_size),interpolation=cv2.INTER_LINEAR)\n            \n            if self.transforms:\n                image, mask = self.transforms(image, mask)\n            \n            data_info ={}\n            data_info['index']= index\n            data_info['id'] = file_name\n            data_info['organ'] = torch.tensor([organ], dtype=torch.long)\n            image = image[:,:,::-1].transpose(2,0,1) # BGR => RGB\n            image = np.ascontiguousarray(image)\n            data_info['image'] = torch.tensor(image, dtype=torch.float)\n            data_info['mask' ] = torch.tensor(mask>0.5, dtype=torch.float)\n            return data_info\n        \n        else:\n            if self.transforms:\n                data = self.transforms(image=img)\n                img  = data['image']\n                \n            img = np.transpose(img, (2, 0, 1))\n            return torch.tensor(img), img_height, img_width, id_","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:50.498098Z","iopub.execute_input":"2022-08-25T10:33:50.498774Z","iopub.status.idle":"2022-08-25T10:33:50.517461Z","shell.execute_reply.started":"2022-08-25T10:33:50.498737Z","shell.execute_reply":"2022-08-25T10:33:50.515592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_dataloader(df, fold, data_transforms, CFG):\n    \n    train_df = df[df.fold != fold].reset_index(drop=True)\n    valid_df = df[df.fold == fold].reset_index(drop=True)\n    \n    train_dataset = build_dataset(train_df, label=True, transforms=data_transforms['train'], CFG=CFG, flag='train')\n    valid_dataset = build_dataset(valid_df, label=True, transforms=data_transforms['valid_test'], CFG=CFG, flag='train')\n    \n    def null_collate(batch):\n        d = {}\n        key = batch[0].keys()\n        for k in key:\n            v = [b[k] for b in batch]\n            if k in ['mask', 'image', 'organ']:\n                v = torch.stack(v)\n            d[k] = v\n\n        d['mask'] = d['mask'].unsqueeze(1)\n        d['organ'] = d['organ'].reshape(-1)\n        return d\n    \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, collate_fn = null_collate)\n    valid_loader = DataLoader(valid_dataset, batch_size=CFG.valid_bs, num_workers=CFG.num_worker, \n                              shuffle=False, pin_memory=True, collate_fn = null_collate)\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#     model = Net()\n    \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\n    model.to(CFG.device)\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:50.519007Z","iopub.execute_input":"2022-08-25T10:33:50.519639Z","iopub.status.idle":"2022-08-25T10:33:50.532156Z","shell.execute_reply.started":"2022-08-25T10:33:50.519600Z","shell.execute_reply":"2022-08-25T10:33:50.531135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"additionals\"><center>Build Loss</center></h3>","metadata":{}},{"cell_type":"code","source":"def build_loss():\n    BCELoss     = smp.losses.SoftBCEWithLogitsLoss()\n    DiceLoss    = smp.losses.DiceLoss(mode='binary')\n    return {\"BCELoss\":BCELoss, \"DiceLoss\":DiceLoss}\n\n###############################################################\n##### >>>>>>> part4: build_metric <<<<<<\n###############################################################\ndef compute_dice_score(probability, mask):\n    N = len(probability)\n    p = probability.reshape(N,-1)\n    t = mask.reshape(N,-1)\n\n    p = p>0.5\n    t = t>0.5\n    uion = p.sum(-1) + t.sum(-1)\n    overlap = (p*t).sum(-1)\n    dice = 2*overlap/(uion+0.0001)\n    return dice\n\ndef dice_coef(y_pred, y_true,  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","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:50.533783Z","iopub.execute_input":"2022-08-25T10:33:50.534183Z","iopub.status.idle":"2022-08-25T10:33:50.546519Z","shell.execute_reply.started":"2022-08-25T10:33:50.534146Z","shell.execute_reply":"2022-08-25T10:33:50.545187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"additionals\"><center>Train</center></h3>","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, dice_all = 0, 0, 0\n    \n    pbar = tqdm(enumerate(train_loader), total=len(train_loader), desc='Train ')\n    for _, batch in pbar:\n        # batch: dict_keys(['index', 'id', 'organ', 'image', 'mask'])\n        optimizer.zero_grad()\n\n        images = batch['image'].to(CFG.device, dtype=torch.float) # [b, c, w, h]\n        masks  = batch['mask'].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 = 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() / images.shape[0]\n        bce_all += bce_loss.item() / images.shape[0]\n        # dice_all += dice_loss.item() / images.shape[0]\n        dice_all += 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}, dice_all: {:.3f}\".format(losses_all, bce_all, dice_all), flush=True)\n    ","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:50.551153Z","iopub.execute_input":"2022-08-25T10:33:50.551817Z","iopub.status.idle":"2022-08-25T10:33:50.562541Z","shell.execute_reply.started":"2022-08-25T10:33:50.551789Z","shell.execute_reply":"2022-08-25T10:33:50.561365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"additionals\"><center>Validation</center></h3>","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef valid_one_epoch(model, valid_loader, CFG):\n    model.eval()\n    \n    valid_image = []\n    valid_probability = []\n    valid_mask = []\n    valid_id = []\n    \n    pbar = tqdm(enumerate(valid_loader), total=len(valid_loader), desc='Valid ')\n    for _, batch in pbar:\n        images = batch['image'].to(CFG.device, dtype=torch.float) # [b, c, w, h]\n        masks = batch['mask'].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            y_preds   = torch.nn.Sigmoid()(y_preds) # [b, c, w, h]\n        \n        valid_image.append(images)\n        valid_probability.append(y_preds)\n        valid_mask.append(masks)\n        valid_id.extend(batch['id'])\n        \n    images = torch.cat(valid_image)\n    probabilitys = torch.cat(valid_probability)\n    masks = torch.cat(valid_mask)\n    val_dice = dice_coef(probabilitys, masks)\n        \n    print(\"val_dice: {:.4f}\".format(val_dice), flush=True)\n    \n    return val_dice, images, probabilitys, masks, valid_id","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:50.564151Z","iopub.execute_input":"2022-08-25T10:33:50.565404Z","iopub.status.idle":"2022-08-25T10:33:50.575793Z","shell.execute_reply.started":"2022-08-25T10:33:50.565367Z","shell.execute_reply":"2022-08-25T10:33:50.574733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef test_one_epoch(ckpt_paths, test_loader, CFG):\n    pred_ids = []\n    pred_rles = []\n    \n    pbar = tqdm(enumerate(test_loader), total=len(test_loader), desc='Test: ')\n    for _, (images, heights, widths, ids) 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], CFG.num_classes, 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 \n        for idx in range(masks.shape[0]):\n            height = heights[idx].item()\n            width = widths[idx].item()\n            id_ = ids[idx].item()\n            msk = cv2.resize(masks[idx].squeeze(), dsize=(width, height), interpolation=cv2.INTER_NEAREST)\n            rle = rle_encode(msk)\n            pred_rles.append(rle)\n            pred_ids.append(id_)\n    \n    return pred_ids, pred_rles","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:50.578396Z","iopub.execute_input":"2022-08-25T10:33:50.578673Z","iopub.status.idle":"2022-08-25T10:33:50.592863Z","shell.execute_reply.started":"2022-08-25T10:33:50.578635Z","shell.execute_reply":"2022-08-25T10:33:50.591822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"additionals\"><center>Visualization</center></h3>","metadata":{}},{"cell_type":"code","source":"from matplotlib import pyplot as plt\ndef plot_visual(image, mask, pred, image_id, cmap):\n    plt.figure(figsize=(16, 10))\n\n    plt.subplot(1, 3, 1)\n    plt.imshow(image.transpose(1, 2, 0))\n    plt.grid(visible=False)\n    plt.title(\"image\", fontsize=10)\n    plt.axis(\"off\")\n\n    plt.subplot(1, 3, 2)\n    plt.imshow(image.transpose(1, 2, 0))\n    plt.imshow(mask.transpose(1, 2, 0), 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.transpose(1, 2, 0))\n    plt.imshow(pred.transpose(1, 2, 0), cmap=cmap, alpha=0.5)\n    plt.title(f\"pred\", fontsize=10)    \n    plt.axis(\"off\")\n\n    plt.savefig(f\"./result/{image_id}\")\n    # plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:50.594939Z","iopub.execute_input":"2022-08-25T10:33:50.595965Z","iopub.status.idle":"2022-08-25T10:33:50.606672Z","shell.execute_reply.started":"2022-08-25T10:33:50.595927Z","shell.execute_reply":"2022-08-25T10:33:50.605673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h3 class=\"list-group-item list-group-item-action active\" data-toggle=\"list\" style=\"color:#c3448b; background:#efe9e9; border:1px dashed #efe50b;\" role=\"tab\" aria-controls=\"additionals\"><center>TRAIN</center></h3>","metadata":{}},{"cell_type":"code","source":"###############################################################\n##### >>>>>>> config <<<<<<\n###############################################################\nclass CFG:\n    # step1: hyper-parameter\n    seed = 42 \n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    num_worker = 0 # 0 if debug. 16 if train by \"htop\" check\n    data_path = \"../input/hubmap-organ-segmentation\"\n    third_data_path = \"../input/hubmap-2022-256x256\"\n    ckpt_path = \"../output/ckpt-path/resnet_img768_bs8_fold2\" # for submit\n    # step2: data\n    n_fold = 2\n    img_size = 768\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 = 2\n    lr = 1e-3\n    wd = 1e-5\n    lr_drop = 8\n    # step5: infer\n    thr = 0.4\n\nset_seed(CFG.seed)\nif not os.path.exists(CFG.ckpt_path):\n    os.makedirs(CFG.ckpt_path)\n\ntrain_val_flag = True\nif 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\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 = {'train':train_augment, 'valid_test': valid_augment}\n        train_loader, valid_loader = build_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) # optimizer\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, images, probabilitys, masks, valid_id = valid_one_epoch(model, valid_loader, CFG)\n\n            ###############################################################\n            ##### >>>>>>> step4: save best & last 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}.pth\"\n                if os.path.isfile(save_path):\n                    os.remove(save_path) \n                torch.save(model.state_dict(), save_path)\n\n            ###############################################################\n            ##### >>>>>>> step5: visual last epoch pred <<<<<<\n            ###############################################################\n            if epoch == CFG.epoch:\n                save_path = f\"{CFG.ckpt_path}/last_fold{fold}.pth\"\n                torch.save(model.state_dict(), save_path)\n\n                val_img_num = images.shape[0]\n#                 for idx in range(10):\n#                     plot_visual(images[idx].cpu().numpy(), masks[idx].cpu().numpy()*255, probabilitys[idx].cpu().numpy(), valid_id[idx], \"bwr\")\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","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:33:50.610356Z","iopub.execute_input":"2022-08-25T10:33:50.610895Z","iopub.status.idle":"2022-08-25T10:39:18.468704Z","shell.execute_reply.started":"2022-08-25T10:33:50.610860Z","shell.execute_reply":"2022-08-25T10:39:18.466213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_flag = True\nif test_flag:\n    set_seed(CFG.seed)\n    ###############################################################\n    ##### part0: data preprocess\n    ###############################################################\n    df = pd.read_csv(os.path.join(CFG.data_path, \"test.csv\"))\n    df['image_path'] = df['id'].apply(lambda x: os.path.join(CFG.data_path, 'test_images', str(x) + '.tiff'))\n\n    data_transforms = build_transforms(CFG)\n    test_dataset = build_dataset(df, label=False, transforms=data_transforms['valid_test'], flag='test')\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'{CFG.ckpt_path}/best*')\n    assert len(ckpt_paths) == CFG.n_fold, \"ckpt path error!\"\n    \n    print(type(ckpt_paths))\n    pred_ids, pred_rles = test_one_epoch(ckpt_paths, test_loader, CFG)\n\n    ###############################################################\n    ##### step3: submit\n    ###############################################################\n    pred_df = pd.DataFrame({\n        \"id\":pred_ids,\n        \"rle\":pred_rles\n    })\n    pred_df.to_csv('submission.csv',index=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-25T10:39:18.473064Z","iopub.execute_input":"2022-08-25T10:39:18.476813Z","iopub.status.idle":"2022-08-25T10:39:20.317110Z","shell.execute_reply.started":"2022-08-25T10:39:18.476768Z","shell.execute_reply":"2022-08-25T10:39:20.315501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}