{"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":"!pip install segmentation_models_pytorch\n!pip install warmup_scheduler\nimport os\nans=os.environ.get('KAGGLE_KERNEL_RUN_TYPE', 'Localhost')\n\nif ans=='Interactive':\n    TEST_RUN=True\n    \nelif ans=='Batch':\n    TEST_RUN=False\nelse:\n    TEST_RUN=False\nprint(ans)","metadata":{"papermill":{"duration":30.474783,"end_time":"2023-05-25T20:22:40.433686","exception":false,"start_time":"2023-05-25T20:22:09.958903","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:01:00.023076Z","iopub.execute_input":"2023-06-05T19:01:00.024140Z","iopub.status.idle":"2023-06-05T19:01:32.977383Z","shell.execute_reply.started":"2023-06-05T19:01:00.024071Z","shell.execute_reply":"2023-06-05T19:01:32.975751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{"papermill":{"duration":0.012219,"end_time":"2023-05-25T20:22:40.458687","exception":false,"start_time":"2023-05-25T20:22:40.446468","status":"completed"},"tags":[]}},{"cell_type":"code","source":"NOTES='3 inchans model with mit_b2'\nprint(\"test run:- \", TEST_RUN)","metadata":{"papermill":{"duration":0.022065,"end_time":"2023-05-25T20:22:40.492682","exception":false,"start_time":"2023-05-25T20:22:40.470617","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:01:32.980243Z","iopub.execute_input":"2023-06-05T19:01:32.980657Z","iopub.status.idle":"2023-06-05T19:01:32.986300Z","shell.execute_reply.started":"2023-06-05T19:01:32.980615Z","shell.execute_reply":"2023-06-05T19:01:32.985344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport scipy as sp\nfrom sklearn.metrics import roc_auc_score, accuracy_score, f1_score, log_loss\nimport matplotlib.pyplot as plt\nimport sys\nimport os\nimport gc\nimport sys\nimport pickle\nimport warnings\nimport math\nimport time\nimport random\nimport argparse\nimport importlib\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.optim import Adam, SGD, AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\nimport segmentation_models_pytorch as smp\nfrom warmup_scheduler import GradualWarmupScheduler\nimport cv2\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\nimport datetime\nfrom kaggle_secrets import UserSecretsClient\nimport wandb","metadata":{"papermill":{"duration":8.017263,"end_time":"2023-05-25T20:22:48.521956","exception":false,"start_time":"2023-05-25T20:22:40.504693","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:01:32.987656Z","iopub.execute_input":"2023-06-05T19:01:32.987988Z","iopub.status.idle":"2023-06-05T19:01:40.153720Z","shell.execute_reply.started":"2023-06-05T19:01:32.987954Z","shell.execute_reply":"2023-06-05T19:01:40.152771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import ssl\nssl._create_default_https_context = ssl._create_unverified_context","metadata":{"papermill":{"duration":0.019178,"end_time":"2023-05-25T20:22:48.553901","exception":false,"start_time":"2023-05-25T20:22:48.534723","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:01:40.156409Z","iopub.execute_input":"2023-06-05T19:01:40.156740Z","iopub.status.idle":"2023-06-05T19:01:40.162348Z","shell.execute_reply.started":"2023-06-05T19:01:40.156708Z","shell.execute_reply":"2023-06-05T19:01:40.160514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Main Configuration Class","metadata":{"papermill":{"duration":0.011925,"end_time":"2023-05-25T20:22:48.578369","exception":false,"start_time":"2023-05-25T20:22:48.566444","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class CFG:\n    # ============== comp exp name =============\n    comp_name = 'vesuvius'\n    exp_name = 'sgm_seresnext'\n    comp_dir_path = '/kaggle/input/'\n    comp_folder_name = 'vesuvius-dataset-split-5'\n    comp_dataset_path = f'{comp_dir_path}{comp_folder_name}/'\n\n    # ============== pred target =============\n    target_size = 1\n\n    # ============== model cfg =============\n    model_name = 'Unet'\n    backbone = 'mit_b2' # 'efficientnet-b0'\n\n    start_chans=27\n    end_chans=30\n#     chans_to_choose=[27,28,29]\n    in_chans=end_chans-start_chans\n    # ============== training cfg =============\n    size = 224\n    tile_size = 224\n    stride = tile_size // 2\n\n    train_batch_size = 32 # 32\n    accumulation_steps=1\n    valid_batch_size = train_batch_size * 1\n    use_amp = True\n\n    scheduler = 'GradualWarmupSchedulerV2' # 'CosineAnnealingLR'\n    epochs = 15 # 30\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 = [1, 2, 3]\n\n    # objective_cv = 'binary'  # 'binary', 'multiclass', 'regression'\n    metric_direction = 'maximize'  # maximize, 'minimize'\n    loss_func= 'binary_cross_entropy'#\"LovaszLoss\"#'Mixed'#'diceloss' #'binary_cross_entropy' #diceloss\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 = 2\n\n    seed = 310\n\n    # ============== set dataset path =============\n    print('set dataset path')\n\n    outputs_path = f'/kaggle/working/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 + \\\n        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.Resize(size, size),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.75),\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.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    \nrun_config_dict = {\n    'model_name': CFG.model_name,\n    'backbone': CFG.backbone,\n    'in_chans': CFG.in_chans,\n    'train_batch_size': CFG.train_batch_size,\n    'valid_batch_size': CFG.valid_batch_size,\n    'scheduler': CFG.scheduler,\n    'epochs': CFG.epochs,\n    'loss_func': CFG.loss_func,\n    'pretrained': CFG.pretrained\n}\n\nRUN_NAME=f\"random rot 90 and HV flip {CFG.loss_func}| {run_config_dict['model_name']} |{run_config_dict['backbone']} |in-chans={run_config_dict['in_chans']} |epochs={run_config_dict['epochs']}\"\nprint(RUN_NAME)\n    ","metadata":{"papermill":{"duration":0.030661,"end_time":"2023-05-25T20:22:48.621331","exception":false,"start_time":"2023-05-25T20:22:48.590670","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:01:40.163554Z","iopub.execute_input":"2023-06-05T19:01:40.163893Z","iopub.status.idle":"2023-06-05T19:01:40.183386Z","shell.execute_reply.started":"2023-06-05T19:01:40.163861Z","shell.execute_reply":"2023-06-05T19:01:40.182253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TEST_RUN:\n    mode='offline'\n    CFG.epochs=1\n    \nelse:\n    user_secrets = UserSecretsClient()\n    secret_value_0 = user_secrets.get_secret(\"wandb\")\n    wandb.login(key=secret_value_0)\n    mode='online'\n    \nwandb.init(\n    project=\"vesuvius-endgame\",\n    mode=mode,\n    config=run_config_dict,\n    notes=NOTES,\n    name=RUN_NAME\n)","metadata":{"papermill":{"duration":34.882158,"end_time":"2023-05-25T20:23:23.515721","exception":false,"start_time":"2023-05-25T20:22:48.633563","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:01:40.185733Z","iopub.execute_input":"2023-06-05T19:01:40.186026Z","iopub.status.idle":"2023-06-05T19:02:15.649407Z","shell.execute_reply.started":"2023-06-05T19:01:40.186002Z","shell.execute_reply":"2023-06-05T19:02:15.648516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Helper Functions (as usual)","metadata":{"papermill":{"duration":0.019378,"end_time":"2023-05-25T20:23:23.554932","exception":false,"start_time":"2023-05-25T20:23:23.535554","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class 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","metadata":{"papermill":{"duration":0.032568,"end_time":"2023-05-25T20:23:23.607875","exception":false,"start_time":"2023-05-25T20:23:23.575307","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.650428Z","iopub.execute_input":"2023-06-05T19:02:15.650731Z","iopub.status.idle":"2023-06-05T19:02:15.671958Z","shell.execute_reply.started":"2023-06-05T19:02:15.650702Z","shell.execute_reply":"2023-06-05T19:02:15.671063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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 = 310\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","metadata":{"papermill":{"duration":0.034801,"end_time":"2023-05-25T20:23:23.661974","exception":false,"start_time":"2023-05-25T20:23:23.627173","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.672915Z","iopub.execute_input":"2023-06-05T19:02:15.673227Z","iopub.status.idle":"2023-06-05T19:02:15.700953Z","shell.execute_reply.started":"2023-06-05T19:02:15.673198Z","shell.execute_reply":"2023-06-05T19:02:15.699891Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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)","metadata":{"papermill":{"duration":0.0299,"end_time":"2023-05-25T20:23:23.711485","exception":false,"start_time":"2023-05-25T20:23:23.681585","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.702215Z","iopub.execute_input":"2023-06-05T19:02:15.702760Z","iopub.status.idle":"2023-06-05T19:02:15.713386Z","shell.execute_reply.started":"2023-06-05T19:02:15.702727Z","shell.execute_reply":"2023-06-05T19:02:15.712135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def cfg_init(cfg, mode='train'):\n    set_seed(cfg.seed)\n\n    if mode == 'train':\n        make_dirs(cfg)","metadata":{"papermill":{"duration":0.029707,"end_time":"2023-05-25T20:23:23.760989","exception":false,"start_time":"2023-05-25T20:23:23.731282","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.727580Z","iopub.execute_input":"2023-06-05T19:02:15.732332Z","iopub.status.idle":"2023-06-05T19:02:15.738886Z","shell.execute_reply.started":"2023-06-05T19:02:15.732299Z","shell.execute_reply":"2023-06-05T19:02:15.737871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg_init(CFG)\n\nLogger = init_logger(log_file=CFG.log_path)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"papermill":{"duration":0.107867,"end_time":"2023-05-25T20:23:23.888257","exception":false,"start_time":"2023-05-25T20:23:23.780390","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.740342Z","iopub.execute_input":"2023-06-05T19:02:15.740670Z","iopub.status.idle":"2023-06-05T19:02:15.794317Z","shell.execute_reply.started":"2023-06-05T19:02:15.740640Z","shell.execute_reply":"2023-06-05T19:02:15.793382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Image and Mask (by mask, it means the inklabels.png)\n\n### The mask will act as the label for segmentation.","metadata":{"papermill":{"duration":0.019076,"end_time":"2023-05-25T20:23:23.927429","exception":false,"start_time":"2023-05-25T20:23:23.908353","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def read_image_mask(fragment_id):\n\n    images = []\n    \n    start = CFG.start_chans\n    end=CFG.end_chans\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\", 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\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","metadata":{"papermill":{"duration":0.035491,"end_time":"2023-05-25T20:23:23.982179","exception":false,"start_time":"2023-05-25T20:23:23.946688","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.795270Z","iopub.execute_input":"2023-06-05T19:02:15.795573Z","iopub.status.idle":"2023-06-05T19:02:15.820965Z","shell.execute_reply.started":"2023-06-05T19:02:15.795545Z","shell.execute_reply":"2023-06-05T19:02:15.820064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_train_valid_dataset(valid_id):\n    train_images = []\n    train_masks = []\n\n    valid_images = []\n    valid_masks = []\n    valid_xyxys = []\n\n    for fragment_id in range(1, 7):\n\n        image, mask = 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 == 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                    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","metadata":{"papermill":{"duration":0.035811,"end_time":"2023-05-25T20:23:24.037698","exception":false,"start_time":"2023-05-25T20:23:24.001887","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.823310Z","iopub.execute_input":"2023-06-05T19:02:15.827490Z","iopub.status.idle":"2023-06-05T19:02:15.850969Z","shell.execute_reply.started":"2023-06-05T19:02:15.827457Z","shell.execute_reply":"2023-06-05T19:02:15.850061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset Classes","metadata":{"papermill":{"duration":0.019024,"end_time":"2023-05-25T20:23:24.076265","exception":false,"start_time":"2023-05-25T20:23:24.057241","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def 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    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.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)\n            image = data['image']\n            label = data['mask']\n\n        return image, label","metadata":{"papermill":{"duration":0.033951,"end_time":"2023-05-25T20:23:24.129577","exception":false,"start_time":"2023-05-25T20:23:24.095626","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.855558Z","iopub.execute_input":"2023-06-05T19:02:15.857825Z","iopub.status.idle":"2023-06-05T19:02:15.875020Z","shell.execute_reply.started":"2023-06-05T19:02:15.857793Z","shell.execute_reply":"2023-06-05T19:02:15.874153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model (finally the good stuff)","metadata":{"papermill":{"duration":0.018988,"end_time":"2023-05-25T20:23:24.167870","exception":false,"start_time":"2023-05-25T20:23:24.148882","status":"completed"},"tags":[]}},{"cell_type":"code","source":"  \nclass CustomModel(nn.Module):\n    def __init__(self, cfg, weight=\"imagenet\"):\n        super().__init__()\n        self.cfg = cfg\n#         self.encoder01 = nn.Conv2d(self.inchans, 3, kernel_size=1)\n\n        self.encoder02=smp.Unet(\n            encoder_name=\"mit_b2\", \n            encoder_weights=weight,\n            in_channels=3,\n            classes=1,\n            activation=None,\n        )\n        self.stacked_unet=nn.Sequential(self.encoder02)\n\n    def forward(self, image):\n\n        output=self.stacked_unet(image)\n        # output = output.squeeze(-1)\n        return output","metadata":{"papermill":{"duration":0.031838,"end_time":"2023-05-25T20:23:24.218859","exception":false,"start_time":"2023-05-25T20:23:24.187021","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.880533Z","iopub.execute_input":"2023-06-05T19:02:15.883315Z","iopub.status.idle":"2023-06-05T19:02:15.899411Z","shell.execute_reply.started":"2023-06-05T19:02:15.883283Z","shell.execute_reply":"2023-06-05T19:02:15.898508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def build_model(cfg, weight=\"imagenet\"):\n    print('model_name', cfg.model_name)\n    print('backbone', cfg.backbone)\n#     if TEST_RUN:\n#         weight=None\n    weight=\"imagenet\"\n    model = CustomModel(cfg, weight)\n\n    return model","metadata":{"papermill":{"duration":0.028512,"end_time":"2023-05-25T20:23:24.266676","exception":false,"start_time":"2023-05-25T20:23:24.238164","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.904108Z","iopub.execute_input":"2023-06-05T19:02:15.906726Z","iopub.status.idle":"2023-06-05T19:02:15.917409Z","shell.execute_reply.started":"2023-06-05T19:02:15.906694Z","shell.execute_reply":"2023-06-05T19:02:15.916498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Scheduler","metadata":{"papermill":{"duration":0.022229,"end_time":"2023-05-25T20:23:24.400670","exception":false,"start_time":"2023-05-25T20:23:24.378441","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class 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()\n","metadata":{"papermill":{"duration":0.034727,"end_time":"2023-05-25T20:23:24.455285","exception":false,"start_time":"2023-05-25T20:23:24.420558","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.922107Z","iopub.execute_input":"2023-06-05T19:02:15.925384Z","iopub.status.idle":"2023-06-05T19:02:15.956232Z","shell.execute_reply.started":"2023-06-05T19:02:15.925351Z","shell.execute_reply":"2023-06-05T19:02:15.955108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Initializing the Model, optimizer and the scheduler","metadata":{"papermill":{"duration":0.018998,"end_time":"2023-05-25T20:23:24.493630","exception":false,"start_time":"2023-05-25T20:23:24.474632","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# if TEST_RUN:\n#     model = build_model(CFG)\n#     model.to(device)\n#     torch.cuda.empty_cache()\n#     gc.collect()\n#     print(\"MODEL IS WORKING\")","metadata":{"papermill":{"duration":0.027221,"end_time":"2023-05-25T20:23:24.540167","exception":false,"start_time":"2023-05-25T20:23:24.512946","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.957547Z","iopub.execute_input":"2023-06-05T19:02:15.958040Z","iopub.status.idle":"2023-06-05T19:02:15.964543Z","shell.execute_reply.started":"2023-06-05T19:02:15.958008Z","shell.execute_reply":"2023-06-05T19:02:15.962966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Losses","metadata":{"papermill":{"duration":0.018917,"end_time":"2023-05-25T20:23:24.578255","exception":false,"start_time":"2023-05-25T20:23:24.559338","status":"completed"},"tags":[]}},{"cell_type":"code","source":"alpha = 0.5\nbeta = 1 - alpha\n\nDiceLoss = smp.losses.DiceLoss(mode='binary')\nBCELoss = smp.losses.SoftBCEWithLogitsLoss()\n","metadata":{"papermill":{"duration":0.028449,"end_time":"2023-05-25T20:23:24.625869","exception":false,"start_time":"2023-05-25T20:23:24.597420","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.965471Z","iopub.execute_input":"2023-06-05T19:02:15.965962Z","iopub.status.idle":"2023-06-05T19:02:15.974247Z","shell.execute_reply.started":"2023-06-05T19:02:15.965924Z","shell.execute_reply":"2023-06-05T19:02:15.973109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def criterion(y_pred, y_true):\n    # return 0.5 * BCELoss(y_pred, y_true) + 0.5 * DiceLoss(y_pred, y_true)\n    # return 0.5 * BCELoss(y_pred, y_true) + 0.5 * TverskyLoss(y_pred, y_true)\n    if CFG.loss_func=='binary_cross_entropy':\n        return BCELoss(y_pred, y_true)\n    \n    elif CFG.loss_func=='diceloss':\n        return DiceLoss(y_pred, y_true)\n    elif CFG.loss_func=='Mixed':\n        return BCELoss(y_pred, y_true)+DiceLoss(y_pred, y_true)\n\n    elif CFG.loss_func=='LovaszLoss':\n        LovaszLoss=smp.losses.LovaszLoss(mode='binary')\n        return LovaszLoss(y_pred, y_true)\n    elif CFG.loss_func=='JaccardLoss':\n        JaccardLoss=smp.losses.JaccardLoss(mode='binary')\n        return JaccardLoss(y_pred, y_true)\n    else:\n        return BCELoss(y_pred, y_true)\n","metadata":{"papermill":{"duration":0.029032,"end_time":"2023-05-25T20:23:24.674484","exception":false,"start_time":"2023-05-25T20:23:24.645452","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.975659Z","iopub.execute_input":"2023-06-05T19:02:15.976059Z","iopub.status.idle":"2023-06-05T19:02:15.988093Z","shell.execute_reply.started":"2023-06-05T19:02:15.976029Z","shell.execute_reply":"2023-06-05T19:02:15.987078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training and Validation Function","metadata":{"papermill":{"duration":0.019053,"end_time":"2023-05-25T20:23:24.712691","exception":false,"start_time":"2023-05-25T20:23:24.693638","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def train_fn(train_loader, model, criterion, optimizer, device):\n    model.train()\n    accumulation_steps=CFG.accumulation_steps\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        # Accumulate gradients\n        scaler.scale(loss / accumulation_steps).backward()\n\n        if (step + 1) % accumulation_steps == 0:\n            # Update gradients every accumulation_steps iterations\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        losses.update(loss.detach().item(), batch_size)\n\n    # Handle leftover gradients if accumulation_steps doesn't divide len(train_loader) evenly\n    if (step + 1) % accumulation_steps != 0:\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\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","metadata":{"papermill":{"duration":0.041682,"end_time":"2023-05-25T20:23:24.773907","exception":false,"start_time":"2023-05-25T20:23:24.732225","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:15.989386Z","iopub.execute_input":"2023-06-05T19:02:15.990416Z","iopub.status.idle":"2023-06-05T19:02:16.025905Z","shell.execute_reply.started":"2023-06-05T19:02:15.990381Z","shell.execute_reply":"2023-06-05T19:02:16.024949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Metrics for the competition","metadata":{"papermill":{"duration":0.018925,"end_time":"2023-05-25T20:23:24.812246","exception":false,"start_time":"2023-05-25T20:23:24.793321","status":"completed"},"tags":[]}},{"cell_type":"code","source":"from 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):\n    mask = mask.astype(int).flatten()\n    mask_pred = mask_pred.flatten()\n\n    best_th = 0\n    best_dice = 0\n    \n    for th in np.array(range(5, 100+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        \n        print(f'th: {th}, fbeta: {dice}')\n        \n        torch.cuda.empty_cache()\n        gc.collect()\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","metadata":{"papermill":{"duration":0.034821,"end_time":"2023-05-25T20:23:24.866555","exception":false,"start_time":"2023-05-25T20:23:24.831734","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:16.027033Z","iopub.execute_input":"2023-06-05T19:02:16.027613Z","iopub.status.idle":"2023-06-05T19:02:16.040091Z","shell.execute_reply.started":"2023-06-05T19:02:16.027579Z","shell.execute_reply":"2023-06-05T19:02:16.039218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Main","metadata":{"papermill":{"duration":0.019174,"end_time":"2023-05-25T20:23:24.905159","exception":false,"start_time":"2023-05-25T20:23:24.885985","status":"completed"},"tags":[]}},{"cell_type":"code","source":"for fold in [1,2,3,4,5, 6]:\n    \n    train_images, train_masks, valid_images, valid_masks, valid_xyxys = get_train_valid_dataset(fold)\n    valid_xyxys = np.stack(valid_xyxys)\n    \n    \n    train_dataset = CustomDataset(train_images, CFG, labels=train_masks,\n                                  transform=get_transforms(data='train', cfg=CFG))\n\n    valid_dataset = CustomDataset(valid_images, CFG, labels=valid_masks,\n                                  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\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    \n    valid_mask_gt = cv2.imread(CFG.comp_dataset_path + f\"train/{fold}/inklabels.png\", 0)\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    \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    \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    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        torch.cuda.empty_cache()\n        gc.collect()\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#         scheduler.step()\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        current_lr=scheduler.get_lr()\n        gpu_mem_data = f\"Mem : {torch.cuda.memory_reserved() / 1E9:.3g}GB\"\n        gpu_mem=torch.cuda.memory_reserved() / 1E9\n\n        data_to_log={\"Epoch \":epoch+1, \n                     \"avg train loss\":avg_loss, \n                     \"avg val loss\":avg_val_loss, \n                     \"SCORE \": score,\n                     \"best_thresh \": best_th,\n                     \"learning rate \":current_lr,\n                     \"gpu memory \":gpu_mem}\n            \n        wandb.log(data_to_log)\n        print(data_to_log)\n        if update_best:\n            print(\"SAVING THE BEST MODEL\")\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            path_name=f\"unet all aug with mitb2_{CFG.loss_func}_{fold}_best.pt\"\n\n            torch.save(model, path_name)\n#             torch.save({\"valid_mask_gt\":valid_mask_gt, \"mask_pred\":mask_pred},f\"predictions_{fold}.pth\")\n\n\n        torch.cuda.empty_cache()\n        gc.collect()\n    del model, train_loader\n    torch.cuda.empty_cache()\n    gc.collect()\n        \nwandb.finish()\n","metadata":{"papermill":{"duration":14872.3831,"end_time":"2023-05-26T00:31:17.307484","exception":false,"start_time":"2023-05-25T20:23:24.924384","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-05T19:02:16.041734Z","iopub.execute_input":"2023-06-05T19:02:16.042445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"papermill":{"duration":0.203323,"end_time":"2023-05-26T00:31:17.760465","exception":false,"start_time":"2023-05-26T00:31:17.557142","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}