{"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"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":47317,"databundleVersionId":5799376,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!/usr/bin/env python\nimport ssl\nssl._create_default_https_context = ssl._create_unverified_context\nimport os\n# os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"1\"\nimport warnings\nimport torch.nn as nn\nimport torch\nimport time\nimport sys\nimport shutil\nimport segmentation_models_pytorch as smp\nimport scipy as sp\nimport random\nimport pickle\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport math\nimport importlib\nimport gc\nimport datetime\nimport cv2\nimport argparse\nimport albumentations as A\nfrom tqdm.auto import tqdm\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.cuda.amp import autocast, GradScaler\nfrom sklearn.metrics import roc_auc_score, accuracy_score, f1_score, log_loss\nfrom pathlib import Path\nfrom functools import partial\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nfrom albumentations.pytorch import ToTensorV2\n\nreplace_fold2fragmentid = {\n    1:1,\n    2:2,\n    3:2,\n    4:2,\n    5:3\n}\n\nclass CFG:\n    # ============== comp exp name =============\n    comp_name = 'vesuvius'\n\n    # comp_dir_path = './'\n    comp_dir_path = '../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    filename = exp_name = os.path.splitext(os.path.basename(__file__))[0]\n\n    # ============== pred target =============\n    target_size = 1\n\n    # ============== model cfg =============\n    model_name = 'Unet'\n    backbone = 'se_resnext50_32x4d'\n    # backbone = 'se_resnext50_32x4d'\n\n    in_chans = 6 # 65\n    # ============== training cfg =============\n\n    dont_train_zero_image = True\n    image_channel_normalize = True\n    size = 224\n    tile_size = 224\n    stride = tile_size // 24\n    valid_stride = tile_size // 8\n    \n    tif_input_dir = comp_dir_path + f'tifdir_size{size}_tilesize{tile_size}_stride{stride}_validstride{valid_stride}_inchans{in_chans}'\n\n    train_batch_size = 256 # 32\n    valid_batch_size = train_batch_size\n    use_amp = True\n\n    scheduler = 'GradualWarmupSchedulerV2'\n    # scheduler = 'CosineAnnealingLR'\n    epochs = 15 # 30\n\n    PATIENCE = 5 #earlystopping\n\n    # adamW warmupあり\n    warmup_factor = 10\n    # lr = 1e-4 / warmup_factor\n    lr = 1e-4 / warmup_factor\n\n    # ============== fold =============\n    valid_id = 2\n    valid_ids = [1,2,3,4,5]\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-6\n    weight_decay = 1e-6\n    max_grad_norm = 1000\n\n    print_freq = 50\n    num_workers = 4\n\n    seed = 42\n\n    # ============== set dataset path =============\n    print('set dataset path')\n\n    outputs_path = f'../outputs/{comp_name}/{exp_name}/'\n\n    submission_dir = outputs_path + 'submissions/'\n    submission_path = submission_dir + f'submission_{exp_name}.csv'\n\n    model_dir = outputs_path + f'{comp_name}-models/'\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.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomBrightnessContrast(p=0.75),\n        A.ShiftScaleRotate(p=0.75),\n        A.OneOf([\n                A.GaussNoise(var_limit=[10, 50]),\n                A.GaussianBlur(),\n                A.MotionBlur(),\n                ], p=0.4),\n        #A.GridDistortion(num_steps=5, distort_limit=0.3, p=0.5),\n        A.CoarseDropout(max_holes=1, max_width=int(size * 0.3), max_height=int(size * 0.3), \n                        mask_fill_value=0, p=0.5),\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] * in_chans,\n        #     std= [1] * 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] * in_chans,\n        #     std= [1] * in_chans\n        # ),\n        ToTensorV2(transpose_mask=True),\n    ]\n\n\n# ## helper\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\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\ndef make_dirs(cfg):\n    for dir in [cfg.model_dir, cfg.figures_dir, cfg.submission_dir, cfg.log_dir]:\n        os.makedirs(dir, exist_ok=True)\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\ndef to_pickle(var, OUTPUT_PATH):\n    with open(OUTPUT_PATH, 'wb') as f:\n        pickle.dump(var, f)\n    print(f'pickle saved {OUTPUT_PATH}')\n\ndef read_pickle(PICKLE_PATH):\n    with open(PICKLE_PATH, 'rb') as f:\n        loaded_pickle = pickle.load(f)  # 復元\n    return loaded_pickle\n\n\n# ## image, mask\n\ndef crop_iamge_fragment2(image, fold):\n    replace_inputfragment2imageshape = {\n        2:[0, 6113],\n        3:[6113, 10609],\n        4:[10609, 14830],\n    }\n    image_shape = replace_inputfragment2imageshape[fold]\n    image = image[image_shape[0]:image_shape[1]]\n    return image\n\ndef read_image_mask_5fold(fold):\n\n    fragment_id = replace_fold2fragmentid[fold]\n    print(f'fold:{fold} fragment_id:{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 = range(start, end)\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\", cv2.IMREAD_UNCHANGED)\n        if CFG.image_channel_normalize:\n            Logger.info('画像channelごとに標準化して読み込みます')\n            transform = A.Compose([\n                A.ToFloat(max_value=65535.0),\n                A.FromFloat(max_value=65535.0),\n                A.Normalize(mean=[0], std=[1]),\n                A.Normalize(mean=[0], std=[1]),\n            ])\n            image = transform(image=image)['image']\n        if fragment_id==2:\n            image = crop_iamge_fragment2(image, fold)\n\n        print(f'image shape:{image.shape}')\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    '''\n    images.shape\n    (8288, 6496, 6)\n    '''\n\n    mask = cv2.imread(CFG.comp_dataset_path + f\"train/{fragment_id}/inklabels.png\", 0)\n    if fragment_id==2:\n        mask = crop_iamge_fragment2(mask, fold)\n    print(f'mask shape:{mask.shape}')\n    mask = np.pad(mask, [(0, pad0), (0, pad1)], constant_values=0)\n\n    mask = mask.astype('float32')\n    mask /= 255.0\n    \n    return images, mask\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 = range(start, end)\n\n    #midからin_chans分幅を持って画像を読みにいく設定\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\", cv2.IMREAD_UNCHANGED)\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        #tileの合計枚数が合うように0paddingしている\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\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    return images, mask\n\n\ndef get_train_valid_dataset():\n\n    def check_not_zero_image(image):\n        if image.sum() != 0:\n            return True\n        else:\n            return False\n\n    train_images = []\n    train_masks = []\n\n    valid_images = []\n    valid_masks = []\n    valid_xyxys = []\n\n    count_zero_image = 0\n    count_not_zero_image = 0\n    pathes = [\n        os.path.join(CFG.tif_input_dir, f'train_images_valid_fragment_{CFG.valid_id}.npy'), \n        os.path.join(CFG.tif_input_dir, f'train_masks_valid_fragment_{CFG.valid_id}.npy'), \n        os.path.join(CFG.tif_input_dir, f'valid_images_valid_fragment_{CFG.valid_id}.npy'),\n        os.path.join(CFG.tif_input_dir, f'valid_masks_valid_fragment_{CFG.valid_id}.npy'),\n        os.path.join(CFG.tif_input_dir, f'valid_xyxys_valid_fragment_{CFG.valid_id}.npy')\n    ]\n    if not os.path.isfile(pathes[0]):\n        for fragment_id in range(1, 6):\n\n            image, mask= read_image_mask_5fold(fragment_id)\n            print(image.shape, mask.shape)\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            for y1 in y1_list:\n                for x1 in x1_list:\n                    \n                    y2 = y1 + CFG.tile_size\n                    x2 = x1 + CFG.tile_size\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                    else:\n                        if CFG.dont_train_zero_image:\n                            if check_not_zero_image(image[y1:y2, x1:x2]):\n                                train_images.append(image[y1:y2, x1:x2])\n                                train_masks.append(mask[y1:y2, x1:x2, None])\n                                count_not_zero_image += 1\n                            else:\n                                count_zero_image += 1\n                        else:\n                            train_images.append(image[y1:y2, x1:x2])\n                            train_masks.append(mask[y1:y2, x1:x2, None])\n        print(count_not_zero_image, count_zero_image)\n        os.makedirs(CFG.tif_input_dir, exist_ok=True)\n\n        # np.save(pathes[0], np.array(train_images))\n        # np.save(pathes[1], np.array(train_masks))\n        # np.save(pathes[2], np.array(valid_images))\n        # np.save(pathes[3], np.array(valid_masks))\n        # np.save(pathes[4], np.array(valid_xyxys))\n    else:\n        train_images = list(np.load(pathes[0]))\n        train_masks = list(np.load(pathes[1]))\n        valid_images = list(np.load(pathes[2]))\n        valid_masks = list(np.load(pathes[3]))\n        valid_xyxys = list(np.load(pathes[4]))\n\n    return train_images, train_masks, valid_images, valid_masks, valid_xyxys\n\n\n# ## dataset\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, transform=None):\n        self.images = images\n        self.cfg = cfg\n        self.labels = labels\n        self.transform = transform\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        label = self.labels[idx]\n\n        if self.transform:\n            data = self.transform(image=image, mask=label)#ここのmaskの意味https://albumentations.ai/docs/getting_started/mask_augmentation/\n            image = data['image']\n            label = data['mask']\n\n        return image, label\n\n\n# ## model\n\nclass CustomModel(nn.Module):\n    def __init__(self, cfg, weight=None):\n        super().__init__()\n        self.cfg = cfg\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    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\nclass EarlyStopping:\n    def __init__(self, patience=7, verbose=False, delta=0, fold=\"\"):\n        self.patience = patience\n        self.verbose = verbose\n        self.counter = 0\n        self.best_score = None\n        self.early_stop = False\n        self.val_loss_min = np.Inf\n        self.delta = delta\n\n    def __call__(self, val_loss, model):\n\n        score = -val_loss\n\n        if self.best_score is None:\n            self.best_score = score\n            self.save_checkpoint(val_loss, model)\n        elif score < self.best_score + self.delta:\n            self.counter += 1\n            Logger.info(f\"EarlyStopping counter: {self.counter} out of {self.patience}\")\n            if self.counter >= self.patience:\n                self.early_stop = True\n        else:\n            self.best_score = score\n            self.save_checkpoint(val_loss, model)\n            self.counter = 0\n\n    def save_checkpoint(self, val_loss, model):\n        \"\"\"Saves model when validation loss decrease.\"\"\"\n        if self.verbose:\n            Logger.info(\n                f\"Validation loss decreased ({self.val_loss_min:.6f} --> {val_loss:.6f}).  Saving model ...\"\n            )\n        torch.save(model.state_dict(), CFG.model_dir + f'{CFG.model_name}_fold{fold}_checkpoint.pth')\n        self.val_loss_min = val_loss\n\n# ## scheduler\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\nfrom warmup_scheduler import GradualWarmupScheduler\n\n\nclass GradualWarmupSchedulerV2(GradualWarmupScheduler):\n    \"\"\"\n    https://www.kaggle.com/code/underwearfitting/single-fold-training-of-resnet200d-lb0-965\n    \"\"\"\n    def __init__(self, optimizer, multiplier, total_epoch, after_scheduler=None):\n        super(GradualWarmupSchedulerV2, self).__init__(\n            optimizer, multiplier, total_epoch, after_scheduler)\n\n    def get_lr(self):\n        if self.last_epoch > self.total_epoch:\n            if self.after_scheduler:\n                if not self.finished:\n                    self.after_scheduler.base_lrs = [\n                        base_lr * self.multiplier for base_lr in self.base_lrs]\n                    self.finished = True\n                return self.after_scheduler.get_lr()\n            return [base_lr * self.multiplier for base_lr in self.base_lrs]\n        if self.multiplier == 1.0:\n            return [base_lr * (float(self.last_epoch) / self.total_epoch) for base_lr in self.base_lrs]\n        else:\n            return [base_lr * ((self.multiplier - 1.) * self.last_epoch / self.total_epoch + 1.) for base_lr in self.base_lrs]\n\ndef get_scheduler(cfg, optimizer):\n    scheduler_cosine = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer, cfg.epochs, eta_min=1e-7)\n    scheduler = GradualWarmupSchedulerV2(\n        optimizer, multiplier=10, total_epoch=1, after_scheduler=scheduler_cosine)\n\n    return scheduler\n\ndef scheduler_step(scheduler, avg_val_loss, epoch):\n    scheduler.step(epoch)\n\n\n# ## loss\nDiceLoss = smp.losses.DiceLoss(mode='binary')\nBCELoss = smp.losses.SoftBCEWithLogitsLoss()\n\nalpha = 0.5\nbeta = 1 - alpha\nTverskyLoss = smp.losses.TverskyLoss(\n    mode='binary', log_loss=False, alpha=alpha, beta=beta)\n\ndef criterion(y_pred, y_true):\n    # return 0.5 * BCELoss(y_pred, y_true) + 0.5 * DiceLoss(y_pred, y_true)\n    return BCELoss(y_pred, y_true)\n    # return 0.5 * BCELoss(y_pred, y_true) + 0.5 * TverskyLoss(y_pred, y_true)\n\n\n# ## train, val\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        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\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) in tqdm(enumerate(valid_loader), total=len(valid_loader)):\n        images = images.to(device)\n        labels = labels.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.sigmoid(y_preds).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    print(f'mask_count_min: {mask_count.min()}')\n    mask_pred /= mask_count\n    return losses.avg, mask_pred\n\n\n# ## metrics\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\n\ndef calc_fbeta(mask, mask_pred ,min_th=45, max_th=60):\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(min_th, max_th+1, 5)) / 100:\n        \n        # dice = fbeta_score(mask, (mask_pred >= th).astype(int), beta=0.5)\n        dice = 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    Logger.info(f'best_th: {best_th}, fbeta: {best_dice}')\n    return best_dice, best_th\n\n\ndef calc_cv(mask_gt, mask_pred):\n    best_dice, best_th = calc_fbeta(mask_gt, mask_pred)\n\n    return best_dice, best_th\n\ndef flatten_valid_mask_3d(fallten_list):\n    output_list = []\n    for a in fallten_list:\n        for b in a:\n            output_list.append(b)\n    return output_list\n\n# ## main\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\nbest_dices = []\nbest_thes = []\n\nvalid_ans = []\nmask_ans = []\n#for all fold\nfor valid_id in CFG.valid_ids:\n    CFG.valid_id = valid_id\n    Logger.info(f'\\n\\n-------- fold{valid_id}-----------------')\n    # Logger.info(datetime.datetime.now().strftime('%Y年%m月%d日 %H:%M:%S'))\n    train_images, train_masks, valid_images, valid_masks, valid_xyxys = get_train_valid_dataset()\n\n    train_dataset = CustomDataset(\n        train_images, CFG, labels=train_masks, transform=get_transforms(data='train', cfg=CFG))\n    valid_dataset = CustomDataset(\n        valid_images, CFG, labels=valid_masks, transform=get_transforms(data='valid', cfg=CFG))\n\n    train_loader = DataLoader(train_dataset,\n                            batch_size=CFG.train_batch_size,\n                            shuffle=True,\n                            num_workers=CFG.num_workers, pin_memory=True, drop_last=True,\n                            )\n    valid_loader = DataLoader(valid_dataset,\n                            batch_size=CFG.valid_batch_size,\n                            shuffle=False,\n                            num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n\n    transform = CFG.train_aug_list\n    transform = A.Compose(\n        [t for t in transform if not isinstance(t, (A.Normalize, ToTensorV2))])\n\n    model = build_model(CFG)\n    model.to(device)\n\n    optimizer = AdamW(model.parameters(), lr=CFG.lr)\n    scheduler = get_scheduler(CFG, optimizer)\n\n    fragment_id = replace_fold2fragmentid[CFG.valid_id]\n\n    valid_mask_gt = cv2.imread(CFG.comp_dataset_path + f\"train/{fragment_id}/inklabels.png\", 0)\n    if CFG.valid_id in [2,3,4]:\n        valid_mask_gt = crop_iamge_fragment2(valid_mask_gt, CFG.valid_id)\n\n    valid_mask_gt = valid_mask_gt / 255\n    pad0 = (CFG.tile_size - valid_mask_gt.shape[0] % CFG.tile_size)\n    pad1 = (CFG.tile_size - valid_mask_gt.shape[1] % CFG.tile_size)\n    valid_mask_gt = np.pad(valid_mask_gt, [(0, pad0), (0, pad1)], constant_values=0)\n\n    fold = CFG.valid_id\n\n    if CFG.metric_direction == 'minimize':\n        best_score = np.inf\n    elif CFG.metric_direction == 'maximize':\n        best_score = -1\n\n    best_loss = np.inf\n    early_stopping = EarlyStopping(\n            patience=CFG.PATIENCE, verbose=True, fold=fold\n        )\n\n    for epoch in range(CFG.epochs):\n\n        start_time = time.time()\n\n        # train\n        avg_loss = train_fn(train_loader, model, criterion, optimizer, device)\n\n        # eval\n        avg_val_loss, mask_pred = valid_fn(\n            valid_loader, model, criterion, device, valid_xyxys, valid_mask_gt)\n\n        scheduler_step(scheduler, avg_val_loss, epoch)\n\n        best_dice, best_th = 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                        'preds': mask_pred},\n                        CFG.model_dir + f'{CFG.model_name}_fold{fold}_best.pth')\n\n        early_stopping(-score, model)\n\n        if early_stopping.early_stop:\n            Logger.info(\"Early stopping\")\n            break\n\n    check_point = torch.load(\n        CFG.model_dir + f'{CFG.model_name}_fold{fold}_{CFG.inf_weight}.pth', map_location=torch.device('cpu'))\n    mask_pred = check_point['preds']\n\n    best_dice, best_th  = calc_fbeta(valid_mask_gt, mask_pred)\n    valid_ans.append(np.ravel(valid_mask_gt))\n    mask_ans.append(np.ravel(mask_pred))\n    Logger.info(f'best_dice:{best_dice}')\n    Logger.info(f'best_th:{best_th}')\n\n    best_dices.append(best_dice)\n    best_thes.append(best_th)\n\nfor valid_id, best_dice, best_th in zip(CFG.valid_ids, best_dices, best_thes):\n    Logger.info(f'fold{valid_id} best_dice:{best_dice} best_th:{best_th}')\n\nall_valid = flatten_valid_mask_3d(valid_ans)\nall_mask = flatten_valid_mask_3d(mask_ans)\ncalc_fbeta(np.ravel(all_valid), np.ravel(all_mask),min_th=20, max_th=80)","metadata":{"_uuid":"40c5734b-b574-4e5a-968e-456af62ff143","_cell_guid":"b45c429a-def2-4419-b968-c27cb7674457","collapsed":false,"jupyter":{"outputs_hidden":false},"trusted":true},"execution_count":null,"outputs":[]}]}