{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Summary\n#### [<u>Initial version</u>](https://www.kaggle.com/code/egortrushin/gr-icrgw-pytorch-lightning-baseline-unet-resnest)\n- Baseline written using Pytorch Lightning\n- resnest26d as encoder\n- Dataset is taken from: https://www.kaggle.com/datasets/shashwatraman/contrails-images-ash-color\n\n#### [<u>Improved version</u>](https://www.kaggle.com/code/egortrushin/gr-icrgw-pl-pipeline-improved)\n- Option to change image size\n- Mixed precision training (only useful with T4x2, on P100 this slows down training). This helps to use GPU memory more efficiently\n- Training using 2 GPUs - with 2 GPUs we have more memory and higher speed\n- Other numerous small changes\n\n#### <u>Present version</u>\n- Training with 4-folds\n- LR scheduler: cosine with warmup\n- Use of CSVLogger with consequent visualization of the optimization process. Since I train without internet, I am limited to *local* CSVLogger or TensorBoardLogger. Alternatively you can train with internet and WanddbLogger.\n- Submission part is rewritten to make it cleaner and to allow easy work with multi-fold models","metadata":{}},{"cell_type":"markdown","source":"### Training part","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-output":true,"execution":{"iopub.status.busy":"2023-07-04T15:10:27.553577Z","iopub.execute_input":"2023-07-04T15:10:27.553924Z","iopub.status.idle":"2023-07-04T15:10:30.347567Z","shell.execute_reply.started":"2023-07-04T15:10:27.553897Z","shell.execute_reply":"2023-07-04T15:10:30.346488Z"},"trusted":true},"execution_count":2,"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-output":true,"execution":{"iopub.status.busy":"2023-07-04T15:10:30.349274Z","iopub.execute_input":"2023-07-04T15:10:30.349662Z","iopub.status.idle":"2023-07-04T15:10:33.614378Z","shell.execute_reply.started":"2023-07-04T15:10:30.349634Z","shell.execute_reply":"2023-07-04T15:10:33.613023Z"},"trusted":true},"execution_count":3,"outputs":[]},{"cell_type":"code","source":"%%writefile config.yaml\n\ndata_path: \"/kaggle/input/contrails-images-ash-color\"\noutput_dir: \"models\"\n\nfolds:\n    n_splits: 4\n    random_state: 42\ntrain_folds: [0, 1, 2, 3]\n    \nseed: 42\n\ntrain_bs: 48\nvalid_bs: 128\nworkers: 2\n\nprogress_bar_refresh_rate: 1\n\nearly_stop:\n    monitor: \"val_loss\"\n    mode: \"min\"\n    patience: 999\n    verbose: 1\n\ntrainer:\n    max_epochs: 20\n    min_epochs: 20\n    enable_progress_bar: True\n    precision: \"16-mixed\"\n    devices: 2\n\nmodel:\n    seg_model: \"Unet\"\n    encoder_name: \"timm-resnest26d\"\n    loss_smooth: 1.0\n    image_size: 384\n    optimizer_params:\n        lr: 0.0005\n        weight_decay: 0.0\n    scheduler:\n        name: \"cosine_with_hard_restarts_schedule_with_warmup\"\n        params:\n            cosine_with_hard_restarts_schedule_with_warmup:\n                num_warmup_steps: 350\n                num_training_steps: 3150\n                num_cycles: 1","metadata":{"execution":{"iopub.status.busy":"2023-07-04T15:10:33.617041Z","iopub.execute_input":"2023-07-04T15:10:33.617356Z","iopub.status.idle":"2023-07-04T15:10:33.626167Z","shell.execute_reply.started":"2023-07-04T15:10:33.617326Z","shell.execute_reply":"2023-07-04T15:10:33.62514Z"},"trusted":true},"execution_count":4,"outputs":[{"name":"stdout","text":"Writing config.yaml\n","output_type":"stream"}]},{"cell_type":"code","source":"# Dataset\n\nimport torch\nimport numpy as np\nimport torchvision.transforms as T\n\nclass ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, df, image_size=256, train=True):\n\n        self.df = df\n        self.trn = train\n        self.normalize_image = T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n        self.image_size = image_size\n        if image_size != 256:\n            self.resize_image = T.transforms.Resize(image_size)\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        con_path = row.path\n        con = np.load(str(con_path))\n\n        img = con[..., :-1]\n        label = con[..., -1]\n\n        label = torch.tensor(label)\n\n        img = torch.tensor(np.reshape(img, (256, 256, 3))).to(torch.float32).permute(2, 0, 1)\n\n        if self.image_size != 256:\n            img = self.resize_image(img)\n\n        img = self.normalize_image(img)\n\n        return img.float(), label.float()\n\n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-07-04T15:10:40.472823Z","iopub.execute_input":"2023-07-04T15:10:40.473174Z","iopub.status.idle":"2023-07-04T15:10:40.48326Z","shell.execute_reply.started":"2023-07-04T15:10:40.473141Z","shell.execute_reply":"2023-07-04T15:10:40.481929Z"},"trusted":true},"execution_count":6,"outputs":[]},{"cell_type":"code","source":"# Lightning module\n\nimport torch\nimport pytorch_lightning as pl\nimport segmentation_models_pytorch as smp\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, ReduceLROnPlateau\nfrom torch.optim import AdamW\nimport torch.nn as nn\nfrom torchmetrics.functional import dice\nfrom transformers import get_cosine_with_hard_restarts_schedule_with_warmup\n\nseg_models = {\n    \"Unet\": smp.Unet,\n    \"Unet++\": smp.UnetPlusPlus,\n    \"MAnet\": smp.MAnet,\n    \"Linknet\": smp.Linknet,\n    \"FPN\": smp.FPN,\n    \"PSPNet\": smp.PSPNet,\n    \"PAN\": smp.PAN,\n    \"DeepLabV3\": smp.DeepLabV3,\n    \"DeepLabV3+\": smp.DeepLabV3Plus,\n}\n\n\nclass LightningModule(pl.LightningModule):\n    def __init__(self, config):\n        super().__init__()\n        self.config = config\n        self.model = model = seg_models[config[\"seg_model\"]](\n            encoder_name=config[\"encoder_name\"],\n            encoder_weights=\"imagenet\",\n            in_channels=3,\n            classes=1,\n            activation=None,\n        )\n        self.loss_module = smp.losses.DiceLoss(mode=\"binary\", smooth=config[\"loss_smooth\"])\n        self.val_step_outputs = []\n        self.val_step_labels = []\n\n    def forward(self, batch):\n        imgs = batch\n        preds = self.model(imgs)\n        return preds\n\n    def configure_optimizers(self):\n        optimizer = AdamW(self.parameters(), **self.config[\"optimizer_params\"])\n\n        if self.config[\"scheduler\"][\"name\"] == \"CosineAnnealingLR\":\n            scheduler = CosineAnnealingLR(\n                optimizer,\n                **self.config[\"scheduler\"][\"params\"][\"CosineAnnealingLR\"],\n            )\n            lr_scheduler_dict = {\"scheduler\": scheduler, \"interval\": \"step\"}\n            return {\"optimizer\": optimizer, \"lr_sched 77/175 [0uler\": lr_scheduler_dict}\n        elif self.config[\"scheduler\"][\"name\"] == \"ReduceLROnPlateau\":\n            scheduler = ReduceLROnPlateau(\n                optimizer,\n                **self.config[\"scheduler\"][\"params\"][\"ReduceLROnPlateau\"],\n            )\n            lr_scheduler = {\"scheduler\": scheduler, \"monitor\": \"val_loss\"}\n            return {\"optimizer\": optimizer, \"lr_scheduler\": lr_scheduler}\n        elif self.config[\"scheduler\"][\"name\"] == \"cosine_with_hard_restarts_schedule_with_warmup\":\n            scheduler = get_cosine_with_hard_restarts_schedule_with_warmup(\n                optimizer,\n                **self.config[\"scheduler\"][\"params\"][self.config[\"scheduler\"][\"name\"]],\n            )\n            lr_scheduler_dict = {\"scheduler\": scheduler, \"interval\": \"step\"}\n            return {\"optimizer\": optimizer, \"lr_scheduler\": lr_scheduler_dict}\n\n    def training_step(self, batch, batch_idx):\n        imgs, labels = batch\n        preds = self.model(imgs)\n        if self.config[\"image_size\"] != 256:\n            preds = torch.nn.functional.interpolate(preds, size=256, mode='bilinear')\n        loss = self.loss_module(preds, labels)\n        self.log(\"train_loss\", loss, on_step=True, on_epoch=True, prog_bar=True, batch_size=16)\n\n        for param_group in self.trainer.optimizers[0].param_groups:\n            lr = param_group[\"lr\"]\n        self.log(\"lr\", lr, on_step=True, on_epoch=False, prog_bar=True)\n\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        imgs, labels = batch\n        preds = self.model(imgs)\n        if self.config[\"image_size\"] != 256:\n            preds = torch.nn.functional.interpolate(preds, size=256, mode='bilinear')\n        loss = self.loss_module(preds, labels)\n        self.log(\"val_loss\", loss, on_step=False, on_epoch=True, prog_bar=True)\n        self.val_step_outputs.append(preds)\n        self.val_step_labels.append(labels)\n\n    def on_validation_epoch_end(self):\n        all_preds = torch.cat(self.val_step_outputs)\n        all_labels = torch.cat(self.val_step_labels)\n        all_preds = torch.sigmoid(all_preds)\n        self.val_step_outputs.clear()\n        self.val_step_labels.clear()\n        val_dice = dice(all_preds, all_labels.long())\n        self.log(\"val_dice\", val_dice, on_step=False, on_epoch=True, prog_bar=True)\n        if self.trainer.global_rank == 0:\n            print(f\"\\nEpoch: {self.current_epoch}\", flush=True)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-07-04T15:10:48.3669Z","iopub.execute_input":"2023-07-04T15:10:48.367285Z","iopub.status.idle":"2023-07-04T15:10:48.404472Z","shell.execute_reply.started":"2023-07-04T15:10:48.367254Z","shell.execute_reply":"2023-07-04T15:10:48.403341Z"},"trusted":true},"execution_count":7,"outputs":[]},{"cell_type":"code","source":"# Actual training\n\nimport warnings\n\nwarnings.filterwarnings(\"ignore\")\n\nimport gc\nimport os\nimport torch\nimport yaml\nimport pandas as pd\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping, TQDMProgressBar\nfrom torch.utils.data import DataLoader\nfrom sklearn.model_selection import KFold\nfrom pytorch_lightning.loggers import CSVLogger\n\ntorch.set_float32_matmul_precision(\"medium\")\n\nwith open(\"config.yaml\", \"r\") as file_obj:\n    config = yaml.safe_load(file_obj)\n\npl.seed_everything(config[\"seed\"])\n\ngc.enable()\n\ncontrails = os.path.join(config[\"data_path\"], \"contrails/\")\ntrain_path = os.path.join(config[\"data_path\"], \"train_df.csv\")\nvalid_path = os.path.join(config[\"data_path\"], \"valid_df.csv\")\n\ntrain_df = pd.read_csv(train_path)\nvalid_df = pd.read_csv(valid_path)\n\ntrain_df[\"path\"] = contrails + train_df[\"record_id\"].astype(str) + \".npy\"\nvalid_df[\"path\"] = contrails + valid_df[\"record_id\"].astype(str) + \".npy\"\n\ndf = pd.concat([train_df, valid_df]).reset_index()\n\nFold = KFold(shuffle=True, **config[\"folds\"])\nfor n, (trn_index, val_index) in enumerate(Fold.split(df)):\n    df.loc[val_index, \"kfold\"] = int(n)\ndf[\"kfold\"] = df[\"kfold\"].astype(int)\n\nfor fold in config[\"train_folds\"]:\n    print(f\"\\n###### Fold {fold}\")\n    print(df[df.kfold != fold])\n    trn_df = df[df.kfold != fold].reset_index(drop=True)\n    print(trn_df)\n    print(\"Before exit\")\n    exit(1)\n    print(\"After exit\")  # This line will not be executed\n    vld_df = df[df.kfold == fold].reset_index(drop=True)\n\n    dataset_train = ContrailsDataset(trn_df, config[\"model\"][\"image_size\"], train=True)\n    dataset_validation = ContrailsDataset(vld_df, config[\"model\"][\"image_size\"], train=False)\n\n    data_loader_train = DataLoader(\n        dataset_train,\n        batch_size=config[\"train_bs\"],\n        shuffle=True,\n        num_workers=config[\"workers\"],\n    )\n    data_loader_validation = DataLoader(\n        dataset_validation,\n        batch_size=config[\"valid_bs\"],\n        shuffle=False,\n        num_workers=config[\"workers\"],\n    )\n\n    checkpoint_callback = ModelCheckpoint(\n        save_weights_only=True,\n        monitor=\"val_dice\",\n        dirpath=config[\"output_dir\"],\n        mode=\"max\",\n        filename=f\"model-f{fold}-{{val_dice:.4f}}\",\n        save_top_k=1,\n        verbose=1,\n    )\n\n    progress_bar_callback = TQDMProgressBar(\n        refresh_rate=config[\"progress_bar_refresh_rate\"]\n    )\n\n    early_stop_callback = EarlyStopping(**config[\"early_stop\"])\n\n\n    trainer = pl.Trainer(\n        callbacks=[checkpoint_callback, early_stop_callback, progress_bar_callback],\n        logger=CSVLogger(save_dir=f'logs_f{fold}/'),\n        **config[\"trainer\"],\n    )\n\n    model = LightningModule(config[\"model\"])\n\n    trainer.fit(model, data_loader_train, data_loader_validation)\n\n    del (\n        dataset_train,\n        dataset_validation,\n        data_loader_train,\n        data_loader_validation,\n        model,\n        trainer,\n        checkpoint_callback,\n        progress_bar_callback,\n        early_stop_callback,\n    )\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import seaborn as sn\nimport matplotlib.pyplot as plt\n\nfor fold in config[\"train_folds\"]:\n    metrics = pd.read_csv(f\"/kaggle/working/logs_f{fold}/lightning_logs/version_0/metrics.csv\")\n    del metrics[\"step\"]\n    del metrics[\"lr\"]\n    del metrics[\"train_loss_step\"]\n    metrics.set_index(\"epoch\", inplace=True)\n    g = sn.relplot(data=metrics, kind=\"line\")\n    plt.title(f\"Fold {fold}\")\n    plt.gcf().set_size_inches(15, 5)\n    plt.grid()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-06-19T16:28:05.101431Z","iopub.execute_input":"2023-06-19T16:28:05.101875Z","iopub.status.idle":"2023-06-19T16:28:08.685494Z","shell.execute_reply.started":"2023-06-19T16:28:05.101841Z","shell.execute_reply":"2023-06-19T16:28:08.684405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\ntest=np.load('/kaggle/input/contrails-images-ash-color/contrails/1000216489776414077.npy',encoding = \"latin1\")  #加载文件\n\nprint(test)\nprint(test.shape)\n# doc = open('1.txt', 'a')  #打开一个存储文件，并依次写入\n# print(test, file=doc)  #将打印内容写入文件中\n","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:53:52.421607Z","iopub.execute_input":"2023-07-04T14:53:52.422074Z","iopub.status.idle":"2023-07-04T14:53:52.435408Z","shell.execute_reply.started":"2023-07-04T14:53:52.422041Z","shell.execute_reply":"2023-07-04T14:53:52.434246Z"},"trusted":true},"execution_count":3,"outputs":[{"name":"stdout","text":"[[[0.09564  0.704    0.783    0.      ]\n  [0.07886  0.7114   0.802    0.      ]\n  [0.0816   0.7017   0.7983   0.      ]\n  ...\n  [0.0797   0.5776   0.6904   0.      ]\n  [0.0968   0.602    0.6914   0.      ]\n  [0.11597  0.612    0.691    0.      ]]\n\n [[0.1509   0.7305   0.744    0.      ]\n  [0.1188   0.7095   0.777    0.      ]\n  [0.10986  0.6987   0.7944   0.      ]\n  ...\n  [0.03583  0.5825   0.7515   0.      ]\n  [0.06097  0.573    0.7236   0.      ]\n  [0.1017   0.583    0.688    0.      ]]\n\n [[0.197    0.736    0.7153   0.      ]\n  [0.1682   0.7246   0.7446   0.      ]\n  [0.1353   0.709    0.7856   0.      ]\n  ...\n  [0.009125 0.588    0.769    0.      ]\n  [0.01287  0.5684   0.752    0.      ]\n  [0.03534  0.5557   0.7285   0.      ]]\n\n ...\n\n [[0.       0.6333   0.8164   0.      ]\n  [0.       0.6025   0.7983   0.      ]\n  [0.       0.6147   0.801    0.      ]\n  ...\n  [0.01249  0.5176   0.669    0.      ]\n  [0.01206  0.526    0.6772   0.      ]\n  [0.       0.5015   0.667    0.      ]]\n\n [[0.       0.6353   0.821    0.      ]\n  [0.       0.5815   0.7847   0.      ]\n  [0.       0.553    0.7656   0.      ]\n  ...\n  [0.08356  0.536    0.6445   0.      ]\n  [0.08167  0.5522   0.6606   0.      ]\n  [0.01918  0.51     0.666    0.      ]]\n\n [[0.       0.6367   0.8286   0.      ]\n  [0.       0.6      0.797    0.      ]\n  [0.       0.5576   0.768    0.      ]\n  ...\n  [0.1399   0.525    0.6255   0.      ]\n  [0.117    0.5757   0.66     0.      ]\n  [0.01845  0.509    0.683    0.      ]]]\n(256, 256, 4)\n","output_type":"stream"}]},{"cell_type":"code","source":"\nimport numpy as np\ntest=np.load('/kaggle/input/google-research-identify-contrails-reduce-global-warming/test/1000834164244036115/band_08.npy',encoding = \"latin1\")  #加载文件\n\nprint(test)\nprint(test.shape)","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:54:03.447218Z","iopub.execute_input":"2023-07-04T14:54:03.447569Z","iopub.status.idle":"2023-07-04T14:54:03.4869Z","shell.execute_reply.started":"2023-07-04T14:54:03.44754Z","shell.execute_reply":"2023-07-04T14:54:03.485716Z"},"trusted":true},"execution_count":4,"outputs":[{"name":"stdout","text":"[[[228.16599 230.11641 232.30144 ... 227.04819 226.94821 224.60127]\n  [228.46814 231.45676 231.5972  ... 226.4471  226.93706 223.59825]\n  [229.74626 232.34785 230.9957  ... 228.28891 226.3236  221.57967]\n  ...\n  [232.10223 232.00272 231.96104 ... 228.8069  228.71698 229.70537]\n  [232.09761 231.82973 231.19353 ... 228.90335 229.1747  229.6833 ]\n  [232.13109 231.14166 230.91747 ... 229.24785 229.30815 229.80779]]\n\n [[229.61823 230.14343 232.46259 ... 228.36987 227.9069  225.01343]\n  [229.3326  231.71704 232.00809 ... 228.1524  228.69707 224.37016]\n  [230.0375  232.64369 231.7687  ... 229.31224 228.21123 222.44333]\n  ...\n  [232.39787 232.03084 231.81126 ... 229.38757 228.67177 229.73633]\n  [232.2574  231.87279 231.07628 ... 229.03514 228.96396 229.73584]\n  [232.20694 231.36177 230.76237 ... 228.97862 228.85626 229.84192]]\n\n [[230.69583 230.22292 232.66068 ... 229.5458  229.02426 225.6884 ]\n  [230.14822 232.1164  232.48479 ... 229.64914 230.41295 225.27357]\n  [230.41557 232.91576 232.34334 ... 230.0503  230.16658 223.58644]\n  ...\n  [232.63504 232.02701 231.80135 ... 229.69577 228.9005  229.60191]\n  [232.32439 231.9565  231.02863 ... 229.2465  229.04553 229.50507]\n  [232.18988 231.72932 230.62218 ... 228.8974  228.85199 229.7048 ]]\n\n ...\n\n [[236.28958 236.72324 236.7262  ... 236.56082 237.07793 237.56606]\n  [236.33115 236.8313  236.80396 ... 236.37117 236.93774 237.296  ]\n  [236.34879 236.77763 236.86697 ... 236.26852 236.88264 237.07005]\n  ...\n  [236.03775 236.20012 236.28616 ... 236.05426 235.69064 235.36497]\n  [236.20158 236.11989 236.3438  ... 235.94507 235.69888 235.45892]\n  [236.25728 236.16095 236.31267 ... 235.88634 235.68674 235.48671]]\n\n [[236.4842  236.66554 236.6746  ... 236.51634 237.02873 237.32884]\n  [236.48668 236.73912 236.72406 ... 236.38594 236.80792 237.10056]\n  [236.44003 236.7603  236.73436 ... 236.25058 236.67209 237.12044]\n  ...\n  [236.12813 236.22325 236.28383 ... 236.00609 235.66048 235.52374]\n  [236.18613 236.18365 236.36021 ... 235.99529 235.6738  235.54088]\n  [236.24799 236.19357 236.27217 ... 235.91237 235.65007 235.58456]]\n\n [[236.50494 236.52325 236.67972 ... 236.48268 237.08765 237.29504]\n  [236.56744 236.65845 236.73439 ... 236.42085 236.84238 237.12547]\n  [236.62561 236.71498 236.66243 ... 236.20366 236.54704 237.19414]\n  ...\n  [236.26196 236.21678 236.35585 ... 236.0237  235.69116 235.68306]\n  [236.2472  236.20012 236.30869 ... 236.05968 235.68422 235.62274]\n  [236.25978 236.20012 236.29765 ... 236.04135 235.64194 235.61906]]]\n(256, 256, 8)\n","output_type":"stream"}]},{"cell_type":"markdown","source":"### Submission part","metadata":{}},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport gc\nimport os\nimport glob\n\nimport numpy as np\nimport pandas as pd\n\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\n\nimport pytorch_lightning as pl\nimport torchvision.transforms as T\nimport yaml","metadata":{"execution":{"iopub.status.busy":"2023-06-19T16:28:18.650212Z","iopub.execute_input":"2023-06-19T16:28:18.650618Z","iopub.status.idle":"2023-06-19T16:28:18.656403Z","shell.execute_reply.started":"2023-06-19T16:28:18.650587Z","shell.execute_reply":"2023-06-19T16:28:18.655478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 32\nnum_workers = 1\nTHR = 0.5\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndata = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'\ndata_root = '/kaggle/input/google-research-identify-contrails-reduce-global-warming/test/'\nsubmission = pd.read_csv(os.path.join(data, 'sample_submission.csv'), index_col='record_id')","metadata":{"execution":{"iopub.status.busy":"2023-06-19T16:28:21.588978Z","iopub.execute_input":"2023-06-19T16:28:21.589487Z","iopub.status.idle":"2023-06-19T16:28:21.605792Z","shell.execute_reply.started":"2023-06-19T16:28:21.58943Z","shell.execute_reply":"2023-06-19T16:28:21.604733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"filenames = os.listdir(data_root)\ntest_df = pd.DataFrame(filenames, columns=['record_id'])\ntest_df['path'] = data_root + test_df['record_id'].astype(str)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T16:28:25.552322Z","iopub.execute_input":"2023-06-19T16:28:25.553302Z","iopub.status.idle":"2023-06-19T16:28:25.563746Z","shell.execute_reply.started":"2023-06-19T16:28:25.553234Z","shell.execute_reply":"2023-06-19T16:28:25.562542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, df, image_size=256, train=True):\n        \n        self.df = df\n        self.trn = train\n        self.df_idx: pd.DataFrame = pd.DataFrame({'idx': os.listdir(f'/kaggle/input/google-research-identify-contrails-reduce-global-warming/test')})\n        self.normalize_image = T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n        self.image_size = image_size\n        if image_size != 256:\n            self.resize_image = T.transforms.Resize(image_size)\n    \n    def read_record(self, directory):\n        record_data = {}\n        for x in [\n            \"band_11\", \n            \"band_14\", \n            \"band_15\"\n        ]:\n\n            record_data[x] = np.load(os.path.join(directory, x + \".npy\"))\n\n        return record_data\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, index):\n        row = self.df.iloc[index]\n        con_path = row.path\n        data = self.read_record(con_path)    \n        \n        img = self.get_false_color(data)\n        \n        img = torch.tensor(np.reshape(img, (256, 256, 3))).to(torch.float32).permute(2, 0, 1)\n        \n        if self.image_size != 256:\n            img = self.resize_image(img)\n        \n        img = self.normalize_image(img)\n        \n        image_id = int(self.df_idx.iloc[index]['idx'])\n            \n        return img.float(), torch.tensor(image_id)\n    \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T16:28:28.391329Z","iopub.execute_input":"2023-06-19T16:28:28.391814Z","iopub.status.idle":"2023-06-19T16:28:28.415525Z","shell.execute_reply.started":"2023-06-19T16:28:28.391777Z","shell.execute_reply":"2023-06-19T16:28:28.414613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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\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-06-19T16:28:31.918781Z","iopub.execute_input":"2023-06-19T16:28:31.919159Z","iopub.status.idle":"2023-06-19T16:28:31.927124Z","shell.execute_reply.started":"2023-06-19T16:28:31.919129Z","shell.execute_reply":"2023-06-19T16:28:31.92614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LightningModule(pl.LightningModule):\n\n    def __init__(self, config):\n        super().__init__()\n        self.model = smp.Unet(encoder_name=config[\"encoder_name\"],\n                              encoder_weights=None,\n                              in_channels=3,\n                              classes=1,\n                              activation=None,\n                              )\n\n    def forward(self, batch):\n        return self.model(batch)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T16:29:52.1835Z","iopub.execute_input":"2023-06-19T16:29:52.183902Z","iopub.status.idle":"2023-06-19T16:29:52.190559Z","shell.execute_reply.started":"2023-06-19T16:29:52.183873Z","shell.execute_reply":"2023-06-19T16:29:52.189333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODEL_PATH = \"/kaggle/working/models/\"\n#with open(os.path.join(MODEL_PATH, \"config.yaml\"), \"r\") as file_obj:\n#    config = yaml.safe_load(file_obj)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T16:28:56.051775Z","iopub.execute_input":"2023-06-19T16:28:56.052165Z","iopub.status.idle":"2023-06-19T16:28:56.056861Z","shell.execute_reply.started":"2023-06-19T16:28:56.052134Z","shell.execute_reply":"2023-06-19T16:28:56.055785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ds = ContrailsDataset(\n        test_df,\n        config[\"model\"][\"image_size\"],\n        train = False\n    )\n \ntest_dl = DataLoader(test_ds, batch_size=batch_size, num_workers = num_workers)","metadata":{"execution":{"iopub.status.busy":"2023-06-19T16:28:58.944979Z","iopub.execute_input":"2023-06-19T16:28:58.946168Z","iopub.status.idle":"2023-06-19T16:28:58.969457Z","shell.execute_reply.started":"2023-06-19T16:28:58.946104Z","shell.execute_reply":"2023-06-19T16:28:58.968329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.enable()\n\nall_preds = {}\n\nfor i, model_path in enumerate(glob.glob(MODEL_PATH + '*.ckpt')):\n    print(model_path)\n    model = LightningModule(config[\"model\"]).load_from_checkpoint(model_path, config=config[\"model\"])\n    model.to(device)\n    model.eval()\n\n    model_preds = {}\n    \n    for _, data in enumerate(test_dl):\n        images, image_id = data\n    \n        images = images.to(device)\n        \n        with torch.no_grad():\n            predicted_mask = model(images[:, :, :, :])\n        if config[\"model\"][\"image_size\"] != 256:\n            predicted_mask = torch.nn.functional.interpolate(predicted_mask, size=256, mode='bilinear')\n        predicted_mask = torch.sigmoid(predicted_mask).cpu().detach().numpy()\n                \n        for img_num in range(0, images.shape[0]):\n            current_mask = predicted_mask[img_num, :, :, :]\n            current_image_id = image_id[img_num].item()\n            model_preds[current_image_id] = current_mask\n    all_preds[f\"f{i}\"] = model_preds\n    \n    del model    \n    torch.cuda.empty_cache()\n    gc.collect() ","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-19T16:29:57.362351Z","iopub.execute_input":"2023-06-19T16:29:57.362741Z","iopub.status.idle":"2023-06-19T16:30:17.81402Z","shell.execute_reply.started":"2023-06-19T16:29:57.36271Z","shell.execute_reply":"2023-06-19T16:30:17.812765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for index in submission.index.tolist():\n    for i in range(len(glob.glob(MODEL_PATH + '*.ckpt'))):\n        if i == 0:\n            predicted_mask = all_preds[f\"f{i}\"][index]\n        else:\n            predicted_mask += all_preds[f\"f{i}\"][index]\n    predicted_mask = predicted_mask / len(glob.glob(MODEL_PATH + '*.ckpt'))\n    predicted_mask_with_threshold = np.zeros((256, 256))\n    predicted_mask_with_threshold[predicted_mask[0, :, :] < THR] = 0\n    predicted_mask_with_threshold[predicted_mask[0, :, :] > THR] = 1\n    submission.loc[int(index), 'encoded_pixels'] = list_to_string(rle_encode(predicted_mask_with_threshold))","metadata":{"execution":{"iopub.status.busy":"2023-06-19T16:30:22.372528Z","iopub.execute_input":"2023-06-19T16:30:22.373336Z","iopub.status.idle":"2023-06-19T16:30:22.393671Z","shell.execute_reply.started":"2023-06-19T16:30:22.373247Z","shell.execute_reply":"2023-06-19T16:30:22.392508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2023-06-19T16:30:24.828167Z","iopub.execute_input":"2023-06-19T16:30:24.829156Z","iopub.status.idle":"2023-06-19T16:30:24.84522Z","shell.execute_reply.started":"2023-06-19T16:30:24.829124Z","shell.execute_reply":"2023-06-19T16:30:24.844178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-06-19T16:30:27.076471Z","iopub.execute_input":"2023-06-19T16:30:27.076848Z","iopub.status.idle":"2023-06-19T16:30:27.085225Z","shell.execute_reply.started":"2023-06-19T16:30:27.076817Z","shell.execute_reply":"2023-06-19T16:30:27.084284Z"},"trusted":true},"execution_count":null,"outputs":[]}]}