{"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":"# Chemtrails Composer Model Train\n---\n\n### <a href='#hyperparameters'> ⚙️ Hyperparameters </a> | <a href='#data-processing'> 📦️ Data Processing </a> | <a href='#model-training'> 🔥️ Model Training </a>\n","metadata":{}},{"cell_type":"code","source":"\"\"\"\nTODO\n- Implement in Pure Pytorch for inference\n- Implement the actual evaluation metric as loss function\n- Implement Mask2Former, IntermImage-H backbones\n- Read more segmentation papers\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:02:28.307923Z","iopub.execute_input":"2023-07-05T18:02:28.308711Z","iopub.status.idle":"2023-07-05T18:02:28.322751Z","shell.execute_reply.started":"2023-07-05T18:02:28.308678Z","shell.execute_reply":"2023-07-05T18:02:28.321878Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Setup & Imports","metadata":{}},{"cell_type":"code","source":"# Installations, Setup and Imports\n!pip install -q mosaicml\n!pip install -q omegaconf\n!pip install -q segmentation-models-pytorch\n\n# Commonly Used Libraries\nimport pandas as pd\nimport numpy as np\nfrom pathlib import Path\nimport collections\nimport termcolor\nimport functools \nimport random\nimport pickle\nimport os\nimport re\nimport gc\n\nfrom tqdm.auto import tqdm\ntqdm.pandas()\n\nimport omegaconf\nimport wandb\n# !wandb login '3xxxxxxxxxxxxxxxxxxxxxxxxxd'\n\nimport transformers\nimport datasets\nimport sklearn\nimport sklearn.metrics\n\n## PyTorch CV Imports ##\nimport torch\nimport torchmetrics\nfrom torch import nn\n\nimport albumentations\nimport composer\n\n\nfrom IPython.core.magic import register_line_cell_magic\n@register_line_cell_magic\ndef hyperparameters(hp_var_name, cell):\n    with open('experiment.yaml', 'w') as f:\n        f.write(cell)\n    HP = omegaconf.OmegaConf.load('experiment.yaml')\n    get_ipython().user_ns[hp_var_name] = HP","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:02:28.331007Z","iopub.execute_input":"2023-07-05T18:02:28.331546Z","iopub.status.idle":"2023-07-05T18:03:34.086597Z","shell.execute_reply.started":"2023-07-05T18:02:28.331499Z","shell.execute_reply":"2023-07-05T18:03:34.085647Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ⚙️ Hyperparameters\n\n<a name='hyperparameters'/>","metadata":{}},{"cell_type":"code","source":"%%hyperparameters args\n\n## Backbone ## \nbackbone_name: 'nvidia/mit-b5' # b0 to b5\n\n## Model Training ##\nnum_train_epochs: 10\ntrain_batch_size: 16\nimg_size: 256\n\n## Submission ##\ntest_img_size: 256\ntest_batch_size: 64\n\n## Cosine One Cycle Schedule\nwarmup_ratio: 0.10\npeak_lr: 1e-4\n\n## AdamW Optimizer ##\nweight_decay: 1e-5\nmax_grad_norm: 100.00\n\n## Loss Function ##\ndebug: False","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:03:34.088702Z","iopub.execute_input":"2023-07-05T18:03:34.089045Z","iopub.status.idle":"2023-07-05T18:03:34.102747Z","shell.execute_reply.started":"2023-07-05T18:03:34.089013Z","shell.execute_reply":"2023-07-05T18:03:34.101763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import transformers\nbackbone = transformers.AutoModelForSemanticSegmentation.from_pretrained(\n    args.backbone_name, \n    id2label={0: 'chemtrail'}, \n    label2id={'chemtrail': 0},\n)","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:03:34.103969Z","iopub.execute_input":"2023-07-05T18:03:34.104884Z","iopub.status.idle":"2023-07-05T18:03:49.083646Z","shell.execute_reply.started":"2023-07-05T18:03:34.104852Z","shell.execute_reply":"2023-07-05T18:03:49.082679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 📦️ Data Processing\n\n<a name='data-processing'>","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv('/kaggle/input/contrails-images-ash-color/train_df.csv')\nvalid_df = pd.read_csv('/kaggle/input/contrails-images-ash-color/valid_df.csv')\n\nimg_dir = '/kaggle/input/contrails-images-ash-color/contrails/'\ntrain_df['img_path'] = img_dir + train_df.record_id.astype(str) + '.npy'\nvalid_df['img_path'] = img_dir + valid_df.record_id.astype(str) + '.npy'\n\ntrain_df","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:03:49.086456Z","iopub.execute_input":"2023-07-05T18:03:49.086852Z","iopub.status.idle":"2023-07-05T18:03:49.159359Z","shell.execute_reply.started":"2023-07-05T18:03:49.086797Z","shell.execute_reply":"2023-07-05T18:03:49.1584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision.transforms as T \n\nclass ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, img_paths):\n        self.img_paths = img_paths\n        self.normalize_img = T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n        \n    def __getitem__(self, idx):\n        img_path = self.img_paths[idx]\n        img = np.load(img_path)\n        img, label = img[..., :-1], img[..., -1]\n        img = np.reshape(img, (256, 256, 3)) \n        img = np.transpose(img, (2, 0, 1))\n        img = torch.tensor(img, dtype=torch.float32)\n        img = self.normalize_img(img)\n        return {\n            'img': img,\n            'label': torch.tensor(label, dtype=torch.float32),\n        }\n    \n    def __len__(self):\n        return len(self.img_paths)\n\ntrain_img_paths = train_df.img_path.values\nvalid_img_paths = valid_df.img_path.values\n\ntrain_dataset = ContrailsDataset(train_img_paths)\nvalid_dataset = ContrailsDataset(valid_img_paths)\n\ntrain_dataloader = torch.utils.data.DataLoader(\n    dataset=train_dataset,\n    batch_size=args.train_batch_size,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True,\n    drop_last=True,\n    #persistent_workers=True,\n)\nvalid_dataloader = torch.utils.data.DataLoader(\n    dataset=valid_dataset,\n    batch_size=args.test_batch_size,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True,\n    drop_last=False,\n    #persistent_workers=True,\n)\nfor batch in tqdm(valid_dataloader): \n    pass","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:03:49.168862Z","iopub.execute_input":"2023-07-05T18:03:49.169459Z","iopub.status.idle":"2023-07-05T18:04:07.639774Z","shell.execute_reply.started":"2023-07-05T18:03:49.169412Z","shell.execute_reply":"2023-07-05T18:04:07.638425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import segmentation_models_pytorch as smp\n\nclass LossMetric(torchmetrics.Metric):\n    def __init__(self, loss_fn, dist_sync_on_step=False):\n        super().__init__(dist_sync_on_step=dist_sync_on_step)\n        self.loss_fn = loss_fn\n        self.add_state('sum_loss', default=torch.tensor(0.), dist_reduce_fx='sum')\n        self.add_state('total_batches', default=torch.tensor(0), dist_reduce_fx='sum')\n    \n    def update(self, y_pred, y_true):\n        self.sum_loss += self.loss_fn(y_pred, y_true)\n        self.total_batches += 1\n    \n    def compute(self):\n        return self.sum_loss / self.total_batches\n\n\nclass ChemtrailsComposerModel(composer.models.ComposerModel):\n    \n    def __init__(self, backbone):\n        super().__init__()\n        self.backbone = backbone\n        self.criterion = smp.losses.DiceLoss(mode=\"binary\", smooth=1)\n        self.img_size = 256\n        \n    def forward(self, batch):\n        encoder_outputs = self.backbone.segformer(\n            pixel_values=batch['img'],\n            output_hidden_states=True, # Fed to decoder\n        )\n        encoder_hidden_states = encoder_outputs.hidden_states\n        logits = self.backbone.decode_head(encoder_hidden_states)\n        \n        # Upsample the logits to the images' original size\n        upsampled_logits = nn.functional.interpolate(\n            logits, \n            size=(self.img_size, self.img_size), \n            mode='bilinear', \n            align_corners=False,\n        )\n        upsampled_logits = upsampled_logits.squeeze(1)\n        # print('upsampled_logits:', upsampled_logits.shape)\n        return upsampled_logits\n        \n    def loss(self, logits, batch):\n        labels = batch['label']\n        loss = self.criterion(logits, labels.float())\n        return loss.mean()\n\n\nmodel = ChemtrailsComposerModel(backbone=backbone)\noptimizer = composer.optim.DecoupledAdamW(model.parameters(), lr=args.peak_lr, weight_decay=args.weight_decay)\n# optimizer = torch.optim.lr_scheduler.CosineAnnealingLR(\n#     optimizer, \n#     T_max=len(valid_dataloader)*args.num_train_epochs, \n#     eta_min=0, \n#     last_epoch=args.num_train_epochs, \n#     verbose=True,\n# )","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:08:08.125698Z","iopub.execute_input":"2023-07-05T18:08:08.126084Z","iopub.status.idle":"2023-07-05T18:08:08.146471Z","shell.execute_reply.started":"2023-07-05T18:08:08.126053Z","shell.execute_reply":"2023-07-05T18:08:08.145432Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import composer.algorithms\n\n# Changes memory format to torch.channels_last()\nchannels_last_algo = composer.algorithms.ChannelsLast()\n\n# Predict with moving average of model weights\nema_algo = composer.algorithms.EMA(half_life='1ep', ema_start='1ep')\n\n# Replace every nn.LayerNorm with fused implementation\n#fused_layer_norm_algo = composer.algorithms.FusedLayerNorm()\n\n# Gradient Clipping\ngrad_clipping_algo = composer.algorithms.GradientClipping(clipping_type='norm', clipping_threshold=args.max_grad_norm)\n\n# Progressively freeze the layers of the network during training\nlayer_freezing_algo = composer.algorithms.LayerFreezing(freeze_start=0.5, freeze_level=1.0)\n\n# Progressive Resizing skipped for segmentation\n# Sam Optimizer \nsam_optimizer_algo = composer.algorithms.SAM(rho=0.05, epsilon=1e-12, interval=4) \n\n# TODO: Couple SWA + Cosine Decay Schedule (?)\nswa_averaging_algo = composer.algorithms.SWA()\n\nalgorithms = [\n    channels_last_algo,\n    #ema_algo,\n    #fused_layer_norm_algo,\n    grad_clipping_algo,\n    #layer_freezing_algo,\n    sam_optimizer_algo,\n    #swa_averaging_algo,\n]","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:08:08.285209Z","iopub.execute_input":"2023-07-05T18:08:08.285547Z","iopub.status.idle":"2023-07-05T18:08:08.295379Z","shell.execute_reply.started":"2023-07-05T18:08:08.2855Z","shell.execute_reply":"2023-07-05T18:08:08.294418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 🔥️ Model Training \n\n<a name='model-training'>","metadata":{}},{"cell_type":"code","source":"trainer = composer.Trainer(\n    model=model,\n    train_dataloader=train_dataloader,\n    eval_dataloader=valid_dataloader,\n    max_duration=f'{args.num_train_epochs}ep',\n    optimizers=optimizer,\n    algorithms=algorithms,\n    device='gpu',\n    precision='amp_fp16',\n)\ntrainer.fit() ","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:08:10.719374Z","iopub.execute_input":"2023-07-05T18:08:10.720367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"test_files_root = '/kaggle/input/google-research-identify-contrails-reduce-global-warming/test/'\ntest_files = os.listdir(test_files_root)\ntest_df = pd.DataFrame(test_files, columns=['record_id'])\ntest_df['path'] = test_files_root + test_df.record_id.astype(str)\ntest_df","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:04:21.424178Z","iopub.status.idle":"2023-07-05T18:04:21.424831Z","shell.execute_reply.started":"2023-07-05T18:04:21.424598Z","shell.execute_reply":"2023-07-05T18:04:21.42462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision\nclass ContrailsTestDataset(torch.utils.data.Dataset):\n    def __init__(self, img_paths, record_ids):\n        self.img_paths = img_paths\n        self.record_ids = record_ids\n        self.normalize_img = torchvision.transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n        self.resize_image = torchvision.transforms.Resize(args.test_img_size)\n    \n    def normalize_range(self, data, bounds):\n        \"\"\"Maps data to the range [0, 1].\"\"\"\n        return (data - bounds[0]) / (bounds[1] - bounds[0])\n    \n    def get_false_color(self, record_data):\n        _T11_BOUNDS = (243, 303)\n        _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n        _TDIFF_BOUNDS = (-4, 2)\n        \n        N_TIMES_BEFORE = 4\n\n        r = self.normalize_range(record_data[\"band_15\"] - record_data[\"band_14\"], _TDIFF_BOUNDS)\n        g = self.normalize_range(record_data[\"band_14\"] - record_data[\"band_11\"], _CLOUD_TOP_TDIFF_BOUNDS)\n        b = self.normalize_range(record_data[\"band_14\"], _T11_BOUNDS)\n        false_color = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n        img = false_color[..., N_TIMES_BEFORE]\n\n        return img\n    \n    def __getitem__(self, idx):\n        img_dir = Path(self.img_paths[idx])\n        \n        record_data = {\n            'band_11': np.load(img_dir/'band_11.npy'),\n            'band_14': np.load(img_dir/'band_14.npy'),\n            'band_15': np.load(img_dir/'band_15.npy'),\n        }\n        img = self.get_false_color(record_data)\n        img = torch.tensor(np.reshape(img, (256, 256, 3))).to(torch.float32).permute(2, 0, 1)\n        img = self.normalize_img(img)\n        \n        return {\n            'img': img.float(),# , dtype=torch.float32),\n        }\n    \n    def __len__(self):\n        return len(self.record_ids)\n\ntest_dataset = ContrailsTestDataset(\n    img_paths=test_df.path.values, \n    record_ids=test_df.record_id.values\n)\ntest_dataloader = torch.utils.data.DataLoader(\n    dataset=test_dataset,\n    batch_size=args.test_batch_size,\n    num_workers=2,\n    pin_memory=True,\n)","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:04:21.426639Z","iopub.status.idle":"2023-07-05T18:04:21.427102Z","shell.execute_reply.started":"2023-07-05T18:04:21.426861Z","shell.execute_reply":"2023-07-05T18:04:21.426883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import scipy.special\n\nall_test_logits = trainer.predict(test_dataloader)\nall_test_logits = np.concatenate(all_test_logits)\nall_test_preds = scipy.special.expit(all_test_logits)","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:04:21.428724Z","iopub.status.idle":"2023-07-05T18:04:21.42917Z","shell.execute_reply.started":"2023-07-05T18:04:21.428942Z","shell.execute_reply":"2023-07-05T18:04:21.428962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_encode(segmented_img):\n    \"\"\"\n    segmented_img: 1 -> mask, 0 -> background\n    Returns run length as a list\n    \"\"\"\n    # 1 -> mask, 0 -> background\n    dots = np.where(segmented_img.T.flatten()==1)[0]\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\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","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:04:21.430714Z","iopub.status.idle":"2023-07-05T18:04:21.431165Z","shell.execute_reply.started":"2023-07-05T18:04:21.430935Z","shell.execute_reply":"2023-07-05T18:04:21.430957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.read_csv('/kaggle/input/google-research-identify-contrails-reduce-global-warming/sample_submission.csv', index_col='record_id')\n\nfor idx, record_id in enumerate(test_df.record_id.values):\n    segmentation_mask = np.where(all_test_preds[idx]>0.50, 1, 0)\n    sub_df.loc[int(record_id), 'encoded_pixels'] = list_to_string(rle_encode(segmentation_mask))\n    \nsub_df.to_csv('submission.csv')\nsub_df","metadata":{"execution":{"iopub.status.busy":"2023-07-05T18:04:21.432751Z","iopub.status.idle":"2023-07-05T18:04:21.433197Z","shell.execute_reply.started":"2023-07-05T18:04:21.432969Z","shell.execute_reply":"2023-07-05T18:04:21.43299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}