{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpuV5e8","dataSources":[{"sourceId":112899,"databundleVersionId":13449579,"sourceType":"competition"},{"sourceId":113002,"databundleVersionId":13471427,"isSourceIdPinned":false,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nos.environ.setdefault(\"TF_CPP_MIN_LOG_LEVEL\", \"3\")\n# Mixed precision on TPU\nos.environ.setdefault(\"XLA_USE_BF16\", \"1\")  # bf16 compute\nos.environ.setdefault(\"XLA_IR_DEBUG\", \"0\")\n\nimport cv2\nimport math\nimport numpy as np\nimport pandas as pd\nimport random\nfrom dataclasses import dataclass\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader, Subset\nimport torchvision.transforms as T\nfrom tqdm import tqdm\n\nimport torch_xla\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\n\nfrom sklearn.metrics import roc_auc_score","metadata":{"_uuid":"debd3eb3-4459-49f6-a8bd-cfadab0f842f","_cell_guid":"d5eb5bde-a92a-43f3-8b9d-dd3b2ebf4b24","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:20:18.830529Z","iopub.execute_input":"2025-10-12T06:20:18.83089Z","iopub.status.idle":"2025-10-12T06:20:18.836267Z","shell.execute_reply.started":"2025-10-12T06:20:18.83086Z","shell.execute_reply":"2025-10-12T06:20:18.835327Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LABEL_COLUMNS = [\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]\nNUM_LABELS = len(LABEL_COLUMNS)\n\n# -------------------------\n# Reproducibility\n# -------------------------\ndef seed_everything(seed=1337):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)","metadata":{"_uuid":"0bef644d-aeef-44e0-91c0-5c95a0e44709","_cell_guid":"108650af-ec79-4d77-9635-c2f62c303114","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:20:19.031349Z","iopub.execute_input":"2025-10-12T06:20:19.031722Z","iopub.status.idle":"2025-10-12T06:20:19.035578Z","shell.execute_reply.started":"2025-10-12T06:20:19.031697Z","shell.execute_reply":"2025-10-12T06:20:19.034582Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from PIL import Image, ImageFile\nImageFile.LOAD_TRUNCATED_IMAGES = True\n\nclass XRayDataset(Dataset):\n    def __init__(self, df, image_dir, transform):\n        self.df = df.reset_index(drop=True)\n        self.image_dir = image_dir\n        self.transform = transform\n        self.bad_files = []\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        p = os.path.join(self.image_dir, str(row[\"Image_name\"]))\n        img = Image.open(p).convert(\"RGB\")\n        img = self.transform(img)\n        labels = torch.from_numpy(row[LABEL_COLUMNS].to_numpy(dtype=np.float32, na_value=0.0))\n        return img, labels\n\n    def __len__(self): return len(self.df)","metadata":{"_uuid":"ecadfe4f-607c-40d4-bbca-d1b0f885f84d","_cell_guid":"aa2574a4-4450-4b63-af9c-13edac49afdf","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:20:19.241021Z","iopub.execute_input":"2025-10-12T06:20:19.241318Z","iopub.status.idle":"2025-10-12T06:20:19.247076Z","shell.execute_reply.started":"2025-10-12T06:20:19.241293Z","shell.execute_reply":"2025-10-12T06:20:19.246276Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, alpha=None, gamma=2.0, reduction='mean'):\n        super().__init__()\n        self.register_buffer(\"alpha\", alpha if isinstance(alpha, torch.Tensor) else None)\n        self.gamma = gamma\n        self.reduction = reduction\n\n    def forward(self, inputs, targets):\n        bce = nn.functional.binary_cross_entropy_with_logits(inputs, targets, reduction='none')\n        pt = torch.exp(-bce)\n        if self.alpha is not None:\n            bce = self.alpha * bce  # shape [C] broadcasts to [B,C]\n        loss = (1 - pt) ** self.gamma * bce\n        return loss.mean() if self.reduction == 'mean' else loss.sum()","metadata":{"_uuid":"6bb65f7e-f999-42ac-a2ad-2e072a75def5","_cell_guid":"30adff1e-558e-47b0-80b0-971f8baff962","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:20:19.416948Z","iopub.execute_input":"2025-10-12T06:20:19.41724Z","iopub.status.idle":"2025-10-12T06:20:19.421828Z","shell.execute_reply.started":"2025-10-12T06:20:19.417216Z","shell.execute_reply":"2025-10-12T06:20:19.420918Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@dataclass\nclass TrainConfig:\n    train_csv: str = \"/kaggle/input/grand-xray-slam-division-b/train2.csv\"\n    train_dir: str = \"/kaggle/input/grand-xray-slam-division-b/train2\"\n    val_split: float = 0.1\n    batch_size: int = 64              # ViT-L is large; tune up/down\n    num_workers: int = 8\n    lr: float = 3e-5                  # typical for ViT-L finetune\n    weight_decay: float = 0.05\n    epochs: int = 10\n    warmup_pct: float = 0.05\n    grad_accum_steps: int = 2         # effective batch = batch_size * accum\n    use_focal: bool = True            # set False to use BCEWithLogitsLoss(pos_weight)\n    use_checkpointing: bool = True    # helps memory\n    seed: int = 1337\n    save_path: str = \"/kaggle/working/vit_l16_multilabel_tpu.pth\"","metadata":{"_uuid":"f6071196-cbfc-43ee-9a0f-16775b4b1ef8","_cell_guid":"7386c616-bf8d-4311-90a0-72d4342173fc","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:20:19.614313Z","iopub.execute_input":"2025-10-12T06:20:19.614605Z","iopub.status.idle":"2025-10-12T06:20:19.619244Z","shell.execute_reply.started":"2025-10-12T06:20:19.614587Z","shell.execute_reply":"2025-10-12T06:20:19.61843Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_transforms(weights):\n    # We’ll use the weights’ eval transforms as base (ensures resize/crop/normalize @384)\n    eval_tf = weights.transforms(antialias=True)\n    img_size = eval_tf.crop_size[0]\n\n    # For train: light, CXR-friendly augmentation.\n    # Avoid horizontal flips (left/right matters for CXR).\n    train_tf = T.Compose([\n        T.Resize((img_size, img_size), antialias=True),\n        T.ToTensor(),\n        T.Normalize(mean=eval_tf.mean, std=eval_tf.std),\n    ])\n    # For val/test: just use weights’ eval pipeline, adapted to numpy input\n    val_tf = T.Compose([\n        T.Resize((img_size, img_size), antialias=True),\n        T.ToTensor(),\n        T.Normalize(mean=eval_tf.mean, std=eval_tf.std),\n    ])\n    return train_tf, val_tf","metadata":{"_uuid":"2a26c87e-1a8b-4538-807f-58d89fc2760d","_cell_guid":"20b67599-6d7f-4cf0-92ea-41d55dbbf944","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:39.574773Z","iopub.execute_input":"2025-10-12T06:22:39.575122Z","iopub.status.idle":"2025-10-12T06:22:39.580571Z","shell.execute_reply.started":"2025-10-12T06:22:39.575096Z","shell.execute_reply":"2025-10-12T06:22:39.579428Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def make_loaders(df, cfg, train_tf, val_tf):\n    # stratified split for multilabel is non-trivial; use random split with seed\n    n = len(df)\n    idx = np.arange(n)\n    np.random.shuffle(idx)\n    split = int(n * (1 - cfg.val_split))\n    tr_idx, va_idx = idx[:split], idx[split:]\n\n    train_ds = XRayDataset(df.iloc[tr_idx], cfg.train_dir, transform=train_tf)\n    val_ds   = XRayDataset(df.iloc[va_idx], cfg.train_dir, transform=val_tf)\n\n    train_dl = DataLoader(train_ds, batch_size=cfg.batch_size, shuffle=True,\n                          num_workers=cfg.num_workers, persistent_workers=True, drop_last=True)\n    val_dl   = DataLoader(val_ds, batch_size=cfg.batch_size, shuffle=False,\n                          num_workers=cfg.num_workers, persistent_workers=True)\n    return train_dl, val_dl","metadata":{"_uuid":"98d69899-e515-4065-aced-0c3fba56dc66","_cell_guid":"569df7a2-0775-4679-bae9-448d4c4b75d1","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:39.767661Z","iopub.execute_input":"2025-10-12T06:22:39.767901Z","iopub.status.idle":"2025-10-12T06:22:39.772356Z","shell.execute_reply.started":"2025-10-12T06:22:39.767883Z","shell.execute_reply":"2025-10-12T06:22:39.771346Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_model():\n    from torchvision.models import vit_l_16, ViT_L_16_Weights\n    # Prefer E2E weights for full fine-tuning\n    weights = ViT_L_16_Weights.IMAGENET1K_SWAG_LINEAR_V1\n    model = vit_l_16(weights=weights)\n    in_features = model.heads.head.in_features\n    model.heads.head = nn.Linear(in_features, NUM_LABELS)\n    return model, weights","metadata":{"_uuid":"56cbce88-1052-43cd-bae1-8717665684b8","_cell_guid":"daf81758-e6ad-439d-ae56-ee6bdc38e582","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:39.958151Z","iopub.execute_input":"2025-10-12T06:22:39.958463Z","iopub.status.idle":"2025-10-12T06:22:39.961671Z","shell.execute_reply.started":"2025-10-12T06:22:39.958449Z","shell.execute_reply":"2025-10-12T06:22:39.960939Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, scheduler, criterion, device, cfg: TrainConfig):\n    model.train()\n    total_loss = 0.0\n    step = 0\n\n    for imgs, labels in loader:\n        imgs = imgs.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        with torch.autocast(device_type=\"xla\", dtype=torch.bfloat16):\n            logits = model(imgs)\n            loss = criterion(logits, labels) / cfg.grad_accum_steps\n\n        loss.backward()\n        step += 1\n\n        if step % cfg.grad_accum_steps == 0:\n            xm.optimizer_step(optimizer)\n            optimizer.zero_grad(set_to_none=True)\n            if scheduler is not None:\n                scheduler.step()\n        total_loss += loss.item() * cfg.grad_accum_steps  # track real loss\n\n    # Only one print from master\n    avg_loss = total_loss / len(loader)\n    xm.master_print(f\"train loss: {avg_loss:.4f}\")\n    return avg_loss","metadata":{"_uuid":"95e993d7-e8c5-4ca7-bbeb-1c6e353e0bf6","_cell_guid":"fb71fe0b-d1fe-494d-a6ab-1ab6d46197d6","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:40.167463Z","iopub.execute_input":"2025-10-12T06:22:40.167636Z","iopub.status.idle":"2025-10-12T06:22:40.173274Z","shell.execute_reply.started":"2025-10-12T06:22:40.167621Z","shell.execute_reply":"2025-10-12T06:22:40.172147Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef evaluate(model, loader, device):\n    model.eval()\n    all_logits = []\n    all_labels = []\n\n    for imgs, labels in loader:\n        imgs = imgs.to(device, non_blocking=True)\n        labels = labels.to(device, non_blocking=True)\n\n        with torch.autocast(device_type=\"xla\", dtype=torch.bfloat16):\n            logits = model(imgs)\n\n        all_logits.append(logits.float().cpu())\n        all_labels.append(labels.float().cpu())\n\n    all_logits = torch.cat(all_logits, dim=0).numpy()\n    all_labels = torch.cat(all_labels, dim=0).numpy()\n\n    # probs for AUC\n    probs = 1.0 / (1.0 + np.exp(-all_logits))\n\n    per_class_auc = []\n    for c in range(NUM_LABELS):\n        y_true = all_labels[:, c]\n        y_pred = probs[:, c]\n        if np.unique(y_true).size < 2:\n            per_class_auc.append(np.nan)  # undefined\n        else:\n            per_class_auc.append(roc_auc_score(y_true, y_pred))\n\n    macro_auc = np.nanmean(per_class_auc)\n    return macro_auc, per_class_auc","metadata":{"_uuid":"838db771-0bb8-4525-a090-fc0db04b7dd1","_cell_guid":"b29617b1-427e-49be-acf1-e649f906244b","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:42.410286Z","iopub.execute_input":"2025-10-12T06:22:42.410611Z","iopub.status.idle":"2025-10-12T06:22:42.416379Z","shell.execute_reply.started":"2025-10-12T06:22:42.410594Z","shell.execute_reply":"2025-10-12T06:22:42.415275Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_scheduler(optimizer, cfg: TrainConfig, steps_per_epoch: int):\n    total_steps = cfg.epochs * math.ceil(steps_per_epoch / cfg.grad_accum_steps)\n    warmup_steps = int(cfg.warmup_pct * total_steps)\n    def lr_lambda(step):\n        if step < warmup_steps:\n            return float(step) / max(1, warmup_steps)\n        # cosine decay to 10% of base LR\n        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n        return 0.1 + 0.9 * (1.0 + math.cos(math.pi * progress)) / 2.0\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=lr_lambda)","metadata":{"_uuid":"94e20c29-919e-4e1e-9833-8e5d140b7933","_cell_guid":"fd601610-1ca8-4449-aa59-c1c2840e1015","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:42.746833Z","iopub.execute_input":"2025-10-12T06:22:42.747161Z","iopub.status.idle":"2025-10-12T06:22:42.751064Z","shell.execute_reply.started":"2025-10-12T06:22:42.747135Z","shell.execute_reply":"2025-10-12T06:22:42.750243Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seed_everything()\ncfg = TrainConfig()\n\n# Device / TPU\ndevice = xm.xla_device()\nxm.master_print(\"Using device: {}\".format(device))","metadata":{"_uuid":"16f305a8-6d9f-41d9-a302-b56153c7113a","_cell_guid":"a26c6224-812e-4471-901b-159273a95282","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:44.063935Z","iopub.execute_input":"2025-10-12T06:22:44.064237Z","iopub.status.idle":"2025-10-12T06:22:44.225918Z","shell.execute_reply.started":"2025-10-12T06:22:44.064216Z","shell.execute_reply":"2025-10-12T06:22:44.22478Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\nfrom PIL import Image, UnidentifiedImageError\nfrom tqdm import tqdm\n\n# -------------------------\n# Load CSV and convert labels to float\n# -------------------------\nlabel_dtypes = {c: \"float32\" for c in LABEL_COLUMNS}\ndf = pd.read_csv(cfg.train_csv, dtype=label_dtypes)\ndf[LABEL_COLUMNS] = df[LABEL_COLUMNS].apply(pd.to_numeric, errors='coerce').fillna(0).astype(\"float32\")\n\n# -------------------------\n# Filter out images that do not exist or are unreadable\n# -------------------------\nimage_root = Path(cfg.train_dir)\nvalid_indices = []\n\nprint(\"[INFO] Checking image file validity...\")\n\nfor i, row in tqdm(df.iterrows()):\n    img_path = image_root / str(row[\"Image_name\"])\n    try:\n        # Fast path existence check\n        if not img_path.exists():\n            continue\n        # Quick open test (ensures not corrupted)\n        with Image.open(img_path) as im:\n            im.verify()  # lightweight consistency check\n        valid_indices.append(i)\n    except (FileNotFoundError, UnidentifiedImageError, OSError):\n        continue\n\ndropped = len(df) - len(valid_indices)\nif dropped > 0:\n    print(f\"[INFO] Dropping {dropped} missing or unreadable images from training data.\")\nelse:\n    print(\"[INFO] All images verified successfully.\")\n\ndf = df.loc[valid_indices].reset_index(drop=True)","metadata":{"_uuid":"c2fb2885-d46e-45bc-87f0-ca953169f148","_cell_guid":"0e511017-2843-41bf-95b5-1fac7444dfaa","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T07:05:17.099684Z","iopub.execute_input":"2025-10-12T07:05:17.100002Z","iopub.status.idle":"2025-10-12T07:05:25.530521Z","shell.execute_reply.started":"2025-10-12T07:05:17.099984Z","shell.execute_reply":"2025-10-12T07:05:25.529181Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Class weights (for BCE pos_weight) and Focal alpha\npos_counts = df[LABEL_COLUMNS].sum()\nneg_counts = len(df) - pos_counts\npos_weight_vec = (neg_counts / (pos_counts + 1e-6)).values.astype(np.float32)\nalpha_vec = torch.tensor(pos_weight_vec, dtype=torch.float32)  # reuse for focal scaling","metadata":{"_uuid":"1e42e95a-24ed-471d-b9f0-e9c2a45a0693","_cell_guid":"40912715-5bae-431a-b852-37f8582a1b1f","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:45.862222Z","iopub.execute_input":"2025-10-12T06:22:45.862574Z","iopub.status.idle":"2025-10-12T06:22:45.872416Z","shell.execute_reply.started":"2025-10-12T06:22:45.862553Z","shell.execute_reply":"2025-10-12T06:22:45.871506Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Model + transforms\nmodel, weights = build_model()\nmodel = model.to(device)","metadata":{"_uuid":"2b46aa37-e994-4a38-8637-086935131065","_cell_guid":"3a3ce73d-dd14-4f2c-94b5-99106bed011c","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:46.066111Z","iopub.execute_input":"2025-10-12T06:22:46.066492Z","iopub.status.idle":"2025-10-12T06:22:49.438325Z","shell.execute_reply.started":"2025-10-12T06:22:46.066475Z","shell.execute_reply":"2025-10-12T06:22:49.437072Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_tf, val_tf = build_transforms(weights)\ntrain_loader, val_loader = make_loaders(df, cfg, train_tf, val_tf)\n\n# XLA device loader\ntrain_loader = pl.MpDeviceLoader(train_loader, device)\nval_loader   = pl.MpDeviceLoader(val_loader, device)","metadata":{"_uuid":"367e8d1a-d8fb-4a02-b653-3a5a686211fb","_cell_guid":"c6fc433b-9975-42d8-8d5d-afe4143ecc5e","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:49.439049Z","iopub.execute_input":"2025-10-12T06:22:49.439243Z","iopub.status.idle":"2025-10-12T06:22:49.484148Z","shell.execute_reply.started":"2025-10-12T06:22:49.439225Z","shell.execute_reply":"2025-10-12T06:22:49.483033Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay, betas=(0.9, 0.999))\nif cfg.use_focal:\n    criterion = FocalLoss(alpha=alpha_vec.to(device), gamma=2.0)\nelse:\n    # BCE with positive class weights for imbalance\n    criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor(pos_weight_vec, device=device))\n\nscheduler = build_scheduler(optimizer, cfg, steps_per_epoch=len(train_loader))","metadata":{"_uuid":"3682894b-5389-4de2-a3ca-bb302128e65e","_cell_guid":"5521ec7f-f081-485f-bcaa-b71262712f71","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:49.484837Z","iopub.execute_input":"2025-10-12T06:22:49.485018Z","iopub.status.idle":"2025-10-12T06:22:49.505826Z","shell.execute_reply.started":"2025-10-12T06:22:49.485001Z","shell.execute_reply":"2025-10-12T06:22:49.504862Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nimgs, _ = next(iter(val_loader))\nwith torch.no_grad(), torch.autocast(device_type=\"xla\", dtype=torch.bfloat16):\n    out = model(imgs)\nmodel.train()","metadata":{"_uuid":"ab83fa72-f80f-404e-90c4-c319f53bb54f","_cell_guid":"72488a4f-fc99-4c4e-bef6-8d53213e4ad7","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:49.507357Z","iopub.execute_input":"2025-10-12T06:22:49.507558Z","iopub.status.idle":"2025-10-12T06:22:55.634521Z","shell.execute_reply.started":"2025-10-12T06:22:49.507542Z","shell.execute_reply":"2025-10-12T06:22:55.633323Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_auc = -1.0\nfor epoch in range(1, cfg.epochs + 1):\n    xm.master_print(f\"\\n===== Epoch {epoch}/{cfg.epochs} =====\")\n    train_one_epoch(model, train_loader, optimizer, scheduler, criterion, device, cfg)\n    val_auc, per_class_auc = evaluate(model, val_loader, device)\n\n    # Reduce across devices if you later use multiple cores (safe on single-core too)\n    xm.master_print(f\"val macro ROC-AUC: {val_auc:.4f}\")\n    # Optional: print a few class AUCs\n    top_show = min(5, NUM_LABELS)\n    show_pairs = list(zip(LABEL_COLUMNS[:top_show], per_class_auc[:top_show]))\n    xm.master_print(\"sample per-class AUCs: \" + \", \".join(f\"{n}: {a:.3f}\" for n, a in show_pairs if a==a))\n\n    if val_auc > best_auc:\n        best_auc = val_auc\n        xm.master_print(f\"New best AUC {best_auc:.4f}. Saving to {cfg.save_path}\")\n        xm.save(model.state_dict(), cfg.save_path)\n\nxm.master_print(f\"Training complete. Best macro ROC-AUC: {best_auc:.4f}\")","metadata":{"_uuid":"cac63882-8b58-4d11-887b-c548c3cee546","_cell_guid":"0f6ed8fb-c754-496b-b981-b7613dde2a63","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:22:57.603739Z","iopub.execute_input":"2025-10-12T06:22:57.604051Z","iopub.status.idle":"2025-10-12T06:34:35.793639Z","shell.execute_reply.started":"2025-10-12T06:22:57.604027Z","shell.execute_reply":"2025-10-12T06:34:35.792411Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, image_dir, transform):\n        self.image_dir = image_dir\n        self.transform = transform\n        # Case-insensitive filter; keep deterministic order\n        self.images = sorted([f for f in os.listdir(image_dir) if f.lower().endswith(\".jpg\")])\n        \n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        name = self.images[idx]\n        path = os.path.join(self.image_dir, name)\n        img = Image.open(path).convert(\"RGB\")\n        img = self.transform(img)\n        return name, img","metadata":{"_uuid":"8b7ad1f6-fbc3-460c-b20a-4aaf2456c65c","_cell_guid":"089a76e3-ee56-4b2c-9c27-3de58046146a","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:40:26.014942Z","iopub.execute_input":"2025-10-12T06:40:26.015236Z","iopub.status.idle":"2025-10-12T06:40:26.019726Z","shell.execute_reply.started":"2025-10-12T06:40:26.015216Z","shell.execute_reply":"2025-10-12T06:40:26.018835Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TEST_DIR = \"/kaggle/input/grand-xray-slam-division-b/test2\"\n\ntest_dataset = TestDataset(TEST_DIR, transform=val_tf)\ntest_loader = DataLoader(test_dataset, batch_size=64, shuffle=False, num_workers=4)\ntest_loader = pl.MpDeviceLoader(test_loader, device)","metadata":{"_uuid":"da3fd1d4-a5d7-4868-8d90-2fde1a5e7d56","_cell_guid":"0c515eda-23ff-46b3-a82c-133149418c26","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:40:39.389825Z","iopub.execute_input":"2025-10-12T06:40:39.39013Z","iopub.status.idle":"2025-10-12T06:40:40.176669Z","shell.execute_reply.started":"2025-10-12T06:40:39.390109Z","shell.execute_reply":"2025-10-12T06:40:40.175375Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nsubmission = []\n\nwith torch.no_grad():\n    for batch in tqdm(test_loader, desc=\"Inference\"):\n        img_names, imgs = batch\n        imgs = imgs.to(device)\n        outputs = model(imgs)\n        probs = torch.sigmoid(outputs).cpu().numpy()\n\n        for name, prob in zip(img_names, probs):\n            row = [name] + prob.tolist()\n            submission.append(row)\n\n# -------------------------\n# Save Submission\n# -------------------------\nsubmission_df = pd.DataFrame(submission, columns=[\"Image_name\"] + LABEL_COLUMNS)\nSUBMISSION_CSV = \"/kaggle/working/submission.csv\"\nsubmission_df.to_csv(SUBMISSION_CSV, index=False)\nprint(f\"Submission file saved to {SUBMISSION_CSV}\")","metadata":{"_uuid":"323605f2-22d0-4504-aeaf-0227c99b222f","_cell_guid":"371182be-7560-4067-8ae3-608f3df67836","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-10-12T06:41:07.851938Z","iopub.execute_input":"2025-10-12T06:41:07.852306Z","iopub.status.idle":"2025-10-12T06:52:02.315521Z","shell.execute_reply.started":"2025-10-12T06:41:07.852285Z","shell.execute_reply":"2025-10-12T06:52:02.314282Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"_uuid":"19d808c6-4653-4b8a-973d-c7e8ae359fc5","_cell_guid":"7be674cf-ee52-4c54-8cfc-5df51336f6ea","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}