{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":24800,"datasetId":1042002,"databundleVersionId":1831594},{"sourceType":"datasetVersion","sourceId":2169393,"datasetId":1302315,"databundleVersionId":2210641},{"sourceType":"datasetVersion","sourceId":1799615,"datasetId":1069544,"databundleVersionId":1837072}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nfrom PIL import Image\nfrom tqdm import tqdm\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score, classification_report, f1_score, precision_recall_curve, multilabel_confusion_matrix\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision import transforms, models\nfrom torch.amp import GradScaler, autocast\n\nclass CFG:\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    amp_device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    seed = 42\n    img_size = 512\n    batch_size = 8\n    num_workers = 4\n    lr = 3e-4\n    lr_backbone = 1e-5\n    epochs = 30\n    freeze_epochs = 5\n    num_classes = 5\n    beta_threshold = 1.0\n    bbox_pad = 0.6\n    bbox_ratio = 0.8\n    lung_opacity_pad = 1.5\n    lung_opacity_weight = 5.0\n    max_no_finding = 6000\n\ntorch.manual_seed(CFG.seed)\nnp.random.seed(CFG.seed)\n\nJPG_BASE = \"/kaggle/input/datasets/raddar/vinbigdata-competition-jpg-data-3x-downsampled\"\nTRAIN_CSV = os.path.join(JPG_BASE, \"train_downsampled.csv\")\nIMG_DIR = os.path.join(JPG_BASE, \"train\", \"train\")\n\nCHEXPERT_DIR = \"/kaggle/input/datasets/ashery/chexpert\"\nCHEX_CKPT = \"/kaggle/working/chexpert_features.pth\"\nBEST_MODEL_PATH = \"/kaggle/working/best_vin_bbox_model.pth\"\n\nSELECTED_CLASSES = [7, 10, 11, 13, 14]\n\nCLASS_NAMES = [\n    \"Lung Opacity\",\n    \"Pleural Effusion\",\n    \"Pleural Thickening\",\n    \"Pulmonary Fibrosis\",\n    \"No Finding\"\n]\n\nprint(\"Device:\", CFG.device)\nprint(\"TRAIN_CSV:\", TRAIN_CSV)\nprint(\"IMG_DIR:\", IMG_DIR)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv(TRAIN_CSV)\n\nprint(train_df.head())\nprint(\"Số dòng train.csv:\", len(train_df))\nprint(\"Số ảnh unique:\", train_df[\"image_id\"].nunique())\nprint(\"Các file trong IMG_DIR:\", os.listdir(IMG_DIR)[:5])\n\nsample_id = train_df[\"image_id\"].iloc[0]\nsample_path = os.path.join(IMG_DIR, f\"{sample_id}.jpg\")\n\nimg = cv2.imread(sample_path, cv2.IMREAD_GRAYSCALE)\n\nprint(\"Sample path:\", sample_path)\nprint(\"Image shape:\", None if img is None else img.shape)\n\nplt.figure(figsize=(5, 5))\nplt.imshow(img, cmap=\"gray\")\nplt.title(\"Sample X-ray\")\nplt.axis(\"off\")\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lung_classes = {\n    0: \"Aortic enlargement\",\n    1: \"Atelectasis\",\n    2: \"Calcification\",\n    3: \"Cardiomegaly\",\n    4: \"Consolidation\",\n    5: \"ILD\",\n    6: \"Infiltration\",\n    7: \"Lung Opacity\",\n    8: \"Nodule/Mass\",\n    9: \"Other lesion\",\n    10: \"Pleural Effusion\",\n    11: \"Pleural Thickening\",\n    12: \"Pneumothorax\",\n    13: \"Pulmonary Fibrosis\",\n    14: \"No Finding\"\n}\n\ngrouped = train_df.groupby(\"image_id\")\ncounts = {name: 0 for name in lung_classes.values()}\n\nfor img_id, group in grouped:\n    class_ids = group[\"class_id\"].values\n    for cid, name in lung_classes.items():\n        if cid in class_ids:\n            counts[name] += 1\n\nprint(\"=== SỐ LƯỢNG ẢNH THEO BỆNH ===\")\nfor k, v in counts.items():\n    print(f\"{k}: {v}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"all_ids = sorted(train_df[\"image_id\"].unique())\n\ndef build_label_matrix(df, ids):\n    labels = np.zeros((len(ids), len(SELECTED_CLASSES)), dtype=np.float32)\n    grouped = df.groupby(\"image_id\")[\"class_id\"].apply(set).to_dict()\n\n    for i, img_id in enumerate(ids):\n        cls_set = grouped[img_id]\n\n        for j, cid in enumerate(SELECTED_CLASSES):\n            if cid in cls_set:\n                labels[i, j] = 1.0\n\n        if labels[i, :4].sum() > 0:\n            labels[i, 4] = 0.0\n        else:\n            labels[i, 4] = 1.0\n\n    return np.array(ids), labels\n\nall_ids, all_labels = build_label_matrix(train_df, all_ids)\n\nprint(\"Tổng ảnh có JPG:\", len(all_ids))\nprint(pd.DataFrame(all_labels, columns=CLASS_NAMES).sum().astype(int))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tr_ids, temp_ids, tr_labs, temp_labs = train_test_split(all_ids, all_labels, test_size=0.2, random_state=CFG.seed, shuffle=True)\nval_ids, test_ids, val_labs, test_labs = train_test_split(temp_ids, temp_labs, test_size=0.5, random_state=CFG.seed, shuffle=True)\n\nprint(\"Train:\", len(tr_ids))\nprint(\"Val:\", len(val_ids))\nprint(\"Test:\", len(test_ids))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def undersample_no_finding(ids, labels, max_no_finding=3500, seed=42):\n    rng = np.random.default_rng(seed)\n    no_find_idx = np.where(labels[:, 4] == 1)[0]\n    disease_idx = np.where(labels[:, :4].sum(axis=1) > 0)[0]\n    keep_nf = rng.choice(no_find_idx, size=min(max_no_finding, len(no_find_idx)), replace=False)\n    keep = np.concatenate([disease_idx, keep_nf])\n    rng.shuffle(keep)\n    return ids[keep], labels[keep]\n\ntr_ids, tr_labs = undersample_no_finding(tr_ids, tr_labs, max_no_finding=CFG.max_no_finding, seed=CFG.seed)\n\nprint(\"Sau undersample train:\")\nprint(\"Train:\", len(tr_ids))\nprint(pd.DataFrame(tr_labs, columns=CLASS_NAMES).sum().astype(int))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_records(df, ids, labels, use_full=True, use_bbox=True, bbox_ratio=0.3):\n    records = []\n    label_map = {img_id: labels[i] for i, img_id in enumerate(ids)}\n    df_sub = df[df[\"image_id\"].isin(ids)].copy()\n\n    if use_full:\n        for img_id in ids:\n            records.append({\n                \"image_id\": img_id,\n                \"label\": label_map[img_id],\n                \"type\": \"full\",\n                \"x_min\": np.nan,\n                \"y_min\": np.nan,\n                \"x_max\": np.nan,\n                \"y_max\": np.nan\n            })\n\n    if use_bbox:\n        bbox_df = df_sub[\n            df_sub[\"x_min\"].notna() &\n            df_sub[\"y_min\"].notna() &\n            df_sub[\"x_max\"].notna() &\n            df_sub[\"y_max\"].notna()\n        ].copy()\n\n        bbox_df = bbox_df[bbox_df[\"class_id\"].isin(SELECTED_CLASSES[:4])]\n\n        if len(bbox_df) > 0:\n            lung_bbox = bbox_df[bbox_df[\"class_id\"] == 7]\n            other_bbox = bbox_df[bbox_df[\"class_id\"] != 7]\n            bbox_df = pd.concat([lung_bbox, other_bbox.sample(frac=bbox_ratio, random_state=CFG.seed)])\n\n        for _, row in bbox_df.iterrows():\n            records.append({\n                \"image_id\": row[\"image_id\"],\n                \"label\": label_map[row[\"image_id\"]],\n                \"type\": \"bbox\",\n                \"x_min\": row[\"x_min\"],\n                \"y_min\": row[\"y_min\"],\n                \"x_max\": row[\"x_max\"],\n                \"y_max\": row[\"y_max\"]\n            })\n\n    return records\n\ntrain_records = build_records(train_df, tr_ids, tr_labs, use_full=True, use_bbox=True, bbox_ratio=CFG.bbox_ratio)\nval_records = build_records(train_df, val_ids, val_labs, use_full=True, use_bbox=False)\ntest_records = build_records(train_df, test_ids, test_labs, use_full=True, use_bbox=False)\n\nprint(\"Train records:\", len(train_records))\nprint(\"Val records:\", len(val_records))\nprint(\"Test records:\", len(test_records))\nprint(pd.Series([r[\"type\"] for r in train_records]).value_counts())","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class VinDualInputDataset(Dataset):\n    def __init__(self, ids, labels, df, full_transform=None, bbox_transform=None, return_raw=False):\n        self.ids = ids\n        self.labels = labels\n        self.df = df\n        self.full_transform = full_transform\n        self.bbox_transform = bbox_transform\n        self.return_raw = return_raw\n        self.clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n        self.bbox_map = self.build_bbox_map()\n\n    def build_bbox_map(self):\n        bbox_map = {}\n        bbox_df = self.df[\n            self.df[\"x_min\"].notna() &\n            self.df[\"y_min\"].notna() &\n            self.df[\"x_max\"].notna() &\n            self.df[\"y_max\"].notna() &\n            self.df[\"class_id\"].isin(SELECTED_CLASSES[:4])\n        ]\n\n        for img_id, group in bbox_df.groupby(\"image_id\"):\n            boxes = group[[\"x_min\", \"y_min\", \"x_max\", \"y_max\", \"class_id\"]].values.astype(float)\n            bbox_map[img_id] = boxes\n\n        return bbox_map\n\n    def __len__(self):\n        return len(self.ids)\n\n    def read_image(self, img_id):\n        img_path = os.path.join(IMG_DIR, f\"{img_id}.jpg\")\n        img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n\n        if img is None:\n            print(\"Không đọc được ảnh:\", img_path)\n            img = np.zeros((CFG.img_size, CFG.img_size), dtype=np.uint8)\n\n        return img\n\n    def make_full_image(self, img):\n        img = cv2.resize(img, (CFG.img_size, CFG.img_size))\n        img = self.clahe.apply(img)\n        img = np.stack([img, img, img], axis=-1)\n        return img\n\n    def make_bbox_image(self, img, img_id):\n        h, w = img.shape[:2]\n    \n        if img_id not in self.bbox_map:\n            return self.make_full_image(img)\n    \n        boxes_all = self.bbox_map[img_id]\n    \n        lung_boxes = boxes_all[boxes_all[:, 4] == 7]\n    \n        if len(lung_boxes) > 0:\n            boxes = lung_boxes[:, :4]\n    \n            areas = (boxes[:, 2] - boxes[:, 0]) * (boxes[:, 3] - boxes[:, 1])\n            best_idx = np.argmax(areas)\n    \n            x1, y1, x2, y2 = boxes[best_idx]\n            pad = CFG.lung_opacity_pad\n    \n        else:\n            boxes = boxes_all[:, :4]\n    \n            x1 = np.min(boxes[:, 0])\n            y1 = np.min(boxes[:, 1])\n            x2 = np.max(boxes[:, 2])\n            y2 = np.max(boxes[:, 3])\n            pad = CFG.bbox_pad\n    \n        x1 = int(x1)\n        y1 = int(y1)\n        x2 = int(x2)\n        y2 = int(y2)\n    \n        if x2 <= x1 or y2 <= y1:\n            return self.make_full_image(img)\n    \n        bw = x2 - x1\n        bh = y2 - y1\n    \n        pad_x = int(bw * pad)\n        pad_y = int(bh * pad)\n    \n        x1 = max(0, x1 - pad_x)\n        y1 = max(0, y1 - pad_y)\n        x2 = min(w, x2 + pad_x)\n        y2 = min(h, y2 + pad_y)\n    \n        crop = img[y1:y2, x1:x2]\n    \n        if crop.size == 0:\n            return self.make_full_image(img)\n    \n        crop = cv2.resize(crop, (CFG.img_size, CFG.img_size))\n        crop = self.clahe.apply(crop)\n        crop = np.stack([crop, crop, crop], axis=-1)\n    \n        return crop\n\n    def __getitem__(self, idx):\n        img_id = self.ids[idx]\n        label = self.labels[idx].copy()\n\n        img = self.read_image(img_id)\n\n        full_img = self.make_full_image(img)\n        bbox_img = self.make_bbox_image(img, img_id)\n\n        raw_full = full_img.copy()\n        raw_bbox = bbox_img.copy()\n\n        if self.full_transform:\n            full_img = self.full_transform(full_img)\n\n        if self.bbox_transform:\n            bbox_img = self.bbox_transform(bbox_img)\n\n        label = torch.tensor(label, dtype=torch.float32)\n\n        if self.return_raw:\n            return full_img, bbox_img, label, raw_full, raw_bbox\n\n        return full_img, bbox_img, label","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_tf = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.RandomResizedCrop(CFG.img_size, scale=(0.92, 1.0)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomRotation(3),\n    transforms.ColorJitter(brightness=0.10, contrast=0.10),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\nval_tf = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((CFG.img_size, CFG.img_size)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_weighted_sampler(labels):\n    class_counts = labels.sum(axis=0)\n    class_counts = np.maximum(class_counts, 1)\n    class_weights = 1.0 / class_counts\n\n    sample_weights = []\n\n    for lab in labels:\n        active_disease = lab[:4].astype(bool)\n\n        if lab[0] == 1:\n            sample_weights.append(class_weights[0] * CFG.lung_opacity_weight)\n        elif active_disease.any():\n            sample_weights.append(class_weights[:4][active_disease].max())\n        else:\n            sample_weights.append(class_weights[4])\n\n    sample_weights = np.array(sample_weights, dtype=np.float32)\n    sample_weights = np.maximum(sample_weights, 1e-6)\n\n    return WeightedRandomSampler(\n        weights=torch.from_numpy(sample_weights),\n        num_samples=len(sample_weights),\n        replacement=True\n    )","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds = VinDualInputDataset(tr_ids, tr_labs, train_df, full_transform=train_tf, bbox_transform=train_tf)\nval_ds = VinDualInputDataset(val_ids, val_labs, train_df, full_transform=val_tf, bbox_transform=val_tf)\ntest_ds = VinDualInputDataset(test_ids, test_labs, train_df, full_transform=val_tf, bbox_transform=val_tf)\n\nsampler = make_weighted_sampler(tr_labs)\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG.batch_size, sampler=sampler, num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\nval_loader = DataLoader(val_ds, batch_size=CFG.batch_size, shuffle=False, num_workers=CFG.num_workers, pin_memory=True)\ntest_loader = DataLoader(test_ds, batch_size=CFG.batch_size, shuffle=False, num_workers=CFG.num_workers, pin_memory=True)\n\nfull_img, bbox_img, label, raw_full, raw_bbox = VinDualInputDataset(tr_ids, tr_labs, train_df, full_transform=val_tf, bbox_transform=val_tf, return_raw=True)[0]\n\nprint(\"Full tensor:\", full_img.shape)\nprint(\"BBox tensor:\", bbox_img.shape)\nprint(\"Label:\", label)\n\nplt.figure(figsize=(10, 5))\nplt.subplot(1, 2, 1)\nplt.imshow(raw_full, cmap=\"gray\")\nplt.title(\"Full image\")\nplt.axis(\"off\")\n\nplt.subplot(1, 2, 2)\nplt.imshow(raw_bbox, cmap=\"gray\")\nplt.title(\"BBox image\")\nplt.axis(\"off\")\n\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CHEXPERT_CLASSES = [\n    \"No Finding\",\n    \"Enlarged Cardiomediastinum\",\n    \"Cardiomegaly\",\n    \"Lung Opacity\",\n    \"Lung Lesion\",\n    \"Edema\",\n    \"Consolidation\",\n    \"Pneumonia\",\n    \"Atelectasis\",\n    \"Pneumothorax\",\n    \"Pleural Effusion\",\n    \"Pleural Other\",\n    \"Fracture\",\n    \"Support Devices\"\n]\n\nclass CheXpertDataset(Dataset):\n    def __init__(self, csv_path, transform=None, policy=\"zeros\"):\n        self.df = pd.read_csv(csv_path)\n        self.transform = transform\n        self.policy = policy\n        self.labels = self.df[CHEXPERT_CLASSES].fillna(0).values.astype(np.float32)\n\n        if policy == \"zeros\":\n            self.labels[self.labels == -1] = 0.0\n        else:\n            self.labels[self.labels == -1] = 1.0\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        raw_path = self.df.iloc[idx][\"Path\"]\n        clean_path = raw_path.replace(\"CheXpert-v1.0-small/\", \"\")\n        img_path = os.path.join(CHEXPERT_DIR, clean_path)\n\n        try:\n            img = Image.open(img_path).convert(\"RGB\")\n        except:\n            img = Image.fromarray(np.zeros((224, 224, 3), dtype=np.uint8))\n\n        if self.transform:\n            img = self.transform(img)\n\n        return img, torch.tensor(self.labels[idx], dtype=torch.float32)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"chex_train_tf = transforms.Compose([\n    transforms.Resize((256, 256)),\n    transforms.RandomCrop(224),\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\nchex_val_tf = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\n\nchex_train_ds = CheXpertDataset(os.path.join(CHEXPERT_DIR, \"train.csv\"), transform=chex_train_tf, policy=\"zeros\")\nchex_val_ds = CheXpertDataset(os.path.join(CHEXPERT_DIR, \"valid.csv\"), transform=chex_val_tf, policy=\"zeros\")\n\nchex_train_loader = DataLoader(chex_train_ds, batch_size=32, shuffle=True, num_workers=CFG.num_workers, pin_memory=True)\nchex_val_loader = DataLoader(chex_val_ds, batch_size=64, shuffle=False, num_workers=CFG.num_workers, pin_memory=True)\n\nprint(\"CheXpert train:\", len(chex_train_ds))\nprint(\"CheXpert val:\", len(chex_val_ds))\n\nsample_img, sample_label = chex_train_ds[0]\n\nprint(\"CheXpert image shape:\", sample_img.shape)\nprint(\"Pixel mean:\", sample_img.mean().item())\nprint(\"Label shape:\", sample_label.shape)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AsymmetricLoss(nn.Module):\n    def __init__(self, gamma_neg=4, gamma_pos=1, clip=0.05, eps=1e-6):\n        super().__init__()\n        self.gamma_neg = gamma_neg\n        self.gamma_pos = gamma_pos\n        self.clip = clip\n        self.eps = eps\n\n    def forward(self, logits, targets):\n        xs_pos = torch.sigmoid(logits).clamp(min=self.eps, max=1 - self.eps)\n        xs_neg = (1 - xs_pos).clamp(min=self.eps, max=1 - self.eps)\n\n        if self.clip is not None:\n            xs_neg = (xs_neg + self.clip).clamp(max=1)\n\n        loss_pos = targets * torch.log(xs_pos)\n        loss_neg = (1 - targets) * torch.log(xs_neg)\n        loss = loss_pos + loss_neg\n\n        pt = xs_pos * targets + xs_neg * (1 - targets)\n        gamma = self.gamma_pos * targets + self.gamma_neg * (1 - targets)\n        loss = loss * torch.pow(1 - pt, gamma)\n\n        return -loss.mean()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EfficientNetCheXpert(nn.Module):\n    def __init__(self):\n        super().__init__()\n        backbone = models.efficientnet_v2_s(weights=models.EfficientNet_V2_S_Weights.IMAGENET1K_V1)\n        self.features = backbone.features\n        self.head = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.BatchNorm1d(1280),\n            nn.Dropout(0.5),\n            nn.Linear(1280, 512),\n            nn.GELU(),\n            nn.Dropout(0.4),\n            nn.Linear(512, 14)\n        )\n\n    def forward(self, x):\n        return self.head(self.features(x))\n\ndef compute_auc(model, loader):\n    model.eval()\n    all_probs = []\n    all_targets = []\n\n    with torch.no_grad():\n        for x, y in loader:\n            x = x.to(CFG.device)\n            probs = torch.sigmoid(model(x)).cpu().numpy()\n            all_probs.append(probs)\n            all_targets.append(y.numpy())\n\n    probs = np.vstack(all_probs)\n    targets = np.vstack(all_targets)\n\n    aucs = []\n\n    for i in range(14):\n        if len(np.unique(targets[:, i])) > 1:\n            aucs.append(roc_auc_score(targets[:, i], probs[:, i]))\n\n    return float(np.mean(aucs))\n\nCHEX_EPOCHS = 5\n\nchex_model = EfficientNetCheXpert().to(CFG.device)\nchex_criterion = AsymmetricLoss(gamma_neg=2, gamma_pos=1, clip=0.03).to(CFG.device)\n\nchex_optimizer = optim.AdamW([\n    {\"params\": chex_model.features.parameters(), \"lr\": 5e-5},\n    {\"params\": chex_model.head.parameters(), \"lr\": 2e-4}\n], weight_decay=1e-3)\n\nchex_scheduler = optim.lr_scheduler.CosineAnnealingLR(chex_optimizer, T_max=CHEX_EPOCHS)\nchex_scaler = GradScaler(\"cuda\", enabled=torch.cuda.is_available())\n\nbest_auc = 0.0\n\nprint(\"=== PRETRAIN CHEXPERT ===\")\n\nfor epoch in range(CHEX_EPOCHS):\n    chex_model.train()\n    total_loss = 0.0\n    valid_batches = 0\n\n    for x, y in tqdm(chex_train_loader, desc=f\"CheXpert Epoch {epoch}\"):\n        x = x.to(CFG.device)\n        y = y.to(CFG.device)\n\n        chex_optimizer.zero_grad()\n\n        with autocast(device_type=CFG.amp_device, enabled=torch.cuda.is_available()):\n            pred = chex_model(x)\n            loss = chex_criterion(pred, y)\n\n        if torch.isnan(loss) or torch.isinf(loss):\n            continue\n\n        chex_scaler.scale(loss).backward()\n        chex_scaler.unscale_(chex_optimizer)\n        torch.nn.utils.clip_grad_norm_(chex_model.parameters(), max_norm=1.0)\n        chex_scaler.step(chex_optimizer)\n        chex_scaler.update()\n\n        total_loss += loss.item()\n        valid_batches += 1\n\n    chex_scheduler.step()\n\n    avg_loss = total_loss / max(valid_batches, 1)\n    mean_auc = compute_auc(chex_model, chex_val_loader)\n\n    print(f\"Epoch {epoch:02d} | Loss: {avg_loss:.4f} | Mean AUC: {mean_auc:.4f}\")\n\n    if mean_auc > best_auc:\n        best_auc = mean_auc\n        torch.save(chex_model.features.state_dict(), CHEX_CKPT)\n        print(\"Lưu CheXpert backbone:\", CHEX_CKPT)\n\ndel chex_model\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EfficientNetV2DualInput(nn.Module):\n    def __init__(self, num_classes=5, chexpert_weights_path=CHEX_CKPT):\n        super().__init__()\n\n        backbone = models.efficientnet_v2_s(weights=None)\n        self.features = backbone.features\n\n        chex_state = torch.load(chexpert_weights_path, map_location=\"cpu\")\n        self.features.load_state_dict(chex_state)\n\n        self.pool = nn.AdaptiveAvgPool2d(1)\n\n        self.classifier = nn.Sequential(\n            nn.BatchNorm1d(1280 * 2),\n            nn.Dropout(0.5),\n            nn.Linear(1280 * 2, 768),\n            nn.GELU(),\n            nn.Dropout(0.4),\n            nn.Linear(768, 256),\n            nn.GELU(),\n            nn.Dropout(0.3),\n            nn.Linear(256, num_classes)\n        )\n\n    def extract_feature(self, x):\n        x = self.features(x)\n        x = self.pool(x)\n        x = torch.flatten(x, 1)\n        return x\n\n    def forward(self, full_img, bbox_img):\n        feat_full = self.extract_feature(full_img)\n        feat_bbox = self.extract_feature(bbox_img)\n        feat = torch.cat([feat_full, feat_bbox], dim=1)\n        out = self.classifier(feat)\n        return out\n\n    def freeze_backbone(self):\n        for p in self.features.parameters():\n            p.requires_grad = False\n        print(\"Backbone frozen\")\n\n    def unfreeze_backbone(self):\n        for p in self.features.parameters():\n            p.requires_grad = True\n        print(\"Backbone unfrozen\")\n\nmodel = EfficientNetV2DualInput(num_classes=CFG.num_classes, chexpert_weights_path=CHEX_CKPT).to(CFG.device)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_optimizer(model, freeze=True):\n    if freeze:\n        return optim.AdamW(\n            model.classifier.parameters(),\n            lr=CFG.lr,\n            weight_decay=1e-3\n        )\n\n    return optim.AdamW([\n        {\"params\": model.features.parameters(), \"lr\": CFG.lr_backbone},\n        {\"params\": model.classifier.parameters(), \"lr\": CFG.lr}\n    ], weight_decay=1e-3)\n\n\ndef collect_probs_targets(model, loader):\n    model.eval()\n\n    all_probs = []\n    all_targets = []\n\n    with torch.no_grad():\n        for full_img, bbox_img, y in loader:\n\n            full_img = full_img.to(CFG.device)\n            bbox_img = bbox_img.to(CFG.device)\n\n            with autocast(\n                device_type=CFG.amp_device,\n                enabled=torch.cuda.is_available()\n            ):\n                logits = model(full_img, bbox_img)\n\n            probs = torch.sigmoid(logits).cpu().numpy()\n\n            all_probs.append(probs)\n            all_targets.append(y.numpy())\n\n    return np.vstack(all_probs), np.vstack(all_targets)\n\n\ndef find_thresholds_precision_target(probs, targets, min_precision=0.70, min_recall=0.60):\n    thresholds = []\n\n    for i in range(CFG.num_classes):\n        p, r, t = precision_recall_curve(targets[:, i], probs[:, i])\n\n        if len(t) == 0:\n            thresholds.append(0.5)\n            continue\n\n        p = p[:-1]\n        r = r[:-1]\n\n        valid = np.where((p >= min_precision) & (r >= min_recall))[0]\n\n        if len(valid) > 0:\n            best_idx = valid[np.argmax(r[valid])]\n        else:\n            valid_p = np.where(p >= min_precision)[0]\n            if len(valid_p) > 0:\n                best_idx = valid_p[np.argmax(r[valid_p])]\n            else:\n                f1s = 2 * p * r / (p + r + 1e-8)\n                best_idx = np.argmax(f1s)\n\n        thresholds.append(float(t[best_idx]))\n\n    return thresholds\n\n\ndef evaluate_f1(model, loader):\n\n    probs, targets = collect_probs_targets(model, loader)\n\n    thresholds = find_thresholds_precision_target(\n        probs,\n        targets,\n        min_precision=0.70,\n        min_recall=0.70\n    )\n\n    preds = (probs > np.array(thresholds)).astype(int)\n\n    return f1_score(\n        targets,\n        preds,\n        average=\"macro\",\n        zero_division=0\n    )\n\n\ncriterion = AsymmetricLoss(\n    gamma_neg=4,\n    gamma_pos=1,\n    clip=0.05\n).to(CFG.device)\n\nscaler = GradScaler(\n    \"cuda\",\n    enabled=torch.cuda.is_available()\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history = {\"train_loss\": [], \"val_f1\": []}\nbest_val_f1 = 0.0\npatience = 7\npatience_count = 0\n\nmodel.freeze_backbone()\n\noptimizer = get_optimizer(model, freeze=True)\nscheduler = optim.lr_scheduler.OneCycleLR(optimizer, max_lr=CFG.lr, steps_per_epoch=len(train_loader), epochs=CFG.freeze_epochs)\n\nprint(\"PHASE 1 — Train classifier only\")\n\nfor epoch in range(CFG.freeze_epochs):\n    model.train()\n    total_loss = 0.0\n    valid_batches = 0\n\n    for full_img, bbox_img, y in tqdm(train_loader, desc=f\"[Frozen] Epoch {epoch}\"):\n        full_img = full_img.to(CFG.device)\n        bbox_img = bbox_img.to(CFG.device)\n        y = y.to(CFG.device)\n\n        optimizer.zero_grad()\n\n        with autocast(device_type=CFG.amp_device, enabled=torch.cuda.is_available()):\n            pred = model(full_img, bbox_img)\n            loss = criterion(pred, y)\n\n        if torch.isnan(loss) or torch.isinf(loss):\n            continue\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        scheduler.step()\n\n        total_loss += loss.item()\n        valid_batches += 1\n\n    avg_loss = total_loss / max(valid_batches, 1)\n    val_f1 = evaluate_f1(model, val_loader)\n\n    history[\"train_loss\"].append(avg_loss)\n    history[\"val_f1\"].append(val_f1)\n\n    print(f\"Epoch {epoch:02d} | Loss: {avg_loss:.4f} | Val F1: {val_f1:.4f}\")\n\n    if val_f1 > best_val_f1:\n        best_val_f1 = val_f1\n        patience_count = 0\n        torch.save(model.state_dict(), BEST_MODEL_PATH)\n        print(\"Saved best model:\", best_val_f1)\n    else:\n        patience_count += 1","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.unfreeze_backbone()\n\noptimizer = get_optimizer(model, freeze=False)\nremaining_epochs = CFG.epochs - CFG.freeze_epochs\nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=remaining_epochs, eta_min=1e-7)\n\nprint(\"PHASE 2 — Fine-tune full dual-input model\")\n\nfor epoch in range(CFG.freeze_epochs, CFG.epochs):\n    model.train()\n    total_loss = 0.0\n    valid_batches = 0\n\n    for full_img, bbox_img, y in tqdm(train_loader, desc=f\"[FineTune] Epoch {epoch}\"):\n        full_img = full_img.to(CFG.device)\n        bbox_img = bbox_img.to(CFG.device)\n        y = y.to(CFG.device)\n\n        optimizer.zero_grad()\n\n        with autocast(device_type=CFG.amp_device, enabled=torch.cuda.is_available()):\n            pred = model(full_img, bbox_img)\n            loss = criterion(pred, y)\n\n        if torch.isnan(loss) or torch.isinf(loss):\n            continue\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        scaler.step(optimizer)\n        scaler.update()\n\n        total_loss += loss.item()\n        valid_batches += 1\n\n    scheduler.step()\n\n    avg_loss = total_loss / max(valid_batches, 1)\n    val_f1 = evaluate_f1(model, val_loader)\n\n    history[\"train_loss\"].append(avg_loss)\n    history[\"val_f1\"].append(val_f1)\n\n    print(f\"Epoch {epoch:02d} | Loss: {avg_loss:.4f} | Val F1: {val_f1:.4f}\")\n\n    if val_f1 > best_val_f1:\n        best_val_f1 = val_f1\n        patience_count = 0\n        torch.save(model.state_dict(), BEST_MODEL_PATH)\n        print(\"Saved best model:\", best_val_f1)\n    else:\n        patience_count += 1\n\n    if patience_count >= patience:\n        print(\"Early stopping\")\n        break","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(\n    torch.load(\n        BEST_MODEL_PATH,\n        map_location=CFG.device\n    )\n)\n\nmodel.to(CFG.device)\nmodel.eval()\n\nval_probs, val_targets = collect_probs_targets(\n    model,\n    val_loader\n)\n\nbest_thresholds = find_thresholds_precision_target(\n    val_probs,\n    val_targets,\n    min_precision=0.8,\n    min_recall=0.65\n)\n\nbest_thresholds[-1] = max(best_thresholds[-1], 0.35)\nbest_thresholds[0] = 0.75\n\nprint(\"Best thresholds:\")\n\nfor name, thr in zip(CLASS_NAMES, best_thresholds):\n    print(f\"{name}: {thr:.3f}\")\n\ntest_probs, test_targets = collect_probs_targets(\n    model,\n    test_loader\n)\n\ntest_preds = (\n    test_probs > np.array(best_thresholds)\n).astype(int)\n\nprint(\"=\" * 60)\nprint(\"CLASSIFICATION REPORT — TEST SET\")\nprint(\"=\" * 60)\n\nprint(\n    classification_report(\n        test_targets,\n        test_preds,\n        target_names=CLASS_NAMES,\n        zero_division=0\n    )\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"mcm = multilabel_confusion_matrix(test_targets, test_preds)\n\nfig, axes = plt.subplots(1, 5, figsize=(25, 5))\n\nfor i in range(5):\n    tn, fp, fn, tp = mcm[i].ravel()\n    precision = tp / (tp + fp + 1e-8)\n    recall = tp / (tp + fn + 1e-8)\n\n    sns.heatmap(\n        mcm[i],\n        annot=True,\n        fmt=\"d\",\n        cmap=\"Blues\",\n        ax=axes[i],\n        cbar=False,\n        xticklabels=[\"Pred Neg\", \"Pred Pos\"],\n        yticklabels=[\"True Neg\", \"True Pos\"]\n    )\n\n    axes[i].set_title(f\"{CLASS_NAMES[i]}\\nP={precision:.2f} R={recall:.2f}\")\n    axes[i].set_xlabel(\"Pred\")\n    axes[i].set_ylabel(\"True\")\n\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/confusion_matrices_dual_input.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(12, 5))\n\nplt.subplot(1, 2, 1)\nplt.plot(history[\"train_loss\"], marker=\"o\")\nplt.title(\"Train Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.grid(True, alpha=0.3)\n\nplt.subplot(1, 2, 2)\nplt.plot(history[\"val_f1\"], marker=\"o\")\nplt.axhline(best_val_f1, linestyle=\"--\")\nplt.title(\"Validation F1\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"F1 Macro\")\nplt.grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/training_curves_dual_input.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def predict_with_tta(model, ids, labels, df, n_aug=5):\n#     tta_tf = transforms.Compose([\n#         transforms.ToPILImage(),\n#         transforms.Resize((CFG.img_size, CFG.img_size)),\n#         transforms.RandomHorizontalFlip(p=0.5),\n#         transforms.RandomRotation(5),\n#         transforms.ToTensor(),\n#         transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n#     ])\n\n#     all_probs = []\n\n#     for run in range(n_aug):\n#         dataset = VinDualInputDataset(\n#             ids,\n#             labels,\n#             df,\n#             full_transform=tta_tf,\n#             bbox_transform=tta_tf\n#         )\n\n#         loader = DataLoader(\n#             dataset,\n#             batch_size=CFG.batch_size,\n#             shuffle=False,\n#             num_workers=CFG.num_workers,\n#             pin_memory=True\n#         )\n\n#         probs, _ = collect_probs_targets(model, loader)\n#         all_probs.append(probs)\n\n#         print(f\"TTA run {run + 1}/{n_aug} done\")\n\n#     return np.mean(all_probs, axis=0)\n\n\n# tta_probs = predict_with_tta(model, test_ids, test_labs, train_df, n_aug=5)\n# tta_preds = (tta_probs > np.array(best_thresholds)).astype(int)\n\n# print(\"=\" * 60)\n# print(\"CLASSIFICATION REPORT — TEST SET WITH TTA\")\n# print(\"=\" * 60)\n# print(classification_report(test_targets, tta_preds, target_names=CLASS_NAMES, zero_division=0))","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}