{"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":113002,"databundleVersionId":13471427,"sourceType":"competition"},{"sourceId":13373222,"sourceType":"datasetVersion","datasetId":8468715},{"sourceId":13586125,"sourceType":"datasetVersion","datasetId":8631720}],"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,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-11-02T17:39:02.309532Z","iopub.execute_input":"2025-11-02T17:39:02.309802Z","iopub.status.idle":"2025-11-02T17:39:20.898354Z","shell.execute_reply.started":"2025-11-02T17:39:02.309781Z","shell.execute_reply":"2025-11-02T17:39:20.897119Z"}},"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-11-02T17:39:20.898851Z","iopub.execute_input":"2025-11-02T17:39:20.899185Z","iopub.status.idle":"2025-11-02T17:39:20.902704Z","shell.execute_reply.started":"2025-11-02T17:39:20.899167Z","shell.execute_reply":"2025-11-02T17:39:20.901879Z"},"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-11-02T17:39:20.903168Z","iopub.execute_input":"2025-11-02T17:39:20.903316Z","iopub.status.idle":"2025-11-02T17:39:20.919333Z","shell.execute_reply.started":"2025-11-02T17:39:20.903303Z","shell.execute_reply":"2025-11-02T17:39:20.918426Z"},"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()\n\n\nclass AsymmetricLossOptimized(nn.Module):\n    ''' Notice - optimized version, minimizes memory allocation and gpu uploading,\n    favors inplace operations'''\n\n    def __init__(self, gamma_neg=4, gamma_pos=1, clip=0.05, eps=1e-8, disable_torch_grad_focal_loss=False):\n        super(AsymmetricLossOptimized, self).__init__()\n\n        self.gamma_neg = gamma_neg\n        self.gamma_pos = gamma_pos\n        self.clip = clip\n        self.disable_torch_grad_focal_loss = disable_torch_grad_focal_loss\n        self.eps = eps\n\n        # prevent memory allocation and gpu uploading every iteration, and encourages inplace operations\n        self.targets = self.anti_targets = self.xs_pos = self.xs_neg = self.asymmetric_w = self.loss = None\n\n    def forward(self, x, y):\n        \"\"\"\"\n        Parameters\n        ----------\n        x: input logits\n        y: targets (multi-label binarized vector)\n        \"\"\"\n\n        self.targets = y\n        self.anti_targets = 1 - y\n\n        # Calculating Probabilities\n        self.xs_pos = torch.sigmoid(x)\n        self.xs_neg = 1.0 - self.xs_pos\n\n        # Asymmetric Clipping\n        if self.clip is not None and self.clip > 0:\n            self.xs_neg.add_(self.clip).clamp_(max=1)\n\n        # Basic CE calculation\n        self.loss = self.targets * torch.log(self.xs_pos.clamp(min=self.eps))\n        self.loss.add_(self.anti_targets * torch.log(self.xs_neg.clamp(min=self.eps)))\n\n        # Asymmetric Focusing\n        if self.gamma_neg > 0 or self.gamma_pos > 0:\n            if self.disable_torch_grad_focal_loss:\n                torch.set_grad_enabled(False)\n            self.xs_pos = self.xs_pos * self.targets\n            self.xs_neg = self.xs_neg * self.anti_targets\n            self.asymmetric_w = torch.pow(1 - self.xs_pos - self.xs_neg,\n                                          self.gamma_pos * self.targets + self.gamma_neg * self.anti_targets)\n            if self.disable_torch_grad_focal_loss:\n                torch.set_grad_enabled(True)\n            self.loss *= self.asymmetric_w\n\n        return -self.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-11-02T17:57:10.070995Z","iopub.execute_input":"2025-11-02T17:57:10.071257Z","iopub.status.idle":"2025-11-02T17:57:10.078442Z","shell.execute_reply.started":"2025-11-02T17:57:10.071238Z","shell.execute_reply":"2025-11-02T17:57:10.077603Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@dataclass\nclass TrainConfig:\n    train_csv: str = \"/kaggle/input/division-b-tpu-csv/cleaned_df.csv\"\n    train_dir: str = \"/kaggle/input/grand-xray-slam-division-b/train2\"\n    pretrained = \"/kaggle/input/eva-x-chestxray/eva_x_base_patch16_merged520k_mim_cxr14_ft.pth\"\n    img_size: int = 224\n    val_split: float = 0.1\n    batch_size: int = 64\n    num_workers: int = 8\n    lr: float = 1e-4\n    weight_decay: float = 0.05\n    epochs: int = 4\n    warmup_pct: float = 0.05\n    grad_accum_steps: int = 2\n    use_focal: bool = True\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-11-02T17:39:20.929574Z","iopub.execute_input":"2025-11-02T17:39:20.929723Z","iopub.status.idle":"2025-11-02T17:39:20.943548Z","shell.execute_reply.started":"2025-11-02T17:39:20.92971Z","shell.execute_reply":"2025-11-02T17:39:20.942955Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_transforms(img_size):\n    mean=(0.49185243, 0.49185243, 0.49185243)\n    std=(0.28509309, 0.28509309, 0.28509309)\n    \n    train_tf = T.Compose([\n        T.Resize((img_size, img_size), antialias=True),\n        T.ToTensor(),\n        T.Normalize(mean=mean, std=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=mean, std=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-11-02T17:39:20.943907Z","iopub.execute_input":"2025-11-02T17:39:20.944052Z","iopub.status.idle":"2025-11-02T17:39:20.953858Z","shell.execute_reply.started":"2025-11-02T17:39:20.944039Z","shell.execute_reply":"2025-11-02T17:39:20.953243Z"},"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-11-02T17:39:20.95447Z","iopub.execute_input":"2025-11-02T17:39:20.954617Z","iopub.status.idle":"2025-11-02T17:39:20.963145Z","shell.execute_reply.started":"2025-11-02T17:39:20.954604Z","shell.execute_reply":"2025-11-02T17:39:20.962264Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom timm.models.eva import Eva\nfrom timm.layers import resample_abs_pos_embed, resample_patch_embed\n\ndef checkpoint_filter_fn(\n        state_dict,\n        model,\n        interpolation='bicubic',\n        antialias=True,\n):\n    \"\"\" convert patch embedding weight from manual patchify + linear proj to conv\"\"\"\n    out_dict = {}\n    state_dict = state_dict.get('model_ema', state_dict)\n    state_dict = state_dict.get('model', state_dict)\n    state_dict = state_dict.get('module', state_dict)\n    state_dict = state_dict.get('state_dict', state_dict)\n    # prefix for loading OpenCLIP compatible weights\n    if 'visual.trunk.pos_embed' in state_dict:\n        prefix = 'visual.trunk.'\n    elif 'visual.pos_embed' in state_dict:\n        prefix = 'visual.'\n    else:\n        prefix = ''\n    mim_weights = prefix + 'mask_token' in state_dict\n    no_qkv = prefix + 'blocks.0.attn.q_proj.weight' in state_dict\n\n    len_prefix = len(prefix)\n    for k, v in state_dict.items():\n        if prefix:\n            if k.startswith(prefix):\n                k = k[len_prefix:]\n            else:\n                continue\n\n        if 'rope' in k:\n            # fixed embedding no need to load buffer from checkpoint\n            continue\n\n        if 'patch_embed.proj.weight' in k:\n            _, _, H, W = model.patch_embed.proj.weight.shape\n            if v.shape[-1] != W or v.shape[-2] != H:\n                v = resample_patch_embed(\n                    v,\n                    (H, W),\n                    interpolation=interpolation,\n                    antialias=antialias,\n                    verbose=True,\n                )\n        elif k == 'pos_embed' and v.shape[1] != model.pos_embed.shape[1]:\n            num_prefix_tokens = 0 if getattr(model, 'no_embed_class', False) else getattr(model, 'num_prefix_tokens', 1)\n            v = resample_abs_pos_embed(\n                v,\n                new_size=model.patch_embed.grid_size,\n                num_prefix_tokens=num_prefix_tokens,\n                interpolation=interpolation,\n                antialias=antialias,\n                verbose=True,\n            )\n\n        k = k.replace('mlp.ffn_ln', 'mlp.norm')\n        k = k.replace('attn.inner_attn_ln', 'attn.norm')\n        k = k.replace('mlp.w12', 'mlp.fc1')\n        k = k.replace('mlp.w1', 'mlp.fc1_g')\n        k = k.replace('mlp.w2', 'mlp.fc1_x')\n        k = k.replace('mlp.w3', 'mlp.fc2')\n        if no_qkv:\n            k = k.replace('q_bias', 'q_proj.bias')\n            k = k.replace('v_bias', 'v_proj.bias')\n\n        if mim_weights and k in ('mask_token', 'lm_head.weight', 'lm_head.bias', 'norm.weight', 'norm.bias'):\n            if k == 'norm.weight' or k == 'norm.bias':\n                # try moving norm -> fc norm on fine-tune, probably a better starting point than new init\n                k = k.replace('norm', 'fc_norm')\n            else:\n                # skip pretrain mask token & head weights\n                continue\n\n        out_dict[k] = v\n\n    return out_dict\n\nclass EVA_X(Eva):\n    def __init__(self, **kwargs):\n        super(EVA_X, self).__init__(**kwargs)\n\n    def forward_features(self, x):\n        x = self.patch_embed(x)\n        x, rot_pos_embed = self._pos_embed(x)\n        for blk in self.blocks:\n            x = blk(x, rope=rot_pos_embed)\n        x = self.norm(x)\n        return x\n\n    def forward_head(self, x, pre_logits: bool = False):\n        if self.global_pool:\n            x = x[:, self.num_prefix_tokens:].mean(dim=1) if self.global_pool == 'avg' else x[:, 0]\n        x = self.fc_norm(x)\n        x = self.head_drop(x)\n        return x if pre_logits else self.head(x)\n\n    def forward(self, x):\n        x = self.forward_features(x)\n        x = self.forward_head(x)\n        return x\n\ndef eva_x_base_patch16(pretrained=False):\n    model = EVA_X(\n        img_size=224,\n        patch_size=16,\n        embed_dim=768,\n        depth=12,\n        num_heads=12,\n        qkv_fused=False,\n        mlp_ratio=4 * 2 / 3,\n        swiglu_mlp=True,\n        scale_mlp=True,\n        use_rot_pos_emb=True,\n        ref_feat_shape=(14, 14),  # 224/16\n    )\n    in_features = model.head.in_features\n    model.head = nn.Linear(in_features, 14)\n    if isinstance(pretrained, str):\n        print(f\"Loading pretrained weights from: {pretrained}\")\n        eva_ckpt = checkpoint_filter_fn(torch.load(pretrained, map_location='cpu', weights_only=False), model)\n        msg = model.load_state_dict(eva_ckpt, strict=False)\n        print(msg)\n    else:\n        print(\"No pretrained weights loaded.\")\n    return model\n\ndef build_model(NUM_LABELS, pretrained):\n    model = eva_x_base_patch16(pretrained)\n    return model","metadata":{"_uuid":"56cbce88-1052-43cd-bae1-8717665684b8","_cell_guid":"daf81758-e6ad-439d-ae56-ee6bdc38e582","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-11-02T17:41:26.723053Z","iopub.execute_input":"2025-11-02T17:41:26.723409Z","iopub.status.idle":"2025-11-02T17:41:26.733191Z","shell.execute_reply.started":"2025-11-02T17:41:26.723382Z","shell.execute_reply":"2025-11-02T17:41:26.732283Z"},"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-11-02T17:39:22.453427Z","iopub.execute_input":"2025-11-02T17:39:22.453599Z","iopub.status.idle":"2025-11-02T17:39:22.457799Z","shell.execute_reply.started":"2025-11-02T17:39:22.453584Z","shell.execute_reply":"2025-11-02T17:39:22.45692Z"},"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-11-02T17:39:22.458262Z","iopub.execute_input":"2025-11-02T17:39:22.45843Z","iopub.status.idle":"2025-11-02T17:39:22.47591Z","shell.execute_reply.started":"2025-11-02T17:39:22.458417Z","shell.execute_reply":"2025-11-02T17:39:22.475153Z"},"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-11-02T17:39:22.476505Z","iopub.execute_input":"2025-11-02T17:39:22.476654Z","iopub.status.idle":"2025-11-02T17:39:22.486674Z","shell.execute_reply.started":"2025-11-02T17:39:22.476641Z","shell.execute_reply":"2025-11-02T17:39:22.485941Z"},"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-11-02T17:39:22.487384Z","iopub.execute_input":"2025-11-02T17:39:22.48758Z","iopub.status.idle":"2025-11-02T17:39:33.732259Z","shell.execute_reply.started":"2025-11-02T17:39:22.487565Z","shell.execute_reply":"2025-11-02T17:39:33.731086Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_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\")","metadata":{"_uuid":"c2fb2885-d46e-45bc-87f0-ca953169f148","_cell_guid":"0e511017-2843-41bf-95b5-1fac7444dfaa","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-11-02T17:48:05.906982Z","iopub.execute_input":"2025-11-02T17:48:05.907247Z","iopub.status.idle":"2025-11-02T17:48:06.030888Z","shell.execute_reply.started":"2025-11-02T17:48:05.907219Z","shell.execute_reply":"2025-11-02T17:48:06.029896Z"},"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)","metadata":{"_uuid":"1e42e95a-24ed-471d-b9f0-e9c2a45a0693","_cell_guid":"40912715-5bae-431a-b852-37f8582a1b1f","trusted":true,"collapsed":false,"execution":{"iopub.status.busy":"2025-11-02T17:53:00.046143Z","iopub.execute_input":"2025-11-02T17:53:00.046424Z","iopub.status.idle":"2025-11-02T17:53:00.054083Z","shell.execute_reply.started":"2025-11-02T17:53:00.046382Z","shell.execute_reply":"2025-11-02T17:53:00.053279Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Model + transforms\nmodel = build_model(NUM_LABELS, cfg.pretrained)\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-11-02T17:41:30.11446Z","iopub.execute_input":"2025-11-02T17:41:30.114713Z","iopub.status.idle":"2025-11-02T17:41:31.38574Z","shell.execute_reply.started":"2025-11-02T17:41:30.114698Z","shell.execute_reply":"2025-11-02T17:41:31.384857Z"},"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_tf, val_tf = build_transforms(cfg.img_size)\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-11-02T17:59:36.990128Z","iopub.execute_input":"2025-11-02T17:59:36.990383Z","iopub.status.idle":"2025-11-02T17:59:37.017579Z","shell.execute_reply.started":"2025-11-02T17:59:36.990365Z","shell.execute_reply":"2025-11-02T17:59:37.016612Z"},"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))\n# if cfg.use_focal:\n#     criterion = FocalLoss(alpha=alpha_vec.to(device), gamma=2.0)\n# else:\n#     criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor(pos_weight_vec, device=device))\n\ncriterion = AsymmetricLossOptimized()\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-11-02T17:59:37.163939Z","iopub.execute_input":"2025-11-02T17:59:37.164159Z","iopub.status.idle":"2025-11-02T17:59:37.168366Z","shell.execute_reply.started":"2025-11-02T17:59:37.164141Z","shell.execute_reply":"2025-11-02T17:59:37.16763Z"},"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":true,"execution":{"iopub.status.busy":"2025-11-02T17:59:39.691072Z","iopub.execute_input":"2025-11-02T17:59:39.691309Z","iopub.status.idle":"2025-11-02T17:59:47.226506Z","shell.execute_reply.started":"2025-11-02T17:59:39.691292Z","shell.execute_reply":"2025-11-02T17:59:47.225499Z"},"jupyter":{"outputs_hidden":true}},"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-11-02T17:59:52.861685Z","iopub.execute_input":"2025-11-02T17:59:52.862004Z","execution_failed":"2025-11-02T18:15:46.292Z"},"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-11-02T11:58:46.921675Z","iopub.status.idle":"2025-11-02T11:58:46.922124Z","shell.execute_reply.started":"2025-11-02T11:58:46.921763Z","shell.execute_reply":"2025-11-02T11:58:46.921771Z"},"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-11-02T11:58:46.922478Z","iopub.status.idle":"2025-11-02T11:58:46.922927Z","shell.execute_reply.started":"2025-11-02T11:58:46.922562Z","shell.execute_reply":"2025-11-02T11:58:46.922569Z"},"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-11-02T11:58:46.923535Z","iopub.status.idle":"2025-11-02T11:58:46.924013Z","shell.execute_reply.started":"2025-11-02T11:58:46.923622Z","shell.execute_reply":"2025-11-02T11:58:46.92363Z"},"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}]}