{"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":"none","dataSources":[{"sourceType":"competition","sourceId":14420,"databundleVersionId":868327,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":15577364,"datasetId":9966298,"databundleVersionId":16508828},{"sourceType":"datasetVersion","sourceId":15580803,"datasetId":9968595,"databundleVersionId":16512579}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport pandas as pd\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms\nfrom PIL import Image\nfrom sklearn.model_selection import GroupKFold\nfrom tqdm.auto import tqdm\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2026-04-01T16:25:07.261Z","iopub.execute_input":"2026-04-01T16:25:07.261739Z","iopub.status.idle":"2026-04-01T16:25:07.267445Z","shell.execute_reply.started":"2026-04-01T16:25:07.261695Z","shell.execute_reply":"2026-04-01T16:25:07.266543Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\nIMG_SIZE = 256\nBATCH_SIZE = 32\nNUM_CLASSES = 1108\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T16:25:09.269757Z","iopub.execute_input":"2026-04-01T16:25:09.270086Z","iopub.status.idle":"2026-04-01T16:25:09.274804Z","shell.execute_reply.started":"2026-04-01T16:25:09.270057Z","shell.execute_reply":"2026-04-01T16:25:09.273782Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Preprocessing","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(f'{DATA_DIR}/train.csv')\ntest_df = pd.read_csv(f'{DATA_DIR}/test.csv')\ntrain_df['group'] = train_df['experiment']\nif train_df['sirna'].dtype == 'O':\n    train_df['sirna'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\nprint(f\"Unique siRNA values: {train_df['sirna'].nunique()}\")\nprint(f\"Sample values: {train_df['sirna'].head().values}\")\nprint(f\"Total training samples: {len(train_df)}\")\nprint(f\"Unique experiments: {train_df['experiment'].nunique()}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T16:25:17.414278Z","iopub.execute_input":"2026-04-01T16:25:17.415022Z","iopub.status.idle":"2026-04-01T16:25:17.49554Z","shell.execute_reply.started":"2026-04-01T16:25:17.414986Z","shell.execute_reply":"2026-04-01T16:25:17.494844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RxRxDataset(Dataset):\n    def __init__(self, df, mode='train', transform=None):\n        self.df = df\n        self.mode = mode\n        self.transform = transform\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_id = f\"{row['experiment']}/Plate{row['plate']}/{row['well']}_s1\"\n        \n        channels = []\n        for i in range(1, 7):\n            path = f\"{DATA_DIR}/{self.mode}/{img_id}_w{i}.png\"\n            img = Image.open(path)\n            channels.append(np.array(img))\n        \n        x = np.stack(channels, axis=-1) \n        \n        if self.transform:\n            x = self.transform(x)\n        else:\n            x = torch.from_numpy(x.transpose(2, 0, 1)).float() / 255.0\n            \n        if self.mode == 'train':\n            return x, row['sirna']\n        return x, row['id_code']\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T16:25:19.333718Z","iopub.execute_input":"2026-04-01T16:25:19.33432Z","iopub.status.idle":"2026-04-01T16:25:19.340624Z","shell.execute_reply.started":"2026-04-01T16:25:19.334286Z","shell.execute_reply":"2026-04-01T16:25:19.339816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"min_val = train_df['sirna'].min()\nmax_val = train_df['sirna'].max()\nunique_count = train_df['sirna'].nunique()\nprint(f\"Min siRNA ID: {min_val}\")\nprint(f\"Max siRNA ID: {max_val}\")\nprint(f\"Unique siRNA count: {unique_count}\")\nif max_val >= 1108:\n    print(\"🚨 DANGER: Max ID is too high for a model with 1108 outputs.\")\nif min_val < 0:\n    print(\"🚨 DANGER: Negative labels detected.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T16:25:24.569579Z","iopub.execute_input":"2026-04-01T16:25:24.570359Z","iopub.status.idle":"2026-04-01T16:25:24.576625Z","shell.execute_reply.started":"2026-04-01T16:25:24.570324Z","shell.execute_reply":"2026-04-01T16:25:24.575949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_CLASSES = 1108 \ndef get_model(num_classes=NUM_CLASSES):\n    model = models.resnet18(pretrained=True)\n    \n    w = model.conv1.weight\n    model.conv1 = nn.Conv2d(6, 64, kernel_size=7, stride=2, padding=3, bias=False)\n    with torch.no_grad():\n        model.conv1.weight[:, :3] = w\n        model.conv1.weight[:, 3:] = w\n        \n    model.fc = nn.Linear(model.fc.in_features, num_classes)\n    return model\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = get_model(num_classes=NUM_CLASSES).to(device)\nprint(\"Model re-initialized and ready.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T16:25:34.644669Z","iopub.execute_input":"2026-04-01T16:25:34.645287Z","iopub.status.idle":"2026-04-01T16:25:35.351183Z","shell.execute_reply.started":"2026-04-01T16:25:34.645238Z","shell.execute_reply":"2026-04-01T16:25:35.350505Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.preprocessing import LabelEncoder\nencoder = LabelEncoder()\ntrain_df['sirna'] = encoder.fit_transform(train_df['sirna'])\nnew_max = train_df['sirna'].max()\nnew_unique = train_df['sirna'].nunique()\nprint(f\"New Max siRNA ID: {new_max}\")\nprint(f\"New Unique count: {new_unique}\")\nnp.save('sirna_classes.npy', encoder.classes_)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T16:25:25.574269Z","iopub.execute_input":"2026-04-01T16:25:25.575121Z","iopub.status.idle":"2026-04-01T16:25:25.584597Z","shell.execute_reply.started":"2026-04-01T16:25:25.575078Z","shell.execute_reply":"2026-04-01T16:25:25.583774Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss(label_smoothing=0.1)\noptimizer = optim.Adam(model.parameters(), lr=3e-4)\ndef train_one_epoch(model, loader, optimizer, criterion):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, total=len(loader))\n    for inputs, labels in pbar:\n        inputs, labels = inputs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_description(f\"Loss: {running_loss/len(loader):.4f} Acc: {100.*correct/total:.2f}%\")\ntest_ds = RxRxDataset(train_df.head(100))\ntest_loader = DataLoader(test_ds, batch_size=8, shuffle=True)\ntrain_one_epoch(model, test_loader, optimizer, criterion)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T16:25:37.797788Z","iopub.execute_input":"2026-04-01T16:25:37.798428Z","iopub.status.idle":"2026-04-01T16:25:48.695806Z","shell.execute_reply.started":"2026-04-01T16:25:37.798396Z","shell.execute_reply":"2026-04-01T16:25:48.695164Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models, transforms\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import GroupKFold\nDATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\nBATCH_SIZE = 16\nNUM_EPOCHS = 10\nLR = 3e-4\nIMG_SIZE = 256\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Loading and cleaning data...\")\ntrain_df = pd.read_csv(f'{DATA_DIR}/train.csv')\ntest_df = pd.read_csv(f'{DATA_DIR}/test.csv')\nif train_df['sirna'].dtype == 'O':\n    train_df['sirna'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\nle = LabelEncoder()\ntrain_df['sirna'] = le.fit_transform(train_df['sirna'])\nNUM_CLASSES = len(le.classes_)\nnp.save('label_mapping.npy', le.classes_)\ngkf = GroupKFold(n_splits=5)\ntrain_idx, val_idx = next(gkf.split(train_df, groups=train_df['experiment']))\ntrain_split = train_df.iloc[train_idx].reset_index(drop=True)\nval_split = train_df.iloc[val_idx].reset_index(drop=True)\nprint(f\"Classes: {NUM_CLASSES} | Train: {len(train_split)} | Val: {len(val_split)}\")\nclass RxRxDataset(Dataset):\n    def __init__(self, df, mode='train', transform=None):\n        self.df = df\n        self.mode = mode\n        self.transform = transform\n    def __len__(self):\n        return len(self.df)\n    def _get_img_path(self, experiment, plate, well, site, channel):\n        return os.path.join(DATA_DIR, self.mode, experiment, f\"Plate{plate}\", f\"{well}_s{site}_w{channel}.png\")\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        \n        channels = []\n        for i in range(1, 7):\n            path = self._get_img_path(row['experiment'], row['plate'], row['well'], 1, i)\n            \n            if not os.path.exists(path):\n                img = Image.new('L', (IMG_SIZE, IMG_SIZE))\n            else:\n                img = Image.open(path).resize((IMG_SIZE, IMG_SIZE))\n                \n            channels.append(np.array(img))\n        \n        x = np.stack(channels, axis=-1)\n        \n        x = torch.from_numpy(x.transpose(2, 0, 1)).float() / 255.0\n        \n        if self.mode == 'train':\n            return x, row['sirna']\n        return x, row['id_code']\ndef get_model():\n    model = models.resnet50(pretrained=True)\n    \n    original_conv = model.conv1\n    model.conv1 = nn.Conv2d(6, 64, kernel_size=7, stride=2, padding=3, bias=False)\n    \n    with torch.no_grad():\n        model.conv1.weight[:, :3] = original_conv.weight\n        model.conv1.weight[:, 3:] = original_conv.weight\n        \n    model.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)\n    return model.to(DEVICE)\ntrain_loader = DataLoader(RxRxDataset(train_split), batch_size=BATCH_SIZE, shuffle=True, num_workers=2)\nval_loader = DataLoader(RxRxDataset(val_split), batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\nmodel = get_model()\noptimizer = optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\ncriterion = nn.CrossEntropyLoss(label_smoothing=0.1)\nscheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=len(train_loader)*NUM_EPOCHS)\nbest_acc = 0.0\nfor epoch in range(NUM_EPOCHS):\n    model.train()\n    train_loss = 0\n    train_correct = 0\n    t_total = 0\n    \n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1} Train\")\n    for inputs, labels in pbar:\n        inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n        \n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n        \n        train_loss += loss.item()\n        _, pred = outputs.max(1)\n        train_correct += pred.eq(labels).sum().item()\n        t_total += labels.size(0)\n        pbar.set_description(f\"Loss: {train_loss/len(train_loader):.4f} Acc: {100.*train_correct/t_total:.2f}%\")\n    model.eval()\n    val_correct = 0\n    v_total = 0\n    with torch.no_grad():\n        for inputs, labels in tqdm(val_loader, desc=f\"Epoch {epoch+1} Val\"):\n            inputs, labels = inputs.to(DEVICE), labels.to(DEVICE)\n            outputs = model(inputs)\n            _, pred = outputs.max(1)\n            val_correct += pred.eq(labels).sum().item()\n            v_total += labels.size(0)\n    \n    val_acc = 100. * val_correct / v_total\n    print(f\"Validation Accuracy: {val_acc:.2f}%\")\n    \n    if val_acc > best_acc:\n        best_acc = val_acc\n        torch.save(model.state_dict(), 'best_model.pth')\n        print(\"⭐ New best model saved!\")\nprint(\"\\nStarting Prediction on Test Set...\")\nmodel.load_state_dict(torch.load('best_model.pth'))\nmodel.eval()\ntest_loader = DataLoader(RxRxDataset(test_df, mode='test'), batch_size=BATCH_SIZE, shuffle=False)\nid_codes, predictions = [], []\nwith torch.no_grad():\n    for inputs, ids in tqdm(test_loader, desc=\"Predicting\"):\n        inputs = inputs.to(DEVICE)\n        outputs = model(inputs)\n        _, pred = outputs.max(1)\n        \n        original_sirna = le.inverse_transform(pred.cpu().numpy())\n        predictions.extend(original_sirna)\n        id_codes.extend(ids)\nsubmission = pd.DataFrame({'id_code': id_codes, 'sirna': predictions})\nsubmission.to_csv('submission.csv', index=False)\nprint(\"Submission file 'submission.csv' generated!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-01T16:44:33.273198Z","iopub.execute_input":"2026-04-01T16:44:33.273792Z","iopub.status.idle":"2026-04-01T20:22:52.491223Z","shell.execute_reply.started":"2026-04-01T16:44:33.273761Z","shell.execute_reply":"2026-04-01T20:22:52.490029Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models\nfrom sklearn.preprocessing import LabelEncoder\nfrom sklearn.model_selection import GroupKFold\nDATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\nIMG_SIZE = 128\nBATCH_SIZE = 32\nEPOCHS = 3\nLR = 3e-4\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ntrain_df = pd.read_csv(f'{DATA_DIR}/train.csv')\ntest_df = pd.read_csv(f'{DATA_DIR}/test.csv')\nprint(\"Test rows:\", len(test_df))\nle = LabelEncoder()\ntrain_df['sirna'] = le.fit_transform(train_df['sirna'])\nNUM_CLASSES = train_df['sirna'].nunique()\ngkf = GroupKFold(n_splits=5)\ntrain_idx, val_idx = next(gkf.split(train_df, groups=train_df['experiment']))\ntrain_split = train_df.iloc[train_idx].reset_index(drop=True)\nval_split   = train_df.iloc[val_idx].reset_index(drop=True)\nclass RxRxDataset(Dataset):\n    def __init__(self, df, mode='train'):\n        self.df = df.reset_index(drop=True)\n        self.mode = mode\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        imgs = []\n        for c in range(1, 7):\n            path = os.path.join(\n                DATA_DIR,\n                self.mode,\n                row['experiment'],\n                f\"Plate{row['plate']}\",\n                f\"{row['well']}_s1_w{c}.png\"\n            )\n            if os.path.exists(path):\n                img = Image.open(path).resize((IMG_SIZE, IMG_SIZE))\n                img = np.array(img)\n            else:\n                img = np.zeros((IMG_SIZE, IMG_SIZE), dtype=np.uint8)\n            imgs.append(img)\n        x = np.stack(imgs, axis=0)\n        x = torch.tensor(x, dtype=torch.float32) / 255.0\n        if self.mode == 'test':\n            return x, row['id_code']\n        else:\n            return x, row['sirna']\ndef get_model():\n    model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)\n    pretrained_weights = model.conv1.weight.clone()\n    model.conv1 = nn.Conv2d(6, 64, kernel_size=7, stride=2, padding=3, bias=False)\n    with torch.no_grad():\n        model.conv1.weight[:, :3] = pretrained_weights\n        model.conv1.weight[:, 3:] = pretrained_weights\n    model.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)\n    return model.to(DEVICE)\nmodel = get_model()\ntrain_loader = DataLoader(RxRxDataset(train_split), batch_size=BATCH_SIZE, shuffle=True, num_workers=2)\nval_loader = DataLoader(RxRxDataset(val_split), batch_size=BATCH_SIZE, shuffle=False, num_workers=2)\noptimizer = optim.AdamW(model.parameters(), lr=LR)\ncriterion = nn.CrossEntropyLoss()\nscaler = torch.cuda.amp.GradScaler()\nbest_acc = 0\nfor epoch in range(EPOCHS):\n    model.train()\n    total, correct = 0, 0\n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}\")\n    for x, y in pbar:\n        x, y = x.to(DEVICE), y.to(DEVICE)\n        optimizer.zero_grad()\n        with torch.cuda.amp.autocast():\n            out = model(x)\n            loss = criterion(out, y)\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        pred = out.argmax(1)\n        correct += (pred == y).sum().item()\n        total += y.size(0)\n        pbar.set_description(f\"Loss {loss.item():.4f} Acc {100*correct/total:.2f}%\")\n    model.eval()\n    total, correct = 0, 0\n    with torch.no_grad():\n        for x, y in val_loader:\n            x, y = x.to(DEVICE), y.to(DEVICE)\n            out = model(x)\n            pred = out.argmax(1)\n            correct += (pred == y).sum().item()\n            total += y.size(0)\n    acc = 100 * correct / total\n    print(f\"Val Acc: {acc:.2f}%\")\n    if acc > best_acc:\n        best_acc = acc\n        torch.save(model.state_dict(), \"best.pth\")\nmodel.load_state_dict(torch.load(\"best.pth\"))\nmodel.eval()\ntest_loader = DataLoader(RxRxDataset(test_df, mode='test'),\n                         batch_size=BATCH_SIZE,\n                         shuffle=False)\nids, preds = [], []\nwith torch.no_grad():\n    for x, id_code in tqdm(test_loader):\n        x = x.to(DEVICE)\n        out = model(x)\n        pred = out.argmax(1).cpu().numpy()\n        preds.extend(pred)\n        ids.extend(id_code)\npreds = le.inverse_transform(preds)\nsubmission = pd.DataFrame({\n    \"id_code\": ids,\n    \"sirna\": preds\n})\nprint(\"Submission rows:\", len(submission))\nsubmission.to_csv(\"submission.csv\", index=False)\nprint(\"submission.csv saved!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-06T18:16:01.850348Z","iopub.execute_input":"2026-04-06T18:16:01.850673Z","iopub.status.idle":"2026-04-06T19:18:28.352015Z","shell.execute_reply.started":"2026-04-06T18:16:01.850646Z","shell.execute_reply":"2026-04-06T19:18:28.351251Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'resnet50'  \n    IMG_SIZE = 320  \n    BATCH_SIZE = 32\n    EPOCHS = 10  \n    LR = 3e-4\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    CELL_TYPES = ['HUVEC']  \n    \n    CONVERT_SIRNA = True\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train', transform=None):\n        self.df = df\n        self.data_dir = data_dir\n        self.mode = mode\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        \n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'train':\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                channels.append(img)\n            else:\n                channels.append(np.zeros((512, 512), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        \n        img = img.astype(np.float32) / 255.0\n        \n        if self.transform:\n            img = torch.from_numpy(img).permute(2, 0, 1)\n        else:\n            img = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train':\n            label = row['label']\n            return img, label\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'conv1'):\n            old_conv = self.backbone.conv1\n            self.backbone.conv1 = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.conv1.weight[:, :3] = old_conv.weight\n                self.backbone.conv1.weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.conv1.bias = old_conv.bias\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    train_df = train_df[train_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    test_df = test_df[test_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types in train: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    \n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {sirna: idx for idx, sirna in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    print(f\"Number of unique sirnas: {len(unique_sirnas)}\")\n    print(f\"Label range: 0 to {train_df['label'].max()}\")\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    \n    train_data, val_data = train_test_split(train_df, test_size=0.15, random_state=Config.SEED, stratify=train_df['label'])\n    \n    train_dataset = CellularDataset(train_data, Config.DATA_DIR, mode='train')\n    val_dataset = CellularDataset(val_data, Config.DATA_DIR, mode='train')\n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, \n                              shuffle=True, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, \n                            shuffle=False, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    \n    print(f\"Creating model: {Config.MODEL_NAME}...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_resnet50.pth')\n            print(f\"Saved best model with accuracy: {best_acc:.2f}%\")\n    \n    print(\"\\nLoading best model for inference...\")\n    model.load_state_dict(torch.load('best_resnet50.pth'))\n    \n    print(\"Generating predictions...\")\n    test_dataset = CellularDataset(test_df, Config.DATA_DIR, mode='test')\n    test_loader = DataLoader(test_dataset, batch_size=Config.BATCH_SIZE, \n                             shuffle=False, num_workers=Config.NUM_WORKERS)\n    \n    model.eval()\n    predictions = []\n    ids = []\n    \n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Inference'):\n            imgs = imgs.to(device)\n            outputs = model(imgs)\n            _, preds = outputs.max(1)\n            \n            predictions.extend(preds.cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {idx: sirna for sirna, idx in sirna_to_label.items()}\n    predictions_sirna = [label_to_sirna[pred] for pred in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions_sirna\n    })\n    submission.to_csv('submission.csv', index=False)\n    print(\"\\nSubmission saved to submission.csv\")\n    print(f\"Best validation accuracy: {best_acc:.2f}%\")\n    print(\"Sample predictions:\")\n    print(submission.head(10))\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-07T14:00:10.865322Z","iopub.execute_input":"2026-04-07T14:00:10.866104Z","iopub.status.idle":"2026-04-07T15:13:31.985593Z","shell.execute_reply.started":"2026-04-07T14:00:10.866035Z","shell.execute_reply":"2026-04-07T15:13:31.984703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'densenet121'\n    IMG_SIZE = 320  \n    BATCH_SIZE = 32\n    EPOCHS = 10  \n    LR = 3e-4\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    CELL_TYPES = ['HUVEC']  \n    \n    CONVERT_SIRNA = True\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train', transform=None):\n        self.df = df\n        self.data_dir = data_dir\n        self.mode = mode\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        \n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'train':\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                channels.append(img)\n            else:\n                channels.append(np.zeros((512, 512), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        \n        img = img.astype(np.float32) / 255.0\n        \n        if self.transform:\n            img = torch.from_numpy(img).permute(2, 0, 1)\n        else:\n            img = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train':\n            label = row['label']\n            return img, label\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'features') and hasattr(self.backbone.features, 'conv0'):\n            old_conv = self.backbone.features.conv0\n            self.backbone.features.conv0 = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.features.conv0.weight[:, :3] = old_conv.weight\n                self.backbone.features.conv0.weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.features.conv0.bias = old_conv.bias\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    train_df = train_df[train_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    test_df = test_df[test_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types in train: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    \n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {sirna: idx for idx, sirna in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    print(f\"Number of unique sirnas: {len(unique_sirnas)}\")\n    print(f\"Label range: 0 to {train_df['label'].max()}\")\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    \n    train_data, val_data = train_test_split(train_df, test_size=0.15, random_state=Config.SEED, stratify=train_df['label'])\n    \n    train_dataset = CellularDataset(train_data, Config.DATA_DIR, mode='train')\n    val_dataset = CellularDataset(val_data, Config.DATA_DIR, mode='train')\n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, \n                              shuffle=True, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, \n                            shuffle=False, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    \n    print(f\"Creating model: {Config.MODEL_NAME}...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_densenet121.pth')\n            print(f\"Saved best model with accuracy: {best_acc:.2f}%\")\n    \n    print(\"\\nLoading best model for inference...\")\n    model.load_state_dict(torch.load('best_densenet121.pth'))\n    \n    print(\"Generating predictions...\")\n    test_dataset = CellularDataset(test_df, Config.DATA_DIR, mode='test')\n    test_loader = DataLoader(test_dataset, batch_size=Config.BATCH_SIZE, \n                             shuffle=False, num_workers=Config.NUM_WORKERS)\n    \n    model.eval()\n    predictions = []\n    ids = []\n    \n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Inference'):\n            imgs = imgs.to(device)\n            outputs = model(imgs)\n            _, preds = outputs.max(1)\n            \n            predictions.extend(preds.cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {idx: sirna for sirna, idx in sirna_to_label.items()}\n    predictions_sirna = [label_to_sirna[pred] for pred in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions_sirna\n    })\n    submission.to_csv('submission.csv', index=False)\n    print(\"\\nSubmission saved to submission.csv\")\n    print(f\"Best validation accuracy: {best_acc:.2f}%\")\n    print(\"Sample predictions:\")\n    print(submission.head(10))\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-07T15:25:28.669604Z","iopub.execute_input":"2026-04-07T15:25:28.669933Z","iopub.status.idle":"2026-04-07T16:29:17.749568Z","shell.execute_reply.started":"2026-04-07T15:25:28.669896Z","shell.execute_reply":"2026-04-07T16:29:17.748478Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-cellular-test/test(1).csv'\n    \n    MODEL_NAME = 'densenet121'  \n    IMG_SIZE = 320  \n    BATCH_SIZE = 32\n    EPOCHS = 15\n    LR = 3e-4\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    CELL_TYPES = ['HUVEC']  \n    \n    AUGMENT = True\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df\n        self.data_dir = data_dir\n        self.mode = mode\n        \n        self.train_transforms = transforms.Compose([\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomVerticalFlip(p=0.5),\n            transforms.RandomRotation(degrees=15),\n        ])\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        \n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'train':\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                channels.append(img)\n            else:\n                channels.append(np.zeros((512, 512), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        \n        img = img.astype(np.float32)\n        \n        for c in range(6):\n            mean = img[:, :, c].mean()\n            std = img[:, :, c].std()\n            if std > 0:\n                img[:, :, c] = (img[:, :, c] - mean) / std\n            else:\n                img[:, :, c] = 0.0\n                \n        img_tensor = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train' and Config.AUGMENT:\n            img_tensor = self.train_transforms(img_tensor)\n        \n        if self.mode == 'train':\n            label = row['label']\n            return img_tensor, label\n        else:\n            return img_tensor, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'features') and hasattr(self.backbone.features, 'conv0'):\n            old_conv = self.backbone.features.conv0\n            self.backbone.features.conv0 = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.features.conv0.weight[:, :3] = old_conv.weight\n                self.backbone.features.conv0.weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.features.conv0.bias = old_conv.bias\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.4),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    train_df = train_df[train_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    test_df = test_df[test_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types in train: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    \n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {sirna: idx for idx, sirna in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    print(f\"Number of unique sirnas: {len(unique_sirnas)}\")\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    \n    train_data, val_data = train_test_split(train_df, test_size=0.15, random_state=Config.SEED, stratify=train_df['label'])\n    \n    train_dataset = CellularDataset(train_data, Config.DATA_DIR, mode='train')\n    val_dataset = CellularDataset(val_data, Config.DATA_DIR, mode='val') \n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, \n                              shuffle=True, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, \n                            shuffle=False, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    \n    print(f\"Creating model: {Config.MODEL_NAME} with Augmentations...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_densenet121_aug.pth')\n            print(f\"Saved best model with accuracy: {best_acc:.2f}%\")\n    \n    print(\"\\nLoading best model for inference...\")\n    model.load_state_dict(torch.load('best_densenet121_aug.pth'))\n    \n    print(\"Generating predictions...\")\n    test_dataset = CellularDataset(test_df, Config.DATA_DIR, mode='test')\n    test_loader = DataLoader(test_dataset, batch_size=Config.BATCH_SIZE, \n                             shuffle=False, num_workers=Config.NUM_WORKERS)\n    \n    model.eval()\n    predictions = []\n    ids = []\n    \n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Inference'):\n            imgs = imgs.to(device)\n            outputs = model(imgs)\n            _, preds = outputs.max(1)\n            \n            predictions.extend(preds.cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {idx: sirna for sirna, idx in sirna_to_label.items()}\n    predictions_sirna = [label_to_sirna[pred] for pred in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions_sirna\n    })\n    submission.to_csv('submission_aug.csv', index=False)\n    print(\"\\nSubmission saved to submission_aug.csv\")\n    print(f\"Best validation accuracy: {best_acc:.2f}%\")\n    print(\"Sample predictions:\")\n    print(submission.head(10))\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-07T16:41:03.804378Z","iopub.execute_input":"2026-04-07T16:41:03.804717Z","iopub.status.idle":"2026-04-07T16:48:31.553786Z","shell.execute_reply.started":"2026-04-07T16:41:03.804685Z","shell.execute_reply":"2026-04-07T16:48:31.552632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df\n        self.data_dir = data_dir\n        self.mode = mode\n        \n        self.train_transforms = transforms.Compose([\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomVerticalFlip(p=0.5),\n            transforms.RandomRotation(degrees=15),\n        ])\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        \n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode in ['train', 'val']:\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                channels.append(img)\n            else:\n                channels.append(np.zeros((320, 320), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        \n        img = img.astype(np.float32)\n        \n        for c in range(6):\n            mean = img[:, :, c].mean()\n            std = img[:, :, c].std()\n            if std > 0:\n                img[:, :, c] = (img[:, :, c] - mean) / std\n            else:\n                img[:, :, c] = 0.0\n                \n        img_tensor = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train' and Config.AUGMENT:\n            img_tensor = self.train_transforms(img_tensor)\n        \n        if self.mode in ['train', 'val']:\n            label = row['label']\n            return img_tensor, label\n        else:\n            return img_tensor, row['id_code']\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df\n        self.data_dir = data_dir\n        self.mode = mode\n        \n        self.train_transforms = transforms.Compose([\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomVerticalFlip(p=0.5),\n            transforms.RandomRotation(degrees=15),\n        ])\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        \n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'train':\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                channels.append(img)\n            else:\n                channels.append(np.zeros((512, 512), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        \n        img = img.astype(np.float32)\n        \n        for c in range(6):\n            mean = img[:, :, c].mean()\n            std = img[:, :, c].std()\n            if std > 0:\n                img[:, :, c] = (img[:, :, c] - mean) / std\n            else:\n                img[:, :, c] = 0.0\n                \n        img_tensor = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train' and Config.AUGMENT:\n            img_tensor = self.train_transforms(img_tensor)\n        \n        if self.mode == 'train':\n            label = row['label']\n            return img_tensor, label\n        else:\n            return img_tensor, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'features') and hasattr(self.backbone.features, 'conv0'):\n            old_conv = self.backbone.features.conv0\n            self.backbone.features.conv0 = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.features.conv0.weight[:, :3] = old_conv.weight\n                self.backbone.features.conv0.weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.features.conv0.bias = old_conv.bias\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.4),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    train_df = train_df[train_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    test_df = test_df[test_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types in train: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    \n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {sirna: idx for idx, sirna in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    print(f\"Number of unique sirnas: {len(unique_sirnas)}\")\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    \n    train_data, val_data = train_test_split(train_df, test_size=0.15, random_state=Config.SEED, stratify=train_df['label'])\n    \n    train_dataset = CellularDataset(train_data, Config.DATA_DIR, mode='train')\n    val_dataset = CellularDataset(val_data, Config.DATA_DIR, mode='val') \n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, \n                              shuffle=True, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, \n                            shuffle=False, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    \n    print(f\"Creating model: {Config.MODEL_NAME} with Augmentations...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_densenet121_aug.pth')\n            print(f\"Saved best model with accuracy: {best_acc:.2f}%\")\n    \n    print(\"\\nLoading best model for inference...\")\n    model.load_state_dict(torch.load('best_densenet121_aug.pth'))\n    \n    print(\"Generating predictions...\")\n    test_dataset = CellularDataset(test_df, Config.DATA_DIR, mode='test')\n    test_loader = DataLoader(test_dataset, batch_size=Config.BATCH_SIZE, \n                             shuffle=False, num_workers=Config.NUM_WORKERS)\n    \n    model.eval()\n    predictions = []\n    ids = []\n    \n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Inference'):\n            imgs = imgs.to(device)\n            outputs = model(imgs)\n            _, preds = outputs.max(1)\n            \n            predictions.extend(preds.cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {idx: sirna for sirna, idx in sirna_to_label.items()}\n    predictions_sirna = [label_to_sirna[pred] for pred in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions_sirna\n    })\n    submission.to_csv('submission_aug.csv', index=False)\n    print(\"\\nSubmission saved to submission_aug.csv\")\n    print(f\"Best validation accuracy: {best_acc:.2f}%\")\n    print(\"Sample predictions:\")\n    print(submission.head(10))\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-07T16:51:29.281055Z","iopub.execute_input":"2026-04-07T16:51:29.281837Z","iopub.status.idle":"2026-04-07T16:58:33.219669Z","shell.execute_reply.started":"2026-04-07T16:51:29.281791Z","shell.execute_reply":"2026-04-07T16:58:33.218621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'convnext_tiny'\n    IMG_SIZE = 320  \n    BATCH_SIZE = 32\n    EPOCHS = 15  \n    LR = 3e-4\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    CELL_TYPES = ['HUVEC']  \n    \n    AUGMENT = True\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df\n        self.data_dir = data_dir\n        self.mode = mode\n        \n        self.train_transforms = transforms.Compose([\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomVerticalFlip(p=0.5),\n            transforms.RandomRotation(degrees=15),\n        ])\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        \n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode in ['train', 'val']:\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                channels.append(img)\n            else:\n                channels.append(np.zeros((320, 320), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        \n        img = img.astype(np.float32)\n        \n        for c in range(6):\n            mean = img[:, :, c].mean()\n            std = img[:, :, c].std()\n            if std > 0:\n                img[:, :, c] = (img[:, :, c] - mean) / std\n            else:\n                img[:, :, c] = 0.0\n                \n        img_tensor = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train' and Config.AUGMENT:\n            img_tensor = self.train_transforms(img_tensor)\n        \n        if self.mode in ['train', 'val']:\n            label = row['label']\n            return img_tensor, label\n        else:\n            return img_tensor, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'stem') and isinstance(self.backbone.stem[0], nn.Conv2d):\n            old_conv = self.backbone.stem[0]\n            self.backbone.stem[0] = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.stem[0].weight[:, :3] = old_conv.weight\n                self.backbone.stem[0].weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.stem[0].bias = old_conv.bias\n                    \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.4), \n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    train_df = train_df[train_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    test_df = test_df[test_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {sirna: idx for idx, sirna in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    \n    train_data, val_data = train_test_split(train_df, test_size=0.15, random_state=Config.SEED, stratify=train_df['label'])\n    \n    train_dataset = CellularDataset(train_data, Config.DATA_DIR, mode='train')\n    val_dataset = CellularDataset(val_data, Config.DATA_DIR, mode='val') \n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, \n                              shuffle=True, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, \n                            shuffle=False, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    \n    print(f\"Creating model: {Config.MODEL_NAME}...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_convnext_tiny.pth')\n            print(f\"Saved best model with accuracy: {best_acc:.2f}%\")\n    \n    print(\"\\nLoading best model for inference...\")\n    model.load_state_dict(torch.load('best_convnext_tiny.pth'))\n    \n    print(\"Generating predictions...\")\n    test_dataset = CellularDataset(test_df, Config.DATA_DIR, mode='test')\n    test_loader = DataLoader(test_dataset, batch_size=Config.BATCH_SIZE, \n                             shuffle=False, num_workers=Config.NUM_WORKERS)\n    \n    model.eval()\n    predictions = []\n    ids = []\n    \n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Inference'):\n            imgs = imgs.to(device)\n            outputs = model(imgs)\n            _, preds = outputs.max(1)\n            \n            predictions.extend(preds.cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {idx: sirna for sirna, idx in sirna_to_label.items()}\n    predictions_sirna = [label_to_sirna[pred] for pred in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions_sirna\n    })\n    submission.to_csv('submission_convnext.csv', index=False)\n    print(\"\\nSubmission saved to submission_convnext.csv\")\n    print(f\"Best validation accuracy: {best_acc:.2f}%\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-07T17:02:26.390716Z","iopub.execute_input":"2026-04-07T17:02:26.391055Z","iopub.status.idle":"2026-04-07T17:09:19.064435Z","shell.execute_reply.started":"2026-04-07T17:02:26.391023Z","shell.execute_reply":"2026-04-07T17:09:19.063457Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'convnext_tiny'\n    IMG_SIZE = 320  \n    BATCH_SIZE = 32\n    EPOCHS = 15  \n    LR = 3e-4\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    AUGMENT = True\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df\n        self.data_dir = data_dir\n        self.mode = mode\n        \n        self.train_transforms = transforms.Compose([\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomVerticalFlip(p=0.5),\n            transforms.RandomRotation(degrees=15),\n        ])\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        \n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode in ['train', 'val']:\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                channels.append(img)\n            else:\n                channels.append(np.zeros((320, 320), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        \n        img = img.astype(np.float32)\n        \n        for c in range(6):\n            mean = img[:, :, c].mean()\n            std = img[:, :, c].std()\n            if std > 0:\n                img[:, :, c] = (img[:, :, c] - mean) / std\n            else:\n                img[:, :, c] = 0.0\n                \n        img_tensor = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train' and Config.AUGMENT:\n            img_tensor = self.train_transforms(img_tensor)\n        \n        if self.mode in ['train', 'val']:\n            label = row['label']\n            return img_tensor, label\n        else:\n            return img_tensor, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'stem') and isinstance(self.backbone.stem[0], nn.Conv2d):\n            old_conv = self.backbone.stem[0]\n            self.backbone.stem[0] = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.stem[0].weight[:, :3] = old_conv.weight\n                self.backbone.stem[0].weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.stem[0].bias = old_conv.bias\n                    \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.4), \n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types in train: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {sirna: idx for idx, sirna in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    \n    train_data, val_data = train_test_split(train_df, test_size=0.15, random_state=Config.SEED, stratify=train_df['label'])\n    \n    train_dataset = CellularDataset(train_data, Config.DATA_DIR, mode='train')\n    val_dataset = CellularDataset(val_data, Config.DATA_DIR, mode='val') \n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, \n                              shuffle=True, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, \n                            shuffle=False, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    \n    print(f\"Creating model: {Config.MODEL_NAME}...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_convnext_tiny.pth')\n            print(f\"Saved best model with accuracy: {best_acc:.2f}%\")\n    \n    print(\"\\nLoading best model for inference...\")\n    model.load_state_dict(torch.load('best_convnext_tiny.pth'))\n    \n    print(\"Generating predictions...\")\n    test_dataset = CellularDataset(test_df, Config.DATA_DIR, mode='test')\n    test_loader = DataLoader(test_dataset, batch_size=Config.BATCH_SIZE, \n                             shuffle=False, num_workers=Config.NUM_WORKERS)\n    \n    model.eval()\n    predictions = []\n    ids = []\n    \n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Inference'):\n            imgs = imgs.to(device)\n            outputs = model(imgs)\n            _, preds = outputs.max(1)\n            \n            predictions.extend(preds.cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {idx: sirna for sirna, idx in sirna_to_label.items()}\n    predictions_sirna = [label_to_sirna[pred] for pred in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions_sirna\n    })\n    submission.to_csv('submission_convnext.csv', index=False)\n    print(\"\\nSubmission saved to submission_convnext.csv\")\n    print(f\"Best validation accuracy: {best_acc:.2f}%\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-07T17:40:36.836795Z","iopub.execute_input":"2026-04-07T17:40:36.837665Z","iopub.status.idle":"2026-04-07T17:41:08.429302Z","shell.execute_reply.started":"2026-04-07T17:40:36.83761Z","shell.execute_reply":"2026-04-07T17:41:08.428117Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'resnet50'  \n    IMG_SIZE = 320  \n    BATCH_SIZE = 32\n    EPOCHS = 10  \n    LR = 3e-4\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    CELL_TYPES = ['HUVEC']  \n    \n    CONVERT_SIRNA = True\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train', transform=None):\n        self.df = df\n        self.data_dir = data_dir\n        self.mode = mode\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        \n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'train':\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                channels.append(img)\n            else:\n                channels.append(np.zeros((512, 512), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        \n        img = img.astype(np.float32) / 255.0\n        \n        if self.transform:\n            img = torch.from_numpy(img).permute(2, 0, 1)\n        else:\n            img = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train':\n            label = row['label']\n            return img, label\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'conv1'):\n            old_conv = self.backbone.conv1\n            self.backbone.conv1 = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.conv1.weight[:, :3] = old_conv.weight\n                self.backbone.conv1.weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.conv1.bias = old_conv.bias\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    train_df = train_df[train_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types in train: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    \n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {sirna: idx for idx, sirna in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    print(f\"Number of unique sirnas: {len(unique_sirnas)}\")\n    print(f\"Label range: 0 to {train_df['label'].max()}\")\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    \n    train_data, val_data = train_test_split(train_df, test_size=0.15, random_state=Config.SEED, stratify=train_df['label'])\n    \n    train_dataset = CellularDataset(train_data, Config.DATA_DIR, mode='train')\n    val_dataset = CellularDataset(val_data, Config.DATA_DIR, mode='train')\n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, \n                              shuffle=True, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, \n                            shuffle=False, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    \n    print(f\"Creating model: {Config.MODEL_NAME}...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_resnet50.pth')\n            print(f\"Saved best model with accuracy: {best_acc:.2f}%\")\n    \n    print(\"\\nLoading best model for inference...\")\n    model.load_state_dict(torch.load('best_resnet50.pth'))\n    \n    print(\"Generating predictions...\")\n    test_dataset = CellularDataset(test_df, Config.DATA_DIR, mode='test')\n    test_loader = DataLoader(test_dataset, batch_size=Config.BATCH_SIZE, \n                             shuffle=False, num_workers=Config.NUM_WORKERS)\n    \n    model.eval()\n    predictions = []\n    ids = []\n    \n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Inference'):\n            imgs = imgs.to(device)\n            outputs = model(imgs)\n            _, preds = outputs.max(1)\n            \n            predictions.extend(preds.cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {idx: sirna for sirna, idx in sirna_to_label.items()}\n    predictions_sirna = [label_to_sirna[pred] for pred in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions_sirna\n    })\n    submission.to_csv('submission.csv', index=False)\n    print(\"\\nSubmission saved to submission.csv\")\n    print(f\"Best validation accuracy: {best_acc:.2f}%\")\n    print(\"Sample predictions:\")\n    print(submission.head(10))\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-07T17:42:30.873608Z","iopub.execute_input":"2026-04-07T17:42:30.873931Z","iopub.status.idle":"2026-04-07T19:03:37.347661Z","shell.execute_reply.started":"2026-04-07T17:42:30.873902Z","shell.execute_reply":"2026-04-07T19:03:37.346755Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'densenet121'\n    IMG_SIZE = 320  \n    BATCH_SIZE = 32\n    EPOCHS = 10  \n    LR = 3e-4\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    CELL_TYPES = ['HUVEC']  \ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df\n        self.data_dir = data_dir\n        self.mode = mode\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'train':\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                channels.append(img)\n            else:\n                channels.append(np.zeros((512, 512), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        \n        img = img.astype(np.float32) / 255.0\n        img = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train':\n            return img, row['label']\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        \n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'features') and hasattr(self.backbone.features, 'conv0'):\n            old_conv = self.backbone.features.conv0\n            \n            self.backbone.features.conv0 = nn.Conv2d(\n                in_channels,\n                old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            \n            with torch.no_grad():\n                new_weight = old_conv.weight.mean(dim=1, keepdim=True)\n                new_weight = new_weight.repeat(1, in_channels, 1, 1) / in_channels\n                self.backbone.features.conv0.weight.copy_(new_weight)\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer):\n    model.train()\n    running_loss, correct, total = 0, 0, 0\n    \n    for imgs, labels in tqdm(loader, desc='Training'):\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, preds = outputs.max(1)\n        total += labels.size(0)\n        correct += preds.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion):\n    model.eval()\n    running_loss, correct, total = 0, 0, 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            \n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, preds = outputs.max(1)\n            total += labels.size(0)\n            correct += preds.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    test_df = test_df.sort_values('id_code').reset_index(drop=True)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    train_df = train_df[train_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    \n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    \n    train_data, val_data = train_test_split(\n        train_df, test_size=0.15, stratify=train_df['label'], random_state=Config.SEED\n    )\n    \n    train_loader = DataLoader(\n        CellularDataset(train_data, Config.DATA_DIR),\n        batch_size=Config.BATCH_SIZE,\n        shuffle=True,\n        num_workers=Config.NUM_WORKERS,\n        pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        CellularDataset(val_data, Config.DATA_DIR),\n        batch_size=Config.BATCH_SIZE,\n        shuffle=False,\n        num_workers=Config.NUM_WORKERS,\n        pin_memory=True\n    )\n    \n    print(\"Creating DenseNet121 model...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    \n    best_acc = 0\n    \n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        \n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer)\n        val_loss, val_acc = validate(model, val_loader, criterion)\n        \n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_densenet121.pth')\n            print(f\"Saved best model: {best_acc:.2f}%\")\n    \n    print(\"\\nInference...\")\n    model.load_state_dict(torch.load('best_densenet121.pth'))\n    \n    test_loader = DataLoader(\n        CellularDataset(test_df, Config.DATA_DIR, mode='test'),\n        batch_size=Config.BATCH_SIZE,\n        shuffle=False,\n        num_workers=Config.NUM_WORKERS\n    )\n    \n    model.eval()\n    preds, ids = [], []\n    \n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader):\n            imgs = imgs.to(device)\n            outputs = model(imgs)\n            _, p = outputs.max(1)\n            \n            preds.extend(p.cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    preds = [label_to_sirna[p] for p in preds]\n    \n    submission = pd.DataFrame({\n        'id_code': test_df['id_code'].values,\n        'sirna': preds\n    })\n    print(test_df.head())\n    print(submission.head())\n    \n    submission.to_csv('submission.csv', index=False)\n    print(\"Saved submission.csv\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T03:16:50.01503Z","iopub.execute_input":"2026-04-08T03:16:50.015385Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'convnext_tiny'\n    IMG_SIZE = 320  \n    BATCH_SIZE = 32\n    EPOCHS = 10\n    LR = 3e-4\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    CELL_TYPES = ['HUVEC']  \n    \n    AUGMENT = True\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df\n        self.data_dir = data_dir\n        self.mode = mode\n        \n        self.train_transforms = transforms.Compose([\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomVerticalFlip(p=0.5),\n            transforms.RandomRotation(degrees=15),\n        ])\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        \n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode in ['train', 'val']:\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                channels.append(img)\n            else:\n                channels.append(np.zeros((320, 320), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        \n        img = img.astype(np.float32)\n        \n        for c in range(6):\n            mean = img[:, :, c].mean()\n            std = img[:, :, c].std()\n            if std > 0:\n                img[:, :, c] = (img[:, :, c] - mean) / std\n            else:\n                img[:, :, c] = 0.0\n                \n        img_tensor = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train' and Config.AUGMENT:\n            img_tensor = self.train_transforms(img_tensor)\n        \n        if self.mode in ['train', 'val']:\n            label = row['label']\n            return img_tensor, label\n        else:\n            return img_tensor, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'stem') and isinstance(self.backbone.stem[0], nn.Conv2d):\n            old_conv = self.backbone.stem[0]\n            self.backbone.stem[0] = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.stem[0].weight[:, :3] = old_conv.weight\n                self.backbone.stem[0].weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.stem[0].bias = old_conv.bias\n                    \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.4), \n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    train_df = train_df[train_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    test_df = test_df[test_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {sirna: idx for idx, sirna in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    \n    train_data, val_data = train_test_split(train_df, test_size=0.15, random_state=Config.SEED, stratify=train_df['label'])\n    \n    train_dataset = CellularDataset(train_data, Config.DATA_DIR, mode='train')\n    val_dataset = CellularDataset(val_data, Config.DATA_DIR, mode='val') \n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, \n                              shuffle=True, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, \n                            shuffle=False, num_workers=Config.NUM_WORKERS, pin_memory=True)\n    \n    print(f\"Creating model: {Config.MODEL_NAME}...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_convnext_tiny.pth')\n            print(f\"Saved best model with accuracy: {best_acc:.2f}%\")\n    \n    print(\"\\nLoading best model for inference...\")\n    model.load_state_dict(torch.load('best_convnext_tiny.pth'))\n    \n    print(\"Generating predictions...\")\n    test_dataset = CellularDataset(test_df, Config.DATA_DIR, mode='test')\n    test_loader = DataLoader(test_dataset, batch_size=Config.BATCH_SIZE, \n                             shuffle=False, num_workers=Config.NUM_WORKERS)\n    \n    model.eval()\n    predictions = []\n    ids = []\n    \n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Inference'):\n            imgs = imgs.to(device)\n            outputs = model(imgs)\n            _, preds = outputs.max(1)\n            \n            predictions.extend(preds.cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {idx: sirna for sirna, idx in sirna_to_label.items()}\n    predictions_sirna = [label_to_sirna[pred] for pred in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions_sirna\n    })\n    submission.to_csv('submission_convnext.csv', index=False)\n    print(\"\\nSubmission saved to submission_convnext.csv\")\n    print(f\"Best validation accuracy: {best_acc:.2f}%\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-23T17:00:29.987076Z","iopub.execute_input":"2026-04-23T17:00:29.98753Z","iopub.status.idle":"2026-04-23T18:00:54.626658Z","shell.execute_reply.started":"2026-04-23T17:00:29.987494Z","shell.execute_reply":"2026-04-23T18:00:54.625592Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport math\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'resnet50'\n    IMG_SIZE = 384\n    BATCH_SIZE = 24\n    EPOCHS = 10\n    LR = 3e-4\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    ARCFACE_S = 30.0\n    ARCFACE_M = 0.5\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train', site=None):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode = mode\n        self.site = site\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row, site):\n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'test':\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s{site}_w'\n        else:\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s{site}_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            else:\n                img = np.zeros((512, 512), dtype=np.uint8)\n            channels.append(img)\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        return img.astype(np.float32) / 255.0\n    \n    def augment(self, img):\n        if self.mode != 'train':\n            return img\n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        if np.random.rand() > 0.5:\n            k = np.random.randint(1, 4)\n            img = np.rot90(img, k=k).copy()\n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        \n        if self.mode == 'train':\n            site = np.random.choice(['1', '2'])\n        else:\n            site = self.site if self.site is not None else '1'\n            \n        img = self.load_image(row, site)\n        img = self.augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode in ['train', 'val']:\n            return img, row['label']\n        else:\n            return img, row['id_code']\nclass ArcMarginProduct(nn.Module):\n    \n    def __init__(self, in_features, out_features, s=30.0, m=0.50):\n        super().__init__()\n        self.in_features = in_features\n        self.out_features = out_features\n        self.s = s\n        self.m = m\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        nn.init.xavier_uniform_(self.weight)\n        \n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n        self.th = math.cos(math.pi - m)\n        self.mm = math.sin(math.pi - m) * m\n    def forward(self, input, label):\n        cosine = F.linear(F.normalize(input), F.normalize(self.weight))\n        sine = torch.sqrt(1.0 - torch.pow(cosine, 2) + 1e-6)\n        phi = cosine * self.cos_m - sine * self.sin_m\n        phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n        one_hot = torch.zeros(cosine.size(), device=input.device)\n        one_hot.scatter_(1, label.view(-1, 1).long(), 1)\n        output = (one_hot * phi) + ((1.0 - one_hot) * cosine)\n        output *= self.s\n        return output\nclass ArcFaceModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6, s=30.0, m=0.5):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'conv1'):\n            old_conv = self.backbone.conv1\n            self.backbone.conv1 = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.conv1.weight[:, :3] = old_conv.weight\n                self.backbone.conv1.weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.conv1.bias = old_conv.bias\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.dropout = nn.Dropout(0.3)\n        self.arc_margin = ArcMarginProduct(n_features, num_classes, s=s, m=m)\n    \n    def forward(self, x, label=None):\n        features = self.backbone(x)\n        features = self.dropout(features)\n        if label is not None:\n            return self.arc_margin(features, label)\n        else:\n            cosine = F.linear(F.normalize(features), F.normalize(self.arc_margin.weight))\n            return cosine\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs, labels)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef predict_with_tta(model, df, data_dir, device):\n    model.eval()\n    configs = [('1', None), ('1', 'h'), ('1', 'v'), ('2', None), ('2', 'h'), ('2', 'v')]\n    all_probs = []\n    ids = None\n    \n    for site, flip in configs:\n        dataset = CellularDataset(df, data_dir, mode='test', site=site)\n        loader = DataLoader(dataset, batch_size=Config.BATCH_SIZE,\n                          shuffle=False, num_workers=Config.NUM_WORKERS)\n        \n        probs = []\n        batch_ids = []\n        with torch.no_grad():\n            for imgs, img_ids in tqdm(loader, desc=f'Site{site} {flip or \"\"}'):\n                imgs = imgs.to(device)\n                if flip == 'h':\n                    imgs = torch.flip(imgs, dims=[3])\n                elif flip == 'v':\n                    imgs = torch.flip(imgs, dims=[2])\n                outputs = model(imgs)\n                probs.append(outputs.softmax(dim=1).cpu())\n                batch_ids.extend(img_ids)\n        \n        probs = torch.cat(probs, dim=0)\n        all_probs.append(probs)\n        if ids is None:\n            ids = batch_ids\n    \n    avg_probs = torch.stack(all_probs).mean(dim=0)\n    preds = avg_probs.argmax(dim=1).numpy()\n    return preds, ids\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types in train: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    print(f\"Number of classes: {Config.NUM_CLASSES}\")\n    \n    train_data, val_data = train_test_split(\n        train_df, test_size=0.15, random_state=Config.SEED,\n        stratify=train_df['label']\n    )\n    \n    train_loader = DataLoader(\n        CellularDataset(train_data, Config.DATA_DIR, mode='train'),\n        batch_size=Config.BATCH_SIZE, shuffle=True,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        CellularDataset(val_data, Config.DATA_DIR, mode='val', site='1'),\n        batch_size=Config.BATCH_SIZE, shuffle=False,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    \n    print(f\"Creating ArcFace model: {Config.MODEL_NAME}...\")\n    model = ArcFaceModel(\n        Config.MODEL_NAME, Config.NUM_CLASSES,\n        s=Config.ARCFACE_S, m=Config.ARCFACE_M\n    ).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optimizer, T_0=Config.EPOCHS, T_mult=1\n    )\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_arcface_resnet50.pth')\n            print(f\"Saved best model: {best_acc:.2f}%\")\n    \n    print(\"\\nInference with TTA (dual site + flips)...\")\n    model.load_state_dict(torch.load('best_arcface_resnet50.pth'))\n    \n    predictions, ids = predict_with_tta(model, test_df, Config.DATA_DIR, device)\n    \n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    predictions = [label_to_sirna[p] for p in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions\n    })\n    submission.to_csv('submission_arcface_resnet50.csv', index=False)\n    print(f\"Saved submission_arcface_resnet50.csv | Best val acc: {best_acc:.2f}%\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-23T18:15:04.845621Z","iopub.execute_input":"2026-04-23T18:15:04.846315Z","iopub.status.idle":"2026-04-23T19:04:36.029562Z","shell.execute_reply.started":"2026-04-23T18:15:04.846279Z","shell.execute_reply":"2026-04-23T19:04:36.028506Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nimport albumentations as A\nfrom sklearn.model_selection import GroupKFold\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    IMG_SIZE = 320\n    BATCH_SIZE = 16\n    EPOCHS = 15\n    LR = 1e-4\n    NUM_WORKERS = 2\n    N_SPLITS = 5\n    SEED = 42\ntrain_tfms = A.Compose([\n    A.RandomRotate90(),\n    A.RandomRotate90(p=0.5),\n    A.Transpose(),\n    A.RandomBrightnessContrast(p=0.5),\n    A.GaussianBlur(p=0.3),\n    A.CoarseDropout(max_holes=8, max_height=32, max_width=32, p=0.5),\n])\nvalid_tfms = A.Compose([])\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, transforms=None, mode='train'):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.transforms = transforms\n        self.mode = mode\n    def load_img(self, row, site):\n        exp, plate, well = row['experiment'], row['plate'], row['well']\n        base = f\"{self.data_dir}/{self.mode}/{exp}/Plate{plate}/{well}_s{site}_w\"\n        channels = []\n        for i in range(1, 7):\n            path = f\"{base}{i}.png\"\n            img = cv2.imread(path, cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                img = np.zeros((512, 512), dtype=np.uint8)\n            channels.append(img)\n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        img = img.astype(np.float32)\n        img = (img - img.mean()) / (img.std() + 1e-6)\n        return img\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img1 = self.load_img(row, site=1)\n        img2 = self.load_img(row, site=2)\n        img = np.concatenate([img1, img2], axis=-1)\n        if self.transforms:\n            img = self.transforms(image=img)['image']\n        img = torch.tensor(img).permute(2, 0, 1).float()\n        if self.mode == 'train':\n            return img, row['label']\n        else:\n            return img, row['id_code']\n    def __len__(self):\n        return len(self.df)\nclass GeM(nn.Module):\n    def __init__(self, p=3):\n        super().__init__()\n        self.p = nn.Parameter(torch.ones(1) * p)\n    def forward(self, x):\n        return F.avg_pool2d(x.clamp(min=1e-6).pow(self.p),\n                            (x.size(-2), x.size(-1))).pow(1./self.p)\nclass ChannelAttention(nn.Module):\n    def __init__(self, c):\n        super().__init__()\n        self.fc = nn.Sequential(\n            nn.Linear(c, c // 2),\n            nn.ReLU(),\n            nn.Linear(c // 2, c),\n            nn.Sigmoid()\n        )\n    def forward(self, x):\n        b, c, _, _ = x.shape\n        y = x.mean((2, 3))\n        y = self.fc(y).view(b, c, 1, 1)\n        return x * y\nclass ArcMarginProduct(nn.Module):\n    def __init__(self, in_f, out_f):\n        super().__init__()\n        self.weight = nn.Parameter(torch.randn(out_f, in_f))\n    def forward(self, x, label):\n        cosine = F.linear(F.normalize(x), F.normalize(self.weight))\n        phi = cosine - 0.5\n        one_hot = torch.zeros_like(cosine)\n        one_hot.scatter_(1, label.view(-1, 1), 1)\n        return (one_hot * phi + (1 - one_hot) * cosine) * 30\nclass Model(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        self.backbone = timm.create_model(\n            'tf_efficientnet_b4',\n            pretrained=True,\n            in_chans=12,\n            features_only=True\n        )\n        self.pool = GeM()\n        self.attn = ChannelAttention(1792)\n        self.fc = nn.Linear(1792, 512)\n        self.bn = nn.BatchNorm1d(512)\n        self.arc = ArcMarginProduct(512, num_classes)\n    def forward(self, x, labels=None):\n        x = self.backbone(x)[-1]\n        x = self.attn(x)\n        x = self.pool(x).flatten(1)\n        x = self.fc(x)\n        x = self.bn(x)\n        if labels is not None:\n            return self.arc(x, labels)\n        return x\ndef train_fn(loader, model, optimizer, scaler):\n    model.train()\n    total_loss = 0\n    for x, y in tqdm(loader):\n        x, y = x.to(device), y.to(device)\n        optimizer.zero_grad()\n        with torch.cuda.amp.autocast():\n            logits = model(x, y)\n            loss = F.cross_entropy(logits, y)\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        total_loss += loss.item()\n    return total_loss / len(loader)\ndef valid_fn(loader, model):\n    model.eval()\n    correct, total = 0, 0\n    with torch.no_grad():\n        for x, y in loader:\n            x, y = x.to(device), y.to(device)\n            logits = model(x, y)\n            pred = logits.argmax(1)\n            correct += (pred == y).sum().item()\n            total += y.size(0)\n    return 100 * correct / total\ndef run():\n    df = pd.read_csv(Config.TRAIN_CSV)\n    df['sirna_id'] = df['sirna'].str.replace('sirna_', '').astype(int)\n    label_map = {v: i for i, v in enumerate(df['sirna_id'].unique())}\n    df['label'] = df['sirna_id'].map(label_map)\n    gkf = GroupKFold(n_splits=Config.N_SPLITS)\n    for fold, (tr, va) in enumerate(gkf.split(df, df['label'], df['experiment'])):\n        print(f\"\\nFOLD {fold}\")\n        train_ds = CellularDataset(df.iloc[tr], Config.DATA_DIR, train_tfms)\n        val_ds = CellularDataset(df.iloc[va], Config.DATA_DIR, valid_tfms)\n        train_loader = DataLoader(train_ds, batch_size=Config.BATCH_SIZE, shuffle=True)\n        val_loader = DataLoader(val_ds, batch_size=Config.BATCH_SIZE)\n        model = Model(df['label'].nunique()).to(device)\n        optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR)\n        scaler = torch.cuda.amp.GradScaler()\n        best = 0\n        for epoch in range(Config.EPOCHS):\n            loss = train_fn(train_loader, model, optimizer, scaler)\n            acc = valid_fn(val_loader, model)\n            print(f\"Epoch {epoch} | Loss: {loss:.4f} | Acc: {acc:.2f}\")\n            if acc > best:\n                best = acc\n                torch.save(model.state_dict(), f\"fold{fold}.pth\")\ndef inference():\n    test_df = pd.read_csv(Config.TEST_CSV)\n    test_ds = CellularDataset(test_df, Config.DATA_DIR, mode='test')\n    loader = DataLoader(test_ds, batch_size=16)\n    model = Model(1108).to(device)\n    model.load_state_dict(torch.load(\"fold0.pth\"))\n    model.eval()\n    preds, ids = [], []\n    with torch.no_grad():\n        for x, id_ in tqdm(loader):\n            x = x.to(device)\n            p1 = model(x)\n            p2 = model(torch.flip(x, dims=[-1]))\n            p = (p1 + p2) / 2\n            pred = p.argmax(1)\n            preds.extend(pred.cpu().numpy())\n            ids.extend(id_)\n    sub = pd.DataFrame({'id_code': ids, 'sirna': preds})\n    sub.to_csv(\"submission.csv\", index=False)\nif __name__ == \"__main__\":\n    run()\n    inference()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-23T19:05:56.730464Z","iopub.execute_input":"2026-04-23T19:05:56.731217Z","iopub.status.idle":"2026-04-23T19:06:07.057583Z","shell.execute_reply.started":"2026-04-23T19:05:56.731185Z","shell.execute_reply":"2026-04-23T19:06:07.056401Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'efficientnet_b0'\n    IMG_SIZE = 384\n    BATCH_SIZE = 32\n    EPOCHS = 12\n    LR = 3e-4\n    LABEL_SMOOTHING = 0.1\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train', site=None):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode = mode\n        self.site = site\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row, site):\n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'test':\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s{site}_w'\n        else:\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s{site}_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            else:\n                img = np.zeros((512, 512), dtype=np.uint8)\n            channels.append(img)\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        return img.astype(np.float32) / 255.0\n    \n    def augment(self, img):\n        if self.mode != 'train':\n            return img\n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        if np.random.rand() > 0.5:\n            k = np.random.randint(1, 4)\n            img = np.rot90(img, k=k).copy()\n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        \n        if self.mode == 'train':\n            site = np.random.choice(['1', '2'])\n        else:\n            site = self.site if self.site is not None else '1'\n            \n        img = self.load_image(row, site)\n        img = self.augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode in ['train', 'val']:\n            return img, row['label']\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'conv_stem'):\n            old_conv = self.backbone.conv_stem\n            self.backbone.conv_stem = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.conv_stem.weight[:, :3] = old_conv.weight\n                self.backbone.conv_stem.weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.conv_stem.bias = old_conv.bias\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef predict_with_tta(model, df, data_dir, device):\n    model.eval()\n    configs = [('1', None), ('1', 'h'), ('1', 'v'), ('2', None), ('2', 'h'), ('2', 'v')]\n    all_probs = []\n    ids = None\n    \n    for site, flip in configs:\n        dataset = CellularDataset(df, data_dir, mode='test', site=site)\n        loader = DataLoader(dataset, batch_size=Config.BATCH_SIZE,\n                          shuffle=False, num_workers=Config.NUM_WORKERS)\n        \n        probs = []\n        batch_ids = []\n        with torch.no_grad():\n            for imgs, img_ids in tqdm(loader, desc=f'Site{site} {flip or \"\"}'):\n                imgs = imgs.to(device)\n                if flip == 'h':\n                    imgs = torch.flip(imgs, dims=[3])\n                elif flip == 'v':\n                    imgs = torch.flip(imgs, dims=[2])\n                outputs = model(imgs)\n                probs.append(outputs.softmax(dim=1).cpu())\n                batch_ids.extend(img_ids)\n        \n        probs = torch.cat(probs, dim=0)\n        all_probs.append(probs)\n        if ids is None:\n            ids = batch_ids\n    \n    avg_probs = torch.stack(all_probs).mean(dim=0)\n    preds = avg_probs.argmax(dim=1).numpy()\n    return preds, ids\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types in train: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    print(f\"Number of classes: {Config.NUM_CLASSES}\")\n    \n    train_data, val_data = train_test_split(\n        train_df, test_size=0.15, random_state=Config.SEED,\n        stratify=train_df['label']\n    )\n    \n    train_loader = DataLoader(\n        CellularDataset(train_data, Config.DATA_DIR, mode='train'),\n        batch_size=Config.BATCH_SIZE, shuffle=True,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        CellularDataset(val_data, Config.DATA_DIR, mode='val', site='1'),\n        batch_size=Config.BATCH_SIZE, shuffle=False,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    \n    print(f\"Creating model: {Config.MODEL_NAME}...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss(label_smoothing=Config.LABEL_SMOOTHING)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optimizer, T_0=Config.EPOCHS, T_mult=1\n    )\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_efficientnet_b0.pth')\n            print(f\"Saved best model: {best_acc:.2f}%\")\n    \n    print(\"\\nInference with TTA (dual site + fl)...\")\n    model.load_state_dict(torch.load('best_efficientnet_b0.pth'))\n    \n    predictions, ids = predict_with_tta(model, test_df, Config.DATA_DIR, device)\n    \n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    predictions = [label_to_sirna[p] for p in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions\n    })\n    submission.to_csv('submission_efficientnet_b0.csv', index=False)\n    print(f\"Saved submission_efficientnet_b0.csv | Best val acc: {best_acc:.2f}%\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-23T19:22:16.508121Z","iopub.execute_input":"2026-04-23T19:22:16.508773Z","iopub.status.idle":"2026-04-23T19:24:19.523254Z","shell.execute_reply.started":"2026-04-23T19:22:16.5087Z","shell.execute_reply":"2026-04-23T19:24:19.522142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'efficientnet_b0'\n    IMG_SIZE = 384\n    BATCH_SIZE = 32\n    EPOCHS = 10\n    LR = 3e-4\n    LABEL_SMOOTHING = 0.1\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train', site=None):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode = mode\n        self.site = site\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row, site):\n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'test':\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s{site}_w'\n        else:\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s{site}_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            else:\n                img = np.zeros((512, 512), dtype=np.uint8)\n            channels.append(img)\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        return img.astype(np.float32) / 255.0\n    \n    def augment(self, img):\n        if self.mode != 'train':\n            return img\n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        if np.random.rand() > 0.5:\n            k = np.random.randint(1, 4)\n            img = np.rot90(img, k=k).copy()\n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        \n        if self.mode == 'train':\n            site = np.random.choice(['1', '2'])\n        else:\n            site = self.site if self.site is not None else '1'\n            \n        img = self.load_image(row, site)\n        img = self.augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode in ['train', 'val']:\n            return img, row['label']\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'conv_stem'):\n            old_conv = self.backbone.conv_stem\n            self.backbone.conv_stem = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.conv_stem.weight[:, :3] = old_conv.weight\n                self.backbone.conv_stem.weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.conv_stem.bias = old_conv.bias\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef predict_with_tta(model, df, data_dir, device):\n    model.eval()\n    configs = [('1', None), ('1', 'h'), ('1', 'v'), ('2', None), ('2', 'h'), ('2', 'v')]\n    all_probs = []\n    ids = None\n    \n    for site, flip in configs:\n        dataset = CellularDataset(df, data_dir, mode='test', site=site)\n        loader = DataLoader(dataset, batch_size=Config.BATCH_SIZE,\n                          shuffle=False, num_workers=Config.NUM_WORKERS)\n        \n        probs = []\n        batch_ids = []\n        with torch.no_grad():\n            for imgs, img_ids in tqdm(loader, desc=f'Site{site} {flip or \"\"}'):\n                imgs = imgs.to(device)\n                if flip == 'h':\n                    imgs = torch.flip(imgs, dims=[3])\n                elif flip == 'v':\n                    imgs = torch.flip(imgs, dims=[2])\n                outputs = model(imgs)\n                probs.append(outputs.softmax(dim=1).cpu())\n                batch_ids.extend(img_ids)\n        \n        probs = torch.cat(probs, dim=0)\n        all_probs.append(probs)\n        if ids is None:\n            ids = batch_ids\n    \n    avg_probs = torch.stack(all_probs).mean(dim=0)\n    preds = avg_probs.argmax(dim=1).numpy()\n    return preds, ids\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types in train: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    print(f\"Number of classes: {Config.NUM_CLASSES}\")\n    \n    train_data, val_data = train_test_split(\n        train_df, test_size=0.15, random_state=Config.SEED,\n        stratify=train_df['label']\n    )\n    \n    train_loader = DataLoader(\n        CellularDataset(train_data, Config.DATA_DIR, mode='train'),\n        batch_size=Config.BATCH_SIZE, shuffle=True,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        CellularDataset(val_data, Config.DATA_DIR, mode='val', site='1'),\n        batch_size=Config.BATCH_SIZE, shuffle=False,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    \n    print(f\"Creating model: {Config.MODEL_NAME}...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss(label_smoothing=Config.LABEL_SMOOTHING)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optimizer, T_0=Config.EPOCHS, T_mult=1\n    )\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_efficientnet_b0.pth')\n            print(f\"Saved best model: {best_acc:.2f}%\")\n    \n    print(\"\\nInference with TTA (dual site + fl)...\")\n    model.load_state_dict(torch.load('best_efficientnet_b0.pth'))\n    \n    predictions, ids = predict_with_tta(model, test_df, Config.DATA_DIR, device)\n    \n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    predictions = [label_to_sirna[p] for p in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions\n    })\n    submission.to_csv('submission_efficientnet_b0.csv', index=False)\n    print(f\"Saved submission_efficientnet_b0.csv | Best val acc: {best_acc:.2f}%\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-23T19:25:54.646964Z","iopub.execute_input":"2026-04-23T19:25:54.64744Z","iopub.status.idle":"2026-04-23T23:25:36.702645Z","shell.execute_reply.started":"2026-04-23T19:25:54.647404Z","shell.execute_reply":"2026-04-23T23:25:36.701811Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'resnext50_32x4d'\n    IMG_SIZE = 384\n    BATCH_SIZE = 24\n    EPOCHS = 10\n    LR = 3e-4\n    LABEL_SMOOTHING = 0.1\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train', site=None):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode = mode\n        self.site = site\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row, site):\n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'test':\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s{site}_w'\n        else:\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s{site}_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            else:\n                img = np.zeros((512, 512), dtype=np.uint8)\n            channels.append(img)\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        return img.astype(np.float32) / 255.0\n    \n    def augment(self, img):\n        if self.mode != 'train':\n            return img\n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        if np.random.rand() > 0.5:\n            k = np.random.randint(1, 4)\n            img = np.rot90(img, k=k).copy()\n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        \n        if self.mode == 'train':\n            site = np.random.choice(['1', '2'])\n        else:\n            site = self.site if self.site is not None else '1'\n            \n        img = self.load_image(row, site)\n        img = self.augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode in ['train', 'val']:\n            return img, row['label']\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'conv1'):\n            old_conv = self.backbone.conv1\n            self.backbone.conv1 = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            with torch.no_grad():\n                self.backbone.conv1.weight[:, :3] = old_conv.weight\n                self.backbone.conv1.weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.conv1.bias = old_conv.bias\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': running_loss/len(loader), 'acc': 100.*correct/total})\n    \n    return running_loss/len(loader), 100.*correct/total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss/len(loader), 100.*correct/total\ndef predict_with_tta(model, df, data_dir, device):\n    model.eval()\n    configs = [('1', None), ('1', 'h'), ('1', 'v'), ('2', None), ('2', 'h'), ('2', 'v')]\n    all_probs = []\n    ids = None\n    \n    for site, flip in configs:\n        dataset = CellularDataset(df, data_dir, mode='test', site=site)\n        loader = DataLoader(dataset, batch_size=Config.BATCH_SIZE,\n                          shuffle=False, num_workers=Config.NUM_WORKERS)\n        \n        probs = []\n        batch_ids = []\n        with torch.no_grad():\n            for imgs, img_ids in tqdm(loader, desc=f'Site{site} {flip or \"\"}'):\n                imgs = imgs.to(device)\n                if flip == 'h':\n                    imgs = torch.flip(imgs, dims=[3])\n                elif flip == 'v':\n                    imgs = torch.flip(imgs, dims=[2])\n                outputs = model(imgs)\n                probs.append(outputs.softmax(dim=1).cpu())\n                batch_ids.extend(img_ids)\n        \n        probs = torch.cat(probs, dim=0)\n        all_probs.append(probs)\n        if ids is None:\n            ids = batch_ids\n    \n    avg_probs = torch.stack(all_probs).mean(dim=0)\n    preds = avg_probs.argmax(dim=1).numpy()\n    return preds, ids\ndef main():\n    print(\"Loading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types in train: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    print(f\"Number of classes: {Config.NUM_CLASSES}\")\n    \n    train_data, val_data = train_test_split(\n        train_df, test_size=0.15, random_state=Config.SEED,\n        stratify=train_df['label']\n    )\n    \n    train_loader = DataLoader(\n        CellularDataset(train_data, Config.DATA_DIR, mode='train'),\n        batch_size=Config.BATCH_SIZE, shuffle=True,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        CellularDataset(val_data, Config.DATA_DIR, mode='val', site='1'),\n        batch_size=Config.BATCH_SIZE, shuffle=False,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    \n    print(f\"Creating model: {Config.MODEL_NAME}...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss(label_smoothing=Config.LABEL_SMOOTHING)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n        optimizer, T_0=Config.EPOCHS, T_mult=1\n    )\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_resnext50_32x4d.pth')\n            print(f\"Saved best model: {best_acc:.2f}%\")\n    \n    print(\"\\nInference with TTA (dual site + flips)...\")\n    model.load_state_dict(torch.load('best_resnext50_32x4d.pth'))\n    \n    predictions, ids = predict_with_tta(model, test_df, Config.DATA_DIR, device)\n    \n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    predictions = [label_to_sirna[p] for p in predictions]\n    \n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': predictions\n    })\n    submission.to_csv('submission_resnext50_32x4d.csv', index=False)\n    print(f\"Saved submission_resnext50_32x4d.csv | Best val acc: {best_acc:.2f}%\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T16:39:23.502779Z","iopub.execute_input":"2026-04-24T16:39:23.50319Z","iopub.status.idle":"2026-04-24T20:01:23.541811Z","shell.execute_reply.started":"2026-04-24T16:39:23.503146Z","shell.execute_reply":"2026-04-24T20:01:23.540126Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'resnet50'\n    IMG_SIZE = 320\n    BATCH_SIZE = 32\n    EPOCHS = 20\n    LR = 3e-4\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    CELL_TYPES = ['HUVEC']\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode = mode\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        \n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'test':\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            else:\n                img = np.zeros((512, 512), dtype=np.uint8)\n            channels.append(img)\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        return img.astype(np.float32) / 255.0\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        img = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train':\n            return img, row['label']\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        \n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'conv1'):\n            old_conv = self.backbone.conv1\n            self.backbone.conv1 = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            \n            with torch.no_grad():\n                self.backbone.conv1.weight[:, :3] = old_conv.weight\n                self.backbone.conv1.weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.conv1.bias = old_conv.bias\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': f'{running_loss/len(loader):.4f}', 'acc': f'{100.*correct/total:.2f}%'})\n    \n    return running_loss / len(loader), 100. * correct / total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss / len(loader), 100. * correct / total\ndef main():\n    print(\"=\" * 60)\n    print(\"ResNet50 Baseline - No Augmentation\")\n    print(\"=\" * 60)\n    \n    print(\"\\nLoading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    train_df = train_df[train_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    print(f\"Number of classes: {Config.NUM_CLASSES}\")\n    \n    train_data, val_data = train_test_split(\n        train_df, test_size=0.15, random_state=Config.SEED, stratify=train_df['label']\n    )\n    print(f\"Train: {len(train_data)}, Val: {len(val_data)}\")\n    \n    train_loader = DataLoader(\n        CellularDataset(train_data, Config.DATA_DIR, mode='train'),\n        batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    val_loader = DataLoader(\n        CellularDataset(val_data, Config.DATA_DIR, mode='train'),\n        batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    \n    print(f\"\\nCreating model: {Config.MODEL_NAME}...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        \n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_resnet50_baseline.pth')\n            print(f\"Saved best model: {best_acc:.2f}%\")\n    \n    print(\"\\nInference...\")\n    model.load_state_dict(torch.load('best_resnet50_baseline.pth'))\n    \n    test_loader = DataLoader(\n        CellularDataset(test_df, Config.DATA_DIR, mode='test'),\n        batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=Config.NUM_WORKERS\n    )\n    \n    model.eval()\n    predictions, ids = [], []\n    \n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Inference'):\n            imgs = imgs.to(device)\n            outputs = model(imgs)\n            _, preds = outputs.max(1)\n            predictions.extend(preds.cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    predictions = [label_to_sirna[p] for p in predictions]\n    \n    submission = pd.DataFrame({'id_code': ids, 'sirna': predictions})\n    submission.to_csv('submission_resnet50_baseline.csv', index=False)\n    \n    print(f\"\\n✓ Complete! Best Val Acc: {best_acc:.2f}%\")\n    print(\"Saved submission_resnet50_baseline.csv\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T20:03:16.951437Z","iopub.execute_input":"2026-04-24T20:03:16.951805Z","iopub.status.idle":"2026-04-24T21:55:46.665402Z","shell.execute_reply.started":"2026-04-24T20:03:16.951765Z","shell.execute_reply":"2026-04-24T21:55:46.664142Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\n    \n    MODEL_NAME = 'convnextv2_nano'\n    IMG_SIZE = 256\n    BATCH_SIZE = 16\n    EPOCHS = 10\n    LR = 3e-4\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    CELL_TYPES = ['HUVEC']\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\nset_seed(Config.SEED)\nclass CellularDataset(Dataset):\n    \n    \n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode = mode\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def load_image(self, row):\n        \n        exp = row['experiment']\n        plate = row['plate']\n        well = row['well']\n        \n        if self.mode == 'test':\n            path_template = f'{self.data_dir}/test/{exp}/Plate{plate}/{well}_s1_w'\n        else:\n            path_template = f'{self.data_dir}/train/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img_path = f'{path_template}{i}.png'\n            if os.path.exists(img_path):\n                img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            else:\n                img = np.zeros((512, 512), dtype=np.uint8)\n            channels.append(img)\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        return img.astype(np.float32) / 255.0\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.load_image(row)\n        img = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train':\n            return img, row['label']\n        else:\n            return img, row['id_code']\nclass ConvNeXtV2Model(nn.Module):\n    \n    \n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        \n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3, num_classes=0)\n        \n        if hasattr(self.backbone, 'stem'):\n            old_conv = self.backbone.stem[0]\n            \n            new_conv = nn.Conv2d(\n                in_channels,\n                old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            \n            with torch.no_grad():\n                new_weight = old_conv.weight.mean(dim=1, keepdim=True)\n                new_weight = new_weight.repeat(1, in_channels, 1, 1) / in_channels\n                new_conv.weight.copy_(new_weight)\n                if old_conv.bias is not None:\n                    new_conv.bias = old_conv.bias\n            \n            self.backbone.stem[0] = new_conv\n        \n        with torch.no_grad():\n            dummy = torch.randn(1, in_channels, Config.IMG_SIZE, Config.IMG_SIZE)\n            out = self.backbone(dummy)\n            if len(out.shape) == 4:\n                n_features = out.shape[1]\n            else:\n                n_features = out.shape[1]\n            del dummy, out\n            torch.cuda.empty_cache()\n        \n        print(f\"Detected feature dimension: {n_features}\")\n        \n        self.pool = nn.AdaptiveAvgPool2d(1)\n        \n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Dropout(0.3),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        \n        if len(features.shape) == 4:\n            features = self.pool(features)\n        \n        return self.classifier(features)\ndef train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n        \n        pbar.set_postfix({'loss': f'{running_loss/len(loader):.4f}', 'acc': f'{100.*correct/total:.2f}%'})\n    \n    return running_loss / len(loader), 100. * correct / total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return running_loss / len(loader), 100. * correct / total\ndef main():\n    torch.cuda.empty_cache()\n    \n    print(\"=\" * 60)\n    print(\"ConvNeXt V2 - Modern ConvNet (No Augmentation)\")\n    print(\"=\" * 60)\n    \n    print(\"\\nLoading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    train_df = train_df[train_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(unique_sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    print(f\"Number of classes: {Config.NUM_CLASSES}\")\n    \n    train_data, val_data = train_test_split(\n        train_df, test_size=0.15, random_state=Config.SEED, stratify=train_df['label']\n    )\n    print(f\"Train: {len(train_data)}, Val: {len(val_data)}\")\n    \n    train_loader = DataLoader(\n        CellularDataset(train_data, Config.DATA_DIR, mode='train'),\n        batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    val_loader = DataLoader(\n        CellularDataset(val_data, Config.DATA_DIR, mode='train'),\n        batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    \n    print(f\"\\nCreating model: {Config.MODEL_NAME}...\")\n    model = ConvNeXtV2Model(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n    \n    best_acc = 0\n    for epoch in range(Config.EPOCHS):\n        torch.cuda.empty_cache()\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        \n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}, Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_convnextv2.pth')\n            print(f\"Saved best model: {best_acc:.2f}%\")\n    \n    print(\"\\nInference...\")\n    model.load_state_dict(torch.load('best_convnextv2.pth'))\n    \n    test_loader = DataLoader(\n        CellularDataset(test_df, Config.DATA_DIR, mode='test'),\n        batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=Config.NUM_WORKERS\n    )\n    \n    model.eval()\n    predictions, ids = [], []\n    \n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Inference'):\n            imgs = imgs.to(device)\n            outputs = model(imgs)\n            _, preds = outputs.max(1)\n            predictions.extend(preds.cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    predictions = [label_to_sirna[p] for p in predictions]\n    \n    submission = pd.DataFrame({'id_code': ids, 'sirna': predictions})\n    submission.to_csv('submission_convnextv2.csv', index=False)\n    \n    print(f\"\\n✓ Complete! Best Val Acc: {best_acc:.2f}%\")\n    print(\"Saved submission_convnextv2.csv\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T22:05:16.678336Z","iopub.execute_input":"2026-04-24T22:05:16.678681Z","iopub.status.idle":"2026-04-24T22:34:57.22756Z","shell.execute_reply.started":"2026-04-24T22:05:16.678647Z","shell.execute_reply":"2026-04-24T22:34:57.226304Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nDATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\nTRAIN_CSV = f'{DATA_DIR}/train.csv'\nTEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\nMODEL_NAME = 'densenet121'\nIMG_SIZE = 320\nBATCH_SIZE = 32\nEPOCHS = 20\nLR = 3e-4\nNUM_WORKERS = 2\nSEED = 42\nCELL_TYPES = ['HUVEC']\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\nclass SimpleDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode = mode\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        exp, plate, well = row['experiment'], row['plate'], row['well']\n        \n        prefix = 'test' if self.mode == 'test' else 'train'\n        base = f'{self.data_dir}/{prefix}/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img = cv2.imread(f'{base}{i}.png', cv2.IMREAD_GRAYSCALE)\n            channels.append(img if img is not None else np.zeros((512, 512), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE)).astype(np.float32) / 255.0\n        img = torch.from_numpy(img).permute(2, 0, 1)\n        \n        if self.mode == 'train':\n            return img, row['label']\n        return img, row['id_code']\ndef create_model(num_classes):\n    model = timm.create_model(MODEL_NAME, pretrained=True, in_chans=3)\n    \n    old_conv = model.features.conv0\n    \n    out_ch = int(old_conv.out_channels)\n    kernel = int(old_conv.kernel_size[0]) if hasattr(old_conv.kernel_size, '__len__') else int(old_conv.kernel_size)\n    stride = int(old_conv.stride[0]) if hasattr(old_conv.stride, '__len__') else int(old_conv.stride)\n    padding = int(old_conv.padding[0]) if hasattr(old_conv.padding, '__len__') else int(old_conv.padding)\n    bias = old_conv.bias is not None\n    \n    model.features.conv0 = nn.Conv2d(6, out_ch, kernel, stride, padding, bias=bias)\n    \n    with torch.no_grad():\n        w = old_conv.weight.mean(dim=1, keepdim=True).repeat(1, 6, 1, 1) / 6\n        model.features.conv0.weight.copy_(w)\n    \n    n_features = model.get_classifier().in_features\n    model.reset_classifier(0)\n    model.classifier = nn.Sequential(nn.Dropout(0.3), nn.Linear(n_features, num_classes))\n    \n    return model\ndef train_epoch(model, loader, criterion, optimizer):\n    model.train()\n    loss_sum, correct, total = 0, 0, 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        out = model(imgs)\n        loss = criterion(out, labels)\n        loss.backward()\n        optimizer.step()\n        \n        loss_sum += loss.item()\n        correct += (out.argmax(1) == labels).sum().item()\n        total += labels.size(0)\n        \n        pbar.set_postfix({\n            'loss': f'{loss_sum/len(loader):.4f}', \n            'acc': f'{100.*correct/total:.2f}%'\n        })\n    \n    return loss_sum / len(loader), 100 * correct / total\ndef validate(model, loader, criterion):\n    model.eval()\n    loss_sum, correct, total = 0, 0, 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Val'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            out = model(imgs)\n            loss = criterion(out, labels)\n            \n            loss_sum += loss.item()\n            correct += (out.argmax(1) == labels).sum().item()\n            total += labels.size(0)\n    \n    return loss_sum / len(loader), 100 * correct / total\ndef main():\n    print(\"=\" * 60)\n    print(\"DenseNet121 - No Augmentation (Procedural)\")\n    print(\"=\" * 60)\n    \n    print(\"\\nLoading data...\")\n    train_df = pd.read_csv(TRAIN_CSV)\n    test_df = pd.read_csv(TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    train_df = train_df[train_df['cell_type'].isin(CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    num_classes = len(sirnas)\n    print(f\"Classes: {num_classes}\")\n    \n    train_data, val_data = train_test_split(train_df, test_size=0.15, \n                                            stratify=train_df['label'], random_state=SEED)\n    \n    train_loader = DataLoader(SimpleDataset(train_data, DATA_DIR), \n                              batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True)\n    val_loader = DataLoader(SimpleDataset(val_data, DATA_DIR), \n                            batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True)\n    \n    print(f\"\\nCreating {MODEL_NAME}...\")\n    model = create_model(num_classes).to(device)\n    \n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n    \n    best_acc, patience, counter = 0, 3, 0\n    \n    for epoch in range(EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer)\n        val_loss, val_acc = validate(model, val_loader, criterion)\n        scheduler.step()\n        \n        print(f\"Train: {train_loss:.4f}, {train_acc:.2f}% | Val: {val_loss:.4f}, {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            counter = 0\n            torch.save(model.state_dict(), 'best_densenet121.pth')\n            print(f\"  → Saved best: {best_acc:.2f}%\")\n        else:\n            counter += 1\n            print(f\"  → No improvement ({counter}/{patience})\")\n            if counter >= patience:\n                print(\"Early stopping!\")\n                break\n    \n    print(\"\\nInference...\")\n    model.load_state_dict(torch.load('best_densenet121.pth'))\n    model.eval()\n    \n    test_loader = DataLoader(SimpleDataset(test_df, DATA_DIR, mode='test'),\n                             batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS)\n    \n    preds, ids = [], []\n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Test'):\n            out = model(imgs.to(device))\n            preds.extend(out.argmax(1).cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    submission = pd.DataFrame({'id_code': ids, 'sirna': [label_to_sirna[p] for p in preds]})\n    submission.to_csv('submission_densenet121.csv', index=False)\n    \n    print(f\"\\n✓ Done! Best Val Acc: {best_acc:.2f}%\")\n    print(f\"Saved submission_densenet121.csv\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-24T22:37:39.126098Z","iopub.execute_input":"2026-04-24T22:37:39.126456Z","iopub.status.idle":"2026-04-25T00:26:03.466951Z","shell.execute_reply.started":"2026-04-24T22:37:39.126421Z","shell.execute_reply":"2026-04-25T00:26:03.466088Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nDATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\nTRAIN_CSV = f'{DATA_DIR}/train.csv'\nTEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\nMODEL_NAME = 'vit_base_patch16_224'\nIMG_SIZE = 224\nBATCH_SIZE = 32\nEPOCHS = 20\nLR = 1e-4\nNUM_WORKERS = 2\nSEED = 42\nCELL_TYPES = ['HUVEC']\nGRAD_CLIP = 1.0\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\nIMAGENET_MEAN = torch.tensor([0.485, 0.456, 0.406, 0.485, 0.456, 0.406]).view(6, 1, 1)\nIMAGENET_STD = torch.tensor([0.229, 0.224, 0.225, 0.229, 0.224, 0.225]).view(6, 1, 1)\nclass SimpleDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode = mode\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def augment(self, img):\n        \n        if self.mode != 'train':\n            return img\n        \n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        \n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        \n        if np.random.rand() > 0.5:\n            k = np.random.randint(1, 4)\n            img = np.rot90(img, k=k).copy()\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        exp, plate, well = row['experiment'], row['plate'], row['well']\n        \n        prefix = 'test' if self.mode == 'test' else 'train'\n        base = f'{self.data_dir}/{prefix}/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img = cv2.imread(f'{base}{i}.png', cv2.IMREAD_GRAYSCALE)\n            channels.append(img if img is not None else np.zeros((512, 512), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        \n        img = self.augment(img)\n        \n        img = torch.from_numpy(img).permute(2, 0, 1).float() / 255.0\n        \n        img = (img - IMAGENET_MEAN) / IMAGENET_STD\n        \n        if self.mode == 'train':\n            return img, row['label']\n        return img, row['id_code']\nclass ViTWrapper(nn.Module):\n    \n    def __init__(self, backbone, num_classes):\n        super().__init__()\n        self.backbone = backbone\n        n_features = backbone.num_features\n        self.classifier = nn.Sequential(\n            nn.LayerNorm(n_features),\n            nn.Dropout(0.2),\n            nn.Linear(n_features, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\ndef create_model(num_classes):\n    backbone = timm.create_model(MODEL_NAME, pretrained=True, in_chans=3, num_classes=0)\n    \n    old_conv = backbone.patch_embed.proj\n    \n    out_ch = int(old_conv.out_channels)\n    kernel = int(old_conv.kernel_size[0]) if hasattr(old_conv.kernel_size, '__len__') else int(old_conv.kernel_size)\n    stride = int(old_conv.stride[0]) if hasattr(old_conv.stride, '__len__') else int(old_conv.stride)\n    padding = int(old_conv.padding[0]) if hasattr(old_conv.padding, '__len__') else int(old_conv.padding)\n    \n    backbone.patch_embed.proj = nn.Conv2d(6, out_ch, kernel, stride, padding)\n    \n    with torch.no_grad():\n        w = old_conv.weight.mean(dim=1, keepdim=True).repeat(1, 6, 1, 1)\n        backbone.patch_embed.proj.weight.copy_(w)\n        if old_conv.bias is not None:\n            backbone.patch_embed.proj.bias.copy_(old_conv.bias)\n    \n    model = ViTWrapper(backbone, num_classes)\n    return model\ndef train_epoch(model, loader, criterion, optimizer, scaler, num_classes):\n    model.train()\n    loss_sum, correct, total = 0, 0, 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        \n        with autocast():\n            out = model(imgs)\n            loss = criterion(out, labels)\n        \n        scaler.scale(loss).backward()\n        \n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n        \n        scaler.step(optimizer)\n        scaler.update()\n        \n        loss_sum += loss.item()\n        correct += (out.argmax(1) == labels).sum().item()\n        total += labels.size(0)\n        \n        pbar.set_postfix({\n            'loss': f'{loss_sum/len(loader):.4f}',\n            'acc': f'{100.*correct/total:.2f}%'\n        })\n    \n    return loss_sum / len(loader), 100 * correct / total\ndef validate(model, loader, criterion):\n    model.eval()\n    loss_sum, correct, total = 0, 0, 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            \n            with autocast():\n                out = model(imgs)\n                loss = criterion(out, labels)\n            \n            loss_sum += loss.item()\n            correct += (out.argmax(1) == labels).sum().item()\n            total += labels.size(0)\n    \n    return loss_sum / len(loader), 100 * correct / total\ndef main():\n    print(\"=\" * 60)\n    print(\"ViT-Base/16 - Optimized (With Augmentation + Mixed Precision)\")\n    print(\"=\" * 60)\n    \n    print(\"\\nLoading data...\")\n    train_df = pd.read_csv(TRAIN_CSV)\n    test_df = pd.read_csv(TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    train_df = train_df[train_df['cell_type'].isin(CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    num_classes = len(sirnas)\n    \n    print(f\"Classes: {num_classes}\")\n    \n    train_data, val_data = train_test_split(\n        train_df, test_size=0.15, random_state=SEED\n    )\n    \n    train_loader = DataLoader(\n        SimpleDataset(train_data, DATA_DIR, mode='train'),\n        batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True\n    )\n    val_loader = DataLoader(\n        SimpleDataset(val_data, DATA_DIR, mode='train'),\n        batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True\n    )\n    \n    print(f\"\\nCreating {MODEL_NAME}...\")\n    model = create_model(num_classes).to(device)\n    \n    criterion = nn.CrossEntropyLoss(label_smoothing=0.1)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n    scaler = GradScaler()\n    \n    best_acc, patience, counter = 0, 3, 0\n    \n    for epoch in range(EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, scaler, num_classes)\n        val_loss, val_acc = validate(model, val_loader, criterion)\n        scheduler.step()\n        \n        print(f\"Train: {train_loss:.4f}, {train_acc:.2f}% | Val: {val_loss:.4f}, {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            counter = 0\n            torch.save(model.state_dict(), 'best_vit_base.pth')\n            print(f\"  → Saved best: {best_acc:.2f}%\")\n        else:\n            counter += 1\n            print(f\"  → No improvement ({counter}/{patience})\")\n            if counter >= patience:\n                print(\"Early stopping!\")\n                break\n    \n    print(\"\\nInference...\")\n    model.load_state_dict(torch.load('best_vit_base.pth'))\n    model.eval()\n    \n    test_loader = DataLoader(\n        SimpleDataset(test_df, DATA_DIR, mode='test'),\n        batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS\n    )\n    \n    preds, ids = [], []\n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Test'):\n            imgs = imgs.to(device)\n            with autocast():\n                out = model(imgs)\n            preds.extend(out.argmax(1).cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': [label_to_sirna[p] for p in preds]\n    })\n    submission.to_csv('submission_vit_base.csv', index=False)\n    \n    print(f\"\\n✓ Done! Best Val Acc: {best_acc:.2f}%\")\n    print(\"Saved submission_vit_base.csv\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T06:15:15.294039Z","iopub.execute_input":"2026-04-25T06:15:15.295254Z","iopub.status.idle":"2026-04-25T06:37:50.601995Z","shell.execute_reply.started":"2026-04-25T06:15:15.2952Z","shell.execute_reply":"2026-04-25T06:37:50.599546Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\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\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nDATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\nTRAIN_CSV = f'{DATA_DIR}/train.csv'\nTEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\nMODEL_NAME = 'densenet121'\nIMG_SIZE = 320\nBATCH_SIZE = 32\nEPOCHS = 20\nLR = 3e-4\nNUM_WORKERS = 2\nSEED = 42\nCELL_TYPES = ['HUVEC']\nGRAD_CLIP = 1.0\nCHAOTIC_MAP = 'skew_tent'\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\nclass SimpleDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode = mode\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def augment(self, img):\n        \n        if self.mode != 'train':\n            return img\n        \n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        exp, plate, well = row['experiment'], row['plate'], row['well']\n        \n        prefix = 'test' if self.mode == 'test' else 'train'\n        base = f'{self.data_dir}/{prefix}/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img = cv2.imread(f'{base}{i}.png', cv2.IMREAD_GRAYSCALE)\n            channels.append(img if img is not None else np.zeros((512, 512), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = self.augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1).float() / 255.0\n        \n        if self.mode == 'train':\n            return img, row['label']\n        return img, row['id_code']\nclass ChaoticTransform(nn.Module):\n    \n    def __init__(self, map_type='logistic', r=4.0, p=0.499):\n        super().__init__()\n        self.map_type = map_type\n        self.r = r\n        self.p = p\n    \n    def logistic_map(self, x):\n        \n        return self.r * x * (1 - x)\n    \n    def skew_tent_map(self, x):\n        \n        mask = x < self.p\n        result = torch.zeros_like(x)\n        result[mask] = x[mask] / self.p\n        result[~mask] = (1 - x[~mask]) / (1 - self.p)\n        return result\n    \n    def sine_map(self, x):\n        \n        return torch.sin(np.pi * x)\n    \n    def forward(self, x):\n        x_min = x.min(dim=-1, keepdim=True)[0].min(dim=-2, keepdim=True)[0]\n        x_max = x.max(dim=-1, keepdim=True)[0].max(dim=-2, keepdim=True)[0]\n        x_norm = (x - x_min) / (x_max - x_min + 1e-8)\n        \n        if self.map_type == 'logistic':\n            return self.logistic_map(x_norm)\n        elif self.map_type == 'skew_tent':\n            return self.skew_tent_map(x_norm)\n        elif self.map_type == 'sine':\n            return self.sine_map(x_norm)\n        else:\n            return x_norm\nclass ChaoticCNN(nn.Module):\n    \n    def __init__(self, num_classes, map_type='logistic'):\n        super().__init__()\n        \n        self.backbone = timm.create_model(MODEL_NAME, pretrained=True, in_chans=3, num_classes=0)\n        \n        old_conv = self.backbone.features.conv0\n        self.backbone.features.conv0 = nn.Conv2d(\n            6, old_conv.out_channels,\n            kernel_size=old_conv.kernel_size,\n            stride=old_conv.stride,\n            padding=old_conv.padding,\n            bias=old_conv.bias is not None\n        )\n        \n        with torch.no_grad():\n            w = old_conv.weight.mean(dim=1, keepdim=True).repeat(1, 6, 1, 1)\n            self.backbone.features.conv0.weight.copy_(w)\n        \n        with torch.no_grad():\n            dummy = torch.randn(1, 6, 32, 32)\n            out = self.backbone(dummy)\n            if len(out.shape) == 4:\n                self.n_features = out.shape[1]\n            else:\n                self.n_features = out.shape[1]\n            print(f\"Detected feature dimension: {self.n_features}\")\n        \n        self.chaotic = ChaoticTransform(map_type=map_type)\n        \n        self.classifier = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.Dropout(0.3),\n            nn.Linear(self.n_features, num_classes)\n        )\n        \n        print(f\"ChaoticCNN initialized: n_features={self.n_features}, map={map_type}\")\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        \n        if len(features.shape) == 2:\n            features = features.unsqueeze(-1).unsqueeze(-1)\n        \n        features = self.chaotic(features)\n        \n        out = self.classifier(features)\n        return out\ndef create_model(num_classes):\n    return ChaoticCNN(num_classes=num_classes, map_type=CHAOTIC_MAP).to(device)\ndef train_epoch(model, loader, criterion, optimizer, scaler):\n    model.train()\n    loss_sum, correct, total = 0, 0, 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        \n        with autocast():\n            out = model(imgs)\n            loss = criterion(out, labels)\n        \n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n        scaler.step(optimizer)\n        scaler.update()\n        \n        loss_sum += loss.item()\n        correct += (out.argmax(1) == labels).sum().item()\n        total += labels.size(0)\n        \n        pbar.set_postfix({\n            'loss': f'{loss_sum/len(loader):.4f}',\n            'acc': f'{100.*correct/total:.2f}%'\n        })\n    \n    return loss_sum / len(loader), 100 * correct / total\ndef validate(model, loader, criterion):\n    model.eval()\n    loss_sum, correct, total = 0, 0, 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            \n            with autocast():\n                out = model(imgs)\n                loss = criterion(out, labels)\n            \n            loss_sum += loss.item()\n            correct += (out.argmax(1) == labels).sum().item()\n            total += labels.size(0)\n    \n    return loss_sum / len(loader), 100 * correct / total\ndef main():\n    print(\"=\" * 60)\n    print(f\"Chaotic CNN - DenseNet + {CHAOTIC_MAP.upper()} Map\")\n    print(\"Based on arXiv:2604.14645\")\n    print(\"=\" * 60)\n    \n    print(\"\\nLoading data...\")\n    train_df = pd.read_csv(TRAIN_CSV)\n    test_df = pd.read_csv(TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    train_df = train_df[train_df['cell_type'].isin(CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    num_classes = len(sirnas)\n    print(f\"Classes: {num_classes}\")\n    \n    train_data, val_data = train_test_split(\n        train_df, test_size=0.15, random_state=SEED\n    )\n    \n    train_loader = DataLoader(\n        SimpleDataset(train_data, DATA_DIR, mode='train'),\n        batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True\n    )\n    val_loader = DataLoader(\n        SimpleDataset(val_data, DATA_DIR, mode='train'),\n        batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True\n    )\n    \n    print(f\"\\nCreating Chaotic CNN with {CHAOTIC_MAP} map...\")\n    model = create_model(num_classes)\n    \n    criterion = nn.CrossEntropyLoss(label_smoothing=0.1)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n    scaler = GradScaler()\n    \n    best_acc, patience, counter = 0, 3, 0\n    \n    for epoch in range(EPOCHS):\n        torch.cuda.empty_cache()\n        print(f\"\\nEpoch {epoch+1}/{EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, scaler)\n        val_loss, val_acc = validate(model, val_loader, criterion)\n        scheduler.step()\n        \n        print(f\"Train: {train_loss:.4f}, {train_acc:.2f}% | Val: {val_loss:.4f}, {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            counter = 0\n            torch.save(model.state_dict(), f'best_chaotic_cnn_{CHAOTIC_MAP}.pth')\n            print(f\"  → Saved best: {best_acc:.2f}%\")\n        else:\n            counter += 1\n            print(f\"  → No improvement ({counter}/{patience})\")\n            if counter >= patience:\n                print(\"Early stopping!\")\n                break\n    \n    print(\"\\nInference...\")\n    model.load_state_dict(torch.load(f'best_chaotic_cnn_{CHAOTIC_MAP}.pth'))\n    model.eval()\n    \n    test_loader = DataLoader(\n        SimpleDataset(test_df, DATA_DIR, mode='test'),\n        batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS\n    )\n    \n    preds, ids = [], []\n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Test'):\n            with autocast():\n                out = model(imgs.to(device))\n            preds.extend(out.argmax(1).cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': [label_to_sirna[p] for p in preds]\n    })\n    submission.to_csv(f'submission_chaotic_cnn_{CHAOTIC_MAP}.csv', index=False)\n    \n    print(f\"\\n✓ Done! Best Val Acc: {best_acc:.2f}%\")\n    print(f\"Saved submission_chaotic_cnn_{CHAOTIC_MAP}.csv\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T07:07:36.951577Z","iopub.execute_input":"2026-04-25T07:07:36.952055Z","iopub.status.idle":"2026-04-25T07:07:37.966966Z","shell.execute_reply.started":"2026-04-25T07:07:36.952019Z","shell.execute_reply":"2026-04-25T07:07:37.965827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\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\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nDATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\nTRAIN_CSV = f'{DATA_DIR}/train.csv'\nTEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\nMODEL_NAME = 'densenet121'\nIMG_SIZE = 320\nBATCH_SIZE = 32\nEPOCHS = 20\nLR = 3e-4\nNUM_WORKERS = 2\nSEED = 42\nCELL_TYPES = ['HUVEC']\nGRAD_CLIP = 1.0\nCHAOTIC_MAP = 'skew_tent'\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\nclass SimpleDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode = mode\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def augment(self, img):\n        \n        if self.mode != 'train':\n            return img\n        \n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        \n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        exp, plate, well = row['experiment'], row['plate'], row['well']\n        \n        prefix = 'test' if self.mode == 'test' else 'train'\n        base = f'{self.data_dir}/{prefix}/{exp}/Plate{plate}/{well}_s1_w'\n        \n        channels = []\n        for i in range(1, 7):\n            img = cv2.imread(f'{base}{i}.png', cv2.IMREAD_GRAYSCALE)\n            channels.append(img if img is not None else np.zeros((512, 512), dtype=np.uint8))\n        \n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = self.augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1).float() / 255.0\n        \n        if self.mode == 'train':\n            return img, row['label']\n        return img, row['id_code']\nclass ChaoticTransform(nn.Module):\n    \n    def __init__(self, map_type='logistic', r=4.0, p=0.499):\n        super().__init__()\n        self.map_type = map_type\n        self.r = r\n        self.p = p\n    \n    def logistic_map(self, x):\n        \n        return self.r * x * (1 - x)\n    \n    def skew_tent_map(self, x):\n        \n        mask = x < self.p\n        result = torch.zeros_like(x)\n        result[mask] = x[mask] / self.p\n        result[~mask] = (1 - x[~mask]) / (1 - self.p)\n        return result\n    \n    def sine_map(self, x):\n        \n        return torch.sin(np.pi * x)\n    \n    def forward(self, x):\n        \n        if len(x.shape) == 2:\n            x_min = x.min(dim=1, keepdim=True)[0]\n            x_max = x.max(dim=1, keepdim=True)[0]\n        else:\n            B, C, H, W = x.shape\n            x_flat = x.view(B, C, -1)\n            x_min = x_flat.min(dim=2, keepdim=True)[0].view(B, C, 1, 1)\n            x_max = x_flat.max(dim=2, keepdim=True)[0].view(B, C, 1, 1)\n        \n        x_norm = (x - x_min) / (x_max - x_min + 1e-8)\n        \n        if self.map_type == 'logistic':\n            return self.logistic_map(x_norm)\n        elif self.map_type == 'skew_tent':\n            return self.skew_tent_map(x_norm)\n        elif self.map_type == 'sine':\n            return self.sine_map(x_norm)\n        else:\n            return x_norm\nclass ChaoticCNN(nn.Module):\n    \n    def __init__(self, num_classes, map_type='logistic'):\n        super().__init__()\n        \n        self.backbone = timm.create_model(MODEL_NAME, pretrained=True, in_chans=3, num_classes=0)\n        \n        old_conv = self.backbone.features.conv0\n        self.backbone.features.conv0 = nn.Conv2d(\n            6, old_conv.out_channels,\n            kernel_size=old_conv.kernel_size,\n            stride=old_conv.stride,\n            padding=old_conv.padding,\n            bias=old_conv.bias is not None\n        )\n        \n        with torch.no_grad():\n            w = old_conv.weight.mean(dim=1, keepdim=True).repeat(1, 6, 1, 1)\n            self.backbone.features.conv0.weight.copy_(w)\n        \n        with torch.no_grad():\n            dummy = torch.randn(1, 6, 32, 32)\n            out = self.backbone(dummy)\n            if len(out.shape) == 4:\n                self.n_features = out.shape[1]\n            else:\n                self.n_features = out.shape[1]\n            print(f\"Detected feature dimension: {self.n_features}\")\n        \n        self.chaotic = ChaoticTransform(map_type=map_type)\n        \n        self.classifier = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.Dropout(0.3),\n            nn.Linear(self.n_features, num_classes)\n        )\n        \n        print(f\"ChaoticCNN initialized: n_features={self.n_features}, map={map_type}\")\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        \n        features = self.chaotic(features)\n        \n        if len(features.shape) == 2:\n            out = self.classifier(features)\n        else:\n            out = self.classifier(features)\n        return out\ndef create_model(num_classes):\n    return ChaoticCNN(num_classes=num_classes, map_type=CHAOTIC_MAP).to(device)\ndef train_epoch(model, loader, criterion, optimizer, scaler):\n    model.train()\n    loss_sum, correct, total = 0, 0, 0\n    \n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        \n        with autocast():\n            out = model(imgs)\n            loss = criterion(out, labels)\n        \n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n        scaler.step(optimizer)\n        scaler.update()\n        \n        loss_sum += loss.item()\n        correct += (out.argmax(1) == labels).sum().item()\n        total += labels.size(0)\n        \n        pbar.set_postfix({\n            'loss': f'{loss_sum/len(loader):.4f}',\n            'acc': f'{100.*correct/total:.2f}%'\n        })\n    \n    return loss_sum / len(loader), 100 * correct / total\ndef validate(model, loader, criterion):\n    model.eval()\n    loss_sum, correct, total = 0, 0, 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Validation'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            \n            with autocast():\n                out = model(imgs)\n                loss = criterion(out, labels)\n            \n            loss_sum += loss.item()\n            correct += (out.argmax(1) == labels).sum().item()\n            total += labels.size(0)\n    \n    return loss_sum / len(loader), 100 * correct / total\ndef main():\n    print(\"=\" * 60)\n    print(f\"Chaotic CNN - DenseNet + {CHAOTIC_MAP.upper()} Map\")\n    print(\"Based on arXiv:2604.14645\")\n    print(\"=\" * 60)\n    \n    print(\"\\nLoading data...\")\n    train_df = pd.read_csv(TRAIN_CSV)\n    test_df = pd.read_csv(TEST_CSV)\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    train_df = train_df[train_df['cell_type'].isin(CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"Training samples: {len(train_df)}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    num_classes = len(sirnas)\n    print(f\"Classes: {num_classes}\")\n    \n    train_data, val_data = train_test_split(\n        train_df, test_size=0.15, random_state=SEED\n    )\n    \n    train_loader = DataLoader(\n        SimpleDataset(train_data, DATA_DIR, mode='train'),\n        batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True\n    )\n    val_loader = DataLoader(\n        SimpleDataset(val_data, DATA_DIR, mode='train'),\n        batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True\n    )\n    \n    print(f\"\\nCreating Chaotic CNN with {CHAOTIC_MAP} map...\")\n    model = create_model(num_classes)\n    \n    criterion = nn.CrossEntropyLoss(label_smoothing=0.1)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n    scaler = GradScaler()\n    \n    best_acc, patience, counter = 0, 3, 0\n    \n    for epoch in range(EPOCHS):\n        torch.cuda.empty_cache()\n        print(f\"\\nEpoch {epoch+1}/{EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, scaler)\n        val_loss, val_acc = validate(model, val_loader, criterion)\n        scheduler.step()\n        \n        print(f\"Train: {train_loss:.4f}, {train_acc:.2f}% | Val: {val_loss:.4f}, {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            counter = 0\n            torch.save(model.state_dict(), f'best_chaotic_cnn_{CHAOTIC_MAP}.pth')\n            print(f\"  → Saved best: {best_acc:.2f}%\")\n        else:\n            counter += 1\n            print(f\"  → No improvement ({counter}/{patience})\")\n            if counter >= patience:\n                print(\"Early stopping!\")\n                break\n    \n    print(\"\\nInference...\")\n    model.load_state_dict(torch.load(f'best_chaotic_cnn_{CHAOTIC_MAP}.pth'))\n    model.eval()\n    \n    test_loader = DataLoader(\n        SimpleDataset(test_df, DATA_DIR, mode='test'),\n        batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS\n    )\n    \n    preds, ids = [], []\n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Test'):\n            with autocast():\n                out = model(imgs.to(device))\n            preds.extend(out.argmax(1).cpu().numpy())\n            ids.extend(img_ids)\n    \n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    submission = pd.DataFrame({\n        'id_code': ids,\n        'sirna': [label_to_sirna[p] for p in preds]\n    })\n    submission.to_csv(f'submission_chaotic_cnn_{CHAOTIC_MAP}.csv', index=False)\n    \n    print(f\"\\n✓ Done! Best Val Acc: {best_acc:.2f}%\")\n    print(f\"Saved submission_chaotic_cnn_{CHAOTIC_MAP}.csv\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T07:10:22.23708Z","iopub.execute_input":"2026-04-25T07:10:22.237932Z","iopub.status.idle":"2026-04-25T07:10:22.708455Z","shell.execute_reply.started":"2026-04-25T07:10:22.237892Z","shell.execute_reply":"2026-04-25T07:10:22.707348Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\nDATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\nTRAIN_CSV = f'{DATA_DIR}/train.csv'\nTEST_CSV = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\nMODEL_NAME = 'densenet121'\nIMG_SIZE = 320\nBATCH_SIZE = 32\nEPOCHS = 20\nLR = 1e-4\nNUM_WORKERS = 2\nSEED = 42\nCELL_TYPES = ['HUVEC']\nGRAD_CLIP = 1.0\nCHAOTIC_MAP = 'skew_tent'\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\nclass SimpleDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode = mode\n    def __len__(self):\n        return len(self.df)\n    def augment(self, img):\n        if self.mode != 'train':\n            return img\n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        return img\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        exp, plate, well = row['experiment'], row['plate'], row['well']\n        prefix = 'test' if self.mode == 'test' else 'train'\n        base = f'{self.data_dir}/{prefix}/{exp}/Plate{plate}/{well}_s1_w'\n        channels = []\n        for i in range(1, 7):\n            img = cv2.imread(f'{base}{i}.png', cv2.IMREAD_GRAYSCALE)\n            if img is None:\n                img = np.zeros((512, 512), dtype=np.uint8)\n            channels.append(img)\n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        img = self.augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1).float() / 255.0\n        mean = torch.tensor([0.5]*6).view(6,1,1)\n        std = torch.tensor([0.25]*6).view(6,1,1)\n        img = (img - mean) / std\n        if self.mode == 'train':\n            return img, row['label']\n        return img, row['id_code']\nclass ChaoticTransform(nn.Module):\n    def __init__(self, map_type='logistic', r=3.8, p=0.499):\n        super().__init__()\n        self.map_type = map_type\n        self.r = r\n        self.p = p\n    def forward(self, x):\n        B, C, H, W = x.shape\n        x_flat = x.view(B, C, -1)\n        x_min = x_flat.min(dim=2)[0].view(B, C, 1, 1)\n        x_max = x_flat.max(dim=2)[0].view(B, C, 1, 1)\n        x = (x - x_min) / (x_max - x_min + 1e-6)\n        x = torch.clamp(x, 1e-6, 1 - 1e-6)\n        if self.map_type == 'logistic':\n            out = self.r * x * (1 - x)\n            return torch.clamp(out, -5, 5)\n        elif self.map_type == 'skew_tent':\n            eps = 1e-6\n            p_safe = max(self.p, eps)\n            p_safe = min(p_safe, 1 - eps)\n            \n            mask = x < p_safe\n            out = torch.zeros_like(x)\n            out[mask] = x[mask] / p_safe\n            out[~mask] = (1 - x[~mask]) / (1 - p_safe)\n            return torch.clamp(out, -5, 5)\n        elif self.map_type == 'sine':\n            return torch.sin(np.pi * x)\n        return x\nclass ChaoticCNN(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        self.backbone = timm.create_model(\n            MODEL_NAME,\n            pretrained=True,\n            in_chans=6,\n            features_only=True\n        )\n        self.n_features = self.backbone.feature_info[-1]['num_chs']\n        self.chaotic = ChaoticTransform(CHAOTIC_MAP)\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.head = nn.Sequential(\n            nn.Flatten(),\n            nn.LayerNorm(self.n_features),\n            nn.Dropout(0.3),\n            nn.Linear(self.n_features, num_classes)\n        )\n    def forward(self, x):\n        features = self.backbone(x)[-1]\n        features = self.chaotic(features)\n        features = self.pool(features)\n        out = self.head(features)\n        return out\ndef train_epoch(model, loader, criterion, optimizer, scaler):\n    model.train()\n    loss_sum, correct, total = 0, 0, 0\n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        with autocast():\n            out = model(imgs)\n            loss = criterion(out, labels)\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n        scaler.step(optimizer)\n        scaler.update()\n        loss_sum += loss.item()\n        correct += (out.argmax(1) == labels).sum().item()\n        total += labels.size(0)\n        pbar.set_postfix({\n            'loss': f'{loss_sum/len(loader):.4f}',\n            'acc': f'{100.*correct/total:.2f}%'\n        })\n    return loss_sum / len(loader), 100 * correct / total\ndef validate(model, loader, criterion):\n    model.eval()\n    loss_sum, correct, total = 0, 0, 0\n    with torch.no_grad():\n        pbar = tqdm(loader, desc='Validation')\n        for imgs, labels in pbar:\n            imgs, labels = imgs.to(device), labels.to(device)\n            with autocast():\n                out = model(imgs)\n                loss = criterion(out, labels)\n            loss_sum += loss.item()\n            correct += (out.argmax(1) == labels).sum().item()\n            total += labels.size(0)\n            pbar.set_postfix({\n                'loss': f'{loss_sum/len(loader):.4f}',\n                'acc': f'{100.*correct/total:.2f}%'\n            })\n    return loss_sum / len(loader), 100 * correct / total\ndef main():\n    train_df = pd.read_csv(TRAIN_CSV)\n    test_df = pd.read_csv(TEST_CSV)\n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    train_df = train_df[train_df['cell_type'].isin(CELL_TYPES)]\n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    num_classes = len(sirnas)\n    train_data, val_data = train_test_split(train_df, test_size=0.15, random_state=SEED)\n    train_loader = DataLoader(\n        SimpleDataset(train_data, DATA_DIR, 'train'),\n        batch_size=BATCH_SIZE,\n        shuffle=True,\n        num_workers=NUM_WORKERS,\n        pin_memory=True,\n        drop_last=True\n    )\n    val_loader = DataLoader(\n        SimpleDataset(val_data, DATA_DIR, 'train'),\n        batch_size=BATCH_SIZE,\n        shuffle=False,\n        num_workers=NUM_WORKERS,\n        pin_memory=True\n    )\n    model = ChaoticCNN(num_classes).to(device)\n    criterion = nn.CrossEntropyLoss(label_smoothing=0.1)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n    scaler = GradScaler()\n    best_acc = 0\n    for epoch in range(EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, scaler)\n        val_loss, val_acc = validate(model, val_loader, criterion)\n        scheduler.step()\n        print(f\"Train: {train_loss:.4f}, {train_acc:.2f}%\")\n        print(f\"Val:   {val_loss:.4f}, {val_acc:.2f}%\")\n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), \"best_model.pth\")\n            print(\"Saved best model\")\n    print(f\"\\nBest Val Acc: {best_acc:.2f}%\")\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T07:15:48.698756Z","iopub.execute_input":"2026-04-25T07:15:48.699224Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nDATA_DIR   = '/kaggle/input/competitions/recursion-cellular-image-classification'\nTRAIN_CSV  = f'{DATA_DIR}/train.csv'\nTEST_CSV   = '/kaggle/input/datasets/himanshusardana2/corrected-test-csv-recurrence-cellular/test.csv'\nMODEL_NAME  = 'densenet121'\nIMG_SIZE    = 320\nBATCH_SIZE  = 32\nEPOCHS      = 20\nLR          = 3e-4\nNUM_WORKERS = 2\nSEED        = 42\nCELL_TYPES  = ['HUVEC']\nCHAOTIC_MAP = 'sine'\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\nclass ChaoticTransform(nn.Module):\n    \n    VALID_MAPS = ('logistic', 'skew_tent', 'sine')\n    def __init__(self, map_type: str = 'skew_tent'):\n        super().__init__()\n        assert map_type in self.VALID_MAPS, \\\n            f\"map_type must be one of {self.VALID_MAPS}, got '{map_type}'\"\n        self.map_type = map_type\n        self.r = 4.0\n        self.p = 0.499\n    def _normalize(self, f: torch.Tensor) -> torch.Tensor:\n        \n        f_min = f.min(dim=1, keepdim=True).values\n        f_max = f.max(dim=1, keepdim=True).values\n        denom = (f_max - f_min).clamp(min=1e-8)\n        return (f - f_min) / denom\n    def _logistic(self, f: torch.Tensor) -> torch.Tensor:\n        return self.r * f * (1.0 - f)\n    def _skew_tent(self, f: torch.Tensor) -> torch.Tensor:\n        p = self.p\n        return torch.where(f < p, f / p, (1.0 - f) / (1.0 - p))\n    def _sine(self, f: torch.Tensor) -> torch.Tensor:\n        return torch.sin(torch.pi * f)\n    def forward(self, f: torch.Tensor) -> torch.Tensor:\n        \n        f_norm = self._normalize(f)\n        if self.map_type == 'logistic':\n            f_star = self._logistic(f_norm)\n        elif self.map_type == 'skew_tent':\n            f_star = self._skew_tent(f_norm)\n        else:\n            f_star = self._sine(f_norm)\n        return f_star\nclass SimpleDataset(Dataset):\n    def __init__(self, df, data_dir, mode='train'):\n        self.df       = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.mode     = mode\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        exp, plate, well = row['experiment'], row['plate'], row['well']\n        prefix = 'test' if self.mode == 'test' else 'train'\n        base   = f'{self.data_dir}/{prefix}/{exp}/Plate{plate}/{well}_s1_w'\n        channels = []\n        for i in range(1, 7):\n            img = cv2.imread(f'{base}{i}.png', cv2.IMREAD_GRAYSCALE)\n            channels.append(img if img is not None else np.zeros((512, 512), dtype=np.uint8))\n        img = np.stack(channels, axis=-1)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE)).astype(np.float32) / 255.0\n        img = torch.from_numpy(img).permute(2, 0, 1)\n        if self.mode == 'train':\n            return img, row['label']\n        return img, row['id_code']\nclass ChaoticDenseNet(nn.Module):\n    \n    def __init__(self, num_classes: int, map_type: str = 'skew_tent'):\n        super().__init__()\n        backbone = timm.create_model(MODEL_NAME, pretrained=True, in_chans=3)\n        old_conv = backbone.features.conv0\n        out_ch   = int(old_conv.out_channels)\n        kernel   = int(old_conv.kernel_size[0]) if hasattr(old_conv.kernel_size, '__len__') else int(old_conv.kernel_size)\n        stride   = int(old_conv.stride[0])      if hasattr(old_conv.stride,      '__len__') else int(old_conv.stride)\n        padding  = int(old_conv.padding[0])     if hasattr(old_conv.padding,     '__len__') else int(old_conv.padding)\n        backbone.features.conv0 = nn.Conv2d(6, out_ch, kernel, stride, padding, bias=False)\n        with torch.no_grad():\n            w = old_conv.weight.mean(dim=1, keepdim=True).repeat(1, 6, 1, 1) / 6\n            backbone.features.conv0.weight.copy_(w)\n        n_features = backbone.get_classifier().in_features\n        backbone.reset_classifier(0)\n        self.backbone = backbone\n        self.chaotic = ChaoticTransform(map_type=map_type)\n        self.classifier = nn.Sequential(\n            nn.Dropout(0.3),\n            nn.Linear(n_features, num_classes),\n        )\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        f      = self.backbone(x)\n        f_star = self.chaotic(f)\n        return self.classifier(f_star)\ndef train_epoch(model, loader, criterion, optimizer):\n    model.train()\n    loss_sum, correct, total = 0, 0, 0\n    pbar = tqdm(loader, desc='Training')\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        out  = model(imgs)\n        loss = criterion(out, labels)\n        loss.backward()\n        optimizer.step()\n        loss_sum += loss.item()\n        correct  += (out.argmax(1) == labels).sum().item()\n        total    += labels.size(0)\n        pbar.set_postfix({\n            'loss': f'{loss_sum / len(loader):.4f}',\n            'acc':  f'{100. * correct / total:.2f}%',\n        })\n    return loss_sum / len(loader), 100 * correct / total\ndef validate(model, loader, criterion):\n    model.eval()\n    loss_sum, correct, total = 0, 0, 0\n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Val'):\n            imgs, labels = imgs.to(device), labels.to(device)\n            out  = model(imgs)\n            loss = criterion(out, labels)\n            loss_sum += loss.item()\n            correct  += (out.argmax(1) == labels).sum().item()\n            total    += labels.size(0)\n    return loss_sum / len(loader), 100 * correct / total\ndef main():\n    print(\"=\" * 65)\n    print(f\"DenseNet121 + Chaotic CNN  [{CHAOTIC_MAP} map]\")\n    print(\"=\" * 65)\n    print(\"\\nLoading data...\")\n    train_df = pd.read_csv(TRAIN_CSV)\n    test_df  = pd.read_csv(TEST_CSV)\n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    train_df = train_df[train_df['cell_type'].isin(CELL_TYPES)].reset_index(drop=True)\n    print(f\"Training samples: {len(train_df)}\")\n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    sirnas          = sorted(train_df['sirna_id'].unique())\n    sirna_to_label  = {s: i for i, s in enumerate(sirnas)}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    num_classes     = len(sirnas)\n    print(f\"Classes: {num_classes}\")\n    train_data, val_data = train_test_split(\n        train_df, test_size=0.15,\n        stratify=train_df['label'], random_state=SEED,\n    )\n    train_loader = DataLoader(\n        SimpleDataset(train_data, DATA_DIR),\n        batch_size=BATCH_SIZE, shuffle=True,\n        num_workers=NUM_WORKERS, pin_memory=True,\n    )\n    val_loader = DataLoader(\n        SimpleDataset(val_data, DATA_DIR),\n        batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS, pin_memory=True,\n    )\n    print(f\"\\nCreating ChaoticDenseNet121 [{CHAOTIC_MAP}]...\")\n    model = ChaoticDenseNet(num_classes=num_classes, map_type=CHAOTIC_MAP).to(device)\n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n    best_acc, patience, counter = 0, 3, 0\n    save_path = f'best_densenet121_chaotic_{CHAOTIC_MAP}.pth'\n    for epoch in range(EPOCHS):\n        print(f\"\\nEpoch {epoch + 1}/{EPOCHS}\")\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer)\n        val_loss,   val_acc   = validate(model, val_loader, criterion)\n        scheduler.step()\n        print(\n            f\"Train: loss={train_loss:.4f}  acc={train_acc:.2f}%  |  \"\n            f\"Val: loss={val_loss:.4f}  acc={val_acc:.2f}%\"\n        )\n        if val_acc > best_acc:\n            best_acc = val_acc\n            counter  = 0\n            torch.save(model.state_dict(), save_path)\n            print(f\"  → Saved best model: {best_acc:.2f}%\")\n        else:\n            counter += 1\n            print(f\"  → No improvement ({counter}/{patience})\")\n            if counter >= patience:\n                print(\"Early stopping!\")\n                break\n    print(\"\\nRunning inference...\")\n    model.load_state_dict(torch.load(save_path))\n    model.eval()\n    test_loader = DataLoader(\n        SimpleDataset(test_df, DATA_DIR, mode='test'),\n        batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS,\n    )\n    preds, ids = [], []\n    with torch.no_grad():\n        for imgs, img_ids in tqdm(test_loader, desc='Test'):\n            out = model(imgs.to(device))\n            preds.extend(out.argmax(1).cpu().numpy())\n            ids.extend(img_ids)\n    label_to_sirna = {v: k for k, v in sirna_to_label.items()}\n    submission     = pd.DataFrame({\n        'id_code': ids,\n        'sirna':   [label_to_sirna[p] for p in preds],\n    })\n    out_csv = f'submission_densenet121_chaotic_{CHAOTIC_MAP}.csv'\n    submission.to_csv(out_csv, index=False)\n    print(f\"\\n✓ Done!  Best Val Acc: {best_acc:.2f}%\")\n    print(f\"Saved {out_csv}\")\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T21:41:52.956859Z","iopub.execute_input":"2026-04-25T21:41:52.957622Z","iopub.status.idle":"2026-04-25T23:54:42.623086Z","shell.execute_reply.started":"2026-04-25T21:41:52.957585Z","shell.execute_reply":"2026-04-25T23:54:42.621677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!lscpu\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T07:39:44.895479Z","iopub.execute_input":"2026-04-26T07:39:44.896215Z","iopub.status.idle":"2026-04-26T07:39:45.11508Z","shell.execute_reply.started":"2026-04-26T07:39:44.896182Z","shell.execute_reply":"2026-04-26T07:39:45.113996Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport timm\nfrom tqdm import tqdm\nimport cv2\nimport warnings\nwarnings.filterwarnings('ignore')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR    = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV   = f'{DATA_DIR}/train.csv'\n    TEST_CSV    = f'{DATA_DIR}/test.csv'\n    PIXEL_STATS = f'{DATA_DIR}/pixel_stats.csv'\n    MODEL_NAME   = 'resnext50_32x4d'\n    IMG_SIZE     = 384\n    BATCH_SIZE   = 16\n    EPOCHS       = 30\n    WARMUP_EP    = 1\n    LR           = 3e-4\n    MIN_LR       = 1e-6\n    WEIGHT_DECAY = 1e-4\n    LABEL_SMOOTH = 0.1\n    NUM_WORKERS  = 4\n    SEED         = 42\n    NUM_CLASSES  = 1108\n    VAL_EXP_FRAC = 0.15\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark     = False\nset_seed(Config.SEED)\ndef build_stats_lookup(pixel_stats_path):\n    \n    stats  = pd.read_csv(pixel_stats_path)\n    lookup = {}\n    for _, row in stats.iterrows():\n        key = (row['experiment'], int(row['plate']),\n               row['well'],       int(row['site']),\n               int(row['channel']))\n        lookup[key] = (float(row['mean']), float(row['std']))\n    return lookup\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, stats_lookup, mode='train', site=None):\n        self.df           = df.reset_index(drop=True)\n        self.data_dir     = data_dir\n        self.stats_lookup = stats_lookup\n        self.mode         = mode\n        self.site         = site\n    def __len__(self):\n        return len(self.df)\n    def _load_channels(self, row, site):\n        \n        exp   = row['experiment']\n        plate = int(row['plate'])\n        well  = row['well']\n        split = 'test' if self.mode == 'test' else 'train'\n        base  = f'{self.data_dir}/{split}/{exp}/Plate{plate}/{well}_s{site}_w'\n        channels = []\n        for ch in range(1, 7):\n            path = f'{base}{ch}.png'\n            img  = cv2.imread(path, cv2.IMREAD_GRAYSCALE) if os.path.exists(path) \\\n                   else np.zeros((512, 512), dtype=np.uint8)\n            img  = img.astype(np.float32)\n            key        = (exp, plate, well, int(site), ch)\n            mean, std  = self.stats_lookup.get(key, (127.5, 64.0))\n            img        = (img - mean) / (std + 1e-6)\n            channels.append(img)\n        stacked = np.stack(channels, axis=-1)\n        stacked = cv2.resize(stacked, (Config.IMG_SIZE, Config.IMG_SIZE))\n        return stacked\n    def _augment(self, img):\n        \n        if self.mode != 'train':\n            return img\n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.rot90(img, k=np.random.randint(1, 4)).copy()\n        return img\n    def __getitem__(self, idx):\n        row  = self.df.iloc[idx]\n        site = np.random.choice(['1', '2']) if self.mode == 'train' \\\n               else (self.site or '1')\n        img = self._load_channels(row, site)\n        img = self._augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1).float()\n        if self.mode in ('train', 'val'):\n            return img, int(row['label'])\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    \n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        if hasattr(self.backbone, 'conv1'):\n            old = self.backbone.conv1\n            new = nn.Conv2d(\n                in_channels, old.out_channels,\n                kernel_size=old.kernel_size,\n                stride=old.stride,\n                padding=old.padding,\n                bias=old.bias is not None\n            )\n            with torch.no_grad():\n                new.weight[:, :3, ...] = old.weight\n                new.weight[:, 3:, ...] = old.weight\n                if old.bias is not None:\n                    new.bias = nn.Parameter(old.bias.clone())\n            self.backbone.conv1 = new\n        n_feat = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        self.head = nn.Sequential(\n            nn.BatchNorm1d(n_feat),\n            nn.Dropout(0.4),\n            nn.Linear(n_feat, 512),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.3),\n            nn.Linear(512, num_classes)\n        )\n    def forward(self, x):\n        return self.head(self.backbone(x))\n    def freeze_backbone(self):\n        for p in self.backbone.parameters():\n            p.requires_grad = False\n    def unfreeze_backbone(self):\n        for p in self.backbone.parameters():\n            p.requires_grad = True\ndef build_scheduler(optimizer, total_steps, warmup_steps):\n    def lr_lambda(step):\n        if step < warmup_steps:\n            return float(step) / max(1, warmup_steps)\n        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n        cosine   = 0.5 * (1.0 + np.cos(np.pi * progress))\n        return max(Config.MIN_LR / Config.LR, cosine)\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\ndef train_epoch(model, loader, criterion, optimizer, scheduler, scaler, device):\n    model.train()\n    total_loss = correct = total = 0\n    pbar = tqdm(loader, desc='Train', leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        with autocast():\n            out  = model(imgs)\n            loss = criterion(out, labels)\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        scheduler.step()\n        total_loss += loss.item()\n        _, pred     = out.max(1)\n        correct    += pred.eq(labels).sum().item()\n        total      += labels.size(0)\n        pbar.set_postfix(loss=f'{total_loss/len(loader):.4f}',\n                         acc=f'{100.*correct/total:.2f}%')\n    return total_loss / len(loader), 100. * correct / total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    total_loss = correct = total = 0\n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Val  ', leave=False):\n            imgs, labels = imgs.to(device), labels.to(device)\n            with autocast():\n                out  = model(imgs)\n                loss = criterion(out, labels)\n            total_loss += loss.item()\n            _, pred     = out.max(1)\n            correct    += pred.eq(labels).sum().item()\n            total      += labels.size(0)\n    return total_loss / len(loader), 100. * correct / total\ndef predict_tta(model, test_df, data_dir, stats_lookup, device):\n    \n    model.eval()\n    tta_configs = [\n        ('1', None), ('1', 'h'), ('1', 'v'),\n        ('2', None), ('2', 'h'), ('2', 'v'),\n    ]\n    all_probs = []\n    ids       = None\n    for site, flip in tta_configs:\n        ds     = CellularDataset(test_df, data_dir, stats_lookup,\n                                 mode='test', site=site)\n        loader = DataLoader(ds, batch_size=Config.BATCH_SIZE,\n                            shuffle=False, num_workers=Config.NUM_WORKERS,\n                            pin_memory=True)\n        probs_list = []\n        id_list    = []\n        with torch.no_grad():\n            for imgs, img_ids in tqdm(loader,\n                                      desc=f'TTA site{site} {flip or \"orig\"}',\n                                      leave=False):\n                imgs = imgs.to(device)\n                if flip == 'h':\n                    imgs = torch.flip(imgs, dims=[3])\n                elif flip == 'v':\n                    imgs = torch.flip(imgs, dims=[2])\n                with autocast():\n                    out = model(imgs)\n                probs_list.append(out.softmax(dim=1).cpu())\n                id_list.extend(img_ids)\n        all_probs.append(torch.cat(probs_list, dim=0))\n        if ids is None:\n            ids = id_list\n    avg_probs = torch.stack(all_probs).mean(dim=0)\n    preds     = avg_probs.argmax(dim=1).numpy()\n    return preds, ids\ndef main():\n    print(\"Loading CSVs...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df  = pd.read_csv(Config.TEST_CSV)\n    if train_df['sirna'].dtype == object:\n        train_df['sirna_id'] = train_df['sirna'].str.extract(r'(\\d+)').astype(int)\n    else:\n        train_df['sirna_id'] = train_df['sirna'].astype(int)\n    unique_sirnas      = sorted(train_df['sirna_id'].unique())\n    sirna_to_label     = {s: i for i, s in enumerate(unique_sirnas)}\n    label_to_sirna     = {i: s for s, i in sirna_to_label.items()}\n    train_df['label']  = train_df['sirna_id'].map(sirna_to_label)\n    Config.NUM_CLASSES = len(unique_sirnas)\n    print(f\"Classes      : {Config.NUM_CLASSES}\")\n    print(f\"Train wells  : {len(train_df)}\")\n    print(f\"Test wells   : {len(test_df)}\")\n    rng      = np.random.default_rng(Config.SEED)\n    exps     = train_df['experiment'].unique()\n    n_val    = max(1, int(len(exps) * Config.VAL_EXP_FRAC))\n    val_exps = set(rng.choice(exps, size=n_val, replace=False))\n    val_df   = train_df[train_df['experiment'].isin(val_exps)].reset_index(drop=True)\n    tr_df    = train_df[~train_df['experiment'].isin(val_exps)].reset_index(drop=True)\n    print(f\"Train exps   : {len(exps)-n_val}  ({len(tr_df)} wells)\")\n    print(f\"Val exps     : {n_val}  ({len(val_df)} wells)\")\n    print(\"Building pixel-stats lookup...\")\n    stats_lookup = build_stats_lookup(Config.PIXEL_STATS)\n    train_loader = DataLoader(\n        CellularDataset(tr_df,  Config.DATA_DIR, stats_lookup, mode='train'),\n        batch_size=Config.BATCH_SIZE, shuffle=True,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    val_loader = DataLoader(\n        CellularDataset(val_df, Config.DATA_DIR, stats_lookup,\n                        mode='val', site='1'),\n        batch_size=Config.BATCH_SIZE, shuffle=False,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    print(f\"Building {Config.MODEL_NAME}...\")\n    model     = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    criterion = nn.CrossEntropyLoss(label_smoothing=Config.LABEL_SMOOTH)\n    optimizer = torch.optim.AdamW(model.parameters(),\n                                  lr=Config.LR, weight_decay=Config.WEIGHT_DECAY)\n    scaler    = GradScaler()\n    total_steps  = Config.EPOCHS * len(train_loader)\n    warmup_steps = Config.WARMUP_EP * len(train_loader)\n    scheduler    = build_scheduler(optimizer, total_steps, warmup_steps)\n    best_acc  = 0.0\n    best_path = 'best_resnext50_32x4d.pth'\n    print(f\"\\nPhase 1 — backbone frozen for {Config.WARMUP_EP} epoch(s)\")\n    model.freeze_backbone()\n    for epoch in range(Config.EPOCHS):\n        if epoch == Config.WARMUP_EP:\n            model.unfreeze_backbone()\n            print(f\"\\nPhase 2 — full fine-tune from epoch {epoch+1}\")\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}  \"\n              f\"lr={optimizer.param_groups[0]['lr']:.2e}\")\n        tr_loss, tr_acc = train_epoch(model, train_loader, criterion,\n                                      optimizer, scheduler, scaler, device)\n        vl_loss, vl_acc = validate(model, val_loader, criterion, device)\n        print(f\"  Train  loss={tr_loss:.4f}  acc={tr_acc:.2f}%\")\n        print(f\"  Val    loss={vl_loss:.4f}  acc={vl_acc:.2f}%\")\n        if vl_acc > best_acc:\n            best_acc = vl_acc\n            torch.save(model.state_dict(), best_path)\n            print(f\"  ✓ Saved best model (val acc={best_acc:.2f}%)\")\n    print(f\"\\nLoading best weights (val acc={best_acc:.2f}%)...\")\n    model.load_state_dict(torch.load(best_path, map_location=device))\n    print(\"Running TTA inference...\")\n    preds, ids = predict_tta(model, test_df, Config.DATA_DIR, stats_lookup, device)\n    sirna_preds = [label_to_sirna[int(p)] for p in preds]\n    submission  = pd.DataFrame({'id_code': ids, 'sirna': sirna_preds})\n    submission.to_csv('submission_resnext50_32x4d.csv', index=False)\n    print(f\"\\nDone. submission_resnext50_32x4d.csv saved  |  Best val acc: {best_acc:.2f}%\")\n    print(submission.head())\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T07:40:08.685054Z","iopub.execute_input":"2026-04-26T07:40:08.685563Z","iopub.status.idle":"2026-04-26T13:50:41.956075Z","shell.execute_reply.started":"2026-04-26T07:40:08.68553Z","shell.execute_reply":"2026-04-26T13:50:41.955048Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport timm\nfrom tqdm import tqdm\nimport cv2\nimport warnings\nwarnings.filterwarnings('ignore')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR    = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV   = f'{DATA_DIR}/train.csv'\n    TEST_CSV    = f'{DATA_DIR}/test.csv'\n    PIXEL_STATS = f'{DATA_DIR}/pixel_stats.csv'\n    MODEL_NAME   = 'resnext50_32x4d'\n    IMG_SIZE     = 384\n    BATCH_SIZE   = 16\n    EPOCHS       = 30\n    WARMUP_EP    = 1\n    LR           = 3e-4\n    MIN_LR       = 1e-6\n    WEIGHT_DECAY = 1e-4\n    LABEL_SMOOTH = 0.1\n    NUM_WORKERS  = 4\n    SEED         = 42\n    NUM_CLASSES  = 1108\n    VAL_EXP_FRAC = 0.15\nCHAOTIC_MAP = 'logistic'\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark     = False\nset_seed(Config.SEED)\ndef build_stats_lookup(pixel_stats_path):\n    \n    stats  = pd.read_csv(pixel_stats_path)\n    lookup = {}\n    for _, row in stats.iterrows():\n        key = (row['experiment'], int(row['plate']),\n               row['well'],       int(row['site']),\n               int(row['channel']))\n        lookup[key] = (float(row['mean']), float(row['std']))\n    return lookup\nclass ChaoticTransform(nn.Module):\n    \n    VALID_MAPS = ('logistic', 'skew_tent', 'sine')\n    def __init__(self, map_type: str = 'sine'):\n        super().__init__()\n        assert map_type in self.VALID_MAPS, \\\n            f\"map_type must be one of {self.VALID_MAPS}, got '{map_type}'\"\n        self.map_type = map_type\n        self.r = 4.0\n        self.p = 0.499\n    def _normalize(self, f: torch.Tensor) -> torch.Tensor:\n        \n        f_min = f.min(dim=1, keepdim=True).values\n        f_max = f.max(dim=1, keepdim=True).values\n        denom = (f_max - f_min).clamp(min=1e-8)\n        return (f - f_min) / denom\n    def _logistic(self, f: torch.Tensor) -> torch.Tensor:\n        return self.r * f * (1.0 - f)\n    def _skew_tent(self, f: torch.Tensor) -> torch.Tensor:\n        p = self.p\n        return torch.where(f < p, f / p, (1.0 - f) / (1.0 - p))\n    def _sine(self, f: torch.Tensor) -> torch.Tensor:\n        return torch.sin(torch.pi * f)\n    def forward(self, f: torch.Tensor) -> torch.Tensor:\n        \n        f_norm = self._normalize(f)\n        if self.map_type == 'logistic':\n            f_star = self._logistic(f_norm)\n        elif self.map_type == 'skew_tent':\n            f_star = self._skew_tent(f_norm)\n        else:\n            f_star = self._sine(f_norm)\n        return f_star\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, stats_lookup, mode='train', site=None):\n        self.df           = df.reset_index(drop=True)\n        self.data_dir     = data_dir\n        self.stats_lookup = stats_lookup\n        self.mode         = mode\n        self.site         = site\n    def __len__(self):\n        return len(self.df)\n    def _load_channels(self, row, site):\n        \n        exp   = row['experiment']\n        plate = int(row['plate'])\n        well  = row['well']\n        split = 'test' if self.mode == 'test' else 'train'\n        base  = f'{self.data_dir}/{split}/{exp}/Plate{plate}/{well}_s{site}_w'\n        channels = []\n        for ch in range(1, 7):\n            path = f'{base}{ch}.png'\n            img  = cv2.imread(path, cv2.IMREAD_GRAYSCALE) if os.path.exists(path) \\\n                   else np.zeros((512, 512), dtype=np.uint8)\n            img  = img.astype(np.float32)\n            key        = (exp, plate, well, int(site), ch)\n            mean, std  = self.stats_lookup.get(key, (127.5, 64.0))\n            img        = (img - mean) / (std + 1e-6)\n            channels.append(img)\n        stacked = np.stack(channels, axis=-1)\n        stacked = cv2.resize(stacked, (Config.IMG_SIZE, Config.IMG_SIZE))\n        return stacked\n    def _augment(self, img):\n        \n        if self.mode != 'train':\n            return img\n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.rot90(img, k=np.random.randint(1, 4)).copy()\n        return img\n    def __getitem__(self, idx):\n        row  = self.df.iloc[idx]\n        site = np.random.choice(['1', '2']) if self.mode == 'train' \\\n               else (self.site or '1')\n        img = self._load_channels(row, site)\n        img = self._augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1).float()\n        if self.mode in ('train', 'val'):\n            return img, int(row['label'])\n        else:\n            return img, row['id_code']\nclass ChaoticResNeXt(nn.Module):\n    \n    def __init__(self, model_name, num_classes, in_channels=6, map_type='sine'):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        if hasattr(self.backbone, 'conv1'):\n            old = self.backbone.conv1\n            new = nn.Conv2d(\n                in_channels, old.out_channels,\n                kernel_size=old.kernel_size,\n                stride=old.stride,\n                padding=old.padding,\n                bias=old.bias is not None\n            )\n            with torch.no_grad():\n                new.weight[:, :3, ...] = old.weight\n                new.weight[:, 3:, ...] = old.weight\n                if old.bias is not None:\n                    new.bias = nn.Parameter(old.bias.clone())\n            self.backbone.conv1 = new\n        n_feat = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        self.chaotic = ChaoticTransform(map_type=map_type)\n        self.head = nn.Sequential(\n            nn.BatchNorm1d(n_feat),\n            nn.Dropout(0.4),\n            nn.Linear(n_feat, 512),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.3),\n            nn.Linear(512, num_classes)\n        )\n    def forward(self, x):\n        features = self.backbone(x)\n        features = self.chaotic(features)\n        return self.head(features)\n    def freeze_backbone(self):\n        for p in self.backbone.parameters():\n            p.requires_grad = False\n    def unfreeze_backbone(self):\n        for p in self.backbone.parameters():\n            p.requires_grad = True\ndef build_scheduler(optimizer, total_steps, warmup_steps):\n    def lr_lambda(step):\n        if step < warmup_steps:\n            return float(step) / max(1, warmup_steps)\n        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n        cosine   = 0.5 * (1.0 + np.cos(np.pi * progress))\n        return max(Config.MIN_LR / Config.LR, cosine)\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\ndef train_epoch(model, loader, criterion, optimizer, scheduler, scaler, device):\n    model.train()\n    total_loss = correct = total = 0\n    pbar = tqdm(loader, desc='Train', leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        with autocast():\n            out  = model(imgs)\n            loss = criterion(out, labels)\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        scheduler.step()\n        total_loss += loss.item()\n        _, pred     = out.max(1)\n        correct    += pred.eq(labels).sum().item()\n        total      += labels.size(0)\n        pbar.set_postfix(loss=f'{total_loss/len(loader):.4f}',\n                         acc=f'{100.*correct/total:.2f}%')\n    return total_loss / len(loader), 100. * correct / total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    total_loss = correct = total = 0\n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Val  ', leave=False):\n            imgs, labels = imgs.to(device), labels.to(device)\n            with autocast():\n                out  = model(imgs)\n                loss = criterion(out, labels)\n            total_loss += loss.item()\n            _, pred     = out.max(1)\n            correct    += pred.eq(labels).sum().item()\n            total      += labels.size(0)\n    return total_loss / len(loader), 100. * correct / total\ndef predict_tta(model, test_df, data_dir, stats_lookup, device):\n    \n    model.eval()\n    tta_configs = [\n        ('1', None), ('1', 'h'), ('1', 'v'),\n        ('2', None), ('2', 'h'), ('2', 'v'),\n    ]\n    all_probs = []\n    ids       = None\n    for site, flip in tta_configs:\n        ds     = CellularDataset(test_df, data_dir, stats_lookup,\n                                 mode='test', site=site)\n        loader = DataLoader(ds, batch_size=Config.BATCH_SIZE,\n                            shuffle=False, num_workers=Config.NUM_WORKERS,\n                            pin_memory=True)\n        probs_list = []\n        id_list    = []\n        with torch.no_grad():\n            for imgs, img_ids in tqdm(loader,\n                                      desc=f'TTA site{site} {flip or \"orig\"}',\n                                      leave=False):\n                imgs = imgs.to(device)\n                if flip == 'h':\n                    imgs = torch.flip(imgs, dims=[3])\n                elif flip == 'v':\n                    imgs = torch.flip(imgs, dims=[2])\n                with autocast():\n                    out = model(imgs)\n                probs_list.append(out.softmax(dim=1).cpu())\n                id_list.extend(img_ids)\n        all_probs.append(torch.cat(probs_list, dim=0))\n        if ids is None:\n            ids = id_list\n    avg_probs = torch.stack(all_probs).mean(dim=0)\n    preds     = avg_probs.argmax(dim=1).numpy()\n    return preds, ids\ndef main():\n    print(\"=\" * 70)\n    print(f\"ResNeXt50-32x4d + Chaotic CNN  [{CHAOTIC_MAP} map]\")\n    print(\"=\" * 70)\n    print(\"\\nLoading CSVs...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df  = pd.read_csv(Config.TEST_CSV)\n    if train_df['sirna'].dtype == object:\n        train_df['sirna_id'] = train_df['sirna'].str.extract(r'(\\d+)').astype(int)\n    else:\n        train_df['sirna_id'] = train_df['sirna'].astype(int)\n    unique_sirnas      = sorted(train_df['sirna_id'].unique())\n    sirna_to_label     = {s: i for i, s in enumerate(unique_sirnas)}\n    label_to_sirna     = {i: s for s, i in sirna_to_label.items()}\n    train_df['label']  = train_df['sirna_id'].map(sirna_to_label)\n    Config.NUM_CLASSES = len(unique_sirnas)\n    print(f\"Classes      : {Config.NUM_CLASSES}\")\n    print(f\"Train wells  : {len(train_df)}\")\n    print(f\"Test wells   : {len(test_df)}\")\n    rng      = np.random.default_rng(Config.SEED)\n    exps     = train_df['experiment'].unique()\n    n_val    = max(1, int(len(exps) * Config.VAL_EXP_FRAC))\n    val_exps = set(rng.choice(exps, size=n_val, replace=False))\n    val_df   = train_df[train_df['experiment'].isin(val_exps)].reset_index(drop=True)\n    tr_df    = train_df[~train_df['experiment'].isin(val_exps)].reset_index(drop=True)\n    print(f\"Train exps   : {len(exps)-n_val}  ({len(tr_df)} wells)\")\n    print(f\"Val exps     : {n_val}  ({len(val_df)} wells)\")\n    print(\"Building pixel-stats lookup...\")\n    stats_lookup = build_stats_lookup(Config.PIXEL_STATS)\n    train_loader = DataLoader(\n        CellularDataset(tr_df,  Config.DATA_DIR, stats_lookup, mode='train'),\n        batch_size=Config.BATCH_SIZE, shuffle=True,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    val_loader = DataLoader(\n        CellularDataset(val_df, Config.DATA_DIR, stats_lookup,\n                        mode='val', site='1'),\n        batch_size=Config.BATCH_SIZE, shuffle=False,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    print(f\"Building ChaoticResNeXt50 [{CHAOTIC_MAP}]...\")\n    model = ChaoticResNeXt(\n        Config.MODEL_NAME, \n        Config.NUM_CLASSES, \n        in_channels=6,\n        map_type=CHAOTIC_MAP\n    ).to(device)\n    \n    criterion = nn.CrossEntropyLoss(label_smoothing=Config.LABEL_SMOOTH)\n    optimizer = torch.optim.AdamW(model.parameters(),\n                                  lr=Config.LR, weight_decay=Config.WEIGHT_DECAY)\n    scaler    = GradScaler()\n    total_steps  = Config.EPOCHS * len(train_loader)\n    warmup_steps = Config.WARMUP_EP * len(train_loader)\n    scheduler    = build_scheduler(optimizer, total_steps, warmup_steps)\n    best_acc  = 0.0\n    best_path = f'best_resnext50_chaotic_{CHAOTIC_MAP}.pth'\n    print(f\"\\nPhase 1 — backbone frozen for {Config.WARMUP_EP} epoch(s)\")\n    model.freeze_backbone()\n    for epoch in range(Config.EPOCHS):\n        if epoch == Config.WARMUP_EP:\n            model.unfreeze_backbone()\n            print(f\"\\nPhase 2 — full fine-tune from epoch {epoch+1}\")\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}  \"\n              f\"lr={optimizer.param_groups[0]['lr']:.2e}\")\n        tr_loss, tr_acc = train_epoch(model, train_loader, criterion,\n                                      optimizer, scheduler, scaler, device)\n        vl_loss, vl_acc = validate(model, val_loader, criterion, device)\n        print(f\"  Train  loss={tr_loss:.4f}  acc={tr_acc:.2f}%\")\n        print(f\"  Val    loss={vl_loss:.4f}  acc={vl_acc:.2f}%\")\n        if vl_acc > best_acc:\n            best_acc = vl_acc\n            torch.save(model.state_dict(), best_path)\n            print(f\"  ✓ Saved best model (val acc={best_acc:.2f}%)\")\n    print(f\"\\nLoading best weights (val acc={best_acc:.2f}%)...\")\n    model.load_state_dict(torch.load(best_path, map_location=device))\n    print(\"Running TTA inference...\")\n    preds, ids = predict_tta(model, test_df, Config.DATA_DIR, stats_lookup, device)\n    sirna_preds = [label_to_sirna[int(p)] for p in preds]\n    submission  = pd.DataFrame({'id_code': ids, 'sirna': sirna_preds})\n    submission.to_csv(f'submission_resnext50_chaotic_{CHAOTIC_MAP}.csv', index=False)\n    print(f\"\\nDone. submission_resnext50_chaotic_{CHAOTIC_MAP}.csv saved  |  Best val acc: {best_acc:.2f}%\")\n    print(submission.head())\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T18:56:28.291811Z","iopub.execute_input":"2026-04-26T18:56:28.292719Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport timm\nfrom tqdm import tqdm\nimport cv2\nimport warnings\nwarnings.filterwarnings('ignore')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR    = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV   = f'{DATA_DIR}/train.csv'\n    TEST_CSV    = f'{DATA_DIR}/test.csv'\n    PIXEL_STATS = f'{DATA_DIR}/pixel_stats.csv'\n    TRAIN_CONTROLS_CSV = f'{DATA_DIR}/train_controls.csv'\n    TEST_CONTROLS_CSV  = f'{DATA_DIR}/test_controls.csv'\n    MODEL_NAME   = 'resnext50_32x4d'\n    IMG_SIZE     = 384\n    BATCH_SIZE   = 16\n    EPOCHS       = 30\n    WARMUP_EP    = 1\n    LR           = 3e-4\n    MIN_LR       = 1e-6\n    WEIGHT_DECAY = 1e-4\n    LABEL_SMOOTH = 0.1\n    NUM_WORKERS  = 4\n    SEED         = 42\n    NUM_CLASSES  = 1108\n    VAL_EXP_FRAC = 0.15\n    USE_PLATE_NORMALIZATION    = True\n    USE_CONTROLS_AS_EXTRA_DATA = True\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark     = False\nset_seed(Config.SEED)\ndef build_stats_lookup(pixel_stats_path):\n    \n    stats  = pd.read_csv(pixel_stats_path)\n    lookup = {}\n    for _, row in stats.iterrows():\n        key = (row['experiment'], int(row['plate']),\n               row['well'],       int(row['site']),\n               int(row['channel']))\n        lookup[key] = (float(row['mean']), float(row['std']))\n    return lookup\ndef compute_plate_control_stats(controls_df, data_dir, mode='train'):\n    \n    stats = {}\n    grouped = controls_df.groupby(['experiment', 'plate'])\n    print(f\"Computing plate control stats from {mode}_controls.csv ...\")\n    for (exp, plate), group in tqdm(grouped, desc='Plate stats'):\n        accum = None\n        count = 0\n        for _, row in group.iterrows():\n            well = row['well']\n            path_template = f'{data_dir}/{mode}/{exp}/Plate{plate}/{well}_s1_w'\n            channels = []\n            for ch in range(1, 7):\n                img_path = f'{path_template}{ch}.png'\n                if os.path.exists(img_path):\n                    img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                else:\n                    img = np.zeros((512, 512), dtype=np.uint8)\n                channels.append(img)\n            img = np.stack(channels, axis=-1).astype(np.float32) / 255.0\n            if accum is None:\n                accum = img\n            else:\n                accum += img\n            count += 1\n        if count > 0:\n            stats[(exp, plate)] = accum / count\n    print(f\"  Computed stats for {len(stats)} plates\")\n    return stats\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, stats_lookup, mode='train',\n                 site=None, plate_stats=None):\n        self.df           = df.reset_index(drop=True)\n        self.data_dir     = data_dir\n        self.stats_lookup = stats_lookup\n        self.mode         = mode\n        self.site         = site\n        self.plate_stats  = plate_stats\n    def __len__(self):\n        return len(self.df)\n    def _load_channels(self, row, site):\n        \n        exp   = row['experiment']\n        plate = int(row['plate'])\n        well  = row['well']\n        split = 'test' if self.mode == 'test' else 'train'\n        base  = f'{self.data_dir}/{split}/{exp}/Plate{plate}/{well}_s{site}_w'\n        channels = []\n        for ch in range(1, 7):\n            path = f'{base}{ch}.png'\n            img  = cv2.imread(path, cv2.IMREAD_GRAYSCALE) if os.path.exists(path) \\\n                   else np.zeros((512, 512), dtype=np.uint8)\n            img  = img.astype(np.float32)\n            key        = (exp, plate, well, int(site), ch)\n            mean, std  = self.stats_lookup.get(key, (127.5, 64.0))\n            img        = (img - mean) / (std + 1e-6)\n            channels.append(img)\n        stacked = np.stack(channels, axis=-1)\n        stacked = cv2.resize(stacked, (Config.IMG_SIZE, Config.IMG_SIZE))\n        if self.plate_stats is not None:\n            ctrl_key = (exp, plate)\n            if ctrl_key in self.plate_stats:\n                ctrl_mean = cv2.resize(\n                    self.plate_stats[ctrl_key],\n                    (Config.IMG_SIZE, Config.IMG_SIZE)\n                )\n                stacked = stacked - ctrl_mean\n        return stacked\n    def _augment(self, img):\n        \n        if self.mode != 'train':\n            return img\n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.rot90(img, k=np.random.randint(1, 4)).copy()\n        return img\n    def __getitem__(self, idx):\n        row  = self.df.iloc[idx]\n        site = np.random.choice(['1', '2']) if self.mode == 'train' \\\n               else (self.site or '1')\n        img = self._load_channels(row, site)\n        img = self._augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1).float()\n        if self.mode in ('train', 'val'):\n            return img, int(row['label'])\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    \n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        if hasattr(self.backbone, 'conv1'):\n            old = self.backbone.conv1\n            new = nn.Conv2d(\n                in_channels, old.out_channels,\n                kernel_size=old.kernel_size,\n                stride=old.stride,\n                padding=old.padding,\n                bias=old.bias is not None\n            )\n            with torch.no_grad():\n                new.weight[:, :3, ...] = old.weight\n                new.weight[:, 3:, ...] = old.weight\n                if old.bias is not None:\n                    new.bias = nn.Parameter(old.bias.clone())\n            self.backbone.conv1 = new\n        n_feat = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        self.head = nn.Sequential(\n            nn.BatchNorm1d(n_feat),\n            nn.Dropout(0.4),\n            nn.Linear(n_feat, 512),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.3),\n            nn.Linear(512, num_classes)\n        )\n    def forward(self, x):\n        return self.head(self.backbone(x))\n    def freeze_backbone(self):\n        for p in self.backbone.parameters():\n            p.requires_grad = False\n    def unfreeze_backbone(self):\n        for p in self.backbone.parameters():\n            p.requires_grad = True\ndef build_scheduler(optimizer, total_steps, warmup_steps):\n    def lr_lambda(step):\n        if step < warmup_steps:\n            return float(step) / max(1, warmup_steps)\n        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n        cosine   = 0.5 * (1.0 + np.cos(np.pi * progress))\n        return max(Config.MIN_LR / Config.LR, cosine)\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\ndef train_epoch(model, loader, criterion, optimizer, scheduler, scaler, device):\n    model.train()\n    total_loss = correct = total = 0\n    pbar = tqdm(loader, desc='Train', leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        with autocast():\n            out  = model(imgs)\n            loss = criterion(out, labels)\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        scheduler.step()\n        total_loss += loss.item()\n        _, pred     = out.max(1)\n        correct    += pred.eq(labels).sum().item()\n        total      += labels.size(0)\n        pbar.set_postfix(loss=f'{total_loss/len(loader):.4f}',\n                         acc=f'{100.*correct/total:.2f}%')\n    return total_loss / len(loader), 100. * correct / total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    total_loss = correct = total = 0\n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Val  ', leave=False):\n            imgs, labels = imgs.to(device), labels.to(device)\n            with autocast():\n                out  = model(imgs)\n                loss = criterion(out, labels)\n            total_loss += loss.item()\n            _, pred     = out.max(1)\n            correct    += pred.eq(labels).sum().item()\n            total      += labels.size(0)\n    return total_loss / len(loader), 100. * correct / total\ndef predict_tta(model, test_df, data_dir, stats_lookup, device,\n                plate_stats=None):\n    \n    model.eval()\n    tta_configs = [\n        ('1', None), ('1', 'h'), ('1', 'v'),\n        ('2', None), ('2', 'h'), ('2', 'v'),\n    ]\n    all_probs = []\n    ids       = None\n    for site, flip in tta_configs:\n        ds     = CellularDataset(test_df, data_dir, stats_lookup,\n                                 mode='test', site=site,\n                                 plate_stats=plate_stats)\n        loader = DataLoader(ds, batch_size=Config.BATCH_SIZE,\n                            shuffle=False, num_workers=Config.NUM_WORKERS,\n                            pin_memory=True)\n        probs_list = []\n        id_list    = []\n        with torch.no_grad():\n            for imgs, img_ids in tqdm(loader,\n                                      desc=f'TTA site{site} {flip or \"orig\"}',\n                                      leave=False):\n                imgs = imgs.to(device)\n                if flip == 'h':\n                    imgs = torch.flip(imgs, dims=[3])\n                elif flip == 'v':\n                    imgs = torch.flip(imgs, dims=[2])\n                with autocast():\n                    out = model(imgs)\n                probs_list.append(out.softmax(dim=1).cpu())\n                id_list.extend(img_ids)\n        all_probs.append(torch.cat(probs_list, dim=0))\n        if ids is None:\n            ids = id_list\n    avg_probs = torch.stack(all_probs).mean(dim=0)\n    preds     = avg_probs.argmax(dim=1).numpy()\n    return preds, ids\ndef main():\n    print(\"Loading CSVs...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df  = pd.read_csv(Config.TEST_CSV)\n    train_controls_df = pd.read_csv(Config.TRAIN_CONTROLS_CSV)\n    test_controls_df  = pd.read_csv(Config.TEST_CONTROLS_CSV)\n    print(f\"  train_controls.csv: {len(train_controls_df)} rows, columns: {list(train_controls_df.columns)}\")\n    print(f\"  test_controls.csv:  {len(test_controls_df)} rows, columns: {list(test_controls_df.columns)}\")\n    if train_df['sirna'].dtype == object:\n        train_df['sirna_id'] = train_df['sirna'].str.extract(r'(\\d+)').astype(int)\n    else:\n        train_df['sirna_id'] = train_df['sirna'].astype(int)\n    unique_sirnas      = sorted(train_df['sirna_id'].unique())\n    sirna_to_label     = {s: i for i, s in enumerate(unique_sirnas)}\n    label_to_sirna     = {i: s for s, i in sirna_to_label.items()}\n    train_df['label']  = train_df['sirna_id'].map(sirna_to_label)\n    Config.NUM_CLASSES = len(unique_sirnas)\n    print(f\"Classes      : {Config.NUM_CLASSES}\")\n    print(f\"Train wells  : {len(train_df)}\")\n    print(f\"Test wells   : {len(test_df)}\")\n    control_df = pd.DataFrame()\n    if Config.USE_CONTROLS_AS_EXTRA_DATA:\n        print(\"\\n--- Adding control wells as extra training data ---\")\n        print(f\"  Controls sirna values (sample): {train_controls_df['sirna'].unique()[:10]}\")\n        neg_mask = train_controls_df['sirna'].astype(str).str.contains(\n            'negative', case=False, na=False\n        )\n        if neg_mask.sum() == 0:\n            neg_mask = pd.Series([True] * len(train_controls_df))\n        neg_controls = train_controls_df[neg_mask].copy()\n        print(f\"  Found {len(neg_controls)} negative control wells\")\n        required_cols = ['experiment', 'plate', 'well']\n        if all(c in neg_controls.columns for c in required_cols) and len(neg_controls) > 0:\n            rng = np.random.default_rng(Config.SEED)\n            neg_controls = neg_controls[required_cols].copy()\n            neg_controls['sirna'] = 'negative_control'\n            neg_controls['sirna_id'] = rng.choice(unique_sirnas, size=len(neg_controls))\n            neg_controls['label'] = neg_controls['sirna_id'].map(sirna_to_label)\n            frac = min(1.0, len(train_df) / max(1, len(neg_controls)))\n            control_df = neg_controls.sample(frac=frac, random_state=Config.SEED)\n            print(f\"  Adding {len(control_df)} control samples to training\")\n            train_df = pd.concat([train_df, control_df], ignore_index=True)\n            train_df = train_df.sample(frac=1, random_state=Config.SEED).reset_index(drop=True)\n            print(f\"  New training size: {len(train_df)}\")\n    plate_stats_train = None\n    plate_stats_test  = None\n    if Config.USE_PLATE_NORMALIZATION:\n        plate_stats_train = compute_plate_control_stats(\n            train_controls_df, Config.DATA_DIR, mode='train'\n        )\n        plate_stats_test = compute_plate_control_stats(\n            test_controls_df, Config.DATA_DIR, mode='test'\n        )\n        for key, val in plate_stats_train.items():\n            if key not in plate_stats_test:\n                plate_stats_test[key] = val\n    real_train_mask = train_df['sirna'] != 'negative_control'\n    real_train_df   = train_df[real_train_mask]\n    rng      = np.random.default_rng(Config.SEED)\n    exps     = real_train_df['experiment'].unique()\n    n_val    = max(1, int(len(exps) * Config.VAL_EXP_FRAC))\n    val_exps = set(rng.choice(exps, size=n_val, replace=False))\n    val_df   = real_train_df[real_train_df['experiment'].isin(val_exps)].reset_index(drop=True)\n    tr_df    = real_train_df[~real_train_df['experiment'].isin(val_exps)].reset_index(drop=True)\n    if len(control_df) > 0:\n        tr_df = pd.concat([tr_df, control_df], ignore_index=True)\n        tr_df = tr_df.sample(frac=1, random_state=Config.SEED).reset_index(drop=True)\n    print(f\"Train exps   : {len(exps)-n_val}  ({len(tr_df)} wells)\")\n    print(f\"Val exps     : {n_val}  ({len(val_df)} wells)\")\n    print(\"Building pixel-stats lookup...\")\n    stats_lookup = build_stats_lookup(Config.PIXEL_STATS)\n    train_loader = DataLoader(\n        CellularDataset(tr_df,  Config.DATA_DIR, stats_lookup,\n                        mode='train', plate_stats=plate_stats_train),\n        batch_size=Config.BATCH_SIZE, shuffle=True,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    val_loader = DataLoader(\n        CellularDataset(val_df, Config.DATA_DIR, stats_lookup,\n                        mode='val', site='1',\n                        plate_stats=plate_stats_train),\n        batch_size=Config.BATCH_SIZE, shuffle=False,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    print(f\"Building {Config.MODEL_NAME}...\")\n    model     = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    criterion = nn.CrossEntropyLoss(label_smoothing=Config.LABEL_SMOOTH)\n    optimizer = torch.optim.AdamW(model.parameters(),\n                                  lr=Config.LR, weight_decay=Config.WEIGHT_DECAY)\n    scaler    = GradScaler()\n    total_steps  = Config.EPOCHS * len(train_loader)\n    warmup_steps = Config.WARMUP_EP * len(train_loader)\n    scheduler    = build_scheduler(optimizer, total_steps, warmup_steps)\n    best_acc  = 0.0\n    best_path = 'best_resnext50_32x4d_controls.pth'\n    print(f\"\\nPhase 1 — backbone frozen for {Config.WARMUP_EP} epoch(s)\")\n    model.freeze_backbone()\n    for epoch in range(Config.EPOCHS):\n        if epoch == Config.WARMUP_EP:\n            model.unfreeze_backbone()\n            print(f\"\\nPhase 2 — full fine-tune from epoch {epoch+1}\")\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}  \"\n              f\"lr={optimizer.param_groups[0]['lr']:.2e}\")\n        tr_loss, tr_acc = train_epoch(model, train_loader, criterion,\n                                      optimizer, scheduler, scaler, device)\n        vl_loss, vl_acc = validate(model, val_loader, criterion, device)\n        print(f\"  Train  loss={tr_loss:.4f}  acc={tr_acc:.2f}%\")\n        print(f\"  Val    loss={vl_loss:.4f}  acc={vl_acc:.2f}%\")\n        if vl_acc > best_acc:\n            best_acc = vl_acc\n            torch.save(model.state_dict(), best_path)\n            print(f\"  ✓ Saved best model (val acc={best_acc:.2f}%)\")\n    print(f\"\\nLoading best weights (val acc={best_acc:.2f}%)...\")\n    model.load_state_dict(torch.load(best_path, map_location=device))\n    print(\"Running TTA inference...\")\n    preds, ids = predict_tta(model, test_df, Config.DATA_DIR, stats_lookup,\n                             device, plate_stats=plate_stats_test)\n    sirna_preds = [label_to_sirna[int(p)] for p in preds]\n    submission  = pd.DataFrame({'id_code': ids, 'sirna': sirna_preds})\n    submission.to_csv('submission_resnext50_32x4d_controls.csv', index=False)\n    print(f\"\\nDone. submission_resnext50_32x4d_controls.csv saved  |  Best val acc: {best_acc:.2f}%\")\n    print(submission.head())\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-28T22:25:40.426318Z","iopub.execute_input":"2026-04-28T22:25:40.426681Z","iopub.status.idle":"2026-04-28T22:25:48.143198Z","shell.execute_reply.started":"2026-04-28T22:25:40.426653Z","shell.execute_reply":"2026-04-28T22:25:48.142037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport cv2\nimport warnings\nwarnings.filterwarnings('ignore')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV = f'{DATA_DIR}/train.csv'\n    TEST_CSV = f'{DATA_DIR}/test.csv'\n    PIXEL_STATS = f'{DATA_DIR}/pixel_stats.csv'\n    TRAIN_CONTROLS_CSV = f'{DATA_DIR}/train_controls.csv'\n    TEST_CONTROLS_CSV = f'{DATA_DIR}/test_controls.csv'\n    \n    MODEL_NAME = 'resnet50'\n    IMG_SIZE = 320\n    BATCH_SIZE = 32\n    EPOCHS = 50\n    WARMUP_EP = 1\n    LR = 3e-4\n    MIN_LR = 1e-6\n    WEIGHT_DECAY = 1e-4\n    LABEL_SMOOTH = 0.1\n    \n    NUM_WORKERS = 2\n    SEED = 42\n    NUM_CLASSES = 1108\n    \n    CELL_TYPES = ['HUVEC']\n    \n    USE_PLATE_NORMALIZATION = True\n    USE_CONTROLS_AS_EXTRA_DATA = True\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\nset_seed(Config.SEED)\ndef build_stats_lookup(pixel_stats_path):\n    \n    stats = pd.read_csv(pixel_stats_path)\n    lookup = {}\n    for _, row in stats.iterrows():\n        key = (row['experiment'], int(row['plate']),\n               row['well'], int(row['site']),\n               int(row['channel']))\n        lookup[key] = (float(row['mean']), float(row['std']))\n    return lookup\ndef compute_plate_control_stats(controls_df, data_dir, mode='train'):\n    \n    stats = {}\n    grouped = controls_df.groupby(['experiment', 'plate'])\n    \n    print(f\"Computing plate control stats from {mode}_controls.csv ...\")\n    for (exp, plate), group in tqdm(grouped, desc='Plate stats'):\n        accum = None\n        count = 0\n        for _, row in group.iterrows():\n            well = row['well']\n            path_template = f'{data_dir}/{mode}/{exp}/Plate{plate}/{well}_s1_w'\n            channels = []\n            for ch in range(1, 7):\n                img_path = f'{path_template}{ch}.png'\n                if os.path.exists(img_path):\n                    img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                else:\n                    img = np.zeros((512, 512), dtype=np.uint8)\n                channels.append(img)\n            img = np.stack(channels, axis=-1).astype(np.float32) / 255.0\n            if accum is None:\n                accum = img\n            else:\n                accum += img\n            count += 1\n        \n        if count > 0:\n            stats[(exp, plate)] = accum / count\n    \n    print(f\"  Computed stats for {len(stats)} plates\")\n    return stats\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, stats_lookup, mode='train',\n                 site=None, plate_stats=None):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.stats_lookup = stats_lookup\n        self.mode = mode\n        self.site = site\n        self.plate_stats = plate_stats\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def _load_channels(self, row, site):\n        \n        exp = row['experiment']\n        plate = int(row['plate'])\n        well = row['well']\n        split = 'test' if self.mode == 'test' else 'train'\n        base = f'{self.data_dir}/{split}/{exp}/Plate{plate}/{well}_s{site}_w'\n        \n        channels = []\n        for ch in range(1, 7):\n            path = f'{base}{ch}.png'\n            img = cv2.imread(path, cv2.IMREAD_GRAYSCALE) if os.path.exists(path) \\\n                  else np.zeros((512, 512), dtype=np.uint8)\n            img = img.astype(np.float32)\n            \n            key = (exp, plate, well, int(site), ch)\n            mean, std = self.stats_lookup.get(key, (127.5, 64.0))\n            img = (img - mean) / (std + 1e-6)\n            channels.append(img)\n        \n        stacked = np.stack(channels, axis=-1)\n        stacked = cv2.resize(stacked, (Config.IMG_SIZE, Config.IMG_SIZE))\n        \n        if self.plate_stats is not None:\n            ctrl_key = (exp, plate)\n            if ctrl_key in self.plate_stats:\n                ctrl_mean = cv2.resize(\n                    self.plate_stats[ctrl_key],\n                    (Config.IMG_SIZE, Config.IMG_SIZE)\n                )\n                stacked = stacked - (ctrl_mean * 255.0 - 127.5) / 64.0\n        \n        return stacked\n    \n    def _augment(self, img):\n        \n        if self.mode != 'train':\n            return img\n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        return img\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        site = np.random.choice(['1', '2']) if self.mode == 'train' \\\n               else (self.site or '1')\n        \n        img = self._load_channels(row, site)\n        img = self._augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1).float()\n        \n        if self.mode in ('train', 'val'):\n            return img, int(row['label'])\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        \n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        \n        if hasattr(self.backbone, 'conv1'):\n            old_conv = self.backbone.conv1\n            self.backbone.conv1 = nn.Conv2d(\n                in_channels, old_conv.out_channels,\n                kernel_size=old_conv.kernel_size,\n                stride=old_conv.stride,\n                padding=old_conv.padding,\n                bias=old_conv.bias is not None\n            )\n            \n            with torch.no_grad():\n                self.backbone.conv1.weight[:, :3] = old_conv.weight\n                self.backbone.conv1.weight[:, 3:] = old_conv.weight\n                if old_conv.bias is not None:\n                    self.backbone.conv1.bias = nn.Parameter(old_conv.bias.clone())\n        \n        n_features = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        \n        self.classifier = nn.Sequential(\n            nn.BatchNorm1d(n_features),\n            nn.Dropout(0.4),\n            nn.Linear(n_features, 512),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.3),\n            nn.Linear(512, num_classes)\n        )\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        return self.classifier(features)\n    \n    def freeze_backbone(self):\n        for p in self.backbone.parameters():\n            p.requires_grad = False\n    \n    def unfreeze_backbone(self):\n        for p in self.backbone.parameters():\n            p.requires_grad = True\ndef build_scheduler(optimizer, total_steps, warmup_steps):\n    def lr_lambda(step):\n        if step < warmup_steps:\n            return float(step) / max(1, warmup_steps)\n        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n        cosine = 0.5 * (1.0 + np.cos(np.pi * progress))\n        return max(Config.MIN_LR / Config.LR, cosine)\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\ndef train_epoch(model, loader, criterion, optimizer, scheduler, scaler, device):\n    model.train()\n    total_loss = correct = total = 0\n    \n    pbar = tqdm(loader, desc='Train', leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        \n        with autocast():\n            out = model(imgs)\n            loss = criterion(out, labels)\n        \n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        scheduler.step()\n        \n        total_loss += loss.item()\n        _, pred = out.max(1)\n        correct += pred.eq(labels).sum().item()\n        total += labels.size(0)\n        pbar.set_postfix(loss=f'{total_loss/len(loader):.4f}',\n                        acc=f'{100.*correct/total:.2f}%')\n    \n    return total_loss / len(loader), 100. * correct / total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    total_loss = correct = total = 0\n    \n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Val  ', leave=False):\n            imgs, labels = imgs.to(device), labels.to(device)\n            with autocast():\n                out = model(imgs)\n                loss = criterion(out, labels)\n            total_loss += loss.item()\n            _, pred = out.max(1)\n            correct += pred.eq(labels).sum().item()\n            total += labels.size(0)\n    \n    return total_loss / len(loader), 100. * correct / total\ndef predict_tta(model, test_df, data_dir, stats_lookup, device, plate_stats=None):\n    \n    model.eval()\n    tta_configs = [\n        ('1', None), ('1', 'h'), ('1', 'v'),\n        ('2', None), ('2', 'h'), ('2', 'v'),\n    ]\n    all_probs = []\n    ids = None\n    \n    for site, flip in tta_configs:\n        ds = CellularDataset(test_df, data_dir, stats_lookup,\n                            mode='test', site=site, plate_stats=plate_stats)\n        loader = DataLoader(ds, batch_size=Config.BATCH_SIZE,\n                           shuffle=False, num_workers=Config.NUM_WORKERS,\n                           pin_memory=True)\n        \n        probs_list = []\n        id_list = []\n        \n        with torch.no_grad():\n            for imgs, img_ids in tqdm(loader,\n                                     desc=f'TTA site{site} {flip or \"orig\"}',\n                                     leave=False):\n                imgs = imgs.to(device)\n                if flip == 'h':\n                    imgs = torch.flip(imgs, dims=[3])\n                elif flip == 'v':\n                    imgs = torch.flip(imgs, dims=[2])\n                with autocast():\n                    out = model(imgs)\n                probs_list.append(out.softmax(dim=1).cpu())\n                id_list.extend(img_ids)\n        \n        all_probs.append(torch.cat(probs_list, dim=0))\n        if ids is None:\n            ids = id_list\n    \n    avg_probs = torch.stack(all_probs).mean(dim=0)\n    preds = avg_probs.argmax(dim=1).numpy()\n    return preds, ids\ndef main():\n    print(\"=\" * 60)\n    print(\"ResNet50 with Controls and Plate Normalization\")\n    print(\"=\" * 60)\n    \n    print(\"\\nLoading data...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df = pd.read_csv(Config.TEST_CSV)\n    train_controls_df = pd.read_csv(Config.TRAIN_CONTROLS_CSV)\n    test_controls_df = pd.read_csv(Config.TEST_CONTROLS_CSV)\n    \n    print(f\"  train_controls.csv: {len(train_controls_df)} rows\")\n    print(f\"  test_controls.csv: {len(test_controls_df)} rows\")\n    \n    train_df['cell_type'] = train_df['experiment'].str.split('-').str[0]\n    test_df['cell_type'] = test_df['experiment'].str.split('-').str[0]\n    \n    train_df = train_df[train_df['cell_type'].isin(Config.CELL_TYPES)].reset_index(drop=True)\n    \n    print(f\"\\nTraining samples: {len(train_df)}\")\n    print(f\"Test samples: {len(test_df)}\")\n    print(f\"Cell types: {train_df['cell_type'].unique()}\")\n    \n    train_df['sirna_id'] = train_df['sirna'].str.replace('sirna_', '').astype(int)\n    unique_sirnas = sorted(train_df['sirna_id'].unique())\n    sirna_to_label = {s: i for i, s in enumerate(unique_sirnas)}\n    label_to_sirna = {i: s for s, i in sirna_to_label.items()}\n    train_df['label'] = train_df['sirna_id'].map(sirna_to_label)\n    \n    Config.NUM_CLASSES = len(unique_sirnas)\n    print(f\"Number of classes: {Config.NUM_CLASSES}\")\n    \n    control_df = pd.DataFrame()\n    if Config.USE_CONTROLS_AS_EXTRA_DATA:\n        print(\"\\n--- Adding control wells as extra training data ---\")\n        \n        neg_mask = train_controls_df['sirna'].astype(str).str.contains(\n            'negative', case=False, na=False\n        )\n        if neg_mask.sum() == 0:\n            neg_mask = pd.Series([True] * len(train_controls_df))\n        \n        neg_controls = train_controls_df[neg_mask].copy()\n        print(f\"  Found {len(neg_controls)} negative control wells\")\n        \n        required_cols = ['experiment', 'plate', 'well']\n        if all(c in neg_controls.columns for c in required_cols) and len(neg_controls) > 0:\n            rng = np.random.default_rng(Config.SEED)\n            neg_controls = neg_controls[required_cols].copy()\n            \n            neg_controls['sirna'] = 'negative_control'\n            neg_controls['sirna_id'] = rng.choice(unique_sirnas, size=len(neg_controls))\n            neg_controls['label'] = neg_controls['sirna_id'].map(sirna_to_label)\n            \n            frac = min(1.0, len(train_df) / max(1, len(neg_controls)))\n            control_df = neg_controls.sample(frac=frac, random_state=Config.SEED)\n            print(f\"  Adding {len(control_df)} control samples to training\")\n            \n            train_df = pd.concat([train_df, control_df], ignore_index=True)\n            train_df = train_df.sample(frac=1, random_state=Config.SEED).reset_index(drop=True)\n            print(f\"  New training size: {len(train_df)}\")\n    \n    plate_stats_train = None\n    plate_stats_test = None\n    \n    if Config.USE_PLATE_NORMALIZATION:\n        print(\"\\n--- Computing plate control statistics ---\")\n        plate_stats_train = compute_plate_control_stats(\n            train_controls_df, Config.DATA_DIR, mode='train'\n        )\n        plate_stats_test = compute_plate_control_stats(\n            test_controls_df, Config.DATA_DIR, mode='test'\n        )\n        for key, val in plate_stats_train.items():\n            if key not in plate_stats_test:\n                plate_stats_test[key] = val\n    \n    real_train_mask = train_df['sirna'] != 'negative_control'\n    real_train_df = train_df[real_train_mask]\n    \n    train_data, val_data = train_test_split(\n        real_train_df, test_size=0.15, random_state=Config.SEED,\n        stratify=real_train_df['label']\n    )\n    \n    if len(control_df) > 0:\n        train_data = pd.concat([train_data, control_df], ignore_index=True)\n        train_data = train_data.sample(frac=1, random_state=Config.SEED).reset_index(drop=True)\n    \n    print(f\"\\nTrain: {len(train_data)}, Val: {len(val_data)}\")\n    \n    print(\"\\nBuilding pixel-stats lookup...\")\n    stats_lookup = build_stats_lookup(Config.PIXEL_STATS)\n    \n    train_loader = DataLoader(\n        CellularDataset(train_data, Config.DATA_DIR, stats_lookup,\n                       mode='train', plate_stats=plate_stats_train),\n        batch_size=Config.BATCH_SIZE, shuffle=True,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    val_loader = DataLoader(\n        CellularDataset(val_data, Config.DATA_DIR, stats_lookup,\n                       mode='val', site='1', plate_stats=plate_stats_train),\n        batch_size=Config.BATCH_SIZE, shuffle=False,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    \n    print(f\"\\nCreating model: {Config.MODEL_NAME}...\")\n    model = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    \n    criterion = nn.CrossEntropyLoss(label_smoothing=Config.LABEL_SMOOTH)\n    optimizer = torch.optim.AdamW(model.parameters(),\n                                 lr=Config.LR, weight_decay=Config.WEIGHT_DECAY)\n    scaler = GradScaler()\n    \n    total_steps = Config.EPOCHS * len(train_loader)\n    warmup_steps = Config.WARMUP_EP * len(train_loader)\n    scheduler = build_scheduler(optimizer, total_steps, warmup_steps)\n    \n    best_acc = 0\n    \n    print(f\"\\nPhase 1 - backbone frozen for {Config.WARMUP_EP} epoch(s)\")\n    model.freeze_backbone()\n    \n    for epoch in range(Config.EPOCHS):\n        if epoch == Config.WARMUP_EP:\n            model.unfreeze_backbone()\n            print(f\"\\nPhase 2 - full fine-tune from epoch {epoch+1}\")\n        \n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}  \"\n              f\"lr={optimizer.param_groups[0]['lr']:.2e}\")\n        \n        train_loss, train_acc = train_epoch(model, train_loader, criterion,\n                                           optimizer, scheduler, scaler, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        \n        print(f\"Train Loss: {train_loss:.4f}, Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_resnet50_baseline.pth')\n            print(f\"Saved best model: {best_acc:.2f}%\")\n    \n    print(f\"\\n{'='*60}\")\n    print(f\"Loading best weights (val acc={best_acc:.2f}%)...\")\n    model.load_state_dict(torch.load('best_resnet50_baseline.pth', map_location=device))\n    \n    print(\"Running TTA inference...\")\n    preds, ids = predict_tta(model, test_df, Config.DATA_DIR, stats_lookup,\n                            device, plate_stats=plate_stats_test)\n    \n    predictions = [label_to_sirna[int(p)] for p in preds]\n    submission = pd.DataFrame({'id_code': ids, 'sirna': predictions})\n    submission.to_csv('submission_resnet50_baseline.csv', index=False)\n    \n    print(f\"\\n{'='*60}\")\n    print(f\"Complete! Best Val Acc: {best_acc:.2f}%\")\n    print(\"Saved submission_resnet50_baseline.csv\")\n    print(submission.head())\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-29T10:01:34.598326Z","iopub.execute_input":"2026-04-29T10:01:34.598838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport timm\nfrom tqdm import tqdm\nimport cv2\nimport warnings\nwarnings.filterwarnings('ignore')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nclass Config:\n    DATA_DIR    = '/kaggle/input/competitions/recursion-cellular-image-classification'\n    TRAIN_CSV   = f'{DATA_DIR}/train.csv'\n    TEST_CSV    = f'{DATA_DIR}/test.csv'\n    PIXEL_STATS = f'{DATA_DIR}/pixel_stats.csv'\n    TRAIN_CONTROLS_CSV = f'{DATA_DIR}/train_controls.csv'\n    TEST_CONTROLS_CSV  = f'{DATA_DIR}/test_controls.csv'\n    MODEL_NAME   = 'resnext50_32x4d'\n    IMG_SIZE     = 384\n    BATCH_SIZE   = 16           \n    EPOCHS       = 50\n    WARMUP_EP    = 1           \n    LR           = 3e-4\n    MIN_LR       = 1e-6\n    WEIGHT_DECAY = 1e-4\n    LABEL_SMOOTH = 0.1\n    NUM_WORKERS  = 4\n    SEED         = 42\n    NUM_CLASSES  = 1108         \n    VAL_EXP_FRAC = 0.15        \n    USE_MSAM     = True\n    MSAM_RHO     = 0.05\n    MSAM_MOMENTUM = 0.9\n    USE_PLATE_NORMALIZATION    = True\n    USE_CONTROLS_AS_EXTRA_DATA = True\ndef set_seed(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark     = False\nset_seed(Config.SEED)\ndef build_stats_lookup(pixel_stats_path):\n    \n    stats  = pd.read_csv(pixel_stats_path)\n    lookup = {}\n    for _, row in stats.iterrows():\n        key = (row['experiment'], int(row['plate']),\n               row['well'],       int(row['site']),\n               int(row['channel']))\n        lookup[key] = (float(row['mean']), float(row['std']))\n    return lookup\ndef compute_plate_control_stats(controls_df, data_dir, mode='train'):\n    \n    stats = {}\n    grouped = controls_df.groupby(['experiment', 'plate'])\n    print(f\"Computing plate control stats from {mode}_controls.csv ...\")\n    for (exp, plate), group in tqdm(grouped, desc='Plate stats'):\n        accum = None\n        count = 0\n        for _, row in group.iterrows():\n            well = row['well']\n            path_template = f'{data_dir}/{mode}/{exp}/Plate{plate}/{well}_s1_w'\n            channels = []\n            for ch in range(1, 7):\n                img_path = f'{path_template}{ch}.png'\n                if os.path.exists(img_path):\n                    img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n                else:\n                    img = np.zeros((512, 512), dtype=np.uint8)\n                channels.append(img)\n            img = np.stack(channels, axis=-1).astype(np.float32) / 255.0\n            if accum is None:\n                accum = img\n            else:\n                accum += img\n            count += 1\n        if count > 0:\n            stats[(exp, plate)] = accum / count\n    print(f\"  Computed stats for {len(stats)} plates\")\n    return stats\nclass CellularDataset(Dataset):\n    def __init__(self, df, data_dir, stats_lookup, mode='train',\n                 site=None, plate_stats=None):\n        self.df           = df.reset_index(drop=True)\n        self.data_dir     = data_dir\n        self.stats_lookup = stats_lookup\n        self.mode         = mode\n        self.site         = site\n        self.plate_stats  = plate_stats\n    def __len__(self):\n        return len(self.df)\n    def _load_channels(self, row, site):\n        \n        exp   = row['experiment']\n        plate = int(row['plate'])\n        well  = row['well']\n        split = 'test' if self.mode == 'test' else 'train'\n        base  = f'{self.data_dir}/{split}/{exp}/Plate{plate}/{well}_s{site}_w'\n        channels = []\n        for ch in range(1, 7):\n            path = f'{base}{ch}.png'\n            img  = cv2.imread(path, cv2.IMREAD_GRAYSCALE) if os.path.exists(path) \\\n                   else np.zeros((512, 512), dtype=np.uint8)\n            img  = img.astype(np.float32)\n            key        = (exp, plate, well, int(site), ch)\n            mean, std  = self.stats_lookup.get(key, (127.5, 64.0))\n            img        = (img - mean) / (std + 1e-6)\n            channels.append(img)\n        stacked = np.stack(channels, axis=-1)\n        stacked = cv2.resize(stacked, (Config.IMG_SIZE, Config.IMG_SIZE))\n        if self.plate_stats is not None:\n            ctrl_key = (exp, plate)\n            if ctrl_key in self.plate_stats:\n                ctrl_mean = cv2.resize(\n                    self.plate_stats[ctrl_key],\n                    (Config.IMG_SIZE, Config.IMG_SIZE)\n                )\n                stacked = stacked - ctrl_mean\n        return stacked\n    def _augment(self, img):\n        \n        if self.mode != 'train':\n            return img\n        if np.random.rand() > 0.5:\n            img = np.fliplr(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.flipud(img).copy()\n        if np.random.rand() > 0.5:\n            img = np.rot90(img, k=np.random.randint(1, 4)).copy()\n        return img\n    def __getitem__(self, idx):\n        row  = self.df.iloc[idx]\n        site = np.random.choice(['1', '2']) if self.mode == 'train' \\\n               else (self.site or '1')\n        img = self._load_channels(row, site)\n        img = self._augment(img)\n        img = torch.from_numpy(img).permute(2, 0, 1).float()\n        if self.mode in ('train', 'val'):\n            return img, int(row['label'])\n        else:\n            return img, row['id_code']\nclass CellularModel(nn.Module):\n    \n    def __init__(self, model_name, num_classes, in_channels=6):\n        super().__init__()\n        self.backbone = timm.create_model(model_name, pretrained=True, in_chans=3)\n        if hasattr(self.backbone, 'conv1'):\n            old = self.backbone.conv1\n            new = nn.Conv2d(\n                in_channels, old.out_channels,\n                kernel_size=old.kernel_size,\n                stride=old.stride,\n                padding=old.padding,\n                bias=old.bias is not None\n            )\n            with torch.no_grad():\n                new.weight[:, :3, ...] = old.weight\n                new.weight[:, 3:, ...] = old.weight\n                if old.bias is not None:\n                    new.bias = nn.Parameter(old.bias.clone())\n            self.backbone.conv1 = new\n        n_feat = self.backbone.get_classifier().in_features\n        self.backbone.reset_classifier(0)\n        self.head = nn.Sequential(\n            nn.BatchNorm1d(n_feat),\n            nn.Dropout(0.4),\n            nn.Linear(n_feat, 512),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.3),\n            nn.Linear(512, num_classes)\n        )\n    def forward(self, x):\n        return self.head(self.backbone(x))\n    def freeze_backbone(self):\n        for p in self.backbone.parameters():\n            p.requires_grad = False\n    def unfreeze_backbone(self):\n        for p in self.backbone.parameters():\n            p.requires_grad = True\ndef build_scheduler(optimizer, total_steps, warmup_steps):\n    def lr_lambda(step):\n        if step < warmup_steps:\n            return float(step) / max(1, warmup_steps)\n        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)\n        cosine   = 0.5 * (1.0 + np.cos(np.pi * progress))\n        return max(Config.MIN_LR / Config.LR, cosine)\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\nclass MomentumSAM:\n    \n    def __init__(self, params, lr=3e-4, momentum=0.9, rho=0.05, weight_decay=1e-4):\n        self.params = list(params)\n        self._lr = lr\n        self.momentum = momentum\n        self.rho = rho\n        self.weight_decay = weight_decay\n        \n        self.param_groups = [{'lr': lr, 'params': self.params}]\n        \n        self.momentum_buffer = {}\n        for p in self.params:\n            if p.requires_grad:\n                self.momentum_buffer[p] = torch.zeros_like(p.data)\n        \n        self.perturbed_params = {}\n        self.is_perturbed = False\n    \n    @property\n    def lr(self):\n        return self.param_groups[0]['lr']\n    \n    @lr.setter\n    def lr(self, value):\n        self._lr = value\n        self.param_groups[0]['lr'] = value\n        \n    def _normalize_momentum(self):\n        \n        total_norm = 0.0\n        for p in self.params:\n            if p.requires_grad and p in self.momentum_buffer:\n                total_norm += self.momentum_buffer[p].pow(2).sum().item()\n        return max(total_norm ** 0.5, 1e-12)\n    \n    def perturb_parameters(self):\n        \n        if self.rho == 0:\n            return\n            \n        momentum_norm = self._normalize_momentum()\n        \n        for p in self.params:\n            if p.requires_grad and p in self.momentum_buffer:\n                self.perturbed_params[p] = p.data.clone()\n                p.data.add_(self.momentum_buffer[p], alpha=-self.rho / momentum_norm)\n        \n        self.is_perturbed = True\n    \n    def remove_perturbation(self):\n        \n        if not self.is_perturbed or self.rho == 0:\n            return\n            \n        for p in self.params:\n            if p.requires_grad and p in self.perturbed_params:\n                p.data.copy_(self.perturbed_params[p])\n        \n        self.is_perturbed = False\n    \n    def step(self, closure=None):\n        \n        \n        self.remove_perturbation()\n        \n        momentum_norm = self._normalize_momentum()\n        \n        for p in self.params:\n            if p.requires_grad and p.grad is not None:\n                grad = p.grad.data\n                \n                if self.weight_decay != 0:\n                    grad = grad.add(p.data, alpha=self.weight_decay)\n                \n                buf = self.momentum_buffer[p]\n                buf.mul_(self.momentum).add_(grad)\n                \n                p.data.add_(buf, alpha=-self.lr)\n        \n        self.perturb_parameters()\n        \n    def zero_grad(self):\n        \n        for p in self.params:\n            if p.grad is not None:\n                p.grad.zero_()\n    \n    def state_dict(self):\n        \n        return {\n            'momentum_buffer': {id(p): v.clone() for p, v in self.momentum_buffer.items()},\n            'perturbed_params': {id(p): v.clone() for p, v in self.perturbed_params.items()},\n            'is_perturbed': self.is_perturbed,\n            'lr': self.lr,\n            'momentum': self.momentum,\n            'rho': self.rho,\n            'weight_decay': self.weight_decay\n        }\n    \n    def load_state_dict(self, state_dict):\n        \n        self.lr = state_dict['lr']\n        self.momentum = state_dict['momentum']\n        self.rho = state_dict['rho']\n        self.weight_decay = state_dict['weight_decay']\n        self.is_perturbed = state_dict['is_perturbed']\n        \n        for p in self.params:\n            if id(p) in state_dict['momentum_buffer']:\n                self.momentum_buffer[p] = state_dict['momentum_buffer'][id(p)]\n            if id(p) in state_dict['perturbed_params']:\n                self.perturbed_params[p] = state_dict['perturbed_params'][id(p)]\n    \n    def finalize(self):\n        \n        self.remove_perturbation()\ndef train_epoch(model, loader, criterion, optimizer, scheduler, scaler, device):\n    model.train()\n    total_loss = correct = total = 0\n    pbar = tqdm(loader, desc='Train', leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        with autocast():\n            out  = model(imgs)\n            loss = criterion(out, labels)\n        scaler.scale(loss).backward()\n        \n        if isinstance(optimizer, MomentumSAM):\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            optimizer.step()\n            scaler.update()\n        else:\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            \n        scheduler.step()\n        total_loss += loss.item()\n        _, pred     = out.max(1)\n        correct    += pred.eq(labels).sum().item()\n        total      += labels.size(0)\n        pbar.set_postfix(loss=f'{total_loss/len(loader):.4f}',\n                         acc=f'{100.*correct/total:.2f}%')\n    return total_loss / len(loader), 100. * correct / total\ndef validate(model, loader, criterion, device):\n    model.eval()\n    total_loss = correct = total = 0\n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc='Val  ', leave=False):\n            imgs, labels = imgs.to(device), labels.to(device)\n            with autocast():\n                out  = model(imgs)\n                loss = criterion(out, labels)\n            total_loss += loss.item()\n            _, pred     = out.max(1)\n            correct    += pred.eq(labels).sum().item()\n            total      += labels.size(0)\n    return total_loss / len(loader), 100. * correct / total\ndef predict_tta(model, test_df, data_dir, stats_lookup, device,\n                plate_stats=None):\n    \n    model.eval()\n    tta_configs = [\n        ('1', None), ('1', 'h'), ('1', 'v'),\n        ('2', None), ('2', 'h'), ('2', 'v'),\n    ]\n    all_probs = []\n    ids       = None\n    for site, flip in tta_configs:\n        ds     = CellularDataset(test_df, data_dir, stats_lookup,\n                                 mode='test', site=site,\n                                 plate_stats=plate_stats)\n        loader = DataLoader(ds, batch_size=Config.BATCH_SIZE,\n                            shuffle=False, num_workers=Config.NUM_WORKERS,\n                            pin_memory=True)\n        probs_list = []\n        id_list    = []\n        with torch.no_grad():\n            for imgs, img_ids in tqdm(loader,\n                                      desc=f'TTA site{site} {flip or \"orig\"}',\n                                      leave=False):\n                imgs = imgs.to(device)\n                if flip == 'h':\n                    imgs = torch.flip(imgs, dims=[3])\n                elif flip == 'v':\n                    imgs = torch.flip(imgs, dims=[2])\n                with autocast():\n                    out = model(imgs)\n                probs_list.append(out.softmax(dim=1).cpu())\n                id_list.extend(img_ids)\n        all_probs.append(torch.cat(probs_list, dim=0))\n        if ids is None:\n            ids = id_list\n    avg_probs = torch.stack(all_probs).mean(dim=0)\n    preds     = avg_probs.argmax(dim=1).numpy()\n    return preds, ids\ndef main():\n    print(\"Loading CSVs...\")\n    train_df = pd.read_csv(Config.TRAIN_CSV)\n    test_df  = pd.read_csv(Config.TEST_CSV)\n    train_controls_df = pd.read_csv(Config.TRAIN_CONTROLS_CSV)\n    test_controls_df  = pd.read_csv(Config.TEST_CONTROLS_CSV)\n    print(f\"  train_controls.csv: {len(train_controls_df)} rows, columns: {list(train_controls_df.columns)}\")\n    print(f\"  test_controls.csv:  {len(test_controls_df)} rows, columns: {list(test_controls_df.columns)}\")\n    if train_df['sirna'].dtype == object:\n        train_df['sirna_id'] = train_df['sirna'].str.extract(r'(\\d+)').astype(int)\n    else:\n        train_df['sirna_id'] = train_df['sirna'].astype(int)\n    unique_sirnas      = sorted(train_df['sirna_id'].unique())\n    sirna_to_label     = {s: i for i, s in enumerate(unique_sirnas)}\n    label_to_sirna     = {i: s for s, i in sirna_to_label.items()}\n    train_df['label']  = train_df['sirna_id'].map(sirna_to_label)\n    Config.NUM_CLASSES = len(unique_sirnas)\n    print(f\"Classes      : {Config.NUM_CLASSES}\")\n    print(f\"Train wells  : {len(train_df)}\")\n    print(f\"Test wells   : {len(test_df)}\")\n    control_df = pd.DataFrame()\n    if Config.USE_CONTROLS_AS_EXTRA_DATA:\n        print(\"\\n--- Adding control wells as extra training data ---\")\n        print(f\"  Controls sirna values (sample): {train_controls_df['sirna'].unique()[:10]}\")\n        neg_mask = train_controls_df['sirna'].astype(str).str.contains(\n            'negative', case=False, na=False\n        )\n        if neg_mask.sum() == 0:\n            neg_mask = pd.Series([True] * len(train_controls_df))\n        neg_controls = train_controls_df[neg_mask].copy()\n        print(f\"  Found {len(neg_controls)} negative control wells\")\n        required_cols = ['experiment', 'plate', 'well']\n        if all(c in neg_controls.columns for c in required_cols) and len(neg_controls) > 0:\n            rng = np.random.default_rng(Config.SEED)\n            neg_controls = neg_controls[required_cols].copy()\n            neg_controls['sirna'] = 'negative_control'\n            neg_controls['sirna_id'] = rng.choice(unique_sirnas, size=len(neg_controls))\n            neg_controls['label'] = neg_controls['sirna_id'].map(sirna_to_label)\n            frac = min(1.0, len(train_df) / max(1, len(neg_controls)))\n            control_df = neg_controls.sample(frac=frac, random_state=Config.SEED)\n            print(f\"  Adding {len(control_df)} control samples to training\")\n            train_df = pd.concat([train_df, control_df], ignore_index=True)\n            train_df = train_df.sample(frac=1, random_state=Config.SEED).reset_index(drop=True)\n            print(f\"  New training size: {len(train_df)}\")\n    plate_stats_train = None\n    plate_stats_test  = None\n    if Config.USE_PLATE_NORMALIZATION:\n        plate_stats_train = compute_plate_control_stats(\n            train_controls_df, Config.DATA_DIR, mode='train'\n        )\n        plate_stats_test = compute_plate_control_stats(\n            test_controls_df, Config.DATA_DIR, mode='test'\n        )\n        for key, val in plate_stats_train.items():\n            if key not in plate_stats_test:\n                plate_stats_test[key] = val\n    real_train_mask = train_df['sirna'] != 'negative_control'\n    real_train_df   = train_df[real_train_mask]\n    rng      = np.random.default_rng(Config.SEED)\n    exps     = real_train_df['experiment'].unique()\n    n_val    = max(1, int(len(exps) * Config.VAL_EXP_FRAC))\n    val_exps = set(rng.choice(exps, size=n_val, replace=False))\n    val_df   = real_train_df[real_train_df['experiment'].isin(val_exps)].reset_index(drop=True)\n    tr_df    = real_train_df[~real_train_df['experiment'].isin(val_exps)].reset_index(drop=True)\n    if len(control_df) > 0:\n        tr_df = pd.concat([tr_df, control_df], ignore_index=True)\n        tr_df = tr_df.sample(frac=1, random_state=Config.SEED).reset_index(drop=True)\n    print(f\"Train exps   : {len(exps)-n_val}  ({len(tr_df)} wells)\")\n    print(f\"Val exps     : {n_val}  ({len(val_df)} wells)\")\n    print(\"Building pixel-stats lookup...\")\n    stats_lookup = build_stats_lookup(Config.PIXEL_STATS)\n    train_loader = DataLoader(\n        CellularDataset(tr_df,  Config.DATA_DIR, stats_lookup,\n                        mode='train', plate_stats=plate_stats_train),\n        batch_size=Config.BATCH_SIZE, shuffle=True,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    val_loader = DataLoader(\n        CellularDataset(val_df, Config.DATA_DIR, stats_lookup,\n                        mode='val', site='1',\n                        plate_stats=plate_stats_train),\n        batch_size=Config.BATCH_SIZE, shuffle=False,\n        num_workers=Config.NUM_WORKERS, pin_memory=True\n    )\n    print(f\"Building {Config.MODEL_NAME}...\")\n    model     = CellularModel(Config.MODEL_NAME, Config.NUM_CLASSES).to(device)\n    criterion = nn.CrossEntropyLoss(label_smoothing=Config.LABEL_SMOOTH)\n    \n    if Config.USE_MSAM:\n        print(f\"Using Momentum-SAM optimizer (rho={Config.MSAM_RHO}, momentum={Config.MSAM_MOMENTUM})\")\n        optimizer = MomentumSAM(\n            model.parameters(),\n            lr=Config.LR,\n            momentum=Config.MSAM_MOMENTUM,\n            rho=Config.MSAM_RHO,\n            weight_decay=Config.WEIGHT_DECAY\n        )\n        optimizer.perturb_parameters()\n    else:\n        print(\"Using standard AdamW optimizer\")\n        optimizer = torch.optim.AdamW(\n            model.parameters(),\n            lr=Config.LR,\n            weight_decay=Config.WEIGHT_DECAY\n        )\n    \n    scaler    = GradScaler()\n    total_steps  = Config.EPOCHS * len(train_loader)\n    warmup_steps = Config.WARMUP_EP * len(train_loader)\n    scheduler    = build_scheduler(optimizer, total_steps, warmup_steps)\n    best_acc  = 0.0\n    best_path = 'best_resnext50_32x4d_controls.pth'\n    print(f\"\\nPhase 1 — backbone frozen for {Config.WARMUP_EP} epoch(s)\")\n    model.freeze_backbone()\n    for epoch in range(Config.EPOCHS):\n        if epoch == Config.WARMUP_EP:\n            model.unfreeze_backbone()\n            print(f\"\\nPhase 2 — full fine-tune from epoch {epoch+1}\")\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}  \"\n              f\"lr={optimizer.param_groups[0]['lr']:.2e}\")\n        tr_loss, tr_acc = train_epoch(model, train_loader, criterion,\n                                      optimizer, scheduler, scaler, device)\n        vl_loss, vl_acc = validate(model, val_loader, criterion, device)\n        print(f\"  Train  loss={tr_loss:.4f}  acc={tr_acc:.2f}%\")\n        print(f\"  Val    loss={vl_loss:.4f}  acc={vl_acc:.2f}%\")\n        if vl_acc > best_acc:\n            best_acc = vl_acc\n            if Config.USE_MSAM and isinstance(optimizer, MomentumSAM):\n                optimizer.remove_perturbation()\n                torch.save(model.state_dict(), best_path)\n                optimizer.perturb_parameters()\n            else:\n                torch.save(model.state_dict(), best_path)\n            print(f\"  ✓ Saved best model (val acc={best_acc:.2f}%)\")\n    if Config.USE_MSAM and isinstance(optimizer, MomentumSAM):\n        print(\"\\nFinalizing MSAM - removing parameter perturbation...\")\n        optimizer.finalize()\n    \n    print(f\"\\nLoading best weights (val acc={best_acc:.2f}%)...\")\n    model.load_state_dict(torch.load(best_path, map_location=device))\n    print(\"Running TTA inference...\")\n    preds, ids = predict_tta(model, test_df, Config.DATA_DIR, stats_lookup,\n                             device, plate_stats=plate_stats_test)\n    sirna_preds = [label_to_sirna[int(p)] for p in preds]\n    submission  = pd.DataFrame({'id_code': ids, 'sirna': sirna_preds})\n    submission.to_csv('submission_resnext50_32x4d_controls.csv', index=False)\n    print(f\"\\nDone. submission_resnext50_32x4d_controls.csv saved  |  Best val acc: {best_acc:.2f}%\")\n    print(submission.head())\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-29T15:34:02.671Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}