{"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":36363,"databundleVersionId":4050810,"sourceType":"competition"}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# 📌 Cell 0  –  chỉ cài thứ pipeline cần, KHÔNG đụng torch/numpy\n!pip install -q --no-deps \\\n    iterative-stratification \\\n    kornia==0.7.2 \\\n    albumentations==1.4.3 \\\n    pydicom pylibjpeg pylibjpeg-libjpeg pylibjpeg-openjpeg \\\n    nibabel tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T15:54:43.44626Z","iopub.execute_input":"2025-05-13T15:54:43.446492Z","iopub.status.idle":"2025-05-13T15:54:48.491265Z","shell.execute_reply.started":"2025-05-13T15:54:43.446474Z","shell.execute_reply":"2025-05-13T15:54:48.490335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy, scipy, sklearn\nprint(\"NumPy :\", numpy.__version__)\nprint(\"SciPy :\", scipy.__version__)\nprint(\"sklearn:\", sklearn.__version__)\n# Kỳ vọng: 1.26.x  /  1.11.x  /  1.3.x\nfrom iterstrat.ml_stratifiers import MultilabelStratifiedShuffleSplit\nprint(\"iterstrat OK ✓\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T15:54:48.492147Z","iopub.execute_input":"2025-05-13T15:54:48.492365Z","iopub.status.idle":"2025-05-13T15:54:49.029029Z","shell.execute_reply.started":"2025-05-13T15:54:48.492344Z","shell.execute_reply":"2025-05-13T15:54:49.028214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════╗\n# 📌 Cell 1 – Config dùng dữ liệu /kaggle/input\n# ╚══════════════════════════════════════════════╝\n\nimport torch, torchvision, nibabel as nib, numpy as np, pandas as pd, random\nfrom pathlib import Path\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom tqdm.auto import tqdm\nfrom tqdm.notebook import tqdm\n\n# ĐƯỜNG DẪN GỐC ĐÃ GẮN (readonly – không tốn quota 20 GB)\nROOT = Path(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection\")\nCSV  = ROOT/\"train.csv\"\nIMG_DIR = ROOT/\"train_images\"        # chứa thư mục UID/*.dcm\nMASK_DIR = ROOT/\"segmentations\"      # đã giải nén sẵn\n\n# CẤU HÌNH GIỐNG BẢN COLAB\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\n\nPERCENT    = 0.25      # giữ 30 % study để train nhanh\nVAL_SPLIT  = 0.15\nIMG_SIZE   = 256\nROI_SIZE   = 224\nBS_SEG     = 8\nBS_CLS     = 16\nEPOCH_SEG  = 5\nEPOCH_CLS  = 10\ndevice_name = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nDEVICE = torch.device(device_name)\n\nprint(\"Device: \", DEVICE)\nprint(\"✅  Dataset root:\", ROOT)\n\nn_mask = len(list(MASK_DIR.glob(\"*.nii\")))\nprint(\"🟢  Số file mask:\", n_mask)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T15:54:49.030856Z","iopub.execute_input":"2025-05-13T15:54:49.031175Z","iopub.status.idle":"2025-05-13T15:54:57.173161Z","shell.execute_reply.started":"2025-05-13T15:54:49.031157Z","shell.execute_reply":"2025-05-13T15:54:57.172464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Cell 3 – Chọn 15 % Study & Phân tầng đa nhãn (với mask luôn vào train)\nfrom iterstrat.ml_stratifiers import MultilabelStratifiedShuffleSplit\n\ndf = pd.read_csv(CSV)\nlabel_cols = [f\"C{i}\" for i in range(1,8)]\nY = df[label_cols].values\nuids = df[\"StudyInstanceUID\"].values\n\n# Tập UID đã có mask\nmask_uids = {p.stem for p in MASK_DIR.glob(\"*.nii\")}\n\n# Tách 2 nhóm: có mask ↴ không có mask\nuids_with_mask    = [uid for uid in uids if uid in mask_uids]\nuids_without_mask = [uid for uid in uids if uid not in mask_uids]\nY_without_mask    = df.loc[df[\"StudyInstanceUID\"].isin(uids_without_mask), label_cols].values\n\n# Lấy 15% stratified trên nhóm không có mask\nsss15 = MultilabelStratifiedShuffleSplit(n_splits=1, test_size=1-PERCENT, random_state=SEED)\nidx_without, _ = next(sss15.split(uids_without_mask, Y_without_mask))\nsampled_uids_without = {uids_without_mask[i] for i in idx_without}\n\n# Ghép chung: all mask_uids + sampled non-mask_uids\nsub_uids = set(uids_with_mask) | sampled_uids_without\ndf_sub   = df[df[\"StudyInstanceUID\"].isin(sub_uids)].reset_index(drop=True)\nprint(f\"💾  Selected {len(df_sub)} / {len(df)} study \"\n      f\"({len(uids_with_mask)} mask + {len(sampled_uids_without)} sampled)\")\n\n# Rồi chia tiếp train/val stratified trên df_sub\nsss20 = MultilabelStratifiedShuffleSplit(n_splits=1, test_size=VAL_SPLIT, random_state=SEED)\ntrain_idx, val_idx = next(sss20.split(df_sub[\"StudyInstanceUID\"], df_sub[label_cols]))\ntrain_uids = set(df_sub.loc[train_idx, \"StudyInstanceUID\"])\nval_uids   = set(df_sub.loc[val_idx,   \"StudyInstanceUID\"])\nprint(f\"Train splits: {len(train_uids)} studies (bao gồm {len(mask_uids & train_uids)} mask)\")\n\n# **Thêm**: in ra số lượng validation và phân phối\nprint(f\"Validation splits: {len(val_uids)} studies (bao gồm {len(mask_uids & val_uids)} mask)\")\n\n# Nếu cần xem thêm positives trên validation thì in tiếp:\npos_val = df_sub.loc[val_idx, label_cols].sum().to_dict()\nprint(\"Positives trên validation:\", pos_val)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T15:54:57.174008Z","iopub.execute_input":"2025-05-13T15:54:57.174472Z","iopub.status.idle":"2025-05-13T15:54:57.243572Z","shell.execute_reply.started":"2025-05-13T15:54:57.174446Z","shell.execute_reply":"2025-05-13T15:54:57.242859Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Cell X – Tạo DataLoader với WeightedRandomSampler để cân bằng multi-label\nimport numpy as np\nfrom torch.utils.data import DataLoader, WeightedRandomSampler\nimport kornia.augmentation as K\naug_cls = K.AugmentationSequential(\n    K.RandomHorizontalFlip(p=0.5),\n    K.RandomAffine(degrees=10, scale=(0.9, 1.1), translate=(0.05, 0.05)),\n    K.Normalize(mean=torch.tensor([0.5]), std=torch.tensor([0.5])),\n    data_keys=[\"input\"]\n)\nclass ClsStudyDS(Dataset):\n    def __init__(self, uids, transform=None):\n        self.uids = list(uids)\n        self.transform = transform\n        # df_sub là DataFrame con đã filter ở Cell 3, có cột StudyInstanceUID và C1–C7\n        self.df = df_sub.set_index(\"StudyInstanceUID\")\n\n    def _get_roi_sagittal(self, uid):\n        # --- 1) Nếu có segmentation mask: load .nii.gz ---\n        nii_path = MASK_DIR / f\"{uid}.nii.gz\"\n        if nii_path.exists():\n            mask = nib.load(str(nii_path)).get_fdata().astype(np.uint8)\n            mask = np.transpose(mask, (2,1,0))  # từ sagittal -> axial-aligned\n        else:\n            # --- 2) Nếu không, dùng U-Net để predict mask cho toàn bộ study ---\n            study_dir = IMG_DIR / uid\n            slices = sorted(study_dir.glob(\"*.dcm\"), key=lambda p: int(p.stem))\n            mask = np.zeros((len(slices), IMG_SIZE, IMG_SIZE), dtype=np.uint8)\n            for iz, p in enumerate(slices):\n                img_arr = cv2.resize(load_dcm(p), (IMG_SIZE, IMG_SIZE))\n                with torch.no_grad():\n                    t = torch.tensor(img_arr, dtype=torch.float32).unsqueeze(0).unsqueeze(0).to(DEVICE)\n                    pred = unet(t)                     # [1, n_cls, H, W]\n                    lab = pred.argmax(dim=1)[0].cpu().numpy().astype(np.uint8)\n                mask[iz] = lab\n\n        # --- 3) Tìm bounding box quanh tất cả label 1–7 ---\n        pos = np.where((mask >= 1) & (mask <= 7))\n        if len(pos[0]) == 0:\n            return None\n\n        zmin, zmax = pos[0].min(), pos[0].max()\n        ymin, ymax = pos[1].min(), pos[1].max()\n        xmin, xmax = pos[2].min(), pos[2].max()\n\n        # --- 4) Tạo sagittal MIP view (max projection) ---\n        sag = mask[:, :, xmin:xmax+1].max(axis=2)  # kết quả shape [Z, Y]\n        sag = (sag > 0).astype(np.float32)\n\n        # --- 5) Resize về ROI_SIZE × ROI_SIZE ---\n        sag = cv2.resize(sag, (ROI_SIZE, ROI_SIZE), interpolation=cv2.INTER_NEAREST)\n        return sag, zmin, zmax\n\n    def __len__(self):\n        return len(self.uids)\n\n    def __getitem__(self, idx):\n        uid = self.uids[idx]\n        roi_res = self._get_roi_sagittal(uid)\n\n        if roi_res is None:\n            # fallback ROI trống\n            roi = np.zeros((ROI_SIZE, ROI_SIZE), dtype=np.float32)\n        else:\n            roi, _, _ = roi_res\n\n        # --- Chuyển numpy ROI [H, W] -> torch.Tensor [C=1, H, W] ---\n        img = torch.tensor(roi[np.newaxis, :, :], dtype=torch.float32)\n\n        # --- Nếu có Kornia transform, thêm batch dim trước khi apply ---\n        if self.transform:\n            img = img.unsqueeze(0)            # [B=1, C=1, H, W]\n            img = self.transform(img)         # Kornia trả về [1,1,H,W]\n            img = img.squeeze(0)              # về [C=1, H, W]\n\n        # --- Lấy label multi-hot vector length=7 ---\n        label = torch.tensor(\n            self.df.loc[uid, label_cols].values.astype(np.float32),\n            dtype=torch.float32\n        )\n        return img, label\n# 1) Khởi tạo dataset như trước\ntrain_cls_ds = ClsStudyDS(train_uids, transform=aug_cls)\n\n# 2) Tính trọng số cho mỗi lớp (1 / số lượng positive trong train set)\n#    và gán weight cho từng sample bằng trung bình weight của các label=1 của nó\ndf_train = df_sub[df_sub[\"StudyInstanceUID\"].isin(train_uids)].reset_index(drop=True)\nlabel_counts = df_train[label_cols].sum().values  # shape (7,)\nclass_weights = 1.0 / (label_counts + 1e-6)        # tránh chia 0\n\nsample_weights = []\nfor uid in train_uids:\n    y = df_train.loc[df_train[\"StudyInstanceUID\"]==uid, label_cols].values.flatten()\n    # chỉ lấy những class mà sample thực sự positive\n    pos_idx = np.where(y==1)[0]\n    if len(pos_idx)==0:\n        # với sample không positive nào, cho weight = average của tất cả lớp\n        sample_weights.append(class_weights.mean())\n    else:\n        sample_weights.append(class_weights[pos_idx].mean())\n\n# 3) Tạo WeightedRandomSampler\nsampler = WeightedRandomSampler(weights=sample_weights,\n                                num_samples=len(sample_weights),\n                                replacement=True)\n\n# 4) Tạo DataLoader dùng sampler\ntrain_cls_ld = DataLoader(\n    train_cls_ds,\n    batch_size=BS_CLS,\n    sampler=sampler,\n    num_workers=0,\n    pin_memory=True\n)\n\n# 5) Tạo val loader như cũ\nval_cls_ld = DataLoader(\n    ClsStudyDS(val_uids, transform=K.AugmentationSequential(\n        K.Normalize(mean=torch.tensor([0.5]), std=torch.tensor([0.5])),\n        data_keys=[\"input\"]\n    )),\n    batch_size=BS_CLS,\n    shuffle=False,\n    num_workers=0,\n    pin_memory=True\n)\n\nprint(f\"✔️  train sampler với {len(sample_weights)} samples, sum weights={sum(sample_weights):.2f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T15:54:57.244383Z","iopub.execute_input":"2025-05-13T15:54:57.244707Z","iopub.status.idle":"2025-05-13T15:54:58.742566Z","shell.execute_reply.started":"2025-05-13T15:54:57.244681Z","shell.execute_reply":"2025-05-13T15:54:58.741838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Cell 4 – Dataset Segmentation 2D (U-Net)\nimport albumentations as A\nimport cv2\n\nprint(train_uids)\n\ndef load_dcm(path):\n    import pydicom, numpy as np\n    ds = pydicom.dcmread(str(path))\n    img = ds.pixel_array.astype(np.float32)\n    img = (img - img.min()) / (img.max() - img.min() + 1e-5)\n    return img\n\nclass SegSliceDS(Dataset):\n    def __init__(self, uids, transform=None):\n        self.samples=[]\n        for uid in tqdm(uids, desc=\"Build seg-slice list\"):\n            nii_path = MASK_DIR/f\"{uid}.nii\"\n            if not nii_path.exists():\n              print(f\"skipping {uid}\")\n              continue\n            print(f\"continuing with {uid}\")\n            mask = nib.load(str(nii_path)).get_fdata().astype(np.uint8)  # sagittal\n            mask = np.transpose(mask, (2,1,0))   # ≈ align axial\n            for z in range(mask.shape[0]):\n                if mask[z].max()==0: continue\n                dcm_path = IMG_DIR/uid/f\"{z+1}.dcm\"\n                if not dcm_path.exists(): continue\n                self.samples.append((dcm_path, mask[z]))\n        self.transform=transform\n    def __len__(self): return len(self.samples)\n    def __getitem__(self, idx):\n        dcm_path, m = self.samples[idx]\n        img = load_dcm(dcm_path)\n        m   = cv2.resize(m,(IMG_SIZE,IMG_SIZE),interpolation=cv2.INTER_NEAREST)\n        m[m>7] = 0\n        img = cv2.resize(img,(IMG_SIZE,IMG_SIZE))\n        if self.transform:\n            aug = self.transform(image=img, mask=m)\n            img, m = aug[\"image\"], aug[\"mask\"]\n        img = torch.tensor(img).unsqueeze(0).float()       # [1,H,W]\n        m   = torch.tensor(m).long()                       # CE Loss đa lớp\n        return img, m\n\naug_seg = A.Compose([A.HorizontalFlip(p=0.5)], additional_targets={'mask':'mask'})\n\nmask_uids = {p.stem for p in MASK_DIR.glob(\"*.nii\")}   # set 900 UID\ntrain_mask_uids = train_uids & mask_uids\nval_mask_uids   = val_uids   & mask_uids\nprint(f\"✅ train seg UID: {len(train_mask_uids)},  val seg UID: {len(val_mask_uids)}\")\n\n# ---- Dataset Segmentation 2D ----\ntrain_seg_ds = SegSliceDS(train_mask_uids, transform=aug_seg)\nval_seg_ds   = SegSliceDS(val_mask_uids,   transform=None)\n\n# Nếu vẫn =0 → bỏ qua bước U‑Net\nif len(train_seg_ds)==0:\n    print(\"⚠️ Không có study có mask → bỏ qua Cell 5,6 liên quan segmentation\")\nelse:\n    train_seg_ld = DataLoader(train_seg_ds, batch_size=BS_SEG, shuffle=True,\n                              num_workers=0, pin_memory=True)\n    val_seg_ld   = DataLoader(val_seg_ds,   batch_size=BS_SEG, shuffle=False,\n                              num_workers=0, pin_memory=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T15:54:58.743283Z","iopub.execute_input":"2025-05-13T15:54:58.743661Z","iopub.status.idle":"2025-05-13T15:58:39.476104Z","shell.execute_reply.started":"2025-05-13T15:54:58.743635Z","shell.execute_reply":"2025-05-13T15:58:39.475363Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Cell 5 – ResNetUNet\nimport torch.nn as nn, torch.nn.functional as F\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision.models import resnet18\n\nclass ResNetUNet(nn.Module):\n    def __init__(self, n_cls=8):  # 0 + 1..7\n        super().__init__()\n        base_model = resnet18(pretrained=True)\n        base_model.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        self.input_conv = nn.Sequential(\n            base_model.conv1,  # 64\n            base_model.bn1,\n            base_model.relu,\n        )\n        self.pool = base_model.maxpool\n        self.enc1 = base_model.layer1  # 64\n        self.enc2 = base_model.layer2  # 128\n        self.enc3 = base_model.layer3  # 256\n        self.enc4 = base_model.layer4  # 512\n\n        def up_blk(c_in, c_out):\n            return nn.Sequential(\n                nn.Conv2d(c_in, c_out, 3, padding=1),\n                nn.BatchNorm2d(c_out),\n                nn.ReLU(inplace=True),\n                nn.Conv2d(c_out, c_out, 3, padding=1),\n                nn.BatchNorm2d(c_out),\n                nn.ReLU(inplace=True),\n            )\n\n        self.up3 = up_blk(512 + 256, 256)\n        self.up2 = up_blk(256 + 128, 128)\n        self.up1 = up_blk(128 + 64, 64)\n        self.up0 = up_blk(64 + 64, 32)  # skip from input_conv\n\n        self.final = nn.Conv2d(32, n_cls, kernel_size=1)\n    \n    def forward(self, x):\n        input_size = x.shape[2:]            # Save original H, W for final upsample\n    \n        x0 = self.input_conv(x)            # [B,64,H/2,W/2]\n        x1 = self.pool(x0)                 # [B,64,H/4,W/4]\n        x2 = self.enc1(x1)                 # [B,64,H/4,W/4]\n        x3 = self.enc2(x2)                 # [B,128,H/8,W/8]\n        x4 = self.enc3(x3)                 # [B,256,H/16,W/16]\n        x5 = self.enc4(x4)                 # [B,512,H/32,W/32]\n    \n        u3 = F.interpolate(x5, scale_factor=2, mode='bilinear', align_corners=False)\n        u3 = self.up3(torch.cat([u3, x4], dim=1))\n    \n        u2 = F.interpolate(u3, scale_factor=2, mode='bilinear', align_corners=False)\n        u2 = self.up2(torch.cat([u2, x3], dim=1))\n    \n        u1 = F.interpolate(u2, scale_factor=2, mode='bilinear', align_corners=False)\n        u1 = self.up1(torch.cat([u1, x2], dim=1))\n    \n        u0 = F.interpolate(u1, scale_factor=2, mode='bilinear', align_corners=False)\n        u0 = self.up0(torch.cat([u0, x0], dim=1))\n    \n        out = self.final(u0)              # [B, n_cls, H/2, W/2]\n        out = F.interpolate(out, size=input_size, mode='bilinear', align_corners=False)  # ⬅ Fix here\n        return out\n\n\nunet = ResNetUNet().to(DEVICE)\nopt_seg = torch.optim.Adam(unet.parameters(),1e-3)\nce = nn.CrossEntropyLoss()\n\ndef run_seg_epoch(loader, training=True):\n    unet.train(training)\n    tot, n = 0, 0\n    for img, m in tqdm(loader, leave=False):\n        img, m = img.to(DEVICE), m.to(DEVICE)\n        pred = unet(img)\n        loss = ce(pred, m)\n        if training:\n            opt_seg.zero_grad(); loss.backward(); opt_seg.step()\n        tot += loss.item() * img.size(0)\n        n += img.size(0)\n    return tot / n\n\nfor ep in range(1, EPOCH_SEG+1):\n    tr = run_seg_epoch(train_seg_ld, True)\n    # Bỏ qua validation để train nhanh hơn 1 tí.\n    # vl = run_seg_epoch(val_seg_ld, False)\n\ntorch.save(unet.state_dict(), \"/kaggle/working/unet_cervical.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T15:58:39.476982Z","iopub.execute_input":"2025-05-13T15:58:39.47724Z","iopub.status.idle":"2025-05-13T16:39:03.34825Z","shell.execute_reply.started":"2025-05-13T15:58:39.477221Z","shell.execute_reply":"2025-05-13T16:39:03.347451Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#Cell 6 – Tạo ROI & Dataset Phân loại (với load U-Net checkpoint)\nimport torch\nimport torch.nn as nn\nimport numpy as np\nimport cv2\nimport nibabel as nib\nfrom pathlib import Path\nfrom torch.utils.data import Dataset\n\n# ── 0) Load U-Net checkpoint trước khi inference ROI ──\ncheckpoint_path = \"/kaggle/working/unet_cervical.pth\"\nif Path(checkpoint_path).exists():\n    unet.load_state_dict(torch.load(checkpoint_path, map_location=DEVICE))\n    unet.to(DEVICE).eval()\n    print(f\"✅ Loaded U-Net weights from {checkpoint_path}\")\nelse:\n    print(f\"⚠️  Không tìm thấy checkpoint U-Net tại {checkpoint_path}, dùng weights mặc định\")\n\n# augmentation classification (Kornia)\n\n# Tạo Dataset & DataLoader\ntrain_cls_ds = ClsStudyDS(train_uids, transform=aug_cls)\nval_cls_ds   = ClsStudyDS(val_uids, transform=K.AugmentationSequential(\n    K.Normalize(mean=torch.tensor([0.5]), std=torch.tensor([0.5])),\n    data_keys=[\"input\"]\n))\n\nfrom torch.utils.data import DataLoader\ntrain_cls_ld = DataLoader(train_cls_ds, batch_size=BS_CLS, shuffle=True,  num_workers=0, pin_memory=True)\nval_cls_ld   = DataLoader(val_cls_ds,   batch_size=BS_CLS, shuffle=False, num_workers=0, pin_memory=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T16:39:24.895152Z","iopub.execute_input":"2025-05-13T16:39:24.895364Z","iopub.status.idle":"2025-05-13T16:39:24.974352Z","shell.execute_reply.started":"2025-05-13T16:39:24.895348Z","shell.execute_reply":"2025-05-13T16:39:24.973394Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === Cell 7 – ResNet-18 + Weighted BCE + Metric fix ===\nimport torch\nimport torch.nn as nn\nimport torchvision.models as models\nfrom tqdm.auto import tqdm\n\n# 1) Build model\nresnet = models.resnet18(weights=None)\nresnet.conv1 = nn.Conv2d(1, 64, 7, 2, 3, bias=False)\nresnet.fc    = nn.Linear(resnet.fc.in_features, 7)\nresnet = resnet.to(DEVICE)\n\n# 2) Tính pos_weight từ df_sub (tập con sau Cell 3)\n#    pos = số samples positive cho mỗi class, neg = tổng – pos\npos = df_sub[label_cols].sum().values\nneg = len(df_sub) - pos\npos_weight = torch.tensor(neg/pos, dtype=torch.float32, device=DEVICE)\n\n# 3) Loss + Optimizer\ncriterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\noptimizer = torch.optim.Adam(resnet.parameters(), lr=1e-4)\n\n# 4) Hàm train/val mỗi epoch\ndef run_cls_epoch(dataloader, train=True):\n    resnet.train(train)\n    total_loss = 0.0\n    correct = 0\n    total_labels = 0\n    with torch.set_grad_enabled(train):\n        for imgs, labels in tqdm(dataloader, desc=\"train\" if train else \"val\"):\n            imgs   = imgs.to(DEVICE).float()\n            labels = labels.to(DEVICE)\n            logits = resnet(imgs)                     # [B,7]\n            loss   = criterion(logits, labels)\n            if train:\n                optimizer.zero_grad()\n                loss.backward()\n                optimizer.step()\n            total_loss += loss.item() * imgs.size(0)\n\n            # --- Tính metric: dùng sigmoid + threshold=0.5 ---\n            probs = torch.sigmoid(logits)\n            preds = (probs > 0.5).float()\n            correct += (preds == labels).sum().item()\n            total_labels += labels.numel()\n\n    avg_loss = total_loss / len(dataloader.dataset)\n    acc      = correct / total_labels\n    return avg_loss, acc\n\n# 5) Vòng huấn luyện với Early Stopping gợi ý\nbest_val_loss = float(\"inf\")\npatience = 2   # dừng nếu Val loss không giảm sau 2 epoch\nwait = 0\n\nfor ep in range(1, EPOCH_CLS+1):\n    tr_loss, tr_acc = run_cls_epoch(train_cls_ld, True)\n    # Bỏ qua validation để train nhanh hơn 1 tí; hardware của Kaggle khá chậm.\n    # vl_loss, vl_acc = run_cls_epoch(val_cls_ld,   False)\n    # print(f\"Ep{ep:02d}  TL={tr_loss:.3f}  VL={vl_loss:.3f}  \"\n    #      f\"Tacc={tr_acc:.4f}  Vacc={vl_acc:.4f}\")\n    print(f\"Ep{ep:02d}  TL={tr_loss:.3f}  Tacc={tr_acc:.4f}\")\n\n    # Save best\n    # if vl_loss < best_val_loss:\n    #     best_val_loss = vl_loss\n    #     torch.save(resnet.state_dict(), \"/kaggle/working/resnet18_best.pth\")\n    #     wait = 0\n    # else:\n    #     wait += 1\n    #     print(f\"  ⚠️ Val loss ↑  (wait {wait}/{patience})\")\n    #     if wait >= patience:\n    #         print(\"🚨 Early stopping!\")\n    #         break\n\n    # Lưu từng bước để có thể backup\n    torch.save(resnet.state_dict(), \"/kaggle/working/resnet18_best.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-13T16:39:24.975191Z","iopub.execute_input":"2025-05-13T16:39:24.975447Z"}},"outputs":[],"execution_count":null}]}