{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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"},"accelerator":"GPU","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":"markdown","source":"# Recursion Cellular Image Classification — Final Kaggle Notebook\n\n- uses all 6 microscopy channels\n- uses both sites `s1` and `s2`\n- averages validation/prediction probabilities across sites\n- uses a CNN backbone modified for 6-channel input\n- uses experiment-aware validation split\n- reports Accuracy, Top-5 Accuracy, Macro F1, Weighted F1, and Micro F1\n","metadata":{}},{"cell_type":"markdown","source":"## 1. Imports and config","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nimport gc\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast, GradScaler\n\nimport torchvision.models as models\nfrom torchvision.models import EfficientNet_B0_Weights\n\nfrom sklearn.model_selection import GroupKFold\nfrom sklearn.metrics import accuracy_score, f1_score\nfrom sklearn.preprocessing import LabelEncoder\n\nfrom tqdm.auto import tqdm\n\npd.set_option(\"display.max_columns\", 100)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:34.366349Z","iopub.execute_input":"2026-04-29T05:49:34.366581Z","iopub.status.idle":"2026-04-29T05:49:49.951535Z","shell.execute_reply.started":"2026-04-29T05:49:34.366555Z","shell.execute_reply":"2026-04-29T05:49:49.950366Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CFG:\n    seed = 42\n\n    data_path = \"/kaggle/input/competitions/recursion-cellular-image-classification\"\n\n    fallback_data_path = \"/kaggle/input/recursion-cellular-image-classification\"\n\n    img_size = 384\n    batch_size = 16\n    epochs = 10\n    lr = 3e-4\n    weight_decay = 1e-4\n    num_workers = 2\n\n    n_folds = 5\n    fold = 0\n\n    use_pretrained = True\n    use_amp = True\n    tta = 4\n\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\nprint(\"Device:\", CFG.device)\nif CFG.device == \"cuda\":\n    print(\"GPU:\", torch.cuda.get_device_name(0))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:49.953949Z","iopub.execute_input":"2026-04-29T05:49:49.954490Z","iopub.status.idle":"2026-04-29T05:49:49.961542Z","shell.execute_reply.started":"2026-04-29T05:49:49.954455Z","shell.execute_reply":"2026-04-29T05:49:49.960632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(CFG.seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:49.962619Z","iopub.execute_input":"2026-04-29T05:49:49.962946Z","iopub.status.idle":"2026-04-29T05:49:49.991877Z","shell.execute_reply.started":"2026-04-29T05:49:49.962902Z","shell.execute_reply":"2026-04-29T05:49:49.990809Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Data path setup\n\nThis keeps your original `DATA_PATH` style but also checks the common Kaggle path.","metadata":{}},{"cell_type":"code","source":"def resolve_data_path():\n    if os.path.exists(CFG.data_path):\n        return CFG.data_path\n    if os.path.exists(CFG.fallback_data_path):\n        return CFG.fallback_data_path\n    raise FileNotFoundError(\n        \"Could not find competition data. Add the Recursion Cellular Image Classification dataset to this notebook.\"\n    )\n\nDATA_PATH = resolve_data_path()\nprint(\"DATA_PATH:\", DATA_PATH)\nprint(\"Files:\", os.listdir(DATA_PATH)[:20])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:49.992957Z","iopub.execute_input":"2026-04-29T05:49:49.993561Z","iopub.status.idle":"2026-04-29T05:49:50.001467Z","shell.execute_reply.started":"2026-04-29T05:49:49.993520Z","shell.execute_reply":"2026-04-29T05:49:50.000072Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Load CSV files","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv(os.path.join(DATA_PATH, \"train.csv\"))\ntest = pd.read_csv(os.path.join(DATA_PATH, \"test.csv\"))\nsample = pd.read_csv(os.path.join(DATA_PATH, \"sample_submission.csv\"))\n\nprint(\"Train shape:\", train.shape)\nprint(\"Test shape:\", test.shape)\nprint(\"Sample shape:\", sample.shape)\nprint(\"Train columns:\", train.columns.tolist())\nprint(\"Test columns:\", test.columns.tolist())\n\ndisplay(train.head())\ndisplay(test.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:50.003156Z","iopub.execute_input":"2026-04-29T05:49:50.003465Z","iopub.status.idle":"2026-04-29T05:49:50.178093Z","shell.execute_reply.started":"2026-04-29T05:49:50.003426Z","shell.execute_reply":"2026-04-29T05:49:50.177009Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Encode labels and create site rows\n\nThe CSV does **not** contain a `site` column. Images exist as `s1` and `s2`, so we manually create rows for both sites.","metadata":{}},{"cell_type":"code","source":"label_encoder = LabelEncoder()\ntrain[\"label\"] = label_encoder.fit_transform(train[\"sirna\"])\nNUM_CLASSES = train[\"label\"].nunique()\n\nprint(\"Number of classes:\", NUM_CLASSES)\n\ntrain_s1 = train.copy()\ntrain_s1[\"site\"] = 1\ntrain_s2 = train.copy()\ntrain_s2[\"site\"] = 2\ntrain_site = pd.concat([train_s1, train_s2], ignore_index=True)\n\ntest_s1 = test.copy()\ntest_s1[\"site\"] = 1\ntest_s2 = test.copy()\ntest_s2[\"site\"] = 2\ntest_site = pd.concat([test_s1, test_s2], ignore_index=True)\n\nprint(\"Train site rows:\", train_site.shape)\nprint(\"Test site rows:\", test_site.shape)\ntrain_site.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:50.180322Z","iopub.execute_input":"2026-04-29T05:49:50.180719Z","iopub.status.idle":"2026-04-29T05:49:50.231636Z","shell.execute_reply.started":"2026-04-29T05:49:50.180688Z","shell.execute_reply":"2026-04-29T05:49:50.230752Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Path helper and image check\n\nif a channel image is missing, it uses a black image","metadata":{}},{"cell_type":"code","source":"def get_image_path(row, mode, site=None, channel=1):\n    if site is None:\n        site = int(row[\"site\"])\n    else:\n        site = int(site)\n\n    return os.path.join(\n        DATA_PATH,\n        mode,\n        row[\"experiment\"],\n        f\"Plate{int(row['plate'])}\",\n        f\"{row['well']}_s{site}_w{channel}.png\"\n    )\n\nexample_row = train.iloc[0]\nfor site in [1, 2]:\n    example_path = get_image_path(example_row, \"train\", site=site, channel=1)\n    print(f\"site {site} example path:\", example_path)\n    print(\"exists:\", os.path.exists(example_path))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:50.232853Z","iopub.execute_input":"2026-04-29T05:49:50.233189Z","iopub.status.idle":"2026-04-29T05:49:50.252703Z","shell.execute_reply.started":"2026-04-29T05:49:50.233157Z","shell.execute_reply":"2026-04-29T05:49:50.251817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_6_channel_image(row, mode):\n    site = int(row[\"site\"])\n    channels = []\n\n    for channel in range(1, 7):\n        path = get_image_path(row, mode, site=site, channel=channel)\n\n        if os.path.exists(path):\n            img = Image.open(path).convert(\"L\")\n        else:\n            # Same idea as your attached notebook: do not crash on missing paths.\n            img = Image.new(\"L\", (512, 512), 0)\n\n        img = img.resize((CFG.img_size, CFG.img_size))\n        img = np.asarray(img, dtype=np.float32) / 255.0\n        channels.append(img)\n\n    image = np.stack(channels, axis=0)\n    return image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:50.254020Z","iopub.execute_input":"2026-04-29T05:49:50.254454Z","iopub.status.idle":"2026-04-29T05:49:50.260745Z","shell.execute_reply.started":"2026-04-29T05:49:50.254421Z","shell.execute_reply":"2026-04-29T05:49:50.259653Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Experiment-aware validation split\n\nWe split by experiment to reduce leakage between train and validation.","metadata":{}},{"cell_type":"code","source":"gkf = GroupKFold(n_splits=CFG.n_folds)\ntrain[\"fold\"] = -1\n\nfor fold, (_, val_idx) in enumerate(gkf.split(train, train[\"label\"], groups=train[\"experiment\"])):\n    train.loc[val_idx, \"fold\"] = fold\n\ntrain_site = train_site.merge(train[[\"id_code\", \"fold\"]], on=\"id_code\", how=\"left\")\n\ntrn_df = train_site[train_site[\"fold\"] != CFG.fold].reset_index(drop=True)\nval_df = train_site[train_site[\"fold\"] == CFG.fold].reset_index(drop=True)\n\nprint(\"Train rows:\", trn_df.shape)\nprint(\"Val rows:\", val_df.shape)\nprint(\"Unique val samples:\", val_df[\"id_code\"].nunique())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:50.262309Z","iopub.execute_input":"2026-04-29T05:49:50.262644Z","iopub.status.idle":"2026-04-29T05:49:50.397841Z","shell.execute_reply.started":"2026-04-29T05:49:50.262612Z","shell.execute_reply":"2026-04-29T05:49:50.396949Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Dataset and dataloaders","metadata":{}},{"cell_type":"code","source":"class RecursionDataset(Dataset):\n    def __init__(self, df, mode=\"train\", augment=False):\n        self.df = df.reset_index(drop=True)\n        self.mode = mode\n        self.augment = augment\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image = load_6_channel_image(row, self.mode)\n\n        if self.augment:\n            if random.random() < 0.5:\n                image = image[:, :, ::-1].copy()\n            if random.random() < 0.5:\n                image = image[:, ::-1, :].copy()\n            if random.random() < 0.5:\n                k = random.randint(0, 3)\n                image = np.rot90(image, k, axes=(1, 2)).copy()\n\n        image = torch.tensor(image, dtype=torch.float32)\n\n        if self.mode == \"test\":\n            return image, row[\"id_code\"]\n\n        label = torch.tensor(int(row[\"label\"]), dtype=torch.long)\n        return image, label, row[\"id_code\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:50.398989Z","iopub.execute_input":"2026-04-29T05:49:50.399370Z","iopub.status.idle":"2026-04-29T05:49:50.407981Z","shell.execute_reply.started":"2026-04-29T05:49:50.399334Z","shell.execute_reply":"2026-04-29T05:49:50.407135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds = RecursionDataset(trn_df, mode=\"train\", augment=True)\nval_ds = RecursionDataset(val_df, mode=\"train\", augment=False)\n\ntrain_loader = DataLoader(\n    train_ds,\n    batch_size=CFG.batch_size,\n    shuffle=True,\n    num_workers=CFG.num_workers,\n    pin_memory=True,\n    drop_last=True,\n)\n\nval_loader = DataLoader(\n    val_ds,\n    batch_size=CFG.batch_size * 2,\n    shuffle=False,\n    num_workers=CFG.num_workers,\n    pin_memory=True,\n)\n\nprint(\"Batches:\", len(train_loader), len(val_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:50.409319Z","iopub.execute_input":"2026-04-29T05:49:50.409686Z","iopub.status.idle":"2026-04-29T05:49:50.438150Z","shell.execute_reply.started":"2026-04-29T05:49:50.409641Z","shell.execute_reply":"2026-04-29T05:49:50.437209Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Model\n\nEfficientNet-B0 is adapted from 3-channel RGB input to 6 microscopy channels.","metadata":{}},{"cell_type":"code","source":"class RecursionModel(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n\n        weights = None\n        if CFG.use_pretrained:\n            try:\n                weights = EfficientNet_B0_Weights.IMAGENET1K_V1\n            except Exception:\n                weights = None\n\n        self.model = models.efficientnet_b0(weights=weights)\n\n        old_conv = self.model.features[0][0]\n        new_conv = nn.Conv2d(\n            6,\n            old_conv.out_channels,\n            kernel_size=old_conv.kernel_size,\n            stride=old_conv.stride,\n            padding=old_conv.padding,\n            bias=False,\n        )\n\n        with torch.no_grad():\n            if old_conv.weight.shape[1] == 3:\n                new_conv.weight[:, :3] = old_conv.weight\n                new_conv.weight[:, 3:] = old_conv.weight\n                new_conv.weight *= 0.5\n            else:\n                nn.init.kaiming_normal_(new_conv.weight, mode=\"fan_out\", nonlinearity=\"relu\")\n\n        self.model.features[0][0] = new_conv\n        in_features = self.model.classifier[1].in_features\n        self.model.classifier[1] = nn.Linear(in_features, num_classes)\n\n    def forward(self, x):\n        return self.model(x)\n\nmodel = RecursionModel(NUM_CLASSES).to(CFG.device)\nprint(\"Model ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:50.439309Z","iopub.execute_input":"2026-04-29T05:49:50.439698Z","iopub.status.idle":"2026-04-29T05:49:50.992073Z","shell.execute_reply.started":"2026-04-29T05:49:50.439651Z","shell.execute_reply":"2026-04-29T05:49:50.990928Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 9. Loss, optimizer, scheduler","metadata":{}},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG.epochs)\nscaler = GradScaler(\"cuda\", enabled=(CFG.device == \"cuda\" and CFG.use_amp))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:50.993229Z","iopub.execute_input":"2026-04-29T05:49:50.993525Z","iopub.status.idle":"2026-04-29T05:49:51.000280Z","shell.execute_reply.started":"2026-04-29T05:49:50.993495Z","shell.execute_reply":"2026-04-29T05:49:50.999295Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Metrics helpers\n\nValidation metrics are computed after averaging predictions for `s1` and `s2` of the same `id_code`.","metadata":{}},{"cell_type":"code","source":"def topk_accuracy_from_probs(y_true, probs, k=5):\n    topk = np.argsort(probs, axis=1)[:, -k:]\n    return np.mean([y in topk[i] for i, y in enumerate(y_true)])\n\n\ndef compute_metrics(y_true, probs):\n    preds = probs.argmax(axis=1)\n    return {\n        \"accuracy\": accuracy_score(y_true, preds),\n        \"top5_accuracy\": topk_accuracy_from_probs(y_true, probs, k=5),\n        \"macro_f1\": f1_score(y_true, preds, average=\"macro\", zero_division=0),\n        \"weighted_f1\": f1_score(y_true, preds, average=\"weighted\", zero_division=0),\n        \"micro_f1\": f1_score(y_true, preds, average=\"micro\", zero_division=0),\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:51.002377Z","iopub.execute_input":"2026-04-29T05:49:51.002712Z","iopub.status.idle":"2026-04-29T05:49:51.028282Z","shell.execute_reply.started":"2026-04-29T05:49:51.002680Z","shell.execute_reply":"2026-04-29T05:49:51.027079Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 11. Training and validation functions","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    for images, labels, _ in tqdm(loader, leave=False):\n        images = images.to(CFG.device, non_blocking=True)\n        labels = labels.to(CFG.device, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        with autocast(\"cuda\", enabled=(CFG.device == \"cuda\" and CFG.use_amp)):\n            logits = model(images)\n            loss = criterion(logits, labels)\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item() * images.size(0)\n        preds = logits.argmax(1)\n        correct += (preds == labels).sum().item()\n        total += labels.size(0)\n\n    return running_loss / total, correct / total\n\n\n@torch.no_grad()\ndef validate_site_rows(model, loader):\n    model.eval()\n    total_loss = 0.0\n    total_rows = 0\n\n    all_ids = []\n    all_labels = []\n    all_probs = []\n\n    for images, labels, ids in tqdm(loader, leave=False):\n        images = images.to(CFG.device, non_blocking=True)\n        labels = labels.to(CFG.device, non_blocking=True)\n\n        with autocast(\"cuda\", enabled=(CFG.device == \"cuda\" and CFG.use_amp)):\n            logits = model(images)\n            loss = criterion(logits, labels)\n            probs = torch.softmax(logits, dim=1)\n\n        total_loss += loss.item() * images.size(0)\n        total_rows += images.size(0)\n\n        all_ids.extend(list(ids))\n        all_labels.extend(labels.cpu().numpy())\n        all_probs.append(probs.cpu().numpy())\n\n    probs = np.concatenate(all_probs, axis=0)\n\n    pred_df = pd.DataFrame({\"id_code\": all_ids, \"label\": all_labels})\n    prob_df = pd.DataFrame(probs, columns=[f\"p{i}\" for i in range(NUM_CLASSES)])\n    pred_df = pd.concat([pred_df, prob_df], axis=1)\n\n    grouped_probs = pred_df.groupby(\"id_code\")[[f\"p{i}\" for i in range(NUM_CLASSES)]].mean().values\n    grouped_labels = pred_df.groupby(\"id_code\")[\"label\"].first().values\n\n    metrics = compute_metrics(grouped_labels, grouped_probs)\n    metrics[\"loss\"] = total_loss / total_rows\n    return metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:51.029516Z","iopub.execute_input":"2026-04-29T05:49:51.029849Z","iopub.status.idle":"2026-04-29T05:49:51.058622Z","shell.execute_reply.started":"2026-04-29T05:49:51.029807Z","shell.execute_reply":"2026-04-29T05:49:51.057519Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 12. Train model","metadata":{}},{"cell_type":"code","source":"best_acc = -1\nhistory = []\n\nfor epoch in range(1, CFG.epochs + 1):\n    train_loss, train_acc = train_one_epoch(model, train_loader)\n    val_metrics = validate_site_rows(model, val_loader)\n    scheduler.step()\n\n    row = {\n        \"epoch\": epoch,\n        \"train_loss\": train_loss,\n        \"train_accuracy\": train_acc,\n        **{f\"val_{k}\": v for k, v in val_metrics.items()},\n    }\n    history.append(row)\n\n    print(\n        f\"Epoch {epoch}/{CFG.epochs} | \"\n        f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | \"\n        f\"Val Loss: {val_metrics['loss']:.4f} | Val Acc: {val_metrics['accuracy']:.4f} | \"\n        f\"Val Top5: {val_metrics['top5_accuracy']:.4f} | \"\n        f\"Val Macro F1: {val_metrics['macro_f1']:.4f}\"\n    )\n\n    if val_metrics[\"accuracy\"] > best_acc:\n        best_acc = val_metrics[\"accuracy\"]\n        torch.save(model.state_dict(), \"/kaggle/working/best_model.pth\")\n        print(\"Best model saved\")\n\nhistory_df = pd.DataFrame(history)\ndisplay(history_df)\nhistory_df.to_csv(\"/kaggle/working/training_history.csv\", index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T05:49:51.059673Z","iopub.execute_input":"2026-04-29T05:49:51.060023Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 13. Plot training history","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nplt.figure(figsize=(8, 5))\nplt.plot(history_df[\"epoch\"], history_df[\"train_accuracy\"], label=\"train accuracy\")\nplt.plot(history_df[\"epoch\"], history_df[\"val_accuracy\"], label=\"val accuracy\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy\")\nplt.legend()\nplt.grid(True)\nplt.show()\n\nplt.figure(figsize=(8, 5))\nplt.plot(history_df[\"epoch\"], history_df[\"train_loss\"], label=\"train loss\")\nplt.plot(history_df[\"epoch\"], history_df[\"val_loss\"], label=\"val loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.legend()\nplt.grid(True)\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 14. TTA prediction\n\nPredictions are averaged across TTA transforms and across `s1`/`s2` sites.","metadata":{}},{"cell_type":"code","source":"test_ds = RecursionDataset(test_site, mode=\"test\", augment=False)\ntest_loader = DataLoader(\n    test_ds,\n    batch_size=CFG.batch_size * 2,\n    shuffle=False,\n    num_workers=CFG.num_workers,\n    pin_memory=True,\n)\n\n\ndef apply_tta(images, tta_id):\n    if tta_id == 0:\n        return images\n    if tta_id == 1:\n        return torch.flip(images, dims=[3])\n    if tta_id == 2:\n        return torch.flip(images, dims=[2])\n    if tta_id == 3:\n        return torch.rot90(images, k=1, dims=[2, 3])\n    return images\n\n\n@torch.no_grad()\ndef predict(model, loader):\n    model.eval()\n    all_ids = []\n    all_probs = []\n\n    for images, ids in tqdm(loader):\n        images = images.to(CFG.device, non_blocking=True)\n        probs_sum = 0\n\n        for tta_id in range(CFG.tta):\n            aug_images = apply_tta(images, tta_id)\n            with autocast(\"cuda\", enabled=(CFG.device == \"cuda\" and CFG.use_amp)):\n                logits = model(aug_images)\n                probs_sum += torch.softmax(logits, dim=1)\n\n        probs = probs_sum / CFG.tta\n        all_ids.extend(list(ids))\n        all_probs.append(probs.cpu().numpy())\n\n    probs = np.concatenate(all_probs, axis=0)\n    pred_df = pd.DataFrame({\"id_code\": all_ids})\n    prob_df = pd.DataFrame(probs, columns=[f\"p{i}\" for i in range(NUM_CLASSES)])\n    pred_df = pd.concat([pred_df, prob_df], axis=1)\n\n    grouped = pred_df.groupby(\"id_code\")[[f\"p{i}\" for i in range(NUM_CLASSES)]].mean()\n    return grouped\n\nmodel.load_state_dict(torch.load(\"/kaggle/working/best_model.pth\", map_location=CFG.device))\ntest_probs_df = predict(model, test_loader)\ntest_probs_df.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 15. Create submission.csv","metadata":{}},{"cell_type":"code","source":"pred_labels = test_probs_df.values.argmax(axis=1)\npred_sirna = label_encoder.inverse_transform(pred_labels)\n\nsubmission = pd.DataFrame({\n    \"id_code\": test_probs_df.index,\n    \"sirna\": pred_sirna,\n})\n\n# Keep original sample order.\nsubmission = sample[[\"id_code\"]].merge(submission, on=\"id_code\", how=\"left\")\n\n# Fill missing predictions with most common class, but keep as string.\nsubmission[\"sirna\"] = submission[\"sirna\"].fillna(train[\"sirna\"].mode()[0])\n\nsubmission.to_csv(\"/kaggle/working/submission.csv\", index=False)\n\nprint(submission.shape)\ndisplay(submission.head())\nprint(submission[\"sirna\"].head())","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}