{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":14420,"databundleVersionId":868327}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nRxRx1 Cellular Image Classification - Kaggle Notebook (FIXED)\n=============================================================\nBUGS FIXED from previous version:\n  1. CRITICAL: augmentation was clamping to [0,1] AFTER z-score normalization,\n     destroying the signal. Now: spatial augments only after normalization.\n  2. Intensity augments moved BEFORE normalization (on raw [0,1] data).\n  3. Using both sites for training (more data = better accuracy).\n  4. Added label smoothing for better generalization.\n  5. Proper submission format handling.\n\nUpload as Kaggle notebook with GPU T4 + Internet ON + competition dataset attached.\nExpected: ~3-4 hours, ~0.50+ leaderboard accuracy.\n\"\"\"\n\nimport os, sys, math, time, random, gc\nfrom collections import defaultdict\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom sklearn.model_selection import StratifiedKFold\nfrom PIL import Image\nfrom tqdm import tqdm\n\nos.system(\"pip install -q efficientnet_pytorch\")\nfrom efficientnet_pytorch import EfficientNet\n\n# ══════════════════════════════════════════════════════════════════════════════\n# CONFIG\n# ══════════════════════════════════════════════════════════════════════════════\nOUTPUT_DIR = \"/kaggle/working\"\n\ndef find_data_dir():\n    base = \"/kaggle/input\"\n    if os.path.isdir(base):\n        for root, dirs, files in os.walk(base):\n            if root.replace(base, \"\").count(os.sep) > 3:\n                continue\n            if \"train.csv\" in files:\n                print(f\"[DATA] Found train.csv at: {root}\")\n                return root\n        print(f\"[DEBUG] /kaggle/input: {os.listdir(base)}\")\n        for d in os.listdir(base):\n            full = os.path.join(base, d)\n            if os.path.isdir(full):\n                print(f\"[DEBUG]   {d}/: {os.listdir(full)[:10]}\")\n    return os.path.join(base, \"recursion-cellular-image-classification\")\n\ndef find_image_dir(csv_dir, split=\"train\"):\n    obvious = os.path.join(csv_dir, split)\n    if os.path.isdir(obvious):\n        subs = os.listdir(obvious)\n        if any(s.startswith((\"HEPG2\",\"HUVEC\",\"RPE\",\"U2OS\")) for s in subs):\n            print(f\"[DATA] {split}/ at: {obvious} ({len(subs)} folders)\")\n            return obvious\n    base = \"/kaggle/input\"\n    print(f\"[DATA] {split}/ not at {obvious}, searching...\")\n    for root, dirs, files in os.walk(base):\n        if root.replace(base, \"\").count(os.sep) > 4:\n            continue\n        if os.path.basename(root) == split:\n            subs = os.listdir(root)\n            if any(s.startswith((\"HEPG2\",\"HUVEC\",\"RPE\",\"U2OS\")) for s in subs):\n                print(f\"[DATA] {split}/ found at: {root}\")\n                return root\n        if any(d.startswith((\"HEPG2\",\"HUVEC\",\"RPE\",\"U2OS\")) for d in dirs):\n            marker = \"HEPG2-01\" if split == \"train\" else \"HEPG2-08\"\n            if marker in dirs:\n                print(f\"[DATA] {split}/ images (flat) at: {root}\")\n                return root\n    print(f\"[WARN] Could not find {split}/ images, using: {obvious}\")\n    return obvious\n\nINPUT_DIR = find_data_dir()\nTRAIN_DIR = find_image_dir(INPUT_DIR, \"train\")\nTEST_DIR = find_image_dir(INPUT_DIR, \"test\")\nTRAIN_CSV = os.path.join(INPUT_DIR, \"train.csv\")\nTEST_CSV = os.path.join(INPUT_DIR, \"test.csv\")\nPIXEL_STATS_CSV = os.path.join(INPUT_DIR, \"pixel_stats.csv\")\n\nfor p, n in [(TRAIN_CSV,\"train.csv\"),(TEST_CSV,\"test.csv\"),(PIXEL_STATS_CSV,\"pixel_stats.csv\")]:\n    print(f\"  {'✓' if os.path.isfile(p) else '✗'} {n}: {p}\")\n\nNUM_CHANNELS = 6\nCROP_SIZE = 256\nNUM_CLASSES = 1108\nCELL_TYPES = [\"HEPG2\", \"HUVEC\", \"RPE\", \"U2OS\"]\nCELL_TYPE_TO_IDX = {ct: i for i, ct in enumerate(CELL_TYPES)}\nNUM_CELL_TYPES = len(CELL_TYPES)\n\nBACKBONE = \"efficientnet-b2\"\nFEATURE_DIM = 1408\nHEAD_HIDDEN = 512\nDROPOUT = 0.3\n\nBATCH_SIZE = 32\nNUM_WORKERS = 4\nMAX_EPOCHS = 45\nLEARNING_RATE = 3e-4\nWEIGHT_DECAY = 1e-4\nGRADIENT_CLIP = 1.0\nLABEL_SMOOTHING = 0.1\nEARLY_STOP_PATIENCE = 7\nSEED = 42\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nUSE_AMP = torch.cuda.is_available()\nprint(f\"Device: {DEVICE}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# PIXEL STATS\n# ══════════════════════════════════════════════════════════════════════════════\nclass PixelStatsLookup:\n    def __init__(self, csv_path=PIXEL_STATS_CSV):\n        self._lookup = {}\n        if not os.path.isfile(csv_path):\n            print(\"[WARN] pixel_stats.csv not found, using default normalization\")\n            return\n        df = pd.read_csv(csv_path)\n        g = df.groupby([\"experiment\",\"plate\",\"channel\"]).agg(\n            {\"mean\":\"mean\",\"std\":\"mean\"}).reset_index()\n        for _, r in g.iterrows():\n            key = (r[\"experiment\"], int(r[\"plate\"]), int(r[\"channel\"]))\n            self._lookup[key] = (float(r[\"mean\"]), max(float(r[\"std\"]), 1e-6))\n        print(f\"[DATA] Loaded pixel stats: {len(self._lookup)} entries\")\n\n    def get(self, experiment, plate):\n        \"\"\"Return (means, stds) arrays of shape (6,) in [0,255] scale.\"\"\"\n        means = np.zeros(NUM_CHANNELS, dtype=np.float32)\n        stds = np.ones(NUM_CHANNELS, dtype=np.float32) * 40.0  # reasonable default\n        for ch in range(NUM_CHANNELS):\n            key = (experiment, plate, ch + 1)\n            if key in self._lookup:\n                means[ch] = self._lookup[key][0]\n                stds[ch] = self._lookup[key][1]\n        return means, stds\n\npixel_stats = PixelStatsLookup()\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# SIRNA LABEL MAPPING\n# ══════════════════════════════════════════════════════════════════════════════\ndef build_sirna_mapping(df):\n    \"\"\"Build consistent sirna -> integer index mapping.\n    Handles both string ('sirna_250') and integer (250) formats.\n    \"\"\"\n    raw_values = sorted(df[\"sirna\"].unique().tolist())\n    sirna_to_idx = {s: i for i, s in enumerate(raw_values)}\n    idx_to_sirna = {i: s for s, i in sirna_to_idx.items()}\n    print(f\"[DATA] siRNA mapping: {len(sirna_to_idx)} classes\")\n    print(f\"[DATA] Sample labels: {raw_values[:5]} ... {raw_values[-3:]}\")\n    return sirna_to_idx, idx_to_sirna\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# DATASET — FIXED augmentation order\n# ══════════════════════════════════════════════════════════════════════════════\ndef load_6ch(base_dir, experiment, plate, well, site):\n    \"\"\"Load 6-channel image as (6, H, W) float32 in [0, 1].\"\"\"\n    channels = []\n    for ch in range(1, NUM_CHANNELS + 1):\n        path = os.path.join(base_dir, experiment, f\"Plate{plate}\",\n                            f\"{well}_s{site}_w{ch}.png\")\n        img = Image.open(path)\n        channels.append(np.array(img, dtype=np.float32) / 255.0)\n    return np.stack(channels, axis=0)\n\ndef random_crop(img, size=CROP_SIZE):\n    _, h, w = img.shape\n    if h <= size or w <= size:\n        return img\n    top = random.randint(0, h - size)\n    left = random.randint(0, w - size)\n    return img[:, top:top+size, left:left+size]\n\ndef center_crop(img, size=CROP_SIZE):\n    _, h, w = img.shape\n    if h <= size or w <= size:\n        return img\n    top = (h - size) // 2\n    left = (w - size) // 2\n    return img[:, top:top+size, left:left+size]\n\ndef augment_raw(img_np):\n    \"\"\"Intensity augmentation on RAW [0,1] data BEFORE normalization.\n    This is where brightness/contrast changes belong.\n    \"\"\"\n    # Brightness\n    if random.random() < 0.5:\n        factor = 1.0 + random.uniform(-0.15, 0.15)\n        img_np = img_np * factor\n    # Contrast\n    if random.random() < 0.3:\n        factor = 1.0 + random.uniform(-0.15, 0.15)\n        for ch in range(img_np.shape[0]):\n            m = img_np[ch].mean()\n            img_np[ch] = (img_np[ch] - m) * factor + m\n    # Channel dropout\n    if random.random() < 0.1:\n        ch = random.randint(0, img_np.shape[0] - 1)\n        img_np[ch] = 0.0\n    return np.clip(img_np, 0.0, 1.0)  # clamp is OK here (raw data)\n\ndef augment_spatial(tensor):\n    \"\"\"Spatial augmentation AFTER normalization — no value clamping needed.\n    Flips and rotations don't change value ranges.\n    \"\"\"\n    if random.random() < 0.5:\n        tensor = torch.flip(tensor, [2])  # H-flip\n    if random.random() < 0.5:\n        tensor = torch.flip(tensor, [1])  # V-flip\n    k = random.choice([0, 1, 2, 3])\n    if k > 0:\n        tensor = torch.rot90(tensor, k, [1, 2])\n    return tensor  # NO CLAMP — normalized values must stay as-is\n\n\nclass RxRxDataset(Dataset):\n    \"\"\"\n    Dataset with FIXED augmentation pipeline:\n      1. Load raw image [0,1]\n      2. Crop (random for train, center for val/test)\n      3. Intensity augment on raw data (train only) — clamp OK here\n      4. Z-score normalize with pixel_stats\n      5. Spatial augment (train only) — NO clamp after this\n    \"\"\"\n    def __init__(self, df, base_dir, sirna_to_idx, training=True, sites=(1,)):\n        self.base_dir = base_dir\n        self.training = training\n        self.sirna_to_idx = sirna_to_idx\n        self.samples = []\n        has_label = \"sirna\" in df.columns\n        skipped = 0\n\n        for _, row in df.iterrows():\n            exp = row[\"experiment\"]\n            plate = int(row[\"plate\"])\n            well = row[\"well\"]\n            sirna = sirna_to_idx[row[\"sirna\"]] if has_label else -1\n            id_code = row.get(\"id_code\", \"\")\n            ct = exp.split(\"-\")[0]\n\n            for site in sites:\n                img_path = os.path.join(base_dir, exp, f\"Plate{plate}\",\n                                        f\"{well}_s{site}_w1.png\")\n                if not os.path.isfile(img_path):\n                    skipped += 1\n                    continue\n                self.samples.append({\n                    \"experiment\": exp, \"plate\": plate, \"well\": well,\n                    \"site\": site, \"cell_type\": ct, \"sirna\": sirna,\n                    \"id_code\": id_code,\n                })\n\n        if skipped:\n            print(f\"[Dataset] WARNING: Skipped {skipped} samples (missing images)\")\n        print(f\"[Dataset] {len(self.samples)} samples ready \"\n              f\"({'train' if training else 'val/test'})\")\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        s = self.samples[idx]\n\n        # 1. Load raw image\n        img = load_6ch(self.base_dir, s[\"experiment\"], s[\"plate\"],\n                       s[\"well\"], s[\"site\"])  # (6, 512, 512), [0, 1]\n\n        # 2. Crop\n        if self.training:\n            img = random_crop(img)\n        else:\n            img = center_crop(img)\n\n        # 3. Intensity augmentation on RAW data (before normalization)\n        if self.training:\n            img = augment_raw(img)\n\n        # 4. Z-score normalization (values now in ~[-3, 3])\n        means, stds = pixel_stats.get(s[\"experiment\"], s[\"plate\"])\n        means_s = means / 255.0\n        stds_s = stds / 255.0\n        for ch in range(NUM_CHANNELS):\n            img[ch] = (img[ch] - means_s[ch]) / stds_s[ch]\n\n        tensor = torch.from_numpy(img).float()\n\n        # 5. Spatial augmentation AFTER normalization (NO clamping!)\n        if self.training:\n            tensor = augment_spatial(tensor)\n\n        # Cell-type one-hot\n        cell_oh = torch.zeros(NUM_CELL_TYPES, dtype=torch.float32)\n        cell_oh[CELL_TYPE_TO_IDX.get(s[\"cell_type\"], 0)] = 1.0\n\n        return tensor, cell_oh, s[\"sirna\"]\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# MODEL\n# ══════════════════════════════════════════════════════════════════════════════\nclass RxRxModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = EfficientNet.from_pretrained(BACKBONE)\n\n        # Replace 3-channel conv with 6-channel, keeping pretrained weights\n        old = self.backbone._conv_stem\n        new = nn.Conv2d(NUM_CHANNELS, old.out_channels,\n                        kernel_size=old.kernel_size,\n                        stride=old.stride, padding=old.padding, bias=False)\n        with torch.no_grad():\n            new.weight[:, :3] = old.weight\n            # Initialize extra channels by duplicating RGB weights\n            new.weight[:, 3:] = old.weight\n        self.backbone._conv_stem = new\n\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        feat_dim = FEATURE_DIM + NUM_CELL_TYPES\n        self.head = nn.Sequential(\n            nn.BatchNorm1d(feat_dim),\n            nn.Dropout(DROPOUT),\n            nn.Linear(feat_dim, HEAD_HIDDEN),\n            nn.ReLU(inplace=True),\n            nn.BatchNorm1d(HEAD_HIDDEN),\n            nn.Dropout(DROPOUT / 2),\n            nn.Linear(HEAD_HIDDEN, NUM_CLASSES),\n        )\n        for m in self.head.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.xavier_uniform_(m.weight)\n                if m.bias is not None:\n                    nn.init.zeros_(m.bias)\n\n    def extract_features(self, x):\n        f = self.backbone.extract_features(x)\n        return self.pool(f).flatten(1)\n\n    def forward(self, x, cell_type):\n        f = self.extract_features(x)\n        f = torch.cat([f, cell_type], dim=1)\n        return self.head(f)\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# TRAINING\n# ══════════════════════════════════════════════════════════════════════════════\ndef train():\n    torch.manual_seed(SEED)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(SEED)\n    random.seed(SEED)\n    np.random.seed(SEED)\n\n    # Load data\n    df = pd.read_csv(TRAIN_CSV)\n    sirna_to_idx, idx_to_sirna = build_sirna_mapping(df)\n\n    # Stratified split\n    skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=SEED)\n    train_idx, val_idx = list(skf.split(df, df[\"sirna\"]))[0]\n    train_df = df.iloc[train_idx].reset_index(drop=True)\n    val_df = df.iloc[val_idx].reset_index(drop=True)\n    print(f\"Split: train={len(train_df)}, val={len(val_df)}\")\n\n    # Datasets — use BOTH sites for training (more data)\n    train_ds = RxRxDataset(train_df, TRAIN_DIR, sirna_to_idx,\n                           training=True, sites=(1, 2))\n    val_ds = RxRxDataset(val_df, TRAIN_DIR, sirna_to_idx,\n                         training=False, sites=(1,))\n\n    if len(train_ds) == 0:\n        print(\"ERROR: No training samples available! Check image paths.\")\n        return None, None\n\n    train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,\n                              num_workers=NUM_WORKERS, pin_memory=True,\n                              drop_last=True)\n    val_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False,\n                            num_workers=NUM_WORKERS, pin_memory=True)\n\n    # Model\n    model = RxRxModel().to(DEVICE)\n    print(f\"Params: {sum(p.numel() for p in model.parameters()):,}\")\n\n    # Loss with label smoothing\n    criterion = nn.CrossEntropyLoss(label_smoothing=LABEL_SMOOTHING)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE,\n                                  weight_decay=WEIGHT_DECAY)\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        optimizer, max_lr=LEARNING_RATE, epochs=MAX_EPOCHS,\n        steps_per_epoch=len(train_loader), pct_start=0.1,\n    )\n    scaler = torch.amp.GradScaler('cuda', enabled=USE_AMP)\n\n    best_acc = 0.0\n    best_epoch = 0\n    no_improve = 0\n    ckpt_path = os.path.join(OUTPUT_DIR, \"best_model.pth\")\n\n    for epoch in range(MAX_EPOCHS):\n        t0 = time.time()\n\n        # ── Train ─────────────────────────────────────────────────────────\n        model.train()\n        total_loss, n_batches = 0.0, 0\n        for imgs, cts, labels in tqdm(train_loader,\n                                       desc=f\"Ep{epoch} train\", leave=False):\n            imgs = imgs.to(DEVICE, non_blocking=True)\n            cts = cts.to(DEVICE, non_blocking=True)\n            labels = labels.to(DEVICE, non_blocking=True)\n\n            optimizer.zero_grad(set_to_none=True)\n            with torch.amp.autocast('cuda', enabled=USE_AMP):\n                logits = model(imgs, cts)\n                loss = criterion(logits, labels)\n\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm_(model.parameters(), GRADIENT_CLIP)\n            scaler.step(optimizer)\n            scaler.update()\n            scheduler.step()\n\n            total_loss += loss.item()\n            n_batches += 1\n\n        avg_loss = total_loss / max(n_batches, 1)\n\n        # ── Validate ──────────────────────────────────────────────────────\n        model.eval()\n        correct, total = 0, 0\n        with torch.no_grad():\n            for imgs, cts, labels in tqdm(val_loader,\n                                           desc=f\"Ep{epoch} val\", leave=False):\n                imgs = imgs.to(DEVICE, non_blocking=True)\n                cts = cts.to(DEVICE, non_blocking=True)\n                labels = labels.to(DEVICE, non_blocking=True)\n\n                with torch.amp.autocast('cuda', enabled=USE_AMP):\n                    logits = model(imgs, cts)\n                preds = logits.argmax(dim=1)\n                correct += (preds == labels).sum().item()\n                total += labels.size(0)\n\n        val_acc = correct / max(total, 1)\n        lr = optimizer.param_groups[0][\"lr\"]\n        elapsed = time.time() - t0\n\n        print(f\"Ep {epoch:2d} | loss={avg_loss:.4f} | \"\n              f\"val_acc={val_acc:.4f} | lr={lr:.2e} | {elapsed:.0f}s\")\n\n        if val_acc > best_acc:\n            best_acc = val_acc\n            best_epoch = epoch\n            no_improve = 0\n            torch.save({\n                \"model_state_dict\": model.state_dict(),\n                \"idx_to_sirna\": idx_to_sirna,\n                \"val_acc\": val_acc,\n                \"epoch\": epoch,\n            }, ckpt_path)\n            print(f\"  → Best model saved (acc={val_acc:.4f})\")\n        else:\n            no_improve += 1\n            if no_improve >= EARLY_STOP_PATIENCE:\n                print(f\"Early stop at ep {epoch}. \"\n                      f\"Best={best_acc:.4f} @ ep {best_epoch}\")\n                break\n\n    print(f\"\\nTraining done. Best val_acc = {best_acc:.4f}\")\n    return ckpt_path, idx_to_sirna\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# INFERENCE — both sites + 4 TTA variants\n# ══════════════════════════════════════════════════════════════════════════════\ndef tta_variants(img):\n    \"\"\"4 TTA variants: original, H-flip, V-flip, 180° rotation.\"\"\"\n    return [img,\n            torch.flip(img, [2]),\n            torch.flip(img, [1]),\n            torch.rot90(img, 2, [1, 2])]\n\ndef inference(ckpt_path, idx_to_sirna):\n    if ckpt_path is None:\n        print(\"No checkpoint, skipping inference.\")\n        return None\n\n    print(\"\\n\" + \"=\" * 60)\n    print(\"Inference (both sites + 4× TTA)...\")\n\n    ckpt = torch.load(ckpt_path, map_location=DEVICE, weights_only=False)\n    if idx_to_sirna is None:\n        idx_to_sirna = ckpt[\"idx_to_sirna\"]\n\n    model = RxRxModel().to(DEVICE)\n    model.load_state_dict(ckpt[\"model_state_dict\"])\n    model.eval()\n\n    test_df = pd.read_csv(TEST_CSV)\n    sirna_to_idx, _ = build_sirna_mapping(pd.read_csv(TRAIN_CSV))\n\n    # Both sites for inference\n    test_ds = RxRxDataset(test_df, TEST_DIR, sirna_to_idx,\n                          training=False, sites=(1, 2))\n    test_loader = DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=False,\n                             num_workers=NUM_WORKERS, pin_memory=True)\n\n    id_probs = defaultdict(lambda: np.zeros(NUM_CLASSES))\n    id_counts = defaultdict(int)\n\n    with torch.no_grad():\n        for batch_idx, (imgs, cts, _) in enumerate(\n                tqdm(test_loader, desc=\"Inference\")):\n            start_idx = batch_idx * BATCH_SIZE\n            for i in range(imgs.size(0)):\n                si = start_idx + i\n                if si >= len(test_ds.samples):\n                    break\n                id_code = test_ds.samples[si][\"id_code\"]\n\n                ct = cts[i].unsqueeze(0).to(DEVICE)\n                variants = tta_variants(imgs[i])\n                batch_v = torch.stack(variants).to(DEVICE)\n                ct_exp = ct.expand(len(variants), -1)\n\n                with torch.amp.autocast('cuda', enabled=USE_AMP):\n                    logits = model(batch_v, ct_exp)\n                probs = F.softmax(logits, dim=1).mean(0).cpu().numpy()\n\n                id_probs[id_code] += probs\n                id_counts[id_code] += 1\n\n    # Build submission\n    results = []\n    for id_code in sorted(id_probs.keys()):\n        avg = id_probs[id_code] / id_counts[id_code]\n        pred_idx = int(np.argmax(avg))\n        pred_sirna = idx_to_sirna[pred_idx]\n        results.append({\"id_code\": id_code, \"sirna\": pred_sirna})\n\n    sub_df = pd.DataFrame(results)\n    sub_path = os.path.join(OUTPUT_DIR, \"submission.csv\")\n    sub_df.to_csv(sub_path, index=False)\n    print(f\"\\nSubmission saved: {sub_path} ({len(sub_df)} rows)\")\n    print(sub_df.head(10))\n    return sub_path\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# RUN\n# ══════════════════════════════════════════════════════════════════════════════\nif __name__ == \"__main__\":\n    ckpt_path, idx_to_sirna = train()\n    sub_path = inference(ckpt_path, idx_to_sirna)\n    if sub_path:\n        print(f\"\\nDone! Submit {sub_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-27T19:44:32.801341Z","iopub.execute_input":"2026-04-27T19:44:32.802053Z"}},"outputs":[],"execution_count":null}]}