{"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":"markdown","source":"C","metadata":{}},{"cell_type":"code","source":"# =========================================================================================\n# AIRR-ML-25: Adaptive Immune Profiling Challenge - Professional Winning Solution\n# Model: Gated Attention-based Multiple Instance Learning (MIL) with Deep Sequence Encoding\n# =========================================================================================\n\n\"\"\"\nSOLUTION OVERVIEW:\nThis solution utilizes a Deep Learning approach tailored for Multiple Instance Learning (MIL).\nIn this context, a patient's Repertoire is a \"Bag\", and the immune sequences are \"Instances\".\n\n1.  **Sequence Encoder:** Each amino acid sequence (junction_aa) is embedded and processed \n    via a Bi-directional GRU to capture structural/functional context.\n2.  **Gated Attention Mechanism:** (Ilse et al., 2018) The network learns an 'attention weight' \n    for every sequence in a repertoire. High weights indicate sequences highly correlated \n    with the label (Disease).\n3.  **Aggregation:** Sequence representations are aggregated into a single 'Repertoire Vector'\n    using the learned attention weights.\n4.  **Classification:** A final classifier predicts the immune state (Healthy/Disease) based \n    on the Repertoire Vector.\n\nStrengths:\n-   **Task 1 (Prediction):** High accuracy by leveraging deep sequence features.\n-   **Task 2 (Interpretation):** The attention weights directly provide the \"importance score\" \n    required to rank the contributing sequences.\n\"\"\"\n# =========================================================================================\n# AIRR-ML-25: Professional Solution - FINAL VERSION (Auto-Path & Bug Fix)\n# =========================================================================================\n\nimport os\nimport sys\nimport numpy as np\nimport pandas as pd\nimport glob\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 sklearn.metrics import roc_auc_score\nimport random\nimport warnings\n\n# Suppress minor warnings for clean output\nwarnings.filterwarnings('ignore')\n\n# --- Reproducibility Setup ---\nSEED = 42\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(SEED)\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"✅ Using device: {DEVICE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T20:27:51.322246Z","iopub.execute_input":"2025-12-15T20:27:51.322652Z","iopub.status.idle":"2025-12-15T20:27:51.335101Z","shell.execute_reply.started":"2025-12-15T20:27:51.322622Z","shell.execute_reply":"2025-12-15T20:27:51.334079Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================================================\n# 1. ROBUST PATH CONFIGURATION\n# =========================================================================================\n\nBASE_DIR = \"/kaggle/input/adaptive-immune-profiling-challenge-2025\"\n\n# 1. Detect Train/Test Directories (Handles 'train' vs 'train_datasets' naming convention)\nif os.path.exists(os.path.join(BASE_DIR, \"train_datasets\")):\n    TRAIN_DIR = os.path.join(BASE_DIR, \"train_datasets\")\n    TEST_DIR = os.path.join(BASE_DIR, \"test_datasets\")\nelse:\n    TRAIN_DIR = os.path.join(BASE_DIR, \"train\")\n    TEST_DIR = os.path.join(BASE_DIR, \"test\")\n\nprint(f\"📂 Detected Repertoires Directory: {TRAIN_DIR}\")\n\n# 2. Detect Metadata File (Crucial Fix for FileNotFoundError)\nMETADATA_PATH = None\n# Check all common Kaggle path structures\npotential_paths = [\n    os.path.join(BASE_DIR, \"metadata.csv\"),\n    os.path.join(TRAIN_DIR, \"metadata.csv\"),\n]\nfor p in potential_paths:\n    if os.path.exists(p):\n        METADATA_PATH = p\n        break\n\nif METADATA_PATH is None:\n    # Fallback search if standard paths fail\n    found_metas = glob.glob(os.path.join(BASE_DIR, \"**\", \"metadata.csv\"), recursive=True)\n    if found_metas:\n        METADATA_PATH = found_metas[0]\n\nif METADATA_PATH is None:\n    # Raise error if still not found, but we catch it later\n    print(f\"❌ CRITICAL ERROR: metadata.csv NOT FOUND in {BASE_DIR} structure.\")\nelse:\n    print(f\"✅ Found Metadata Path: {METADATA_PATH}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T20:27:51.386687Z","iopub.execute_input":"2025-12-15T20:27:51.387046Z","iopub.status.idle":"2025-12-15T20:27:54.110891Z","shell.execute_reply.started":"2025-12-15T20:27:51.387015Z","shell.execute_reply":"2025-12-15T20:27:54.109519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================================================\n# 2. Data Processing & Helper Functions\n# =========================================================================================\n\n# Amino Acid Vocabulary for encoding\nAA_VOCAB = \"ACDEFGHIKLMNPQRSTVWY\"\nAA_TO_INT = {aa: i + 1 for i, aa in enumerate(AA_VOCAB)} # 0 reserved for padding\nVOCAB_SIZE = len(AA_VOCAB) + 1\nMAX_SEQ_LEN = 30 \n\ndef encode_sequence(seq, max_len=MAX_SEQ_LEN):\n    \"\"\"Encodes amino acid string to integer list.\"\"\"\n    if pd.isna(seq): return [0] * max_len\n    seq = seq[:max_len]\n    encoded = [AA_TO_INT.get(aa, 0) for aa in seq]\n    padding = [0] * (max_len - len(encoded))\n    return encoded + padding\n\nclass AIRRDataset(Dataset):\n    \"\"\"Dataset for MIL, handling variable-sized bags (Repertoires).\"\"\"\n    def __init__(self, repertoires_data, labels_map=None, is_train=True, max_seqs_per_bag=10000):\n        self.repertoires_data = repertoires_data\n        self.rep_ids = list(repertoires_data.keys())\n        self.labels_map = labels_map\n        self.is_train = is_train\n        self.max_seqs_per_bag = max_seqs_per_bag\n\n    def __len__(self):\n        return len(self.rep_ids)\n\n    def __getitem__(self, idx):\n        rep_id = self.rep_ids[idx]\n        df = self.repertoires_data[rep_id]\n        \n        # Subsampling for memory efficiency and regularization\n        if self.is_train and len(df) > self.max_seqs_per_bag:\n            df = df.sample(n=self.max_seqs_per_bag, random_state=SEED)\n            \n        sequences = [encode_sequence(seq) for seq in df['junction_aa'].values]\n        seq_tensor = torch.tensor(sequences, dtype=torch.long)\n        \n        # Get Label\n        label = torch.tensor(0.0, dtype=torch.float)\n        if self.is_train and self.labels_map:\n            val = self.labels_map.get(str(rep_id))\n            if val is not None:\n                label = torch.tensor(val, dtype=torch.float)\n                \n        # Return seqs, label, ID (as string), and raw sequences/V/J calls\n        return seq_tensor, label, str(rep_id), df[['junction_aa', 'v_call', 'j_call']].reset_index(drop=True)\n\ndef collate_bags(batch):\n    \"\"\"Custom collate function for DataLoader.\"\"\"\n    seqs, labels, rep_ids, raw_dfs = zip(*batch)\n    labels = torch.stack(labels)\n    return seqs, labels, rep_ids, raw_dfs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T20:27:54.112607Z","iopub.execute_input":"2025-12-15T20:27:54.112890Z","iopub.status.idle":"2025-12-15T20:27:54.123878Z","shell.execute_reply.started":"2025-12-15T20:27:54.112869Z","shell.execute_reply":"2025-12-15T20:27:54.123038Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================================================\n# 3. Model Architecture: Gated Attention MIL\n# =========================================================================================\n\nclass AttentionMILModel(nn.Module):\n    def __init__(self, vocab_size, embedding_dim=128, hidden_dim=256, mlp_dim=128):\n        super().__init__()\n        \n        # 1. Sequence Encoder (Instance Encoder)\n        self.embedding = nn.Embedding(vocab_size, embedding_dim, padding_idx=0)\n        self.seq_encoder = nn.GRU(embedding_dim, hidden_dim, batch_first=True, bidirectional=True)\n        seq_out_dim = hidden_dim * 2\n\n        # 2. Gated Attention Mechanism\n        self.attention_V = nn.Sequential(nn.Linear(seq_out_dim, mlp_dim), nn.Tanh())\n        self.attention_U = nn.Sequential(nn.Linear(seq_out_dim, mlp_dim), nn.Sigmoid())\n        self.attention_weights = nn.Linear(mlp_dim, 1)\n\n        # 3. Bag Classifier\n        self.classifier = nn.Sequential(\n            nn.Linear(seq_out_dim, mlp_dim),\n            nn.ReLU(),\n            nn.Dropout(0.25),\n            nn.Linear(mlp_dim, 1)\n        )\n\n    def forward(self, bag_seqs):\n        # Feature Extraction\n        embedded = self.embedding(bag_seqs) \n        _, hidden = self.seq_encoder(embedded)\n        hidden = torch.cat((hidden[-2], hidden[-1]), dim=1) # (N, Dim)\n\n        # Attention Scores\n        A = self.attention_weights(self.attention_V(hidden) * self.attention_U(hidden)) # (N, 1)\n        A = torch.softmax(torch.transpose(A, 1, 0), dim=1) # (1, N)\n\n        # Aggregation\n        bag_rep = torch.mm(A, hidden) # (1, Dim)\n\n        # Classification\n        return self.classifier(bag_rep).squeeze(1), A.squeeze(0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T20:27:54.125178Z","iopub.execute_input":"2025-12-15T20:27:54.125539Z","iopub.status.idle":"2025-12-15T20:27:54.151090Z","shell.execute_reply.started":"2025-12-15T20:27:54.125509Z","shell.execute_reply":"2025-12-15T20:27:54.150022Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================================================\n# 4. Predictor Engine (Handles Data Loading and Training)\n# =========================================================================================\n\nclass ImmuneStatePredictor:\n    def __init__(self):\n        self.device = DEVICE\n        self.model = None\n        self.train_metadata = {}\n\n    def _load_files(self, directory):\n        \"\"\"Loads all parquet files in a directory recursively.\"\"\"\n        reps = {}\n        meta_list = []\n        # Robust recursive search for parquets\n        files = glob.glob(os.path.join(directory, \"**\", \"*.parquet\"), recursive=True)\n        \n        for f in tqdm(files, desc=\"Loading Parquets\"):\n            try:\n                rep_id = os.path.basename(f).replace('.parquet', '')\n                dataset_id = os.path.basename(os.path.dirname(f))\n                \n                df = pd.read_parquet(f)\n                # Keep only essential columns to save memory\n                if 'junction_aa' in df.columns:\n                    reps[rep_id] = df[['junction_aa', 'v_call', 'j_call']]\n                    meta_list.append({'repertoire_id': rep_id, 'dataset_id': dataset_id})\n            except: pass\n        return reps, pd.DataFrame(meta_list)\n\n    def fit(self, train_dir, meta_path):\n        \"\"\"Loads data, initializes model, and starts training.\"\"\"\n        if not meta_path or not os.path.exists(meta_path):\n            raise FileNotFoundError(\"Metadata file missing. Cannot train.\")\n            \n        print(\"\\n--- Starting Training Process ---\")\n        train_reps, train_meta = self._load_files(train_dir)\n        if not train_reps: raise ValueError(\"No training repertoires found! Check TRAIN_DIR path.\")\n        \n        # Load Labels\n        labels_df = pd.read_csv(meta_path)\n        labels_df['repertoire_id'] = labels_df['repertoire_id'].astype(str)\n        labels_map = dict(zip(labels_df['repertoire_id'], labels_df['label']))\n        \n        self.train_metadata = {'repertoires': train_reps, 'meta': train_meta}\n        \n        ds = AIRRDataset(train_reps, labels_map, is_train=True)\n        # Use batch_size=1 for standard MIL training\n        loader = DataLoader(ds, batch_size=1, shuffle=True, collate_fn=collate_bags, num_workers=2)\n        \n        # Model and Optimizer setup\n        self.model = AttentionMILModel(VOCAB_SIZE).to(self.device)\n        opt = optim.AdamW(self.model.parameters(), lr=5e-4, weight_decay=1e-4)\n        crit = nn.BCEWithLogitsLoss()\n        \n        self.model.train()\n        EPOCHS = 8 \n        for epoch in range(EPOCHS):\n            total_loss = 0\n            for seqs, labels, _, _ in tqdm(loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\"):\n                # Since batch_size=1, we take the first element from the list\n                seqs, labels = seqs[0].to(self.device), labels[0].to(self.device)\n                \n                opt.zero_grad()\n                logits, _ = self.model(seqs)\n                loss = crit(logits, labels)\n                loss.backward()\n                opt.step()\n                total_loss += loss.item()\n            print(f\"Epoch {epoch+1} finished. Avg Loss: {total_loss/len(loader):.4f}\")\n        print(\"--- Training Completed ---\")\n\n    def predict(self, test_dir):\n        \"\"\"Performs inference on test data (Task 1).\"\"\"\n        print(\"\\n--- Phase 2: Predicting Test Set ---\")\n        test_reps, _ = self._load_files(test_dir)\n        loader = DataLoader(AIRRDataset(test_reps, is_train=False), batch_size=1, collate_fn=collate_bags, num_workers=2)\n        \n        preds = {}\n        self.model.eval()\n        with torch.no_grad():\n            for seqs, _, rep_ids, _ in tqdm(loader, desc=\"Inference\"):\n                logits, _ = self.model(seqs[0].to(self.device))\n                preds[rep_ids[0]] = torch.sigmoid(logits).item()\n        return preds\n\n    def interpret(self):\n        \"\"\"Extracts attention scores for sequence ranking (Task 2).\"\"\"\n        print(\"\\n--- Phase 3: Interpreting Sequences (Attention Scores) ---\")\n        loader = DataLoader(AIRRDataset(self.train_metadata['repertoires'], is_train=False), \n                            batch_size=1, collate_fn=collate_bags, num_workers=2)\n        scores = {} # {dataset_id: { (junc, v, j): max_score }}\n\n        self.model.eval()\n        with torch.no_grad():\n            for seqs, _, rep_ids, dfs in tqdm(loader, desc=\"Scanning Attention\"):\n                \n                # Get dataset_id for the current repertoire\n                ds_row = self.train_metadata['meta'][self.train_metadata['meta']['repertoire_id'] == rep_ids[0]]\n                if ds_row.empty: continue\n                ds_id = ds_row['dataset_id'].values[0]\n                if ds_id not in scores: scores[ds_id] = {}\n                \n                _, attn = self.model(seqs[0].to(self.device))\n                attn = attn.cpu().numpy()\n                df = dfs[0] # Raw dataframe for the bag\n                \n                # Store the max attention score for each unique sequence\n                for i, r in df.iterrows():\n                    key = (r['junction_aa'], r['v_call'], r['j_call'])\n                    if key not in scores[ds_id] or attn[i] > scores[ds_id][key]:\n                        scores[ds_id][key] = attn[i]\n        \n        # Rank top 50k per dataset\n        rows = []\n        print(\"Sorting and ranking top 50,000 sequences...\")\n        for ds, data in scores.items():\n            sorted_seqs = sorted(data.items(), key=lambda x: x[1], reverse=True)[:50000]\n            for rank, (k, s) in enumerate(sorted_seqs, 1):\n                rows.append({'dataset_id': ds, 'junction_aa': k[0], 'v_call': k[1], 'j_call': k[2], 'rank': rank})\n        return pd.DataFrame(rows)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T20:27:54.153070Z","iopub.execute_input":"2025-12-15T20:27:54.153360Z","iopub.status.idle":"2025-12-15T20:27:54.175316Z","shell.execute_reply.started":"2025-12-15T20:27:54.153341Z","shell.execute_reply":"2025-12-15T20:27:54.174366Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================================================\n# 5. Execution Pipeline (FINAL)\n# =========================================================================================\n\n# Initialize\npredictor = ImmuneStatePredictor()\n\ntry:\n    # A. Train\n    if METADATA_PATH:\n        predictor.fit(TRAIN_DIR, METADATA_PATH)\n        \n        # B. Task 1 Predictions\n        preds = predictor.predict(TEST_DIR)\n        df1 = pd.DataFrame(list(preds.items()), columns=['repertoire_id', 'probability'])\n        # Fill required columns with placeholders\n        for c in ['dataset_id', 'junction_aa', 'v_call', 'j_call', 'rank']: df1[c] = -999.0\n        \n        # C. Task 2 Interpretation\n        df2 = predictor.interpret()\n        df2['repertoire_id'] = \"dummy_id\" # Required placeholder\n        df2['probability'] = -999.0\n        df2 = df2[df1.columns] # Ensure column order matches df1\n        \n        # D. Submission\n        final = pd.concat([df1, df2], ignore_index=True)\n        final = final.astype({'repertoire_id': str, 'dataset_id': str, 'junction_aa': str, \n                              'v_call': str, 'j_call': str, 'probability': float, 'rank': float})\n        final.to_csv(\"submission.csv\", index=False)\n        \n        print(\"\\n--- Final Submission Summary ---\")\n        print(f\"✅ Success! Submission saved to submission.csv\")\n        print(f\"Final shape: {final.shape}\")\n        print(\"Head (Predictions):\")\n        print(final[final['repertoire_id'] != 'dummy_id'].head())\n        print(\"Head (Ranked Sequences):\")\n        print(final[final['repertoire_id'] == 'dummy_id'].head())\n        \n    else:\n        # Fallback if metadata was not found (should be caught by the check above)\n        print(\"❌ ERROR: Metadata file not found. Submission cannot be generated.\")\n\nexcept Exception as e:\n    # Catch any runtime errors gracefully without crashing the kernel\n    print(f\"\\n❌ CRITICAL EXECUTION FAILURE: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-15T20:27:54.176233Z","iopub.execute_input":"2025-12-15T20:27:54.176546Z","iopub.status.idle":"2025-12-15T20:27:54.242929Z","shell.execute_reply.started":"2025-12-15T20:27:54.176512Z","shell.execute_reply":"2025-12-15T20:27:54.241914Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"V\n","metadata":{}}]}