{"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":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"},{"sourceId":11334027,"sourceType":"datasetVersion","datasetId":7089850},{"sourceId":11367935,"sourceType":"datasetVersion","datasetId":7116013},{"sourceId":11368083,"sourceType":"datasetVersion","datasetId":7116134},{"sourceId":11368169,"sourceType":"datasetVersion","datasetId":7116196},{"sourceId":11368268,"sourceType":"datasetVersion","datasetId":7116272},{"sourceId":11368499,"sourceType":"datasetVersion","datasetId":7116445},{"sourceId":11368545,"sourceType":"datasetVersion","datasetId":7116479},{"sourceId":11368547,"sourceType":"datasetVersion","datasetId":7116481},{"sourceId":11376433,"sourceType":"datasetVersion","datasetId":7122462},{"sourceId":11376448,"sourceType":"datasetVersion","datasetId":7122476},{"sourceId":11376464,"sourceType":"datasetVersion","datasetId":7122489},{"sourceId":11376742,"sourceType":"datasetVersion","datasetId":7122712},{"sourceId":11376868,"sourceType":"datasetVersion","datasetId":7122812},{"sourceId":11376871,"sourceType":"datasetVersion","datasetId":7122814},{"sourceId":11376872,"sourceType":"datasetVersion","datasetId":7122815},{"sourceId":11376935,"sourceType":"datasetVersion","datasetId":7122866},{"sourceId":11377083,"sourceType":"datasetVersion","datasetId":7122981},{"sourceId":11377231,"sourceType":"datasetVersion","datasetId":7123090},{"sourceId":11377291,"sourceType":"datasetVersion","datasetId":7123138},{"sourceId":11377325,"sourceType":"datasetVersion","datasetId":7123163},{"sourceId":11377334,"sourceType":"datasetVersion","datasetId":7123172},{"sourceId":11377594,"sourceType":"datasetVersion","datasetId":7123380},{"sourceId":11377614,"sourceType":"datasetVersion","datasetId":7123394},{"sourceId":11377741,"sourceType":"datasetVersion","datasetId":7123490},{"sourceId":11377752,"sourceType":"datasetVersion","datasetId":7123499},{"sourceId":11377756,"sourceType":"datasetVersion","datasetId":7123503},{"sourceId":11377935,"sourceType":"datasetVersion","datasetId":7123649},{"sourceId":11377970,"sourceType":"datasetVersion","datasetId":7123675},{"sourceId":11378141,"sourceType":"datasetVersion","datasetId":7123813},{"sourceId":11378162,"sourceType":"datasetVersion","datasetId":7123831},{"sourceId":11378178,"sourceType":"datasetVersion","datasetId":7123842},{"sourceId":11380909,"sourceType":"datasetVersion","datasetId":7125380}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Modified code of [HgNet-Training & Inference](https://www.kaggle.com/code/i2nfinit3y/hgnet-training-inference) Upvote it \n\n**V5**\n* train with all OpenFWI datasets\n* multiple gpus\n* average of l2 + l1 loss","metadata":{}},{"cell_type":"markdown","source":" ![](https://www.googleapis.com/download/storage/v1/b/kaggle-forum-message-attachments/o/inbox%2F761268%2F29abb6c11501446345a147792e4a8f73%2FScreenshot%202025-04-14%20at%201.50.06PM.png?generation=1744618887612399&alt=media)","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nfrom colorama import init, Fore, Style\nfrom torch.utils.data import Dataset, DataLoader\n\ninit(autoreset=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T21:29:46.471419Z","iopub.execute_input":"2025-04-12T21:29:46.471710Z","iopub.status.idle":"2025-04-12T21:29:49.200873Z","shell.execute_reply.started":"2025-04-12T21:29:46.471688Z","shell.execute_reply":"2025-04-12T21:29:49.199869Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_parts = range(2,31)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T21:29:49.202122Z","iopub.execute_input":"2025-04-12T21:29:49.202490Z","iopub.status.idle":"2025-04-12T21:29:49.206238Z","shell.execute_reply.started":"2025-04-12T21:29:49.202467Z","shell.execute_reply":"2025-04-12T21:29:49.205527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\" PP-HGNet (V1 & V2)\n\nReference:\nhttps://github.com/PaddlePaddle/PaddleClas/blob/develop/docs/zh_CN/models/ImageNet1k/PP-HGNetV2.md\nThe Paddle Implement of PP-HGNet (https://github.com/PaddlePaddle/PaddleClas/blob/release/2.5.1/docs/en/models/PP-HGNet_en.md)\nPP-HGNet: https://github.com/PaddlePaddle/PaddleClas/blob/release/2.5.1/ppcls/arch/backbone/legendary_models/pp_hgnet.py\nPP-HGNetv2: https://github.com/PaddlePaddle/PaddleClas/blob/release/2.5.1/ppcls/arch/backbone/legendary_models/pp_hgnet_v2.py\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom timm.data import IMAGENET_DEFAULT_MEAN, IMAGENET_DEFAULT_STD\nfrom timm.layers import SelectAdaptivePool2d, DropPath, create_conv2d\nfrom timm.models._builder import build_model_with_cfg\nfrom timm.models._registry import register_model, generate_default_cfgs\n\n__all__ = ['HighPerfGpuNet']\n\n\nclass LearnableAffineBlock(nn.Module):\n    def __init__(\n            self,\n            scale_value=1.0,\n            bias_value=0.0\n    ):\n        super().__init__()\n        self.scale = nn.Parameter(torch.tensor([scale_value]), requires_grad=True)\n        self.bias = nn.Parameter(torch.tensor([bias_value]), requires_grad=True)\n\n    def forward(self, x):\n        return self.scale * x + self.bias\n\n\nclass ConvBNAct(nn.Module):\n    def __init__(\n            self,\n            in_chs,\n            out_chs,\n            kernel_size,\n            stride=1,\n            groups=1,\n            padding='',\n            use_act=True,\n            use_lab=False\n    ):\n        super().__init__()\n        self.use_act = use_act\n        self.use_lab = use_lab\n        self.conv = create_conv2d(\n            in_chs,\n            out_chs,\n            kernel_size,\n            stride=stride,\n            padding=padding,\n            groups=groups,\n        )\n        self.bn = nn.BatchNorm2d(out_chs)\n        if self.use_act:\n            self.act = nn.ReLU()\n        else:\n            self.act = nn.Identity()\n        if self.use_act and self.use_lab:\n            self.lab = LearnableAffineBlock()\n        else:\n            self.lab = nn.Identity()\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.bn(x)\n        x = self.act(x)\n        x = self.lab(x)\n        return x\n\n\nclass LightConvBNAct(nn.Module):\n    def __init__(\n            self,\n            in_chs,\n            out_chs,\n            kernel_size,\n            groups=1,\n            use_lab=False\n    ):\n        super().__init__()\n        self.conv1 = ConvBNAct(\n            in_chs,\n            out_chs,\n            kernel_size=1,\n            use_act=False,\n            use_lab=use_lab,\n        )\n        self.conv2 = ConvBNAct(\n            out_chs,\n            out_chs,\n            kernel_size=kernel_size,\n            groups=out_chs,\n            use_act=True,\n            use_lab=use_lab,\n        )\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.conv2(x)\n        return x\n\n\nclass EseModule(nn.Module):\n    def __init__(self, chs):\n        super().__init__()\n        self.conv = nn.Conv2d(\n            chs,\n            chs,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n        )\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        identity = x\n        x = x.mean((2, 3), keepdim=True)\n        x = self.conv(x)\n        x = self.sigmoid(x)\n        return torch.mul(identity, x)\n\n\nclass StemV1(nn.Module):\n    # for PP-HGNet\n    def __init__(self, stem_chs):\n        super().__init__()\n        self.stem = nn.Sequential(*[\n            ConvBNAct(\n                stem_chs[i],\n                stem_chs[i + 1],\n                kernel_size=3,\n                stride=2 if i == 0 else 1) for i in range(\n                len(stem_chs) - 1)\n        ])\n        self.pool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n\n    def forward(self, x):\n        x = self.stem(x)\n        x = self.pool(x)\n        return x\n\n\nclass StemV2(nn.Module):\n    # for PP-HGNetv2\n    def __init__(self, in_chs, mid_chs, out_chs, use_lab=False):\n        super().__init__()\n        self.stem1 = ConvBNAct(\n            in_chs,\n            mid_chs,\n            kernel_size=3,\n            stride=2,\n            use_lab=use_lab,\n        )\n        self.stem2a = ConvBNAct(\n            mid_chs,\n            mid_chs // 2,\n            kernel_size=2,\n            stride=1,\n            use_lab=use_lab,\n        )\n        self.stem2b = ConvBNAct(\n            mid_chs // 2,\n            mid_chs,\n            kernel_size=2,\n            stride=1,\n            use_lab=use_lab,\n        )\n        self.stem3 = ConvBNAct(\n            mid_chs * 2,\n            mid_chs,\n            kernel_size=3,\n            stride=2,\n            use_lab=use_lab,\n        )\n        self.stem4 = ConvBNAct(\n            mid_chs,\n            out_chs,\n            kernel_size=1,\n            stride=1,\n            use_lab=use_lab,\n        )\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=1, ceil_mode=True)\n\n    def forward(self, x):\n        x = self.stem1(x)\n        x = F.pad(x, (0, 1, 0, 1))\n        x2 = self.stem2a(x)\n        x2 = F.pad(x2, (0, 1, 0, 1))\n        x2 = self.stem2b(x2)\n        x1 = self.pool(x)\n        x = torch.cat([x1, x2], dim=1)\n        x = self.stem3(x)\n        x = self.stem4(x)\n        return x\n\n\nclass HighPerfGpuBlock(nn.Module):\n    def __init__(\n            self,\n            in_chs,\n            mid_chs,\n            out_chs,\n            layer_num,\n            kernel_size=3,\n            residual=False,\n            light_block=False,\n            use_lab=False,\n            agg='ese',\n            drop_path=0.,\n    ):\n        super().__init__()\n        self.residual = residual\n\n        self.layers = nn.ModuleList()\n        for i in range(layer_num):\n            if light_block:\n                self.layers.append(\n                    LightConvBNAct(\n                        in_chs if i == 0 else mid_chs,\n                        mid_chs,\n                        kernel_size=kernel_size,\n                        use_lab=use_lab,\n                    )\n                )\n            else:\n                self.layers.append(\n                    ConvBNAct(\n                        in_chs if i == 0 else mid_chs,\n                        mid_chs,\n                        kernel_size=kernel_size,\n                        stride=1,\n                        use_lab=use_lab,\n                    )\n                )\n\n        # feature aggregation\n        total_chs = in_chs + layer_num * mid_chs\n        if agg == 'se':\n            aggregation_squeeze_conv = ConvBNAct(\n                total_chs,\n                out_chs // 2,\n                kernel_size=1,\n                stride=1,\n                use_lab=use_lab,\n            )\n            aggregation_excitation_conv = ConvBNAct(\n                out_chs // 2,\n                out_chs,\n                kernel_size=1,\n                stride=1,\n                use_lab=use_lab,\n            )\n            self.aggregation = nn.Sequential(\n                aggregation_squeeze_conv,\n                aggregation_excitation_conv,\n            )\n        else:\n            aggregation_conv = ConvBNAct(\n                total_chs,\n                out_chs,\n                kernel_size=1,\n                stride=1,\n                use_lab=use_lab,\n            )\n            att = EseModule(out_chs)\n            self.aggregation = nn.Sequential(\n                aggregation_conv,\n                att,\n            )\n\n        self.drop_path = DropPath(drop_path) if drop_path else nn.Identity()\n\n    def forward(self, x):\n        identity = x\n        output = [x]\n        for layer in self.layers:\n            x = layer(x)\n            output.append(x)\n        x = torch.cat(output, dim=1)\n        x = self.aggregation(x)\n        if self.residual:\n            x = self.drop_path(x) + identity\n        return x\n\n\nclass HighPerfGpuStage(nn.Module):\n    def __init__(\n            self,\n            in_chs,\n            mid_chs,\n            out_chs,\n            block_num,\n            layer_num,\n            downsample=True,\n            stride=2,\n            light_block=False,\n            kernel_size=3,\n            use_lab=False,\n            agg='ese',\n            drop_path=0.,\n    ):\n        super().__init__()\n        self.downsample = downsample\n        if downsample:\n            self.downsample = ConvBNAct(\n                in_chs,\n                in_chs,\n                kernel_size=3,\n                stride=stride,\n                groups=in_chs,\n                use_act=False,\n                use_lab=use_lab,\n            )\n        else:\n            self.downsample = nn.Identity()\n\n        blocks_list = []\n        for i in range(block_num):\n            blocks_list.append(\n                HighPerfGpuBlock(\n                    in_chs if i == 0 else out_chs,\n                    mid_chs,\n                    out_chs,\n                    layer_num,\n                    residual=False if i == 0 else True,\n                    kernel_size=kernel_size,\n                    light_block=light_block,\n                    use_lab=use_lab,\n                    agg=agg,\n                    drop_path=drop_path[i] if isinstance(drop_path, (list, tuple)) else drop_path,\n                )\n            )\n        self.blocks = nn.Sequential(*blocks_list)\n\n    def forward(self, x):\n        x = self.downsample(x)\n        x = self.blocks(x)\n        return x\n\n\nclass ClassifierHead(nn.Module):\n    def __init__(\n            self,\n            num_features,\n            num_classes,\n            pool_type='avg',\n            drop_rate=0.,\n            use_last_conv=True,\n            class_expand=2048,\n            use_lab=False\n    ):\n        super(ClassifierHead, self).__init__()\n        self.global_pool = SelectAdaptivePool2d(pool_type=pool_type, flatten=False, input_fmt='NCHW')\n        if use_last_conv:\n            last_conv = nn.Conv2d(\n                num_features,\n                class_expand,\n                kernel_size=1,\n                stride=1,\n                padding=0,\n                bias=False,\n            )\n            act = nn.ReLU()\n            if use_lab:\n                lab = LearnableAffineBlock()\n                self.last_conv = nn.Sequential(last_conv, act, lab)\n            else:\n                self.last_conv = nn.Sequential(last_conv, act)\n        else:\n            self.last_conv = nn.Indentity()\n\n        if drop_rate > 0:\n            self.dropout = nn.Dropout(drop_rate)\n        else:\n            self.dropout = nn.Identity()\n\n        self.flatten = nn.Flatten()\n        self.fc = nn.Linear(class_expand if use_last_conv else num_features, num_classes)\n\n    def forward(self, x, pre_logits: bool = False):\n        x = self.global_pool(x)\n        x = self.last_conv(x)\n        x = self.dropout(x)\n        x = self.flatten(x)\n        if pre_logits:\n            return x\n        x = self.fc(x)\n        return x\n\n\nclass HighPerfGpuNet(nn.Module):\n\n    def __init__(\n            self,\n            cfg,\n            in_chans=3,\n            num_classes=1000,\n            global_pool='avg',\n            use_last_conv=True,\n            class_expand=2048,\n            drop_rate=0.,\n            drop_path_rate=0.,\n            use_lab=False,\n            **kwargs,\n    ):\n        super(HighPerfGpuNet, self).__init__()\n        stem_type = cfg[\"stem_type\"]\n        stem_chs = cfg[\"stem_chs\"]\n        stages_cfg = [cfg[\"stage1\"], cfg[\"stage2\"], cfg[\"stage3\"], cfg[\"stage4\"]]\n        self.num_classes = num_classes\n        self.drop_rate = drop_rate\n        self.use_last_conv = use_last_conv\n        self.class_expand = class_expand\n        self.use_lab = use_lab\n\n        assert stem_type in ['v1', 'v2']\n        if stem_type == 'v2':\n            self.stem = StemV2(\n                in_chs=in_chans,\n                mid_chs=stem_chs[0],\n                out_chs=stem_chs[1],\n                use_lab=use_lab)\n        else:\n            self.stem = StemV1([in_chans] + stem_chs)\n\n        current_stride = 4\n\n        stages = []\n        self.feature_info = []\n        block_depths = [c[3] for c in stages_cfg]\n        dpr = [x.tolist() for x in torch.linspace(0, drop_path_rate, sum(block_depths)).split(block_depths)]\n        for i, stage_config in enumerate(stages_cfg):\n            in_chs, mid_chs, out_chs, block_num, downsample, light_block, kernel_size, layer_num = stage_config\n            stages += [HighPerfGpuStage(\n                in_chs=in_chs,\n                mid_chs=mid_chs,\n                out_chs=out_chs,\n                block_num=block_num,\n                layer_num=layer_num,\n                downsample=downsample,\n                light_block=light_block,\n                kernel_size=kernel_size,\n                use_lab=use_lab,\n                agg='ese' if stem_type == 'v1' else 'se',\n                drop_path=dpr[i],\n            )]\n            self.num_features = out_chs\n            if downsample:\n                current_stride *= 2\n            self.feature_info += [dict(num_chs=self.num_features, reduction=current_stride, module=f'stages.{i}')]\n        self.stages = nn.Sequential(*stages)\n\n        if num_classes > 0:\n            self.head = ClassifierHead(\n                self.num_features,\n                num_classes=num_classes,\n                pool_type=global_pool,\n                drop_rate=drop_rate,\n                use_last_conv=use_last_conv,\n                class_expand=class_expand,\n                use_lab=use_lab\n            )\n        else:\n            if global_pool == 'avg':\n                self.head = SelectAdaptivePool2d(pool_type=global_pool, flatten=True)\n            else:\n                self.head = nn.Identity()\n\n        for n, m in self.named_modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.ones_(m.weight)\n                nn.init.zeros_(m.bias)\n            elif isinstance(m, nn.Linear):\n                nn.init.zeros_(m.bias)\n\n    @torch.jit.ignore\n    def group_matcher(self, coarse=False):\n        return dict(\n            stem=r'^stem',\n            blocks=r'^stages\\.(\\d+)' if coarse else r'^stages\\.(\\d+).blocks\\.(\\d+)',\n        )\n\n    @torch.jit.ignore\n    def set_grad_checkpointing(self, enable=True):\n        for s in self.stages:\n            s.grad_checkpointing = enable\n\n    @torch.jit.ignore\n    def get_classifier(self):\n        return self.head.fc\n\n    def reset_classifier(self, num_classes, global_pool='avg'):\n        self.num_classes = num_classes\n        if num_classes > 0:\n            self.head = ClassifierHead(\n                self.num_features,\n                num_classes=num_classes,\n                pool_type=global_pool,\n                drop_rate=self.drop_rate,\n                use_last_conv=self.use_last_conv,\n                class_expand=self.class_expand,\n                use_lab=self.use_lab)\n        else:\n            if global_pool:\n                self.head = SelectAdaptivePool2d(pool_type=global_pool, flatten=True)\n            else:\n                self.head = nn.Identity()\n\n    def forward_features(self, x):\n        x = self.stem(x)\n        return self.stages(x)\n\n    def forward_head(self, x, pre_logits: bool = False):\n        return self.head(x, pre_logits=pre_logits) if pre_logits else self.head(x)\n\n    def forward(self, x):\n        x = self.forward_features(x)\n        x = self.forward_head(x)\n        return x\n\n\nmodel_cfgs = dict(\n    # PP-HGNet\n    hgnet_tiny={\n        \"stem_type\": 'v1',\n        \"stem_chs\": [48, 48, 96],\n        # in_chs, mid_chs, out_chs, blocks, downsample, light_block, kernel_size, layer_num\n        \"stage1\": [96, 96, 224, 1, False, False, 3, 5],\n        \"stage2\": [224, 128, 448, 1, True, False, 3, 5],\n        \"stage3\": [448, 160, 512, 2, True, False, 3, 5],\n        \"stage4\": [512, 192, 768, 1, True, False, 3, 5],\n    },\n    hgnet_small={\n        \"stem_type\": 'v1',\n        \"stem_chs\": [64, 64, 128],\n        # in_chs, mid_chs, out_chs, blocks, downsample, light_block, kernel_size, layer_num\n        \"stage1\": [128, 128, 256, 1, False, False, 3, 6],\n        \"stage2\": [256, 160, 512, 1, True, False, 3, 6],\n        \"stage3\": [512, 192, 768, 2, True, False, 3, 6],\n        \"stage4\": [768, 224, 1024, 1, True, False, 3, 6],\n    },\n    hgnet_base={\n        \"stem_type\": 'v1',\n        \"stem_chs\": [96, 96, 160],\n        # in_chs, mid_chs, out_chs, blocks, downsample, light_block, kernel_size, layer_num\n        \"stage1\": [160, 192, 320, 1, False, False, 3, 7],\n        \"stage2\": [320, 224, 640, 2, True, False, 3, 7],\n        \"stage3\": [640, 256, 960, 3, True, False, 3, 7],\n        \"stage4\": [960, 288, 1280, 2, True, False, 3, 7],\n    },\n    # PP-HGNetv2\n    hgnetv2_b0={\n        \"stem_type\": 'v2',\n        \"stem_chs\": [16, 16],\n        # in_chs, mid_chs, out_chs, blocks, downsample, light_block, kernel_size, layer_num\n        \"stage1\": [16, 16, 64, 1, False, False, 3, 3],\n        \"stage2\": [64, 32, 256, 1, True, False, 3, 3],\n        \"stage3\": [256, 64, 512, 2, True, True, 5, 3],\n        \"stage4\": [512, 128, 1024, 1, True, True, 5, 3],\n    },\n    hgnetv2_b1={\n        \"stem_type\": 'v2',\n        \"stem_chs\": [24, 32],\n        # in_chs, mid_chs, out_chs, blocks, downsample, light_block, kernel_size, layer_num\n        \"stage1\": [32, 32, 64, 1, False, False, 3, 3],\n        \"stage2\": [64, 48, 256, 1, True, False, 3, 3],\n        \"stage3\": [256, 96, 512, 2, True, True, 5, 3],\n        \"stage4\": [512, 192, 1024, 1, True, True, 5, 3],\n    },\n    hgnetv2_b2={\n        \"stem_type\": 'v2',\n        \"stem_chs\": [24, 32],\n        # in_chs, mid_chs, out_chs, blocks, downsample, light_block, kernel_size, layer_num\n        \"stage1\": [32, 32, 96, 1, False, False, 3, 4],\n        \"stage2\": [96, 64, 384, 1, True, False, 3, 4],\n        \"stage3\": [384, 128, 768, 3, True, True, 5, 4],\n        \"stage4\": [768, 256, 1536, 1, True, True, 5, 4],\n    },\n    hgnetv2_b3={\n        \"stem_type\": 'v2',\n        \"stem_chs\": [24, 32],\n        # in_chs, mid_chs, out_chs, blocks, downsample, light_block, kernel_size, layer_num\n        \"stage1\": [32, 32, 128, 1, False, False, 3, 5],\n        \"stage2\": [128, 64, 512, 1, True, False, 3, 5],\n        \"stage3\": [512, 128, 1024, 3, True, True, 5, 5],\n        \"stage4\": [1024, 256, 2048, 1, True, True, 5, 5],\n    },\n    hgnetv2_b4={\n        \"stem_type\": 'v2',\n        \"stem_chs\": [32, 48],\n        # in_chs, mid_chs, out_chs, blocks, downsample, light_block, kernel_size, layer_num\n        \"stage1\": [48, 48, 128, 1, False, False, 3, 6],\n        \"stage2\": [128, 96, 512, 1, True, False, 3, 6],\n        \"stage3\": [512, 192, 1024, 3, True, True, 5, 6],\n        \"stage4\": [1024, 384, 2048, 1, True, True, 5, 6],\n    },\n    hgnetv2_b5={\n        \"stem_type\": 'v2',\n        \"stem_chs\": [32, 64],\n        # in_chs, mid_chs, out_chs, blocks, downsample, light_block, kernel_size, layer_num\n        \"stage1\": [64, 64, 128, 1, False, False, 3, 6],\n        \"stage2\": [128, 128, 512, 2, True, False, 3, 6],\n        \"stage3\": [512, 256, 1024, 5, True, True, 5, 6],\n        \"stage4\": [1024, 512, 2048, 2, True, True, 5, 6],\n    },\n    hgnetv2_b6={\n        \"stem_type\": 'v2',\n        \"stem_chs\": [48, 96],\n        # in_chs, mid_chs, out_chs, blocks, downsample, light_block, kernel_size, layer_num\n        \"stage1\": [96, 96, 192, 2, False, False, 3, 6],\n        \"stage2\": [192, 192, 512, 3, True, False, 3, 6],\n        \"stage3\": [512, 384, 1024, 6, True, True, 5, 6],\n        \"stage4\": [1024, 768, 2048, 3, True, True, 5, 6],\n    },\n)\n\n\ndef _create_hgnet(variant, pretrained=False, **kwargs):\n    out_indices = kwargs.pop('out_indices', (0, 1, 2, 3))\n    return build_model_with_cfg(\n        HighPerfGpuNet,\n        variant,\n        pretrained,\n        model_cfg=model_cfgs[variant],\n        feature_cfg=dict(flatten_sequential=True, out_indices=out_indices),\n        **kwargs,\n    )\n\n\ndef _cfg(url='', **kwargs):\n    return {\n        'url': url,\n        'num_classes': 1000, 'input_size': (3, 224, 224), 'pool_size': (7, 7),\n        'crop_pct': 0.965, 'interpolation': 'bicubic',\n        'mean': IMAGENET_DEFAULT_MEAN, 'std': IMAGENET_DEFAULT_STD,\n        'classifier': 'head.fc', 'first_conv': 'stem.stem1.conv',\n        'test_crop_pct': 1.0, 'test_input_size': (3, 288, 288),\n        **kwargs,\n    }\n\n\ndefault_cfgs = generate_default_cfgs({\n    'hgnet_tiny.paddle_in1k': _cfg(\n        first_conv='stem.stem.0.conv',\n        hf_hub_id='timm/'),\n    'hgnet_tiny.ssld_in1k': _cfg(\n        first_conv='stem.stem.0.conv',\n        hf_hub_id='timm/'),\n    'hgnet_small.paddle_in1k': _cfg(\n        first_conv='stem.stem.0.conv',\n        hf_hub_id='timm/'),\n    'hgnet_small.ssld_in1k': _cfg(\n        first_conv='stem.stem.0.conv',\n        hf_hub_id='timm/'),\n    'hgnet_base.ssld_in1k': _cfg(\n        first_conv='stem.stem.0.conv',\n        hf_hub_id='timm/'),\n    'hgnetv2_b0.ssld_stage2_ft_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b0.ssld_stage1_in22k_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b1.ssld_stage2_ft_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b1.ssld_stage1_in22k_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b2.ssld_stage2_ft_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b2.ssld_stage1_in22k_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b3.ssld_stage2_ft_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b3.ssld_stage1_in22k_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b4.ssld_stage2_ft_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b4.ssld_stage1_in22k_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b5.ssld_stage2_ft_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b5.ssld_stage1_in22k_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b6.ssld_stage2_ft_in1k': _cfg(\n        hf_hub_id='timm/'),\n    'hgnetv2_b6.ssld_stage1_in22k_in1k': _cfg(\n        hf_hub_id='timm/'),\n})\n\n\n@register_model\ndef hgnet_tiny(pretrained=False, **kwargs) -> HighPerfGpuNet:\n    return _create_hgnet('hgnet_tiny', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef hgnet_small(pretrained=False, **kwargs) -> HighPerfGpuNet:\n    return _create_hgnet('hgnet_small', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef hgnet_base(pretrained=False, **kwargs) -> HighPerfGpuNet:\n    return _create_hgnet('hgnet_base', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef hgnetv2_b0(pretrained=False, **kwargs) -> HighPerfGpuNet:\n    return _create_hgnet('hgnetv2_b0', pretrained=pretrained, use_lab=True, **kwargs)\n\n\n@register_model\ndef hgnetv2_b1(pretrained=False, **kwargs) -> HighPerfGpuNet:\n    return _create_hgnet('hgnetv2_b1', pretrained=pretrained, use_lab=True, **kwargs)\n\n\n@register_model\ndef hgnetv2_b2(pretrained=False, **kwargs) -> HighPerfGpuNet:\n    return _create_hgnet('hgnetv2_b2', pretrained=pretrained, use_lab=True, **kwargs)\n\n\n@register_model\ndef hgnetv2_b3(pretrained=False, **kwargs) -> HighPerfGpuNet:\n    return _create_hgnet('hgnetv2_b3', pretrained=pretrained, use_lab=True, **kwargs)\n\n\n@register_model\ndef hgnetv2_b4(pretrained=False, **kwargs) -> HighPerfGpuNet:\n    return _create_hgnet('hgnetv2_b4', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef hgnetv2_b5(pretrained=False, **kwargs) -> HighPerfGpuNet:\n    return _create_hgnet('hgnetv2_b5', pretrained=pretrained, **kwargs)\n\n\n@register_model\ndef hgnetv2_b6(pretrained=False, **kwargs) -> HighPerfGpuNet:\n    return _create_hgnet('hgnetv2_b6', pretrained=pretrained, **kwargs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T21:29:49.207067Z","iopub.execute_input":"2025-04-12T21:29:49.207308Z","iopub.status.idle":"2025-04-12T21:29:52.036284Z","shell.execute_reply.started":"2025-04-12T21:29:49.207291Z","shell.execute_reply":"2025-04-12T21:29:52.035598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_inputs = []\nfor i in train_parts:\n    train_inputs.extend([f\n    for f in\n    Path(f'/kaggle/input/waveform-inversion-{i}').rglob('*.npy')\n    if ('seis' in f.stem) or ('data' in f.stem)\n])\n\nvalid_inputs = [\n    f\n    for f in\n    Path('/kaggle/input/waveform-inversion/train_samples').rglob('*.npy')\n    if ('seis' in f.stem) or ('data' in f.stem)\n]\nvalid_inputs = [valid_inputs[i] for i in range(0, len(valid_inputs), 2)]\nlen(train_inputs), len(valid_inputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T21:29:52.037979Z","iopub.execute_input":"2025-04-12T21:29:52.038235Z","iopub.status.idle":"2025-04-12T21:29:52.090062Z","shell.execute_reply.started":"2025-04-12T21:29:52.038217Z","shell.execute_reply":"2025-04-12T21:29:52.089514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def inputs_files_to_output_files(input_files):\n    return [\n        Path(str(f).replace('seis', 'vel').replace('data', 'model'))\n        for f in input_files\n    ]\n\ntrain_outputs = inputs_files_to_output_files(train_inputs)\nvalid_outputs = inputs_files_to_output_files(valid_inputs)\nlen(train_outputs), len(valid_outputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T21:29:52.090655Z","iopub.execute_input":"2025-04-12T21:29:52.090876Z","iopub.status.idle":"2025-04-12T21:29:52.096256Z","shell.execute_reply.started":"2025-04-12T21:29:52.090856Z","shell.execute_reply":"2025-04-12T21:29:52.095676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SeismicDataset(Dataset):\n    def __init__(self, inputs_files, output_files, n_examples_per_file=500):\n        assert len(inputs_files) == len(output_files)\n        self.inputs_files = inputs_files\n        self.output_files = output_files\n        self.n_examples_per_file = n_examples_per_file\n\n    def __len__(self):\n        return len(self.inputs_files) * self.n_examples_per_file\n\n    def __getitem__(self, idx):\n        # Calculate file offset and sample offset within file\n        file_idx = idx // self.n_examples_per_file\n        sample_idx = idx % self.n_examples_per_file\n\n        X = np.load(self.inputs_files[file_idx], mmap_mode='r')\n        y = np.load(self.output_files[file_idx], mmap_mode='r')\n\n        try:\n            return X[sample_idx].copy(), y[sample_idx].copy()\n        finally:\n            del X, y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T21:29:52.096944Z","iopub.execute_input":"2025-04-12T21:29:52.097179Z","iopub.status.idle":"2025-04-12T21:29:52.106523Z","shell.execute_reply.started":"2025-04-12T21:29:52.097157Z","shell.execute_reply":"2025-04-12T21:29:52.105603Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dstrain = SeismicDataset(train_inputs, train_outputs)\ndltrain = DataLoader(dstrain, batch_size=500, shuffle=True, pin_memory=True, drop_last=True, num_workers=4, persistent_workers=True)\n\ndsvalid = SeismicDataset(valid_inputs, valid_outputs)\ndlvalid = DataLoader(dsvalid, batch_size=500, shuffle=False, pin_memory=True, drop_last=False, num_workers=4, persistent_workers=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T21:29:52.107444Z","iopub.execute_input":"2025-04-12T21:29:52.107715Z","iopub.status.idle":"2025-04-12T21:29:52.118533Z","shell.execute_reply.started":"2025-04-12T21:29:52.107697Z","shell.execute_reply":"2025-04-12T21:29:52.117842Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class hgnet_model(nn.Module):\n    def __init__(self, pool_size=(4, 4), input_chans=5, output_size=70 * 70):\n        super().__init__()\n        self.backbone = hgnetv2_b0(in_chans=input_chans)\n        self.avg = nn.AdaptiveAvgPool2d((4, 4))\n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(1024 * 4 * 4, 2048),\n            nn.GELU(),\n            nn.Dropout(0.2),\n\n            nn.Linear(2048, 1024),\n            nn.GELU(),\n            nn.Dropout(0.2),\n\n            nn.Linear(1024, output_size)\n        )\n\n    def forward(self, x):\n        bs = x.size(0)\n\n        x = self.backbone.forward_features(x)\n        x = self.avg(x)\n\n        x = self.classifier(x)\n        return x.view(bs, 1, 70, 70) * 1000 + 1500","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T21:29:52.119197Z","iopub.execute_input":"2025-04-12T21:29:52.119385Z","iopub.status.idle":"2025-04-12T21:29:52.126102Z","shell.execute_reply.started":"2025-04-12T21:29:52.119370Z","shell.execute_reply":"2025-04-12T21:29:52.125406Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = hgnet_model().to(device)\ncriterion = nn.L1Loss()\noptim = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T21:29:52.126818Z","iopub.execute_input":"2025-04-12T21:29:52.126984Z","iopub.status.idle":"2025-04-12T21:29:52.893691Z","shell.execute_reply.started":"2025-04-12T21:29:52.126971Z","shell.execute_reply":"2025-04-12T21:29:52.893059Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.MSELoss()  \nl1_criterion = nn.L1Loss() \n\nif torch.cuda.device_count() > 1:\n    print(f\"Using {torch.cuda.device_count()} GPUs!\")\n    model = nn.DataParallel(model, device_ids=[0, 1])\nmodel = model.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T21:29:52.895329Z","iopub.execute_input":"2025-04-12T21:29:52.895539Z","iopub.status.idle":"2025-04-12T21:29:52.905104Z","shell.execute_reply.started":"2025-04-12T21:29:52.895522Z","shell.execute_reply":"2025-04-12T21:29:52.904392Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_epochs = 4\npatience = 1\nbest_loss = float('inf')\ncounter = 0\nmin_delta = 1e-4\nhistory = []\n\nfor epoch in range(1, n_epochs + 1):\n    print(f'{Fore.YELLOW}[{epoch:02d}] Begin train{Style.RESET_ALL}')\n\n    # Train\n    model.train()\n    train_losses = []\n    train_l1_losses = []\n    train_l2_losses = []\n    train_avg_losses = []\n    train_prog = tqdm(dltrain, desc='train', leave=False)\n    count = 0\n    total = len(dltrain)\n    for inputs, targets in train_prog:\n        inputs = inputs.to(device).float()  # Cast to float32\n        targets = targets.to(device).float()  # Cast to float32\n        \n        optim.zero_grad()\n        outputs = model(inputs)\n        \n        # Compute losses\n        loss = criterion(outputs, targets)  # L2 loss\n        l1_loss = l1_criterion(outputs, targets)  # L1 loss\n        avg_loss = (loss + l1_loss) / 2  # Average of L1 and L2\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        optim.step()\n\n        train_losses.append(loss.item())\n        train_l1_losses.append(l1_loss.item())\n        train_l2_losses.append(loss.item())  # Same as main loss (L2)\n        train_avg_losses.append(avg_loss.item())\n        train_prog.set_description(f\"Train Loss: {np.mean(train_losses):0.4f} | Comp L1: {np.mean(train_l1_losses):0.4f}\")\n\n        if count%10 == 0:\n            print(f\"{count}/{total} Train Loss: {np.mean(train_losses):0.4f} | Comp L1: {np.mean(train_l1_losses):0.4f}\")\n        count += 1\n    print(f'{Fore.GREEN}Train Results:{Style.RESET_ALL}')\n    print(f'{Fore.GREEN}  Loss: {np.mean(train_losses):.5f}{Style.RESET_ALL}')\n    print(f'{Fore.GREEN}  Comp L1 Loss: {np.mean(train_l1_losses):.5f}{Style.RESET_ALL}')\n    print(f'{Fore.GREEN}  L2 Loss: {np.mean(train_l2_losses):.5f}{Style.RESET_ALL}')\n    #print(f'{Fore.GREEN}  Avg Loss: {np.mean(train_avg_losses):.5f}{Style.RESET_ALL}')\n\n    # Valid\n    model.eval()\n    valid_losses = []\n    valid_l1_losses = []\n    valid_l2_losses = []\n    valid_avg_losses = []\n    valid_prog = tqdm(dlvalid, desc='valid', leave=False)\n    for inputs, targets in valid_prog:\n        inputs = inputs.to(device).float()  # Cast to float32\n        targets = targets.to(device).float()  # Cast to float32\n\n        with torch.inference_mode():\n            outputs = model(inputs)\n        \n        # Compute losses\n        loss = criterion(outputs, targets)  # L2 loss\n        l1_loss = l1_criterion(outputs, targets)  # L1 loss\n        avg_loss = (loss + l1_loss) / 2  # Average of L1 and L2\n\n        valid_losses.append(loss.item())\n        valid_l1_losses.append(l1_loss.item())\n        valid_l2_losses.append(loss.item())  # Same as main loss (L2)\n        valid_avg_losses.append(avg_loss.item())\n        valid_prog.set_description(f\"Valid Loss: {np.mean(valid_losses):0.4f} | Comp L1: {np.mean(valid_l1_losses):0.4f}\")\n    \n    valid_loss_mean = np.mean(valid_losses)\n    \n    print(f'{Fore.CYAN}Valid Results:{Style.RESET_ALL}')\n    print(f'{Fore.CYAN}  Loss: {valid_loss_mean:.5f}{Style.RESET_ALL}')\n    print(f'{Fore.CYAN}  Comp L1 Loss: {np.mean(valid_l1_losses):.5f}{Style.RESET_ALL}')\n    print(f'{Fore.CYAN}  L2 Loss: {np.mean(valid_l2_losses):.5f}{Style.RESET_ALL}')\n    #print(f'{Fore.CYAN}  Avg Loss: {np.mean(valid_avg_losses):.5f}{Style.RESET_ALL}')\n\n    # Update history\n    history.append({\n        'train': {\n            'loss': np.mean(train_losses),\n            'l1': np.mean(train_l1_losses),\n            'l2': np.mean(train_l2_losses),\n            #'avg': np.mean(train_avg_losses)\n        },\n        'valid': {\n            'loss': valid_loss_mean,\n            'l1': np.mean(valid_l1_losses),\n            'l2': np.mean(valid_l2_losses),\n            #'avg': np.mean(valid_avg_losses)\n        }\n    })\n\n    # Early stop\n    if valid_loss_mean < best_loss - min_delta:\n        best_loss = valid_loss_mean\n        counter = 0\n        torch.save(model.module.state_dict(), \"best_model.pt\")\n    else:\n        counter += 1\n        if counter >= patience:\n            print(f'{Fore.RED}Early stopping triggered{Style.RESET_ALL}')\n            break\n\n    # Plot results\n    if epoch % 1 == 0:\n        fig, ax = plt.subplots(2, 3, figsize=(10, 5))\n        fig.suptitle(f'Epoch {epoch} | Valid Loss: {valid_loss_mean:.5f} | Comp L1: {np.mean(valid_l1_losses):.5f}')\n        for i in range(min(3, targets.size(0))):\n            ax[0, i].imshow(targets[i, 0].detach().cpu(), cmap='viridis')\n            ax[0, i].set_title('Ground Truth')\n            ax[1, i].imshow(outputs[i, 0].detach().cpu(), cmap='viridis')\n            ax[1, i].set_title('Prediction')\n        plt.tight_layout()\n        plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-12T21:29:52.905982Z","iopub.execute_input":"2025-04-12T21:29:52.906245Z","iopub.status.idle":"2025-04-12T21:40:45.445591Z","shell.execute_reply.started":"2025-04-12T21:29:52.906215Z","shell.execute_reply":"2025-04-12T21:40:45.444688Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}