{"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":"QQ","metadata":{}},{"cell_type":"code","source":"# =========================================================================================\n# AIRR-ML-25: Professional Solution - Production V.13 (Cell 1/4: Setup & Preprocessing)\n# =========================================================================================\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport glob\nfrom tqdm.auto import tqdm \nfrom tqdm.contrib.concurrent import process_map \nfrom itertools import repeat \n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom sklearn.model_selection import train_test_split\nimport random\nimport warnings\n\n# --- Reproducibility & Environment 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    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n\nseed_everything(SEED)\nwarnings.filterwarnings('ignore')\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# --- Constants & Speed Optimization Configuration ---\nBASE_DIR = \"/kaggle/input/adaptive-immune-profiling-challenge-2025\"\nTRAIN_DIR = os.path.join(BASE_DIR, \"train_datasets\")\nTEST_DIR = os.path.join(BASE_DIR, \"test_datasets\")\n\nNUM_FILE_WORKERS = os.cpu_count() * 2 if os.cpu_count() else 4\nDL_WORKERS = 0 # 🛑 FIX: Use 0 workers for DataLoader to prevent Deadlock/Freezing\n\n# Sequence/Amino Acid Constants\nAA_VOCAB = \"ACDEFGHIKLMNPQRSTVWY\"\nAA_TO_INT = {aa: i + 1 for i, aa in enumerate(AA_VOCAB)} \nVOCAB_SIZE = len(AA_VOCAB) + 1\nMAX_SEQ_LEN = 30 \n# ⚡️ OPTIMIZATION: Reduce max sequences per bag for faster training batches\nMAX_SEQS_PER_BAG = 3000 \n\n# Dynamic Gene Call Encoding Maps (Global, shared state during load)\nV_CALLS_MAP = {}\nJ_CALLS_MAP = {}\nV_VOCAB_SIZE = 1 \nJ_VOCAB_SIZE = 1 \n\n# Metadata Detection\nMETADATA_PATH = None\nfound_metas = glob.glob(os.path.join(BASE_DIR, \"**\", \"metadata.csv\"), recursive=True)\nif found_metas:\n    METADATA_PATH = found_metas[0] \n\n# --- Preprocessing Helper Functions (No Change) ---\n\ndef get_gene_id(gene_call, gene_map, is_v_call):\n    \"\"\"Maps gene call strings to integer IDs dynamically (Used for sequential TRAIN encoding).\"\"\"\n    global V_VOCAB_SIZE, J_VOCAB_SIZE\n    if pd.isna(gene_call) or gene_call == \"\": return 0 \n    \n    gene = gene_call.split('*')[0].split(',')[0].strip() \n    \n    if gene not in gene_map:\n        if is_v_call:\n            gene_map[gene] = V_VOCAB_SIZE\n            V_VOCAB_SIZE += 1\n            return V_VOCAB_SIZE - 1\n        else:\n            gene_map[gene] = J_VOCAB_SIZE\n            J_VOCAB_SIZE += 1\n            return V_VOCAB_SIZE - 1\n    return gene_map[gene]\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\ndef _process_single_file_global(f, is_train_dir, v_map_train, j_map_train):\n    \"\"\"Helper function for parallel loading of a single TSV file.\"\"\"\n    try:\n        rep_id = os.path.basename(f).replace('.tsv', '')\n        # ⚡️ SPEED: Use dtype={'junction_aa': str} to minimize pandas overhead/inference\n        df = pd.read_csv(f, sep='\\t', usecols=['junction_aa', 'v_call', 'j_call'], \n                         dtype={'junction_aa': str, 'v_call': str, 'j_call': str}) \n        \n        df = df[['junction_aa', 'v_call', 'j_call']].dropna(subset=['junction_aa'])\n        \n        if not is_train_dir:\n            # Use fixed maps from training (Test/Inference)\n            v_call_mapper = lambda x: v_map_train.get(x.split('*')[0].split(',')[0].strip(), 0) if pd.notna(x) else 0\n            j_call_mapper = lambda x: j_map_train.get(x.split('*')[0].split(',')[0].strip(), 0) if pd.notna(x) else 0\n            df['v_call_id'] = df['v_call'].apply(v_call_mapper)\n            df['j_call_id'] = df['j_call'].apply(j_call_mapper)\n            return rep_id, df\n        \n        return rep_id, df \n        \n    except Exception as e:\n        print(f\"Error processing file {f}: {e}\")\n        return None, None\n\nprint(f\"✅ Using device: {DEVICE}\")\nprint(f\"⚙️ Using {DL_WORKERS} workers for DataLoader (0 for stability).\")\nprint(f\"✅ Found Metadata Path: {METADATA_PATH if METADATA_PATH else '❌ NOT FOUND'}\")\n# =========================================================================================","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-16T21:00:41.126128Z","iopub.execute_input":"2025-12-16T21:00:41.127416Z","iopub.status.idle":"2025-12-16T21:00:50.091895Z","shell.execute_reply.started":"2025-12-16T21:00:41.12736Z","shell.execute_reply":"2025-12-16T21:00:50.091216Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================================================\n# AIRR-ML-25: Professional Solution - Production V.13 (Cell 2/4: Dataset & Model)\n# =========================================================================================\n\nclass AIRRDataset(Dataset):\n    \"\"\"Dataset for MIL, handling variable-sized bags (Repertoires).\"\"\"\n    def __init__(self, rep_ids, repertoires_data, labels_map=None, is_train=True):\n        self.rep_ids = rep_ids\n        self.repertoires_data = repertoires_data\n        self.labels_map = labels_map\n        self.is_train = is_train\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        if self.is_train and len(df) > MAX_SEQS_PER_BAG:\n            # Random sampling is crucial for MIL robustness\n            df = df.sample(n=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        v_tensor = torch.tensor(df['v_call_id'].values, dtype=torch.long)\n        j_tensor = torch.tensor(df['j_call_id'].values, dtype=torch.long)\n        \n        label = torch.tensor(0.0, dtype=torch.float)\n        if 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        raw_df = df[['junction_aa', 'v_call', 'j_call']].reset_index(drop=True)\n                \n        return seq_tensor, v_tensor, j_tensor, label, str(rep_id), raw_df\n\ndef collate_bags_optimized(batch):\n    \"\"\"\n    ⚡️ OPTIMIZATION: Custom collate function to batch multiple bags (repertoires) \n    by flattening all instances into one large tensor, suitable for MIL.\n    \"\"\"\n    all_seqs = []\n    all_v_calls = []\n    all_j_calls = []\n    bag_labels = []\n    bag_rep_ids = []\n    bag_raw_dfs = []\n    \n    # Track the start/end indices for each bag\n    bag_indices = []\n    current_idx = 0\n    \n    for seqs, v_calls, j_calls, label, rep_id, raw_df in batch:\n        num_instances = len(seqs)\n        \n        all_seqs.append(seqs)\n        all_v_calls.append(v_calls)\n        all_j_calls.append(j_calls)\n        \n        bag_labels.append(label)\n        bag_rep_ids.append(rep_id)\n        bag_raw_dfs.append(raw_df)\n\n        bag_indices.append((current_idx, current_idx + num_instances))\n        current_idx += num_instances\n\n    # Flatten the lists of tensors into one large tensor for the whole batch\n    seqs_batch = torch.cat(all_seqs, dim=0)\n    v_calls_batch = torch.cat(all_v_calls, dim=0)\n    j_calls_batch = torch.cat(all_j_calls, dim=0)\n    labels_batch = torch.stack(bag_labels)\n    \n    return seqs_batch, v_calls_batch, j_calls_batch, labels_batch, bag_rep_ids, bag_raw_dfs, bag_indices\n\n\nclass AttentionMILModel(nn.Module):\n    \"\"\"\n    Multiple Instance Learning model with Gated Attention, integrating\n    CDR3 sequence features (Transformer) and V/J gene call features (Embeddings).\n    \"\"\"\n    def __init__(self, vocab_size, v_vocab_size, j_vocab_size, \n                 embedding_dim=32, hidden_dim=32, mlp_dim=64): # ⚡️ OPTIMIZATION: Reduced dimensions (32/64)\n        super().__init__()\n        \n        # 1. Sequence Encoder (Transformer Instance Encoder)\n        self.seq_embedding = nn.Embedding(vocab_size, embedding_dim, padding_idx=0)\n        \n        transformer_layer = nn.TransformerEncoderLayer(\n            d_model=embedding_dim, nhead=4, dim_feedforward=hidden_dim * 2, dropout=0.1, batch_first=True\n        ) # Reduced nhead\n        self.seq_encoder = nn.TransformerEncoder(transformer_layer, num_layers=2)\n        \n        seq_out_dim = embedding_dim\n        \n        # 2. V/J Embeddings (External Features)\n        self.v_embedding = nn.Embedding(v_vocab_size, embedding_dim // 2, padding_idx=0)\n        self.j_embedding = nn.Embedding(j_vocab_size, embedding_dim // 2, padding_idx=0)\n        vj_out_dim = embedding_dim\n        \n        total_feature_dim = seq_out_dim + vj_out_dim \n\n        # 3. Gated Attention Mechanism \n        self.attention_V = nn.Sequential(nn.Linear(total_feature_dim, mlp_dim), nn.Tanh())\n        self.attention_U = nn.Sequential(nn.Linear(total_feature_dim, mlp_dim), nn.Sigmoid())\n        self.attention_weights = nn.Linear(mlp_dim, 1)\n\n        # 4. Bag Classifier\n        self.classifier = nn.Sequential(\n            nn.Linear(total_feature_dim, mlp_dim),\n            nn.ReLU(),\n            nn.Dropout(0.25),\n            nn.Linear(mlp_dim, 1)\n        )\n        \n    def forward(self, seqs_batch, v_calls_batch, j_calls_batch, bag_indices=None):\n        \n        # 1. Sequence Features (CDR3) via Transformer\n        embedded_seq = self.seq_embedding(seqs_batch) \n        encoded_seq = self.seq_encoder(embedded_seq)\n        instance_seq_features = encoded_seq.mean(dim=1) \n\n        # 2. V/J Features\n        embedded_v = self.v_embedding(v_calls_batch) \n        embedded_j = self.j_embedding(j_calls_batch) \n        instance_vj_features = torch.cat((embedded_v, embedded_j), dim=1)\n        \n        # 3. Concatenate all features (Multimodal Fusion)\n        instance_features = torch.cat((instance_seq_features, instance_vj_features), dim=1)\n        \n        # If running in batch_size > 1 mode, we must process attention per bag\n        if bag_indices and len(bag_indices) > 1:\n            bag_reps = []\n            all_attns = []\n            \n            for start, end in bag_indices:\n                features = instance_features[start:end]\n                \n                # Attention Scores (4)\n                A = self.attention_weights(self.attention_V(features) * self.attention_U(features)) \n                A = torch.softmax(torch.transpose(A, 1, 0), dim=1) \n                \n                # Aggregation (5)\n                bag_rep = torch.mm(A, features) \n                bag_reps.append(bag_rep)\n                all_attns.append(A.squeeze(0))\n                \n            bag_rep_batch = torch.cat(bag_reps, dim=0)\n            logits = self.classifier(bag_rep_batch).squeeze(1)\n            # Return attention only for interpretation/batch_size=1\n            return logits, all_attns \n\n        else: # Standard batch_size=1 processing (or single bag batch)\n            # 4. Attention Scores \n            A = self.attention_weights(self.attention_V(instance_features) * self.attention_U(instance_features)) \n            A = torch.softmax(torch.transpose(A, 1, 0), dim=1) \n\n            # 5. Aggregation (Weighted Average)\n            bag_rep = torch.mm(A, instance_features) \n            \n            # 6. Classification\n            return self.classifier(bag_rep).squeeze(1), A.squeeze(0)\n\n# =========================================================================================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T21:00:50.093013Z","iopub.execute_input":"2025-12-16T21:00:50.093355Z","iopub.status.idle":"2025-12-16T21:00:50.112167Z","shell.execute_reply.started":"2025-12-16T21:00:50.093339Z","shell.execute_reply":"2025-12-16T21:00:50.111166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================================================\n# AIRR-ML-25: Professional Solution - Production V.13 (Cell 3/4: Predictor Engine Class)\n# =========================================================================================\n\nclass ImmuneStatePredictor:\n    def __init__(self):\n        self.device = DEVICE \n        self.model = None\n        self.train_data = {}\n        self.v_map = V_CALLS_MAP.copy() \n        self.j_map = J_CALLS_MAP.copy()\n        self.v_vocab_size = V_VOCAB_SIZE\n        self.j_vocab_size = J_VOCAB_SIZE\n        self.best_val_loss = float('inf')\n        self.best_model_weights = None\n        \n        self.dataloader_workers = DL_WORKERS # 0\n        self.file_workers = NUM_FILE_WORKERS\n        self.BATCH_SIZE = 4 # ⚡️ OPTIMIZATION: Further reduce batch size for low RAM/CPU stability\n        \n        print(f\"⚙️ Predictor Initialized. DataLoader Workers: {self.dataloader_workers}, Batch Size: {self.BATCH_SIZE}\")\n\n\n    def _load_files(self, directory, is_train_dir):\n        \"\"\"Loads all repertoire data files (.tsv) using parallel processing (process_map).\"\"\"\n        files = glob.glob(os.path.join(directory, \"**\", \"*.tsv\"), recursive=True)\n        print(f\"🔍 Found {len(files)} repertoire files (.tsv) via deep scan in {directory}\")\n\n        results = process_map(\n            _process_single_file_global,\n            files,\n            repeat(is_train_dir),\n            repeat(self.v_map),\n            repeat(self.j_map),\n            max_workers=self.file_workers, \n            chunksize=8,\n            desc=\"Loading TSV Files\"\n        )\n        \n        reps = {}\n        for rep_id, df in results:\n            if rep_id:\n                reps[rep_id] = df\n\n        if is_train_dir:\n            global V_CALLS_MAP, J_CALLS_MAP, V_VOCAB_SIZE, J_VOCAB_SIZE\n            V_CALLS_MAP = {}\n            J_CALLS_MAP = {}\n            V_VOCAB_SIZE = 1 \n            J_VOCAB_SIZE = 1 \n            \n            print(\"Encoding V/J genes sequentially after parallel loading...\")\n            for rep_id, df in tqdm(reps.items(), desc=\"Sequential Encoding\"):\n                df['v_call_id'] = df['v_call'].apply(lambda x: get_gene_id(x, V_CALLS_MAP, True))\n                df['j_call_id'] = df['j_call'].apply(lambda x: get_gene_id(x, J_CALLS_MAP, False))\n            \n            self.v_map = V_CALLS_MAP.copy()\n            self.j_map = J_CALLS_MAP.copy()\n            self.v_vocab_size = V_VOCAB_SIZE\n            self.j_vocab_size = J_VOCAB_SIZE\n            \n        return reps, pd.DataFrame()\n\n\n    def fit(self, train_dir, meta_path):\n        \"\"\"Loads data, initializes model, and starts training with Validation and Early Stopping.\"\"\"\n        print(\"\\n--- Starting Training Process ---\")\n        train_reps, _ = self._load_files(train_dir, is_train_dir=True)\n        if not train_reps: raise ValueError(\"No training repertoires found.\")\n        \n        print(f\"📊 V/J Vocab Size: V={self.v_vocab_size}, J={self.j_vocab_size}\")\n        \n        labels_df = pd.read_csv(meta_path)\n        labels_map = dict(zip(labels_df['repertoire_id'].astype(str), labels_df['label_positive']))\n        self.train_data['repertoires'] = train_reps\n        all_rep_ids = list(train_reps.keys())\n        train_ids, val_ids = train_test_split(all_rep_ids, test_size=0.2, random_state=SEED)\n        \n        train_ds = AIRRDataset(train_ids, train_reps, labels_map, is_train=True)\n        val_ds = AIRRDataset(val_ids, train_reps, labels_map, is_train=False) \n        \n        # num_workers=0 and BATCH_SIZE=4\n        train_loader = DataLoader(train_ds, batch_size=self.BATCH_SIZE, shuffle=True, collate_fn=collate_bags_optimized, \n                                  num_workers=self.dataloader_workers, pin_memory=True)\n        val_loader = DataLoader(val_ds, batch_size=self.BATCH_SIZE, shuffle=False, collate_fn=collate_bags_optimized, \n                                num_workers=self.dataloader_workers, pin_memory=True)\n        \n        self.model = AttentionMILModel(VOCAB_SIZE, self.v_vocab_size, self.j_vocab_size).to(self.device)\n        opt = optim.AdamW(self.model.parameters(), lr=5e-5, weight_decay=1e-4)\n        scheduler = CosineAnnealingLR(opt, T_max=20, eta_min=1e-7) \n        crit = nn.BCEWithLogitsLoss()\n        \n        PATIENCE = 5\n        epochs_no_improve = 0\n        EPOCHS = 30 \n\n        for epoch in range(EPOCHS):\n            total_loss = 0\n            self.model.train()\n            for seqs_t, v_calls_t, j_calls_t, labels_t, _, _, bag_indices in tqdm(train_loader, desc=f\"Train Epoch {epoch+1}/{EPOCHS}\"):\n                \n                # Data Transfer\n                seqs_t = seqs_t.to(self.device)\n                v_calls_t = v_calls_t.to(self.device)\n                j_calls_t = j_calls_t.to(self.device)\n                labels_t = labels_t.to(self.device).view(-1) \n                \n                opt.zero_grad()\n                \n                # Pass bag_indices for multi-bag batching in forward pass\n                logits, _ = self.model(seqs_t, v_calls_t, j_calls_t, bag_indices=bag_indices)\n                loss = crit(logits, labels_t)\n                loss.backward()\n                opt.step()\n                total_loss += loss.item()\n            \n            scheduler.step()\n\n            # Validation Loop\n            val_loss = 0\n            self.model.eval()\n            with torch.no_grad():\n                for seqs_t, v_calls_t, j_calls_t, labels_t, _, _, bag_indices in val_loader:\n                    seqs_t = seqs_t.to(self.device)\n                    v_calls_t = v_calls_t.to(self.device)\n                    j_calls_t = j_calls_t.to(self.device)\n                    labels_t = labels_t.to(self.device).view(-1) \n                    \n                    logits, _ = self.model(seqs_t, v_calls_t, j_calls_t, bag_indices=bag_indices)\n                    val_loss += crit(logits, labels_t).item()\n            \n            avg_val_loss = val_loss / len(val_loader)\n            print(f\"Epoch {epoch+1} finished. Train Loss: {total_loss/len(train_loader):.4f} | Val Loss: {avg_val_loss:.4f} | LR: {scheduler.get_last_lr()[0]:.6f}\")\n\n            # Early Stopping Check\n            if avg_val_loss < self.best_val_loss:\n                self.best_val_loss = avg_val_loss\n                self.best_model_weights = self.model.state_dict()\n                epochs_no_improve = 0\n            else:\n                epochs_no_improve += 1\n                if epochs_no_improve == PATIENCE:\n                    print(f\"🛑 Early stopping triggered after {epoch+1} epochs.\")\n                    break\n        \n        if self.best_model_weights:\n            self.model.load_state_dict(self.best_model_weights)\n            print(\"✅ Loaded best model weights.\")\n\n        print(\"--- Training Completed ---\")\n\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, is_train_dir=False) \n        \n        # Use batch size 1 for inference/interpretation (simpler logic, fewer errors)\n        loader = DataLoader(AIRRDataset(list(test_reps.keys()), test_reps, is_train=False), \n                            batch_size=1, collate_fn=collate_bags_optimized, num_workers=self.dataloader_workers, pin_memory=True)\n        \n        preds = {}\n        self.model.eval()\n        with torch.no_grad():\n            for seqs_t, v_calls_t, j_calls_t, _, rep_ids, _, _ in tqdm(loader, desc=\"Inference\"):\n                seqs_t = seqs_t.to(self.device)\n                v_calls_t = v_calls_t.to(self.device)\n                j_calls_t = j_calls_t.to(self.device)\n                \n                logits, _ = self.model(seqs_t, v_calls_t, j_calls_t)\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        \n        rep_id_to_dataset_id = {}\n        for rep_id in self.train_data['repertoires'].keys():\n            try:\n                full_path = glob.glob(os.path.join(TRAIN_DIR, \"**\", f\"{rep_id}.tsv\"), recursive=True)[0]\n                dataset_id = os.path.basename(os.path.dirname(os.path.dirname(full_path)))\n                rep_id_to_dataset_id[rep_id] = dataset_id\n            except:\n                 rep_id_to_dataset_id[rep_id] = \"unknown_dataset\" \n\n\n        loader = DataLoader(AIRRDataset(list(self.train_data['repertoires'].keys()), self.train_data['repertoires'], is_train=False), \n                            batch_size=1, collate_fn=collate_bags_optimized, num_workers=self.dataloader_workers, pin_memory=True)\n        scores = {} \n\n        self.model.eval()\n        with torch.no_grad():\n            for seqs_t, v_calls_t, j_calls_t, _, rep_ids, dfs, _ in tqdm(loader, desc=\"Scanning Attention\"):\n                \n                rep_id = rep_ids[0]\n                ds_id = rep_id_to_dataset_id.get(rep_id)\n                if not ds_id or ds_id == \"unknown_dataset\": continue\n                \n                if ds_id not in scores: scores[ds_id] = {}\n                \n                seqs_t = seqs_t.to(self.device)\n                v_calls_t = v_calls_t.to(self.device)\n                j_calls_t = j_calls_t.to(self.device)\n                \n                _, attn = self.model(seqs_t, v_calls_t, j_calls_t)\n                attn = attn.cpu().numpy()\n                df = dfs[0] \n                \n                for i, r in df.iterrows():\n                    key = (r['junction_aa'], r['v_call'], r['j_call'])\n                    if i < len(attn) and (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)\n\n# =========================================================================================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T21:00:50.113079Z","iopub.execute_input":"2025-12-16T21:00:50.113315Z","iopub.status.idle":"2025-12-16T21:00:50.14114Z","shell.execute_reply.started":"2025-12-16T21:00:50.113297Z","shell.execute_reply":"2025-12-16T21:00:50.140265Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =========================================================================================\n# AIRR-ML-25: Professional Solution - Production V.13 (Cell 4/4: Execution & Submission)\n# =========================================================================================\n\n# Initialize\n# The execution should now be stable and much faster due to Batching and reduced complexity.\npredictor = ImmuneStatePredictor()\n\ntry:\n    if METADATA_PATH:\n        # A. Train\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        \n        # Prepare df1 for concatenation with dummy values for Task 2 columns (-999.0 as float)\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        \n        # Prepare df2 for concatenation with dummy values for Task 1 columns\n        df2['repertoire_id'] = \"dummy_id\"\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        \n        # Ensure correct data types for final submission file (as required by the competition)\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        \n    else:\n        print(\"❌ ERROR: Metadata file not found. Submission cannot be generated.\")\n\nexcept Exception as e:\n    print(f\"\\n❌ CRITICAL EXECUTION FAILURE: {e}\")\n\n# =========================================================================================","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T21:00:50.142567Z","iopub.execute_input":"2025-12-16T21:00:50.142944Z","execution_failed":"2025-12-16T22:04:35.456Z"}},"outputs":[],"execution_count":null}]}