{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!/usr/bin/env python\n# coding: utf-8\n\n# ## summary\n# \n# * 2.5d segmentation\n#     *  segmentation_models_pytorch \n#     *  Unet\n# * use only 6 slices in the middle\n# * slide inference\n\n# In[1]:\n\n# from resnet3d import generate_model\n\nimport sys\nsys.path.append('/kaggle/input/resnet')\nsys.path.append('/kaggle/input/efficient-3dcnns/Efficient-3DCNNs')\n\n\nfrom resnet import generate_model\nimport torch.nn.functional as F\n\nfrom sklearn.metrics import roc_auc_score, accuracy_score, f1_score, log_loss\nimport pickle\nfrom torch.utils.data import DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport warnings\nimport sys\nimport os\nimport gc\nimport sys\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nimport cv2\n\nimport scipy as sp\nimport numpy as np\n\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport argparse\nimport importlib\nimport torch\nimport torch.nn as nn\nfrom torch.optim import Adam, SGD, AdamW\n\nimport datetime\n\n\n\n\n# import segmentation_models_pytorch as smp\n\n\n\nimport numpy as np\nfrom torch.utils.data import DataLoader, Dataset\nimport cv2\nimport datetime\nimport torch\nimport os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\n\n# ## config\n\n# In[7]:\n\n\nimport os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport argparse\n\n\nparser = argparse.ArgumentParser()\nparser.add_argument(\n    \"--valid_id\",\n    type=int,\n    default=1,\n    help='valid_id'\n    )\n\nparser.add_argument(\n    \"--model_name\",\n    type=str,\n    default='Unet',\n    help='model_name'\n    )\n\nparser.add_argument(\n    \"--image_size\",\n    type=int,\n    default=256,\n    help='image_size'\n    )\n\n\nparser.add_argument(\n    \"--backbone_name\",\n    type=str,\n    default='efficientnet-b1',\n    help='model_name'\n    )\n\nparser.add_argument(\n    \"--batch_size\",\n    type=int,\n    default=1,\n    help='batch_size'\n    )\n\nparser.add_argument(\n    \"--epochs\",\n    type=int,\n    default=1,\n    help='epochs'\n    )\n\nparser.add_argument(\n    \"--diceloss\",\n    type=int,\n    default=0,\n    help='diceloss'\n    )\n\nparser.add_argument(\n    \"--tverskyloss\",\n    type=int,\n    default=0,\n    help='tverskyloss'\n    )\n\nparser.add_argument(\n    \"--in_chans\",\n    type=int,\n    default=22,\n    help='in_chans'\n    )\n\nparser.add_argument(\n    \"--resnet_depth\",\n    type=int,\n    default=152,\n    help='resnet_depth'\n    )\n\nparser.add_argument(\n    \"--resnet_weight\",\n    type=str,\n    default='r3d152_KM_200ep.pth',\n    help='resnet_weight'\n    )\n\nparser.add_argument(\n    \"--slicing_num\",\n    type=int,\n    default=4300,\n    help='balance in mask pixel for 4Fold Training'\n    )\n\nparser.add_argument(\n    \"--cropping_num_min\",\n    type=int,\n    default=12,\n    help='balance in mask pixel for 4Fold Training'\n    )\n\nparser.add_argument(\n    \"--cropping_num_max\",\n    type=int,\n    default=22,\n    help='balance in mask pixel for 4Fold Training'\n    )\n\nparser.add_argument(\n    \"--fbeta_gamma\",\n    type=int,\n    default=2,\n    help='balance in mask pixel for 4Fold Training'\n    )\n\nparser.add_argument(\n    \"--clip_min\",\n    type=int,\n    default=50,\n    help='balance in mask pixel for 4Fold Training'\n    )\n\nparser.add_argument(\n    \"--clip_max\",\n    type=int,\n    default=200,\n    help='balance in mask pixel for 4Fold Training'\n    )\n\nparser.add_argument(\n    \"--ls\",\n    type=float,\n    default=0.3,\n    help='balance in mask pixel for 4Fold Training'\n    )\n\n\nparser.add_argument(\n    \"--model\",\n    type=str,\n    default='resnet152',\n    help='choose between resnet152, resnet200, resnext101'\n    )\n\n\nargs = parser.parse_args(args=[])\n\nd = datetime.datetime.now()\nyear, month, day, hour, minute, second = d.year, d.month, d.day, d.hour, d.minute, d.second\nif len(str(month)) == 1:\n    month = '0' + str(month)\nif len(str(day)) == 1:\n    day = '0' + str(day)\nif len(str(hour)) == 1:\n    hour = '0' + str(hour)\n\ncurrent_day = f'{year}_{month}_{day}'\n\nclass CFG:\n    # ============== comp exp name =============\n    comp_name = 'vesuvius'\n\n    # comp_dir_path = './'\n    comp_dir_path = '/kaggle/input'\n    comp_folder_name = 'vesuvius-challenge-ink-detection'\n    # comp_dataset_path = f'{comp_dir_path}datasets/{comp_folder_name}/'\n    comp_dataset_path = f'{comp_dir_path}/{comp_folder_name}/'\n\n\n    # ============== pred target =============\n    target_size = 1\n\n    # ============== model cfg =============\n    model_name = args.model_name\n    backbone = args.backbone_name\n    #backbone = 'se_resnext50_32x4d'\n\n    in_chans = args.in_chans # 65\n    # ============== training cfg =============\n    size = args.image_size\n    tile_size = args.image_size\n    stride_rate = 3\n    stride = tile_size // stride_rate\n\n    # cropping_num = args.cropping_num\n\n    # exp_name = f'{current_day}_{model_name}_{backbone}_{tile_size}_{in_chans}_batchsize{args.batch_size}_diceloss{args.diceloss}_tverskyloss{args.tverskyloss}_3DCNN_depth{args.resnet_depth}'\n    exp_name = f'{current_day}_{tile_size}_{in_chans}_batchsize{args.batch_size}_diceloss{args.diceloss}_tverskyloss{args.tverskyloss}_3DCNN_depth{args.resnet_depth}_stride{stride_rate}_cropping_num{args.cropping_num_min}-{args.cropping_num_max}_ls{args.ls}_clip{args.clip_min}_{args.clip_max}_CosineAnnealingWarmUpRestarts(optimizer, T_0=9, T_mult=1, eta_max=1e-4, T_up=3, gamma=0.5)'\n\n    train_batch_size = args.batch_size # 32\n    valid_batch_size = train_batch_size * 2\n    use_amp = True\n\n    scheduler = 'GradualWarmupSchedulerV2'\n    # scheduler = 'CosineAnnealingLR'\n    epochs = args.epochs # 30\n\n    # adamW warmupあり\n    warmup_factor = 10\n    # lr = 1e-4 / warmup_factor\n    lr = 1e-3 / warmup_factor\n\n    # ============== fold =============\n    valid_id = args.valid_id\n\n    # objective_cv = 'binary'  # 'binary', 'multiclass', 'regression'\n    metric_direction = 'maximize'  # maximize, 'minimize'\n    # metrics = 'dice_coef'\n\n    # ============== fixed =============\n    pretrained = True\n    inf_weight = 'best'  # 'best'\n\n    min_lr = 1e-7\n    weight_decay = 1e-6\n    max_grad_norm = 1000\n\n    print_freq = 50\n    num_workers = 8\n\n    seed = 42\n\n    # ============== set dataset path =============\n    print('set dataset path')\n    output_save_folder = 'checkpoints'\n    outputs_path = f'/kaggle/working/{output_save_folder}/{current_day}/{exp_name}/{args.valid_id}/'\n\n    submission_dir = outputs_path + 'submissions/'\n    submission_path = submission_dir + f'submission_{exp_name}.csv'\n\n    model_dir = outputs_path + \\\n        f'{comp_name}-models/'\n\n    mask_dir = outputs_path + 'mask_pred'\n\n    figures_dir = outputs_path + 'figures/'\n\n    log_dir = outputs_path + 'logs/'\n    log_path = log_dir + f'{exp_name}.txt'\n\n    # ============== augmentation =============\n    train_aug_list = [\n        # A.RandomResizedCrop(\n        #     size, size, scale=(0.85, 1.0)),\n        #A.Resize(size, size),\n        A.HorizontalFlip(p=0.6),\n        A.VerticalFlip(p=0.6),\n        A.RandomGamma(gamma_limit=(50, 150), p=0.6),\n        A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.6),\n        A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=360, interpolation=0, border_mode=0, p=0.6),\n        A.OneOf([\n                A.GaussNoise(var_limit=[10, 30]),\n                A.GaussianBlur(),\n                # A.MotionBlur(),\n                ], p=0.6),\n        # A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.5),\n        A.CoarseDropout(max_holes=4, max_width=int(size * 0.2), max_height=int(size * 0.2),\n                        mask_fill_value=0, p=0.6),\n        # A.Cutout(max_h_size=int(size * 0.6),\n        #          max_w_size=int(size * 0.6), num_holes=1, p=1.0),\n        A.Normalize(\n            mean= [0] * args.in_chans,\n            std= [1] * args.in_chans\n        ),\n        ToTensorV2(transpose_mask=True),\n    ]\n\n    valid_aug_list = [\n        A.Resize(size, size),\n        A.Normalize(\n            mean= [0] * args.in_chans,\n            std= [1] * args.in_chans\n        ),\n        ToTensorV2(transpose_mask=True),\n    ]\n\n\n# ## helper\n\n# In[8]:\n\n\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\n# In[9]:\n\n\ndef init_logger(log_file):\n    from logging import getLogger, INFO, FileHandler, Formatter, StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\ndef set_seed(seed=None, cudnn_deterministic=True):\n    if seed is None:\n        seed = 42\n\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = cudnn_deterministic\n    torch.backends.cudnn.benchmark = False\n\n\n# In[10]:\n\n\ndef make_dirs(cfg):\n    for dir in [cfg.model_dir, cfg.figures_dir, cfg.submission_dir, cfg.log_dir, cfg.mask_dir]:\n        os.makedirs(dir, exist_ok=True)\n\n\n# In[11]:\n\n\ndef cfg_init(cfg, mode='train'):\n    set_seed(cfg.seed)\n    # set_env_name()\n    # set_dataset_path(cfg)\n\n    if mode == 'train':\n        make_dirs(cfg)\n\n\n# In[12]:\n\n\ncfg_init(CFG)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nLogger = init_logger(log_file=CFG.log_path)\n\nLogger.info('\\n\\n-------- exp_info -----------------')\n# Logger.info(datetime.datetime.now().strftime('%Y年%m月%d日 %H:%M:%S'))\n\n\n\n\ndef read_image_mask(fragment_id):\n\n    images = []\n\n    # idxs = range(65)\n    mid = 65 // 2\n    start = mid - CFG.in_chans // 2\n    end = mid + CFG.in_chans // 2\n    idxs = list(range(start, end))\n    # idxs = [14, 15] + idxs\n\n\n    for i in tqdm(idxs):\n        \n        image = cv2.imread(CFG.comp_dataset_path + f\"train/{fragment_id}/surface_volume/{i:02}.tif\", 0)\n\n        pad0 = (CFG.tile_size - image.shape[0] % CFG.tile_size)\n        pad1 = (CFG.tile_size - image.shape[1] % CFG.tile_size)\n\n        image = np.pad(image, [(0, pad0), (0, pad1)], constant_values=0)\n\n        images.append(image)\n    images = np.stack(images, axis=2)\n    #images = np.media`n(images, axis=2)[..., None].astype('uint8')\n\n    \n    mask = cv2.imread(CFG.comp_dataset_path + f\"train/{fragment_id}/inklabels.png\", 0)\n    mask = np.pad(mask, [(0, pad0), (0, pad1)], constant_values=0)\n\n    mask = mask.astype('float32')\n    mask /= 255.0\n\n    mask_prag = cv2.imread(CFG.comp_dataset_path + f\"train/{fragment_id}/mask.png\", 0)\n    mask_prag = np.pad(mask_prag, [(0, pad0), (0, pad1)], constant_values=0)\n\n    mask_prag = mask_prag.astype('float32')\n    mask_prag /= 255.0\n\n    return images, mask, mask_prag\n\n\n# In[14]:\n\n\ndef get_train_valid_dataset():\n    train_images = []\n    train_masks = []\n\n    valid_images = []\n    valid_masks = []\n    valid_xyxys = []\n\n    valid_masks_frag = []\n\n    for fragment_id in range(1, 5):\n\n        if fragment_id == 2:\n            image, mask, mask_frag = read_image_mask(2)\n            image, mask, mask_frag = image[:, :args.slicing_num, :], mask[:, :args.slicing_num], mask_frag[:, :args.slicing_num]\n\n        elif fragment_id == 4:\n            image, mask, mask_frag = read_image_mask(2)\n            image, mask, mask_frag = image[:, args.slicing_num:, :], mask[:, args.slicing_num:], mask_frag[:, args.slicing_num:]\n\n        else:\n            image, mask, mask_frag = read_image_mask(fragment_id)\n\n        x1_list = list(range(0, image.shape[1]-CFG.tile_size+1, CFG.stride))\n        y1_list = list(range(0, image.shape[0]-CFG.tile_size+1, CFG.stride))\n\n        for y1 in y1_list:\n            for x1 in x1_list:\n                y2 = y1 + CFG.tile_size\n                x2 = x1 + CFG.tile_size\n                # xyxys.append((x1, y1, x2, y2))\n        \n                if fragment_id == CFG.valid_id:\n                    valid_images.append(image[y1:y2, x1:x2])\n                    valid_masks.append(mask[y1:y2, x1:x2, None])\n\n                    valid_xyxys.append([x1, y1, x2, y2])\n\n                    valid_masks_frag.append(mask_frag[y1:y2, x1:x2, None])\n                else:\n\n                    tmp = mask_frag[y1:y2, x1:x2]\n                    if tmp.min() == 1:\n                        train_images.append(image[y1:y2, x1:x2])\n                        train_masks.append(mask[y1:y2, x1:x2, None])\n\n    return train_images, train_masks, valid_images, valid_masks, valid_xyxys, valid_masks_frag\n\n\n# In[15]:\n\n\ntrain_images, train_masks, valid_images, valid_masks, valid_xyxys, valid_masks_frag = get_train_valid_dataset()\n\n\n\nvalid_xyxys = np.stack(valid_xyxys)\n\n\n\n\nimport numpy as np\nfrom torch.utils.data import DataLoader, Dataset\nimport cv2\nimport torch\nimport os\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\n\n# In[18]:\n\n\ndef get_transforms(data, cfg):\n    if data == 'train':\n        aug = A.Compose(cfg.train_aug_list)\n    elif data == 'valid':\n        aug = A.Compose(cfg.valid_aug_list)\n\n    # print(aug)\n    return aug\n\nclass CustomDataset(Dataset):\n    def __init__(self, images, cfg, labels=None, mask_frag=None, transform=None, mode=None):\n        self.images = images\n        self.cfg = cfg\n        self.labels = labels\n        self.mask_frag = mask_frag\n        self.transform = transform\n        self.mode = mode\n\n    def __len__(self):\n        # return len(self.df)\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        image = self.images[idx]\n        image = np.clip(image, args.clip_min, args.clip_max)\n\n\n        label = self.labels[idx]\n\n        if self.mask_frag:\n            mask_frag = self.mask_frag[idx]\n\n        image_tmp = np.zeros_like(image)\n\n        if self.transform:\n\n            if self.mode == 'train':\n\n                # cropping_num = CFG.cropping_num\n                cropping_num = random.randint(args.cropping_num_min, args.cropping_num_max)\n\n                start_idx = random.randint(0, args.in_chans - cropping_num)\n                crop_indices = np.arange(start_idx, start_idx + cropping_num)\n\n                start_paste_idx = random.randint(0, args.in_chans - cropping_num)\n\n                tmp = np.arange(start_paste_idx, cropping_num)\n                np.random.shuffle(tmp)\n\n                cutout_idx = random.randint(0, 2)\n                temporal_random_cutout_idx = tmp[:cutout_idx]\n\n                image_tmp[..., start_paste_idx : start_paste_idx + cropping_num] = image[..., crop_indices]\n\n                if random.random() > 0.4:\n                    image_tmp[..., temporal_random_cutout_idx] = 0\n                image = image_tmp\n\n\n        data = self.transform(image=image, mask=label)\n\n        image = data['image'].unsqueeze(0)\n        label = data['mask']\n\n        if self.mode == 'train':\n            return image, label\n        else:\n            return image, label, mask_frag\n\n\ntrain_dataset = CustomDataset(\n    train_images, CFG, labels=train_masks, mask_frag=None, transform=get_transforms(data='train', cfg=CFG), mode='train')\nvalid_dataset = CustomDataset(\n    valid_images, CFG, labels=valid_masks, mask_frag=valid_masks_frag, transform=get_transforms(data='valid', cfg=CFG), mode='valid')\n\ntrain_loader = DataLoader(train_dataset,\n                          batch_size=CFG.train_batch_size,\n                          shuffle=True,\n                          num_workers=CFG.num_workers, drop_last=True,\n                          )\nvalid_loader = DataLoader(valid_dataset,\n                          batch_size=CFG.valid_batch_size,\n                          shuffle=False,\n                          num_workers=CFG.num_workers, drop_last=False)\n\n\n\n\n\nclass CustomModel(nn.Module):\n    def __init__(self, cfg, weight=None):\n        super().__init__()\n        self.cfg = cfg\n\n        if args.model_name == 'Unet':\n\n            self.encoder = smp.Unet(\n                encoder_name=cfg.backbone,\n                encoder_weights=weight,\n                in_channels=cfg.in_chans,\n                classes=cfg.target_size,\n                activation=None,\n            )\n\n        if args.model_name == 'UnetPlusPlus':\n\n            self.encoder = smp.UnetPlusPlus(\n                encoder_name=cfg.backbone,\n                encoder_weights=weight,\n                in_channels=cfg.in_chans,\n                classes=cfg.target_size,\n                activation=None,\n            )\n\n    def forward(self, image):\n        output = self.encoder(image)\n        # output = output.squeeze(-1)\n        return output\n\n\ndef build_model(cfg, weight=\"imagenet\"):\n    print('model_name', cfg.model_name)\n    print('backbone', cfg.backbone)\n\n    model = CustomModel(cfg, weight)\n\n    return model\n\n\nclass Decoder(nn.Module):\n    def __init__(self, encoder_dims, upscale):\n        super().__init__()\n        self.convs = nn.ModuleList([\n            nn.Sequential(\n                nn.Conv2d(encoder_dims[i]+encoder_dims[i-1], encoder_dims[i-1], 3, 1, 1, bias=False),\n                nn.BatchNorm2d(encoder_dims[i-1]),\n                nn.ReLU(inplace=True)\n            ) for i in range(1, len(encoder_dims))])\n\n        self.logit = nn.Conv2d(encoder_dims[0], 2, 1, 1, 0)\n        self.up = nn.Upsample(scale_factor=upscale, mode=\"bilinear\")\n    def forward(self, feature_maps):\n        for i in range(len(feature_maps)-1, 0, -1):\n            f_up = F.interpolate(feature_maps[i], scale_factor=2, mode=\"bilinear\")\n            f = torch.cat([feature_maps[i-1], f_up], dim=1)\n            f_down = self.convs[i-1](f)\n            feature_maps[i-1] = f_down\n\n        x = self.logit(feature_maps[0])\n        mask = self.up(x)\n        return mask\n\n\nclass SegModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # original kaggle code\n        # self.encoder = generate_model(model_depth=args.resnet_depth, n_input_channels=1)\n\n        # original paper code\n        self.encoder = generate_model(model_depth=args.resnet_depth,\n                                      n_input_channels=1,\n                                      shortcut_type='B',\n                                      conv1_t_size=7,\n                                      conv1_t_stride=1,\n                                      widen_factor=1.0,\n                                      n_classes=1039,\n                                      no_max_pool=True)\n\n        # original kaggle code\n        # self.decoder = Decoder(encoder_dims=[64, 128, 256, 512], upscale=4)\n\n        # original paper code\n        self.decoder = Decoder(encoder_dims=[256, 512, 1024, 2048], upscale=4)\n\n\n    def forward(self, x):\n        feat_maps = self.encoder(x)\n        feat_maps_pooled = [torch.mean(f, dim=2) for f in feat_maps]\n        pred_mask = self.decoder(feat_maps_pooled)\n        return pred_mask\n    \n    def load_pretrained_weights(self, state_dict):\n        # Convert 3 channel weights to single channel\n        # ref - https://timm.fast.ai/models#Case-1:-When-the-number-of-input-channels-is-1\n        conv1_weight = state_dict['conv1.weight']\n        state_dict['conv1.weight'] = conv1_weight.sum(dim=1, keepdim=True)\n        print(self.encoder.load_state_dict(state_dict, strict=False))\n\n\n\n# ## scheduler\n\n# In[24]:\n\n\nimport torch.nn as nn\nimport torch\nimport math\nimport time\nimport numpy as np\nimport torch\n\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\n\n\n\nimport math\nfrom torch.optim.lr_scheduler import _LRScheduler\nclass CosineAnnealingWarmUpRestarts(_LRScheduler):\n    def __init__(self, optimizer, T_0, T_mult=1, eta_max=0.1, T_up=0, gamma=1., last_epoch=-1):\n        if T_0 <= 0 or not isinstance(T_0, int):\n            raise ValueError(\"Expected positive integer T_0, but got {}\".format(T_0))\n        if T_mult < 1 or not isinstance(T_mult, int):\n            raise ValueError(\"Expected integer T_mult >= 1, but got {}\".format(T_mult))\n        if T_up < 0 or not isinstance(T_up, int):\n            raise ValueError(\"Expected positive integer T_up, but got {}\".format(T_up))\n        self.T_0 = T_0\n        self.T_mult = T_mult\n        self.base_eta_max = eta_max\n        self.eta_max = eta_max\n        self.T_up = T_up\n        self.T_i = T_0\n        self.gamma = gamma\n        self.cycle = 0\n        self.T_cur = last_epoch\n        super(CosineAnnealingWarmUpRestarts, self).__init__(optimizer, last_epoch)\n\n    def get_lr(self):\n        if self.T_cur == -1:\n            return self.base_lrs\n        elif self.T_cur < self.T_up:\n            return [(self.eta_max - base_lr) * self.T_cur / self.T_up + base_lr for base_lr in self.base_lrs]\n        else:\n            return [base_lr + (self.eta_max - base_lr) * (\n                        1 + math.cos(math.pi * (self.T_cur - self.T_up) / (self.T_i - self.T_up))) / 2\n                    for base_lr in self.base_lrs]\n\n    def step(self, epoch=None):\n        if epoch is None:\n            epoch = self.last_epoch + 1\n            self.T_cur = self.T_cur + 1\n            if self.T_cur >= self.T_i:\n                self.cycle += 1\n                self.T_cur = self.T_cur - self.T_i\n                self.T_i = (self.T_i - self.T_up) * self.T_mult + self.T_up\n        else:\n            if epoch >= self.T_0:\n                if self.T_mult == 1:\n                    self.T_cur = epoch % self.T_0\n                    self.cycle = epoch // self.T_0\n                else:\n                    n = int(math.log((epoch / self.T_0 * (self.T_mult - 1) + 1), self.T_mult))\n                    self.cycle = n\n                    self.T_cur = epoch - self.T_0 * (self.T_mult ** n - 1) / (self.T_mult - 1)\n                    self.T_i = self.T_0 * self.T_mult ** (n)\n            else:\n                self.T_i = self.T_0\n                self.T_cur = epoch\n\n        self.eta_max = self.base_eta_max * (self.gamma ** self.cycle)\n        self.last_epoch = math.floor(epoch)\n        for param_group, lr in zip(self.optimizer.param_groups, self.get_lr()):\n            param_group['lr'] = lr\n\n\n            \n            \n            \n\n\nfrom models.resnext import resnext101\n\nclass SegModel_resnext101(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # original kaggle code\n        # self.encoder = generate_model(model_depth=args.resnet_depth, n_input_channels=1)\n\n        # original paper code\n        self.encoder = resnext101(sample_size=112,\n                                  sample_duration=16,\n                                  shortcut_type='B',\n                                  cardinality=32,\n                                  num_classes=600)\n\n        # original paper code\n        self.decoder = Decoder(encoder_dims=[256, 512, 1024, 2048], upscale=4)\n\n    def forward(self, x):\n        feat_maps = self.encoder(x)\n        feat_maps_pooled = [torch.mean(f, dim=2) for f in feat_maps]\n        pred_mask = self.decoder(feat_maps_pooled)\n        return pred_mask\n\n            \n            \n        \nif args.model == 'resnet152':          \n\n    model = SegModel()\n    model.load_pretrained_weights(torch.load('/kaggle/input/weights/weights/r3d152_KM_200ep.pth')[\"state_dict\"])\n    model = model.to(device)\n    \n    \nelif args.model == 'resnet200':          \n\n    model = SegModel()\n    model.load_pretrained_weights(torch.load('/kaggle/input/weights/weights/r3d200_KM_200ep.pth')[\"state_dict\"])\n    model = model.to(device)\n    \n    \nelif args.model == 'resnext101':          \n\n    model = SegModel_resnext101()\n\n    checkpoint = '/kaggle/input/weights/weights/kinetics_resnext_101_RGB_16_best.pth'\n    state_dict = torch.load(checkpoint)['state_dict']\n\n    from collections import OrderedDict\n    checkpoint_custom = OrderedDict()\n    for key_model, key_checkpoint in zip(model.encoder.state_dict().keys(), state_dict.keys()):\n        checkpoint_custom.update({f'{key_model}' : state_dict[f'{key_checkpoint}']})\n\n    model.encoder.load_state_dict(checkpoint_custom, strict=True)\n    model.encoder.conv1 = nn.Conv3d(1, 64, kernel_size=(7, 7, 7), stride=(1, 2, 2), padding=(3, 3, 3), bias=False)\n    model = model.to(device)\n\n    \n    \n\noptimizer = AdamW(model.parameters(), lr=0)\nscheduler = CosineAnnealingWarmUpRestarts(optimizer, T_0=9, T_mult=1, eta_max=1e-4,  T_up=3, gamma=0.5)\n\n\n# DiceLoss = smp.losses.DiceLoss(mode='binary')\n# BCELoss = smp.losses.SoftBCEWithLogitsLoss()\nCELoss = nn.CrossEntropyLoss(label_smoothing=args.ls)\n\n# alpha = 0.5\n# beta = 1 - alpha\n# TverskyLoss = smp.losses.TverskyLoss(\n#     mode='binary', log_loss=False, alpha=alpha, beta=beta)\n\n\ndef fbeta_loss(preds, targets, beta=0.5, smooth=1e-5):\n    \"\"\"\n    https://www.kaggle.com/competitions/vesuvius-challenge-ink-detection/discussion/397288\n    \"\"\"\n    \n    preds = torch.sigmoid(preds)\n\n    y_true_count = targets.sum()\n    ctp = preds[targets==1].sum()\n    cfp = preds[targets==0].sum()\n    beta_squared = beta * beta\n\n    c_precision = ctp / (ctp + cfp + smooth)\n    c_recall = ctp / (y_true_count + smooth)\n    dice = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall + smooth)\n\n    return (1 - dice)**(args.fbeta_gamma)\n\n\n\ndef focal_tversky_loss(preds, targets, alpha=0.7, beta=0.3, epsilon=1e-6, gamma=1):\n\n    preds = torch.sigmoid(preds)\n\n    preds = preds.reshape(-1)\n    targets = targets.reshape(-1)\n\n    TP = (preds * targets).sum()\n    FP = ((1-targets) * preds).sum()\n    FN = (targets * (1-preds)).sum()\n    Tversky = (TP + epsilon) / (TP + alpha*FP + beta*FN + epsilon)\n    FocalTversky = (1 - Tversky)**gamma\n\n    return FocalTversky\n\n\n\ndef criterion(y_pred, y_true):\n    return CELoss(y_pred, y_true)\n\n\n\ndef rand_bbox(size, lam):\n    W = size[-2]\n    H = size[-1]\n    cut_rat = np.sqrt(1. - lam)\n    cut_w = int(W * cut_rat)\n    cut_h = int(H * cut_rat)\n\n    #uniform\n    cx = np.random.randint(W)\n    cy = np.random.randint(H)\n\n    bbx1 = np.clip(cx - cut_w // 2, 0, W)\n    bby1 = np.clip(cy - cut_h // 2, 0, H)\n    bbx2 = np.clip(cx + cut_w // 2, 0, W)\n    bby2 = np.clip(cy + cut_h // 2, 0, H)\n\n    return bbx1, bby1, bbx2, bby2\n\ncutmix = True\nbeta = 1\n\ndef train_fn(train_loader, model, criterion, optimizer, device):\n    model.train()\n\n    scaler = GradScaler(enabled=CFG.use_amp)\n    losses = AverageMeter()\n\n    for step, (images, labels) in tqdm(enumerate(train_loader), total=len(train_loader)):\n        break\n        images = images.to(device)\n        labels = labels[:, 0, ...].long().to(device)\n        batch_size = labels.size(0)\n            \n        if cutmix and random.random() > 0.4:\n            lam = np.random.beta(beta, beta)\n            rand_index = torch.randperm(images.size()[0]).cuda()\n            bbx1, bby1, bbx2, bby2 = rand_bbox(images.size(), lam)\n\n            images[:, :, :, bbx1:bbx2, bby1:bby2] = images[rand_index, :, :, bbx1:bbx2, bby1:bby2]\n            labels[:, bbx1:bbx2, bby1:bby2] = labels[rand_index, bbx1:bbx2, bby1:bby2]\n\n\n\n        with autocast(CFG.use_amp):\n            y_preds = model(images)\n            loss = criterion(y_preds, labels)\n\n        losses.update(loss.item(), batch_size)\n        scaler.scale(loss).backward()\n\n        grad_norm = torch.nn.utils.clip_grad_norm_(\n            model.parameters(), CFG.max_grad_norm)\n\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad()\n\n    return losses.avg\n\ndef valid_fn(valid_loader, model, criterion, device, valid_xyxys, valid_mask_gt):\n    mask_pred = np.zeros(valid_mask_gt.shape)\n    mask_count = np.zeros(valid_mask_gt.shape)\n\n    model.eval()\n    losses = AverageMeter()\n\n    for step, (images, labels, mask_frag) in tqdm(enumerate(valid_loader), total=len(valid_loader)):\n\n        break\n        \n        if mask_frag.max() == 0:\n            continue\n\n        images = images.to(device)\n        labels = labels[:, 0, ...].long().to(device)\n        batch_size = labels.size(0)\n\n        with torch.no_grad():\n            y_preds = model(images)\n            loss = criterion(y_preds, labels)\n        losses.update(loss.item(), batch_size)\n\n        # make whole mask\n        y_preds = torch.softmax(y_preds, 1)[:, 1, ...].to('cpu').numpy()\n        start_idx = step*CFG.valid_batch_size\n        end_idx = start_idx + batch_size\n        for i, (x1, y1, x2, y2) in enumerate(valid_xyxys[start_idx:end_idx]):\n            mask_pred[y1:y2, x1:x2] += y_preds[i]#.squeeze(0)\n            mask_count[y1:y2, x1:x2] += np.ones((CFG.tile_size, CFG.tile_size))\n\n            \n            \n    # print(f'mask_count_min: {mask_count.min()}')\n    mask_pred /= mask_count\n    return losses.avg, mask_pred\n\n\n\n\nfrom sklearn.metrics import fbeta_score\n\ndef fbeta_numpy(targets, preds, beta=0.5, smooth=1e-5):\n    \"\"\"\n    https://www.kaggle.com/competitions/vesuvius-challenge-ink-detection/discussion/397288\n    \"\"\"\n    y_true_count = targets.sum()\n    ctp = preds[targets==1].sum()\n    cfp = preds[targets==0].sum()\n    beta_squared = beta * beta\n\n    c_precision = ctp / (ctp + cfp + smooth)\n    c_recall = ctp / (y_true_count + smooth)\n    dice = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall + smooth)\n\n    return dice, ctp, cfp, (y_true_count - ctp)\n\n\n\n\ndef calc_fbeta(mask, mask_pred):\n    mask = mask.astype(int).flatten()\n    mask_pred = mask_pred.flatten()\n\n    best_th = 0\n    best_dice = 0\n    for th in np.array(range(10, 95+1, 5)) / 100:\n        \n        # dice = fbeta_score(mask, (mask_pred >= th).astype(int), beta=0.5)\n        dice, ctp, cfp, cfn = fbeta_numpy(mask, (mask_pred >= th).astype(int), beta=0.5)\n        Logger.info(f'th: {th}, fbeta: {dice}')\n\n        if dice > best_dice:\n            best_dice = dice\n            best_th = th\n\n    dice, ctp, cfp, cfn = fbeta_numpy(mask, (mask_pred >= 0.5).astype(int), beta=0.5)\n    Logger.info(f'best_th: {best_th}, fbeta: {best_dice}')\n    return dice, 0.5, ctp, cfp, cfn\n\n\ndef calc_cv(mask_gt, mask_pred):\n    best_dice, best_th, ctp, cfp, cfn = calc_fbeta(mask_gt, mask_pred)\n\n    return best_dice, best_th, ctp, cfp, cfn\n\n\n\nfragment_id = CFG.valid_id\n\nif fragment_id == 2:\n    valid_mask_gt = cv2.imread(CFG.comp_dataset_path + f\"train/2/inklabels.png\", 0)\n    valid_mask_gt = valid_mask_gt[:, :args.slicing_num]\n\nelif fragment_id == 4:\n    valid_mask_gt = cv2.imread(CFG.comp_dataset_path + f\"train/2/inklabels.png\", 0)\n    valid_mask_gt = valid_mask_gt[:, args.slicing_num:]\n\nelse:\n    valid_mask_gt = cv2.imread(CFG.comp_dataset_path + f\"train/{fragment_id}/inklabels.png\", 0)\n\nvalid_mask_gt = valid_mask_gt / 255\npad0 = (CFG.tile_size - valid_mask_gt.shape[0] % CFG.tile_size)\npad1 = (CFG.tile_size - valid_mask_gt.shape[1] % CFG.tile_size)\nvalid_mask_gt = np.pad(valid_mask_gt, [(0, pad0), (0, pad1)], constant_values=0)\n\n\n\n\nif fragment_id == 2:\n    valid_mask_gt_frag = cv2.imread(CFG.comp_dataset_path + f\"train/2/mask.png\", 0)\n    valid_mask_gt_frag = valid_mask_gt_frag[:, :args.slicing_num]\n\nelif fragment_id == 4:\n    valid_mask_gt_frag = cv2.imread(CFG.comp_dataset_path + f\"train/2/mask.png\", 0)\n    valid_mask_gt_frag = valid_mask_gt_frag[:, args.slicing_num:]\n\nelse:\n    valid_mask_gt_frag = cv2.imread(CFG.comp_dataset_path + f\"train/{fragment_id}/mask.png\", 0)\n\nvalid_mask_gt_frag = valid_mask_gt_frag / 255\npad0 = (CFG.tile_size - valid_mask_gt_frag.shape[0] % CFG.tile_size)\npad1 = (CFG.tile_size - valid_mask_gt_frag.shape[1] % CFG.tile_size)\nvalid_mask_gt_frag = np.pad(valid_mask_gt_frag, [(0, pad0), (0, pad1)], constant_values=0)\nvalid_mask_gt_frag = valid_mask_gt_frag.astype('bool')\n\n\n\nfold = CFG.valid_id\n\nif CFG.metric_direction == 'minimize':\n    best_score = np.inf\nelif CFG.metric_direction == 'maximize':\n    best_score = -1\n\nbest_loss = np.inf\n\nfor epoch in range(CFG.epochs):\n\n    scheduler.step()\n    Logger.info(f'Lr : {scheduler.get_lr()}')\n\n\n    start_time = time.time()\n\n    # train\n    avg_loss = train_fn(train_loader, model, criterion, optimizer, device)\n    Logger.info(f'Train Loss : {avg_loss}')\n\n\n    if epoch >= 20:\n        # eval\n        avg_val_loss, mask_pred = valid_fn(\n            valid_loader, model, criterion, device, valid_xyxys, valid_mask_gt)\n\n        mask_pred = mask_pred * valid_mask_gt_frag\n\n\n        best_dice, best_th, ctp, cfp, cfn = calc_cv(valid_mask_gt, mask_pred)\n\n        # score = avg_val_loss\n        score = best_dice\n\n        elapsed = time.time() - start_time\n\n        Logger.info(\n            f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n        # Logger.info(f'Epoch {epoch+1} - avgScore: {avg_score:.4f}')\n        Logger.info(\n            f'Epoch {epoch+1} - avgScore: {score:.4f}')\n\n        if CFG.metric_direction == 'minimize':\n            update_best = score < best_score\n        elif CFG.metric_direction == 'maximize':\n            update_best = score > best_score\n\n        if update_best:\n            best_loss = avg_val_loss\n            best_score = score\n\n            Logger.info(\n                f'Epoch {epoch+1} - Save Best Score: {best_score:.4f} Model')\n            Logger.info(\n                f'Epoch {epoch+1} - Save Best Loss: {best_loss:.4f} Model')\n\n        torch.save({'model': model.state_dict()},\n                    CFG.model_dir + f'{CFG.model_name}_fold{fold}_epoch{epoch+1}_score{score}.pth')\n\n\n        mask_pred_png = (mask_pred > 0.5).astype('uint8')*255\n        cv2.imwrite(f'{CFG.mask_dir}/Fold{args.valid_id}_epoch{epoch+1}_score{score}_ctp{ctp}_cfp{cfp}_cfn{cfn}.png', mask_pred_png)\n\n        Logger.info('\\n')\n\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]}]}