{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":112899,"databundleVersionId":13449579,"sourceType":"competition"},{"sourceId":13183819,"sourceType":"datasetVersion","datasetId":8354746},{"sourceId":13183822,"sourceType":"datasetVersion","datasetId":8354748},{"sourceId":585600,"sourceType":"modelInstanceVersion","modelInstanceId":437575,"modelId":454277},{"sourceId":591800,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":442753,"modelId":459287}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os, gc, math, json, random\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\n\nimport torch, torch.nn as nn, torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast, GradScaler\nfrom torchvision import transforms\nfrom PIL import Image\nimport timm\n\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.metrics import roc_auc_score\n\ndef seed_everything(seed=42):\n    random.seed(seed); np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"]=str(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\nseed_everything(42)\n\n# Paths\nDATA_DIR = \"/kaggle/input/grand-xray-slam-division-a\"\nTRAIN_CSV = f\"{DATA_DIR}/train1.csv\"\nIMG_TRAIN_DIR = f\"{DATA_DIR}/train1\"\nIMG_TEST_DIR  = f\"{DATA_DIR}/test1\"\nSAMPLE_SUB = f\"{DATA_DIR}/sample_submission_1.csv\"\n\n# === Add Data: dataset từ Notebook A (chứa artifacts/) ===\n# Ví dụ: /kaggle/input/<your-view-mtl-output-dataset>/artifacts/...\nVIEW_DATASET_DIR = \"/kaggle/input/grand-a\"  # <-- sửa tên dataset bạn vừa tạo\nART_DIR = Path(f\"{VIEW_DATASET_DIR}/artifacts\")\n\n# Device & Config\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nN_GPU = torch.cuda.device_count()\nprint(\"CUDA:\", torch.cuda.is_available(), \"| GPUs:\", N_GPU, [torch.cuda.get_device_name(i) for i in range(N_GPU)])\n\nMAIN_BACKBONE = \"convnext_tiny\"\nIMG_SIZE_MAIN = 320\nBATCH_MAIN    = 24\nEPOCHS_MAIN   = 10\nLR_MAIN       = 2e-4\nWD_MAIN       = 1e-4\nACCUM_TARGET  = 64\n\nLABELS = [\n    \"Atelectasis\",\"Cardiomegaly\",\"Consolidation\",\"Edema\",\"Enlarged Cardiomediastinum\",\n    \"Fracture\",\"Lung Lesion\",\"Lung Opacity\",\"No Finding\",\"Pleural Effusion\",\n    \"Pleural Other\",\"Pneumonia\",\"Pneumothorax\",\"Support Devices\"\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T14:26:06.484282Z","iopub.execute_input":"2025-09-27T14:26:06.484889Z","iopub.status.idle":"2025-09-27T14:26:16.028602Z","shell.execute_reply.started":"2025-09-27T14:26:06.484865Z","shell.execute_reply":"2025-09-27T14:26:16.027888Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(TRAIN_CSV)\nIMAGE_COL_TRAIN = \"Image_name\" if \"Image_name\" in df.columns else \"Image_Name\"\ndf[\"image_path\"] = df[IMAGE_COL_TRAIN].apply(lambda x: os.path.join(IMG_TRAIN_DIR, x))\n\n# Folds theo bệnh nhân (nếu cần)\nif \"fold\" not in df.columns or (df[\"fold\"]<0).all():\n    df[\"fold\"] = -1\n    gkf = GroupKFold(n_splits=5)\n    for f,(tr,va) in enumerate(gkf.split(df, groups=df[\"Patient_ID\"])):\n        df.loc[va,\"fold\"] = f\n\n# Base meta từ Age/Sex\ndef build_base_meta_train(_df):\n    meta = pd.DataFrame(index=_df.index)\n    age = _df[\"Age\"].astype(float)\n    mu = age.mean(skipna=True); sd = age.std(skipna=True) or 1.0\n    meta[\"Age_z\"] = ((age.fillna(mu) - mu) / sd).astype(np.float32)\n\n    sex = _df[\"Sex\"].astype(str).str.upper()\n    meta[\"Sex_M\"] = (sex==\"MALE\").astype(np.float32)\n    meta[\"Sex_F\"] = (sex==\"FEMALE\").astype(np.float32)\n    meta[\"Sex_U\"] = (~sex.isin([\"MALE\",\"FEMALE\"])).astype(np.float32)\n    return meta\n\ndf_meta_true = build_base_meta_train(df)\nprint(\"df_meta_true:\", df_meta_true.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T14:26:19.802648Z","iopub.execute_input":"2025-09-27T14:26:19.803646Z","iopub.status.idle":"2025-09-27T14:26:20.485443Z","shell.execute_reply.started":"2025-09-27T14:26:19.803584Z","shell.execute_reply":"2025-09-27T14:26:20.484651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load view meta từ dataset A\nmeta_train_view = np.load(ART_DIR/\"meta_train_view.npy\")\nmeta_test_view  = np.load(ART_DIR/\"meta_test_view.npy\")\nwith open(ART_DIR/\"meta_cols_view.json\") as f: meta_cols_view = json.load(f)\nprint(\"view meta shapes:\", meta_train_view.shape, meta_test_view.shape)\n\n# Base train meta\nbase_train_np = df_meta_true.values.astype(np.float32)\nbase_cols = list(df_meta_true.columns)\n\n# Base test meta (Age/Sex không có -> zero/neutral)\nsub = pd.read_csv(SAMPLE_SUB)\nbase_test = pd.DataFrame(0, index=np.arange(len(sub)), columns=base_cols, dtype=np.float32)\nif \"Age_z\" in base_test.columns: base_test[\"Age_z\"] = 0.0\nbase_test_np = base_test.values.astype(np.float32)\n\n# Final META\nmeta_train_np = np.concatenate([base_train_np, meta_train_view], axis=1)\ntest_meta_np  = np.concatenate([base_test_np,  meta_test_view],  axis=1)\nmeta_cols = base_cols + meta_cols_view\nMETA_DIM = meta_train_np.shape[1]\nprint(\"Final META:\", meta_train_np.shape, test_meta_np.shape, \"| META_DIM:\", META_DIM)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T14:26:22.367063Z","iopub.execute_input":"2025-09-27T14:26:22.367336Z","iopub.status.idle":"2025-09-27T14:26:22.964211Z","shell.execute_reply.started":"2025-09-27T14:26:22.367313Z","shell.execute_reply":"2025-09-27T14:26:22.963498Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"assert all(c in df.columns for c in LABELS), \"Thiếu một số cột 14 labels trong train1.csv\"\n\ndef tfms_main(train=True):\n    if train:\n        return transforms.Compose([\n            transforms.Resize((IMG_SIZE_MAIN, IMG_SIZE_MAIN)),\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.ColorJitter(brightness=0.2, contrast=0.2),\n            transforms.ToTensor(),\n            transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225]),\n        ])\n    else:\n        return transforms.Compose([\n            transforms.Resize((IMG_SIZE_MAIN, IMG_SIZE_MAIN)),\n            transforms.ToTensor(),\n            transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225]),\n        ])\n\nclass MainTrainDS(Dataset):\n    def __init__(self, df_idx, meta_np, tfm):\n        self.df = df.loc[df_idx].reset_index(drop=True)\n        self.meta = meta_np[df_idx].astype(np.float32)\n        self.tfm = tfm\n    def __len__(self): return len(self.df)\n    def __getitem__(self, i):\n        row = self.df.loc[i]\n        img = Image.open(row[\"image_path\"]).convert(\"L\").convert(\"RGB\")\n        img = self.tfm(img)\n        y = torch.tensor(row[LABELS].values.astype(np.float32), dtype=torch.float32)\n        m = torch.tensor(self.meta[i], dtype=torch.float32)\n        return img, m, y\n\nclass MainTestDS(Dataset):\n    def __init__(self, paths, meta_np, tfm):\n        self.paths = list(paths); self.meta = meta_np.astype(np.float32); self.tfm=tfm\n    def __len__(self): return len(self.paths)\n    def __getitem__(self, i):\n        img = Image.open(self.paths[i]).convert(\"L\").convert(\"RGB\")\n        img = self.tfm(img)\n        m = torch.tensor(self.meta[i], dtype=torch.float32)\n        return img, m","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T14:26:23.896939Z","iopub.execute_input":"2025-09-27T14:26:23.897193Z","iopub.status.idle":"2025-09-27T14:26:23.906267Z","shell.execute_reply.started":"2025-09-27T14:26:23.89717Z","shell.execute_reply":"2025-09-27T14:26:23.905752Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ConvNextWithMeta(nn.Module):\n    def __init__(self, name=MAIN_BACKBONE, meta_dim=8, n_out=14):\n        super().__init__()\n        self.backbone = timm.create_model(name, pretrained=True, num_classes=0, global_pool=\"avg\")\n        feat = self.backbone.num_features\n        self.meta_bn = nn.BatchNorm1d(meta_dim)\n        self.head = nn.Sequential(\n            nn.Linear(feat+meta_dim, 512),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.2),\n            nn.Linear(512, n_out),\n        )\n    def forward(self, x, meta):\n        f = self.backbone(x)\n        meta = self.meta_bn(meta)\n        z = torch.cat([f, meta], dim=1)\n        return self.head(z)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T14:26:25.696827Z","iopub.execute_input":"2025-09-27T14:26:25.697513Z","iopub.status.idle":"2025-09-27T14:26:25.702292Z","shell.execute_reply.started":"2025-09-27T14:26:25.697492Z","shell.execute_reply":"2025-09-27T14:26:25.701671Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import time \n# # Helpers để format thời gian & lấy LR\n# def fmt_hms(sec: float):\n#     m = int(sec // 60); s = int(sec % 60); h = m // 60; m = m % 60\n#     return f\"{h:d}h{m:02d}m{s:02d}s\" if h>0 else f\"{m:02d}m{s:02d}s\"\n\n# def get_lr(optimizer):\n#     for pg in optimizer.param_groups:\n#         return pg.get(\"lr\", None)\n\n# # ====== REPLACE HÀM NÀY (bản không dùng tqdm) ======\n# def train_one_fold_main(fold, epochs=EPOCHS_MAIN, log_every=50):\n#     tr_idx = df.index[df.fold!=fold].values\n#     va_idx = df.index[df.fold==fold].values\n\n#     dl_tr = DataLoader(\n#         MainTrainDS(tr_idx, meta_train_np, tfms_main(True)),\n#         batch_size=BATCH_MAIN, shuffle=True, num_workers=4,\n#         pin_memory=True, drop_last=True\n#     )\n#     dl_va = DataLoader(\n#         MainTrainDS(va_idx, meta_train_np, tfms_main(False)),\n#         batch_size=BATCH_MAIN*2, shuffle=False, num_workers=4,\n#         pin_memory=True\n#     )\n\n#     model = ConvNextWithMeta(name=MAIN_BACKBONE, meta_dim=META_DIM, n_out=len(LABELS))\n#     if torch.cuda.device_count()>1: model = nn.DataParallel(model)\n#     model = model.to(DEVICE).to(memory_format=torch.channels_last)\n\n#     y_tr = df.loc[tr_idx, LABELS].values.astype(np.float32)\n#     pos  = y_tr.sum(axis=0) + 1e-3\n#     neg  = (y_tr.shape[0] - y_tr.sum(axis=0)) + 1e-3\n#     pos_w = torch.tensor(neg/pos, dtype=torch.float32).to(DEVICE)\n#     crit  = nn.BCEWithLogitsLoss(pos_weight=pos_w)\n#     opt   = optim.AdamW(model.parameters(), lr=LR_MAIN, weight_decay=WD_MAIN)\n#     sch   = optim.lr_scheduler.CosineAnnealingLR(opt, T_max=epochs, eta_min=LR_MAIN*0.1)\n#     scaler = GradScaler(\"cuda\", enabled=torch.cuda.is_available())\n#     grad_accum = max(1, math.ceil(ACCUM_TARGET / BATCH_MAIN))\n\n#     n_train_steps = len(dl_tr)\n#     n_val_steps   = len(dl_va)\n#     print(f\"[MAIN] fold{fold} | train_steps={n_train_steps} | val_steps={n_val_steps} \"\n#           f\"| bs={BATCH_MAIN} | accum={grad_accum} | lr0={LR_MAIN:g}\")\n\n#     best_auc = -1.0\n#     if torch.cuda.is_available(): torch.cuda.reset_peak_memory_stats()\n\n#     for ep in range(1, epochs+1):\n#         # ---------- TRAIN ----------\n#         model.train(); tr_loss = 0.0; seen_imgs = 0\n#         opt.zero_grad(set_to_none=True)\n\n#         ep_start = time.time()\n#         # mốc để tính ips/eta ổn định sau vài bước\n#         tick0 = None\n#         for it, (x, m, y) in enumerate(dl_tr, 1):\n#             if tick0 is None: tick0 = time.time()\n#             x = x.to(DEVICE, non_blocking=True).to(memory_format=torch.channels_last)\n#             m = m.to(DEVICE, non_blocking=True)\n#             y = y.to(DEVICE, non_blocking=True)\n\n#             with autocast(\"cuda\", enabled=torch.cuda.is_available()):\n#                 logits = model(x, m)\n#                 loss   = crit(logits, y) / grad_accum\n\n#             scaler.scale(loss).backward()\n#             if it % grad_accum == 0:\n#                 scaler.step(opt); scaler.update(); opt.zero_grad(set_to_none=True)\n\n#             tr_loss   += (loss.item() * grad_accum) * x.size(0)\n#             seen_imgs += x.size(0)\n\n#             # LOG THỦ CÔNG mỗi log_every step\n#             if (it % log_every == 0) or (it == n_train_steps):\n#                 elapsed = max(1e-6, time.time() - tick0)\n#                 ips     = seen_imgs / elapsed\n#                 remain_imgs = n_train_steps * BATCH_MAIN - seen_imgs\n#                 eta_s   = remain_imgs / max(1e-6, ips)\n#                 avg_loss = tr_loss / max(1, seen_imgs)\n#                 print(f\"[fold{fold}] train ep{ep}/{epochs} | step {it:5d}/{n_train_steps:5d} \"\n#                       f\"| loss {avg_loss:.4f} | lr {get_lr(opt):.2e} | ips {ips:.1f} img/s | eta {fmt_hms(eta_s)}\")\n\n#         # flush nếu còn dư tích lũy\n#         if (it % grad_accum) != 0:\n#             scaler.step(opt); scaler.update(); opt.zero_grad(set_to_none=True)\n\n#         tr_time = time.time() - ep_start\n#         tr_loss = tr_loss / max(1, seen_imgs)\n#         sch.step()\n\n#         # ---------- VALID ----------\n#         model.eval(); preds = []; targs = []\n#         val_start = time.time()\n\n#         val_seen = 0\n#         tickv0 = None\n#         with torch.no_grad(), autocast(\"cuda\", enabled=torch.cuda.is_available()):\n#             for j, (x, m, y) in enumerate(dl_va, 1):\n#                 if tickv0 is None: tickv0 = time.time()\n#                 x = x.to(DEVICE, non_blocking=True).to(memory_format=torch.channels_last)\n#                 m = m.to(DEVICE, non_blocking=True)\n#                 logits = model(x, m).float().cpu().numpy()\n#                 preds.append(1/(1+np.exp(-logits)))\n#                 targs.append(y.numpy())\n\n#                 # LOG thủ công cho valid cũng theo nhịp\n#                 val_seen += x.size(0)\n#                 if (j % log_every == 0) or (j == n_val_steps):\n#                     elapsed = max(1e-6, time.time() - tickv0)\n#                     ips     = val_seen / elapsed\n#                     remain_imgs = n_val_steps * (BATCH_MAIN*2) - val_seen\n#                     eta_s   = remain_imgs / max(1e-6, ips)\n#                     print(f\"[fold{fold}] valid ep{ep}/{epochs} | step {j:5d}/{n_val_steps:5d} \"\n#                           f\"| ips {ips:.1f} img/s | eta {fmt_hms(eta_s)}\")\n\n#         va_time = time.time() - val_start\n#         P = np.vstack(preds); T = np.vstack(targs)\n\n#         aucs = []\n#         for k in range(len(LABELS)):\n#             try:   aucs.append(roc_auc_score(T[:,k], P[:,k]))\n#             except: pass\n#         mean_auc = float(np.mean(aucs)) if len(aucs)>0 else 0.0\n\n#         peak_mem = (torch.cuda.max_memory_allocated()/1024**3) if torch.cuda.is_available() else 0.0\n#         print(f\"[MAIN fold{fold}] ep{ep}/{epochs} | train_loss {tr_loss:.4f} | val mAUC {mean_auc:.4f} \"\n#               f\"| train {fmt_hms(tr_time)} | valid {fmt_hms(va_time)} | peak {peak_mem:.2f} GB\")\n\n#         if mean_auc > best_auc:\n#             best_auc = mean_auc\n#             sp = f\"./main_fold{fold}.pth\"\n#             torch.save(model.module.state_dict() if hasattr(model,\"module\") else model.state_dict(), sp)\n#             print(\"  -> saved\", sp)\n\n#         torch.cuda.empty_cache()\n#         if torch.cuda.is_available(): torch.cuda.reset_peak_memory_stats()\n\n#     return best_auc\n\n\n# # =========== REPLACE: vòng chạy CV có log tổng thể ===========\n# cv_scores = []\n# t0_all = time.perf_counter()\n# for fold in range(5):\n#     print(f\"\\n===== START MAIN FOLD {fold} =====\")\n#     t_fold0 = time.perf_counter()\n#     s = train_one_fold_main(fold); cv_scores.append(s)\n#     fold_time = time.perf_counter() - t_fold0\n#     done = fold + 1\n#     remain = 5 - done\n#     avg_per_fold = (time.perf_counter() - t0_all) / done\n#     eta_total = remain * avg_per_fold\n#     print(f\"[DONE fold{fold}] AUC={s:.4f} | time {fmt_hms(fold_time)} | ETA total ~{fmt_hms(eta_total)}\")\n\n# print(\"\\nMAIN fold AUCs:\", [f\"{x:.4f}\" for x in cv_scores], \"mean:\", float(np.mean(cv_scores)))\n# print(\"TOTAL time:\", fmt_hms(time.perf_counter() - t0_all))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T01:46:41.039482Z","iopub.execute_input":"2025-09-27T01:46:41.040149Z"},"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================== UNIFIED PREDICT CELL (no global collisions) ==================\nimport os, gc, glob\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nimport time\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast\n\ndef run_inference():\n    # ----------- 0) Các phụ thuộc/cấu hình lấy từ global nếu có -----------\n    DEVICE      = globals().get('DEVICE', 'cuda' if torch.cuda.is_available() else 'cpu')\n    BATCH_MAIN  = globals().get('BATCH_MAIN', 16)\n    MAIN_BACKBONE = globals().get('MAIN_BACKBONE', 'convnext_base')  # sửa đúng nếu bạn train khác\n\n    # Class model phải có sẵn (định nghĩa như lúc train)\n    if 'ConvNextWithMeta' not in globals():\n        raise NameError(\"Chưa có class ConvNextWithMeta (hãy dán lại định nghĩa đúng như lúc train).\")\n\n    # ----------- 1) Sample submission + LABELS -----------\n    sample_sub = '/kaggle/input/grand-xray-slam-division-a/sample_submission1.csv'\n    if not os.path.exists(sample_sub):\n        sample_sub = '/kaggle/input/grand-xray-slam-division-a/sample_submission_1.csv'\n    sub = pd.read_csv(sample_sub)\n    LABELS = [c for c in sub.columns if c != 'Image_name']\n\n    # ----------- 2) Resolve test image paths -----------\n    IMG_TEST_DIR = globals().get('IMG_TEST_DIR', None)\n    if IMG_TEST_DIR is None:\n        cands = [\n            '/kaggle/input/grand-xray-slam-division-a/test_images',\n            '/kaggle/input/grand-xray-slam-division-a/test',\n            '/kaggle/input/grand-xray-slam-division-a/images/test',\n            '/kaggle/input/grand-xray-slam-division-a'\n        ]\n        IMG_TEST_DIR = next((p for p in cands if os.path.isdir(p)), cands[-1])\n\n    def resolve_path(root, name):\n        p = os.path.join(root, name)\n        if os.path.exists(p): return p\n        for ext in ['.png', '.jpg', '.jpeg', '.bmp']:\n            q = os.path.join(root, name + ext)\n            if os.path.exists(q): return q\n        hits = glob.glob(os.path.join(root, \"**\", name), recursive=True)\n        if not hits:\n            for ext in ['.png', '.jpg', '.jpeg', '.bmp']:\n                hits = glob.glob(os.path.join(root, \"**\", name + ext), recursive=True)\n                if hits: break\n        return hits[0] if hits else os.path.join(root, name)\n\n    test_paths = [resolve_path(IMG_TEST_DIR, n) for n in sub['Image_name']]\n    miss = sum(not os.path.exists(p) for p in test_paths)\n    if miss:\n        print(f\"[WARN] {miss} ảnh không tìm thấy. Kiểm tra tên file/đuôi ảnh.\")\n\n    # ----------- 3) Transforms -----------\n    try:\n        tfms = globals()['tfms_main'](False)\n    except Exception:\n        import torchvision.transforms as T\n        IMG_SIZE = 384  # sửa đúng size đã train\n        def tfms_main(is_train=False):\n            return T.Compose([\n                T.Resize((IMG_SIZE, IMG_SIZE)),\n                T.ToTensor(),\n                T.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),\n            ])\n        tfms = tfms_main(False)\n\n    # ----------- 4) Chọn weights -----------\n    fold_weights = [f'./main_fold{f}.pth' for f in range(5)]\n    have_5fold = all(os.path.exists(p) for p in fold_weights)\n    single_main_weight = globals().get('SINGLE_MAIN_WEIGHT',\n        '/kaggle/input/model1/pytorch/default/1/main_fold0.pth')  # chỉnh nếu khác\n    weight_paths = fold_weights if have_5fold else [single_main_weight]\n    print(\"[INFO] Weights:\", weight_paths)\n\n    # ----------- 5) Suy META_DIM từ ckpt + tạo meta_test_np -----------\n    def _load_state_dict_any(path):\n        obj = torch.load(path, map_location='cpu')\n        if isinstance(obj, dict) and all(isinstance(k, str) for k in obj.keys()):\n            if 'state_dict' in obj and isinstance(obj['state_dict'], dict):\n                return obj['state_dict']\n            return obj\n        raise RuntimeError(\"Không nhận dạng được định dạng checkpoint.\")\n\n    def infer_meta_dim_from_ckpt(path):\n        sd = _load_state_dict_any(path)\n        keys = list(sd.keys())\n        cands = [k for k in keys if k.endswith('meta_bn.weight') or k.endswith('meta_bn.bias')]\n        if not cands:\n            cands = [k for k in keys if 'meta' in k and (k.endswith('.weight') or k.endswith('.bias'))]\n        if not cands:\n            print(\"[INFO] Không thấy nhánh meta trong ckpt -> META_DIM=0\")\n            return 0\n        t = sd[cands[0]]\n        try: return int(t.shape[0])\n        except: return int(t.numel())\n\n    required_meta_dim = infer_meta_dim_from_ckpt(weight_paths[0])\n    print(f\"[INFO] Checkpoint yêu cầu META_DIM = {required_meta_dim}\")\n\n    # Nếu có sẵn meta_test_np global thì dùng; nếu không, tạo zeros đúng kích thước\n    meta_test_np = globals().get('meta_test_np', None)\n    if meta_test_np is None:\n        n_test = len(sub)\n        meta_test_np = np.zeros((n_test, max(required_meta_dim, 0)), dtype=np.float32)\n        if required_meta_dim > 0:\n            print(\"[WARN] Ckpt có nhánh meta nhưng bạn chưa cung cấp meta -> dùng zeros.\")\n\n    def adjust_meta(meta_np, k):\n        c = meta_np.shape[1]\n        if c == k: return meta_np\n        if k == 0:\n            print(f\"[WARN] Model không dùng meta (k=0). Bỏ {c} cột meta.\")\n            return np.zeros((meta_np.shape[0], 0), dtype=meta_np.dtype)\n        if c < k:\n            pad = np.zeros((meta_np.shape[0], k - c), dtype=meta_np.dtype)\n            print(f\"[WARN] Padded meta {c} → {k}.\")\n            return np.concatenate([meta_np, pad], axis=1)\n        print(f\"[WARN] Sliced meta {c} → {k}. Đảm bảo thứ tự cột trùng lúc train.\")\n        return meta_np[:, :k]\n\n    meta_test_np = adjust_meta(meta_test_np, required_meta_dim)\n    META_DIM = required_meta_dim\n    print(\"[INFO] meta_test_np shape:\", meta_test_np.shape)\n\n    # ----------- 6) Dataset/DataLoader (scope local) -----------\n    class _MainTestDS(Dataset):\n        def __init__(self, paths, meta_np, tfm):\n            self.paths = list(paths)\n            self.meta  = meta_np.astype(np.float32)\n            self.tfm   = tfm\n        def __len__(self): return len(self.paths)\n        def __getitem__(self, i):\n            img = Image.open(self.paths[i]).convert(\"L\").convert(\"RGB\")\n            img = self.tfm(img)\n            m = torch.tensor(self.meta[i], dtype=torch.float32)\n            return img, m\n\n    dl_test = DataLoader(\n        _MainTestDS(test_paths, meta_test_np, tfms),\n        batch_size=BATCH_MAIN*2, shuffle=False, num_workers=4, pin_memory=True\n    )\n    print(f\"[INFO] Test steps: {len(dl_test)} | batch: {BATCH_MAIN*2}\")\n\n    # ----------- 7) Load state_dict linh hoạt -----------\n    def load_state_dict_flexible(model, sd):\n        model_keys = set(model.state_dict().keys())\n        sd_keys = set(sd.keys())\n        if all(k.startswith('module.') for k in sd_keys):\n            sd = {k.replace('module.','',1): v for k,v in sd.items()}\n            sd_keys = set(sd.keys())\n        try:\n            model.load_state_dict(sd, strict=True)\n        except Exception:\n            missing = [k for k in model_keys if k not in sd_keys][:10]\n            unexpected = [k for k in sd_keys if k not in model_keys][:10]\n            print(\"[WARN] strict=False do lệch key.\")\n            if missing: print(\"  missing:\", missing, \"...\")\n            if unexpected: print(\"  unexpected:\", unexpected, \"...\")\n            model.load_state_dict(sd, strict=False)\n        return model\n\n    # ----------- 8) Hàm build model & predict -----------\n       # ----------- 8) Hàm build model & predict -----------\n    def build_model(weight_path):\n        try:\n            m = ConvNextWithMeta(name=MAIN_BACKBONE, meta_dim=META_DIM, n_out=len(LABELS))\n        except Exception:\n            m = ConvNextWithMeta(meta_dim=META_DIM, n_out=len(LABELS))\n        sd = _load_state_dict_any(weight_path)\n        m = load_state_dict_flexible(m, sd)\n        m.eval()\n        if torch.cuda.device_count() > 1:\n            m = nn.DataParallel(m)\n        return m.to(DEVICE)\n\n    preds = np.zeros((len(sub), len(LABELS)), dtype=np.float32)\n    with torch.no_grad():\n        # vòng ngoài: theo dõi từng weight\n        for wi, wpath in enumerate(weight_paths, 1):\n            print(f\"\\n[RUN] Loading weight {wi}/{len(weight_paths)}: {wpath}\")\n            model = build_model(wpath)\n\n            chunk = []\n            seen = 0\n            t0 = time.perf_counter()\n\n            # vòng trong: theo dõi từng batch\n            pbar = tqdm(\n                DataLoader(\n                    _MainTestDS(test_paths, meta_test_np, tfms),\n                    batch_size=BATCH_MAIN*2, shuffle=False,\n                    num_workers=4, pin_memory=True\n                ),\n                desc=f\"[infer {wi}/{len(weight_paths)}]\",\n                leave=False\n            )\n            for x, mmeta in pbar:\n                x = x.to(DEVICE, non_blocking=True)\n                mmeta = mmeta.to(DEVICE, non_blocking=True)\n                with autocast(\"cuda\", enabled=(str(DEVICE).startswith(\"cuda\") and torch.cuda.is_available())):\n                    logits = model(x, mmeta)\n                    probs  = torch.sigmoid(logits).detach().cpu().numpy()\n                chunk.append(probs)\n\n                seen += x.size(0)\n                dt = max(1e-6, time.perf_counter() - t0)\n                pbar.set_postfix(imgs=seen, ips=f\"{seen/dt:.1f}\")\n\n            fold_pred = np.vstack(chunk)\n            if len(weight_paths) == 5:\n                preds += fold_pred / 5.0\n            else:\n                preds = fold_pred\n\n            del model; gc.collect()\n            if DEVICE != \"cpu\":\n                torch.cuda.empty_cache()\n\n    # ----------- 9) Save submission -----------\n    sub_out = pd.DataFrame({'Image_name': sub['Image_name']})\n    for i, lab in enumerate(LABELS):\n        sub_out[lab] = preds[:, i].clip(0,1)\n    sub_out.to_csv(\"submission.csv\", index=False)\n    print(\"\\n[SAVED] submission.csv | shape:\", sub_out.shape)\n    return sub_out\n\n# Chạy:\n_ = run_inference()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-27T14:26:29.393844Z","iopub.execute_input":"2025-09-27T14:26:29.394534Z","execution_failed":"2025-09-27T14:28:32.446Z"},"scrolled":true},"outputs":[],"execution_count":null}]}