{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"papermill":{"default_parameters":{},"duration":118.739056,"end_time":"2023-08-12T07:56:14.876757","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2023-08-12T07:54:16.137701","version":"2.4.0"},"widgets":{"application/vnd.jupyter.widget-state+json":{"state":{"1f3a1ae3455346278b1cb29c7b68e32c":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_e3e427348361400db8393cf0efa8b59e","IPY_MODEL_9d7b07b12a284383a6a76803d3167a65","IPY_MODEL_4f69a940478640ef877a9fdedc0c13b5"],"layout":"IPY_MODEL_2d19eb43c9d64334a3ec294ce159058c"}},"24e814a6f3da4f1889a0323c8d5e0ba4":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"2d19eb43c9d64334a3ec294ce159058c":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"37830213025e4044999544fca9149295":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"4f69a940478640ef877a9fdedc0c13b5":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_abd9e8c96fc94557aa8fa341cb9da696","placeholder":"​","style":"IPY_MODEL_24e814a6f3da4f1889a0323c8d5e0ba4","value":" 2/2 [00:00&lt;00:00,  2.61it/s]"}},"7881caf14de84c66b1dee7b30dd554a6":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"9d7b07b12a284383a6a76803d3167a65":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ProgressView","bar_style":"success","description":"","description_tooltip":null,"layout":"IPY_MODEL_f4e719a69073435ea1e7b0ac1285f2a6","max":2,"min":0,"orientation":"horizontal","style":"IPY_MODEL_37830213025e4044999544fca9149295","value":2}},"abd9e8c96fc94557aa8fa341cb9da696":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"e3e427348361400db8393cf0efa8b59e":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_7881caf14de84c66b1dee7b30dd554a6","placeholder":"​","style":"IPY_MODEL_eafb5dbe22b34fe9a4a6cc4fc8eee41c","value":"100%"}},"eafb5dbe22b34fe9a4a6cc4fc8eee41c":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"f4e719a69073435ea1e7b0ac1285f2a6":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}}},"version_major":2,"version_minor":0}}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# U-Net model with symmetric label and rotation augmentation\n\n\n### Background\nThis is a semantic segmentation task, marking contrails in satellite images. This U-Net model is the basis of my two final models in the [1st-place solution](https://www.kaggle.com/competitions/google-research-identify-contrails-reduce-global-warming/discussion/430618).\n\n1. The problem is that the segmentation label is shifted 0.5 pixels and we cannot use rotation augmentation as usual.\n2. I shift the label y by 0.5 pixels and called it the symmetric label y_sym,\n3. train U-Net with y_sym using rotation augmentation.\n4. A tiny convolution, asym_conv, shifts 0.5 pixels back to the required prediction y_pred.\n\n\n### What this notebook does\n\n1. Demonstrate training for maxvit_tiny but only with 100 data. Full training takes too long.\n2. Load trained maxvit_tiny, predict, and submit\n3. 1 fold x 8 TTA score is validation 0.692, public 0.709, private 0.704\n\nWhen you do real training please change the configuration:\n- pretrained: True\n- batch_size: 8\n- and don't forget to save the model weights.\n\nBatch size 8 is assuming 24GB of GPU RAM and not possible with 16GB; you can try amp or gradient accumulation but I experience a little drop in the score.\n\n**Version 7**\n\nSimplified the model using the 'tu-' prefix in the segmentation_models_pytorch model, which allows us to use various timm models as the encoder.","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","papermill":{"duration":0.007548,"end_time":"2023-08-12T07:54:26.977963","exception":false,"start_time":"2023-08-12T07:54:26.970415","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Install for the internet-off code competitions\n! pip install -q --no-index --find-links=file:///kaggle/input/contrails-base-model/pip \\\n    segmentation_models_pytorch\n! cp -r /kaggle/input/contrails-base-model/src ./src","metadata":{"papermill":{"duration":13.969015,"end_time":"2023-08-12T07:54:40.953845","exception":false,"start_time":"2023-08-12T07:54:26.98483","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:57:33.280123Z","iopub.execute_input":"2023-09-08T02:57:33.280488Z","iopub.status.idle":"2023-09-08T02:57:46.822962Z","shell.execute_reply.started":"2023-09-08T02:57:33.280459Z","shell.execute_reply":"2023-09-08T02:57:46.821637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport sys\nimport glob\nimport time\nimport yaml\nfrom tqdm.auto import tqdm\nfrom sklearn.model_selection import KFold\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn.modules.loss import _Loss\n\nimport torchvision.transforms as T\nimport albumentations as A\nimport segmentation_models_pytorch as smp\n\nsys.path.append('/kaggle/working/src')\nimport util\nfrom lr_scheduler import Scheduler\nfrom submit import write_submission\n\n\ndi = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'\ndevice = torch.device('cuda')","metadata":{"papermill":{"duration":8.224344,"end_time":"2023-08-12T07:54:49.18573","exception":false,"start_time":"2023-08-12T07:54:40.961386","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:57:46.82629Z","iopub.execute_input":"2023-09-08T02:57:46.827072Z","iopub.status.idle":"2023-09-08T02:57:54.494068Z","shell.execute_reply.started":"2023-09-08T02:57:46.827031Z","shell.execute_reply":"2023-09-08T02:57:54.493071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{"papermill":{"duration":0.007298,"end_time":"2023-08-12T07:54:49.201004","exception":false,"start_time":"2023-08-12T07:54:49.193706","status":"completed"},"tags":[]}},{"cell_type":"code","source":"cfg = yaml.safe_load(\"\"\"\ndata:\n  resize: 512          # 256 or 512. 512 for maxvit_tiny_tf_512\n  augment: rotation    # d4 or rotation\n  augment_prob: 0.95   # probability of applying augmentation \n\nmodel:\n  encoder: tu-maxvit_tiny_tf_512.in1k  # timm-resnest26d\n  pretrained: False    # Use True! False due to internet connection\n  decoder_channels: [256, 128, 64, 32, 16]\n\nkfold:\n  k: 10\n  folds: 0  # 0,1,2,3,4\n\ntrain:\n  batch_size: 4        # I use 8 for RTX3090 (24GB) but reduce for Kaggle notebook\n  weight_decay: 1e-2\n  clip_grad_norm: 1000.0\n  num_workers: 2\n\nval:\n  per_epoch: 1         # number of val evaluation per epoch (only int >= 1)\n  th: 0.45             # threshold for k-fold out-of-fold val\n\ntest:\n  th: 0.45             # for the validation data\n\nscheduler:\n  - linear:\n      lr_start: 1e-8\n      lr_end: 8e-4\n      epoch_end: 0.5\n  - cosine:\n      lr_end: 1e-6\n      epoch_end: 40    # 50 for resnest26d\n\"\"\")","metadata":{"papermill":{"duration":0.021632,"end_time":"2023-08-12T07:54:49.230044","exception":false,"start_time":"2023-08-12T07:54:49.208412","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:57:54.495724Z","iopub.execute_input":"2023-09-08T02:57:54.496387Z","iopub.status.idle":"2023-09-08T02:57:54.516589Z","shell.execute_reply.started":"2023-09-08T02:57:54.496353Z","shell.execute_reply":"2023-09-08T02:57:54.515163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data\n\nI am reading the original npy data but this is inefficient because we have to read all T=8 data and become the bottleneck.  Locally, I converted the data to HDF5 format, in which I can load only t=4 from multidimensional array. If you are using your local machine, M.2 NVMe SSD can also change the training time for small models.","metadata":{"papermill":{"duration":0.006697,"end_time":"2023-08-12T07:54:49.243238","exception":false,"start_time":"2023-08-12T07:54:49.236541","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# Regular grid for sampling y_sym from y\n# See torch grid_sample() for the convention\ndef create_grid(nc: int, offset=0.5) -> torch.Tensor:\n    \"\"\"\n    Create xy values of nc x nc grid\n    \n    Arg:\n      nc (int): number of grid points per dimension (e.g. 512)\n      offset (float): offset in sampling in units of original 256 x 256 image\n                      offset 0 gives identity mapping\n                      Use offset 0.5 for shifted contrail label\n\n    Returns: grid (Tensor)\n      grid points in [-1, 1] for torch grid_sample()\n    \"\"\"\n    grid = np.zeros((nc, nc, 2), dtype=np.float32)\n    for ix in range(nc):\n        for iy in range(nc):\n            grid[ix, iy, 1] = -1 + 2 * (ix + 0.5) / nc + offset / 128\n            grid[ix, iy, 0] = -1 + 2 * (iy + 0.5) / nc + offset / 128\n    grid = torch.from_numpy(grid).unsqueeze(0)\n    return grid\n\n\n# Augmentation\ndef augmentation(aug: str):\n    if aug == 'd4':  # Dihedral group D4\n        return A.Compose([\n            A.RandomRotate90(p=1),\n            A.HorizontalFlip(p=0.5),\n        ])\n    elif aug == 'rotation':\n        return A.Compose([\n            A.RandomRotate90(p=1),\n            A.HorizontalFlip(p=0.5),\n            A.ShiftScaleRotate(rotate_limit=30, scale_limit=0.2, p=0.75)\n        ])\n    else:\n        raise ValueError\n\n        \n# Create list of data\ndef create_df(data_type: str):\n    assert data_type in ['train', 'validation', 'test']\n\n    filenames = glob.glob('%s/%s/*' % (di, data_type))\n    filenames.sort()\n    assert filenames\n\n    file_ids = []\n    for filename in filenames:\n        file_id = os.path.basename(filename).replace('.h5', '')\n        file_ids.append(file_id)\n\n    return pd.DataFrame({'file_id': file_ids, 'filename': filenames})\n\n\n# Load from file\ndef load_data(path: str, t=4) -> dict:\n    \"\"\"\n    Arg:\n      path (str): directory including data npy\n\n    Returns: dict\n      x (array): (3, 256, 256) for bands 11, 14, 15\n      y (array): (1, 256, 256) for ground truth label\n      annotation_mean (Optional[array])\n    \"\"\"\n    file_id = path.split('/')[-1]\n\n    # Load input images\n    bands = []\n    for k in [11, 14, 15]:\n        a = np.load('%s/band_%02d.npy' % (path, k))\n        bands.append(a)\n\n    x = np.stack(bands, axis=0)  # (C, H, W, T)\n    \n    ret = {'file_id': file_id,\n           'x': x[:, :, :, t].copy()}\n\n    # Ground truth label\n    filename = '%s/human_pixel_masks.npy' % path\n    if os.path.exists(filename):\n        y = np.load(filename)\n        y_sum = np.sum(y)\n        ret['label'] = y.reshape(1, 256, 256).astype(np.float32)\n\n    # Soft label: mean of individual annotations\n    filename = '%s/human_individual_masks.npy' % path\n    if os.path.exists(filename):\n        annot = np.load(filename)\n        annot = annot.astype(np.float32)  # (256, 256, 1, A)\n        annot = np.mean(annot, axis=3).reshape(1, 256, 256)\n        \n        ret['annotation_mean'] = annot\n            \n    return ret\n\n\n# Ash false color\ndef rescale_range(x, f_min, f_max):\n    # Rescale [f_min, f_max] to [0, 1]\n    return (x - f_min) / (f_max - f_min)\n\n\ndef ash_color(x):\n    \"\"\"\n    False color for contrail annotation\n    x (array): (3, H, W) -> (3, H, W)\n    \"\"\"\n    r = rescale_range(x[2] - x[1], -4, 2)  # band 15 - 11\n    g = rescale_range(x[1] - x[0], -4, 5)  # band 14 - 11\n    b = rescale_range(x[1], 243, 303)      # band 14\n\n    x = torch.stack([r, g, b], axis=0)\n    x = 1 - x\n\n    return x\n\n\nclass Dataset(torch.utils.data.Dataset):\n    def __init__(self, df, cfg, *, augment=False):\n        self.df = df\n        \n        # Augmentation\n        self.augment = None\n        if augment and cfg['data']['augment']:\n            self.augment = augmentation(cfg['data']['augment'])\n\n        # Resize input image\n        nc = cfg['data']['resize']\n        self.resize = nn.Identity() if nc == 256 else T.Resize(nc, antialias=False)\n\n        # Sample y_sym from y on this 0.5-pixel shifted grid\n        self.grid = create_grid(nc, offset=0.5)\n\n        # The probability to apply augmentation (0.95)\n        self.augment_prob = cfg['data']['augment_prob']\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n        r = self.df.iloc[i]\n        file_id = r['file_id']\n        \n        d = load_data(r['filename'])\n\n        x = torch.from_numpy(d['x'])\n        y = label = None\n        if 'annotation_mean' in d:\n            y = torch.from_numpy(d['annotation_mean'])\n        \n        if 'label' in d:\n            label = torch.from_numpy(d['label'])\n\n        # Create false-color image\n        x = ash_color(x)\n        x = self.resize(x)\n\n        # Create symmetric label y_sym\n        if y is not None:\n            y_sym = F.grid_sample(y.unsqueeze(0), self.grid,\n                                  mode='bilinear', padding_mode='border',\n                                  align_corners=False).squeeze(0)\n\n        # Augment\n        w_original = 1.0\n        if self.augment is not None and np.random.random() < self.augment_prob:\n            assert y is not None  # y and y_sym always exist if you want to augment\n            w_original = 0.0      # augmented\n\n            x = x.permute(1, 2, 0).numpy()  # => (H, W, C)\n            y_sym = y_sym.permute(1, 2, 0).numpy()\n\n            aug = self.augment(image=x, mask=y_sym)\n\n            x = torch.from_numpy(aug['image'].transpose(2, 0, 1))     # => (C, H, W)\n            y_sym = torch.from_numpy(aug['mask'].transpose(2, 0, 1))  # array (1, 256, 256)\n\n        # Return values\n        ret = {'file_id': file_id,\n               'x': x,\n               'w': np.float32(w_original)}  # 1 if y is not augmented\n\n        # y is target for loss (soft label)\n        if y is not None:\n            ret['y_sym'] = y_sym\n            ret['y'] = y\n\n        # label is ground truth for score\n        if label is not None:\n            ret['label'] = label\n        return ret\n\n\nclass Data:\n    def __init__(self, data_type, *, debug=False):\n        # Load filename list\n        df = create_df(data_type)\n        if debug:\n            df = df.iloc[:100]\n\n        self.df = df\n\n    def __len__(self):\n        return len(self.df)\n\n    def dataset(self, idx, cfg, augment):\n        df = self.df.iloc[idx] if idx is not None else self.df\n\n        return Dataset(df, cfg, augment=augment)\n\n    def loader(self, idx, cfg, *, augment=False, batch_size=None, shuffle=False, drop_last=False):\n        batch_size = batch_size if batch_size is not None else cfg['train']['batch_size']\n        num_workers = cfg['train']['num_workers']\n\n        ds = self.dataset(idx, cfg, augment)\n        return torch.utils.data.DataLoader(ds,\n                                           batch_size=batch_size,\n                                           num_workers=num_workers,\n                                           shuffle=shuffle,\n                                           drop_last=drop_last)","metadata":{"papermill":{"duration":0.043273,"end_time":"2023-08-12T07:54:49.294501","exception":false,"start_time":"2023-08-12T07:54:49.251228","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:57:54.522784Z","iopub.execute_input":"2023-09-08T02:57:54.525478Z","iopub.status.idle":"2023-09-08T02:57:54.574906Z","shell.execute_reply.started":"2023-09-08T02:57:54.525443Z","shell.execute_reply":"2023-09-08T02:57:54.573793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test-Time Augmentation","metadata":{"papermill":{"duration":0.006429,"end_time":"2023-08-12T07:54:49.307919","exception":false,"start_time":"2023-08-12T07:54:49.30149","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\"\"\"\n- d4 means 8 patterns of rot90 and flip\n- rot is 4 patterns of rot90\n\n- logit averages the logit y_sym\n- prob averages y_sym.sigmoid()\n\nI thought prob is slightly better, but the difference is small.\n\"\"\"\n_tta_config = {'d4prob': (8, True), 'd4logit': (8, False),\n               'rotprob': (4, True), 'rotlogit': (4, False), 'none': (1, False)}\n\ndef _tta_stack(x: torch.Tensor, n: int) -> torch.Tensor:\n    \"\"\"\n    Increate input x by n TTA patterns\n    batch_size -> n * batch_size\n    \"\"\"\n    stack = []\n    for k in range(4):\n        xa = torch.rot90(x, k, dims=[2, 3])\n        stack.append(xa)\n        if n == 8:\n            stack.append(torch.flip(xa, dims=[3, ]))\n    return torch.cat(stack, dim=0)\n\n\ndef _tta_average(y_pred: torch.Tensor, n: int, prob: bool) -> torch.Tensor:\n    \"\"\"\n    Average TTA augmented predictions\n    n * batch_size -> batch_size\n    \"\"\"\n    batch_size, nch, H, W = y_pred.shape\n    y_pred = y_pred.view(n, batch_size // n, nch, H, W)\n    batch_size = batch_size // n\n    y_avg = torch.zeros((batch_size, 1, H, W), dtype=torch.float32, device=y_pred.device)\n\n    if prob:\n        y_pred = y_pred.sigmoid()\n\n    if n == 4:\n        for k in range(4):\n            y_avg += (1 / n) * torch.rot90(y_pred[k], -k, dims=[2, 3])\n    if n == 8:\n        for k in range(4):\n            y_avg += (1 / n) * torch.rot90(y_pred[2 * k], -k, dims=[2, 3])\n            y_avg += (1 / n) * torch.rot90(y_pred[2 * k + 1].flip(dims=[3, ]), -k, dims=(2, 3))\n\n    if prob:\n        y_avg = y_avg.clamp(1e-6, 1 - 1e-6)  # logit() is NaN at 0 and 1\n        return y_avg.logit()\n    return y_avg\n\n\nclass TTA:\n    \"\"\"\n    Test-Time Augment input and average output\n    \n    Args:\n      tta_str: d4prob, d4logit, rotprob, rotlogit\n    \"\"\"\n    def __init__(self, tta_str: str):        \n        self.n, self.prob = _tta_config[tta_str]\n    \n    def __repr__(self):\n        return 'TTA(n={}, prob={})'.format(self.n, self.prob)\n\n    def stack(self, x):\n        if self.n == 1:\n            return x\n        else:\n            return _tta_stack(x, self.n)\n\n    def average(self, y):\n        if self.n == 1:\n            return y\n        else:\n            return _tta_average(y, self.n, self.prob)","metadata":{"papermill":{"duration":0.024968,"end_time":"2023-08-12T07:54:49.339599","exception":false,"start_time":"2023-08-12T07:54:49.314631","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:57:54.58015Z","iopub.execute_input":"2023-09-08T02:57:54.58245Z","iopub.status.idle":"2023-09-08T02:57:54.603273Z","shell.execute_reply.started":"2023-09-08T02:57:54.582402Z","shell.execute_reply":"2023-09-08T02:57:54.602339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model\n\nU-Net from segmentation_models_pytorch and an additional tiny conv to shift back 0.5 pixels.","metadata":{"execution":{"iopub.execute_input":"2023-08-11T09:18:19.018233Z","iopub.status.busy":"2023-08-11T09:18:19.016763Z","iopub.status.idle":"2023-08-11T09:18:19.331103Z","shell.execute_reply":"2023-08-11T09:18:19.328946Z","shell.execute_reply.started":"2023-08-11T09:18:19.01819Z"},"papermill":{"duration":0.006661,"end_time":"2023-08-12T07:54:49.352924","exception":false,"start_time":"2023-08-12T07:54:49.346263","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\ndef get_asym_conv(nc):\n    \"\"\"\n    Final tiny convolution from y_sym_pred to y_pred\n    - expected to shift 0.5 pixel back\n    - also reduce from 512 to 256 when y_sym is 512x512\n    \"\"\"\n    if nc == 256:\n        hidden_size = 9\n        asym_conv = nn.Sequential(\n            nn.Conv2d(1, hidden_size, kernel_size=(3, 3), padding=1, padding_mode='replicate'),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(hidden_size, 1, kernel_size=1),\n        )\n    elif nc == 512:\n        hidden_size = 25\n        asym_conv = nn.Sequential(\n            nn.Conv2d(1, hidden_size, kernel_size=(5, 5), padding=2, stride=2),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(hidden_size, 1, kernel_size=1),\n        )\n    else:\n        raise NotImplementedError\n\n    return asym_conv\n\n\nclass Model(nn.Module):\n    def __init__(self, cfg, pretrained=True, tta=None):\n        super().__init__()\n        name = cfg['model']['encoder']\n        pretrained = 'imagenet' if (pretrained and cfg['model']['pretrained']) else None\n        decoder_channels = cfg['model']['decoder_channels']  # (256, 128, 64, 32, 16)\n\n        self.unet = smp.Unet(name,\n                             encoder_weights=pretrained,\n                             classes=1,\n                             decoder_channels=decoder_channels,\n        )\n\n        self.asym_conv = get_asym_conv(cfg['data']['resize'])        \n        self.tta = TTA(tta) if tta is not None else None\n\n    def forward(self, x):\n        if self.tta is not None:\n            x = self.tta.stack(x)  # TTA input\n\n        y_sym = self.unet(x)\n\n        if self.tta is not None:\n            y_sym = self.tta.average(y_sym)\n            \n        # Tiny conv from y_sym_pred -> y_pred\n        y_pred = self.asym_conv(y_sym)\n\n        return y_sym, y_pred","metadata":{"papermill":{"duration":0.026678,"end_time":"2023-08-12T07:54:49.417328","exception":false,"start_time":"2023-08-12T07:54:49.39065","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:57:54.608181Z","iopub.execute_input":"2023-09-08T02:57:54.611406Z","iopub.status.idle":"2023-09-08T02:57:54.62802Z","shell.execute_reply.started":"2023-09-08T02:57:54.611371Z","shell.execute_reply":"2023-09-08T02:57:54.626947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss and evaluate functions","metadata":{"papermill":{"duration":0.006682,"end_time":"2023-08-12T07:54:49.431107","exception":false,"start_time":"2023-08-12T07:54:49.424425","status":"completed"},"tags":[]}},{"cell_type":"code","source":"\"\"\"\nLoss: criterion = BCEWithLogitsLoss()\n\nif augmented:  # w = 0\n  # Always train y_sym\n  loss = criterion(y_sym_pred, y_sym)  \nelse:  # w = 1\n  # Also train y when no augmentation (random probability 0.05)\n  # Augmentation cannot be applied to y due to lack of symmetry\n  loss = criterion(y_sym_pred, y_sym) + criterion(y_pred, y)\n\"\"\"\nclass BCELoss(_Loss):\n    def __init__(self):\n        super().__init__()\n        self.criterion = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self,\n                y_sym_pred: torch.Tensor, y_sym: torch.Tensor,\n                y_pred: torch.Tensor, y: torch.Tensor, w) -> torch.Tensor:\n        loss_sym = self.criterion(y_sym_pred, y_sym).mean(dim=(1, 2, 3))\n        loss_original = self.criterion(y_pred, y).mean(dim=(1, 2, 3))\n        loss = loss_sym + w * loss_original\n\n        return loss.mean()  # mean of batch\n\n\ndef evaluate(model, loader_val, *, th=0.45):\n    \"\"\"\n    Compute validation loss and score\n    \"\"\"\n    tb = time.time()\n\n    was_training = model.training\n    model.eval()\n\n    n_sum = 0\n    loss_sum = 0.0\n    dice_sum = 0.0\n    tp = 0\n    positives_pred = 0\n    positives_true = 0\n\n    for d in loader_val:\n        x = d['x'].to(device)      # input image (3, H, W); may be upscaled\n        if 'y' in d:\n            y = d['y'].to(device)  # soft label: (1, H, W)\n            y_sym = d['y_sym'].to(device)\n        else:\n            y = None\n\n        label = d['label'].to(device)  # ground-truth segmentation mask (always 256 x 256)\n        batch_size = len(x)\n\n        # Predict\n        with torch.no_grad():\n            y_sym_pred, y_pred = model(x)  # (batch_size, 1, H, W)\n\n        # Compute validation loss\n        if y is not None:\n            # No augmentation in validation but rescale criteion(y_pred, y) like training time\n            w = 1 - augment_fraction\n            loss = criterion(y_sym_pred, y_sym, y_pred, y, w)\n            loss_dice = dice(y_pred, y)\n\n            n_sum += batch_size\n            loss_sum += loss.item() * batch_size\n            dice_sum += loss_dice.item() * batch_size\n\n        # Compute score\n        y_pred = y_pred.sigmoid() > th\n        tp += (y_pred * label).sum().item()\n        positives_pred += y_pred.sum().item()\n        positives_true += label.sum().item()\n\n    global_dice = 2 * tp / (positives_pred + positives_true)\n\n    model.train(was_training)\n\n    dt = time.time() - tb\n    ret = {'score': global_dice,\n           'dt': dt}\n\n    if y is not None:\n        ret['loss'] = loss_sum / n_sum\n        ret['dice'] = 1 - dice_sum / n_sum\n    return ret","metadata":{"papermill":{"duration":0.02421,"end_time":"2023-08-12T07:54:49.462096","exception":false,"start_time":"2023-08-12T07:54:49.437886","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:57:54.63302Z","iopub.execute_input":"2023-09-08T02:57:54.635894Z","iopub.status.idle":"2023-09-08T02:57:54.656743Z","shell.execute_reply.started":"2023-09-08T02:57:54.635861Z","shell.execute_reply":"2023-09-08T02:57:54.655797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training\n\nOnly use 100 data because real training takes too much time on Kaggle notebook.","metadata":{"papermill":{"duration":0.006723,"end_time":"2023-08-12T07:54:49.475735","exception":false,"start_time":"2023-08-12T07:54:49.469012","status":"completed"},"tags":[]}},{"cell_type":"code","source":"debug = True  # Set it to False for real training, but takes too long on Kaggle notebook\n\n# Data\ndata = Data('train', debug=debug)\ndata_test = Data('validation', debug=debug)\nprint('Data', len(data), len(data_test))\n\n# Kfold\nnfolds = cfg['kfold']['k']\nfolds = util.as_list(cfg['kfold']['folds'])  # list[int]\nkfold = KFold(n_splits=nfolds, shuffle=True, random_state=42)\nprint('folds', folds, '/', nfolds)\n\n# Loss\ncriterion = BCELoss()\ndice = smp.losses.DiceLoss('binary', from_logits=True)  # just as metric in validation\n\n# Training parameters\nweight_decay = float(cfg['train']['weight_decay'])\nclip_grad_norm = float(cfg['train']['clip_grad_norm'])\n\naugment_fraction = cfg['data']['augment_prob']\nsteps_per_epoch = cfg['val']['per_epoch']\nth_val = cfg['val']['th']\nth_test = cfg['test']['th']\n\n# Train loop\nlog = {}\nepochs_log = []\nlosses_train = []\nlosses_val = []\nfinal_scores = []\n\nn_sum = 0\nloss_sum = 0.0\nlrs = []\n\nprint('Pretrained: ', cfg['model']['pretrained'])\n\ntb_global = time.time()\nfor ifold, (idx_train, idx_val) in enumerate(kfold.split(data.df)):\n    if ifold not in folds:\n        continue\n\n    # Data\n    loader_train = data.loader(idx_train, cfg, augment=True, drop_last=True, shuffle=True)\n    loader_val = data.loader(idx_val, cfg)\n    loader_test = data_test.loader(None, cfg)\n\n    nbatch = len(loader_train)\n\n    # Model\n    model = Model(cfg, pretrained=True)\n    model.to(device)\n    model.train()\n\n    # Optimizer\n    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4,\n                                  weight_decay=weight_decay)\n\n    scheduler = Scheduler(optimizer, cfg['scheduler'])\n    epochs = 4 if debug else len(scheduler)\n    print('%d epochs' % epochs)\n\n    tb = time.time()\n    dt_val = 0.0\n    print('KFold %d/%d' % (ifold, nfolds))\n    print('Epoch        loss          dice  score         lr       time')\n    for iepoch in range(epochs):\n        # This is for validating multiple time per epoch\n        icheck = [nbatch * (i + 1) // steps_per_epoch - 1 for i in range(steps_per_epoch)]\n\n        for ibatch, d in enumerate(loader_train):\n            x = d['x'].to(device)          # input image\n            y = d['y'].to(device)          # segmentation label\n            y_sym = d['y_sym'].to(device)  # symmetric label\n            w = d['w'].to(device)          # w=0 if augmented, 1 if x and y are original \n            batch_size = len(x)\n\n            optimizer.zero_grad()\n\n            # Predict\n            y_sym_pred, y_pred = model(x)  # (batch_size, 1, 256, 256)\n            loss = criterion(y_sym_pred, y_sym, y_pred, y, w)\n\n            # Backpropagate\n            loss.backward()\n\n            n_sum += batch_size\n            loss_sum += batch_size * loss.item()\n\n            nn.utils.clip_grad_norm_(model.parameters(), clip_grad_norm)\n            optimizer.step()\n\n            ep = iepoch + (ibatch + 1) / nbatch\n            lr = optimizer.param_groups[0]['lr']\n            lrs.append((ep, lr))\n\n            # Validation\n            if ibatch == icheck[0]:\n                icheck.pop(0)\n\n                epochs_log.append(iepoch + (ibatch + 1) / nbatch)\n                loss_train = loss_sum / n_sum\n                losses_train.append(loss_train)\n\n                val = evaluate(model, loader_val, th=th_val)\n                test = evaluate(model, loader_test, th=th_test)\n\n                losses_val.append(val['loss'])\n                dt = time.time() - tb\n                dt_val += val['dt'] + test['dt']\n\n                print('Epoch %5.2f %6.3f %6.3f  %.3f %.3f %.3f  %5.1e %5.1f %5.1f min' % (ep,\n                      10 * loss_train, 10 * val['loss'],\n                      val['dice'], val['score'], test['score'],\n                      lr, dt_val / 60, dt / 60))\n\n                # Reset train loss\n                n_sum = 0\n                loss_sum = 0.0\n\n            scheduler.step(ep)\n\n    # Save model\n    #model.eval()\n    #ofilename = 'model%d.pytorch' % ifold\n    #torch.save(model.state_dict(), ofilename)\n    #print(ofilename, 'written')\n    del model\n\n# Kfolds done\ndt = time.time() - tb_global\nprint('Total time: %.2f min' % (dt / 60))","metadata":{"papermill":{"duration":78.129004,"end_time":"2023-08-12T07:56:07.611574","exception":false,"start_time":"2023-08-12T07:54:49.48257","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:57:54.661667Z","iopub.execute_input":"2023-09-08T02:57:54.66456Z","iopub.status.idle":"2023-09-08T02:59:08.144701Z","shell.execute_reply.started":"2023-09-08T02:57:54.664526Z","shell.execute_reply":"2023-09-08T02:59:08.143442Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The score is ~0 because this is a debug run with 100 data, and also pretrained=False due to internet off.\n\nWith `debug=False`, it should train 40 epochs:\n\n```text\nEpoch 10.00  0.208  0.203  0.405 0.668 0.645  6.9e-04  12.4 148.8 min\nEpoch 20.00  0.196  0.199  0.414 0.671 0.670  4.1e-04  24.6 300.4 min\nEpoch 30.00  0.187  0.192  0.435 0.690 0.686  1.2e-04  36.9 449.2 min\nEpoch 40.00  0.184  0.191  0.441 0.696 0.687  1.0e-06  49.1 599.5 min\nTotal time: 599.50 min (with RTX3090)\n```\n\nThe validation score varies due to randomness in the model initial conditions and training:\n0.688 ± 0.004","metadata":{"papermill":{"duration":0.007534,"end_time":"2023-08-12T07:56:07.626536","exception":false,"start_time":"2023-08-12T07:56:07.619002","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"# Inference","metadata":{"papermill":{"duration":0.007091,"end_time":"2023-08-12T07:56:07.640929","exception":false,"start_time":"2023-08-12T07:56:07.633838","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def initial_preds(file_ids):\n    \"\"\"\n    Initialize prediction with zeros\n    \"\"\"\n    preds = []\n    for file_id in file_ids:\n        pred = {'file_id': file_id,\n                'y_pred': np.zeros((1, 256, 256), dtype=np.float32)}\n        preds.append(pred)\n\n    return preds\n\n\ndef predict1(model, loader, w: float, device, preds: list):\n    \"\"\"\n    Predict with one model and add to preds\n    \"\"\"\n    i = 0\n    for d in tqdm(loader):\n        x = d['x'].to(device)  # input image\n\n        with torch.no_grad():\n            _, y_pred = model(x)\n\n        y_pred = y_pred.sigmoid().cpu().numpy()\n\n        for y_pred1 in y_pred:\n            preds[i]['y_pred'] += w * y_pred1\n            i += 1","metadata":{"papermill":{"duration":0.01947,"end_time":"2023-08-12T07:56:07.667684","exception":false,"start_time":"2023-08-12T07:56:07.648214","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:59:08.146437Z","iopub.execute_input":"2023-09-08T02:59:08.147567Z","iopub.status.idle":"2023-09-08T02:59:08.156426Z","shell.execute_reply.started":"2023-09-08T02:59:08.147528Z","shell.execute_reply":"2023-09-08T02:59:08.155252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_type = 'test'  # 'validation' gives validation score with TTA, \nfolds = [0, ]\nw = 1 / len(folds)\n\ndata = Data(data_type, debug=False)\nloader_test = data.loader(None, cfg, augment=False, batch_size=1)  # TTA increases batch_size 1 -> 8\n\n# Initalize prediction with 0\npreds = initial_preds(data.df.file_id.values)\n    \nfor ifold in folds:\n    model = Model(cfg, pretrained=False, tta='d4prob')\n\n    model_filename = '/kaggle/input/contrails-base-model/weights/maxvit_tiny_simple/model%d.pytorch' % ifold\n    model.load_state_dict(torch.load(model_filename))\n    model.to(device)\n    model.eval()\n\n    print('Load %s %.4f' % (model_filename, w))\n    \n    predict1(model, loader_test, w, device, preds)  # Add prediction","metadata":{"papermill":{"duration":3.269898,"end_time":"2023-08-12T07:56:10.944743","exception":false,"start_time":"2023-08-12T07:56:07.674845","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:59:08.160198Z","iopub.execute_input":"2023-09-08T02:59:08.160534Z","iopub.status.idle":"2023-09-08T02:59:12.324493Z","shell.execute_reply.started":"2023-09-08T02:59:08.160504Z","shell.execute_reply":"2023-09-08T02:59:12.323177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"th = 0.45\nsubmit = write_submission(preds, th, 'submission.csv')","metadata":{"papermill":{"duration":0.028088,"end_time":"2023-08-12T07:56:10.981108","exception":false,"start_time":"2023-08-12T07:56:10.95302","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:59:12.326741Z","iopub.execute_input":"2023-09-08T02:59:12.327054Z","iopub.status.idle":"2023-09-08T02:59:12.344403Z","shell.execute_reply.started":"2023-09-08T02:59:12.327026Z","shell.execute_reply":"2023-09-08T02:59:12.343425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! tail submission.csv\n\nif data_type == 'validation':\n    ! python3 src/score.py submission.csv\n    # => validation score 0.691894 with TTA, 0.686770 without TTA","metadata":{"papermill":{"duration":0.016169,"end_time":"2023-08-12T07:56:11.005238","exception":false,"start_time":"2023-08-12T07:56:10.989069","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-09-08T02:59:12.345803Z","iopub.execute_input":"2023-09-08T02:59:12.34634Z","iopub.status.idle":"2023-09-08T02:59:13.348148Z","shell.execute_reply.started":"2023-09-08T02:59:12.346306Z","shell.execute_reply":"2023-09-08T02:59:13.346885Z"},"trusted":true},"execution_count":null,"outputs":[]}],"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"}}