{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# IC2RGW:PyTorch Baseline Train&Inference","metadata":{}},{"cell_type":"code","source":"TRAIN_MODE = True","metadata":{"execution":{"iopub.status.busy":"2023-06-24T07:25:14.308122Z","iopub.execute_input":"2023-06-24T07:25:14.308678Z","iopub.status.idle":"2023-06-24T07:25:14.320974Z","shell.execute_reply.started":"2023-06-24T07:25:14.308646Z","shell.execute_reply":"2023-06-24T07:25:14.320192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Libraries","metadata":{}},{"cell_type":"code","source":"import os\nimport copy\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm\n\nimport torch\nimport torch.nn as nn\nfrom timm.scheduler import CosineLRScheduler","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":3.421775,"end_time":"2023-05-15T00:35:16.571754","exception":false,"start_time":"2023-05-15T00:35:13.149979","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-24T07:25:14.350193Z","iopub.execute_input":"2023-06-24T07:25:14.350747Z","iopub.status.idle":"2023-06-24T07:25:20.314656Z","shell.execute_reply.started":"2023-06-24T07:25:14.350717Z","shell.execute_reply":"2023-06-24T07:25:20.313711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utilities","metadata":{}},{"cell_type":"code","source":"def get_device():\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    print(f\"Using {device} device\")\n    return device\n\ndevice = get_device()","metadata":{"papermill":{"duration":0.064303,"end_time":"2023-05-15T00:35:16.6412","exception":false,"start_time":"2023-05-15T00:35:16.576897","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-24T07:25:20.316531Z","iopub.execute_input":"2023-06-24T07:25:20.318019Z","iopub.status.idle":"2023-06-24T07:25:20.350913Z","shell.execute_reply.started":"2023-06-24T07:25:20.317985Z","shell.execute_reply":"2023-06-24T07:25:20.349834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    # Data\n    base_dir = \"../input/google-research-identify-contrails-reduce-global-warming\"\n    train_path = os.path.join(base_dir,\"train\")\n    val_path = os.path.join(base_dir,\"validation\")\n    \n    # Train\n    num_epochs = 10\n    batch_size = 48\n    num_workers = 2\n    \n    # Optimizer & Scheduler\n    lr_max = 3e-4\n    epochs_warmup = 5\n    warmup_lr_init = 5e-4\n    lr_min = 1e-6\n    scheduler_name = \"CosineAnnealingLR\"\n    \n    # threshold\n    threshold = 0.4","metadata":{"papermill":{"duration":0.012225,"end_time":"2023-05-15T00:35:16.658379","exception":false,"start_time":"2023-05-15T00:35:16.646154","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-24T07:25:20.354711Z","iopub.execute_input":"2023-06-24T07:25:20.355039Z","iopub.status.idle":"2023-06-24T07:25:20.363556Z","shell.execute_reply.started":"2023-06-24T07:25:20.355004Z","shell.execute_reply":"2023-06-24T07:25:20.36241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"markdown","source":"### Reference\n[Visualizing Contrails](https://www.kaggle.com/code/inversion/visualizing-contrails)","metadata":{}},{"cell_type":"code","source":"_T11_BOUNDS = (243, 303)\n_CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n_TDIFF_BOUNDS = (-4, 2)\n\ndef normalize_range(data, bounds):\n    \"\"\"Maps data to the range [0, 1].\"\"\"\n    return (data - bounds[0]) / (bounds[1] - bounds[0])\n\ndef normalize_std(spec):\n    return (spec- np.mean(spec))/np.std(spec)\n\nclass Dataset(torch.utils.data.Dataset):\n    def __init__(self, data_path, mode='train'):\n        self.data_path = data_path\n        self.file_name = os.listdir(data_path)\n        self.mode = mode\n        \n\n    def __len__(self):\n        return len(self.file_name)\n\n    def __getitem__(self, i):\n        \n        band11 = np.load(os.path.join(self.data_path, self.file_name[i], 'band_11.npy'))\n        band14 = np.load(os.path.join(self.data_path, self.file_name[i], 'band_14.npy'))\n        band15 = np.load(os.path.join(self.data_path, self.file_name[i], 'band_15.npy'))\n        \n        r = normalize_range(band15 - band14, _TDIFF_BOUNDS)\n        g = normalize_range(band14 - band11, _CLOUD_TOP_TDIFF_BOUNDS)\n        b = normalize_range(band14, _T11_BOUNDS)\n        x = np.transpose(np.clip(np.stack([r, g, b], axis=2), 0, 1)[:,:,:,4],(2,0,1))\n        x = normalize_std(x)\n        \n        if self.mode == 'train':\n            y = np.load(os.path.join(self.data_path, self.file_name[i], 'human_pixel_masks.npy')).astype(np.float32).transpose(2,0,1)\n        elif self.mode == 'test':\n            y = self.file_name[i]\n        \n        return x, y","metadata":{"papermill":{"duration":0.017522,"end_time":"2023-05-15T00:35:16.699306","exception":false,"start_time":"2023-05-15T00:35:16.681784","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-24T07:25:20.369963Z","iopub.execute_input":"2023-06-24T07:25:20.370392Z","iopub.status.idle":"2023-06-24T07:25:20.382799Z","shell.execute_reply.started":"2023-06-24T07:25:20.370368Z","shell.execute_reply":"2023-06-24T07:25:20.38183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = Dataset(CFG.train_path)\na, b = dataset[3]\n\nplt.figure(figsize=(12, 6))\nax = plt.subplot(1, 2, 1)\nax.imshow(np.transpose(a,(1,2,0)))\nax = plt.subplot(1, 2, 2)\nax.imshow(np.transpose(b,(1,2,0)), interpolation='none') \nplt.show()","metadata":{"papermill":{"duration":0.911202,"end_time":"2023-05-15T00:35:17.615302","exception":false,"start_time":"2023-05-15T00:35:16.7041","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-24T07:25:20.384114Z","iopub.execute_input":"2023-06-24T07:25:20.384684Z","iopub.status.idle":"2023-06-24T07:25:21.255855Z","shell.execute_reply.started":"2023-06-24T07:25:20.384653Z","shell.execute_reply":"2023-06-24T07:25:21.255019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"markdown","source":"### Reference\n[[GR-ICRGW] Pytorch Lightning baseline UNet+resnest](https://www.kaggle.com/code/egortrushin/gr-icrgw-pytorch-lightning-baseline-unet-resnest)","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append(\"../input/pretrained-models-pytorch\")\nsys.path.append(\"../input/efficientnet-pytorch\")\nsys.path.append(\"/kaggle/input/smp-github/segmentation_models.pytorch-master\")\nsys.path.append(\"/kaggle/input/timm-pretrained-resnest/resnest/\")\nimport segmentation_models_pytorch as smp","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-06-24T07:25:21.256944Z","iopub.execute_input":"2023-06-24T07:25:21.257292Z","iopub.status.idle":"2023-06-24T07:25:22.957842Z","shell.execute_reply.started":"2023-06-24T07:25:21.257261Z","shell.execute_reply":"2023-06-24T07:25:22.956624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints/\n!cp /kaggle/input/timm-pretrained-resnest/resnest/gluon_resnest26-50eb607c.pth /root/.cache/torch/hub/checkpoints/gluon_resnest26-50eb607c.pth","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-06-24T07:25:22.959297Z","iopub.execute_input":"2023-06-24T07:25:22.959643Z","iopub.status.idle":"2023-06-24T07:25:25.901646Z","shell.execute_reply.started":"2023-06-24T07:25:22.959615Z","shell.execute_reply":"2023-06-24T07:25:25.900306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAIN_MODE == True:\n    model = smp.Unet(encoder_name='timm-resnest26d',\n                encoder_weights=\"imagenet\",\n                in_channels=3,\n                classes=1,\n                activation=None,\n            ).to(device);\nelse:\n    best_model = torch.load('/kaggle/input/ic2rgw-pytorch-samplemodel/unet-smp-sample.pth').to(device);\n    \nsigmoid = nn.Sigmoid()","metadata":{"papermill":{"duration":6.104629,"end_time":"2023-05-15T00:35:23.727688","exception":false,"start_time":"2023-05-15T00:35:17.623059","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-24T07:25:25.904004Z","iopub.execute_input":"2023-06-24T07:25:25.904755Z","iopub.status.idle":"2023-06-24T07:25:29.324766Z","shell.execute_reply.started":"2023-06-24T07:25:25.904716Z","shell.execute_reply":"2023-06-24T07:25:29.32383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metric","metadata":{}},{"cell_type":"code","source":"class Dice(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(Dice, self).__init__()\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, inputs, targets, smooth=1):\n        \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.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n        \n        return dice\n    \ndice = Dice()","metadata":{"papermill":{"duration":0.018966,"end_time":"2023-05-15T00:35:23.758237","exception":false,"start_time":"2023-05-15T00:35:23.739271","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-24T07:25:29.326071Z","iopub.execute_input":"2023-06-24T07:25:29.326412Z","iopub.status.idle":"2023-06-24T07:25:29.337955Z","shell.execute_reply.started":"2023-06-24T07:25:29.326381Z","shell.execute_reply":"2023-06-24T07:25:29.336827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"if TRAIN_MODE == True:\n    num_epochs = CFG.num_epochs\n    batch_size = CFG.batch_size\n    num_workers = CFG.num_workers\n\n    train_dataset = Dataset(CFG.train_path, mode='train')\n    val_dataset = Dataset(CFG.val_path, mode='train')\n\n    train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers, drop_last=True, pin_memory=True)\n    val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers, drop_last=True, pin_memory=True)\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr_max)\n    criterion = smp.losses.DiceLoss(mode=\"binary\", smooth=1.0)\n    nbatch = len(train_loader)\n    warmup = CFG.epochs_warmup * nbatch\n    nsteps = CFG.num_epochs * nbatch\n    scheduler = CosineLRScheduler(optimizer,\n                          warmup_t=warmup, warmup_lr_init=CFG.warmup_lr_init, warmup_prefix=True,\n                          t_initial=(nsteps - warmup), lr_min=CFG.lr_min) \n    \n    val_best_dice = 0.0\n    i_scheduler = 0\n    \n    for epoch in range(num_epochs):\n        train_loss, val_loss = 0, 0\n        train_dice, val_dice = 0, 0\n        n_train, n_val = 0, 0\n\n        model.train()\n        for i_train, (X, y) in enumerate(tqdm(train_loader)):   \n            n_train += len(y)\n            X = X.to(device)\n            y = y.to(device)\n\n            optimizer.zero_grad()\n            pred = model(X)\n            loss = criterion(pred, y)\n            dice_temp = dice(pred, y)\n\n\n            loss.backward()\n\n            optimizer.step()\n            train_loss += loss.item()\n            train_dice += dice_temp.item()\n            \n            scheduler.step(i_scheduler)\n            i_scheduler +=1\n\n        model.eval()\n        with torch.no_grad():\n            for i_val, (X, y) in enumerate(val_loader):   \n                n_val += len(y)\n                X = X.to(device)\n                y = y.to(device)\n\n                pred = model(X)\n                loss = criterion(pred, y)\n                dice_temp = dice(pred, y)\n\n                val_loss += loss.item()\n                val_dice += dice_temp.item()\n        \n        print (f'Epoch [{(epoch+1)}/{num_epochs}], loss: {train_loss/n_train:.5f}, dice: {(train_dice+1e-23)/i_train:.5f}, val_loss: {val_loss/n_val:.5f}, val_dice: {(val_dice+1e-23)/i_val:.5f}')\n        print(optimizer.param_groups[0][\"lr\"])\n        if val_dice > val_best_dice:\n            val_best_dice = val_dice\n            torch.save(model, 'unet-smp-sample.pth')\n            best_model = copy.deepcopy(model)\n    ","metadata":{"papermill":{"duration":22180.769991,"end_time":"2023-05-15T06:45:04.728989","exception":false,"start_time":"2023-05-15T00:35:23.958998","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-24T07:26:37.252231Z","iopub.execute_input":"2023-06-24T07:26:37.252611Z","iopub.status.idle":"2023-06-24T07:27:51.582397Z","shell.execute_reply.started":"2023-06-24T07:26:37.252582Z","shell.execute_reply":"2023-06-24T07:27:51.579899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict and submit","metadata":{}},{"cell_type":"markdown","source":"### Reference\n[Contrails - RLE Submission](https://www.kaggle.com/code/inversion/contrails-rle-submission)","metadata":{}},{"cell_type":"code","source":"def rle_encode(x, fg_val=1):\n    \"\"\"\n    Args:\n        x:  numpy array of shape (height, width), 1 - mask, 0 - background\n    Returns: run length encoding as list\n    \"\"\"\n\n    dots = np.where(\n        x.T.flatten() == fg_val)[0]  # .T sets Fortran order down-then-right\n    run_lengths = []\n    prev = -2\n    for b in dots:\n        if b > prev + 1:\n            run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n    return run_lengths\n\n\ndef list_to_string(x):\n    \"\"\"\n    Converts list to a string representation\n    Empty list returns '-'\n    \"\"\"\n    if x: # non-empty list\n        s = str(x).replace(\"[\", \"\").replace(\"]\", \"\").replace(\",\", \"\")\n    else:\n        s = '-'\n    return s\n\n\ndef rle_decode(mask_rle, shape=(256, 256)):\n    '''\n    mask_rle: run-length as string formatted (start length)\n              empty predictions need to be encoded with '-'\n    shape: (height, width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n    '''\n\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    if mask_rle != '-': \n        s = mask_rle.split()\n        starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n        starts -= 1\n        ends = starts + lengths\n        for lo, hi in zip(starts, ends):\n            img[lo:hi] = 1\n    return img.reshape(shape, order='F')  # Needed to align to RLE direction","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-06-24T07:25:30.001389Z","iopub.status.idle":"2023-06-24T07:25:30.002111Z","shell.execute_reply.started":"2023-06-24T07:25:30.001856Z","shell.execute_reply":"2023-06-24T07:25:30.001879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_recs = os.listdir(os.path.join(CFG.base_dir,\"test\"))\nprint(test_recs)","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-24T07:25:30.00344Z","iopub.status.idle":"2023-06-24T07:25:30.004157Z","shell.execute_reply.started":"2023-06-24T07:25:30.003898Z","shell.execute_reply":"2023-06-24T07:25:30.00392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 32\nnum_workers = 2\n\ntest_dataset = Dataset('/kaggle/input/google-research-identify-contrails-reduce-global-warming/test', mode='test')\n\ntest_loader = torch.utils.data.DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers, drop_last=False)","metadata":{"execution":{"iopub.status.busy":"2023-06-24T07:25:30.005389Z","iopub.status.idle":"2023-06-24T07:25:30.006097Z","shell.execute_reply.started":"2023-06-24T07:25:30.005852Z","shell.execute_reply":"2023-06-24T07:25:30.005875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv('/kaggle/input/google-research-identify-contrails-reduce-global-warming/sample_submission.csv', index_col='record_id')\nbest_model.eval()\nwith torch.no_grad():\n    for X, rec in test_loader:\n        X = X.to(device)\n        pred = sigmoid(best_model(X)).cpu().detach().numpy().copy()[:,0,:,:] \n        mask = np.zeros((len(rec), 256, 256))\n        mask[pred<CFG.threshold] = 0\n        mask[pred>CFG.threshold] = 1\n        \n        for file_id, file_name in enumerate(rec):\n            submission.loc[int(file_name), 'encoded_pixels'] = list_to_string(rle_encode(mask[file_id,:,:]))\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-24T07:25:30.007429Z","iopub.status.idle":"2023-06-24T07:25:30.008254Z","shell.execute_reply.started":"2023-06-24T07:25:30.007945Z","shell.execute_reply":"2023-06-24T07:25:30.00797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-06-24T07:25:30.009474Z","iopub.status.idle":"2023-06-24T07:25:30.010167Z","shell.execute_reply.started":"2023-06-24T07:25:30.009917Z","shell.execute_reply":"2023-06-24T07:25:30.009939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}