{"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":"# U-Net model with symmetric label and rotation augmentation\n\n\n### Background\nThis is a semantic segmentation task, marking contrails in satellite images. \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### Features\n\nThis model can use many encoders in timm, not limited to the ones in segmention_models_pytorch, including MaxViT which was very strong for this competition\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\n\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.","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"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":{"execution":{"iopub.status.busy":"2023-08-23T01:09:02.660405Z","iopub.execute_input":"2023-08-23T01:09:02.660782Z","iopub.status.idle":"2023-08-23T01:09:15.914783Z","shell.execute_reply.started":"2023-08-23T01:09:02.660751Z","shell.execute_reply":"2023-08-23T01:09:15.913419Z"},"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\n\nimport timm\nimport segmentation_models_pytorch as smp\nfrom segmentation_models_pytorch.base.initialization import initialize_decoder\nfrom segmentation_models_pytorch.base import modules as md\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":{"execution":{"iopub.status.busy":"2023-08-23T01:09:15.917599Z","iopub.execute_input":"2023-08-23T01:09:15.918238Z","iopub.status.idle":"2023-08-23T01:09:23.458437Z","shell.execute_reply.started":"2023-08-23T01:09:15.9182Z","shell.execute_reply":"2023-08-23T01:09:23.457436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"cell_type":"code","source":"cfg = yaml.safe_load(\"\"\"\ndata:\n  resize: 512          # 256 or 512 for maxvit_tiny_tf_512\n  augment: rotation    # d4 or rotation\n  augment_prob: 0.95   # probability of applying augmentation \n\nmodel:\n  # resnest26d, inception_v4 maxvit_tiny_tf_512.in1k,timm-efficientnet-b7 timm-efficientnet-b6\n  encoder: maxvit_tiny_tf_512.in1k  # I also use timm-efficientnet-b7\n  pretrained: False    # Use True! False due to internet connection\n  decoder_channels: [256, 128, 64, 32, 16]\n  dropout: 0.0\n\nkfold:\n  k: 10\n  folds: 0  # 0,1,2,3,4\n\ntrain:\n  weight_decay: 1e-2\n  batch_size: 2        # I use 16 for A100 (40GB) but reduce for Kaggle notebook\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":{"execution":{"iopub.status.busy":"2023-08-23T01:09:23.459973Z","iopub.execute_input":"2023-08-23T01:09:23.460379Z","iopub.status.idle":"2023-08-23T01:09:23.472209Z","shell.execute_reply.started":"2023-08-23T01:09:23.460295Z","shell.execute_reply":"2023-08-23T01:09:23.470632Z"},"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 change the training time.","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2023-08-23T01:09:23.475204Z","iopub.execute_input":"2023-08-23T01:09:23.475843Z","iopub.status.idle":"2023-08-23T01:09:23.509147Z","shell.execute_reply.started":"2023-08-23T01:09:23.475808Z","shell.execute_reply":"2023-08-23T01:09:23.508142Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test-Time Augmentation","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2023-08-23T01:09:23.510771Z","iopub.execute_input":"2023-08-23T01:09:23.51121Z","iopub.status.idle":"2023-08-23T01:09:23.530168Z","shell.execute_reply.started":"2023-08-23T01:09:23.51115Z","shell.execute_reply":"2023-08-23T01:09:23.528908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model\n\nU-Net using timm for encoder and segmentation_models_pytorch for decoder","metadata":{"execution":{"iopub.status.busy":"2023-08-11T09:18:19.016763Z","iopub.execute_input":"2023-08-11T09:18:19.018233Z","iopub.status.idle":"2023-08-11T09:18:19.331103Z","shell.execute_reply.started":"2023-08-11T09:18:19.01819Z","shell.execute_reply":"2023-08-11T09:18:19.328946Z"}}},{"cell_type":"code","source":"\"\"\"\nU-Net decoder from Segmentation Models PyTorch\nhttps://github.com/qubvel/segmentation_models.pytorch\n\"\"\"\nclass DecoderBlock(nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        skip_channels,\n        out_channels,\n        use_batchnorm=True,\n        dropout=0,\n    ):\n        super().__init__()\n\n        conv_in_channels = in_channels + skip_channels\n\n        # Convolve input embedding and upscaled embedding\n        self.conv1 = md.Conv2dReLU(\n            conv_in_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n\n        self.conv2 = md.Conv2dReLU(\n            out_channels,\n            out_channels,\n            kernel_size=3,\n            padding=1,\n            use_batchnorm=use_batchnorm,\n        )\n\n        self.dropout_skip = nn.Dropout(p=dropout)\n\n    def forward(self, x, skip=None):\n        x = F.interpolate(x, scale_factor=2, mode='nearest')\n        if skip is not None:\n            skip = self.dropout_skip(skip)\n            x = torch.cat([x, skip], dim=1)\n\n        x = self.conv1(x)\n        x = self.conv2(x)\n\n        return x\n\n\nclass UnetDecoder(nn.Module):\n    def __init__(\n        self,\n        encoder_channels,\n        decoder_channels,\n        use_batchnorm=True,\n        dropout=0,\n    ):\n        super().__init__()\n\n        encoder_channels = encoder_channels[::-1]\n\n        # Computing blocks input and output channels\n        head_channels = encoder_channels[0]\n        in_channels = [head_channels] + list(decoder_channels[:-1])\n        skip_channels = list(encoder_channels[1:]) + [0]\n        out_channels = decoder_channels\n\n        self.center = nn.Identity()\n\n        # Combine decoder keyword arguments\n        blocks = [\n            DecoderBlock(in_ch, skip_ch, out_ch, use_batchnorm=use_batchnorm, dropout=dropout)\n            for in_ch, skip_ch, out_ch in zip(in_channels, skip_channels, out_channels)\n        ]\n        self.blocks = nn.ModuleList(blocks)\n\n    def forward(self, features):\n        features = features[::-1]  # reverse channels to start from head of encoder\n\n        head = features[0]\n        skips = features[1:]\n\n        x = self.center(head)\n        for i, decoder_block in enumerate(self.blocks):\n            skip = skips[i] if i < len(skips) else None\n            x = decoder_block(x, skip)\n\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-08-23T01:09:23.531892Z","iopub.execute_input":"2023-08-23T01:09:23.53242Z","iopub.status.idle":"2023-08-23T01:09:23.546876Z","shell.execute_reply.started":"2023-08-23T01:09:23.532387Z","shell.execute_reply":"2023-08-23T01:09:23.545823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _check_reduction(reduction_factors):\n    \"\"\"\n    Assume spatial dimensions of the features decrease by factors of two.\n    For example, convnext start with stride=4 cannot be used in my code.\n    \"\"\"\n    r_prev = 1\n    for r in reduction_factors:\n        if r / r_prev != 2:\n            raise AssertionError('Reduction assumed to increase by 2: {}'.format(reduction_factors))\n        r_prev = r\n\n\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    # The main U-Net model\n    # See also TimmUniversalEncoder in Segmentation Models PyTorch\n    def __init__(self, cfg, pretrained=True, tta=None):\n        super().__init__()\n        name = cfg['model']['encoder']\n        dropout = cfg['model']['dropout']\n        pretrained = pretrained and cfg['model']['pretrained']\n\n        self.encoder = timm.create_model(name, features_only=True, pretrained=pretrained)\n        encoder_channels = self.encoder.feature_info.channels()\n\n        _check_reduction(self.encoder.feature_info.reduction())\n\n        decoder_channels = cfg['model']['decoder_channels']  # (256, 128, 64, 32, 16)\n        print('Encoder channels:', name, encoder_channels)\n        print('Decoder channels:', decoder_channels)\n\n        assert len(encoder_channels) == len(decoder_channels)\n\n        self.decoder = UnetDecoder(\n            encoder_channels=encoder_channels,\n            decoder_channels=decoder_channels,\n            dropout=dropout,\n        )\n\n        self.segmentation_head = smp.base.SegmentationHead(\n            in_channels=decoder_channels[-1],\n            out_channels=1, activation=None, kernel_size=3,\n        )\n\n        initialize_decoder(self.decoder)\n\n        self.asym_conv = get_asym_conv(cfg['data']['resize'])\n        \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_pred = unet(x)\n        features = self.encoder(x)\n        decoder_output = self.decoder(features)\n        y_sym = self.segmentation_head(decoder_output)\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":{"execution":{"iopub.status.busy":"2023-08-23T01:09:23.548393Z","iopub.execute_input":"2023-08-23T01:09:23.548764Z","iopub.status.idle":"2023-08-23T01:09:23.565952Z","shell.execute_reply.started":"2023-08-23T01:09:23.548732Z","shell.execute_reply":"2023-08-23T01:09:23.564985Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss and evaluate functions","metadata":{}},{"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\"\"\"\n# class 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\nclass CombinedLoss(_Loss):\n    def __init__(self):\n        super().__init__()\n        self.bce_criterion = nn.BCEWithLogitsLoss(reduction='none')\n\n    def dice_loss(self, pred, target, smooth=1e-7):\n        intersection = (pred * target).sum()\n        return 1 - (2. * intersection + smooth) / (pred.sum() + target.sum() + smooth)\n\n    def forward(self, y_sym_pred, y_sym, y_pred, y, w, alpha=0.5):\n        bce_loss_sym = self.bce_criterion(y_sym_pred, y_sym).mean(dim=(1, 2, 3))\n        bce_loss_original = self.bce_criterion(y_pred, y).mean(dim=(1, 2, 3))\n        \n        dice_loss_sym = self.dice_loss(y_sym_pred.sigmoid(), y_sym)\n        dice_loss_original = self.dice_loss(y_pred.sigmoid(), y)\n        \n        combined_loss_sym = alpha * bce_loss_sym + (1 - alpha) * dice_loss_sym\n        combined_loss_original = alpha * bce_loss_original + (1 - alpha) * dice_loss_original\n        \n        loss = combined_loss_sym + w * combined_loss_original\n\n        return loss.mean()  # mean of batch\n\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":{"execution":{"iopub.status.busy":"2023-08-23T01:09:23.56769Z","iopub.execute_input":"2023-08-23T01:09:23.568053Z","iopub.status.idle":"2023-08-23T01:09:23.58636Z","shell.execute_reply.started":"2023-08-23T01:09:23.568022Z","shell.execute_reply":"2023-08-23T01:09:23.585256Z"},"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":{}},{"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\n#criterion = BCELoss()\ncriterion =  CombinedLoss()\ndice = smp.losses.DiceLoss('binary', from_logits=True)  # just as metric in validation\n\n# Training parameters\nweight_decay = float(cfg['train']['weight_decay'])\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\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_value_(model.parameters(), 1000.0)\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":{"execution":{"iopub.status.busy":"2023-08-23T01:09:23.589028Z","iopub.execute_input":"2023-08-23T01:09:23.589882Z","iopub.status.idle":"2023-08-23T01:10:51.43138Z","shell.execute_reply.started":"2023-08-23T01:09:23.589858Z","shell.execute_reply":"2023-08-23T01:10:51.430156Z"},"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.","metadata":{}},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2023-08-23T01:10:51.435421Z","iopub.execute_input":"2023-08-23T01:10:51.435718Z","iopub.status.idle":"2023-08-23T01:10:51.444718Z","shell.execute_reply.started":"2023-08-23T01:10:51.435691Z","shell.execute_reply":"2023-08-23T01:10:51.442742Z"},"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_wd/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":{"execution":{"iopub.status.busy":"2023-08-23T01:10:51.446002Z","iopub.execute_input":"2023-08-23T01:10:51.446915Z","iopub.status.idle":"2023-08-23T01:10:56.729877Z","shell.execute_reply.started":"2023-08-23T01:10:51.446883Z","shell.execute_reply":"2023-08-23T01:10:56.728624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"th = 0.45\nsubmit = write_submission(preds, th, 'submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-08-23T01:10:56.732461Z","iopub.execute_input":"2023-08-23T01:10:56.732857Z","iopub.status.idle":"2023-08-23T01:10:56.750982Z","shell.execute_reply.started":"2023-08-23T01:10:56.732821Z","shell.execute_reply":"2023-08-23T01:10:56.750084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if data_type == 'validation':\n    ! python3 src/score.py submission.csv\n    # => validation score 0.695404 with TTA, 0.691025 without TTA","metadata":{"execution":{"iopub.status.busy":"2023-08-23T01:10:56.752049Z","iopub.execute_input":"2023-08-23T01:10:56.752301Z","iopub.status.idle":"2023-08-23T01:10:56.757953Z","shell.execute_reply.started":"2023-08-23T01:10:56.752279Z","shell.execute_reply":"2023-08-23T01:10:56.757047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! tail submission.csv","metadata":{"execution":{"iopub.status.busy":"2023-08-23T01:10:56.759157Z","iopub.execute_input":"2023-08-23T01:10:56.759781Z","iopub.status.idle":"2023-08-23T01:10:57.911977Z","shell.execute_reply.started":"2023-08-23T01:10:56.759749Z","shell.execute_reply":"2023-08-23T01:10:57.910795Z"},"trusted":true},"execution_count":null,"outputs":[]}]}