{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":14774,"databundleVersionId":875431,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-03T18:04:40.093712Z","iopub.execute_input":"2025-12-03T18:04:40.09391Z","iopub.status.idle":"2025-12-03T18:04:41.17029Z","shell.execute_reply.started":"2025-12-03T18:04:40.093894Z","shell.execute_reply":"2025-12-03T18:04:41.169549Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================\n# CELL 1 — Config, imports, speed flags, cache setup\n# ================================\nimport os, gc, math, random, warnings, time\nwarnings.filterwarnings(\"ignore\")\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import cohen_kappa_score, f1_score, accuracy_score\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport torchvision\nfrom torchvision import models\n\n# ---- Paths (APTOS 2019 competition dataset) ----\nCOMP_DIR  = \"/kaggle/input/aptos2019-blindness-detection\"\nTRAIN_DIR = f\"{COMP_DIR}/train_images\"\nTEST_DIR  = f\"{COMP_DIR}/test_images\"\nCSV_PATH  = f\"{COMP_DIR}/train.csv\"\n\n# ---- Repro ----\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\n\n# ---- Device ----\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"CUDA Available:\", torch.cuda.is_available())\nif torch.cuda.is_available():\n    print(\"GPU:\", torch.cuda.get_device_name(0))\n\n# ---- Paper-compliant core + speed-optimized training ----\nIMG_SIZE_TRAIN = 512              # training/inference size\nDO_FINETUNE768 = False            # skipping progressive finetune for speed\nUSE_EMA = False                   # disabled for speed\nUSE_TTA = True\nTTA_N = 5\n\n# Preprocessing selection (your Option 3): ROI + Bilateral + CLAHE + Normalize\nUSE_BILATERAL = True\nUSE_CLAHE     = True\n\n# Training knobs (speed profile)\nN_FOLDS    = 3\nEPOCHS_512 = 6\nBATCH_SIZE = 16\nLR = 3e-4\nWEIGHT_DECAY = 1e-4\nLABEL_SMOOTH = 0.05\n\n# Dataloader workers\nNUM_WORKERS = 2  # safer on Kaggle; we’ll keep workers light since we cache\n\n# ---- Speed flags ----\nimport torch.backends.cudnn as cudnn\ntorch.set_float32_matmul_precision(\"high\")\ntorch.backends.cuda.matmul.allow_tf32 = True\ncudnn.allow_tf32 = True\ncudnn.benchmark  = True\nCHANNELS_LAST    = True\n\n# ---- Disk cache (train + test) ----\nCACHE_SIZE = 512   # you chose 512px cache\nCACHE_DIR  = \"/kaggle/temp/aptos_cache_512\"\nCACHE_TRAIN_DIR = os.path.join(CACHE_DIR, \"train\")\nCACHE_TEST_DIR  = os.path.join(CACHE_DIR, \"test\")\nos.makedirs(CACHE_TRAIN_DIR, exist_ok=True)\nos.makedirs(CACHE_TEST_DIR,  exist_ok=True)\n\nos.makedirs(\"checkpoints\", exist_ok=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T17:15:33.272154Z","iopub.execute_input":"2025-12-03T17:15:33.272391Z","iopub.status.idle":"2025-12-03T17:15:33.283361Z","shell.execute_reply.started":"2025-12-03T17:15:33.272377Z","shell.execute_reply":"2025-12-03T17:15:33.282795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================\n# CELL 2 — Preprocessing + cache utilities\n# ================================\ndef crop_center_retina(img):\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    _, th = cv2.threshold(gray, 10, 255, cv2.THRESH_BINARY)\n    contours, _ = cv2.findContours(th, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    if not contours:\n        return img\n    c = max(contours, key=cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(c)\n    h0, w0 = img.shape[:2]\n    x1, y1 = max(0, x), max(0, y)\n    x2, y2 = min(w0, x + w), min(h0, y + h)\n    crop = img[y1:y2, x1:x2]\n    return crop if crop.size else img\n\ndef noise_reduce(img):\n    if USE_BILATERAL:\n        return cv2.bilateralFilter(img, d=9, sigmaColor=75, sigmaSpace=75)\n    return img\n\ndef apply_clahe(img):\n    if not USE_CLAHE:\n        return img\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    l, a, b = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    l2 = clahe.apply(l)\n    img2 = cv2.cvtColor(cv2.merge([l2, a, b]), cv2.COLOR_LAB2BGR)\n    return img2\n\ndef preprocess_core(img_bgr):\n    # Paper pipeline: ROI -> bilateral -> CLAHE -> RGB -> resize to CACHE_SIZE\n    img = crop_center_retina(img_bgr)\n    img = noise_reduce(img)\n    img = apply_clahe(img)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    # resize long side to CACHE_SIZE while keeping aspect, then pad to square\n    h, w = img.shape[:2]\n    scale = CACHE_SIZE / max(h, w)\n    if scale != 1.0:\n        img = cv2.resize(img, (int(w*scale), int(h*scale)), interpolation=cv2.INTER_AREA)\n    h2, w2 = img.shape[:2]\n    top = (CACHE_SIZE - h2) // 2\n    bottom = CACHE_SIZE - h2 - top\n    left = (CACHE_SIZE - w2) // 2\n    right = CACHE_SIZE - w2 - left\n    img = cv2.copyMakeBorder(img, top, bottom, left, right, cv2.BORDER_CONSTANT, value=(0,0,0))\n    return img\n\ndef cache_write_rgb(cache_path, img_rgb):\n    # Save as PNG with light compression\n    img_bgr = cv2.cvtColor(img_rgb, cv2.COLOR_RGB2BGR)\n    cv2.imwrite(cache_path, img_bgr, [cv2.IMWRITE_PNG_COMPRESSION, 3])\n\ndef cache_read_rgb(cache_path):\n    img_bgr = cv2.imread(cache_path, cv2.IMREAD_COLOR)\n    if img_bgr is None:\n        raise FileNotFoundError(cache_path)\n    return cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n\ndef raw_image_path(img_dir, img_id):\n    p = f\"{img_dir}/{img_id}.png\"\n    if not os.path.exists(p):\n        p = f\"{img_dir}/{img_id}.jpg\"\n    return p\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T17:15:43.032425Z","iopub.execute_input":"2025-12-03T17:15:43.032906Z","iopub.status.idle":"2025-12-03T17:15:43.043312Z","shell.execute_reply.started":"2025-12-03T17:15:43.032885Z","shell.execute_reply":"2025-12-03T17:15:43.042508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================\n# CELL 3 — Cached Dataset (train + test) with balanced fast augs\n# ================================\nclass DRDataset(Dataset):\n    def __init__(self, df, img_dir, size=IMG_SIZE_TRAIN, mode=\"train\"):\n        self.df = df.reset_index(drop=True)\n        self.img_dir = img_dir\n        self.size = size\n        self.mode = mode\n\n        # choose cache subdir by source\n        self.cache_dir = CACHE_TRAIN_DIR if os.path.samefile(img_dir, TRAIN_DIR) else CACHE_TEST_DIR\n\n        if mode == \"train\":\n            # Balanced + fast (keeps accuracy, reduces CPU cost)\n            self.tf = A.Compose([\n                A.RandomResizedCrop(size=(size, size), scale=(0.92, 1.0), ratio=(0.95, 1.05)),\n                A.HorizontalFlip(p=0.5),\n                A.ShiftScaleRotate(shift_limit=0.02, scale_limit=0.06, rotate_limit=8, p=0.4),\n                A.RandomBrightnessContrast(p=0.35, brightness_limit=0.12, contrast_limit=0.12),\n                A.Normalize(),\n                ToTensorV2(),\n            ], p=1.0)\n        else:\n            self.tf = A.Compose([\n                A.LongestMaxSize(max_size=size),\n                A.PadIfNeeded(min_height=size, min_width=size, border_mode=cv2.BORDER_CONSTANT, value=0),\n                A.Normalize(),\n                ToTensorV2(),\n            ], p=1.0)\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        r = self.df.iloc[idx]\n        img_id = r[\"id_code\"]\n\n        cache_path = os.path.join(self.cache_dir, f\"{img_id}.png\")\n        if os.path.exists(cache_path):\n            img = cache_read_rgb(cache_path)\n        else:\n            raw_path = raw_image_path(self.img_dir, img_id)\n            img_bgr = cv2.imread(raw_path, cv2.IMREAD_COLOR)\n            if img_bgr is None:\n                raise FileNotFoundError(raw_path)\n            img = preprocess_core(img_bgr)  # ROI + bilateral + CLAHE + RGB + pad/resize to CACHE_SIZE\n            cache_write_rgb(cache_path, img)\n\n        img = self.tf(image=img)[\"image\"]\n\n        if \"diagnosis\" in r:\n            return img, int(r[\"diagnosis\"])\n        else:\n            return img, img_id\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T17:15:49.721968Z","iopub.execute_input":"2025-12-03T17:15:49.722268Z","iopub.status.idle":"2025-12-03T17:15:49.730672Z","shell.execute_reply.started":"2025-12-03T17:15:49.722249Z","shell.execute_reply":"2025-12-03T17:15:49.729847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================\n# CELL 4 — CBAM + ResNet-50 (torchvision) and model builder\n# ================================\nclass ChannelAttention(nn.Module):\n    def __init__(self, in_planes, ratio=16):\n        super().__init__()\n        self.avg = nn.AdaptiveAvgPool2d(1)\n        self.max = nn.AdaptiveMaxPool2d(1)\n        self.fc1 = nn.Conv2d(in_planes, in_planes // ratio, 1, bias=False)\n        self.relu = nn.ReLU(inplace=True)\n        self.fc2 = nn.Conv2d(in_planes // ratio, in_planes, 1, bias=False)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg_out = self.fc2(self.relu(self.fc1(self.avg(x))))\n        max_out = self.fc2(self.relu(self.fc1(self.max(x))))\n        return self.sigmoid(avg_out + max_out)\n\nclass SpatialAttention(nn.Module):\n    def __init__(self, kernel_size=7):\n        super().__init__()\n        self.conv = nn.Conv2d(2, 1, kernel_size, padding=(kernel_size-1)//2, bias=False)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg_out = torch.mean(x, dim=1, keepdim=True)\n        max_out, _ = torch.max(x, dim=1, keepdim=True)\n        x = torch.cat([avg_out, max_out], dim=1)\n        x = self.conv(x)\n        return self.sigmoid(x)\n\nclass CBAM(nn.Module):\n    def __init__(self, channels, ratio=16, k=7):\n        super().__init__()\n        self.ca = ChannelAttention(channels, ratio)\n        self.sa = SpatialAttention(k)\n    def forward(self, x):\n        x = x * self.ca(x)\n        x = x * self.sa(x)\n        return x\n\nclass ResNet50_CBAM_TV(nn.Module):\n    \"\"\"Torchvision ResNet-50 with CBAM after layer3 and layer4; 5-class head.\"\"\"\n    def __init__(self, num_classes=5, try_pretrained=True):\n        super().__init__()\n        weights = None\n        if try_pretrained:\n            try:\n                weights = models.ResNet50_Weights.IMAGENET1K_V1\n                backbone = models.resnet50(weights=weights)\n            except Exception:\n                backbone = models.resnet50(weights=None)\n        else:\n            backbone = models.resnet50(weights=None)\n\n        self.stem = nn.Sequential(backbone.conv1, backbone.bn1, backbone.relu, backbone.maxpool)\n        self.layer1 = backbone.layer1\n        self.layer2 = backbone.layer2\n        self.layer3 = backbone.layer3\n        self.layer4 = backbone.layer4\n\n        self.cbam3 = CBAM(1024)\n        self.cbam4 = CBAM(2048)\n\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.drop = nn.Dropout(0.2)\n        self.fc   = nn.Linear(2048, num_classes)\n\n    def forward(self, x):\n        x = self.stem(x)\n        x = self.layer1(x); x = self.layer2(x)\n        x = self.layer3(x); x = self.cbam3(x)\n        x = self.layer4(x); x = self.cbam4(x)\n        x = self.pool(x).flatten(1)\n        x = self.drop(x)\n        x = self.fc(x)\n        return x\n\ndef build_model(num_classes=5, try_pretrained=True):\n    m = ResNet50_CBAM_TV(num_classes=num_classes, try_pretrained=try_pretrained).to(DEVICE)\n    if CHANNELS_LAST:\n        m = m.to(memory_format=torch.channels_last)\n    try:\n        m = torch.compile(m, mode=\"max-autotune\")\n    except Exception:\n        pass\n    return m\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T17:15:55.752001Z","iopub.execute_input":"2025-12-03T17:15:55.752275Z","iopub.status.idle":"2025-12-03T17:15:55.764576Z","shell.execute_reply.started":"2025-12-03T17:15:55.752256Z","shell.execute_reply":"2025-12-03T17:15:55.763996Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================\n# CELL 5 — Loss, scheduler, metrics, fast helpers\n# ================================\nclass LabelSmoothingCE(nn.Module):\n    def __init__(self, eps=0.05, weight=None):\n        super().__init__()\n        self.eps = eps\n        self.weight = weight\n    def forward(self, logits, target):\n        n = logits.size(-1)\n        log_p = F.log_softmax(logits, dim=-1)\n        with torch.no_grad():\n            true = torch.zeros_like(log_p)\n            true.fill_(self.eps/(n-1))\n            true.scatter_(1, target.unsqueeze(1), 1 - self.eps)\n        loss = (-true * log_p)\n        if self.weight is not None:\n            w = self.weight[target].unsqueeze(1)\n            loss = loss * w\n        return loss.sum(dim=1).mean()\n\ndef cosine_with_warmup(optimizer, steps, warmup):\n    def f(step):\n        if step < warmup:\n            return float(step) / float(max(1, warmup))\n        prog = float(step - warmup) / float(max(1, steps - warmup))\n        return 0.5 * (1.0 + math.cos(math.pi * prog))\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, f)\n\ndef compute_scores(y_true, y_pred):\n    acc = accuracy_score(y_true, y_pred)\n    qwk = cohen_kappa_score(y_true, y_pred, weights=\"quadratic\")\n    f1m = f1_score(y_true, y_pred, average=\"macro\")\n    return acc, qwk, f1m\n\ndef _to_device_fast(x):\n    x = x.to(DEVICE, non_blocking=True)\n    if CHANNELS_LAST:\n        x = x.contiguous(memory_format=torch.channels_last)\n    return x\n\ndef train_one_epoch(model, loader, optimizer, criterion, scaler):\n    model.train()\n    total = 0.0\n    for imgs, targets in loader:\n        imgs = _to_device_fast(imgs)\n        targets = targets.to(DEVICE, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n        with autocast():\n            logits = model(imgs)\n            loss = criterion(logits, targets)\n        scaler.scale(loss).backward()\n        scaler.step(optimizer); scaler.update()\n\n        total += loss.item() * imgs.size(0)\n    return total / len(loader.dataset)\n\n@torch.no_grad()\ndef evaluate(model, loader):\n    model.eval()\n    preds, targs = [], []\n    for imgs, targets in loader:\n        imgs = _to_device_fast(imgs)\n        logits = model(imgs)\n        pred = logits.softmax(1).argmax(1).cpu().numpy()\n        preds.extend(pred); targs.extend(targets.numpy())\n    return compute_scores(np.array(targs), np.array(preds))\n\n@torch.no_grad()\ndef predict_tta(model, loader, tta_n=5):\n    model.eval()\n    preds = []\n    for imgs, _ in loader:\n        imgs = _to_device_fast(imgs)\n        logits_sum = torch.zeros((imgs.size(0), 5), device=DEVICE)\n        for _ in range(tta_n):\n            aug = imgs.flip(-1) if random.random() < 0.5 else imgs\n            logits_sum += model(aug)\n        pred = (logits_sum/tta_n).softmax(1).argmax(1).cpu().numpy()\n        preds.extend(pred)\n    return preds\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T17:16:02.601872Z","iopub.execute_input":"2025-12-03T17:16:02.602622Z","iopub.status.idle":"2025-12-03T17:16:02.614673Z","shell.execute_reply.started":"2025-12-03T17:16:02.602598Z","shell.execute_reply":"2025-12-03T17:16:02.613807Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================\n# CELL 6 — Load CSV + WeightedRandomSampler for Class Imbalance\n# ================================\nfrom torch.utils.data import WeightedRandomSampler\n\ndf = pd.read_csv(CSV_PATH)  # columns: id_code, diagnosis\n\n# Class counts and weights\ncounts = df[\"diagnosis\"].value_counts().sort_index()\nclass_weights = 1.0 / counts\nsample_weights = df[\"diagnosis\"].map(class_weights).values\n\nprint(\"Dataset shape:\", df.shape)\nprint(\"Class counts:\\n\", counts)\nprint(\"\\nClass weights used for sampling:\", class_weights.to_dict())\n\n# We will use WeightedRandomSampler ONLY for training folds in Cell 7\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T17:16:09.111748Z","iopub.execute_input":"2025-12-03T17:16:09.112272Z","iopub.status.idle":"2025-12-03T17:16:09.148485Z","shell.execute_reply.started":"2025-12-03T17:16:09.112247Z","shell.execute_reply":"2025-12-03T17:16:09.147851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================\n# CELL 7 — Balanced Training with WeightedRandomSampler + Cache + TTA OOF\n# ================================\nprint(f\"[FAST/CACHED + BALANCED] folds={N_FOLDS}, epochs={EPOCHS_512}, bs={BATCH_SIZE}, size={IMG_SIZE_TRAIN}, tta={USE_TTA}\")\n\nskf = StratifiedKFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\noof_preds = np.zeros(len(df), dtype=int)\noof_targs = df[\"diagnosis\"].values\nfold_scores = []\n\nfor fold, (tr_idx, va_idx) in enumerate(skf.split(df[\"id_code\"], df[\"diagnosis\"])):\n    print(f\"\\n================ FOLD {fold+1}/{N_FOLDS} ================\")\n    t_fold = time.time()\n\n    df_tr = df.iloc[tr_idx].copy()\n    df_va = df.iloc[va_idx].copy()\n\n    ds_tr = DRDataset(df_tr, TRAIN_DIR, size=IMG_SIZE_TRAIN, mode=\"train\")\n    ds_va = DRDataset(df_va, TRAIN_DIR, size=IMG_SIZE_TRAIN, mode=\"valid\")\n\n    # -------- WeightedRandomSampler for Balanced Training --------\n    train_sampler = WeightedRandomSampler(\n        weights=sample_weights[tr_idx],\n        num_samples=len(tr_idx),\n        replacement=True\n    )\n\n    try:\n        dl_tr = DataLoader(\n            ds_tr,\n            batch_size=BATCH_SIZE,\n            sampler=train_sampler,\n            num_workers=2,\n            pin_memory=True,\n            persistent_workers=False,\n            prefetch_factor=2,\n            drop_last=True\n        )\n        dl_va = DataLoader(\n            ds_va,\n            batch_size=BATCH_SIZE,\n            shuffle=False,\n            num_workers=2,\n            pin_memory=True,\n            persistent_workers=False,\n            prefetch_factor=2\n        )\n    except TypeError:\n        dl_tr = DataLoader(ds_tr, batch_size=BATCH_SIZE, sampler=train_sampler, num_workers=2, pin_memory=True, drop_last=True)\n        dl_va = DataLoader(ds_va, batch_size=BATCH_SIZE, shuffle=False, num_workers=2, pin_memory=True)\n\n    model = build_model(num_classes=5, try_pretrained=True)\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n    sched_steps = EPOCHS_512 * len(dl_tr)\n    scheduler   = cosine_with_warmup(optimizer, steps=sched_steps, warmup=int(0.1 * sched_steps))\n    scaler      = GradScaler()\n    criterion   = LabelSmoothingCE(eps=LABEL_SMOOTH, weight=None)  # No need for class_weights now\n\n    best_qwk = -1.0\n    best_path = f\"checkpoints/fold{fold}_best.pth\"\n\n    # ---- Train @512 only (fast + cached + balanced) ----\n    for ep in range(EPOCHS_512):\n        t_ep = time.time()\n        tr_loss = train_one_epoch(model, dl_tr, optimizer, criterion, scaler)\n        scheduler.step()\n        acc, qwk, f1m = evaluate(model, dl_va)\n\n        print(f\"Ep {ep+1:02d}/{EPOCHS_512} | {time.time()-t_ep:5.1f}s | loss {tr_loss:.4f} | acc {acc:.4f} | qwk {qwk:.4f} | f1 {f1m:.4f}\")\n\n        if qwk > best_qwk:\n            best_qwk = qwk\n            torch.save(model.state_dict(), best_path)\n\n    # ---- OOF with TTA ----\n    model.load_state_dict(torch.load(best_path, map_location=DEVICE))\n    dl_va_eval = DataLoader(\n        DRDataset(df_va, TRAIN_DIR, size=IMG_SIZE_TRAIN, mode=\"valid\"),\n        batch_size=BATCH_SIZE, shuffle=False, num_workers=2, pin_memory=True\n    )\n\n    preds = predict_tta(model, dl_va_eval, tta_n=TTA_N) if USE_TTA else None\n    if preds is None:\n        preds = []\n        model.eval()\n        for imgs, _ in dl_va_eval:\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            if CHANNELS_LAST:\n                imgs = imgs.contiguous(memory_format=torch.channels_last)\n            logits = model(imgs)\n            preds.extend(logits.softmax(1).argmax(1).cpu().numpy())\n\n    oof_preds[va_idx] = np.array(preds)\n    acc, qwk, f1m = compute_scores(df_va[\"diagnosis\"].values, oof_preds[va_idx])\n    print(f\"FOLD {fold+1} DONE in {(time.time()-t_fold)/60:.1f} min | acc {acc:.4f} | qwk {qwk:.4f} | f1 {f1m:.4f}\")\n    fold_scores.append((acc, qwk, f1m))\n\nprint(\"\\n================ OOF RESULTS ================\")\nacc = accuracy_score(oof_targs, oof_preds)\nqwk = cohen_kappa_score(oof_targs, oof_preds, weights=\"quadratic\")\nf1m = f1_score(oof_targs, oof_preds, average=\"macro\")\nprint(f\"OOF | acc {acc:.4f} | qwk {qwk:.4f} | f1 {f1m:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T17:16:14.611801Z","iopub.execute_input":"2025-12-03T17:16:14.612391Z","iopub.status.idle":"2025-12-03T18:04:32.917046Z","shell.execute_reply.started":"2025-12-03T17:16:14.612369Z","shell.execute_reply":"2025-12-03T18:04:32.916049Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nprint(\"Folders:\", os.listdir())\nif \"checkpoints\" in os.listdir():\n    print(\"Checkpoint files:\", os.listdir(\"checkpoints\"))\nelse:\n    print(\"No checkpoints folder found.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T18:05:20.446291Z","iopub.execute_input":"2025-12-03T18:05:20.44678Z","iopub.status.idle":"2025-12-03T18:05:20.45161Z","shell.execute_reply.started":"2025-12-03T18:05:20.446756Z","shell.execute_reply":"2025-12-03T18:05:20.450926Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================\n# CELL 8 — Stage-3 Finetune at 768px (Balanced Base)\n# ================================\nprint(\"[FT768-BAL] Starting Stage-3 Finetune at 768px\")\n\nFT_EPOCHS = 3\nIMG_SIZE_FT = 768\nFT_BS = 8\nFT_LR = 1e-4\nos.makedirs(\"checkpoints\", exist_ok=True)\n\n# Store all fold indices FIRST to avoid generator issue\nfold_indices = list(skf.split(df[\"id_code\"], df[\"diagnosis\"]))\n\nfor fold in range(N_FOLDS):\n    print(f\"\\n[FT768-BAL] Fold {fold} Loading balanced checkpoint...\")\n    tr_idx, va_idx = fold_indices[fold]\n\n    base_path = f\"checkpoints/fold{fold}_best.pth\"   # Balanced checkpoint\n    ft_path   = f\"checkpoints/fold{fold}_best_ft768_bal.pth\"  # New filename\n\n    model = build_model(num_classes=5, try_pretrained=False)\n    model.load_state_dict(torch.load(base_path, map_location=DEVICE))\n    model.to(DEVICE)\n\n    # Datasets for finetune\n    ds_tr = DRDataset(df.iloc[tr_idx].copy(), TRAIN_DIR, size=IMG_SIZE_FT, mode=\"train\")\n    ds_va = DRDataset(df.iloc[va_idx].copy(), TRAIN_DIR, size=IMG_SIZE_FT, mode=\"valid\")\n\n    # Weighted sampling reused for fine-tuning\n    train_sampler = WeightedRandomSampler(\n        weights=sample_weights[tr_idx],\n        num_samples=len(tr_idx),\n        replacement=True\n    )\n\n    dl_tr = DataLoader(ds_tr, batch_size=FT_BS, sampler=train_sampler, num_workers=2, pin_memory=True, drop_last=True)\n    dl_va = DataLoader(ds_va, batch_size=FT_BS, shuffle=False, num_workers=2, pin_memory=True)\n\n    optimizer = torch.optim.AdamW(model.parameters(), lr=FT_LR, weight_decay=WEIGHT_DECAY)\n    sched_steps = FT_EPOCHS * len(dl_tr)\n    scheduler = cosine_with_warmup(optimizer, steps=sched_steps, warmup=int(0.15 * sched_steps))\n    scaler = GradScaler()\n    criterion = LabelSmoothingCE(eps=LABEL_SMOOTH)\n\n    best_qwk = -1\n    print(f\"[FT768-BAL] Fold {fold} training...\")\n\n    for ep in range(FT_EPOCHS):\n        t_ep = time.time()\n        tr_loss = train_one_epoch(model, dl_tr, optimizer, criterion, scaler)\n        scheduler.step()\n        acc, qwk, f1m = evaluate(model, dl_va)\n        print(f\"[FT768-BAL] Fold {fold} Ep {ep+1}/{FT_EPOCHS} | {time.time()-t_ep:5.1f}s | loss {tr_loss:.4f} | acc {acc:.4f} | qwk {qwk:.4f} | f1 {f1m:.4f}\")\n\n        if qwk > best_qwk:\n            best_qwk = qwk\n            torch.save(model.state_dict(), ft_path)\n\n    print(f\"[FT768-BAL] Fold {fold} best QWK after finetune: {best_qwk:.4f}\")\n\nprint(\"[FT768-BAL] Finetune complete.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T18:18:53.491799Z","iopub.execute_input":"2025-12-03T18:18:53.49212Z","iopub.status.idle":"2025-12-03T18:49:12.673447Z","shell.execute_reply.started":"2025-12-03T18:18:53.492096Z","shell.execute_reply":"2025-12-03T18:49:12.672326Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================\n# CELL 9 — Threshold Tuning + Final Predictions (Robust Version)\n# ================================\nprint(\"[TUNING] Running Threshold Optimization for Final Accuracy Boost...\")\n\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score, f1_score\nfrom scipy.special import softmax\nfrom scipy.optimize import minimize\n\ndef infer_all_folds(use_tta=True):\n    \"\"\"Return blended logits from all folds with guaranteed correct shape.\"\"\"\n    all_logits = np.zeros((len(df), 5), dtype=float)\n    fold_indices = list(skf.split(df[\"id_code\"], df[\"diagnosis\"]))\n\n    for fold in range(N_FOLDS):\n        tr_idx, va_idx = fold_indices[fold]\n\n        ckpt = f\"checkpoints/fold{fold}_best_ft768_bal.pth\"\n        print(f\"[TUNING] Loading {ckpt}\")\n\n        model = build_model(num_classes=5, try_pretrained=False)\n        model.load_state_dict(torch.load(ckpt, map_location=DEVICE))\n        model.to(DEVICE)\n        model.eval()\n\n        dl = DataLoader(\n            DRDataset(df, TRAIN_DIR, size=IMG_SIZE_FT, mode=\"valid\"),\n            batch_size=BATCH_SIZE,\n            shuffle=False,\n            num_workers=2,\n            pin_memory=True\n        )\n\n        fold_logits = []\n\n        with torch.no_grad():\n            for imgs, _ in dl:\n                imgs = imgs.to(DEVICE, non_blocking=True)\n                if CHANNELS_LAST:\n                    imgs = imgs.contiguous(memory_format=torch.channels_last)\n\n                if use_tta:\n                    tta_logits = []\n                    for _ in range(TTA_N):\n                        tta_logits.append(model(imgs).cpu().numpy())\n                    batch_logits = np.mean(tta_logits, axis=0)\n                else:\n                    batch_logits = model(imgs).cpu().numpy()\n\n                fold_logits.append(batch_logits)\n\n        fold_logits = np.concatenate(fold_logits, axis=0)\n\n        if fold_logits.shape != (len(df), 5):\n            raise ValueError(f\"Fold {fold} logits shape mismatch: {fold_logits.shape}, expected {(len(df), 5)}\")\n\n        all_logits += fold_logits / N_FOLDS\n\n    return all_logits\n\n\n# 1: Inference\nlogits = infer_all_folds(use_tta=True)\n\n# 2: Probabilities\nprobs = softmax(logits, axis=1)\ntrue_labels = df[\"diagnosis\"].values\n\n\n# 3: Threshold Optimization\ndef qwk_loss(thresholds, probs, y_true):\n    thresholds = np.sort(thresholds)\n    preds = np.digitize(np.argmax(probs, axis=1), thresholds)\n    return -cohen_kappa_score(y_true, preds, weights=\"quadratic\")\n\ninit_thresh = [0.5, 1.5, 2.5, 3.5]\nres = minimize(qwk_loss, init_thresh, args=(probs, true_labels), method='Powell')\nbest_thresh = np.sort(res.x)\n\nprint(\"\\nOptimized Thresholds:\", best_thresh)\n\n# 4: Final predictions\nfinal_preds = np.digitize(np.argmax(probs, axis=1), best_thresh)\n\n# 5: Final scores\nacc = accuracy_score(true_labels, final_preds)\nqwk = cohen_kappa_score(true_labels, final_preds, weights=\"quadratic\")\nf1m = f1_score(true_labels, final_preds, average=\"macro\")\n\nprint(f\"\\n[TUNED RESULTS] Accuracy: {acc:.4f} | QWK: {qwk:.4f} | F1: {f1m:.4f}\")\n\n# 6: Save CSV\ndf_sub = pd.DataFrame({\"id_code\": df[\"id_code\"], \"diagnosis\": final_preds})\ndf_sub.to_csv(\"submission_tuned.csv\", index=False)\n\nprint(\"\\nSaved: submission_tuned.csv ✅\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T18:50:24.774194Z","iopub.execute_input":"2025-12-03T18:50:24.774501Z","iopub.status.idle":"2025-12-03T19:20:55.0955Z","shell.execute_reply.started":"2025-12-03T18:50:24.774475Z","shell.execute_reply":"2025-12-03T19:20:55.094689Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ================================\n# CELL 8 — Test-time inference + submission (cached, TTA, ensemble)\n# ================================\nfrom glob import glob\n\n# Gather test files\ntest_files = []\nif os.path.exists(TEST_DIR):\n    test_files = [f for f in os.listdir(TEST_DIR) if f.lower().endswith(('.png','.jpg','.jpeg'))]\n\nif not len(test_files):\n    print(\"No test images found. Skipping submission generation.\")\nelse:\n    print(f\"Found {len(test_files)} test images. Running inference with TTA={USE_TTA} (N={TTA_N})...\")\n\n    # Build test dataframe\n    test_df = pd.DataFrame({\n        \"id_code\": [os.path.splitext(f)[0] for f in test_files]\n    })\n\n    # Use the same cached dataset class; it will cache test images to CACHE_TEST_DIR\n    test_ds = DRDataset(test_df, TEST_DIR, size=IMG_SIZE_TRAIN, mode=\"valid\")\n    test_dl = DataLoader(\n        test_ds, batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=2, pin_memory=True\n    )\n\n    # Collect per-fold probabilities for soft-voting\n    fold_prob_list = []\n    for fold in range(N_FOLDS):\n        ckpt = f\"checkpoints/fold{fold}_best.pth\"\n        if not os.path.exists(ckpt):\n            print(f\"[WARN] Missing checkpoint for fold {fold}: {ckpt} — skipping this fold.\")\n            continue\n\n        # Build model and load weights\n        model = build_model(num_classes=5, try_pretrained=False)\n        model.load_state_dict(torch.load(ckpt, map_location=DEVICE))\n        model.eval()\n\n        # Inference with probability outputs (softmax), with or without TTA\n        all_probs = []\n        with torch.no_grad():\n            for imgs, _ in test_dl:\n                imgs = imgs.to(DEVICE, non_blocking=True)\n                if CHANNELS_LAST:\n                    imgs = imgs.contiguous(memory_format=torch.channels_last)\n\n                if USE_TTA:\n                    logits_sum = torch.zeros((imgs.size(0), 5), device=DEVICE)\n                    for _ in range(TTA_N):\n                        aug = imgs.flip(-1) if random.random() < 0.5 else imgs\n                        logits_sum += model(aug)\n                    probs = (logits_sum / TTA_N).softmax(1)\n                else:\n                    probs = model(imgs).softmax(1)\n\n                all_probs.append(probs.cpu().numpy())\n\n        fold_probs = np.concatenate(all_probs, axis=0)  # shape: [num_test, 5]\n        fold_prob_list.append(fold_probs)\n        print(f\"Fold {fold}: collected probabilities for {fold_probs.shape[0]} images.\")\n\n    if not len(fold_prob_list):\n        print(\"No fold checkpoints were available. Cannot create submission.\")\n    else:\n        # Soft-vote across folds\n        ens_probs = np.mean(fold_prob_list, axis=0)  # [num_test, 5]\n        ens_pred = ens_probs.argmax(axis=1).astype(int)\n\n        submission = pd.DataFrame({\n            \"id_code\": test_df[\"id_code\"],\n            \"diagnosis\": ens_pred\n        })\n        submission.to_csv(\"submission.csv\", index=False)\n        print(\"Saved submission.csv\")\n\n        # Optional: show quick distribution sanity check\n        counts = pd.Series(ens_pred).value_counts().sort_index()\n        print(\"Submission class distribution:\", dict(counts))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T19:21:07.010596Z","iopub.execute_input":"2025-12-03T19:21:07.011273Z","iopub.status.idle":"2025-12-03T19:30:40.686455Z","shell.execute_reply.started":"2025-12-03T19:21:07.011247Z","shell.execute_reply":"2025-12-03T19:30:40.68573Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save full model\ntorch.save(model, \"final_dr_model_full.pt\")\n\nprint(\"Full model saved to final_dr_model_full.pt\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T19:31:00.978285Z","iopub.execute_input":"2025-12-03T19:31:00.97895Z","iopub.status.idle":"2025-12-03T19:31:01.13489Z","shell.execute_reply.started":"2025-12-03T19:31:00.978918Z","shell.execute_reply":"2025-12-03T19:31:01.134243Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = torch.load(\"final_dr_model_full.pt\", map_location=device, weights_only=False)\nmodel.eval()\nmodel.to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T19:44:00.437569Z","iopub.execute_input":"2025-12-03T19:44:00.437852Z","iopub.status.idle":"2025-12-03T19:44:00.579657Z","shell.execute_reply.started":"2025-12-03T19:44:00.437832Z","shell.execute_reply":"2025-12-03T19:44:00.579017Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.serialization import add_safe_globals\nfrom torch._dynamo.eval_frame import OptimizedModule\n\nadd_safe_globals([OptimizedModule])\n\nmodel = torch.load(\"final_dr_model_full.pt\", map_location=device, weights_only=False)\nmodel.eval()\nmodel.to(device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T19:45:08.782269Z","iopub.execute_input":"2025-12-03T19:45:08.782533Z","iopub.status.idle":"2025-12-03T19:45:08.913557Z","shell.execute_reply.started":"2025-12-03T19:45:08.782516Z","shell.execute_reply":"2025-12-03T19:45:08.912802Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nfrom torchvision import transforms\nfrom PIL import Image\n\n# ---------------------------\n# Load saved FULL model\n# ---------------------------\nmodel_path = \"/kaggle/working/final_dr_model_full.pt\"   # your saved file\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nmodel = torch.load(model_path, map_location=device)\nmodel.eval()\nmodel.to(device)\n\n# ---------------------------\n# Image preprocessing\n# ---------------------------\nIMG_SIZE = 768   # use your finetune size\n\ntransform = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225],\n    )\n])\n\n# ---------------------------\n# Predict function\n# ---------------------------\ndef predict_image(img_path):\n    # Load image\n    img = Image.open(img_path).convert(\"RGB\")\n    img_t = transform(img).unsqueeze(0).to(device)\n\n    # Forward pass\n    with torch.no_grad():\n        logits = model(img_t)\n        probs = F.softmax(logits, dim=1).cpu().numpy()[0]\n\n    # Class names (5 classes)\n    class_names = [\"No DR (0)\", \"Mild (1)\", \"Moderate (2)\", \"Severe (3)\", \"Proliferative DR (4)\"]\n\n    # Get top prediction\n    top_class = probs.argmax()\n    top_prob = probs[top_class] * 100\n\n    print(f\"\\n=== Prediction Result ===\")\n    print(f\"Predicted Class: {class_names[top_class]}  ({top_prob:.2f}%)\\n\")\n\n    print(\"Class Probabilities:\")\n    for i, p in enumerate(probs):\n        print(f\"  {class_names[i]}: {p*100:.2f}%\")\n\n    return top_class, probs\n\n# ---------------------------\n# RUN PREDICTION\n# ---------------------------\n# Example:\n# predict_image(\"/kaggle/input/sample.jpg\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T19:45:14.01246Z","iopub.execute_input":"2025-12-03T19:45:14.012732Z","iopub.status.idle":"2025-12-03T19:45:14.036897Z","shell.execute_reply.started":"2025-12-03T19:45:14.012712Z","shell.execute_reply":"2025-12-03T19:45:14.036016Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from cbam_resnet import ResNet50_CBAM_TV\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T19:47:13.417716Z","iopub.execute_input":"2025-12-03T19:47:13.418304Z","iopub.status.idle":"2025-12-03T19:47:13.434428Z","shell.execute_reply.started":"2025-12-03T19:47:13.418276Z","shell.execute_reply":"2025-12-03T19:47:13.43342Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.serialization import add_safe_globals\n\n# -------------------------\n# 1. Import your model class\n# -------------------------\n# IMPORTANT: this must match EXACT class used during training\nfrom your_model_file import ResNet50_CBAM_TV   # <-- CHANGE to your actual file name\n\n# Example:\n# from model import ResNet50_CBAM_TV\n\n# -------------------------\n# 2. Add model to safe globals\n# -------------------------\nadd_safe_globals([ResNet50_CBAM_TV])\n\n# -------------------------\n# 3. Load FULL model\n# -------------------------\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nmodel_path = \"/kaggle/working/final_dr_model_full.pt\"   # <-- your saved file path\n\nmodel = torch.load(\n    model_path,\n    map_location=device,\n    weights_only=False   # MUST BE FALSE\n)\n\nmodel.to(device)\nmodel.eval()\n\nprint(\"Model loaded successfully!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-03T19:46:24.846673Z","iopub.execute_input":"2025-12-03T19:46:24.847008Z","iopub.status.idle":"2025-12-03T19:46:24.873828Z","shell.execute_reply.started":"2025-12-03T19:46:24.846961Z","shell.execute_reply":"2025-12-03T19:46:24.872888Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}