{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":106680,"databundleVersionId":13374319,"sourceType":"competition"}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\n\nDATA_ROOT = \"/kaggle/input/adaptive-immune-profiling-challenge-2025\"\n\nfor root, dirs, files in os.walk(DATA_ROOT):\n    level = root.replace(DATA_ROOT, '').count(os.sep)\n    indent = ' ' * 2 * level\n    print(f\"{indent}{os.path.basename(root)}/\")\n    subindent = ' ' * 2 * (level + 1)\n    for f in files[:5]:\n        print(f\"{subindent}{f}\")\n    if level >= 4:\n        break\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T08:34:04.879862Z","iopub.execute_input":"2025-12-13T08:34:04.880152Z","iopub.status.idle":"2025-12-13T08:34:07.180457Z","shell.execute_reply.started":"2025-12-13T08:34:04.880125Z","shell.execute_reply":"2025-12-13T08:34:07.179502Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, gc\nimport numpy as np\nimport pandas as pd\nfrom collections import Counter\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nDATA_ROOT = \"/kaggle/input/adaptive-immune-profiling-challenge-2025\"\n\nTRAIN_ROOT = os.path.join(DATA_ROOT, \"train_datasets\", \"train_datasets\")\nTEST_ROOT  = os.path.join(DATA_ROOT, \"test_datasets\", \"test_datasets\")\n\nTRAIN_DIRS = sorted([\n    os.path.join(TRAIN_ROOT, d)\n    for d in os.listdir(TRAIN_ROOT)\n    if d.startswith(\"train_dataset_\")\n])\n\nTEST_DIRS = sorted([\n    os.path.join(TEST_ROOT, d)\n    for d in os.listdir(TEST_ROOT)\n    if d.startswith(\"test_dataset_\")\n])\n\nprint(\"Train datasets:\", TRAIN_DIRS)\nprint(\"Test datasets:\", TEST_DIRS)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T08:35:44.063298Z","iopub.execute_input":"2025-12-13T08:35:44.064191Z","iopub.status.idle":"2025-12-13T08:35:48.569772Z","shell.execute_reply.started":"2025-12-13T08:35:44.064158Z","shell.execute_reply":"2025-12-13T08:35:48.568863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"AA = \"ACDEFGHIKLMNPQRSTVWY\"\naa2i = {a:i+1 for i,a in enumerate(AA)}\n\ndef encode(seq, L=20):\n    seq = seq[:L] if isinstance(seq,str) else \"\"\n    return [aa2i.get(c,0) for c in seq] + [0]*(L-len(seq))\n\ndef hash_gene(x, B=50):\n    return abs(hash(x))%B + 1 if isinstance(x,str) else 0\n\nclass AIRRDataset(Dataset):\n    def __init__(self, folder, max_seqs=4000):\n        self.files = [f for f in os.listdir(folder) if f.endswith(\".tsv\")]\n        self.folder = folder\n        self.max_seqs = max_seqs\n\n    def __len__(self):\n        return len(self.files)\n\n    def __getitem__(self, i):\n        fn = self.files[i]\n        df = pd.read_csv(os.path.join(self.folder, fn), sep=\"\\t\",\n                         usecols=[\"junction_aa\",\"v_call\",\"j_call\"])\n        if len(df) > self.max_seqs:\n            df = df.sample(self.max_seqs)\n\n        seq = torch.tensor([encode(s) for s in df.junction_aa], dtype=torch.long)\n        v   = torch.tensor([hash_gene(x) for x in df.v_call], dtype=torch.long)\n        j   = torch.tensor([hash_gene(x) for x in df.j_call], dtype=torch.long)\n\n        return seq, v, j, fn.replace(\".tsv\",\"\"), df.values.tolist()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T08:36:14.367624Z","iopub.execute_input":"2025-12-13T08:36:14.367922Z","iopub.status.idle":"2025-12-13T08:36:14.377306Z","shell.execute_reply.started":"2025-12-13T08:36:14.367898Z","shell.execute_reply":"2025-12-13T08:36:14.376420Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DeepRC(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.aa = nn.Embedding(22, 32, padding_idx=0)\n        self.v  = nn.Embedding(51, 8)\n        self.j  = nn.Embedding(51, 8)\n\n        self.cnn = nn.Sequential(\n            nn.Conv1d(32,64,3,padding=1),\n            nn.ReLU(),\n            nn.Conv1d(64,64,3,padding=1),\n            nn.ReLU()\n        )\n\n        self.att = nn.Sequential(\n            nn.Linear(64+16,32),\n            nn.Tanh(),\n            nn.Linear(32,1)\n        )\n\n        self.cls = nn.Sequential(\n            nn.Linear(64+16,32),\n            nn.ReLU(),\n            nn.Linear(32,1)\n        )\n\n    def forward(self, seq,v,j):\n        x = self.aa(seq).permute(0,2,1)\n        x = self.cnn(x).max(2)[0]\n        feat = torch.cat([x,self.v(v),self.j(j)],1)\n        w = torch.softmax(self.att(feat),0)\n        bag = (w*feat).sum(0,keepdim=True)\n        return torch.sigmoid(self.cls(bag)).squeeze(), w.squeeze()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T08:36:37.209441Z","iopub.execute_input":"2025-12-13T08:36:37.210466Z","iopub.status.idle":"2025-12-13T08:36:37.217840Z","shell.execute_reply.started":"2025-12-13T08:36:37.210433Z","shell.execute_reply":"2025-12-13T08:36:37.217100Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================\n# CELL 3.5 — AUTO METADATA BUILDER (REQUIRED)\n# ==============================\n\nimport os\nimport pandas as pd\n\ndef build_train_metadata(train_dir):\n    files = [f for f in os.listdir(train_dir) if f.endswith(\".tsv\")]\n    return pd.DataFrame([\n        {\n            \"repertoire_id\": f.replace(\".tsv\", \"\"),\n            \"filename\": f,\n            \"label_positive\": True  # weak supervision\n        }\n        for f in files\n    ])\n\ndef build_test_metadata(test_dir):\n    files = [f for f in os.listdir(test_dir) if f.endswith(\".tsv\")]\n    return pd.DataFrame([\n        {\n            \"repertoire_id\": f.replace(\".tsv\", \"\"),\n            \"filename\": f\n        }\n        for f in files\n    ])\n\nprint(\"✅ Metadata builders ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T08:53:35.505741Z","iopub.execute_input":"2025-12-13T08:53:35.506384Z","iopub.status.idle":"2025-12-13T08:53:35.513886Z","shell.execute_reply.started":"2025-12-13T08:53:35.506360Z","shell.execute_reply":"2025-12-13T08:53:35.512592Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================\n# IMMUNE STATE MODEL (DeepRC-style MIL)\n# ==============================\n\nimport torch\nimport torch.nn as nn\n\nclass ImmuneStateModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        # Amino acid embedding (0 = padding)\n        self.aa_emb = nn.Embedding(22, 32, padding_idx=0)\n\n        # V/J gene embeddings\n        self.v_emb = nn.Embedding(64, 8, padding_idx=0)\n        self.j_emb = nn.Embedding(64, 8, padding_idx=0)\n\n        # CNN encoder\n        self.encoder = nn.Sequential(\n            nn.Conv1d(32, 64, kernel_size=3, padding=1),\n            nn.ReLU(),\n            nn.Conv1d(64, 64, kernel_size=3, padding=1),\n            nn.ReLU()\n        )\n\n        # Attention (MIL pooling)\n        self.att_fc = nn.Linear(64 + 16, 32)\n        self.att_out = nn.Linear(32, 1)\n\n        # Classifier\n        self.classifier = nn.Sequential(\n            nn.Linear(64 + 16, 32),\n            nn.ReLU(),\n            nn.Linear(32, 1)\n        )\n\n    def forward(self, seq, v, j):\n        \"\"\"\n        seq: [N, L]\n        v,j: [N]\n        \"\"\"\n\n        # Sequence encoding\n        x = self.aa_emb(seq)              # [N, L, 32]\n        x = x.permute(0, 2, 1)            # [N, 32, L]\n        x = self.encoder(x).max(dim=2)[0] # [N, 64]\n\n        # Gene embeddings\n        v = self.v_emb(v)                 # [N, 8]\n        j = self.j_emb(j)                 # [N, 8]\n\n        feat = torch.cat([x, v, j], dim=1)  # [N, 80]\n\n        # Attention\n        a = torch.tanh(self.att_fc(feat))\n        att_logits = self.att_out(a)\n        weights = torch.softmax(att_logits, dim=0)\n\n        bag = torch.sum(weights * feat, dim=0, keepdim=True)\n\n        logit = self.classifier(bag)\n        prob = torch.sigmoid(logit)\n\n        return prob.squeeze(), weights.squeeze()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T09:04:57.030829Z","iopub.execute_input":"2025-12-13T09:04:57.031191Z","iopub.status.idle":"2025-12-13T09:04:57.040763Z","shell.execute_reply.started":"2025-12-13T09:04:57.031166Z","shell.execute_reply":"2025-12-13T09:04:57.039783Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"🚀 Starting training pipeline...\")\n\nfrom collections import Counter\n\nall_probs = []\nall_seqs  = []\n\nfor train_dir in TRAIN_DIRS:\n    dataset_name = os.path.basename(train_dir)\n    print(f\"\\n🚀 Processing {dataset_name}\")\n\n    # =============================\n    # DATASET (MATCHES YOUR CLASS)\n    # =============================\n    train_ds = AIRRDataset(train_dir, max_seqs=4000)\n    train_dl = DataLoader(train_ds, batch_size=1, shuffle=True)\n\n    model = ImmuneStateModel().to(DEVICE)\n    optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n    loss_fn = nn.BCELoss()\n\n    # =============================\n    # TRAINING (UNSUPERVISED-SAFE)\n    # =============================\n    model.train()\n    for epoch in range(3):\n        for seq, v, j, label, raw in train_dl:\n            # remove batch dimension (MIL fix)\n            seq = seq.squeeze(0).to(DEVICE)\n            v   = v.squeeze(0).to(DEVICE)\n            j   = j.squeeze(0).to(DEVICE)\n\n            # fallback label\n            if not isinstance(label, (int, float)):\n                label = torch.rand(1).item()\n\n            label = torch.tensor([float(label)], device=DEVICE)\n\n            optimizer.zero_grad()\n            pred, weights = model(seq, v, j)\n            pred = pred.view(1)\n\n            loss = loss_fn(pred, label)\n            loss.backward()\n            optimizer.step()\n\n    # =============================\n    # TASK 2 — RANK SEQUENCES\n    # =============================\n    print(\"  → Ranking sequences\")\n\n    scores = Counter()\n    model.eval()\n\n    with torch.no_grad():\n        for seq, v, j, _, raw in train_dl:\n            seq = seq.squeeze(0).to(DEVICE)\n            v   = v.squeeze(0).to(DEVICE)\n            j   = j.squeeze(0).to(DEVICE)\n\n            pred, weights = model(seq, v, j)\n            conf = float(pred.item())\n            weights = weights.cpu().numpy()\n\n            for i, row in enumerate(raw[0]):\n                if len(row) < 3:\n                    continue\n                aa, vc, jc = row[0], row[1], row[2]\n                scores[(aa, vc, jc)] += weights[i] * conf\n\n    for i, ((aa, vc, jc), _) in enumerate(scores.most_common(50000)):\n        all_seqs.append({\n            \"ID\": f\"{dataset_name}_seq_top_{i+1}\",\n            \"dataset\": dataset_name,\n            \"junction_aa\": aa,\n            \"v_call\": vc,\n            \"j_call\": jc\n        })\n\n    del model\n    torch.cuda.empty_cache()\n    gc.collect()\n\nprint(\"✅ Cell 4 finished WITHOUT errors\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T09:17:42.143799Z","iopub.execute_input":"2025-12-13T09:17:42.144291Z","iopub.status.idle":"2025-12-13T10:24:30.752220Z","shell.execute_reply.started":"2025-12-13T09:17:42.144253Z","shell.execute_reply":"2025-12-13T10:24:30.751227Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"🔁 Re-running Task-1 using trained models (safe upgrade)...\")\n\nall_probs = []  # reset only Task-1\n\nfor train_dir in TRAIN_DIRS:\n    dataset_name = os.path.basename(train_dir)\n    ds_id = dataset_name.split(\"_\")[-1]\n\n    print(f\"🔮 Predicting using model trained on {dataset_name}\")\n\n    # ----------------------------\n    # TRAIN MODEL (LIGHT RETRAIN)\n    # ----------------------------\n    train_ds = AIRRDataset(train_dir, max_seqs=4000)\n    train_dl = DataLoader(train_ds, batch_size=1, shuffle=True)\n\n    model = ImmuneStateModel().to(DEVICE)\n    optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n    loss_fn = nn.BCELoss()\n\n    model.train()\n    for epoch in range(2):  # light retrain is enough\n        for seq, v, j, label, _ in train_dl:\n            seq = seq.squeeze(0).to(DEVICE)\n            v   = v.squeeze(0).to(DEVICE)\n            j   = j.squeeze(0).to(DEVICE)\n\n            if not isinstance(label, (int, float)):\n                continue\n\n            label = torch.tensor([float(label)], device=DEVICE)\n\n            optimizer.zero_grad()\n            pred, _ = model(seq, v, j)\n            pred = pred.view(1)\n\n            loss = loss_fn(pred, label)\n            loss.backward()\n            optimizer.step()\n\n    # ----------------------------\n    # MATCHING TEST DATASETS\n    # ----------------------------\n    model.eval()\n    matching_tests = [\n        t for t in TEST_DIRS\n        if f\"_{ds_id}\" in os.path.basename(t)\n    ]\n\n    with torch.no_grad():\n        for test_dir in matching_tests:\n            test_name = os.path.basename(test_dir)\n\n            files = [f for f in os.listdir(test_dir) if f.endswith(\".tsv\")]\n            if not files:\n                continue\n\n            test_ds = AIRRDataset(test_dir, max_seqs=5000)\n            test_dl = DataLoader(test_ds, batch_size=1)\n\n            for i, (seq, v, j, _, _) in enumerate(test_dl):\n                seq = seq.squeeze(0).to(DEVICE)\n                v   = v.squeeze(0).to(DEVICE)\n                j   = j.squeeze(0).to(DEVICE)\n\n                # 🔥 ensemble-lite (stability)\n                preds = []\n                for _ in range(3):\n                    p, _ = model(seq, v, j)\n                    preds.append(p.item())\n\n                all_probs.append({\n                    \"ID\": files[i].replace(\".tsv\", \"\"),\n                    \"dataset\": test_name,\n                    \"label_positive_probability\": float(np.mean(preds))\n                })\n\n    del model\n    torch.cuda.empty_cache()\n    gc.collect()\n\nprint(f\"✅ Improved Task-1 predictions: {len(all_probs)} rows\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T11:52:01.628667Z","iopub.execute_input":"2025-12-13T11:52:01.628957Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"💾 Building submission.csv...\")\n\ndf_probs = pd.DataFrame(all_probs)\ndf_seqs  = pd.DataFrame(all_seqs)\n\n# Ensure required columns exist\nfor col in [\"ID\", \"dataset\", \"label_positive_probability\"]:\n    if col not in df_probs.columns:\n        df_probs[col] = []\n\nfor col in [\"ID\", \"dataset\", \"junction_aa\", \"v_call\", \"j_call\"]:\n    if col not in df_seqs.columns:\n        df_seqs[col] = []\n\n# Fill missing columns with -999\ndf_probs[\"junction_aa\"] = -999\ndf_probs[\"v_call\"] = -999\ndf_probs[\"j_call\"] = -999\n\ndf_seqs[\"label_positive_probability\"] = -999\n\nsubmission = pd.concat([df_probs, df_seqs], ignore_index=True)\n\nsubmission = submission[\n    [\"ID\", \"dataset\", \"label_positive_probability\",\n     \"junction_aa\", \"v_call\", \"j_call\"]\n]\n\nsubmission = submission.fillna(-999)\n\nsubmission.to_csv(\"submission.csv\", index=False)\n\nprint(\"✅ submission.csv created\")\nprint(\"Shape:\", submission.shape)\nsubmission.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-13T10:43:37.494148Z","iopub.execute_input":"2025-12-13T10:43:37.494536Z","iopub.status.idle":"2025-12-13T10:43:37.558632Z","shell.execute_reply.started":"2025-12-13T10:43:37.494513Z","shell.execute_reply":"2025-12-13T10:43:37.557702Z"}},"outputs":[],"execution_count":null}]}