{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":112899,"databundleVersionId":13449579,"sourceType":"competition"},{"sourceId":13202187,"sourceType":"datasetVersion","datasetId":8366986},{"sourceId":13124570,"sourceType":"datasetVersion","datasetId":8314063}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ===============================\n# Cell 1 — Setup / Dataset / Utils (with light train-time aug)\n# ===============================\nimport os, cv2, math, random, numpy as np, pandas as pd\nfrom typing import List\nimport torch, torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nfrom torchvision.transforms import InterpolationMode\nfrom sklearn.metrics import roc_auc_score\nfrom tqdm import tqdm\nfrom PIL import Image\n\nimport timm\n\n# TPU (PyTorch/XLA)\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\n\n# -------------------------\n# Config\n# -------------------------\nDATA_DIR   = \"/kaggle/input/grand-xray-slam-division-a\"\nTRAIN_CSV  = f\"{DATA_DIR}/train1.csv\"\nTRAIN_DIR  = f\"{DATA_DIR}/train1\"\nTEST_DIR   = f\"{DATA_DIR}/test1\"\nOUT_CSV    = \"/kaggle/working/submission.csv\"\n\nLABEL_COLUMNS: List[str] = [\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]\n\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD  = [0.229, 0.224, 0.225]\nIMG_SIZE = (380, 380)  # EfficientNet-B4 input\nTO_DROP = set([\"00025979_008_001.jpg\", \"00048043_001_002.jpg\"])\n\n# -------------------------\n# Transforms\n# -------------------------\ndef get_train_transforms(img_size=(380, 380)):\n    H, W = img_size\n    return T.Compose([\n        # 안전한 의료형 약증강\n        T.RandomResizedCrop(size=(H, W), scale=(0.85, 1.0), ratio=(0.95, 1.05), interpolation=InterpolationMode.BILINEAR),\n        T.RandomHorizontalFlip(p=0.5),\n        T.RandomRotation(degrees=5, interpolation=InterpolationMode.BILINEAR, fill=0),\n        T.RandomAutocontrast(p=0.2),\n        T.RandomApply([T.GaussianBlur(kernel_size=3)], p=0.1),\n        T.ToTensor(),\n        T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ])\n\ndef get_test_transforms(img_size=(380, 380)):\n    H, W = img_size\n    return T.Compose([\n        T.Resize(size=(H, W), interpolation=InterpolationMode.BILINEAR),\n        T.ToTensor(),\n        T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n    ])\n\n# -------------------------\n# Dataset (now with train-time aug)\n# -------------------------\nclass XRayDataset(Dataset):\n    def __init__(self, df: pd.DataFrame, image_dir: str, img_size=(380, 380), is_train=True):\n        self.df = df.reset_index(drop=True)\n        self.image_dir = image_dir\n        self.img_size = img_size\n        self.is_train = is_train\n        self.tfms = get_train_transforms(img_size) if is_train else get_test_transforms(img_size)\n\n    def __len__(self):\n        return len(self.df)\n\n    def _imread_rgb(self, path):\n        # 원본 해상도 유지(증강이 크기/크롭/회전을 처리)\n        img = cv2.imread(path, cv2.IMREAD_COLOR)\n        if img is None:\n            # 비정상 파일은 검은 이미지로 대체\n            H, W = self.img_size\n            img = np.zeros((H, W, 3), dtype=np.uint8)\n        else:\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        return img\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = os.path.join(self.image_dir, row[\"Image_name\"])\n        img = self._imread_rgb(img_path)\n\n        # PIL로 변환 후 torchvision 증강 수행\n        img = Image.fromarray(img)\n        img = self.tfms(img)  # Tensor [C,H,W]\n\n        if \"split\" in self.df.columns and not self.is_train:\n            return row[\"Image_name\"], img\n\n        labels = torch.tensor(row[LABEL_COLUMNS].values.astype(np.float32))\n        return img, labels\n\nclass TestDataset(Dataset):\n    def __init__(self, image_dir: str, img_size=(380, 380)):\n        self.image_dir = image_dir\n        self.img_size = img_size\n        self.images = sorted([f for f in os.listdir(image_dir) if f.lower().endswith(\".jpg\")])\n        self.tfms = get_test_transforms(img_size)\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        name = self.images[idx]\n        p = os.path.join(self.image_dir, name)\n        img = cv2.imread(p, cv2.IMREAD_COLOR)\n        if img is None:\n            H, W = self.img_size\n            img = np.zeros((H, W, 3), dtype=np.uint8)\n            img = Image.fromarray(img)\n        else:\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n            img = Image.fromarray(img)\n        img = self.tfms(img)\n        return name, img\n\n# (Optional) FocalLoss kept (unused)\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=None, gamma=2.0, reduction='mean'):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n    def forward(self, logits, targets):\n        bce = nn.functional.binary_cross_entropy_with_logits(logits, targets, reduction='none')\n        pt = torch.exp(-bce)\n        loss = (1 - pt) ** self.gamma * bce\n        if self.alpha is not None:\n            loss = loss * self.alpha.to(logits.device)\n        if self.reduction == 'mean':\n            return loss.mean()\n        elif self.reduction == 'sum':\n            return loss.sum()\n        return loss\n\n@torch.no_grad()\ndef compute_macro_auc(y_true: np.ndarray, y_pred: np.ndarray) -> float:\n    try:\n        return roc_auc_score(y_true, y_pred, average=\"macro\")\n    except ValueError:\n        return float(\"nan\")\n\nprint(\"Cell 1 ready ✔ (with light train-time aug)\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===============================\n# Cell 2 — Ensemble Inference (EMA B4 + EMA ViT), TTA hflip x2 (bf16→fp32 fix)\n# ===============================\nimport os, contextlib\nimport torch\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nimport timm\nimport pandas as pd\nimport numpy as np\n\n# TPU (PyTorch/XLA)\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\nfrom torch_xla.amp import autocast as xla_autocast\n\n# ----- 체크포인트 경로 -----\nB4_EMA_PATH  = \"/kaggle/input/929591-pth/ema_b4_final.pth\"\nVIT_EMA_PATH = \"/kaggle/input/vit-data/ema_vitb16_384.pth\"\nassert os.path.exists(B4_EMA_PATH),  f\"Not found: {B4_EMA_PATH}\"\nassert os.path.exists(VIT_EMA_PATH), f\"Not found: {VIT_EMA_PATH}\"\n\n# ----- Device / Loader -----\ndevice = xm.xla_device()\nbatch_size  = 24\nnum_workers = 8\n\n# 셀1의 TestDataset 재사용\ntest_ds = TestDataset(TEST_DIR, img_size=(384, 384))\ntest_loader_cpu = DataLoader(test_ds, batch_size=batch_size, shuffle=False,\n                             num_workers=num_workers, pin_memory=True)\ntest_loader = pl.MpDeviceLoader(test_loader_cpu, device)\n\nxm.master_print(\"Using TPU device:\", device)\n\n# ----- Model build & load -----\n# EfficientNet-B4 (380)\nmodel_b4 = timm.create_model(\n    'tf_efficientnet_b4_ns', pretrained=False,\n    num_classes=len(LABEL_COLUMNS), drop_rate=0.0\n).to(device)\nstate_b4 = torch.load(B4_EMA_PATH, map_location='cpu')\nmodel_b4.load_state_dict(state_b4, strict=False)\nmodel_b4.eval()\n\n# ViT-B/16 (384)\nmodel_vit = timm.create_model(\n    'vit_base_patch16_384', pretrained=False,\n    num_classes=len(LABEL_COLUMNS)\n).to(device)\nstate_vit = torch.load(VIT_EMA_PATH, map_location='cpu')\nmodel_vit.load_state_dict(state_vit, strict=False)\nmodel_vit.eval()\n\n# ----- Helpers -----\ndef resize_batch(x, size_hw):\n    if list(x.shape[-2:]) == list(size_hw):\n        return x\n    return F.interpolate(x, size=size_hw, mode='bilinear', align_corners=False)\n\n@torch.no_grad()\ndef predict_probs_tta(model, imgs, in_size, use_amp=True):\n    \"\"\"\n    Equal-weight TTA = original + hflip (2-pass). Returns float32 probs.\n    \"\"\"\n    preds_accum = 0.0\n    cnt = 0\n    x1 = resize_batch(imgs, in_size)       # original\n    x2 = torch.flip(x1, dims=[3])          # hflip\n\n    for xx in (x1, x2):\n        with xla_autocast(device=device, dtype=torch.bfloat16) if use_amp else contextlib.nullcontext():\n            logits = model(xx)\n            probs = torch.sigmoid(logits)\n        preds_accum += probs\n        cnt += 1\n\n    # ✅ numpy로 보낼 때 bf16 이슈 방지: float32로 캐스팅\n    return (preds_accum / cnt).to(torch.float32)\n\n# ----- Inference: (B4 + ViT)/2, TTA flip x2 -----\nsubmission_rows = []\nuse_amp = True\n\nwith torch.no_grad():\n    pbar = tqdm(test_loader, desc=\"[Ensemble Infer] (B4 + ViT)/2, TTA hflip x2\", unit=\"batch\")\n    for names, imgs in pbar:\n        imgs = imgs.to(device, non_blocking=True)\n\n        probs_b4  = predict_probs_tta(model_b4,  imgs, (380, 380), use_amp=use_amp)\n        probs_vit = predict_probs_tta(model_vit, imgs, (384, 384), use_amp=use_amp)\n\n        probs = 0.5 * (probs_b4 + probs_vit)                # equal-weight\n        probs = probs.cpu().numpy()                          # now safe (fp32)\n\n        for n, p in zip(names, probs):\n            submission_rows.append([n] + p.tolist())\n        xm.mark_step()\n\nsub_df = pd.DataFrame(submission_rows, columns=[\"Image_name\"] + LABEL_COLUMNS)\nsub_df.to_csv(OUT_CSV, index=False)\nxm.master_print(f\"✅ Saved submission: {OUT_CSV}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}