{"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":11367935,"sourceType":"datasetVersion","datasetId":7116013},{"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":11376935,"sourceType":"datasetVersion","datasetId":7122866},{"sourceId":11377083,"sourceType":"datasetVersion","datasetId":7122981},{"sourceId":11569755,"sourceType":"datasetVersion","datasetId":7253661},{"sourceId":11887338,"sourceType":"datasetVersion","datasetId":7377931}],"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-27T02:49:41.922735Z","iopub.execute_input":"2025-05-27T02:49:41.923552Z","iopub.status.idle":"2025-05-27T02:49:41.928532Z","shell.execute_reply.started":"2025-05-27T02:49:41.923528Z","shell.execute_reply":"2025-05-27T02:49:41.92767Z"}},"outputs":[],"execution_count":131},{"cell_type":"code","source":"!pip install kagglehub\n\nimport kagglehub","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T02:49:41.929774Z","iopub.execute_input":"2025-05-27T02:49:41.930013Z","iopub.status.idle":"2025-05-27T02:49:44.841357Z","shell.execute_reply.started":"2025-05-27T02:49:41.929996Z","shell.execute_reply":"2025-05-27T02:49:44.840623Z"}},"outputs":[{"name":"stdout","text":"Requirement already satisfied: kagglehub in /usr/local/lib/python3.11/dist-packages (0.3.11)\nRequirement already satisfied: packaging in /usr/local/lib/python3.11/dist-packages (from kagglehub) (24.2)\nRequirement already satisfied: pyyaml in /usr/local/lib/python3.11/dist-packages (from kagglehub) (6.0.2)\nRequirement already satisfied: requests in /usr/local/lib/python3.11/dist-packages (from kagglehub) (2.32.3)\nRequirement already satisfied: tqdm in /usr/local/lib/python3.11/dist-packages (from kagglehub) (4.67.1)\nRequirement already satisfied: charset-normalizer<4,>=2 in /usr/local/lib/python3.11/dist-packages (from requests->kagglehub) (3.4.1)\nRequirement already satisfied: idna<4,>=2.5 in /usr/local/lib/python3.11/dist-packages (from requests->kagglehub) (3.10)\nRequirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.11/dist-packages (from requests->kagglehub) (2.3.0)\nRequirement already satisfied: certifi>=2017.4.17 in /usr/local/lib/python3.11/dist-packages (from requests->kagglehub) (2025.1.31)\n","output_type":"stream"}],"execution_count":132},{"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\n\ncfg.in_channels = 5       # hoặc 3 nếu ảnh RGB\ncfg.out_channels = 1      # thường là 1 với segmentation\ncfg.encoder_channels = (64, 128, 256, 512)\ncfg.decoder_channels = (256, 128, 64, 32)\ncfg.upsample_mode = \"deconv\"      # hoặc \"bilinear\"\ncfg.attention = None ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T02:49:44.842436Z","iopub.execute_input":"2025-05-27T02:49:44.842673Z","iopub.status.idle":"2025-05-27T02:49:44.848656Z","shell.execute_reply.started":"2025-05-27T02:49:44.84265Z","shell.execute_reply":"2025-05-27T02:49:44.847983Z"}},"outputs":[{"name":"stdout","text":"Overwriting _cfg.py\n","output_type":"stream"}],"execution_count":133},{"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-27T02:49:44.850281Z","iopub.execute_input":"2025-05-27T02:49:44.850568Z","iopub.status.idle":"2025-05-27T02:49:44.86328Z","shell.execute_reply.started":"2025-05-27T02:49:44.850548Z","shell.execute_reply":"2025-05-27T02:49:44.862681Z"}},"outputs":[],"execution_count":134},{"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-27T02:49:44.864027Z","iopub.execute_input":"2025-05-27T02:49:44.86426Z","iopub.status.idle":"2025-05-27T02:49:45.749689Z","shell.execute_reply.started":"2025-05-27T02:49:44.864245Z","shell.execute_reply":"2025-05-27T02:49:45.749104Z"}},"outputs":[{"name":"stdout","text":"Shape Change\n  (500, 5, 1000, 70) -> (500, 5, 72, 72)\n\nMemory Usage\n  700.0 MB -> 25.9 MB\n  (3.7% of original size)\n","output_type":"stream"}],"execution_count":135},{"cell_type":"code","source":"%%writefile _dataset.py\n\nimport os\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport torch\n\nclass CustomDataset(torch.utils.data.Dataset):\n    def __init__(self, cfg, mode=\"train\"):\n        self.cfg = cfg\n        self.mode = mode\n        self.data, self.labels, self.records = self.load_metadata()\n\n    def load_metadata(self):\n        # Load metadata\n        df = pd.read_csv(\"/kaggle/input/openfwi-preprocessed-72x72/folds.csv\")\n        \n        # Chỉ lấy dữ liệu từ CurveFault_A\n        df = df[df[\"dataset\"] == \"CurveFault_A\"]\n\n        # Áp dụng subsample nếu có\n        if self.cfg.subsample is not None:\n            df = df.groupby([\"dataset\", \"fold\"]).head(self.cfg.subsample)\n\n        # Chia train/val\n        if self.mode == \"train\":\n            df = df[df[\"fold\"] != 0]\n        else:\n            df = df[df[\"fold\"] == 0]\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 file dữ liệu và nhãn\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            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        if self.mode == \"train\":\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-27T02:49:45.750379Z","iopub.execute_input":"2025-05-27T02:49:45.750572Z","iopub.status.idle":"2025-05-27T02:49:45.756093Z","shell.execute_reply.started":"2025-05-27T02:49:45.750557Z","shell.execute_reply":"2025-05-27T02:49:45.75536Z"}},"outputs":[{"name":"stdout","text":"Overwriting _dataset.py\n","output_type":"stream"}],"execution_count":136},{"cell_type":"code","source":"# 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])]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T02:49:45.756862Z","iopub.execute_input":"2025-05-27T02:49:45.757634Z","iopub.status.idle":"2025-05-27T02:49:45.772057Z","shell.execute_reply.started":"2025-05-27T02:49:45.757614Z","shell.execute_reply":"2025-05-27T02:49:45.77136Z"}},"outputs":[{"execution_count":137,"output_type":"execute_result","data":{"text/plain":"[torch.Size([36, 36]),\n torch.Size([18, 18]),\n torch.Size([9, 9]),\n torch.Size([9, 9])]"},"metadata":{}}],"execution_count":137},{"cell_type":"code","source":"# %%writefile _model_unet.py\n\n# import torch\n# import torch.nn as nn\n# import torch.nn.functional as F\n# from monai.networks.blocks import UpSample, SubpixelUpsample\n\n# ###################\n# ## Core U-Net Components ##\n# ###################\n\n# class ConvBnAct2d(nn.Module):\n#     def __init__(\n#         self,\n#         in_channels,\n#         out_channels,\n#         kernel_size,\n#         padding=0,\n#         stride=1,\n#         norm_layer=nn.BatchNorm2d,\n#         act_layer=nn.ReLU,\n#     ):\n#         super().__init__()\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# class 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.ReLU(inplace=True),\n#             nn.Conv2d(in_channels // reduction, in_channels, 1),\n#             nn.Sigmoid(),\n#         )\n#         self.sSE = nn.Sequential(nn.Conv2d(in_channels, 1, 1), nn.Sigmoid())\n\n#     def forward(self, x):\n#         return x * self.cSE(x) + x * self.sSE(x)\n\n# class 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(f\"Attention {name} is not implemented\")\n\n#     def forward(self, x):\n#         return self.attention(x)\n\n# class DecoderBlock2d(nn.Module):\n#     def __init__(\n#         self,\n#         in_channels,\n#         skip_channels,\n#         out_channels,\n#         norm_layer=nn.BatchNorm2d,\n#         attention_type=None,\n#         intermediate_conv=False,\n#         upsample_mode=\"deconv\",\n#         scale_factor=2,\n#     ):\n#         super().__init__()\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, norm_layer=norm_layer),\n#                 ConvBnAct2d(c, c, k, k // 2, norm_layer=norm_layer),\n#             )\n#         else:\n#             self.intermediate_conv = None\n\n#         self.attention1 = Attention2d(name=attention_type, in_channels=in_channels + skip_channels)\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#         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(name=attention_type, in_channels=out_channels)\n\n#     def forward(self, x, skip=None):\n#         x = self.upsample(x)\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#         if skip is not None:\n#             x = torch.cat([x, skip], dim=1)\n#             x = self.attention1(x)\n#         x = self.conv1(x)\n#         x = self.conv2(x)\n#         x = self.attention2(x)\n#         return x\n\n# class UnetDecoder2d(nn.Module):\n#     def __init__(\n#         self,\n#         encoder_channels,                 # eg: (64, 128, 256, 512)\n#         decoder_channels=(256, 128, 64, 32),\n#         scale_factors=(2, 2, 2, 2),\n#         norm_layer=nn.BatchNorm2d,\n#         attention_type=None,\n#         intermediate_conv=True,\n#         upsample_mode=\"deconv\",\n#     ):\n#         super().__init__()\n\n#         # đảm bảo đúng chiều: sâu nhất ở đầu\n#         encoder_channels = list(encoder_channels)\n#         encoder_channels = encoder_channels[::-1]   # now: [512, 256, 128, 64]\n\n#         # skip channels là mọi thứ trừ tầng đáy\n#         skip_channels = encoder_channels[1:] + [0]  # [256, 128, 64, 0]\n#         in_channels = [encoder_channels[0]] + list(decoder_channels[:-1])  # [512, 256, 128, 64]\n\n#         self.blocks = nn.ModuleList()\n#         for i, (ic, sc, dc) in enumerate(zip(in_channels, skip_channels, decoder_channels)):\n#             self.blocks.append(\n#                 DecoderBlock2d(\n#                     in_channels=ic,\n#                     skip_channels=sc,\n#                     out_channels=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):   # feats = [bottom] + skip1 + skip2 + ...\n#         x = feats[0]            # đáy encoder\n#         skips = feats[1:]       # các skip connection\n#         res = []\n#         for i, block in enumerate(self.blocks):\n#             skip = skips[i] if i < len(skips) else None\n#             x = block(x, skip)\n#             res.append(x)\n#         return res\n\n\n# class SegmentationHead2d(nn.Module):\n#     def __init__(\n#         self,\n#         in_channels,\n#         out_channels,\n#         scale_factor=2,\n#         kernel_size=3,\n#         mode=\"bilinear\",\n#     ):\n#         super().__init__()\n#         self.conv = nn.Conv2d(\n#             in_channels,\n#             out_channels,\n#             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=\"deconv\",\n#         )\n\n#     def forward(self, x):\n#         x = self.conv(x)\n#         x = self.upsample(x)\n#         return x\n\n# class EncoderBlock(nn.Module):\n#     def __init__(self, in_channels, out_channels, norm_layer=nn.BatchNorm2d):\n#         super().__init__()\n#         self.conv1 = ConvBnAct2d(in_channels, out_channels, kernel_size=3, padding=1, norm_layer=norm_layer)\n#         self.conv2 = ConvBnAct2d(out_channels, out_channels, kernel_size=3, padding=1, norm_layer=norm_layer)\n#         self.pool = nn.MaxPool2d(kernel_size=2, stride=2)\n\n#     def forward(self, x):\n#         x = self.conv1(x)\n#         x = self.conv2(x)\n#         skip = x\n#         x = self.pool(x)\n#         return x, skip\n\n# class Unet(nn.Module):\n#     def __init__(\n#         self,\n#         in_channels=5,\n#         out_channels=1,\n#         encoder_channels=(64, 128, 256, 512),\n#         decoder_channels=(256, 128, 64, 32),\n#         norm_layer=nn.BatchNorm2d,\n#         upsample_mode=\"deconv\",\n#         attention_type=None,\n#     ):\n#         super().__init__()\n#         # Encoder\n#         self.encoder_blocks = nn.ModuleList()\n#         prev_channels = in_channels\n#         for ch in encoder_channels:\n#             self.encoder_blocks.append(EncoderBlock(prev_channels, ch, norm_layer=norm_layer))\n#             prev_channels = ch\n\n#         # Decoder\n#         self.decoder = UnetDecoder2d(\n#             encoder_channels=encoder_channels,\n#             decoder_channels=decoder_channels,\n#             norm_layer=norm_layer,\n#             attention_type=attention_type,\n#             upsample_mode=upsample_mode,\n#         )\n\n#         # Segmentation Head\n#         self.seg_head = SegmentationHead2d(\n#             in_channels=decoder_channels[-1],\n#             out_channels=out_channels,\n#             scale_factor=2,\n#         )\n\n#     def forward(self, x):\n#         # Encoder\n#         skips = []\n#         for block in self.encoder_blocks:\n#             x, skip = block(x)\n#             skips.append(skip)\n\n#         # Chuẩn hóa thứ tự: sâu nhất ở đầu\n#         x = self.decoder([x] + skips[::-1])\n\n#         # Segmentation Head\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#         return x_seg","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T02:49:45.772897Z","iopub.execute_input":"2025-05-27T02:49:45.77321Z","iopub.status.idle":"2025-05-27T02:49:45.784567Z","shell.execute_reply.started":"2025-05-27T02:49:45.773191Z","shell.execute_reply":"2025-05-27T02:49:45.783707Z"}},"outputs":[],"execution_count":138},{"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-27T02:49:45.786352Z","iopub.execute_input":"2025-05-27T02:49:45.786554Z","iopub.status.idle":"2025-05-27T02:49:45.800043Z","shell.execute_reply.started":"2025-05-27T02:49:45.786531Z","shell.execute_reply":"2025-05-27T02:49:45.799484Z"}},"outputs":[{"name":"stdout","text":"Writing _model.py\n","output_type":"stream"}],"execution_count":139},{"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-27T02:49:45.800876Z","iopub.execute_input":"2025-05-27T02:49:45.801124Z","iopub.status.idle":"2025-05-27T02:49:45.814494Z","shell.execute_reply.started":"2025-05-27T02:49:45.801108Z","shell.execute_reply":"2025-05-27T02:49:45.813909Z"}},"outputs":[{"name":"stdout","text":"Overwriting _utils.py\n","output_type":"stream"}],"execution_count":140},{"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-27T02:49:45.815253Z","iopub.execute_input":"2025-05-27T02:49:45.81552Z","iopub.status.idle":"2025-05-27T02:49:45.827969Z","shell.execute_reply.started":"2025-05-27T02:49:45.815496Z","shell.execute_reply":"2025-05-27T02:49:45.827453Z"}},"outputs":[{"name":"stdout","text":"Overwriting _train.py\n","output_type":"stream"}],"execution_count":141},{"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-27T02:49:45.82872Z","iopub.execute_input":"2025-05-27T02:49:45.828966Z"}},"outputs":[{"name":"stdout","text":"Starting training..\n2025-05-27 02:49:56.089462: E external/local_xla/xla/stream_executor/cuda/cuda_fft.cc:477] Unable to register cuFFT factory: Attempting to register factory for plugin cuFFT when one has already been registered\nWARNING: All log messages before absl::InitializeLog() is called are written to STDERR\nE0000 00:00:1748314196.110896     653 cuda_dnn.cc:8310] Unable to register cuDNN factory: Attempting to register factory for plugin cuDNN when one has already been registered\nE0000 00:00:1748314196.117636     653 cuda_blas.cc:1418] Unable to register cuBLAS factory: Attempting to register factory for plugin cuBLAS when one has already been registered\n2025-05-27 02:49:56.394088: E external/local_xla/xla/stream_executor/cuda/cuda_fft.cc:477] Unable to register cuFFT factory: Attempting to register factory for plugin cuFFT when one has already been registered\nWARNING: All log messages before absl::InitializeLog() is called are written to STDERR\nE0000 00:00:1748314196.415719     654 cuda_dnn.cc:8310] Unable to register cuDNN factory: Attempting to register factory for plugin cuDNN when one has already been registered\nE0000 00:00:1748314196.422444     654 cuda_blas.cc:1418] Unable to register cuBLAS factory: Attempting to register factory for plugin cuBLAS when one has already been registered\nRank: 0, World size: 2, GPU memory: 14.74GB\nRank: 1, World size: 2, GPU memory: 14.74GB\n=========================\nLoading data..\n100%|████████████████████████████████████████| 106/106 [00:00<00:00, 212.96it/s]\n100%|█████████████████████████████████████████████| 2/2 [00:00<00:00, 76.47it/s]\nmodel.safetensors: 100%|████████████████████| 45.0M/45.0M [00:00<00:00, 172MB/s]\nInitializing EMA model..\n=========================\nGive me warp 2, Mr. Sulu.\n=========================\n100%|█████████████████████████████████████████████| 4/4 [00:04<00:00,  1.02s/it]\nNew best: 1000000.00 -> 734.29\nSaved weights..\nEpoch 1:     Train MAE: 756.69     Val MAE: 734.29     Time: 0:00:18     Step: 1/53\n100%|█████████████████████████████████████████████| 4/4 [00:01<00:00,  3.25it/s]\nNew best: 734.29 -> 703.26\nSaved weights..\nEpoch 2:     Train MAE: 273.29     Val MAE: 703.26     Time: 0:00:02     Step: 1/53\n100%|█████████████████████████████████████████████| 4/4 [00:01<00:00,  3.27it/s]\nNew best: 703.26 -> 641.98\nSaved weights..\nEpoch 3:     Train MAE: 202.42     Val MAE: 641.98     Time: 0:00:02     Step: 1/53\n100%|█████████████████████████████████████████████| 4/4 [00:01<00:00,  3.34it/s]\nNew best: 641.98 -> 569.99\nSaved weights..\nEpoch 4:     Train MAE: 190.75     Val MAE: 569.99     Time: 0:00:02     Step: 1/53\n100%|█████████████████████████████████████████████| 4/4 [00:01<00:00,  3.30it/s]\nNew best: 569.99 -> 529.64\nSaved weights..\nEpoch 5:     Train MAE: 142.47     Val MAE: 529.64     Time: 0:00:02     Step: 1/53\n100%|█████████████████████████████████████████████| 4/4 [00:01<00:00,  3.23it/s]\nNew best: 529.64 -> 510.14\nSaved weights..\nEpoch 6:     Train MAE: 175.72     Val MAE: 510.14     Time: 0:00:02     Step: 1/53\n100%|█████████████████████████████████████████████| 4/4 [00:01<00:00,  3.21it/s]\nEpoch 7:     Train MAE: 194.85     Val MAE: 541.71     Time: 0:00:02     Step: 1/53\n100%|█████████████████████████████████████████████| 4/4 [00:01<00:00,  3.24it/s]\nEpoch 8:     Train MAE: 145.64     Val MAE: 564.00     Time: 0:00:02     Step: 1/53\n100%|█████████████████████████████████████████████| 4/4 [00:01<00:00,  3.27it/s]\nEpoch 9:     Train MAE: 118.81     Val MAE: 546.78     Time: 0:00:02     Step: 1/53\n100%|█████████████████████████████████████████████| 4/4 [00:01<00:00,  3.33it/s]\nNew best: 510.14 -> 510.07\nSaved weights..\nEpoch 10:     Train MAE: 156.50     Val MAE: 510.07     Time: 0:00:02     Step: 1/53\n100%|█████████████████████████████████████████████| 4/4 [00:01<00:00,  3.24it/s]\nNew best: 510.07 -> 477.35\nSaved weights..\n","output_type":"stream"}],"execution_count":null},{"cell_type":"code","source":"model = Net(backbone=\"resnet18\", pretrained=False)\n\n# In tổng số lượng tham số\ntotal_params = sum(p.numel() for p in model.parameters())\nprint(f\"Tổng số tham số: {total_params:,}\")\n\n# In chi tiết từng tham số\nfor name, param in model.named_parameters():\n    print(f\"{name}: shape={tuple(param.shape)}, requires_grad={param.requires_grad}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T03:06:19.538465Z","iopub.execute_input":"2025-05-27T03:06:19.539049Z","iopub.status.idle":"2025-05-27T03:06:19.735663Z","shell.execute_reply.started":"2025-05-27T03:06:19.539027Z","shell.execute_reply":"2025-05-27T03:06:19.735049Z"}},"outputs":[{"name":"stdout","text":"Tổng số tham số: 16,554,913\nbackbone.conv1.weight: shape=(64, 5, 7, 7), requires_grad=True\nbackbone.bn1.weight: shape=(64,), requires_grad=True\nbackbone.bn1.bias: shape=(64,), requires_grad=True\nbackbone.layer1.0.conv1.weight: shape=(64, 64, 3, 3), requires_grad=True\nbackbone.layer1.0.bn1.weight: shape=(64,), requires_grad=True\nbackbone.layer1.0.bn1.bias: shape=(64,), requires_grad=True\nbackbone.layer1.0.conv2.weight: shape=(64, 64, 3, 3), requires_grad=True\nbackbone.layer1.0.bn2.weight: shape=(64,), requires_grad=True\nbackbone.layer1.0.bn2.bias: shape=(64,), requires_grad=True\nbackbone.layer1.1.conv1.weight: shape=(64, 64, 3, 3), requires_grad=True\nbackbone.layer1.1.bn1.weight: shape=(64,), requires_grad=True\nbackbone.layer1.1.bn1.bias: shape=(64,), requires_grad=True\nbackbone.layer1.1.conv2.weight: shape=(64, 64, 3, 3), requires_grad=True\nbackbone.layer1.1.bn2.weight: shape=(64,), requires_grad=True\nbackbone.layer1.1.bn2.bias: shape=(64,), requires_grad=True\nbackbone.layer2.0.conv1.weight: shape=(128, 64, 3, 3), requires_grad=True\nbackbone.layer2.0.bn1.weight: shape=(128,), requires_grad=True\nbackbone.layer2.0.bn1.bias: shape=(128,), requires_grad=True\nbackbone.layer2.0.conv2.weight: shape=(128, 128, 3, 3), requires_grad=True\nbackbone.layer2.0.bn2.weight: shape=(128,), requires_grad=True\nbackbone.layer2.0.bn2.bias: shape=(128,), requires_grad=True\nbackbone.layer2.0.downsample.0.weight: shape=(128, 64, 1, 1), requires_grad=True\nbackbone.layer2.0.downsample.1.weight: shape=(128,), requires_grad=True\nbackbone.layer2.0.downsample.1.bias: shape=(128,), requires_grad=True\nbackbone.layer2.1.conv1.weight: shape=(128, 128, 3, 3), requires_grad=True\nbackbone.layer2.1.bn1.weight: shape=(128,), requires_grad=True\nbackbone.layer2.1.bn1.bias: shape=(128,), requires_grad=True\nbackbone.layer2.1.conv2.weight: shape=(128, 128, 3, 3), requires_grad=True\nbackbone.layer2.1.bn2.weight: shape=(128,), requires_grad=True\nbackbone.layer2.1.bn2.bias: shape=(128,), requires_grad=True\nbackbone.layer3.0.conv1.weight: shape=(256, 128, 3, 3), requires_grad=True\nbackbone.layer3.0.bn1.weight: shape=(256,), requires_grad=True\nbackbone.layer3.0.bn1.bias: shape=(256,), requires_grad=True\nbackbone.layer3.0.conv2.weight: shape=(256, 256, 3, 3), requires_grad=True\nbackbone.layer3.0.bn2.weight: shape=(256,), requires_grad=True\nbackbone.layer3.0.bn2.bias: shape=(256,), requires_grad=True\nbackbone.layer3.0.downsample.0.weight: shape=(256, 128, 1, 1), requires_grad=True\nbackbone.layer3.0.downsample.1.weight: shape=(256,), requires_grad=True\nbackbone.layer3.0.downsample.1.bias: shape=(256,), requires_grad=True\nbackbone.layer3.1.conv1.weight: shape=(256, 256, 3, 3), requires_grad=True\nbackbone.layer3.1.bn1.weight: shape=(256,), requires_grad=True\nbackbone.layer3.1.bn1.bias: shape=(256,), requires_grad=True\nbackbone.layer3.1.conv2.weight: shape=(256, 256, 3, 3), requires_grad=True\nbackbone.layer3.1.bn2.weight: shape=(256,), requires_grad=True\nbackbone.layer3.1.bn2.bias: shape=(256,), requires_grad=True\nbackbone.layer4.0.conv1.weight: shape=(512, 256, 3, 3), requires_grad=True\nbackbone.layer4.0.bn1.weight: shape=(512,), requires_grad=True\nbackbone.layer4.0.bn1.bias: shape=(512,), requires_grad=True\nbackbone.layer4.0.conv2.weight: shape=(512, 512, 3, 3), requires_grad=True\nbackbone.layer4.0.bn2.weight: shape=(512,), requires_grad=True\nbackbone.layer4.0.bn2.bias: shape=(512,), requires_grad=True\nbackbone.layer4.0.downsample.0.weight: shape=(512, 256, 1, 1), requires_grad=True\nbackbone.layer4.0.downsample.1.weight: shape=(512,), requires_grad=True\nbackbone.layer4.0.downsample.1.bias: shape=(512,), requires_grad=True\nbackbone.layer4.1.conv1.weight: shape=(512, 512, 3, 3), requires_grad=True\nbackbone.layer4.1.bn1.weight: shape=(512,), requires_grad=True\nbackbone.layer4.1.bn1.bias: shape=(512,), requires_grad=True\nbackbone.layer4.1.conv2.weight: shape=(512, 512, 3, 3), requires_grad=True\nbackbone.layer4.1.bn2.weight: shape=(512,), requires_grad=True\nbackbone.layer4.1.bn2.bias: shape=(512,), requires_grad=True\ndecoder.blocks.0.upsample.deconv.weight: shape=(512, 512, 1, 1), requires_grad=True\ndecoder.blocks.0.upsample.deconv.bias: shape=(512,), requires_grad=True\ndecoder.blocks.0.intermediate_conv.0.conv.weight: shape=(256, 256, 3, 3), requires_grad=True\ndecoder.blocks.0.intermediate_conv.1.conv.weight: shape=(256, 256, 3, 3), requires_grad=True\ndecoder.blocks.0.conv1.conv.weight: shape=(256, 768, 3, 3), requires_grad=True\ndecoder.blocks.0.conv2.conv.weight: shape=(256, 256, 3, 3), requires_grad=True\ndecoder.blocks.1.upsample.deconv.weight: shape=(256, 256, 2, 2), requires_grad=True\ndecoder.blocks.1.upsample.deconv.bias: shape=(256,), requires_grad=True\ndecoder.blocks.1.intermediate_conv.0.conv.weight: shape=(128, 128, 3, 3), requires_grad=True\ndecoder.blocks.1.intermediate_conv.1.conv.weight: shape=(128, 128, 3, 3), requires_grad=True\ndecoder.blocks.1.conv1.conv.weight: shape=(128, 384, 3, 3), requires_grad=True\ndecoder.blocks.1.conv2.conv.weight: shape=(128, 128, 3, 3), requires_grad=True\ndecoder.blocks.2.upsample.deconv.weight: shape=(128, 128, 2, 2), requires_grad=True\ndecoder.blocks.2.upsample.deconv.bias: shape=(128,), requires_grad=True\ndecoder.blocks.2.intermediate_conv.0.conv.weight: shape=(64, 64, 3, 3), requires_grad=True\ndecoder.blocks.2.intermediate_conv.1.conv.weight: shape=(64, 64, 3, 3), requires_grad=True\ndecoder.blocks.2.conv1.conv.weight: shape=(64, 192, 3, 3), requires_grad=True\ndecoder.blocks.2.conv2.conv.weight: shape=(64, 64, 3, 3), requires_grad=True\ndecoder.blocks.3.upsample.deconv.weight: shape=(64, 64, 2, 2), requires_grad=True\ndecoder.blocks.3.upsample.deconv.bias: shape=(64,), requires_grad=True\ndecoder.blocks.3.intermediate_conv.0.conv.weight: shape=(64, 64, 3, 3), requires_grad=True\ndecoder.blocks.3.intermediate_conv.1.conv.weight: shape=(64, 64, 3, 3), requires_grad=True\ndecoder.blocks.3.conv1.conv.weight: shape=(32, 128, 3, 3), requires_grad=True\ndecoder.blocks.3.conv2.conv.weight: shape=(32, 32, 3, 3), requires_grad=True\nseg_head.conv.weight: shape=(1, 32, 3, 3), requires_grad=True\nseg_head.conv.bias: shape=(1,), requires_grad=True\n","output_type":"stream"}],"execution_count":144},{"cell_type":"code","source":"import torch\nfrom tqdm import tqdm\nfrom _cfg import cfg\nfrom _dataset import CustomDataset\nfrom _model import Net\n\n# Thiết bị\ncfg.device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n\n# Load mô hình và trọng số đã lưu\nmodel = Net(backbone=cfg.backbone).to(cfg.device)\nmodel.load_state_dict(torch.load(f'best_model_{cfg.seed}.pt', map_location=cfg.device))\nmodel.eval()\n\n# Tập validation\nval_ds = CustomDataset(cfg=cfg, mode=\"valid\")\nval_dl = torch.utils.data.DataLoader(val_ds, batch_size=cfg.batch_size_val)\n\n# Dự đoán\nval_preds, val_targets = [], []\nwith torch.no_grad():\n    for x, y in tqdm(val_dl):\n        x, y = x.to(cfg.device), y.to(cfg.device)\n\n        # Sửa lỗi autocast\n        if cfg.device.type == 'cuda':\n            autocast_ctx = torch.cuda.amp.autocast()\n        else:\n            autocast_ctx = torch.amp.autocast(device_type='cpu')\n\n        with autocast_ctx:\n            out = model(x)\n\n        val_preds.append(out.cpu())\n        val_targets.append(y.cpu())\n\n# Tính MAE\nval_preds = torch.cat(val_preds, dim=0)\nval_targets = torch.cat(val_targets, dim=0)\nmae = torch.nn.L1Loss()(val_preds, val_targets).item()\nprint(f\"Final validation MAE: {mae:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-27T03:06:25.311151Z","iopub.execute_input":"2025-05-27T03:06:25.311986Z","iopub.status.idle":"2025-05-27T03:06:27.920208Z","shell.execute_reply.started":"2025-05-27T03:06:25.311943Z","shell.execute_reply":"2025-05-27T03:06:27.919435Z"}},"outputs":[{"name":"stderr","text":"/tmp/ipykernel_31/1971169604.py:12: FutureWarning: You are using `torch.load` with `weights_only=False` (the current default value), which uses the default pickle module implicitly. It is possible to construct malicious pickle data which will execute arbitrary code during unpickling (See https://github.com/pytorch/pytorch/blob/main/SECURITY.md#untrusted-models for more details). In a future release, the default value for `weights_only` will be flipped to `True`. This limits the functions that could be executed during unpickling. Arbitrary objects will no longer be allowed to be loaded via this mode unless they are explicitly allowlisted by the user via `torch.serialization.add_safe_globals`. We recommend you start setting `weights_only=True` for any use case where you don't have full control of the loaded file. Please open an issue on GitHub for any issues related to this experimental feature.\n  model.load_state_dict(torch.load(f'best_model_{cfg.seed}.pt', map_location=cfg.device))\n100%|██████████| 2/2 [00:00<00:00, 79.62it/s]\n  0%|          | 0/8 [00:00<?, ?it/s]/tmp/ipykernel_31/1971169604.py:27: FutureWarning: `torch.cuda.amp.autocast(args...)` is deprecated. Please use `torch.amp.autocast('cuda', args...)` instead.\n  autocast_ctx = torch.cuda.amp.autocast()\n100%|██████████| 8/8 [00:01<00:00,  4.02it/s]","output_type":"stream"},{"name":"stdout","text":"Final validation MAE: 476.4735\n","output_type":"stream"},{"name":"stderr","text":"\n","output_type":"stream"}],"execution_count":145}]}