{"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 warmup_scheduler segmentation_models_pytorch -q","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:07.468049Z","iopub.execute_input":"2023-07-14T17:23:07.469121Z","iopub.status.idle":"2023-07-14T17:23:29.133993Z","shell.execute_reply.started":"2023-07-14T17:23:07.469087Z","shell.execute_reply":"2023-07-14T17:23:29.132664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nimport gc\nimport os\nimport sys\nimport warnings\nfrom glob import glob\nimport matplotlib.pyplot as plt\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom warmup_scheduler import GradualWarmupScheduler\nimport pandas as pd\nimport cv2\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nimport segmentation_models_pytorch as smp\nfrom torch.cuda import amp\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:29.137006Z","iopub.execute_input":"2023-07-14T17:23:29.137862Z","iopub.status.idle":"2023-07-14T17:23:36.53127Z","shell.execute_reply.started":"2023-07-14T17:23:29.137829Z","shell.execute_reply":"2023-07-14T17:23:36.530322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    debug = False\n    # ============== comp exp name =============\n    comp_name = 'contrail'\n    comp_dir_path = '/kaggle/input/'\n    comp_folder_name = 'google-research-identify-contrails-reduce-global-warming'\n\n    dataset_path = \"/kaggle/input/ashcolor-4labels/dataset_train/ash_color_4labels/\"\n    train_label_path = f\"{dataset_path}/labels/true/\"\n\n    exp_name = \"model\"\n\n    # ============== model cfg =============\n    model_arch = 'Unet'\n    backbone = 'timm-resnest50d'\n    \n    in_chans = 3\n    target_size = 4\n\n    # ============== training cfg =============\n    train_batch_size = 32\n    valid_batch_size = train_batch_size\n\n    epochs = 35\n\n    lr = 1e-4\n    loss = \"DiceLoss\"\n    \n    # ============== fixed =============\n    num_workers = 4\n    seed = 42\n\n    # ============== augmentation =============\n    train_aug_list = [\n        ToTensorV2(transpose_mask=True),\n    ]\n\n    valid_aug_list = [\n        ToTensorV2(transpose_mask=True),\n    ]\n","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:36.53269Z","iopub.execute_input":"2023-07-14T17:23:36.533125Z","iopub.status.idle":"2023-07-14T17:23:36.543619Z","shell.execute_reply.started":"2023-07-14T17:23:36.533089Z","shell.execute_reply":"2023-07-14T17:23:36.541972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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\nwarnings.filterwarnings(\"ignore\")\ntorch.backends.cudnn.benchmark = True\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nset_seed(CFG.seed)\nos.makedirs(f'./{CFG.exp_name}/', exist_ok=True)\npd.options.display.max_colwidth = 300","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:36.546777Z","iopub.execute_input":"2023-07-14T17:23:36.547204Z","iopub.status.idle":"2023-07-14T17:23:36.587664Z","shell.execute_reply.started":"2023-07-14T17:23:36.547174Z","shell.execute_reply":"2023-07-14T17:23:36.586542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(f\"{CFG.dataset_path}/train_df.csv\")\nvalid_df = pd.read_csv(f\"{CFG.dataset_path}/validation_df.csv\")\ntrain_df[\"image_path\"]=train_df[\"image_path\"].str.replace(\"/kaggle/working/dataset_train/ash_color_4labels/\", CFG.dataset_path)\ntrain_df[\"label_path\"]=train_df[\"label_path\"].str.replace(\"/kaggle/working/dataset_train/ash_color_4labels/\", CFG.dataset_path)\nvalid_df[\"image_path\"]=valid_df[\"image_path\"].str.replace(\"/kaggle/working/dataset_train/ash_color_4labels/\", CFG.dataset_path)\nvalid_df[\"label_path\"]=valid_df[\"label_path\"].str.replace(\"/kaggle/working/dataset_train/ash_color_4labels/\", CFG.dataset_path)\nif CFG.debug:\n    train_df=train_df[:2000]\n    valid_df=valid_df[:2000]\nprint(train_df.shape, valid_df.shape)\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:36.590081Z","iopub.execute_input":"2023-07-14T17:23:36.590929Z","iopub.status.idle":"2023-07-14T17:23:36.770154Z","shell.execute_reply.started":"2023-07-14T17:23:36.590897Z","shell.execute_reply":"2023-07-14T17:23:36.769138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrailsDataset(Dataset):\n    def __init__(self, df, transform, mode='train'):\n        self.df = df\n        self.transform = A.Compose(transform)\n        self.mode = mode\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        record_id = row[\"record_id\"]\n\n        if self.mode == 'train':\n            image_path = row[\"image_path\"]\n            label_path = row[\"label_path\"]\n            image = np.load(str(image_path)).astype(\"float32\")\n            label = np.load(str(label_path)).astype(\"float32\")\n            data = self.transform(image=image, mask=label)\n            image = data['image']\n            label = data['mask']\n            image = torch.tensor(image)\n            return image.float(), label.float()\n\n        if self.mode == 'valid':\n            image_path = row[\"image_path\"]\n            label_path = row[\"label_path\"]\n            image = np.load(str(image_path)).astype(\"float32\")\n            label = np.load(str(label_path)).astype(\"float32\")\n            data = self.transform(image=image, mask=label)\n            image = data['image']\n            label = data['mask']\n            image = torch.tensor(image)\n            return image.float(), label.float()\n\n    def __len__(self):\n        return len(self.df)\n\n\ndef show_dataset(idx, dataset):\n    image, label = dataset[idx]\n    label_num = label.shape[0]\n    fig, axes = plt.subplots(nrows=1, ncols=1+label_num, figsize=(20, 4))\n    axes = axes.flatten()\n    fig.tight_layout(pad=0.1)\n    axes[0].imshow(image.permute(1, 2, 0).to(torch.float))\n    axes[0].axis('off')\n    for i in range(label_num):\n        axes[i+1].imshow(label[i].to(torch.float))\n        axes[i+1].axis('off')","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:36.772261Z","iopub.execute_input":"2023-07-14T17:23:36.77293Z","iopub.status.idle":"2023-07-14T17:23:36.785958Z","shell.execute_reply.started":"2023-07-14T17:23:36.772897Z","shell.execute_reply":"2023-07-14T17:23:36.784812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_train = ContrailsDataset(train_df, CFG.train_aug_list, \"train\")\ndataset_valid = ContrailsDataset(valid_df, CFG.valid_aug_list, \"valid\")\n\ndataloader_train = DataLoader(dataset_train, batch_size=CFG.train_batch_size ,shuffle=True, num_workers = CFG.num_workers)\ndataloader_valid = DataLoader(dataset_valid, batch_size=CFG.valid_batch_size, num_workers = CFG.num_workers)\n\nprint(f\"\"\"\n{len(dataset_train) = }\ntrain_image_shape : {dataset_train[0][0].shape}\ntrain_mask_shape  : {dataset_train[0][1].shape}\ntrain_image_dtype : {dataset_train[0][0].dtype}\ntrain_mask_dtype : {dataset_train[0][1].dtype}\n\n{len(dataset_valid) = }\nvalid_image_shape : {dataset_valid[0][0].shape}\nvalid_mask_shape  : {dataset_valid[0][1].shape}\nvalid_image_dtype : {dataset_valid[0][0].dtype}\nvalid_mask_dtype : {dataset_valid[0][1].dtype}\n\"\"\")","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:36.787408Z","iopub.execute_input":"2023-07-14T17:23:36.788014Z","iopub.status.idle":"2023-07-14T17:23:36.905516Z","shell.execute_reply.started":"2023-07-14T17:23:36.78798Z","shell.execute_reply":"2023-07-14T17:23:36.904491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_dataset(12, dataset_train)","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:36.907388Z","iopub.execute_input":"2023-07-14T17:23:36.908084Z","iopub.status.idle":"2023-07-14T17:23:37.655507Z","shell.execute_reply.started":"2023-07-14T17:23:36.90805Z","shell.execute_reply":"2023-07-14T17:23:37.654587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_dataset(76, dataset_valid)","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:37.656558Z","iopub.execute_input":"2023-07-14T17:23:37.656893Z","iopub.status.idle":"2023-07-14T17:23:38.086027Z","shell.execute_reply.started":"2023-07-14T17:23:37.656866Z","shell.execute_reply":"2023-07-14T17:23:38.085164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    def __init__(self, model_arch, backbone, in_chans, target_size, weight):\n        super().__init__()\n\n        self.model = smp.create_model(\n            model_arch,\n            encoder_name=backbone,\n            encoder_weights=weight,\n            in_channels=in_chans,\n            classes=target_size,\n            activation=None,\n        )\n\n    def forward(self, image):\n        output = self.model(image)\n        return output\n\n\ndef build_model(model_arch, backbone, in_chans, target_size, weight=\"imagenet\"):\n    print('model_arch: ', model_arch)\n    print('backbone: ', backbone)\n    model = CustomModel(model_arch, backbone, in_chans, target_size, weight)\n    return model\n\n\ndef load_model(pth_path):\n    pth = torch.load(f'{pth_path}')\n\n    model = build_model(pth[\"model_arch\"], pth[\"backbone\"], pth[\"in_chans\"], pth[\"target_size\"], weight=None, dataparallel=False)\n    model.load_state_dict(pth['model'])\n    thresh = pth['thresh']\n    dice_score = pth['dice_score']\n\n    return model, dice_score, thresh","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:38.089331Z","iopub.execute_input":"2023-07-14T17:23:38.090275Z","iopub.status.idle":"2023-07-14T17:23:38.100776Z","shell.execute_reply.started":"2023-07-14T17:23:38.090229Z","shell.execute_reply":"2023-07-14T17:23:38.099661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_gpus = torch.cuda.device_count()\ndevice_ids = list(range(num_gpus))\n\nmodel = build_model(CFG.model_arch, CFG.backbone, CFG.in_chans, CFG.target_size)\nmodel = nn.DataParallel(model, device_ids=device_ids)\nmodel.to(device);","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:38.102375Z","iopub.execute_input":"2023-07-14T17:23:38.102793Z","iopub.status.idle":"2023-07-14T17:23:52.733345Z","shell.execute_reply.started":"2023-07-14T17:23:38.102762Z","shell.execute_reply":"2023-07-14T17:23:52.732292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GradualWarmupSchedulerV2(GradualWarmupScheduler):\n    \"\"\"\n    https://www.kaggle.com/code/underwearfitting/single-fold-training-of-resnet200d-lb0-965\n    \"\"\"\n\n    def __init__(self, optimizer, multiplier, total_epoch, after_scheduler=None):\n        super(GradualWarmupSchedulerV2, self).__init__(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 = [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\n\ndef get_scheduler(epochs, optimizer, multiplier=10, total_epoch=1, eta_min=1e-7):\n    scheduler_cosine = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, epochs, eta_min=eta_min)\n    scheduler = GradualWarmupSchedulerV2(optimizer, multiplier=multiplier, total_epoch=total_epoch, after_scheduler=scheduler_cosine)\n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:52.734671Z","iopub.execute_input":"2023-07-14T17:23:52.735831Z","iopub.status.idle":"2023-07-14T17:23:52.746639Z","shell.execute_reply.started":"2023-07-14T17:23:52.735796Z","shell.execute_reply":"2023-07-14T17:23:52.745443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"scaler = amp.GradScaler()\ncriterion = smp.losses.DiceLoss(mode=\"multilabel\")\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr)\nscheduler = get_scheduler(CFG.epochs, optimizer)\n\nthresholds_to_test = [round(x * 0.01, 2) for x in range(1, 101, 2)]","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:52.749968Z","iopub.execute_input":"2023-07-14T17:23:52.750241Z","iopub.status.idle":"2023-07-14T17:23:52.765932Z","shell.execute_reply.started":"2023-07-14T17:23:52.750218Z","shell.execute_reply":"2023-07-14T17:23:52.764943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"class Dice(nn.Module):\n    def __init__(self, use_sigmoid=True):\n        super(Dice, self).__init__()\n        self.sigmoid = nn.Sigmoid()\n        self.use_sigmoid = use_sigmoid\n\n    def forward(self, inputs, targets, smooth=1):\n        if self.use_sigmoid:\n            inputs = self.sigmoid(inputs)\n\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n\n        intersection = (inputs * targets).sum()\n        dice = (2.0 * intersection + smooth)/(inputs.sum() + targets.sum() + smooth)\n\n        return dice\n\n\ndef calc_dice_score(pred, true, thresh: float) -> float:\n    dice = Dice(use_sigmoid=False)\n    pred_thresh = np.where(pred > thresh, 1, 0)\n    pred_thresh = torch.flatten(torch.from_numpy(pred_thresh))\n    return dice(true, pred_thresh).item()\n\n\ndef calc_optim_thresh(pred, true, threshs_to_test):\n    best_dice = -1\n    for thresh in threshs_to_test:\n        dice = calc_dice_score(pred, true, thresh)\n        if dice > best_dice:\n            best_dice = dice\n            best_thresh = thresh\n    return best_dice, best_thresh","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:52.767136Z","iopub.execute_input":"2023-07-14T17:23:52.768192Z","iopub.status.idle":"2023-07-14T17:23:52.779008Z","shell.execute_reply.started":"2023-07-14T17:23:52.768156Z","shell.execute_reply":"2023-07-14T17:23:52.777292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.debug:\n    print(\"!!!Debug mode!!!\\n\")\n\ndice_score=0\nfor epoch in range(CFG.epochs):\n    model.train()\n    \n    pbar_train = enumerate(dataloader_train)\n    pbar_train = tqdm(pbar_train, total=len(dataloader_train), bar_format=\"{l_bar}{bar:10}{r_bar}{bar:-0b}\")\n    loss_train, loss_val= 0.0, 0.0\n    for i, (images, masks) in pbar_train:\n        images, masks = images.cuda(), masks.cuda()\n        optimizer.zero_grad()\n        with amp.autocast():\n            preds = model(images)\n            loss = criterion(preds, masks)\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            loss_train += loss.detach().item()\n        \n        lr = f\"LR : {scheduler.get_lr()[0]:.2E}\"\n        gpu_mem = f\"Mem : {torch.cuda.memory_reserved() / 1E9:.3g}GB\"\n        pbar_train.set_description((\"%10s  \" * 3 + \"%10s\") % (f\"Epoch {epoch}/{CFG.epochs}\", gpu_mem, lr,\n                                                                f\"Loss: {loss_train / (i + 1):.4f}\"))\n\n    scheduler.step()\n    model.eval()\n    \n    cum_pred = []\n    cum_true = []\n    pbar_val = enumerate(dataloader_valid)\n    pbar_val = tqdm(pbar_val, total=len(dataloader_valid), bar_format=\"{l_bar}{bar:10}{r_bar}{bar:-10b}\")\n    for i, (images, masks) in pbar_val:\n        images, masks = images.cuda(), masks.cuda()\n        with torch.no_grad():\n            preds = model(images)[:,2]\n            loss_val += criterion(preds, masks).item()\n            preds = torch.sigmoid(preds)\n            cum_pred.append(preds.cpu().detach().numpy())\n            cum_true.append(masks.cpu().detach().numpy())\n\n        pbar_val.set_description((\"%10s\") % (f\"Val Loss: {loss_val / (i+1):.4f}\"))\n    \n    cum_pred = torch.flatten(torch.from_numpy(np.concatenate(cum_pred, axis=0)))\n    cum_true = torch.flatten(torch.from_numpy(np.concatenate(cum_true, axis=0)))\n    \n    dice_score_, thresh = calc_optim_thresh(cum_pred, cum_true, thresholds_to_test)\n    \n    if dice_score_ > dice_score:\n        print(f\"score : {dice_score_:.4f}\\tthresh : {thresh}\\tSAVED MODEL\\n\")\n        epoch_best=epoch\n        dice_score =dice_score_\n        torch.save({'model': model.module.state_dict(), 'dice_score': dice_score, 'thresh': thresh,\n                    \"model_arch\":CFG.model_arch, \"backbone\":CFG.backbone,\"in_chans\":CFG.in_chans,\"target_size\":CFG.target_size,},\n                    f'./{CFG.exp_name}/{CFG.exp_name}.pth')\n    else:\n        print(f\"score : {dice_score_:.4f}\\tthresh : {thresh}\\n\")\n    \n","metadata":{"execution":{"iopub.status.busy":"2023-07-14T17:23:52.780841Z","iopub.execute_input":"2023-07-14T17:23:52.781566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}