{"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":31153,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nos.environ[\"TORCH_USE_CUDA_DSA\"] = \"1\"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T01:33:12.893529Z","iopub.execute_input":"2025-12-16T01:33:12.893917Z","iopub.status.idle":"2025-12-16T01:33:12.897488Z","shell.execute_reply.started":"2025-12-16T01:33:12.893897Z","shell.execute_reply":"2025-12-16T01:33:12.896930Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# AIRR-ML🧬25: Deep CNN + Attention MIL on GPU\n\nimport sys\nimport glob\nfrom collections import defaultdict\nfrom typing import List, Tuple, Iterator, Union\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.metrics import roc_auc_score, balanced_accuracy_score\nfrom sklearn.model_selection import StratifiedKFold, train_test_split\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# =========================\n# Basic utilities\n# =========================\n\nAMINO_ACIDS = list(\"ACDEFGHIKLMNPQRSTVWY\")\nAA_TO_IDX = {aa: i for i, aa in enumerate(AMINO_ACIDS)}\nPAD_IDX = len(AMINO_ACIDS)  # padding token index\nVOCAB_SIZE = len(AMINO_ACIDS) + 1  # +1 for PAD\n\ndef load_data_generator(\n    data_dir: str,\n    metadata_filename: str = \"metadata.csv\"\n) -> Iterator[Union[Tuple[str, pd.DataFrame, bool], Tuple[str, pd.DataFrame]]]:\n    metadata_path = os.path.join(data_dir, metadata_filename)\n    if os.path.exists(metadata_path):\n        metadata_df = pd.read_csv(metadata_path)\n        for row in metadata_df.itertuples(index=False):\n            file_path = os.path.join(data_dir, row.filename)\n            try:\n                df = pd.read_csv(file_path, sep=\"\\t\")\n                yield row.repertoire_id, df, bool(row.label_positive)\n            except FileNotFoundError:\n                print(f\"Warning: missing file '{row.filename}'\")\n    else:\n        for file_path in sorted(glob.glob(os.path.join(data_dir, \"*.tsv\"))):\n            try:\n                df = pd.read_csv(file_path, sep=\"\\t\")\n                filename = os.path.basename(file_path)\n                yield filename, df\n            except Exception as e:\n                print(f\"Warning: error reading '{file_path}': {e}\")\n\n\ndef load_full_dataset(data_dir: str) -> pd.DataFrame:\n    metadata_path = os.path.join(data_dir, \"metadata.csv\")\n    data_loader = load_data_generator(data_dir=data_dir)\n    dfs = []\n    if os.path.exists(metadata_path):\n        metadata_df = pd.read_csv(metadata_path)\n        total = len(metadata_df)\n        for rep_id, df, label in tqdm(data_loader, total=total, desc=\"Loading full dataset\"):\n            df[\"ID\"] = rep_id\n            df[\"label_positive\"] = label\n            dfs.append(df)\n    else:\n        tsv_files = glob.glob(os.path.join(data_dir, \"*.tsv\"))\n        total = len(tsv_files)\n        for fname, df in tqdm(data_loader, total=total, desc=\"Loading full dataset\"):\n            rep_id = os.path.basename(fname).replace(\".tsv\", \"\")\n            df[\"ID\"] = rep_id\n            dfs.append(df)\n    if not dfs:\n        return pd.DataFrame()\n    return pd.concat(dfs, ignore_index=True)\n\n\ndef get_repertoire_ids(data_dir: str) -> List[str]:\n    metadata_path = os.path.join(data_dir, \"metadata.csv\")\n    if os.path.exists(metadata_path):\n        meta = pd.read_csv(metadata_path)\n        return meta[\"repertoire_id\"].tolist()\n    tsv_files = glob.glob(os.path.join(data_dir, \"*.tsv\"))\n    return [os.path.basename(f).replace(\".tsv\", \"\") for f in sorted(tsv_files)]\n\n\ndef validate_dirs_and_files(train_dir: str, test_dirs: List[str], out_dir: str) -> None:\n    assert os.path.isdir(train_dir), f\"Train dir {train_dir} missing\"\n    assert os.path.isfile(os.path.join(train_dir, \"metadata.csv\")), \"metadata.csv missing in train\"\n    assert glob.glob(os.path.join(train_dir, \"*.tsv\")), \"No .tsv in train dir\"\n\n    for td in test_dirs:\n        assert os.path.isdir(td), f\"Test dir {td} missing\"\n        assert glob.glob(os.path.join(td, \"*.tsv\")), f\"No .tsv in test dir {td}\"\n\n    os.makedirs(out_dir, exist_ok=True)\n    tmp = os.path.join(out_dir, \"tmp.test\")\n    with open(tmp, \"w\") as f:\n        f.write(\"ok\")\n    os.remove(tmp)\n\n\ndef get_dataset_pairs(train_root: str, test_root: str) -> List[Tuple[str, List[str]]]:\n    test_groups = defaultdict(list)\n    for tname in sorted(os.listdir(test_root)):\n        if not tname.startswith(\"test_dataset_\"):\n            continue\n        base = tname.replace(\"test_dataset_\", \"\").split(\"_\")[0]\n        test_groups[base].append(os.path.join(test_root, tname))\n    pairs = []\n    for tname in sorted(os.listdir(train_root)):\n        if not tname.startswith(\"train_dataset_\"):\n            continue\n        base = tname.replace(\"train_dataset_\", \"\")\n        train_path = os.path.join(train_root, tname)\n        pairs.append((train_path, test_groups.get(base, [])))\n    return pairs\n\n\ndef save_tsv(df: pd.DataFrame, path: str):\n    os.makedirs(os.path.dirname(path), exist_ok=True)\n    df.to_csv(path, sep=\"\\t\", index=False)\n\n\ndef concatenate_output_files(out_dir: str) -> pd.DataFrame:\n    preds_files = sorted(glob.glob(os.path.join(out_dir, \"*_test_predictions.tsv\")))\n    seq_files = sorted(glob.glob(os.path.join(out_dir, \"*_important_sequences.tsv\")))\n    dfs = []\n    for f in preds_files + seq_files:\n        try:\n            dfs.append(pd.read_csv(f, sep=\"\\t\"))\n        except Exception as e:\n            print(f\"Warning reading {f}: {e}\")\n    if dfs:\n        all_df = pd.concat(dfs, ignore_index=True)\n    else:\n        all_df = pd.DataFrame(columns=[\"ID\",\"dataset\",\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"])\n    for col in [\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"]:\n        if col in all_df.columns:\n            all_df[col] = all_df[col].fillna(-999.0)\n    out_path = os.path.join(out_dir, \"submissions.csv\")\n    all_df.to_csv(out_path, index=False)\n    print(f\"Wrote {out_path} with shape {all_df.shape}\")\n    return all_df\n\n# =========================\n# Sequence encoding\n# =========================\n\ndef encode_sequence(seq: str, max_len: int) -> List[int]:\n    if not isinstance(seq, str):\n        seq = \"\"\n    seq = seq.strip()\n    ids = [AA_TO_IDX.get(ch, PAD_IDX) for ch in seq][:max_len]\n    if len(ids) < max_len:\n        ids += [PAD_IDX] * (max_len - len(ids))\n    return ids\n\n\ndef build_repertoire_tensors(\n    data_dir: str,\n    max_seqs_per_rep: int = 512,\n    max_len: int = 25,\n    for_training: bool = True\n):\n    \"\"\"\n    For each repertoire:\n      - sample up to max_seqs_per_rep sequences\n      - encode junction_aa as integer tokens\n    Returns:\n      rep_ids: list of repertoire IDs\n      X: tensor [N, max_seqs, max_len]\n      y: tensor [N] or None\n      per_rep_seq_lists: list of sequence DataFrames (for later importance scoring)\n    \"\"\"\n    loader = load_data_generator(data_dir=data_dir)\n    metadata_path = os.path.join(data_dir, \"metadata.csv\")\n    has_meta = os.path.exists(metadata_path)\n\n    rep_ids = []\n    labels = []\n    tensors = []\n    rep_seq_dfs = []\n\n    # determine quantile of lengths to set max_len adaptively if desired\n    # (here we use provided max_len for simplicity)\n\n    for item in tqdm(loader, desc=f\"Building repertoires ({'train' if for_training else 'test'})\"):\n        if has_meta:\n            rep_id, df, label = item\n        else:\n            rep_file, df = item\n            rep_id = os.path.basename(rep_file).replace(\".tsv\",\"\")\n            label = None\n\n        # drop missing junction_aa\n        df = df.dropna(subset=[\"junction_aa\"])\n        if df.empty:\n            continue\n\n        # sample or truncate sequences\n        if len(df) > max_seqs_per_rep:\n            df = df.sample(max_seqs_per_rep, random_state=42)\n        else:\n            df = df.sample(len(df), random_state=42)  # shuffle\n\n        # encode each sequence\n        seq_tensor = []\n        for s in df[\"junction_aa\"].tolist():\n            seq_tensor.append(encode_sequence(s, max_len=max_len))\n        # pad with dummy sequences if needed\n        if len(seq_tensor) < max_seqs_per_rep:\n            pad_seq = [PAD_IDX] * max_len\n            seq_tensor += [pad_seq] * (max_seqs_per_rep - len(seq_tensor))\n        seq_tensor = torch.tensor(seq_tensor, dtype=torch.long)  # [max_seqs, max_len]\n\n        rep_ids.append(rep_id)\n        tensors.append(seq_tensor.unsqueeze(0))  # [1, max_seqs, max_len]\n        rep_seq_dfs.append(df[[\"junction_aa\",\"v_call\",\"j_call\"]].reset_index(drop=True))\n\n        if has_meta:\n            labels.append(int(label))\n\n    if not tensors:\n        return [], None, None, []\n\n    X = torch.cat(tensors, dim=0)  # [N, max_seqs, max_len]\n    if has_meta and for_training:\n        y = torch.tensor(labels, dtype=torch.float32)\n    else:\n        y = None\n    return rep_ids, X, y, rep_seq_dfs\n\n# =========================\n# Deep MIL model (CNN + attention)\n# =========================\n\nclass CNNSeqEncoder(nn.Module):\n    def __init__(self, vocab_size, embed_dim=32, num_filters=64, kernel_sizes=(5,7)):\n        super().__init__()\n        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=PAD_IDX)\n        convs = []\n        for k in kernel_sizes:\n            convs.append(nn.Conv1d(embed_dim, num_filters, kernel_size=k, padding=k//2))\n        self.convs = nn.ModuleList(convs)\n        self.activation = nn.ReLU()\n\n        self.out_dim = num_filters * len(kernel_sizes)\n\n    def forward(self, x):\n        \"\"\"\n        x: [B, L] integer tokens\n        return: [B, out_dim]\n        \"\"\"\n        emb = self.embedding(x)           # [B, L, E]\n        emb = emb.transpose(1, 2)         # [B, E, L]\n        conv_outs = []\n        for conv in self.convs:\n            h = conv(emb)                 # [B, C, L]\n            h = self.activation(h)\n            h = torch.max(h, dim=2).values  # max over L -> [B, C]\n            conv_outs.append(h)\n        h_cat = torch.cat(conv_outs, dim=1)  # [B, out_dim]\n        return h_cat\n\n\nclass AttentionMIL(nn.Module):\n    def __init__(self, input_dim, att_dim=64):\n        super().__init__()\n        self.att_mlp = nn.Sequential(\n            nn.Linear(input_dim, att_dim),\n            nn.Tanh(),\n            nn.Linear(att_dim, 1)  # scalar attention logit per sequence\n        )\n\n    def forward(self, seq_repr):\n        \"\"\"\n        seq_repr: [B, S, D]\n        returns: (rep_repr [B, D], att_weights [B, S])\n        \"\"\"\n        B, S, D = seq_repr.shape\n        logits = self.att_mlp(seq_repr)         # [B, S, 1]\n        logits = logits.squeeze(-1)             # [B, S]\n        weights = F.softmax(logits, dim=1)      # [B, S]\n        rep_repr = torch.bmm(weights.unsqueeze(1), seq_repr).squeeze(1)  # [B, D]\n        return rep_repr, weights\n\n\nclass DeepRepertoireNet(nn.Module):\n    def __init__(self, vocab_size=VOCAB_SIZE, embed_dim=32, num_filters=64,\n                 kernel_sizes=(5,7), att_dim=64, hidden_dim=64):\n        super().__init__()\n        self.seq_encoder = CNNSeqEncoder(vocab_size, embed_dim, num_filters, kernel_sizes)\n        self.att_pool = AttentionMIL(self.seq_encoder.out_dim, att_dim=att_dim)\n        self.fc = nn.Sequential(\n            nn.Linear(self.seq_encoder.out_dim, hidden_dim),\n            nn.ReLU(),\n            nn.Linear(hidden_dim, 1)\n        )\n\n    def forward(self, x):\n        \"\"\"\n        x: [B, S, L]\n        Returns: logits [B], att_weights [B, S], seq_repr [B, S, D]\n        \"\"\"\n        B, S, L = x.shape\n        x_flat = x.view(B * S, L)\n        seq_repr = self.seq_encoder(x_flat)      # [B*S, D]\n        D = seq_repr.shape[1]\n        seq_repr = seq_repr.view(B, S, D)        # [B, S, D]\n        rep_repr, att_weights = self.att_pool(seq_repr)  # [B, D], [B, S]\n        logits = self.fc(rep_repr).squeeze(-1)   # [B]\n        return logits, att_weights, seq_repr\n\n\nclass RepertoireDataset(Dataset):\n    def __init__(self, X_tensor, y_tensor=None):\n        self.X = X_tensor\n        self.y = y_tensor\n\n    def __len__(self):\n        return self.X.shape[0]\n\n    def __getitem__(self, idx):\n        if self.y is None:\n            return self.X[idx]\n        return self.X[idx], self.y[idx]\n\n\ndef train_one_model(\n    X, y,\n    num_epochs=15,\n    batch_size=4,\n    lr=1e-3,\n    weight_decay=1e-4,\n    device=\"cpu\",\n    val_ratio=0.2,\n    seed=42,\n    plot_curves=True\n):\n    \"\"\"\n    Train DeepRepertoireNet on X, y (tensors), return model and (optionally) training curves.\n    \"\"\"\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n\n    # train/val split\n    idx = np.arange(len(y))\n    tr_idx, val_idx = train_test_split(idx, test_size=val_ratio, random_state=seed, stratify=y.numpy())\n    X_tr, y_tr = X[tr_idx], y[tr_idx]\n    X_val, y_val = X[val_idx], y[val_idx]\n\n    train_ds = RepertoireDataset(X_tr, y_tr)\n    val_ds = RepertoireDataset(X_val, y_val)\n\n    train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=2, pin_memory=True)\n    val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=2, pin_memory=True)\n\n    model = DeepRepertoireNet().to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)\n\n    best_auc = -np.inf\n    best_state = None\n    history = {\"train_loss\": [], \"val_loss\": [], \"val_auc\": []}\n\n    for epoch in range(1, num_epochs+1):\n        model.train()\n        train_losses = []\n        for xb, yb in train_loader:\n            xb = xb.to(device)\n            yb = yb.to(device)\n            optimizer.zero_grad()\n            logits, _, _ = model(xb)\n            loss = F.binary_cross_entropy_with_logits(logits, yb)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n            optimizer.step()\n            train_losses.append(loss.item())\n\n        # validation\n        model.eval()\n        val_losses = []\n        all_probs = []\n        all_labels = []\n        with torch.no_grad():\n            for xb, yb in val_loader:\n                xb = xb.to(device)\n                yb = yb.to(device)\n                logits, _, _ = model(xb)\n                loss = F.binary_cross_entropy_with_logits(logits, yb)\n                val_losses.append(loss.item())\n                probs = torch.sigmoid(logits).cpu().numpy()\n                all_probs.extend(probs)\n                all_labels.extend(yb.cpu().numpy())\n\n        val_auc = roc_auc_score(all_labels, all_probs)\n        mean_tr = float(np.mean(train_losses))\n        mean_val = float(np.mean(val_losses))\n        history[\"train_loss\"].append(mean_tr)\n        history[\"val_loss\"].append(mean_val)\n        history[\"val_auc\"].append(val_auc)\n\n        print(f\"Epoch {epoch:02d} | train_loss={mean_tr:.4f} | val_loss={mean_val:.4f} | val_auc={val_auc:.4f}\")\n\n        if val_auc > best_auc:\n            best_auc = val_auc\n            best_state = {k: v.cpu() for k, v in model.state_dict().items()}\n\n    if best_state is not None:\n        model.load_state_dict({k: v.to(device) for k, v in best_state.items()})\n\n    if plot_curves:\n        fig, ax1 = plt.subplots(figsize=(6,4))\n        ax1.plot(history[\"train_loss\"], label=\"train loss\")\n        ax1.plot(history[\"val_loss\"], label=\"val loss\")\n        ax1.set_xlabel(\"epoch\")\n        ax1.set_ylabel(\"loss\")\n        ax2 = ax1.twinx()\n        ax2.plot(history[\"val_auc\"], color=\"green\", label=\"val AUC\")\n        ax2.set_ylabel(\"AUC\")\n        ax1.legend(loc=\"upper left\")\n        ax2.legend(loc=\"upper right\")\n        plt.title(\"Training curves\")\n        plt.show()\n\n    return model, history\n\n# =========================\n# ImmuneStatePredictor wrapper\n# =========================\n\nclass ImmuneStatePredictor:\n    \"\"\"\n    Deep CNN+attention MIL model, template-compatible.\n    \"\"\"\n\n    def __init__(self, n_jobs: int = 1, device: str = \"cpu\", **kwargs):\n        self.n_jobs = n_jobs\n        # if device == \"cuda\" and not torch.cuda.is_available():\n        #     print(\"CUDA requested but not available; falling back to CPU.\")\n        #     device = \"cpu\"\n        self.device = device\n        self.model = None\n        self.rep_seq_dfs_ = None  # per-rep sequence DataFrames (train only)\n        self.train_rep_ids_ = None\n        self.max_seqs_per_rep = kwargs.get(\"max_seqs_per_rep\", 512)\n        self.max_len = kwargs.get(\"max_len\", 25)\n        self.num_epochs = kwargs.get(\"num_epochs\", 15)\n        self.batch_size = kwargs.get(\"batch_size\", 4)\n        self.lr = kwargs.get(\"lr\", 1e-3)\n        self.weight_decay = kwargs.get(\"weight_decay\", 1e-4)\n\n    def fit(self, train_dir_path: str):\n        print(f\"Building tensors and training DeepRepertoireNet for {train_dir_path}...\")\n        rep_ids, X, y, rep_seq_dfs = build_repertoire_tensors(\n            train_dir_path,\n            max_seqs_per_rep=self.max_seqs_per_rep,\n            max_len=self.max_len,\n            for_training=True\n        )\n        if X is None or y is None:\n            raise ValueError(\"No training data found.\")\n        self.train_rep_ids_ = rep_ids\n        self.rep_seq_dfs_ = rep_seq_dfs\n\n        device = self.device\n        X = X.to(device)\n        # y stays on CPU; moved in batches\n\n        self.model, _ = train_one_model(\n            X, y,\n            num_epochs=self.num_epochs,\n            batch_size=self.batch_size,\n            lr=self.lr,\n            weight_decay=self.weight_decay,\n            device=self.device,\n            val_ratio=0.2,\n        )\n        # after training, precompute sequence-level attention on full train to score sequences\n        self.important_sequences_ = self._identify_associated_sequences_internal(train_dir_path)\n        print(\"Training complete.\")\n        return self\n\n    def predict_proba(self, test_dir_path: str) -> pd.DataFrame:\n        if self.model is None:\n            raise RuntimeError(\"Model not trained yet.\")\n\n        print(f\"Preparing test repertoires from {test_dir_path}...\")\n        rep_ids, X, _, _ = build_repertoire_tensors(\n            test_dir_path,\n            max_seqs_per_rep=self.max_seqs_per_rep,\n            max_len=self.max_len,\n            for_training=False\n        )\n        if X is None or len(rep_ids) == 0:\n            return pd.DataFrame()\n\n        device = self.device\n        X = X.to(device)\n\n        ds = RepertoireDataset(X)\n        loader = DataLoader(ds, batch_size=self.batch_size, shuffle=False, num_workers=2, pin_memory=True)\n\n        self.model.eval()\n        all_probs = []\n        with torch.no_grad():\n            for xb in loader:\n                xb = xb.to(device)\n                logits, _, _ = self.model(xb)\n                probs = torch.sigmoid(logits).cpu().numpy()\n                all_probs.extend(probs)\n\n        dataset_name = os.path.basename(test_dir_path)\n        preds = pd.DataFrame({\n            \"ID\": rep_ids,\n            \"dataset\": [dataset_name] * len(rep_ids),\n            \"label_positive_probability\": all_probs\n        })\n        preds[\"junction_aa\"] = -999.0\n        preds[\"v_call\"] = -999.0\n        preds[\"j_call\"] = -999.0\n        preds = preds[[\"ID\",\"dataset\",\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"]]\n        print(f\"Predicted {len(preds)} repertoires in {test_dir_path}.\")\n        return preds\n\n    def _identify_associated_sequences_internal(self, train_dir_path: str, top_k: int = 50000) -> pd.DataFrame:\n        \"\"\"\n        Internal: use attention weights * per-sequence contribution to rank sequences.\n        \"\"\"\n        print(\"Scoring sequences for label association...\")\n        device = self.device\n        model = self.model\n        model.eval()\n\n        # Rebuild tensor with per-repertoire sequences in the same order as rep_seq_dfs_\n        rep_ids, X, y, rep_seq_dfs = build_repertoire_tensors(\n            train_dir_path,\n            max_seqs_per_rep=self.max_seqs_per_rep,\n            max_len=self.max_len,\n            for_training=True\n        )\n        X = X.to(device)\n\n        ds = RepertoireDataset(X, y)\n        loader = DataLoader(ds, batch_size=self.batch_size, shuffle=False, num_workers=2, pin_memory=True)\n\n        all_scores = []  # list of dicts: junction_aa, v_call, j_call, score\n\n        with torch.no_grad():\n            idx_offset = 0\n            for xb, yb in loader:\n                xb = xb.to(device)\n                logits, att_weights, seq_repr = model(xb)  # B, [B,S], [B,S,D]\n                probs = torch.sigmoid(logits)              # [B]\n\n                B, S = att_weights.shape\n                att_np = att_weights.cpu().numpy()\n                probs_np = probs.cpu().numpy()\n\n                # quick visualization: distribution of attention weights (optional)\n                # sns.histplot(att_np.flatten(), bins=50); plt.show()\n\n                for i in range(B):\n                    global_idx = idx_offset + i\n                    rep_df = rep_seq_dfs[global_idx]  # actual number of sequences may be < S\n                    num_real = min(len(rep_df), S)\n                    # sequence importance = attention_weight * sign(logit) * |logit|\n                    # (approx contribution)\n                    logit_i = float(logits[i].item())\n                    for j in range(num_real):\n                        score = att_np[i, j] * logit_i\n                        row = rep_df.iloc[j]\n                        all_scores.append({\n                            \"junction_aa\": row[\"junction_aa\"],\n                            \"v_call\": row.get(\"v_call\", np.nan),\n                            \"j_call\": row.get(\"j_call\", np.nan),\n                            \"score\": score\n                        })\n                idx_offset += B\n\n        seq_df = pd.DataFrame(all_scores)\n        # aggregate by unique (junction_aa, v_call, j_call): mean score\n        seq_df = seq_df.groupby([\"junction_aa\",\"v_call\",\"j_call\"], as_index=False)[\"score\"].mean()\n        seq_df = seq_df.sort_values(\"score\", ascending=False).head(top_k)\n\n        dataset_name = os.path.basename(train_dir_path)\n        seq_df[\"dataset\"] = dataset_name\n        seq_df[\"ID\"] = range(1, len(seq_df)+1)\n        seq_df[\"ID\"] = seq_df[\"dataset\"] + \"_seq_top_\" + seq_df[\"ID\"].astype(str)\n        seq_df[\"label_positive_probability\"] = -999.0\n        seq_df = seq_df[[\"ID\",\"dataset\",\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"]]\n\n        # simple viz: top motif lengths\n        plt.figure(figsize=(6,3))\n        seq_df[\"junction_aa\"].str.len().hist(bins=20)\n        plt.title(\"Top sequence length distribution\")\n        plt.xlabel(\"length\")\n        plt.ylabel(\"count\")\n        plt.show()\n\n        return seq_df\n\n    def identify_associated_sequences(self, train_dir_path: str, top_k: int = 50000) -> pd.DataFrame:\n        # Wrapper to comply with template; we already call internal version during fit.\n        return self._identify_associated_sequences_internal(train_dir_path, top_k=top_k)\n\n# =========================\n# Pipeline helpers\n# =========================\n\ndef _train_predictor(predictor: ImmuneStatePredictor, train_dir: str):\n    print(f\"Fitting model on {train_dir} ...\")\n    predictor.fit(train_dir)\n\n\ndef _generate_predictions(predictor: ImmuneStatePredictor, test_dirs: List[str]) -> pd.DataFrame:\n    all_preds = []\n    for td in test_dirs:\n        print(f\"Predicting on {td} ...\")\n        preds = predictor.predict_proba(td)\n        if preds is not None and not preds.empty:\n            all_preds.append(preds)\n    if all_preds:\n        return pd.concat(all_preds, ignore_index=True)\n    return pd.DataFrame()\n\n\ndef _save_predictions(predictions: pd.DataFrame, out_dir: str, train_dir: str):\n    if predictions.empty:\n        raise ValueError(\"No predictions to save.\")\n    path = os.path.join(out_dir, f\"{os.path.basename(train_dir)}_test_predictions.tsv\")\n    save_tsv(predictions, path)\n    print(f\"Saved predictions to {path}\")\n\n\ndef _save_important_sequences(predictor: ImmuneStatePredictor, out_dir: str, train_dir: str):\n    seqs = predictor.important_sequences_\n    if seqs is None or seqs.empty:\n        raise ValueError(\"No important sequences found.\")\n    path = os.path.join(out_dir, f\"{os.path.basename(train_dir)}_important_sequences.tsv\")\n    save_tsv(seqs, path)\n    print(f\"Saved important sequences to {path}\")\n\n\ndef main(train_dir: str, test_dirs: List[str], out_dir: str, n_jobs: int, device: str):\n    validate_dirs_and_files(train_dir, test_dirs, out_dir)\n    predictor = ImmuneStatePredictor(\n        n_jobs=n_jobs,\n        device=device,\n        max_seqs_per_rep=512,\n        max_len=25,\n        num_epochs=15,\n        batch_size=4,\n        lr=1e-3,\n        weight_decay=1e-4,\n    )\n    _train_predictor(predictor, train_dir)\n    preds = _generate_predictions(predictor, test_dirs)\n    _save_predictions(preds, out_dir, train_dir)\n    _save_important_sequences(predictor, out_dir, train_dir)\n\n\ndef run():\n    import argparse\n    parser = argparse.ArgumentParser()\n    parser.add_argument(\"--train_dir\", required=True)\n    parser.add_argument(\"--test_dirs\", required=True, nargs=\"+\")\n    parser.add_argument(\"--out_dir\", required=True)\n    parser.add_argument(\"--n_jobs\", type=int, default=2)\n    parser.add_argument(\"--device\", type=str)\n    args = parser.parse_args()\n    main(args.train_dir, args.test_dirs, args.out_dir, args.n_jobs, args.device)\n\n\n# =========================\n# Kaggle notebook entry\n# =========================\n\n","metadata":{"_uuid":"2114303b-e203-40ba-b227-2130fa1b3fdb","_cell_guid":"257fcd9c-c80c-4790-b346-0dee1d694608","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2025-12-16T01:36:15.617004Z","iopub.execute_input":"2025-12-16T01:36:15.617890Z","iopub.status.idle":"2025-12-16T01:36:23.443837Z","shell.execute_reply.started":"2025-12-16T01:36:15.617861Z","shell.execute_reply":"2025-12-16T01:36:23.443010Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T01:37:16.943752Z","iopub.execute_input":"2025-12-16T01:37:16.944419Z","iopub.status.idle":"2025-12-16T01:37:16.949047Z","shell.execute_reply.started":"2025-12-16T01:37:16.944393Z","shell.execute_reply":"2025-12-16T01:37:16.947801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    PATH_DATASET = \"/kaggle/input/adaptive-immune-profiling-challenge-2025\"\n    TRAIN_ROOT = os.path.join(PATH_DATASET, \"train_datasets\", \"train_datasets\")\n    TEST_ROOT = os.path.join(PATH_DATASET, \"test_datasets\", \"test_datasets\")\n    OUT_ROOT = \"/kaggle/working/results_deep_mil\"\n\n    device = 'cpu'\n    print(device)\n    print(\"Using device:\", device)\n\n    pairs = get_dataset_pairs(TRAIN_ROOT, TEST_ROOT)\n    print(\"Dataset pairs:\", pairs)\n\n    \n    for train_path, test_paths in pairs:\n        if not test_paths:\n            print(f\"No test sets for {train_path}, skipping.\")\n            continue\n        main(train_path, test_paths, OUT_ROOT, n_jobs=2, device=device)\n\n    concatenate_output_files(OUT_ROOT)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T01:37:18.308621Z","iopub.execute_input":"2025-12-16T01:37:18.309247Z","iopub.status.idle":"2025-12-16T02:23:45.968349Z","shell.execute_reply.started":"2025-12-16T01:37:18.309221Z","shell.execute_reply":"2025-12-16T02:23:45.966971Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# AIRR-ML🧬25: Improved Deep CNN + Attention MIL\n\nimport os\nimport sys\nimport glob\nfrom collections import defaultdict\nfrom typing import List, Tuple, Iterator, Union\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import train_test_split\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# =========================\n# Basic utilities\n# =========================\n\nAMINO_ACIDS = list(\"ACDEFGHIKLMNPQRSTVWY\")\nAA_TO_IDX = {aa: i for i, aa in enumerate(AMINO_ACIDS)}\nPAD_IDX = len(AMINO_ACIDS)\nVOCAB_SIZE = len(AMINO_ACIDS) + 1\n\ndef load_data_generator(\n    data_dir: str,\n    metadata_filename: str = \"metadata.csv\"\n) -> Iterator[Union[Tuple[str, pd.DataFrame, bool], Tuple[str, pd.DataFrame]]]:\n    metadata_path = os.path.join(data_dir, metadata_filename)\n    if os.path.exists(metadata_path):\n        metadata_df = pd.read_csv(metadata_path)\n        for row in metadata_df.itertuples(index=False):\n            file_path = os.path.join(data_dir, row.filename)\n            try:\n                df = pd.read_csv(file_path, sep=\"\\t\")\n                yield row.repertoire_id, df, bool(row.label_positive)\n            except FileNotFoundError:\n                print(f\"Warning: missing file '{row.filename}'\")\n    else:\n        for file_path in sorted(glob.glob(os.path.join(data_dir, \"*.tsv\"))):\n            try:\n                df = pd.read_csv(file_path, sep=\"\\t\")\n                filename = os.path.basename(file_path)\n                yield filename, df\n            except Exception as e:\n                print(f\"Warning: error reading '{file_path}': {e}\")\n\n\ndef load_full_dataset(data_dir: str) -> pd.DataFrame:\n    metadata_path = os.path.join(data_dir, \"metadata.csv\")\n    data_loader = load_data_generator(data_dir=data_dir)\n    dfs = []\n    if os.path.exists(metadata_path):\n        metadata_df = pd.read_csv(metadata_path)\n        total = len(metadata_df)\n        for rep_id, df, label in tqdm(data_loader, total=total, desc=\"Loading full dataset\"):\n            df[\"ID\"] = rep_id\n            df[\"label_positive\"] = label\n            dfs.append(df)\n    else:\n        tsv_files = glob.glob(os.path.join(data_dir, \"*.tsv\"))\n        total = len(tsv_files)\n        for fname, df in tqdm(data_loader, total=total, desc=\"Loading full dataset\"):\n            rep_id = os.path.basename(fname).replace(\".tsv\",\"\")\n            df[\"ID\"] = rep_id\n            dfs.append(df)\n    if not dfs:\n        return pd.DataFrame()\n    return pd.concat(dfs, ignore_index=True)\n\n\ndef get_repertoire_ids(data_dir: str) -> List[str]:\n    metadata_path = os.path.join(data_dir, \"metadata.csv\")\n    if os.path.exists(metadata_path):\n        meta = pd.read_csv(metadata_path)\n        return meta[\"repertoire_id\"].tolist()\n    tsv_files = glob.glob(os.path.join(data_dir, \"*.tsv\"))\n    return [os.path.basename(f).replace(\".tsv\", \"\") for f in sorted(tsv_files)]\n\n\ndef validate_dirs_and_files(train_dir: str, test_dirs: List[str], out_dir: str) -> None:\n    assert os.path.isdir(train_dir), f\"Train dir {train_dir} missing\"\n    assert os.path.isfile(os.path.join(train_dir, \"metadata.csv\")), \"metadata.csv missing in train\"\n    assert glob.glob(os.path.join(train_dir, \"*.tsv\")), \"No .tsv in train dir\"\n\n    for td in test_dirs:\n        assert os.path.isdir(td), f\"Test dir {td} missing\"\n        assert glob.glob(os.path.join(td, \"*.tsv\")), f\"No .tsv in test dir {td}\"\n\n    os.makedirs(out_dir, exist_ok=True)\n    tmp = os.path.join(out_dir, \"tmp.test\")\n    with open(tmp, \"w\") as f:\n        f.write(\"ok\")\n    os.remove(tmp)\n\n\ndef get_dataset_pairs(train_root: str, test_root: str) -> List[Tuple[str, List[str]]]:\n    test_groups = defaultdict(list)\n    for tname in sorted(os.listdir(test_root)):\n        if not tname.startswith(\"test_dataset_\"):\n            continue\n        base = tname.replace(\"test_dataset_\", \"\").split(\"_\")[0]\n        test_groups[base].append(os.path.join(test_root, tname))\n    pairs = []\n    for tname in sorted(os.listdir(train_root)):\n        if not tname.startswith(\"train_dataset_\"):\n            continue\n        base = tname.replace(\"train_dataset_\", \"\")\n        train_path = os.path.join(train_root, tname)\n        pairs.append((train_path, test_groups.get(base, [])))\n    return pairs\n\n\ndef save_tsv(df: pd.DataFrame, path: str):\n    os.makedirs(os.path.dirname(path), exist_ok=True)\n    df.to_csv(path, sep=\"\\t\", index=False)\n\n\ndef concatenate_output_files(out_dir: str) -> pd.DataFrame:\n    preds_files = sorted(glob.glob(os.path.join(out_dir, \"*_test_predictions.tsv\")))\n    seq_files = sorted(glob.glob(os.path.join(out_dir, \"*_important_sequences.tsv\")))\n    dfs = []\n    for f in preds_files + seq_files:\n        try:\n            dfs.append(pd.read_csv(f, sep=\"\\t\"))\n        except Exception as e:\n            print(f\"Warning reading {f}: {e}\")\n    if dfs:\n        all_df = pd.concat(dfs, ignore_index=True)\n    else:\n        all_df = pd.DataFrame(columns=[\"ID\",\"dataset\",\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"])\n    for col in [\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"]:\n        if col in all_df.columns:\n            all_df[col] = all_df[col].fillna(-999.0)\n    out_path = os.path.join(out_dir, \"submissions.csv\")\n    all_df.to_csv(out_path, index=False)\n    print(f\"Wrote {out_path} with shape {all_df.shape}\")\n    return all_df\n\n# =========================\n# Sequence encoding\n# =========================\n\ndef encode_sequence(seq: str, max_len: int) -> List[int]:\n    if not isinstance(seq, str):\n        seq = \"\"\n    seq = seq.strip()\n    ids = [AA_TO_IDX.get(ch, PAD_IDX) for ch in seq][:max_len]\n    if len(ids) < max_len:\n        ids += [PAD_IDX] * (max_len - len(ids))\n    return ids\n\n\ndef build_repertoire_tensors(\n    data_dir: str,\n    max_seqs_per_rep: int = 1024,\n    max_len: int = 35,\n    for_training: bool = True\n):\n    loader = load_data_generator(data_dir=data_dir)\n    metadata_path = os.path.join(data_dir, \"metadata.csv\")\n    has_meta = os.path.exists(metadata_path)\n\n    rep_ids = []\n    labels = []\n    tensors = []\n    rep_seq_dfs = []\n\n    for item in tqdm(loader, desc=f\"Building repertoires ({'train' if for_training else 'test'})\"):\n        if has_meta:\n            rep_id, df, label = item\n        else:\n            rep_file, df = item\n            rep_id = os.path.basename(rep_file).replace(\".tsv\",\"\")\n            label = None\n\n        df = df.dropna(subset=[\"junction_aa\"])\n        if df.empty:\n            continue\n\n        # shuffle and subsample up to max_seqs_per_rep\n        if len(df) > max_seqs_per_rep:\n            df = df.sample(max_seqs_per_rep, random_state=42)\n        else:\n            df = df.sample(len(df), random_state=42)\n\n        seq_tensor = [encode_sequence(s, max_len=max_len) for s in df[\"junction_aa\"].tolist()]\n        if len(seq_tensor) < max_seqs_per_rep:\n            pad_seq = [PAD_IDX] * max_len\n            seq_tensor += [pad_seq] * (max_seqs_per_rep - len(seq_tensor))\n        seq_tensor = torch.tensor(seq_tensor, dtype=torch.long)  # [S, L]\n\n        rep_ids.append(rep_id)\n        tensors.append(seq_tensor.unsqueeze(0))  # [1,S,L]\n        rep_seq_dfs.append(df[[\"junction_aa\",\"v_call\",\"j_call\"]].reset_index(drop=True))\n\n        if has_meta:\n            labels.append(int(label))\n\n    if not tensors:\n        return [], None, None, []\n\n    X = torch.cat(tensors, dim=0)  # [N,S,L]\n    if has_meta and for_training:\n        y = torch.tensor(labels, dtype=torch.float32)\n    else:\n        y = None\n    return rep_ids, X, y, rep_seq_dfs\n\n# =========================\n# Deep MIL model (larger CNN + attention)\n# =========================\n\nclass CNNSeqEncoder(nn.Module):\n    def __init__(self, vocab_size, embed_dim=64, num_filters=128, kernel_sizes=(5,7,9), dropout=0.2):\n        super().__init__()\n        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=PAD_IDX)\n        convs = []\n        for k in kernel_sizes:\n            convs.append(nn.Conv1d(embed_dim, num_filters, kernel_size=k, padding=k//2))\n        self.convs = nn.ModuleList(convs)\n        self.activation = nn.GELU()\n        self.dropout = nn.Dropout(dropout)\n        self.out_dim = num_filters * len(kernel_sizes)\n\n    def forward(self, x):\n        # x: [B,L]\n        emb = self.embedding(x)       # [B,L,E]\n        emb = emb.transpose(1, 2)     # [B,E,L]\n        conv_outs = []\n        for conv in self.convs:\n            h = conv(emb)             # [B,C,L]\n            h = self.activation(h)\n            h = F.max_pool1d(h, kernel_size=h.size(2)).squeeze(2)  # [B,C]\n            conv_outs.append(h)\n        h_cat = torch.cat(conv_outs, dim=1)      # [B,out_dim]\n        h_cat = self.dropout(h_cat)\n        return h_cat\n\n\nclass AttentionMIL(nn.Module):\n    def __init__(self, input_dim, att_dim=128, dropout=0.1):\n        super().__init__()\n        self.att_mlp = nn.Sequential(\n            nn.Linear(input_dim, att_dim),\n            nn.Tanh(),\n            nn.Dropout(dropout),\n            nn.Linear(att_dim, 1)\n        )\n\n    def forward(self, seq_repr):\n        # seq_repr: [B,S,D]\n        logits = self.att_mlp(seq_repr).squeeze(-1)     # [B,S]\n        weights = F.softmax(logits, dim=1)              # [B,S]\n        rep_repr = torch.bmm(weights.unsqueeze(1), seq_repr).squeeze(1)  # [B,D]\n        return rep_repr, weights\n\n\nclass DeepRepertoireNet(nn.Module):\n    def __init__(\n        self,\n        vocab_size=VOCAB_SIZE,\n        embed_dim=64,\n        num_filters=128,\n        kernel_sizes=(5,7,9),\n        att_dim=128,\n        hidden_dim=128,\n        dropout=0.3\n    ):\n        super().__init__()\n        self.seq_encoder = CNNSeqEncoder(\n            vocab_size, embed_dim=embed_dim,\n            num_filters=num_filters, kernel_sizes=kernel_sizes,\n            dropout=dropout*0.5\n        )\n        self.att_pool = AttentionMIL(self.seq_encoder.out_dim, att_dim=att_dim, dropout=dropout*0.5)\n        self.fc = nn.Sequential(\n            nn.Linear(self.seq_encoder.out_dim, hidden_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, 1)\n        )\n\n    def forward(self, x):\n        # x: [B,S,L]\n        B,S,L = x.shape\n        x_flat = x.view(B*S, L)\n        seq_repr_flat = self.seq_encoder(x_flat)   # [B*S,D]\n        D = seq_repr_flat.shape[1]\n        seq_repr = seq_repr_flat.view(B,S,D)       # [B,S,D]\n        rep_repr, att_weights = self.att_pool(seq_repr)  # [B,D],[B,S]\n        logits = self.fc(rep_repr).squeeze(-1)     # [B]\n        return logits, att_weights, seq_repr\n\n\nclass RepertoireDataset(Dataset):\n    def __init__(self, X_tensor, y_tensor=None):\n        self.X = X_tensor\n        self.y = y_tensor\n\n    def __len__(self):\n        return self.X.shape[0]\n\n    def __getitem__(self, idx):\n        if self.y is None:\n            return self.X[idx]\n        return self.X[idx], self.y[idx]\n\n\ndef train_one_model(\n    X, y,\n    num_epochs=30,\n    batch_size=8,\n    lr=1e-3,\n    weight_decay=5e-5,\n    device=\"cuda\",\n    val_ratio=0.2,\n    label_smoothing=0.05,\n    seed=42,\n    plot_curves=True\n):\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n\n    idx = np.arange(len(y))\n    tr_idx, val_idx = train_test_split(idx, test_size=val_ratio, random_state=seed, stratify=y.numpy())\n    X_tr, y_tr = X[tr_idx], y[tr_idx]\n    X_val, y_val = X[val_idx], y[val_idx]\n\n    train_ds = RepertoireDataset(X_tr, y_tr)\n    val_ds = RepertoireDataset(X_val, y_val)\n\n    train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=2, pin_memory=True)\n    val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=2, pin_memory=True)\n\n    model = DeepRepertoireNet().to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)\n\n    best_auc = -np.inf\n    best_state = None\n    history = {\"train_loss\": [], \"val_loss\": [], \"val_auc\": []}\n\n    for epoch in range(1, num_epochs+1):\n        model.train()\n        train_losses = []\n        for xb, yb in train_loader:\n            xb = xb.to(device)\n            yb = yb.to(device)\n            optimizer.zero_grad()\n            logits, _, _ = model(xb)\n            # label smoothing: y=(1-eps) for positives\n            y_smooth = yb * (1.0 - label_smoothing) + 0.5 * label_smoothing\n            loss = F.binary_cross_entropy_with_logits(logits, y_smooth)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n            optimizer.step()\n            train_losses.append(loss.item())\n\n        scheduler.step()\n\n        model.eval()\n        val_losses = []\n        all_probs = []\n        all_labels = []\n        with torch.no_grad():\n            for xb, yb in val_loader:\n                xb = xb.to(device)\n                yb = yb.to(device)\n                logits, _, _ = model(xb)\n                y_smooth = yb * (1.0 - label_smoothing) + 0.5 * label_smoothing\n                loss = F.binary_cross_entropy_with_logits(logits, y_smooth)\n                val_losses.append(loss.item())\n                probs = torch.sigmoid(logits).cpu().numpy()\n                all_probs.extend(probs)\n                all_labels.extend(yb.cpu().numpy())\n\n        val_auc = roc_auc_score(all_labels, all_probs)\n        mean_tr = float(np.mean(train_losses))\n        mean_val = float(np.mean(val_losses))\n        history[\"train_loss\"].append(mean_tr)\n        history[\"val_loss\"].append(mean_val)\n        history[\"val_auc\"].append(val_auc)\n        print(f\"Epoch {epoch:02d} | train_loss={mean_tr:.4f} | val_loss={mean_val:.4f} | val_auc={val_auc:.4f}\")\n\n        if val_auc > best_auc:\n            best_auc = val_auc\n            best_state = {k: v.cpu() for k, v in model.state_dict().items()}\n\n    if best_state is not None:\n        model.load_state_dict({k: v.to(device) for k, v in best_state.items()})\n\n    if plot_curves:\n        fig, ax1 = plt.subplots(figsize=(6,4))\n        ax1.plot(history[\"train_loss\"], label=\"train loss\")\n        ax1.plot(history[\"val_loss\"], label=\"val loss\")\n        ax1.set_xlabel(\"epoch\")\n        ax1.set_ylabel(\"loss\")\n        ax2 = ax1.twinx()\n        ax2.plot(history[\"val_auc\"], color=\"green\", label=\"val AUC\")\n        ax2.set_ylabel(\"AUC\")\n        ax1.legend(loc=\"upper left\")\n        ax2.legend(loc=\"upper right\")\n        plt.title(\"Training curves (improved model)\")\n        plt.show()\n\n    return model, history\n\n# =========================\n# ImmuneStatePredictor wrapper\n# =========================\n\nclass ImmuneStatePredictor:\n    \"\"\"\n    Deep CNN+attention MIL model, improved and template-compatible.\n    \"\"\"\n\n    def __init__(self, n_jobs: int = 1, device: str = \"cpu\", **kwargs):\n        self.n_jobs = n_jobs\n        if device == \"cuda\" and not torch.cuda.is_available():\n            print(\"CUDA requested but not available; falling back to CPU.\")\n            device = \"cpu\"\n        self.device = device\n        self.model = None\n        self.rep_seq_dfs_ = None\n        self.train_rep_ids_ = None\n\n        self.max_seqs_per_rep = kwargs.get(\"max_seqs_per_rep\", 1024)\n        self.max_len = kwargs.get(\"max_len\", 35)\n        self.num_epochs = kwargs.get(\"num_epochs\", 30)\n        self.batch_size = kwargs.get(\"batch_size\", 8)\n        self.lr = kwargs.get(\"lr\", 1e-3)\n        self.weight_decay = kwargs.get(\"weight_decay\", 5e-5)\n\n        self.important_sequences_ = None\n\n    def fit(self, train_dir_path: str):\n        print(f\"Building tensors and training DeepRepertoireNet for {train_dir_path}...\")\n        rep_ids, X, y, rep_seq_dfs = build_repertoire_tensors(\n            train_dir_path,\n            max_seqs_per_rep=self.max_seqs_per_rep,\n            max_len=self.max_len,\n            for_training=True\n        )\n        if X is None or y is None:\n            raise ValueError(\"No training data found.\")\n        self.train_rep_ids_ = rep_ids\n        self.rep_seq_dfs_ = rep_seq_dfs\n\n        device = self.device\n        X = X.to(device if device == \"cuda\" else \"cpu\")\n\n        self.model, _ = train_one_model(\n            X, y,\n            num_epochs=self.num_epochs,\n            batch_size=self.batch_size,\n            lr=self.lr,\n            weight_decay=self.weight_decay,\n            device=self.device,\n            val_ratio=0.2,\n            label_smoothing=0.05,\n        )\n\n        self.important_sequences_ = self._identify_associated_sequences_internal(train_dir_path)\n        print(\"Training complete.\")\n        return self\n\n    def predict_proba(self, test_dir_path: str) -> pd.DataFrame:\n        if self.model is None:\n            raise RuntimeError(\"Model not trained yet.\")\n\n        print(f\"Preparing test repertoires from {test_dir_path}...\")\n        rep_ids, X, _, _ = build_repertoire_tensors(\n            test_dir_path,\n            max_seqs_per_rep=self.max_seqs_per_rep,\n            max_len=self.max_len,\n            for_training=False\n        )\n        if X is None or len(rep_ids) == 0:\n            return pd.DataFrame()\n\n        device = self.device\n        X = X.to(device if device == \"cuda\" else \"cpu\")\n\n        ds = RepertoireDataset(X)\n        loader = DataLoader(ds, batch_size=self.batch_size, shuffle=False, num_workers=2, pin_memory=True)\n\n        self.model.eval()\n        all_probs = []\n        with torch.no_grad():\n            for xb in loader:\n                xb = xb.to(device)\n                logits, _, _ = self.model(xb)\n                probs = torch.sigmoid(logits).cpu().numpy()\n                all_probs.extend(probs)\n\n        dataset_name = os.path.basename(test_dir_path)\n        preds = pd.DataFrame({\n            \"ID\": rep_ids,\n            \"dataset\": [dataset_name] * len(rep_ids),\n            \"label_positive_probability\": all_probs\n        })\n        preds[\"junction_aa\"] = -999.0\n        preds[\"v_call\"] = -999.0\n        preds[\"j_call\"] = -999.0\n        preds = preds[[\"ID\",\"dataset\",\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"]]\n        print(f\"Predicted {len(preds)} repertoires in {test_dir_path}.\")\n        return preds\n\n    def _identify_associated_sequences_internal(self, train_dir_path: str, top_k: int = 50000) -> pd.DataFrame:\n        print(\"Scoring sequences for label association (improved model)...\")\n        device = self.device\n        model = self.model\n        model.eval()\n\n        rep_ids, X, y, rep_seq_dfs = build_repertoire_tensors(\n            train_dir_path,\n            max_seqs_per_rep=self.max_seqs_per_rep,\n            max_len=self.max_len,\n            for_training=True\n        )\n        X = X.to(device if device == \"cuda\" else \"cpu\")\n\n        ds = RepertoireDataset(X, y)\n        loader = DataLoader(ds, batch_size=self.batch_size, shuffle=False, num_workers=2, pin_memory=True)\n\n        all_scores = []\n\n        with torch.no_grad():\n            idx_offset = 0\n            for xb, yb in loader:\n                xb = xb.to(device)\n                logits, att_weights, seq_repr = model(xb)\n                B,S = att_weights.shape\n                att_np = att_weights.cpu().numpy()\n                logits_np = logits.cpu().numpy()\n\n                for i in range(B):\n                    global_idx = idx_offset + i\n                    rep_df = rep_seq_dfs[global_idx]\n                    num_real = min(len(rep_df), S)\n                    logit_i = float(logits_np[i])\n                    for j in range(num_real):\n                        score = att_np[i, j] * logit_i\n                        row = rep_df.iloc[j]\n                        all_scores.append({\n                            \"junction_aa\": row[\"junction_aa\"],\n                            \"v_call\": row.get(\"v_call\", np.nan),\n                            \"j_call\": row.get(\"j_call\", np.nan),\n                            \"score\": score\n                        })\n                idx_offset += B\n\n        seq_df = pd.DataFrame(all_scores)\n        seq_df = seq_df.groupby([\"junction_aa\",\"v_call\",\"j_call\"], as_index=False)[\"score\"].mean()\n        seq_df = seq_df.sort_values(\"score\", ascending=False).head(top_k)\n\n        dataset_name = os.path.basename(train_dir_path)\n        seq_df[\"dataset\"] = dataset_name\n        seq_df[\"ID\"] = range(1, len(seq_df)+1)\n        seq_df[\"ID\"] = seq_df[\"dataset\"] + \"_seq_top_\" + seq_df[\"ID\"].astype(str)\n        seq_df[\"label_positive_probability\"] = -999.0\n        seq_df = seq_df[[\"ID\",\"dataset\",\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"]]\n\n        # Visualization: sequence length distribution of top hits\n        plt.figure(figsize=(6,3))\n        seq_df[\"junction_aa\"].str.len().hist(bins=20)\n        plt.title(\"Top sequence length distribution (improved model)\")\n        plt.xlabel(\"length\")\n        plt.ylabel(\"count\")\n        plt.show()\n\n        return seq_df\n\n    def identify_associated_sequences(self, train_dir_path: str, top_k: int = 50000) -> pd.DataFrame:\n        return self._identify_associated_sequences_internal(train_dir_path, top_k=top_k)\n\n# =========================\n# Pipeline helpers\n# =========================\n\ndef _train_predictor(predictor: ImmuneStatePredictor, train_dir: str):\n    print(f\"Fitting model on {train_dir} ...\")\n    predictor.fit(train_dir)\n\n\ndef _generate_predictions(predictor: ImmuneStatePredictor, test_dirs: List[str]) -> pd.DataFrame:\n    all_preds = []\n    for td in test_dirs:\n        print(f\"Predicting on {td} ...\")\n        preds = predictor.predict_proba(td)\n        if preds is not None and not preds.empty:\n            all_preds.append(preds)\n    if all_preds:\n        return pd.concat(all_preds, ignore_index=True)\n    return pd.DataFrame()\n\n\ndef _save_predictions(predictions: pd.DataFrame, out_dir: str, train_dir: str):\n    if predictions.empty:\n        raise ValueError(\"No predictions to save.\")\n    path = os.path.join(out_dir, f\"{os.path.basename(train_dir)}_test_predictions.tsv\")\n    save_tsv(predictions, path)\n    print(f\"Saved predictions to {path}\")\n\n\ndef _save_important_sequences(predictor: ImmuneStatePredictor, out_dir: str, train_dir: str):\n    seqs = predictor.important_sequences_\n    if seqs is None or seqs.empty:\n        raise ValueError(\"No important sequences found.\")\n    path = os.path.join(out_dir, f\"{os.path.basename(train_dir)}_important_sequences.tsv\")\n    save_tsv(seqs, path)\n    print(f\"Saved important sequences to {path}\")\n\n\ndef main(train_dir: str, test_dirs: List[str], out_dir: str, n_jobs: int, device: str):\n    validate_dirs_and_files(train_dir, test_dirs, out_dir)\n    predictor = ImmuneStatePredictor(\n        n_jobs=n_jobs,\n        device=device,\n        max_seqs_per_rep=1024,\n        max_len=35,\n        num_epochs=30,\n        batch_size=8,\n        lr=1e-3,\n        weight_decay=5e-5,\n    )\n    _train_predictor(predictor, train_dir)\n    preds = _generate_predictions(predictor, test_dirs)\n    _save_predictions(preds, out_dir, train_dir)\n    _save_important_sequences(predictor, out_dir, train_dir)\n\n\ndef run():\n    import argparse\n    parser = argparse.ArgumentParser()\n    parser.add_argument(\"--train_dir\", required=True)\n    parser.add_argument(\"--test_dirs\", required=True, nargs=\"+\")\n    parser.add_argument(\"--out_dir\", required=True)\n    parser.add_argument(\"--n_jobs\", type=int, default=1)\n    parser.add_argument(\"--device\", type=str, default=\"cpu\", choices=[\"cpu\",\"cuda\"])\n    args = parser.parse_args()\n    main(args.train_dir, args.test_dirs, args.out_dir, args.n_jobs, args.device)\n\n\nif __name__ == \"__main__\":\n    PATH_DATASET = \"/kaggle/input/adaptive-immune-profiling-challenge-2025\"\n    TRAIN_ROOT = os.path.join(PATH_DATASET, \"train_datasets\", \"train_datasets\")\n    TEST_ROOT = os.path.join(PATH_DATASET, \"test_datasets\", \"test_datasets\")\n    OUT_ROOT = \"/kaggle/working/results_deep_mil_improved\"\n\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    print(\"Using device:\", device)\n\n    pairs = get_dataset_pairs(TRAIN_ROOT, TEST_ROOT)\n    print(\"Dataset pairs:\", pairs)\n\n    for train_path, test_paths in pairs:\n        if not test_paths:\n            print(f\"No test sets for {train_path}, skipping.\")\n            continue\n        main(train_path, test_paths, OUT_ROOT, n_jobs=4, device=device)\n\n    concatenate_output_files(OUT_ROOT)\n# AIRR-ML🧬25: Improved Deep CNN + Attention MIL\n\nimport os\nimport sys\nimport glob\nfrom collections import defaultdict\nfrom typing import List, Tuple, Iterator, Union\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import train_test_split\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# =========================\n# Basic utilities\n# =========================\n\nAMINO_ACIDS = list(\"ACDEFGHIKLMNPQRSTVWY\")\nAA_TO_IDX = {aa: i for i, aa in enumerate(AMINO_ACIDS)}\nPAD_IDX = len(AMINO_ACIDS)\nVOCAB_SIZE = len(AMINO_ACIDS) + 1\n\ndef load_data_generator(\n    data_dir: str,\n    metadata_filename: str = \"metadata.csv\"\n) -> Iterator[Union[Tuple[str, pd.DataFrame, bool], Tuple[str, pd.DataFrame]]]:\n    metadata_path = os.path.join(data_dir, metadata_filename)\n    if os.path.exists(metadata_path):\n        metadata_df = pd.read_csv(metadata_path)\n        for row in metadata_df.itertuples(index=False):\n            file_path = os.path.join(data_dir, row.filename)\n            try:\n                df = pd.read_csv(file_path, sep=\"\\t\")\n                yield row.repertoire_id, df, bool(row.label_positive)\n            except FileNotFoundError:\n                print(f\"Warning: missing file '{row.filename}'\")\n    else:\n        for file_path in sorted(glob.glob(os.path.join(data_dir, \"*.tsv\"))):\n            try:\n                df = pd.read_csv(file_path, sep=\"\\t\")\n                filename = os.path.basename(file_path)\n                yield filename, df\n            except Exception as e:\n                print(f\"Warning: error reading '{file_path}': {e}\")\n\n\ndef load_full_dataset(data_dir: str) -> pd.DataFrame:\n    metadata_path = os.path.join(data_dir, \"metadata.csv\")\n    data_loader = load_data_generator(data_dir=data_dir)\n    dfs = []\n    if os.path.exists(metadata_path):\n        metadata_df = pd.read_csv(metadata_path)\n        total = len(metadata_df)\n        for rep_id, df, label in tqdm(data_loader, total=total, desc=\"Loading full dataset\"):\n            df[\"ID\"] = rep_id\n            df[\"label_positive\"] = label\n            dfs.append(df)\n    else:\n        tsv_files = glob.glob(os.path.join(data_dir, \"*.tsv\"))\n        total = len(tsv_files)\n        for fname, df in tqdm(data_loader, total=total, desc=\"Loading full dataset\"):\n            rep_id = os.path.basename(fname).replace(\".tsv\",\"\")\n            df[\"ID\"] = rep_id\n            dfs.append(df)\n    if not dfs:\n        return pd.DataFrame()\n    return pd.concat(dfs, ignore_index=True)\n\n\ndef get_repertoire_ids(data_dir: str) -> List[str]:\n    metadata_path = os.path.join(data_dir, \"metadata.csv\")\n    if os.path.exists(metadata_path):\n        meta = pd.read_csv(metadata_path)\n        return meta[\"repertoire_id\"].tolist()\n    tsv_files = glob.glob(os.path.join(data_dir, \"*.tsv\"))\n    return [os.path.basename(f).replace(\".tsv\", \"\") for f in sorted(tsv_files)]\n\n\ndef validate_dirs_and_files(train_dir: str, test_dirs: List[str], out_dir: str) -> None:\n    assert os.path.isdir(train_dir), f\"Train dir {train_dir} missing\"\n    assert os.path.isfile(os.path.join(train_dir, \"metadata.csv\")), \"metadata.csv missing in train\"\n    assert glob.glob(os.path.join(train_dir, \"*.tsv\")), \"No .tsv in train dir\"\n\n    for td in test_dirs:\n        assert os.path.isdir(td), f\"Test dir {td} missing\"\n        assert glob.glob(os.path.join(td, \"*.tsv\")), f\"No .tsv in test dir {td}\"\n\n    os.makedirs(out_dir, exist_ok=True)\n    tmp = os.path.join(out_dir, \"tmp.test\")\n    with open(tmp, \"w\") as f:\n        f.write(\"ok\")\n    os.remove(tmp)\n\n\ndef get_dataset_pairs(train_root: str, test_root: str) -> List[Tuple[str, List[str]]]:\n    test_groups = defaultdict(list)\n    for tname in sorted(os.listdir(test_root)):\n        if not tname.startswith(\"test_dataset_\"):\n            continue\n        base = tname.replace(\"test_dataset_\", \"\").split(\"_\")[0]\n        test_groups[base].append(os.path.join(test_root, tname))\n    pairs = []\n    for tname in sorted(os.listdir(train_root)):\n        if not tname.startswith(\"train_dataset_\"):\n            continue\n        base = tname.replace(\"train_dataset_\", \"\")\n        train_path = os.path.join(train_root, tname)\n        pairs.append((train_path, test_groups.get(base, [])))\n    return pairs\n\n\ndef save_tsv(df: pd.DataFrame, path: str):\n    os.makedirs(os.path.dirname(path), exist_ok=True)\n    df.to_csv(path, sep=\"\\t\", index=False)\n\n\ndef concatenate_output_files(out_dir: str) -> pd.DataFrame:\n    preds_files = sorted(glob.glob(os.path.join(out_dir, \"*_test_predictions.tsv\")))\n    seq_files = sorted(glob.glob(os.path.join(out_dir, \"*_important_sequences.tsv\")))\n    dfs = []\n    for f in preds_files + seq_files:\n        try:\n            dfs.append(pd.read_csv(f, sep=\"\\t\"))\n        except Exception as e:\n            print(f\"Warning reading {f}: {e}\")\n    if dfs:\n        all_df = pd.concat(dfs, ignore_index=True)\n    else:\n        all_df = pd.DataFrame(columns=[\"ID\",\"dataset\",\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"])\n    for col in [\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"]:\n        if col in all_df.columns:\n            all_df[col] = all_df[col].fillna(-999.0)\n    out_path = os.path.join(out_dir, \"submissions.csv\")\n    all_df.to_csv(out_path, index=False)\n    print(f\"Wrote {out_path} with shape {all_df.shape}\")\n    return all_df\n\n# =========================\n# Sequence encoding\n# =========================\n\ndef encode_sequence(seq: str, max_len: int) -> List[int]:\n    if not isinstance(seq, str):\n        seq = \"\"\n    seq = seq.strip()\n    ids = [AA_TO_IDX.get(ch, PAD_IDX) for ch in seq][:max_len]\n    if len(ids) < max_len:\n        ids += [PAD_IDX] * (max_len - len(ids))\n    return ids\n\n\ndef build_repertoire_tensors(\n    data_dir: str,\n    max_seqs_per_rep: int = 1024,\n    max_len: int = 35,\n    for_training: bool = True\n):\n    loader = load_data_generator(data_dir=data_dir)\n    metadata_path = os.path.join(data_dir, \"metadata.csv\")\n    has_meta = os.path.exists(metadata_path)\n\n    rep_ids = []\n    labels = []\n    tensors = []\n    rep_seq_dfs = []\n\n    for item in tqdm(loader, desc=f\"Building repertoires ({'train' if for_training else 'test'})\"):\n        if has_meta:\n            rep_id, df, label = item\n        else:\n            rep_file, df = item\n            rep_id = os.path.basename(rep_file).replace(\".tsv\",\"\")\n            label = None\n\n        df = df.dropna(subset=[\"junction_aa\"])\n        if df.empty:\n            continue\n\n        # shuffle and subsample up to max_seqs_per_rep\n        if len(df) > max_seqs_per_rep:\n            df = df.sample(max_seqs_per_rep, random_state=42)\n        else:\n            df = df.sample(len(df), random_state=42)\n\n        seq_tensor = [encode_sequence(s, max_len=max_len) for s in df[\"junction_aa\"].tolist()]\n        if len(seq_tensor) < max_seqs_per_rep:\n            pad_seq = [PAD_IDX] * max_len\n            seq_tensor += [pad_seq] * (max_seqs_per_rep - len(seq_tensor))\n        seq_tensor = torch.tensor(seq_tensor, dtype=torch.long)  # [S, L]\n\n        rep_ids.append(rep_id)\n        tensors.append(seq_tensor.unsqueeze(0))  # [1,S,L]\n        rep_seq_dfs.append(df[[\"junction_aa\",\"v_call\",\"j_call\"]].reset_index(drop=True))\n\n        if has_meta:\n            labels.append(int(label))\n\n    if not tensors:\n        return [], None, None, []\n\n    X = torch.cat(tensors, dim=0)  # [N,S,L]\n    if has_meta and for_training:\n        y = torch.tensor(labels, dtype=torch.float32)\n    else:\n        y = None\n    return rep_ids, X, y, rep_seq_dfs\n\n# =========================\n# Deep MIL model (larger CNN + attention)\n# =========================\n\nclass CNNSeqEncoder(nn.Module):\n    def __init__(self, vocab_size, embed_dim=64, num_filters=128, kernel_sizes=(5,7,9), dropout=0.2):\n        super().__init__()\n        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=PAD_IDX)\n        convs = []\n        for k in kernel_sizes:\n            convs.append(nn.Conv1d(embed_dim, num_filters, kernel_size=k, padding=k//2))\n        self.convs = nn.ModuleList(convs)\n        self.activation = nn.GELU()\n        self.dropout = nn.Dropout(dropout)\n        self.out_dim = num_filters * len(kernel_sizes)\n\n    def forward(self, x):\n        # x: [B,L]\n        emb = self.embedding(x)       # [B,L,E]\n        emb = emb.transpose(1, 2)     # [B,E,L]\n        conv_outs = []\n        for conv in self.convs:\n            h = conv(emb)             # [B,C,L]\n            h = self.activation(h)\n            h = F.max_pool1d(h, kernel_size=h.size(2)).squeeze(2)  # [B,C]\n            conv_outs.append(h)\n        h_cat = torch.cat(conv_outs, dim=1)      # [B,out_dim]\n        h_cat = self.dropout(h_cat)\n        return h_cat\n\n\nclass AttentionMIL(nn.Module):\n    def __init__(self, input_dim, att_dim=128, dropout=0.1):\n        super().__init__()\n        self.att_mlp = nn.Sequential(\n            nn.Linear(input_dim, att_dim),\n            nn.Tanh(),\n            nn.Dropout(dropout),\n            nn.Linear(att_dim, 1)\n        )\n\n    def forward(self, seq_repr):\n        # seq_repr: [B,S,D]\n        logits = self.att_mlp(seq_repr).squeeze(-1)     # [B,S]\n        weights = F.softmax(logits, dim=1)              # [B,S]\n        rep_repr = torch.bmm(weights.unsqueeze(1), seq_repr).squeeze(1)  # [B,D]\n        return rep_repr, weights\n\n\nclass DeepRepertoireNet(nn.Module):\n    def __init__(\n        self,\n        vocab_size=VOCAB_SIZE,\n        embed_dim=64,\n        num_filters=128,\n        kernel_sizes=(5,7,9),\n        att_dim=128,\n        hidden_dim=128,\n        dropout=0.3\n    ):\n        super().__init__()\n        self.seq_encoder = CNNSeqEncoder(\n            vocab_size, embed_dim=embed_dim,\n            num_filters=num_filters, kernel_sizes=kernel_sizes,\n            dropout=dropout*0.5\n        )\n        self.att_pool = AttentionMIL(self.seq_encoder.out_dim, att_dim=att_dim, dropout=dropout*0.5)\n        self.fc = nn.Sequential(\n            nn.Linear(self.seq_encoder.out_dim, hidden_dim),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, 1)\n        )\n\n    def forward(self, x):\n        # x: [B,S,L]\n        B,S,L = x.shape\n        x_flat = x.view(B*S, L)\n        seq_repr_flat = self.seq_encoder(x_flat)   # [B*S,D]\n        D = seq_repr_flat.shape[1]\n        seq_repr = seq_repr_flat.view(B,S,D)       # [B,S,D]\n        rep_repr, att_weights = self.att_pool(seq_repr)  # [B,D],[B,S]\n        logits = self.fc(rep_repr).squeeze(-1)     # [B]\n        return logits, att_weights, seq_repr\n\n\nclass RepertoireDataset(Dataset):\n    def __init__(self, X_tensor, y_tensor=None):\n        self.X = X_tensor\n        self.y = y_tensor\n\n    def __len__(self):\n        return self.X.shape[0]\n\n    def __getitem__(self, idx):\n        if self.y is None:\n            return self.X[idx]\n        return self.X[idx], self.y[idx]\n\n\ndef train_one_model(\n    X, y,\n    num_epochs=30,\n    batch_size=8,\n    lr=1e-3,\n    weight_decay=5e-5,\n    device=\"cuda\",\n    val_ratio=0.2,\n    label_smoothing=0.05,\n    seed=42,\n    plot_curves=True\n):\n    torch.manual_seed(seed)\n    np.random.seed(seed)\n\n    idx = np.arange(len(y))\n    tr_idx, val_idx = train_test_split(idx, test_size=val_ratio, random_state=seed, stratify=y.numpy())\n    X_tr, y_tr = X[tr_idx], y[tr_idx]\n    X_val, y_val = X[val_idx], y[val_idx]\n\n    train_ds = RepertoireDataset(X_tr, y_tr)\n    val_ds = RepertoireDataset(X_val, y_val)\n\n    train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=2, pin_memory=True)\n    val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=2, pin_memory=True)\n\n    model = DeepRepertoireNet().to(device)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)\n\n    best_auc = -np.inf\n    best_state = None\n    history = {\"train_loss\": [], \"val_loss\": [], \"val_auc\": []}\n\n    for epoch in range(1, num_epochs+1):\n        model.train()\n        train_losses = []\n        for xb, yb in train_loader:\n            xb = xb.to(device)\n            yb = yb.to(device)\n            optimizer.zero_grad()\n            logits, _, _ = model(xb)\n            # label smoothing: y=(1-eps) for positives\n            y_smooth = yb * (1.0 - label_smoothing) + 0.5 * label_smoothing\n            loss = F.binary_cross_entropy_with_logits(logits, y_smooth)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)\n            optimizer.step()\n            train_losses.append(loss.item())\n\n        scheduler.step()\n\n        model.eval()\n        val_losses = []\n        all_probs = []\n        all_labels = []\n        with torch.no_grad():\n            for xb, yb in val_loader:\n                xb = xb.to(device)\n                yb = yb.to(device)\n                logits, _, _ = model(xb)\n                y_smooth = yb * (1.0 - label_smoothing) + 0.5 * label_smoothing\n                loss = F.binary_cross_entropy_with_logits(logits, y_smooth)\n                val_losses.append(loss.item())\n                probs = torch.sigmoid(logits).cpu().numpy()\n                all_probs.extend(probs)\n                all_labels.extend(yb.cpu().numpy())\n\n        val_auc = roc_auc_score(all_labels, all_probs)\n        mean_tr = float(np.mean(train_losses))\n        mean_val = float(np.mean(val_losses))\n        history[\"train_loss\"].append(mean_tr)\n        history[\"val_loss\"].append(mean_val)\n        history[\"val_auc\"].append(val_auc)\n        print(f\"Epoch {epoch:02d} | train_loss={mean_tr:.4f} | val_loss={mean_val:.4f} | val_auc={val_auc:.4f}\")\n\n        if val_auc > best_auc:\n            best_auc = val_auc\n            best_state = {k: v.cpu() for k, v in model.state_dict().items()}\n\n    if best_state is not None:\n        model.load_state_dict({k: v.to(device) for k, v in best_state.items()})\n\n    if plot_curves:\n        fig, ax1 = plt.subplots(figsize=(6,4))\n        ax1.plot(history[\"train_loss\"], label=\"train loss\")\n        ax1.plot(history[\"val_loss\"], label=\"val loss\")\n        ax1.set_xlabel(\"epoch\")\n        ax1.set_ylabel(\"loss\")\n        ax2 = ax1.twinx()\n        ax2.plot(history[\"val_auc\"], color=\"green\", label=\"val AUC\")\n        ax2.set_ylabel(\"AUC\")\n        ax1.legend(loc=\"upper left\")\n        ax2.legend(loc=\"upper right\")\n        plt.title(\"Training curves (improved model)\")\n        plt.show()\n\n    return model, history\n\n# =========================\n# ImmuneStatePredictor wrapper\n# =========================\n\nclass ImmuneStatePredictor:\n    \"\"\"\n    Deep CNN+attention MIL model, improved and template-compatible.\n    \"\"\"\n\n    def __init__(self, n_jobs: int = 1, device: str = \"cpu\", **kwargs):\n        self.n_jobs = n_jobs\n        if device == \"cuda\" and not torch.cuda.is_available():\n            print(\"CUDA requested but not available; falling back to CPU.\")\n            device = \"cpu\"\n        self.device = device\n        self.model = None\n        self.rep_seq_dfs_ = None\n        self.train_rep_ids_ = None\n\n        self.max_seqs_per_rep = kwargs.get(\"max_seqs_per_rep\", 1024)\n        self.max_len = kwargs.get(\"max_len\", 35)\n        self.num_epochs = kwargs.get(\"num_epochs\", 30)\n        self.batch_size = kwargs.get(\"batch_size\", 8)\n        self.lr = kwargs.get(\"lr\", 1e-3)\n        self.weight_decay = kwargs.get(\"weight_decay\", 5e-5)\n\n        self.important_sequences_ = None\n\n    def fit(self, train_dir_path: str):\n        print(f\"Building tensors and training DeepRepertoireNet for {train_dir_path}...\")\n        rep_ids, X, y, rep_seq_dfs = build_repertoire_tensors(\n            train_dir_path,\n            max_seqs_per_rep=self.max_seqs_per_rep,\n            max_len=self.max_len,\n            for_training=True\n        )\n        if X is None or y is None:\n            raise ValueError(\"No training data found.\")\n        self.train_rep_ids_ = rep_ids\n        self.rep_seq_dfs_ = rep_seq_dfs\n\n        device = self.device\n        X = X.to(device if device == \"cuda\" else \"cpu\")\n\n        self.model, _ = train_one_model(\n            X, y,\n            num_epochs=self.num_epochs,\n            batch_size=self.batch_size,\n            lr=self.lr,\n            weight_decay=self.weight_decay,\n            device=self.device,\n            val_ratio=0.2,\n            label_smoothing=0.05,\n        )\n\n        self.important_sequences_ = self._identify_associated_sequences_internal(train_dir_path)\n        print(\"Training complete.\")\n        return self\n\n    def predict_proba(self, test_dir_path: str) -> pd.DataFrame:\n        if self.model is None:\n            raise RuntimeError(\"Model not trained yet.\")\n\n        print(f\"Preparing test repertoires from {test_dir_path}...\")\n        rep_ids, X, _, _ = build_repertoire_tensors(\n            test_dir_path,\n            max_seqs_per_rep=self.max_seqs_per_rep,\n            max_len=self.max_len,\n            for_training=False\n        )\n        if X is None or len(rep_ids) == 0:\n            return pd.DataFrame()\n\n        device = self.device\n        X = X.to(device if device == \"cuda\" else \"cpu\")\n\n        ds = RepertoireDataset(X)\n        loader = DataLoader(ds, batch_size=self.batch_size, shuffle=False, num_workers=2, pin_memory=True)\n\n        self.model.eval()\n        all_probs = []\n        with torch.no_grad():\n            for xb in loader:\n                xb = xb.to(device)\n                logits, _, _ = self.model(xb)\n                probs = torch.sigmoid(logits).cpu().numpy()\n                all_probs.extend(probs)\n\n        dataset_name = os.path.basename(test_dir_path)\n        preds = pd.DataFrame({\n            \"ID\": rep_ids,\n            \"dataset\": [dataset_name] * len(rep_ids),\n            \"label_positive_probability\": all_probs\n        })\n        preds[\"junction_aa\"] = -999.0\n        preds[\"v_call\"] = -999.0\n        preds[\"j_call\"] = -999.0\n        preds = preds[[\"ID\",\"dataset\",\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"]]\n        print(f\"Predicted {len(preds)} repertoires in {test_dir_path}.\")\n        return preds\n\n    def _identify_associated_sequences_internal(self, train_dir_path: str, top_k: int = 50000) -> pd.DataFrame:\n        print(\"Scoring sequences for label association (improved model)...\")\n        device = self.device\n        model = self.model\n        model.eval()\n\n        rep_ids, X, y, rep_seq_dfs = build_repertoire_tensors(\n            train_dir_path,\n            max_seqs_per_rep=self.max_seqs_per_rep,\n            max_len=self.max_len,\n            for_training=True\n        )\n        X = X.to(device if device == \"cuda\" else \"cpu\")\n\n        ds = RepertoireDataset(X, y)\n        loader = DataLoader(ds, batch_size=self.batch_size, shuffle=False, num_workers=2, pin_memory=True)\n\n        all_scores = []\n\n        with torch.no_grad():\n            idx_offset = 0\n            for xb, yb in loader:\n                xb = xb.to(device)\n                logits, att_weights, seq_repr = model(xb)\n                B,S = att_weights.shape\n                att_np = att_weights.cpu().numpy()\n                logits_np = logits.cpu().numpy()\n\n                for i in range(B):\n                    global_idx = idx_offset + i\n                    rep_df = rep_seq_dfs[global_idx]\n                    num_real = min(len(rep_df), S)\n                    logit_i = float(logits_np[i])\n                    for j in range(num_real):\n                        score = att_np[i, j] * logit_i\n                        row = rep_df.iloc[j]\n                        all_scores.append({\n                            \"junction_aa\": row[\"junction_aa\"],\n                            \"v_call\": row.get(\"v_call\", np.nan),\n                            \"j_call\": row.get(\"j_call\", np.nan),\n                            \"score\": score\n                        })\n                idx_offset += B\n\n        seq_df = pd.DataFrame(all_scores)\n        seq_df = seq_df.groupby([\"junction_aa\",\"v_call\",\"j_call\"], as_index=False)[\"score\"].mean()\n        seq_df = seq_df.sort_values(\"score\", ascending=False).head(top_k)\n\n        dataset_name = os.path.basename(train_dir_path)\n        seq_df[\"dataset\"] = dataset_name\n        seq_df[\"ID\"] = range(1, len(seq_df)+1)\n        seq_df[\"ID\"] = seq_df[\"dataset\"] + \"_seq_top_\" + seq_df[\"ID\"].astype(str)\n        seq_df[\"label_positive_probability\"] = -999.0\n        seq_df = seq_df[[\"ID\",\"dataset\",\"label_positive_probability\",\"junction_aa\",\"v_call\",\"j_call\"]]\n\n        # Visualization: sequence length distribution of top hits\n        plt.figure(figsize=(6,3))\n        seq_df[\"junction_aa\"].str.len().hist(bins=20)\n        plt.title(\"Top sequence length distribution (improved model)\")\n        plt.xlabel(\"length\")\n        plt.ylabel(\"count\")\n        plt.show()\n\n        return seq_df\n\n    def identify_associated_sequences(self, train_dir_path: str, top_k: int = 50000) -> pd.DataFrame:\n        return self._identify_associated_sequences_internal(train_dir_path, top_k=top_k)\n\n# =========================\n# Pipeline helpers\n# =========================\n\ndef _train_predictor(predictor: ImmuneStatePredictor, train_dir: str):\n    print(f\"Fitting model on {train_dir} ...\")\n    predictor.fit(train_dir)\n\n\ndef _generate_predictions(predictor: ImmuneStatePredictor, test_dirs: List[str]) -> pd.DataFrame:\n    all_preds = []\n    for td in test_dirs:\n        print(f\"Predicting on {td} ...\")\n        preds = predictor.predict_proba(td)\n        if preds is not None and not preds.empty:\n            all_preds.append(preds)\n    if all_preds:\n        return pd.concat(all_preds, ignore_index=True)\n    return pd.DataFrame()\n\n\ndef _save_predictions(predictions: pd.DataFrame, out_dir: str, train_dir: str):\n    if predictions.empty:\n        raise ValueError(\"No predictions to save.\")\n    path = os.path.join(out_dir, f\"{os.path.basename(train_dir)}_test_predictions.tsv\")\n    save_tsv(predictions, path)\n    print(f\"Saved predictions to {path}\")\n\n\ndef _save_important_sequences(predictor: ImmuneStatePredictor, out_dir: str, train_dir: str):\n    seqs = predictor.important_sequences_\n    if seqs is None or seqs.empty:\n        raise ValueError(\"No important sequences found.\")\n    path = os.path.join(out_dir, f\"{os.path.basename(train_dir)}_important_sequences.tsv\")\n    save_tsv(seqs, path)\n    print(f\"Saved important sequences to {path}\")\n\n\ndef main(train_dir: str, test_dirs: List[str], out_dir: str, n_jobs: int, device: str):\n    validate_dirs_and_files(train_dir, test_dirs, out_dir)\n    predictor = ImmuneStatePredictor(\n        n_jobs=n_jobs,\n        device=device,\n        max_seqs_per_rep=1024,\n        max_len=35,\n        num_epochs=30,\n        batch_size=8,\n        lr=1e-3,\n        weight_decay=5e-5,\n    )\n    _train_predictor(predictor, train_dir)\n    preds = _generate_predictions(predictor, test_dirs)\n    _save_predictions(preds, out_dir, train_dir)\n    _save_important_sequences(predictor, out_dir, train_dir)\n\n\ndef run():\n    import argparse\n    parser = argparse.ArgumentParser()\n    parser.add_argument(\"--train_dir\", required=True)\n    parser.add_argument(\"--test_dirs\", required=True, nargs=\"+\")\n    parser.add_argument(\"--out_dir\", required=True)\n    parser.add_argument(\"--n_jobs\", type=int, default=1)\n    parser.add_argument(\"--device\", type=str, default=\"cpu\", choices=[\"cpu\",\"cuda\"])\n    args = parser.parse_args()\n    main(args.train_dir, args.test_dirs, args.out_dir, args.n_jobs, args.device)\n\n\nif __name__ == \"__main__\":\n    PATH_DATASET = \"/kaggle/input/adaptive-immune-profiling-challenge-2025\"\n    TRAIN_ROOT = os.path.join(PATH_DATASET, \"train_datasets\", \"train_datasets\")\n    TEST_ROOT = os.path.join(PATH_DATASET, \"test_datasets\", \"test_datasets\")\n    OUT_ROOT = \"/kaggle/working/results_deep_mil_improved\"\n\n    device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    print(\"Using device:\", device)\n\n    pairs = get_dataset_pairs(TRAIN_ROOT, TEST_ROOT)\n    print(\"Dataset pairs:\", pairs)\n\n    for train_path, test_paths in pairs:\n        if not test_paths:\n            print(f\"No test sets for {train_path}, skipping.\")\n            continue\n        main(train_path, test_paths, OUT_ROOT, n_jobs=4, device=device)\n\n    concatenate_output_files(OUT_ROOT)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T02:55:59.134984Z","iopub.execute_input":"2025-12-16T02:55:59.135436Z","execution_failed":"2025-12-17T20:42:05.990Z"}},"outputs":[],"execution_count":null}]}