{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpuV5e8","dataSources":[{"sourceType":"competition","sourceId":14420,"databundleVersionId":868327}],"dockerImageVersionId":31331,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!/usr/bin/env python\n\"\"\"\nRecursion Cellular Image Classification\n2-Headed EfficientNet with 2-Stage Training (PyTorch version)\n=============================================================\nConverted from the original Keras notebook:\n  - 2 inputs (site1 + site2) → shared EfficientNet → GAP → add → classify\n  - Phase 1: Train on ALL data (10 epochs)\n  - Phase 2: Fine-tune on each cell type separately (10 epochs)\n\nKey fixes from original notebook:\n  - Reads directly from competition images (no separate preprocessed dataset)\n  - Modern PyTorch + timm (not old Keras/TF1)\n  - Proper sirna label handling (sirna_XXX → int → 0-indexed)\n  - Per-plate pixel normalization (competition best practice)\n\"\"\"\n\n# %% [markdown]\n# # Cell 1 — Setup\n\n# %%\nimport subprocess, sys\nsubprocess.check_call([sys.executable, \"-m\", \"pip\", \"install\", \"-q\", \"timm\"])\n\nimport gc, json, time, random, warnings, os\nfrom pathlib import Path\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport timm\nfrom sklearn.model_selection import train_test_split\n\nwarnings.filterwarnings(\"ignore\")\n\n# ─── TPU Detection ───\nUSE_TPU = False\ntry:\n    import torch_xla\n    import torch_xla.core.xla_model as xm\n    import torch_xla.distributed.parallel_loader as pl\n    USE_TPU = True\n    print(\"✓ TPU detected (torch_xla available)\")\nexcept ImportError:\n    print(\"No TPU — using CUDA/CPU\")\n\nclass CFG:\n    DATA_DIR    = Path(\"/kaggle/input/competitions/recursion-cellular-image-classification\")\n    OUTPUT_DIR  = Path(\"/kaggle/working\")\n    IMG_SIZE    = 300\n    N_CLASSES   = 1108\n    DROPOUT     = 0.5\n    EPOCHS_P1   = 10\n    EPOCHS_P2   = 10\n    BATCH_SIZE  = 64 if USE_TPU else 32   # TPU has more memory\n    LR          = 0.0001\n    WEIGHT_DECAY = 1e-5\n    NUM_WORKERS  = 4 if USE_TPU else 2\n    SEED = 2019\n\ndef seed_all(s=2019):\n    random.seed(s); np.random.seed(s); torch.manual_seed(s)\n    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)\n    torch.backends.cudnn.deterministic = True\n\nseed_all(CFG.SEED)\n\n# Device selection: TPU > CUDA > CPU\nif USE_TPU:\n    DEVICE = xm.xla_device()\n    print(f\"Device: TPU ({DEVICE})\")\nelif torch.cuda.is_available():\n    DEVICE = torch.device(\"cuda\")\n    print(f\"Device: {DEVICE} ({torch.cuda.get_device_name(0)})\")\nelse:\n    DEVICE = torch.device(\"cpu\")\n    print(f\"Device: {DEVICE}\")\n\n# %% [markdown]\n# # Cell 2 — Per-Plate Stats\n\n# %%\ndef compute_plate_stats(df, data_dir, n_samples=12):\n    stats = {}\n    for (exp, plate), grp in df.groupby([\"experiment\", \"plate\"]):\n        ch_vals = {c: [] for c in range(1, 7)}\n        sample = grp.sample(min(n_samples, len(grp)), random_state=42)\n        for _, row in sample.iterrows():\n            for site in [1, 2]:\n                for ch in range(1, 7):\n                    p = data_dir/row[\"experiment\"]/f\"Plate{row['plate']}\"/f\"{row['well']}_s{site}_w{ch}.png\"\n                    if p.exists():\n                        ch_vals[ch].append(np.array(Image.open(p), np.float32).mean())\n        m, s = [], []\n        for ch in range(1, 7):\n            m.append(np.mean(ch_vals[ch]) if ch_vals[ch] else 0.0)\n            s.append(max(np.std(ch_vals[ch]), 1e-6) if ch_vals[ch] else 1.0)\n        stats[(exp, plate)] = {\"mean\": m, \"std\": s}\n    return stats\n\n# %% [markdown]\n# # Cell 3 — Dataset (2-headed: site1 + site2)\n\n# %%\nclass TwoSiteDataset(Dataset):\n    \"\"\"\n    Loads BOTH site1 and site2 images separately (like the original notebook).\n    Returns (img_site1, img_site2, label) for train\n    Returns (img_site1, img_site2, id_code) for test\n    \"\"\"\n    def __init__(self, df, data_dir, pstats=None, is_train=True,\n                 img_size=300, swap=False):\n        self.df = df.reset_index(drop=True)\n        self.root = Path(data_dir)\n        self.pstats = pstats\n        self.is_train = is_train\n        self.img_size = img_size\n        self.split = \"train\" if is_train else \"test\"\n        self.swap = swap\n\n    def __len__(self): return len(self.df)\n\n    def _load_site(self, exp, plate, well, site):\n        \"\"\"Load 6-channel image for one site.\"\"\"\n        chs = []\n        for c in range(1, 7):\n            p = self.root/self.split/exp/f\"Plate{plate}\"/f\"{well}_s{site}_w{c}.png\"\n            chs.append(np.array(Image.open(p), np.float32) if p.exists()\n                       else np.zeros((512, 512), np.float32))\n        return np.stack(chs, -1)  # [512, 512, 6]\n\n    def _norm(self, img, exp, plate):\n        if self.pstats and (exp, plate) in self.pstats:\n            s = self.pstats[(exp, plate)]\n            return (img - np.array(s[\"mean\"], np.float32)) / np.array(s[\"std\"], np.float32)\n        return (img - img.mean((0,1))) / (img.std((0,1)) + 1e-6)\n\n    def _resize(self, img):\n        if img.shape[0] != self.img_size:\n            rs = []\n            for c in range(6):\n                ch = Image.fromarray(img[:,:,c]).resize(\n                    (self.img_size, self.img_size), Image.BILINEAR)\n                rs.append(np.array(ch, np.float32))\n            return np.stack(rs, -1)\n        return img\n\n    def _augment(self, img):\n        \"\"\"Random 90° rotation + flip (same as original notebook).\"\"\"\n        if random.random() > 0.5: img = np.flip(img, 0).copy()\n        if random.random() > 0.5: img = np.flip(img, 1).copy()\n        img = np.rot90(img, random.randint(0,3), (0,1)).copy()\n        return img\n\n    def _to_tensor(self, img):\n        return torch.from_numpy(img.transpose(2,0,1).copy()).float()\n\n    def __getitem__(self, idx):\n        r = self.df.iloc[idx]\n        exp, plate, well = r[\"experiment\"], r[\"plate\"], r[\"well\"]\n\n        img1 = self._load_site(exp, plate, well, 1)\n        img2 = self._load_site(exp, plate, well, 2)\n\n        img1 = self._norm(img1, exp, plate)\n        img2 = self._norm(img2, exp, plate)\n\n        img1 = self._resize(img1)\n        img2 = self._resize(img2)\n\n        if self.is_train:\n            img1 = self._augment(img1)\n            img2 = self._augment(img2)\n            # Swap sites randomly (like original notebook)\n            if self.swap and random.random() > 0.5:\n                img1, img2 = img2, img1\n\n        img1 = self._to_tensor(img1)\n        img2 = self._to_tensor(img2)\n\n        if self.is_train and \"sirna\" in r.index:\n            return img1, img2, int(r[\"sirna\"])\n        return img1, img2, r[\"id_code\"]\n\n# %% [markdown]\n# # Cell 4 — 2-Headed EfficientNet Model\n\n# %%\nclass TwoHeadedEfficientNet(nn.Module):\n    \"\"\"\n    2-Headed architecture (same concept as original notebook):\n      - Shared EfficientNet backbone processes site1 and site2\n      - Global Average Pooling on each\n      - ADD the two feature vectors (not concat — same as original)\n      - Dropout → Dense classification\n    \"\"\"\n    def __init__(self, n_classes=1108, in_chans=6, drop=0.5):\n        super().__init__()\n        # Shared backbone — EfficientNet-B2 (original used B3)\n        self.backbone = timm.create_model(\n            \"tf_efficientnet_b2\",\n            pretrained=True,\n            in_chans=in_chans,\n            num_classes=0,       # feature extractor only\n            drop_rate=drop,\n        )\n        n_feat = self.backbone.num_features  # 1408 for B2\n\n        # Classification head\n        self.drop = nn.Dropout(drop)\n        self.fc = nn.Linear(n_feat, n_classes)\n\n    def forward(self, site1, site2):\n        \"\"\"\n        site1, site2: [B, 6, H, W]\n        Both pass through the SAME backbone (shared weights).\n        \"\"\"\n        feat1 = self.backbone(site1)  # [B, n_feat]\n        feat2 = self.backbone(site2)  # [B, n_feat]\n\n        # Add features from both sites (same as original)\n        combined = feat1 + feat2\n\n        out = self.drop(combined)\n        return self.fc(out)\n\n# %% [markdown]\n# # Cell 5 — Train / Validate / Predict\n\n# %%\ndef train_epoch(model, dl, opt, sched, dev):\n    model.train(); tloss = tc = tt = 0\n    for i, (s1, s2, y) in enumerate(dl):\n        s1, s2, y = s1.to(dev), s2.to(dev), y.to(dev)\n        logits = model(s1, s2)\n        loss = F.cross_entropy(logits, y, label_smoothing=0.1)\n        opt.zero_grad(); loss.backward()\n        nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n        if USE_TPU:\n            xm.optimizer_step(opt)     # TPU: sync + step\n            xm.mark_step()             # TPU: execute graph\n        else:\n            opt.step()\n        tloss += loss.item()*s1.size(0)\n        tc += (logits.argmax(1)==y).sum().item(); tt += s1.size(0)\n        if (i+1) % 100 == 0:\n            print(f\"      batch {i+1}/{len(dl)} loss={loss.item():.4f} acc={tc/tt:.4f}\")\n    sched.step()\n    return tloss/tt, tc/tt\n\n@torch.no_grad()\ndef validate(model, dl, dev):\n    model.eval(); c = t = 0\n    for s1, s2, y in dl:\n        s1, s2, y = s1.to(dev), s2.to(dev), y.to(dev)\n        c += (model(s1, s2).argmax(1)==y).sum().item(); t += s1.size(0)\n    return c/t\n\n@torch.no_grad()\ndef predict(model, dl, dev):\n    model.eval(); all_l, all_id = [], []\n    for s1, s2, ids in dl:\n        s1, s2 = s1.to(dev), s2.to(dev)\n        logits = model(s1, s2)\n        # TTA: also swap sites\n        logits_swap = model(s2, s1)\n        all_l.append(((logits + logits_swap) / 2).cpu())\n        all_id.extend(ids if isinstance(ids, (list, tuple)) else\n                      (ids.tolist() if torch.is_tensor(ids) else list(ids)))\n    return torch.cat(all_l, 0), all_id\n\n# %% [markdown]\n# # Cell 6 — Main Pipeline\n\n# %%\ndef main():\n    print(\"=\"*60 + \"\\nSTEP 1: Loading data\\n\" + \"=\"*60)\n    train_csv = pd.read_csv(CFG.DATA_DIR/\"train.csv\")\n    test_csv  = pd.read_csv(CFG.DATA_DIR/\"test.csv\")\n    sample_sub = pd.read_csv(CFG.DATA_DIR/\"sample_submission.csv\")\n\n    # Parse sirna\n    train_csv[\"sirna\"] = train_csv[\"sirna\"].str.replace(\"sirna_\",\"\").astype(int)\n    train_csv[\"category\"] = train_csv[\"experiment\"].str.split(\"-\").str[0]\n    test_csv[\"category\"]  = test_csv[\"experiment\"].str.split(\"-\").str[0]\n\n    # Remap labels to 0-indexed\n    uniq = sorted(train_csv[\"sirna\"].unique())\n    s2i = {s:i for i,s in enumerate(uniq)}\n    i2s = {i:s for s,i in s2i.items()}\n    train_csv[\"sirna\"] = train_csv[\"sirna\"].map(s2i)\n    CFG.N_CLASSES = len(uniq)\n    print(f\"  N_CLASSES={CFG.N_CLASSES}\")\n\n    # Filter test\n    valid = set(sample_sub[\"id_code\"])\n    test_csv = test_csv[test_csv[\"id_code\"].isin(valid)].reset_index(drop=True)\n    print(f\"  Train={len(train_csv)}, Test={len(test_csv)}\")\n    print(f\"  Categories: {train_csv['category'].unique().tolist()}\")\n\n    # Plate stats\n    print(\"\\nComputing plate stats...\")\n    ps_tr = compute_plate_stats(train_csv, CFG.DATA_DIR/\"train\")\n    ps_te = compute_plate_stats(test_csv, CFG.DATA_DIR/\"test\")\n    print(f\"  Train plates={len(ps_tr)}, Test plates={len(ps_te)}\")\n\n    # ═══════════════════════════════════════════════════════════\n    # PHASE 1: Train on ALL data (same as original notebook)\n    # ═══════════════════════════════════════════════════════════\n    print(\"\\n\" + \"=\"*60)\n    print(f\"PHASE 1: Train on ALL data ({CFG.EPOCHS_P1} epochs)\")\n    print(\"=\"*60)\n\n    # 85/15 split (same as original)\n    train_idx, val_idx = train_test_split(\n        train_csv.index, test_size=0.15, random_state=2019)\n    trn = train_csv.loc[train_idx].reset_index(drop=True)\n    val = train_csv.loc[val_idx].reset_index(drop=True)\n    print(f\"  Train={len(trn)}, Val={len(val)}\")\n\n    tr_dl = DataLoader(\n        TwoSiteDataset(trn, CFG.DATA_DIR, ps_tr, True, CFG.IMG_SIZE, swap=True),\n        CFG.BATCH_SIZE, True, num_workers=CFG.NUM_WORKERS,\n        pin_memory=True, drop_last=True)\n    va_dl = DataLoader(\n        TwoSiteDataset(val, CFG.DATA_DIR, ps_tr, True, CFG.IMG_SIZE),\n        CFG.BATCH_SIZE, num_workers=CFG.NUM_WORKERS, pin_memory=True)\n\n    model = TwoHeadedEfficientNet(CFG.N_CLASSES, 6, CFG.DROPOUT).to(DEVICE)\n    print(f\"  Params: {sum(p.numel() for p in model.parameters()):,}\")\n\n    opt = torch.optim.Adam(model.parameters(), CFG.LR, weight_decay=CFG.WEIGHT_DECAY)\n    sched = CosineAnnealingLR(opt, CFG.EPOCHS_P1, 1e-6)\n\n    history = []\n    best_acc, best_state = 0, None\n\n    for ep in range(CFG.EPOCHS_P1):\n        t0 = time.time()\n        tl, ta = train_epoch(model, tr_dl, opt, sched, DEVICE)\n        va = validate(model, va_dl, DEVICE)\n        elapsed = time.time() - t0\n        history.append({\"epoch\": ep+1, \"phase\": \"P1-ALL\", \"train_loss\": tl,\n                        \"train_acc\": ta, \"val_acc\": va, \"time\": elapsed})\n        print(f\"  [P1] Epoch {ep+1}/{CFG.EPOCHS_P1} loss={tl:.4f} \"\n              f\"train={ta:.4f} val={va:.4f} {elapsed:.0f}s\")\n        if va > best_acc:\n            best_acc = va\n            best_state = {k: v.cpu().clone() for k,v in model.state_dict().items()}\n            print(f\"    ★ Best: {va:.4f}\")\n            torch.save(best_state, CFG.OUTPUT_DIR/\"model_phase1.pt\")\n        gc.collect()\n        if USE_TPU:\n            xm.mark_step()\n        elif torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n    print(f\"\\nPhase 1 done. Best val: {best_acc:.4f}\")\n    # Load best phase 1 weights\n    if best_state:\n        model.load_state_dict(best_state)\n\n    # ═══════════════════════════════════════════════════════════\n    # PHASE 2: Train per cell type (same as original notebook)\n    # ═══════════════════════════════════════════════════════════\n    categories = [\"HEPG2\", \"HUVEC\", \"RPE\", \"U2OS\"]\n    output_preds = []\n\n    for cat in categories:\n        print(\"\\n\" + \"=\"*60)\n        print(f\"PHASE 2: Fine-tuning on {cat} ({CFG.EPOCHS_P2} epochs)\")\n        print(\"=\"*60)\n\n        # Filter data by category\n        cat_train = train_csv[train_csv[\"category\"] == cat].copy()\n        cat_test  = test_csv[test_csv[\"category\"] == cat].copy()\n        print(f\"  {cat}: Train={len(cat_train)}, Test={len(cat_test)}\")\n\n        if len(cat_train) == 0 or len(cat_test) == 0:\n            print(f\"  SKIP {cat} — no data\")\n            continue\n\n        # 85/15 split within category\n        cat_tr_idx, cat_va_idx = train_test_split(\n            cat_train.index, test_size=0.15, random_state=2019)\n        cat_trn = cat_train.loc[cat_tr_idx].reset_index(drop=True)\n        cat_val = cat_train.loc[cat_va_idx].reset_index(drop=True)\n\n        cat_tr_dl = DataLoader(\n            TwoSiteDataset(cat_trn, CFG.DATA_DIR, ps_tr, True, CFG.IMG_SIZE, swap=True),\n            CFG.BATCH_SIZE, True, num_workers=CFG.NUM_WORKERS,\n            pin_memory=True, drop_last=True)\n        cat_va_dl = DataLoader(\n            TwoSiteDataset(cat_val, CFG.DATA_DIR, ps_tr, True, CFG.IMG_SIZE),\n            CFG.BATCH_SIZE, num_workers=CFG.NUM_WORKERS, pin_memory=True)\n        cat_te_dl = DataLoader(\n            TwoSiteDataset(cat_test, CFG.DATA_DIR, ps_te, False, CFG.IMG_SIZE),\n            1, num_workers=CFG.NUM_WORKERS, pin_memory=True)\n\n        # Reload Phase 1 best weights\n        model.load_state_dict(best_state)\n        model.to(DEVICE)\n\n        opt = torch.optim.Adam(model.parameters(), CFG.LR, weight_decay=CFG.WEIGHT_DECAY)\n        sched = CosineAnnealingLR(opt, CFG.EPOCHS_P2, 1e-6)\n\n        cat_best_acc, cat_best_state = 0, None\n\n        for ep in range(CFG.EPOCHS_P2):\n            t0 = time.time()\n            tl, ta = train_epoch(model, cat_tr_dl, opt, sched, DEVICE)\n            va = validate(model, cat_va_dl, DEVICE)\n            elapsed = time.time() - t0\n            history.append({\"epoch\": ep+1, \"phase\": f\"P2-{cat}\", \"train_loss\": tl,\n                            \"train_acc\": ta, \"val_acc\": va, \"time\": elapsed})\n            print(f\"  [{cat}] Ep {ep+1}/{CFG.EPOCHS_P2} loss={tl:.4f} \"\n                  f\"train={ta:.4f} val={va:.4f} {elapsed:.0f}s\")\n            if va > cat_best_acc:\n                cat_best_acc = va\n                cat_best_state = {k: v.cpu().clone() for k,v in model.state_dict().items()}\n                print(f\"    ★ {va:.4f}\")\n            gc.collect()\n            if USE_TPU:\n                xm.mark_step()\n            elif torch.cuda.is_available():\n                torch.cuda.empty_cache()\n\n        print(f\"  {cat} best val: {cat_best_acc:.4f}\")\n\n        # Load best per-category weights and predict\n        if cat_best_state:\n            model.load_state_dict(cat_best_state)\n            model.to(DEVICE)\n            torch.save(cat_best_state, CFG.OUTPUT_DIR/f\"model_{cat}.pt\")\n\n        logits, ids = predict(model, cat_te_dl, DEVICE)\n        preds = logits.argmax(1).numpy()\n        preds_original = np.array([i2s[int(p)] for p in preds])\n\n        cat_pred_df = pd.DataFrame({\"id_code\": ids, \"sirna\": preds_original.astype(int)})\n        output_preds.append(cat_pred_df)\n        print(f\"  {cat} predictions: {len(cat_pred_df)}\")\n\n    # ═══════════════════════════════════════════════════════════\n    # SUBMISSION\n    # ═══════════════════════════════════════════════════════════\n    print(\"\\n\" + \"=\"*60 + \"\\nGenerating Submission\\n\" + \"=\"*60)\n\n    # Combine per-category predictions\n    all_preds = pd.concat(output_preds, ignore_index=True)\n    sub = sample_sub[[\"id_code\"]].merge(all_preds, on=\"id_code\", how=\"left\")\n    sub[\"sirna\"] = sub[\"sirna\"].fillna(0).astype(int)\n    sub.to_csv(CFG.OUTPUT_DIR/\"submission.csv\", index=False)\n    print(f\"  ✓ submission.csv ({len(sub)} rows)\")\n    print(f\"  Head:\\n{sub.head()}\")\n\n    # Save history\n    hist_df = pd.DataFrame(history)\n    hist_df.to_csv(CFG.OUTPUT_DIR/\"training_history.csv\", index=False)\n    print(f\"  ✓ training_history.csv saved\")\n\n    print(f\"\\n{'='*60}\\nDONE!\\n{'='*60}\")\n    return sub, hist_df\n\n# %%\nif __name__ == \"__main__\":\n    submission, history = main()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-04-29T13:27:58.188211Z","iopub.execute_input":"2026-04-29T13:27:58.188411Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}