{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":39763,"databundleVersionId":11756775},{"sourceType":"datasetVersion","sourceId":11801643,"datasetId":7377931,"databundleVersionId":12290899},{"sourceType":"datasetVersion","sourceId":11569755,"datasetId":7253661,"databundleVersionId":12026132}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"RUN_TRAIN = True\nRUN_VALID = True\nRUN_TEST  = True\n\nimport torch\nif not torch.cuda.is_available() or torch.cuda.device_count() < 2:\n    raise RuntimeError(\"Requires >= 2 GPUs with CUDA enabled.\")\n\ntry: \n    import monai\nexcept: \n    !pip install --no-deps monai -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:23:31.182136Z","iopub.execute_input":"2025-05-13T22:23:31.182831Z","iopub.status.idle":"2025-05-13T22:23:38.369160Z","shell.execute_reply.started":"2025-05-13T22:23:31.182797Z","shell.execute_reply":"2025-05-13T22:23:38.368134Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# HGNet-V2 Starter Notebook\n\nThis notebook builds on Egor Trushin's great starter notebook [here](https://www.kaggle.com/code/egortrushin/gwi-unet-with-float16-dataset), thanks for sharing. \n\nThe main purpose of this notebook is to show how to use 2 GPUs during model training, maximizing our weekly GPU quota in the Kaggle environment. \n\nIn addition, I provide 3x pretrained model checkpoints that were trained for 150 epochs using this setup. Each model achieved a validation MAE of ~60.\n\nOther additions:                                                                 \n- Flip augmentation\n- Dataset preprocessing\n- EMA (Exponential moving average)\n- Pretrained encoder\n- Monai usample blocks","metadata":{}},{"cell_type":"code","source":"%%writefile _cfg.py\n\nfrom types import SimpleNamespace\nimport torch\n\ncfg= SimpleNamespace()\ncfg.device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ncfg.data_dir = \"/kaggle/input/openfwi-preprocessed-72x72/openfwi_72x72/\"\ncfg.local_rank = 0\ncfg.seed = 123\ncfg.subsample = None\n\ncfg.backbone = \"hgnetv2_b2.ssld_stage2_ft_in1k\"\ncfg.ema = True\ncfg.ema_decay = 0.99\n\ncfg.epochs = 10\ncfg.batch_size = 512\ncfg.batch_size_val = 128\n\ncfg.early_stopping = {\"patience\": 3, \"streak\": 0}\ncfg.logging_steps = 100","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:23:38.370663Z","iopub.execute_input":"2025-05-13T22:23:38.371254Z","iopub.status.idle":"2025-05-13T22:23:38.377203Z","shell.execute_reply.started":"2025-05-13T22:23:38.371234Z","shell.execute_reply":"2025-05-13T22:23:38.376504Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Preprocess\n\nHere we reduce the size of each input to (72,72), and then save as fp16. This reduces the size of each input to roughly 4% of the original size.\n\nThis has already been done for every datapoint in the OpenFWI dataset and can be found [here](https://www.kaggle.com/datasets/brendanartley/openfwi-preprocessed-72x72).","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport numpy as np\n\ndef _preprocess(x):\n    x = F.interpolate(x, size=(70, 70), mode='area')\n    x = F.pad(x, (1,1,1,1), mode='replicate')\n    return x\n\ndef _helper(x, ):\n    before_shape = x.shape\n    before_mem = x.nbytes / 1e6\n    x = torch.from_numpy(x).float()\n\n    # Interpolate and pad\n    x = _preprocess(x)\n    x = x.cpu().numpy().astype(np.float16)\n\n    after_mem = x.nbytes / 1e6\n    percent = 100 - 100 * (before_mem - after_mem) / before_mem if before_mem else 0\n\n    # Log\n    print(\"Shape Change\")\n    print(\"  {} -> {}\".format(before_shape, x.shape))\n    print()\n    print(\"Memory Usage\")\n    print(\"  {:.1f} MB -> {:.1f} MB\".format(before_mem, after_mem))\n    print(\"  ({:.1f}% of original size)\".format(percent))\n    return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:23:38.377929Z","iopub.execute_input":"2025-05-13T22:23:38.378129Z","iopub.status.idle":"2025-05-13T22:23:38.391775Z","shell.execute_reply.started":"2025-05-13T22:23:38.378114Z","shell.execute_reply":"2025-05-13T22:23:38.391060Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Preprocess\nx= np.load(\"/kaggle/input/waveform-inversion/train_samples/CurveFault_A/seis2_1_0.npy\")\nx = _helper(x)\n\n# Sanity check: Confirm preprocessing matches w/ Dataset\nz= np.load(\"/kaggle/input/openfwi-preprocessed-72x72/openfwi_72x72/CurveFault_A/seis2_1_0.npy\")\nassert np.all(z == x)\n\ndel x, z","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:23:38.393386Z","iopub.execute_input":"2025-05-13T22:23:38.393611Z","iopub.status.idle":"2025-05-13T22:23:44.247944Z","shell.execute_reply.started":"2025-05-13T22:23:38.393595Z","shell.execute_reply":"2025-05-13T22:23:44.247400Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Dataset\n\nHere, we introduce a flip augmentation. \n\nUnlike a normal horizontal flip, we have to reverse the source and receiver dimensions. To match this, we reverse the width dimension of the label as well.\n\nWe use this flip as TTA (test-time augmentation) during inference.","metadata":{}},{"cell_type":"code","source":"%%writefile _dataset.py\n\nimport os\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\n\nimport torch\n\nclass CustomDataset(torch.utils.data.Dataset):\n    def __init__(\n        self, \n        cfg,\n        mode = \"train\", \n    ):\n        self.cfg = cfg\n        self.mode = mode\n        \n        self.data, self.labels, self.records = self.load_metadata()\n\n    def load_metadata(self, ):\n\n        # Select rows\n        df= pd.read_csv(\"/kaggle/input/openfwi-preprocessed-72x72/folds.csv\")\n        if self.cfg.subsample is not None:\n            df= df.groupby([\"dataset\", \"fold\"]).head(self.cfg.subsample)\n\n        if self.mode == \"train\":\n            df= df[df[\"fold\"] != 0]\n        else:\n            df= df[df[\"fold\"] == 0]\n\n        \n        data = []\n        labels = []\n        records = []\n        mmap_mode = \"r\" if self.mode == \"train\" else None\n\n        for idx, row in tqdm(df.iterrows(), total=len(df), disable=self.cfg.local_rank != 0):\n            row= row.to_dict()\n\n            # Load\n            farr= os.path.join(self.cfg.data_dir, row[\"data_fpath\"])\n            flbl= os.path.join(self.cfg.data_dir, row[\"label_fpath\"])\n            arr= np.load(farr, mmap_mode=mmap_mode)\n            lbl= np.load(flbl, mmap_mode=mmap_mode)\n\n            # Append\n            data.append(arr)\n            labels.append(lbl)\n            records.append(row[\"dataset\"])\n\n        return data, labels, records\n\n    def __getitem__(self, idx):\n        row_idx= idx // 500\n        col_idx= idx % 500\n\n        d= self.records[row_idx]\n        x= self.data[row_idx][col_idx, ...]\n        y= self.labels[row_idx][col_idx, ...]\n\n        # Augs \n        if self.mode == \"train\":\n            \n            # Temporal flip\n            if np.random.random() < 0.5:\n                x= x[::-1, :, ::-1]\n                y= y[..., ::-1]\n\n        x= x.copy()\n        y= y.copy()\n        \n        return x, y\n\n    def __len__(self, ):\n        return len(self.records) * 500","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:23:44.248624Z","iopub.execute_input":"2025-05-13T22:23:44.248844Z","iopub.status.idle":"2025-05-13T22:23:44.254226Z","shell.execute_reply.started":"2025-05-13T22:23:44.248827Z","shell.execute_reply":"2025-05-13T22:23:44.253411Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model\n\nThis model includes several modifications beyond a standard U-Net architecture.\n\n\n### Encoder\n\nThe model uses the `HgnetV2` backbone from timm as the encoder. We have to make a few modifications for this to work with the Unet. See more info on this backbone [here](https://huggingface.co/timm/hgnetv2_b2.ssld_stage1_in22k_in1k).\n\n\nFirst, we reduce the stride of the stem convolution from (2,2) to (1,1). This increases the size of the feature maps in the backbone. Second, we reduce the stride of the downsample convolution in the deepest block from (2,2) to (1,1). We do this so that upsampling in the decoder can be done without padding.\n\n```python\n# Original feature map\n[torch.Size([18, 18]), torch.Size([9, 9]), torch.Size([5, 5]), torch.Size([3, 3])]\n\n# Updated stem conv\n[torch.Size([36, 36]), torch.Size([18, 18]), torch.Size([9, 9]), torch.Size([5, 5])]\n\n# Updated downsample conv\n[torch.Size([36, 36]), torch.Size([18, 18]), torch.Size([9, 9]), torch.Size([9, 9])]\n```\n\n### Decoder\n\nThe decoder has a few modifications as well. \n\nWe remove all BatchNorm2d layers and add intermediate convolutions to the skip connections. I found that removing the normalization layers increased the convergence speed, and the intermediate convolutions improved the model's predictiveness.\n\n---\n\n### EMA\n\nWe also add an EMA (exponential moving average) class. This is a common strategy used to increase the stability of validation performance between steps/epochs. \n\nThis implementation is from Tereka [here](https://www.kaggle.com/competitions/blood-vessel-segmentation/discussion/475080#2641635).","metadata":{}},{"cell_type":"code","source":"%%writefile _model.py\n\nfrom copy import deepcopy\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport timm\n\nfrom monai.networks.blocks import UpSample, SubpixelUpsample\n\n####################\n## EMA + Ensemble ##\n####################\n\nclass ModelEMA(nn.Module):\n    def __init__(self, model, decay=0.99, device=None):\n        super().__init__()\n        self.module = deepcopy(model)\n        self.module.eval()\n        self.decay = decay\n        self.device = device\n        if self.device is not None:\n            self.module.to(device=device)\n\n    def _update(self, model, update_fn):\n        with torch.no_grad():\n            for ema_v, model_v in zip(self.module.state_dict().values(), model.state_dict().values()):\n                if self.device is not None:\n                    model_v = model_v.to(device=self.device)\n                ema_v.copy_(update_fn(ema_v, model_v))\n\n    def update(self, model):\n        self._update(model, update_fn=lambda e, m: self.decay * e + (1. - self.decay) * m)\n\n    def set(self, model):\n        self._update(model, update_fn=lambda e, m: m)\n\n\nclass EnsembleModel(nn.Module):\n    def __init__(self, models):\n        super().__init__()\n        self.models = nn.ModuleList(models).eval()\n\n    def forward(self, x):\n        output = None\n        \n        for m in self.models:\n            logits= m(x)\n            \n            if output is None:\n                output = logits\n            else:\n                output += logits\n                \n        output /= len(self.models)\n        return output\n        \n\n###################\n## HGNet-V2 Unet ##\n###################\n\nclass ConvBnAct2d(nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        out_channels,\n        kernel_size,\n        padding: int = 0,\n        stride: int = 1,\n        norm_layer: nn.Module = nn.Identity,\n        act_layer: nn.Module = nn.ReLU,\n    ):\n        super().__init__()\n\n        self.conv= nn.Conv2d(\n            in_channels, \n            out_channels,\n            kernel_size,\n            stride=stride, \n            padding=padding, \n            bias=False,\n        )\n        self.norm = norm_layer(out_channels) if norm_layer != nn.Identity else nn.Identity()\n        self.act= act_layer(inplace=True)\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.norm(x)\n        x = self.act(x)\n        return x\n\n\nclass SCSEModule2d(nn.Module):\n    def __init__(self, in_channels, reduction=16):\n        super().__init__()\n        self.cSE = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(in_channels, in_channels // reduction, 1),\n            nn.Tanh(),\n            nn.Conv2d(in_channels // reduction, in_channels, 1),\n            nn.Sigmoid(),\n        )\n        self.sSE = nn.Sequential(\n            nn.Conv2d(in_channels, 1, 1), \n            nn.Sigmoid(),\n            )\n\n    def forward(self, x):\n        return x * self.cSE(x) + x * self.sSE(x)\n\nclass Attention2d(nn.Module):\n    def __init__(self, name, **params):\n        super().__init__()\n        if name is None:\n            self.attention = nn.Identity(**params)\n        elif name == \"scse\":\n            self.attention = SCSEModule2d(**params)\n        else:\n            raise ValueError(\"Attention {} is not implemented\".format(name))\n\n    def forward(self, x):\n        return self.attention(x)\n\nclass DecoderBlock2d(nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        skip_channels,\n        out_channels,\n        norm_layer: nn.Module = nn.Identity,\n        attention_type: str = None,\n        intermediate_conv: bool = False,\n        upsample_mode: str = \"deconv\",\n        scale_factor: int = 2,\n    ):\n        super().__init__()\n\n        # Upsample block\n        if upsample_mode == \"pixelshuffle\":\n            self.upsample= SubpixelUpsample(\n                spatial_dims= 2,\n                in_channels= in_channels,\n                scale_factor= scale_factor,\n            )\n        else:\n            self.upsample = UpSample(\n                spatial_dims= 2,\n                in_channels= in_channels,\n                out_channels= in_channels,\n                scale_factor= scale_factor,\n                mode= upsample_mode,\n            )\n\n        if intermediate_conv:\n            k= 3\n            c= skip_channels if skip_channels != 0 else in_channels\n            self.intermediate_conv = nn.Sequential(\n                ConvBnAct2d(c, c, k, k//2),\n                ConvBnAct2d(c, c, k, k//2),\n                )\n        else:\n            self.intermediate_conv= None\n\n        self.attention1 = Attention2d(\n            name= attention_type, \n            in_channels= in_channels + skip_channels,\n            )\n\n        self.conv1 = ConvBnAct2d(\n            in_channels + skip_channels,\n            out_channels,\n            kernel_size= 3,\n            padding= 1,\n            norm_layer= norm_layer,\n        )\n\n        self.conv2 = ConvBnAct2d(\n            out_channels,\n            out_channels,\n            kernel_size= 3,\n            padding= 1,\n            norm_layer= norm_layer,\n        )\n        self.attention2 = Attention2d(\n            name= attention_type, \n            in_channels= out_channels,\n            )\n\n    def forward(self, x, skip=None):\n        x = self.upsample(x)\n\n        if self.intermediate_conv is not None:\n            if skip is not None:\n                skip = self.intermediate_conv(skip)\n            else:\n                x = self.intermediate_conv(x)\n\n        if skip is not None:\n            # print(x.shape, skip.shape)\n            x = torch.cat([x, skip], dim=1)\n            x = self.attention1(x)\n\n        x = self.conv1(x)\n        x = self.conv2(x)\n        x = self.attention2(x)\n        return x\n\n\nclass UnetDecoder2d(nn.Module):\n    \"\"\"\n    Unet decoder.\n    Source: https://arxiv.org/abs/1505.04597\n    \"\"\"\n    def __init__(\n        self,\n        encoder_channels: tuple[int],\n        skip_channels: tuple[int] = None,\n        decoder_channels: tuple = (256, 128, 64, 32),\n        scale_factors: tuple = (1,2,2,2),\n        norm_layer: nn.Module = nn.Identity,\n        attention_type: str = None,\n        intermediate_conv: bool = True,\n        upsample_mode: str = \"deconv\",\n    ):\n        super().__init__()\n        \n        if len(encoder_channels) == 4:\n            decoder_channels= decoder_channels[1:]\n        self.decoder_channels= decoder_channels\n        \n        if skip_channels is None:\n            skip_channels= list(encoder_channels[1:]) + [0]\n\n        # Build decoder blocks\n        in_channels= [encoder_channels[0]] + list(decoder_channels[:-1])\n        self.blocks = nn.ModuleList()\n\n        for i, (ic, sc, dc) in enumerate(zip(in_channels, skip_channels, decoder_channels)):\n            # print(i, ic, sc, dc)\n            self.blocks.append(\n                DecoderBlock2d(\n                    ic, sc, dc, \n                    norm_layer= norm_layer,\n                    attention_type= attention_type,\n                    intermediate_conv= intermediate_conv,\n                    upsample_mode= upsample_mode,\n                    scale_factor= scale_factors[i],\n                    )\n            )\n\n    def forward(self, feats: list[torch.Tensor]):\n        res= [feats[0]]\n        feats= feats[1:]\n\n        # Decoder blocks\n        for i, b in enumerate(self.blocks):\n            skip= feats[i] if i < len(feats) else None\n            res.append(\n                b(res[-1], skip=skip),\n                )\n            \n        return res\n\nclass SegmentationHead2d(nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        out_channels,\n        scale_factor: tuple[int] = (2,2),\n        kernel_size: int = 3,\n        mode: str = \"nontrainable\",\n    ):\n        super().__init__()\n        self.conv= nn.Conv2d(\n            in_channels, out_channels, kernel_size= kernel_size,\n            padding= kernel_size//2\n        )\n        self.upsample = UpSample(\n            spatial_dims= 2,\n            in_channels= out_channels,\n            out_channels= out_channels,\n            scale_factor= scale_factor,\n            mode= mode,\n        )\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.upsample(x)\n        return x\n\nclass Net(nn.Module):\n    def __init__(\n        self,\n        backbone: str,\n        pretrained: bool = True,\n    ):\n        super().__init__()\n        \n        # Encoder\n        self.backbone= timm.create_model(\n            backbone,\n            in_chans= 5,\n            pretrained= pretrained,\n            features_only= True,\n            drop_path_rate=0.4,\n            )\n        ecs= [_[\"num_chs\"] for _ in self.backbone.feature_info][::-1]\n\n        # Decoder\n        self.decoder= UnetDecoder2d(\n            encoder_channels= ecs,\n        )\n\n        self.seg_head= SegmentationHead2d(\n            in_channels= self.decoder.decoder_channels[-1],\n            out_channels= 1,\n            scale_factor= 2,\n        )\n        self._update_stem(backbone)\n\n    def _update_stem(self, backbone):\n        if backbone.startswith(\"hgnet\"):\n            self.backbone.stem.stem1.conv.stride=(1,1)\n            self.backbone.stages_3.downsample.conv.stride=(1,1)\n        \n        elif backbone in [\"resnet18\"]:\n            self.backbone.layer4[0].downsample[0].stride= (1,1)\n            self.backbone.layer4[0].conv1.stride= (1,1)\n            self.backbone.layer3[0].downsample[0].stride= (1,1)\n            self.backbone.layer3[0].conv1.stride= (1,1)\n\n        else:\n            raise ValueError(\"Custom striding not implemented.\")\n        pass\n\n        \n    def proc_flip(self, x_in):\n        x_in= torch.flip(x_in, dims=[-3, -1])\n        x= self.backbone(x_in)\n        x= x[::-1]\n\n        # Decoder\n        x= self.decoder(x)\n        x_seg= self.seg_head(x[-1])\n        x_seg= x_seg[..., 1:-1, 1:-1]\n        x_seg= torch.flip(x_seg, dims=[-1])\n        x_seg= x_seg * 1500 + 3000\n        return x_seg\n\n    def forward(self, batch):\n        x= batch\n\n        # Encoder\n        x_in = x\n        x= self.backbone(x)\n        # print([_.shape for _ in x])\n        x= x[::-1]\n\n        # Decoder\n        x= self.decoder(x)\n        # print([_.shape for _ in x])\n        x_seg= self.seg_head(x[-1])\n        x_seg= x_seg[..., 1:-1, 1:-1]\n        x_seg= x_seg * 1500 + 3000\n    \n        if self.training:\n            return x_seg\n        else:\n            p1 = self.proc_flip(x_in)\n            x_seg = torch.mean(torch.stack([x_seg, p1]), dim=0)\n            return x_seg","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:23:44.255046Z","iopub.execute_input":"2025-05-13T22:23:44.255271Z","iopub.status.idle":"2025-05-13T22:23:44.270651Z","shell.execute_reply.started":"2025-05-13T22:23:44.255255Z","shell.execute_reply":"2025-05-13T22:23:44.270112Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Utils\n\nSame as Egor's.","metadata":{}},{"cell_type":"code","source":"%%writefile _utils.py\n\nimport datetime\n\ndef format_time(elapsed):\n    elapsed_rounded = int(round((elapsed)))\n    return str(datetime.timedelta(seconds=elapsed_rounded))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:23:44.271276Z","iopub.execute_input":"2025-05-13T22:23:44.271511Z","iopub.status.idle":"2025-05-13T22:23:44.284695Z","shell.execute_reply.started":"2025-05-13T22:23:44.271495Z","shell.execute_reply":"2025-05-13T22:23:44.284095Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train\n\nHere is the main training script. \n\nBy using 2x GPUs, we can use larger batch sizes and speed up model training. No more wasted Quota!\n\nI won't go into the details of the script as there are already many good resources explaining DDP. Here are a couple of good starting points. \n\n- [Run DDP scripts with 2 T 4](https://www.kaggle.com/code/cpmpml/run-ddp-scripts-with-2-t-4) by @CPMP\n- [Getting Started with Distributed Data Parallel\n](https://docs.pytorch.org/tutorials/intermediate/ddp_tutorial.html)\n- [Distributed Data Parallel Docs](https://docs.pytorch.org/docs/stable/notes/ddp.html)","metadata":{}},{"cell_type":"code","source":"%%writefile _train.py\n\nimport os\nimport time \nimport random\nimport numpy as np\nfrom tqdm import tqdm\n\nimport torch\nimport torch.nn as nn\nfrom torch.amp import autocast, GradScaler\n\nimport torch.distributed as dist\nfrom torch.utils.data import DistributedSampler\nfrom torch.nn.parallel import DistributedDataParallel\n\nfrom _cfg import cfg\nfrom _dataset import CustomDataset\nfrom _model import ModelEMA, Net\nfrom _utils import format_time\n\ndef set_seed(seed=1234):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = True\n\ndef setup(rank, world_size):\n    torch.cuda.set_device(rank)\n    dist.init_process_group(\"nccl\", rank=rank, world_size=world_size)\n    return\n\ndef cleanup():\n    dist.barrier()\n    dist.destroy_process_group()\n    return\n\ndef main(cfg):\n\n    # ========== Datasets / Dataloaders ==========\n    if cfg.local_rank == 0:\n        print(\"=\"*25)\n        print(\"Loading data..\")\n    train_ds = CustomDataset(cfg=cfg, mode=\"train\")\n    sampler= DistributedSampler(\n        train_ds, \n        num_replicas=cfg.world_size, \n        rank=cfg.local_rank,\n    )\n    train_dl = torch.utils.data.DataLoader(\n        train_ds, \n        sampler= sampler,\n        batch_size= cfg.batch_size, \n        num_workers= 4,\n    )\n    \n    valid_ds = CustomDataset(cfg=cfg, mode=\"valid\")\n    sampler= DistributedSampler(\n        valid_ds, \n        num_replicas=cfg.world_size, \n        rank=cfg.local_rank,\n    )\n    valid_dl = torch.utils.data.DataLoader(\n        valid_ds, \n        sampler= sampler,\n        batch_size= cfg.batch_size_val, \n        num_workers= 4,\n    )\n\n    # ========== Model / Optim ==========\n    model = Net(backbone=cfg.backbone)\n    model= model.to(cfg.local_rank)\n    if cfg.ema:\n        if cfg.local_rank == 0:\n            print(\"Initializing EMA model..\")\n        ema_model = ModelEMA(\n            model, \n            decay=cfg.ema_decay, \n            device=cfg.local_rank,\n        )\n    else:\n        ema_model = None\n    model= DistributedDataParallel(\n        model, \n        device_ids=[cfg.local_rank], \n        )\n    \n    criterion = nn.L1Loss()\n    optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\n    scaler = GradScaler()\n\n\n    # ========== Training ==========\n    if cfg.local_rank == 0:\n        print(\"=\"*25)\n        print(\"Give me warp {}, Mr. Sulu.\".format(cfg.world_size))\n        print(\"=\"*25)\n    \n    best_loss= 1_000_000\n    val_loss= 1_000_000\n\n    for epoch in range(0, cfg.epochs+1):\n        if epoch != 0:\n            tstart= time.time()\n            train_dl.sampler.set_epoch(epoch)\n    \n            # Train loop\n            model.train()\n            total_loss = []\n            for i, (x, y) in enumerate(train_dl):\n                x = x.to(cfg.local_rank)\n                y = y.to(cfg.local_rank)\n        \n                with autocast(cfg.device.type):\n                    logits = model(x)\n                    \n                loss = criterion(logits, y)\n        \n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n        \n                torch.nn.utils.clip_grad_norm_(model.parameters(), 3.0)\n        \n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n    \n                total_loss.append(loss.item())\n                \n                if ema_model is not None:\n                    ema_model.update(model)\n                    \n                if cfg.local_rank == 0 and (len(total_loss) >= cfg.logging_steps or i == 0):\n                    train_loss = np.mean(total_loss)\n                    total_loss = []\n                    print(\"Epoch {}:     Train MAE: {:.2f}     Val MAE: {:.2f}     Time: {}     Step: {}/{}\".format(\n                        epoch, \n                        train_loss,\n                        val_loss,\n                        format_time(time.time() - tstart),\n                        i+1, \n                        len(train_dl)+1, \n                    ))\n    \n        # ========== Valid ==========\n        model.eval()\n        val_logits = []\n        val_targets = []\n        with torch.no_grad():\n            for x, y in tqdm(valid_dl, disable=cfg.local_rank != 0):\n                x = x.to(cfg.local_rank)\n                y = y.to(cfg.local_rank)\n    \n                with autocast(cfg.device.type):\n                    if ema_model is not None:\n                        out = ema_model.module(x)\n                    else:\n                        out = model(x)\n\n                val_logits.append(out.cpu())\n                val_targets.append(y.cpu())\n\n            val_logits= torch.cat(val_logits, dim=0)\n            val_targets= torch.cat(val_targets, dim=0)\n                \n            loss = criterion(val_logits, val_targets).item()\n\n        # Gather loss\n        v = torch.tensor([loss], device=cfg.local_rank)\n        torch.distributed.all_reduce(v, op=dist.ReduceOp.SUM)\n        val_loss = (v[0] / cfg.world_size).item()\n    \n        # ========== Weights / Early stopping ==========\n        stop_train = torch.tensor([0], device=cfg.local_rank)\n        if cfg.local_rank == 0:\n            es= cfg.early_stopping\n            if val_loss < best_loss:\n                print(\"New best: {:.2f} -> {:.2f}\".format(best_loss, val_loss))\n                print(\"Saved weights..\")\n                best_loss = val_loss\n                if ema_model is not None:\n                    torch.save(ema_model.module.state_dict(), f'best_model_{cfg.seed}.pt')\n                else:\n                    torch.save(model.state_dict(), f'best_model_{cfg.seed}.pt')\n        \n                es[\"streak\"] = 0\n            else:\n                es= cfg.early_stopping\n                es[\"streak\"] += 1\n                if es[\"streak\"] > es[\"patience\"]:\n                    print(\"Ending training (early_stopping).\")\n                    stop_train = torch.tensor([1], device=cfg.local_rank)\n        \n        # Exits training on all ranks\n        dist.broadcast(stop_train, src=0)\n        if stop_train.item() == 1:\n            return\n\n    return\n    \n\n\nif __name__ == \"__main__\":\n\n    # GPU Specs\n    rank = int(os.environ[\"RANK\"])\n    world_size = int(os.environ[\"WORLD_SIZE\"])\n    _, total = torch.cuda.mem_get_info(device=rank)\n\n    # Init\n    setup(rank, world_size)\n    time.sleep(rank)\n    print(f\"Rank: {rank}, World size: {world_size}, GPU memory: {total / 1024**3:.2f}GB\", flush=True)\n    time.sleep(world_size - rank)\n\n    # Seed\n    set_seed(cfg.seed+rank)\n\n    # Run\n    cfg.local_rank= rank\n    cfg.world_size= world_size\n    main(cfg)\n    cleanup()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:23:44.285420Z","iopub.execute_input":"2025-05-13T22:23:44.285648Z","iopub.status.idle":"2025-05-13T22:23:44.299641Z","shell.execute_reply.started":"2025-05-13T22:23:44.285622Z","shell.execute_reply":"2025-05-13T22:23:44.299064Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if RUN_TRAIN:\n    print(\"Starting training..\")\n    !OMP_NUM_THREADS=1 torchrun --nproc_per_node=2 _train.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:23:44.300294Z","iopub.execute_input":"2025-05-13T22:23:44.300537Z","iopub.status.idle":"2025-05-13T22:23:44.312888Z","shell.execute_reply.started":"2025-05-13T22:23:44.300515Z","shell.execute_reply":"2025-05-13T22:23:44.312389Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Pretrained Models\n\nNext, we load in 3x pretrained models. These models were trained with with an effective batch_size of 512 (256 per GPU) and use the B4 variant of the HgnetV2 backbone.","metadata":{}},{"cell_type":"code","source":"import glob\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom _cfg import cfg\nfrom _model import Net, EnsembleModel\n\nif RUN_VALID or RUN_TEST:\n\n    # Load pretrained models\n    models = []\n    for f in sorted(glob.glob(\"/kaggle/input/openfwi-preprocessed-72x72/models/*.pt\")):\n        print(\"Loading: \", f)\n        m = Net(\n            backbone=\"hgnetv2_b4.ssld_stage2_ft_in1k\",\n            pretrained=False,\n        )\n        state_dict= torch.load(f, map_location=cfg.device, weights_only=True)\n        m.load_state_dict(state_dict)\n        models.append(m)\n    \n    # Combine\n    model = EnsembleModel(models)\n    model = model.to(cfg.device)\n    model = model.eval()\n    print(\"n_models: {:_}\".format(len(models)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:23:44.314910Z","iopub.execute_input":"2025-05-13T22:23:44.315100Z","iopub.status.idle":"2025-05-13T22:24:17.847293Z","shell.execute_reply.started":"2025-05-13T22:23:44.315086Z","shell.execute_reply":"2025-05-13T22:24:17.846606Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Valid\n\nNext, we score the ensemble on the validation set.","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm\nimport numpy as np\n\nimport torch\nimport torch.nn as nn\nfrom torch.amp import autocast\n\nfrom _dataset import CustomDataset\n\n\nif RUN_VALID:\n\n    # Dataset / Dataloader\n    valid_ds = CustomDataset(cfg=cfg, mode=\"valid\")\n    sampler = torch.utils.data.SequentialSampler(valid_ds)\n    valid_dl = torch.utils.data.DataLoader(\n        valid_ds, \n        sampler= sampler,\n        batch_size= cfg.batch_size_val, \n        num_workers= 4,\n    )\n\n    # Valid loop\n    criterion = nn.L1Loss()\n    val_logits = []\n    val_targets = []\n    \n    with torch.no_grad():\n        for x, y in tqdm(valid_dl):\n            x = x.to(cfg.device)\n            y = y.to(cfg.device)\n    \n            with autocast(cfg.device.type):\n                out = model(x)\n    \n            val_logits.append(out.cpu())\n            val_targets.append(y.cpu())\n    \n        val_logits= torch.cat(val_logits, dim=0)\n        val_targets= torch.cat(val_targets, dim=0)\n    \n        total_loss= criterion(val_logits, val_targets).item()\n    \n    # Dataset Scores\n    ds_idxs= np.array([valid_ds.records])\n    ds_idxs= np.repeat(ds_idxs, repeats=500)\n    \n    print(\"=\"*25)\n    with torch.no_grad():    \n        for idx in sorted(np.unique(ds_idxs)):\n    \n            # Mask\n            mask = ds_idxs == idx\n            logits_ds = val_logits[mask]\n            targets_ds = val_targets[mask]\n    \n            # Score predictions\n            loss = criterion(val_logits[mask], val_targets[mask]).item()\n            print(\"{:15} {:.2f}\".format(idx, loss))\n    print(\"=\"*25)\n    print(\"Val MAE: {:.2f}\".format(total_loss))\n    print(\"=\"*25)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:24:17.848251Z","iopub.execute_input":"2025-05-13T22:24:17.849046Z","iopub.status.idle":"2025-05-13T22:25:32.421396Z","shell.execute_reply.started":"2025-05-13T22:24:17.849024Z","shell.execute_reply":"2025-05-13T22:25:32.420441Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Test\n\nFinally, we make predictions on the test data.","metadata":{}},{"cell_type":"code","source":"import torch\n\nclass TestDataset(torch.utils.data.Dataset):\n    def __init__(self, test_files):\n        self.test_files = test_files\n\n    def __len__(self):\n        return len(self.test_files)\n\n    def __getitem__(self, i):\n        test_file = self.test_files[i]\n        test_stem = test_file.split(\"/\")[-1].split(\".\")[0]\n        return np.load(test_file), test_stem\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:25:32.422541Z","iopub.execute_input":"2025-05-13T22:25:32.422865Z","iopub.status.idle":"2025-05-13T22:25:32.428302Z","shell.execute_reply.started":"2025-05-13T22:25:32.422829Z","shell.execute_reply":"2025-05-13T22:25:32.427518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import csv\nimport time\nimport glob\nfrom tqdm import tqdm\nimport numpy as np\n\nfrom _utils import format_time\n\n\nif RUN_TEST:\n    row_count = 0\n    t0 = time.time()\n    \n    test_files = sorted(glob.glob(\"/kaggle/input/open-wfi-test/test/*.npy\"))\n    x_cols = [f\"x_{i}\" for i in range(1, 70, 2)]\n    fieldnames = [\"oid_ypos\"] + x_cols\n    \n    test_ds = TestDataset(test_files)\n    test_dl = torch.utils.data.DataLoader(\n        test_ds, \n        sampler=torch.utils.data.SequentialSampler(test_ds),\n        batch_size=cfg.batch_size_val, \n        num_workers=4,\n    )\n    \n    with open(\"submission.csv\", \"wt\", newline=\"\") as csvfile:\n        writer = csv.DictWriter(csvfile, fieldnames=fieldnames)\n        writer.writeheader()\n\n        with torch.inference_mode():\n            with torch.autocast(cfg.device.type):\n                for inputs, oids_test in tqdm(test_dl, total=len(test_dl)):\n                    inputs = inputs.to(cfg.device)\n            \n                    inputs = _preprocess(inputs)\n                    outputs = model(inputs)\n                            \n                    y_preds = outputs[:, 0].cpu().numpy()\n                    \n                    for y_pred, oid_test in zip(y_preds, oids_test):\n                        for y_pos in range(70):\n                            row = dict(zip(x_cols, [y_pred[y_pos, x_pos] for x_pos in range(1, 70, 2)]))\n                            row[\"oid_ypos\"] = f\"{oid_test}_y_{y_pos}\"\n            \n                            writer.writerow(row)\n                            row_count += 1\n\n                            # Clear buffer\n                            if row_count % 100_000 == 0:\n                                csvfile.flush()\n    \n    t1 = format_time(time.time() - t0)\n    print(f\"Inference Time: {t1}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:25:32.429192Z","iopub.execute_input":"2025-05-13T22:25:32.429469Z","iopub.status.idle":"2025-05-13T22:36:59.529536Z","shell.execute_reply.started":"2025-05-13T22:25:32.429449Z","shell.execute_reply":"2025-05-13T22:36:59.528616Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We can also view a few samples to make sure things look reasonable.","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt \n\nif RUN_TEST:\n    # Plot a few samples\n    fig, axes = plt.subplots(3, 5, figsize=(10, 6))\n    axes= axes.flatten()\n\n    n = min(len(outputs), len(axes))\n    \n    for i in range(n):\n        img= outputs[0, 0, ...].cpu().numpy()\n        img = outputs[i, 0].cpu().numpy()\n        idx= oids_test[i]\n    \n        # Plot\n        axes[i].imshow(img, cmap='gray')\n        axes[i].set_title(idx)\n        axes[i].axis('off')\n\n    for i in range(n, len(axes)):\n        axes[i].axis('off')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T22:36:59.534154Z","iopub.execute_input":"2025-05-13T22:36:59.534446Z","iopub.status.idle":"2025-05-13T22:37:01.164185Z","shell.execute_reply.started":"2025-05-13T22:36:59.534421Z","shell.execute_reply":"2025-05-13T22:37:01.163292Z"}},"outputs":[],"execution_count":null}]}