{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"1d91421b","cell_type":"markdown","source":"# Phase 5 — Branch D: Deep MIL Repertoire Model (DeepRC successor, dual-head)\n\nEach repertoire = bag of sequences (top-K by abundance + random tail). Per-sequence encoder: small CNN/Transformer over CDR3 aa + V/J token + abundance scalar. Aggregator: attention pooling. Dual-head outputs: (1) repertoire-level disease score, (2) per-sequence attribution score. Outputs branch_d_oof.csv.\n\n---\n\n## Kaggle inputs to add before running this notebook\n\nAdd the following as Kaggle inputs (via the \"Add Input\" button on the right\npanel of the notebook editor):\n- AIRR-ML competition dataset\n- Phase 1 notebook output\n\n## Workflow\n\n1. Run the **SETUP** cell — it creates `/kaggle/working/project/` and auto-merges\n   any previous-phase notebook outputs found under `/kaggle/input/`.\n2. Run each subsequent cell in order. The **Run** cell executes the phase's\n   training script; the **Inspect** cell prints a quick summary of the outputs.\n3. When the run completes, click **Save Version → Save & Run All (Commit)** so\n   the next phase can pick this phase's outputs up via \"Add Input\".\n","metadata":{}},{"id":"27bb774b","cell_type":"code","source":"# ============================================================\n# SETUP — Initialize project + merge previous-phase inputs\n# ============================================================\n# This cell:\n#   1. Creates /kaggle/working/project/ fresh (idempotent re-runs).\n#   2. Auto-detects previous-phase notebook outputs under /kaggle/input/.\n#   3. Merges their project/ contents (src/, artifacts/, configs/) into\n#      /kaggle/working/project/ so this phase can build on them.\n#\n# Kaggle \"Add Input\" workflow:\n#   - Phase 1: add the AIRR-ML competition dataset only.\n#   - Phase N (N>=2): add the AIRR-ML competition dataset AND the previous\n#     phase notebook output(s). For Phase 9 add ALL of Phases 1..8.\n# ============================================================\n\nimport os\nimport shutil\nfrom pathlib import Path\n\nPROJECT_ROOT = Path(\"/kaggle/working/project\")\n\n# Reset working project dir (idempotent re-runs)\nif PROJECT_ROOT.exists():\n    shutil.rmtree(PROJECT_ROOT)\nPROJECT_ROOT.mkdir(parents=True, exist_ok=True)\n(PROJECT_ROOT / \"src\").mkdir(parents=True, exist_ok=True)\n(PROJECT_ROOT / \"src\" / \"__init__.py\").write_text(\"\", encoding=\"utf-8\")\n(PROJECT_ROOT / \"configs\").mkdir(parents=True, exist_ok=True)\n(PROJECT_ROOT / \"artifacts\").mkdir(parents=True, exist_ok=True)\n\n# Auto-detect and merge previous-phase inputs.\n# Each previous-phase notebook output should contain a top-level `project/`\n# directory (created by Phase 1 and propagated through every later phase).\nprev_inputs_found = []\ninput_root = Path(\"/kaggle/input\")\nif input_root.exists():\n    for entry in sorted(input_root.iterdir()):\n        if not entry.is_dir():\n            continue\n        prev_project = entry / \"project\"\n        if not prev_project.is_dir():\n            continue\n        prev_inputs_found.append(entry.name)\n        for sub in [\"src\", \"artifacts\", \"configs\"]:\n            src_dir = prev_project / sub\n            if not src_dir.exists():\n                continue\n            for path in src_dir.rglob(\"*\"):\n                if path.is_file():\n                    rel = path.relative_to(src_dir)\n                    target = PROJECT_ROOT / sub / rel\n                    target.parent.mkdir(parents=True, exist_ok=True)\n                    shutil.copy2(path, target)\n\nprint(\"PROJECT_ROOT :\", PROJECT_ROOT)\nif prev_inputs_found:\n    print(\"Previous-phase inputs merged:\")\n    for name in prev_inputs_found:\n        print(\"  -\", name)\nelse:\n    print(\"Previous-phase inputs: (none — running fresh)\")\nprint(\"\\nExisting artifacts:\")\nart_dir = PROJECT_ROOT / \"artifacts\"\nif art_dir.exists():\n    for p in sorted(art_dir.glob(\"*\")):\n        print(\"  -\", p.name)\nelse:\n    print(\"  (none)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T11:16:46.396844Z","iopub.execute_input":"2026-07-24T11:16:46.397147Z","iopub.status.idle":"2026-07-24T11:16:46.409404Z","shell.execute_reply.started":"2026-07-24T11:16:46.397122Z","shell.execute_reply":"2026-07-24T11:16:46.408778Z"}},"outputs":[],"execution_count":null},{"id":"efb5bf9b","cell_type":"code","source":"#Cell 1 — branch_d.yaml write করো\nfrom pathlib import Path\n\nPROJECT_ROOT = Path(\"/kaggle/working/project\")\n\nbranch_d_yaml = \"\"\"\nenabled: true\n\npaths:\n  output_root: /kaggle/working/project/artifacts/phase5_branch_d\n\ndictionary:\n  max_sequences_per_file: 30000\n  max_repertoires_per_class: 150\n  min_freq: 0.10\n  enrichment: 3.0\n  top_exact: 3000\n  top_clusters: 5000\n\nbag:\n  top_k_abundant: 256\n  random_tail: 128\n  force_public_top: 32\n  max_seq_len: 32\n\nmodel:\n  aa_emb_dim: 32\n  gene_emb_dim: 8\n  conv_channels: 64\n  hidden_dim: 128\n  dropout: 0.20\n\ntraining:\n  seeds: [42, 52, 62]\n  batch_size: 16\n  epochs: 8\n  lr: 0.001\n  weight_decay: 0.0001\n  lambda_attr: 0.25\n  lambda_sparse: 0.01\n\"\"\"\n\n(PROJECT_ROOT / \"configs\" / \"branch_d.yaml\").write_text(branch_d_yaml.strip() + \"\\n\", encoding=\"utf-8\")\nprint(\"Written:\", PROJECT_ROOT / \"configs\" / \"branch_d.yaml\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T11:16:46.410831Z","iopub.execute_input":"2026-07-24T11:16:46.411143Z","iopub.status.idle":"2026-07-24T11:16:46.426490Z","shell.execute_reply.started":"2026-07-24T11:16:46.411120Z","shell.execute_reply":"2026-07-24T11:16:46.425782Z"}},"outputs":[],"execution_count":null},{"id":"345a3a46","cell_type":"code","source":"#Cell 2 — src/branch_d_mil.py আর train_branch_d.py write করো\nfrom pathlib import Path\nfrom textwrap import dedent\n\nPROJECT_ROOT = Path(\"/kaggle/working/project\")\n\nbranch_d_mil_py = dedent(\"\"\"\nfrom __future__ import annotations\n\nimport random\nfrom pathlib import Path\nfrom typing import Dict, List, Tuple\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset\n\nfrom .branch_a_features import read_repertoire\nfrom .branch_b_cluster import approx_cluster_keys\nfrom .utils import stable_hash\n\nAA_VOCAB = [\"<PAD>\"] + list(\"ARNDCQEGHILKMFPSTWYV\")\nAA_TO_ID = {aa: i for i, aa in enumerate(AA_VOCAB)}\nPAD_ID = 0\n\n\ndef seed_torch(seed: int = 42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\n\ndef gene_family(gene_call: str) -> str:\n    if not isinstance(gene_call, str) or not gene_call:\n        return \"UNK\"\n    return gene_call.split(\"*\")[0].split(\"-\")[0].upper() or \"UNK\"\n\n\ndef gene_hash_id(gene_call: str, mod: int = 512) -> int:\n    fam = gene_family(gene_call)\n    return int(stable_hash(fam, mod=10_000_019) % mod)\n\n\ndef aggregate_unique_sequences(df: pd.DataFrame) -> pd.DataFrame:\n    if df is None or len(df) == 0:\n        return pd.DataFrame(columns=[\"junction_aa\", \"templates\", \"v_call\", \"j_call\"])\n\n    x = df.copy()\n    x[\"junction_aa\"] = x[\"junction_aa\"].fillna(\"\").astype(str)\n    x[\"templates\"] = pd.to_numeric(x[\"templates\"], errors=\"coerce\").fillna(1.0).clip(lower=1.0)\n\n    if \"v_call\" not in x.columns:\n        x[\"v_call\"] = \"\"\n    if \"j_call\" not in x.columns:\n        x[\"j_call\"] = \"\"\n\n    x = x[x[\"junction_aa\"] != \"\"].copy()\n    if len(x) == 0:\n        return pd.DataFrame(columns=[\"junction_aa\", \"templates\", \"v_call\", \"j_call\"])\n\n    grouped = (\n        x.groupby(\"junction_aa\", as_index=False)\n         .agg({\n             \"templates\": \"sum\",\n             \"v_call\": \"first\",\n             \"j_call\": \"first\",\n         })\n         .sort_values(\"templates\", ascending=False)\n         .reset_index(drop=True)\n    )\n    return grouped\n\n\ndef build_repertoire_table_cache(meta_ds: pd.DataFrame, ds_path: Path, max_sequences_per_file: int = 30000, random_state: int = 42):\n    cache = {}\n    for _, row in meta_ds.iterrows():\n        rep_id = row[\"repertoire_id\"]\n        fn = row[\"filename\"]\n        df = read_repertoire(ds_path / fn, max_seqs=max_sequences_per_file, random_state=random_state)\n        agg = aggregate_unique_sequences(df)\n        cache[rep_id] = agg\n    return cache\n\n\ndef weak_target_score(seq: str, bundle: dict) -> float:\n    exact_catalog = bundle.get(\"exact_catalog\", {})\n    cluster_catalog = bundle.get(\"cluster_catalog\", {})\n\n    score = 0.0\n    if seq in exact_catalog:\n        score += float(exact_catalog[seq][\"score\"])\n\n    best_cluster = 0.0\n    for key in approx_cluster_keys(seq):\n        if key in cluster_catalog:\n            best_cluster = max(best_cluster, float(cluster_catalog[key][\"score\"]))\n    score += best_cluster\n    return float(score)\n\n\ndef encode_seq(seq: str, max_len: int = 32) -> np.ndarray:\n    toks = np.zeros(max_len, dtype=np.int64)\n    if not seq:\n        return toks\n    seq = seq[:max_len]\n    for i, aa in enumerate(seq):\n        toks[i] = AA_TO_ID.get(aa, PAD_ID)\n    return toks\n\n\ndef build_bag_from_table(\n    agg_df: pd.DataFrame,\n    bundle: dict,\n    top_k_abundant: int = 256,\n    random_tail: int = 128,\n    force_public_top: int = 32,\n    max_seq_len: int = 32,\n    random_state: int = 42,\n):\n    rng = np.random.RandomState(random_state)\n\n    bag_size = int(top_k_abundant + random_tail)\n\n    if agg_df is None or len(agg_df) == 0:\n        return {\n            \"tokens\": np.zeros((bag_size, max_seq_len), dtype=np.int64),\n            \"v_ids\": np.zeros(bag_size, dtype=np.int64),\n            \"j_ids\": np.zeros(bag_size, dtype=np.int64),\n            \"abundance\": np.zeros(bag_size, dtype=np.float32),\n            \"lengths\": np.zeros(bag_size, dtype=np.float32),\n            \"weak_targets\": np.zeros(bag_size, dtype=np.float32),\n            \"mask\": np.zeros(bag_size, dtype=np.float32),\n            \"seqs\": [\"\"] * bag_size,\n        }\n\n    x = agg_df.copy().reset_index(drop=True)\n    x[\"templates\"] = pd.to_numeric(x[\"templates\"], errors=\"coerce\").fillna(1.0).clip(lower=1.0)\n    x = x.sort_values(\"templates\", ascending=False).reset_index(drop=True)\n\n    weak_scores = np.array([weak_target_score(seq, bundle) for seq in x[\"junction_aa\"].astype(str)], dtype=np.float32)\n    x[\"weak_score\"] = weak_scores\n\n    public_idx = np.where(weak_scores > 0)[0].tolist()\n    public_idx = sorted(public_idx, key=lambda i: (-weak_scores[i], -float(x.iloc[i][\"templates\"])))\n    public_idx = public_idx[:force_public_top]\n\n    top_idx = list(range(min(top_k_abundant, len(x))))\n    selected = []\n    seen = set()\n\n    for idx in public_idx + top_idx:\n        if idx not in seen:\n            selected.append(idx)\n            seen.add(idx)\n\n    remaining = [i for i in range(len(x)) if i not in seen]\n    tail_need = max(0, bag_size - len(selected))\n    if tail_need > 0 and len(remaining) > 0:\n        if len(remaining) <= tail_need:\n            tail_idx = remaining\n        else:\n            tail_idx = rng.choice(remaining, size=tail_need, replace=False).tolist()\n        for idx in tail_idx:\n            if idx not in seen:\n                selected.append(idx)\n                seen.add(idx)\n\n    selected = selected[:bag_size]\n    bag = x.iloc[selected].copy().reset_index(drop=True)\n\n    n_real = len(bag)\n    if n_real < bag_size:\n        pad_rows = pd.DataFrame({\n            \"junction_aa\": [\"\"] * (bag_size - n_real),\n            \"templates\": [0.0] * (bag_size - n_real),\n            \"v_call\": [\"\"] * (bag_size - n_real),\n            \"j_call\": [\"\"] * (bag_size - n_real),\n            \"weak_score\": [0.0] * (bag_size - n_real),\n        })\n        bag = pd.concat([bag, pad_rows], ignore_index=True)\n\n    seqs = bag[\"junction_aa\"].astype(str).tolist()\n    tokens = np.vstack([encode_seq(s, max_len=max_seq_len) for s in seqs]).astype(np.int64)\n\n    total_templates = float(max(1e-9, bag[\"templates\"].sum()))\n    abundance = np.log1p(pd.to_numeric(bag[\"templates\"], errors=\"coerce\").fillna(0.0).values.astype(np.float32))\n    abundance = abundance / max(1e-9, abundance.max() if abundance.max() > 0 else 1.0)\n\n    lengths = np.array([min(len(s), max_seq_len) / max_seq_len for s in seqs], dtype=np.float32)\n    weak = bag[\"weak_score\"].fillna(0.0).astype(np.float32).values\n    weak = weak / max(1e-9, weak.max() if weak.max() > 0 else 1.0)\n\n    v_ids = np.array([gene_hash_id(v, mod=512) for v in bag[\"v_call\"].fillna(\"\").astype(str)], dtype=np.int64)\n    j_ids = np.array([gene_hash_id(j, mod=512) for j in bag[\"j_call\"].fillna(\"\").astype(str)], dtype=np.int64)\n\n    mask = np.array([1.0 if s else 0.0 for s in seqs], dtype=np.float32)\n\n    return {\n        \"tokens\": tokens,\n        \"v_ids\": v_ids,\n        \"j_ids\": j_ids,\n        \"abundance\": abundance.astype(np.float32),\n        \"lengths\": lengths.astype(np.float32),\n        \"weak_targets\": weak.astype(np.float32),\n        \"mask\": mask.astype(np.float32),\n        \"seqs\": seqs,\n    }\n\n\nclass RepertoireBagDataset(Dataset):\n    def __init__(self, bag_records: List[dict]):\n        self.bag_records = bag_records\n\n    def __len__(self):\n        return len(self.bag_records)\n\n    def __getitem__(self, idx):\n        r = self.bag_records[idx]\n        return {\n            \"ID\": r[\"ID\"],\n            \"dataset\": r[\"dataset\"],\n            \"label\": torch.tensor(float(r[\"label_positive\"]), dtype=torch.float32),\n            \"tokens\": torch.tensor(r[\"tokens\"], dtype=torch.long),\n            \"v_ids\": torch.tensor(r[\"v_ids\"], dtype=torch.long),\n            \"j_ids\": torch.tensor(r[\"j_ids\"], dtype=torch.long),\n            \"abundance\": torch.tensor(r[\"abundance\"], dtype=torch.float32),\n            \"lengths\": torch.tensor(r[\"lengths\"], dtype=torch.float32),\n            \"weak_targets\": torch.tensor(r[\"weak_targets\"], dtype=torch.float32),\n            \"mask\": torch.tensor(r[\"mask\"], dtype=torch.float32),\n            \"seqs\": r[\"seqs\"],\n        }\n\n\nclass SequenceEncoder(nn.Module):\n    def __init__(self, aa_vocab_size=21, aa_emb_dim=32, gene_vocab_size=512, gene_emb_dim=8, conv_channels=64, hidden_dim=128, dropout=0.2):\n        super().__init__()\n        self.aa_emb = nn.Embedding(aa_vocab_size, aa_emb_dim, padding_idx=0)\n        self.v_emb = nn.Embedding(gene_vocab_size, gene_emb_dim)\n        self.j_emb = nn.Embedding(gene_vocab_size, gene_emb_dim)\n\n        self.conv1 = nn.Conv1d(aa_emb_dim, conv_channels, kernel_size=3, padding=1)\n        self.conv2 = nn.Conv1d(conv_channels, conv_channels, kernel_size=5, padding=2)\n        self.proj = nn.Sequential(\n            nn.Linear(conv_channels * 2 + gene_emb_dim * 2 + 2, hidden_dim),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.ReLU(),\n        )\n\n    def forward(self, tokens, v_ids, j_ids, abundance, lengths):\n        # tokens: [B, N, L]\n        B, N, L = tokens.shape\n        x = self.aa_emb(tokens)                  # [B, N, L, E]\n        x = x.view(B * N, L, -1).transpose(1, 2)  # [B*N, E, L]\n        x = F.relu(self.conv1(x))\n        x = F.relu(self.conv2(x))\n\n        x_max = F.adaptive_max_pool1d(x, 1).squeeze(-1)\n        x_mean = F.adaptive_avg_pool1d(x, 1).squeeze(-1)\n\n        v = self.v_emb(v_ids.view(B * N))\n        j = self.j_emb(j_ids.view(B * N))\n\n        aux = torch.stack([\n            abundance.view(B * N),\n            lengths.view(B * N),\n        ], dim=1)\n\n        out = torch.cat([x_max, x_mean, v, j, aux], dim=1)\n        out = self.proj(out)\n        out = out.view(B, N, -1)\n        return out\n\n\nclass DeepMILModel(nn.Module):\n    def __init__(self, aa_emb_dim=32, gene_emb_dim=8, conv_channels=64, hidden_dim=128, dropout=0.2):\n        super().__init__()\n        self.encoder = SequenceEncoder(\n            aa_emb_dim=aa_emb_dim,\n            gene_emb_dim=gene_emb_dim,\n            conv_channels=conv_channels,\n            hidden_dim=hidden_dim,\n            dropout=dropout,\n        )\n        self.attn = nn.Sequential(\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.Tanh(),\n            nn.Linear(hidden_dim, 1),\n        )\n        self.attr_head = nn.Linear(hidden_dim, 1)\n        self.cls_head = nn.Sequential(\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, 1),\n        )\n\n    def forward(self, tokens, v_ids, j_ids, abundance, lengths, mask):\n        seq_repr = self.encoder(tokens, v_ids, j_ids, abundance, lengths)  # [B, N, H]\n        attn_logits = self.attn(seq_repr).squeeze(-1)                      # [B, N]\n        attn_logits = attn_logits.masked_fill(mask <= 0, -1e9)\n        attn_weights = torch.softmax(attn_logits, dim=1)                   # [B, N]\n\n        bag_repr = torch.sum(seq_repr * attn_weights.unsqueeze(-1), dim=1)\n        bag_logit = self.cls_head(bag_repr).squeeze(-1)                    # [B]\n\n        attr_logits = self.attr_head(seq_repr).squeeze(-1)                 # [B, N]\n        attr_probs = torch.sigmoid(attr_logits) * mask\n\n        final_seq_score = attr_probs * attn_weights * mask\n        return {\n            \"bag_logit\": bag_logit,\n            \"attn_weights\": attn_weights,\n            \"attr_logits\": attr_logits,\n            \"attr_probs\": attr_probs,\n            \"final_seq_score\": final_seq_score,\n        }\n\n\ndef compute_losses(outputs, labels, weak_targets, mask, lambda_attr=0.25, lambda_sparse=0.01):\n    bag_logit = outputs[\"bag_logit\"]\n    attn = outputs[\"attn_weights\"]\n    attr_logits = outputs[\"attr_logits\"]\n\n    cls_loss = F.binary_cross_entropy_with_logits(bag_logit, labels)\n\n    weak_binary = (weak_targets > 0).float()\n    attr_bce = F.binary_cross_entropy_with_logits(attr_logits, weak_binary, reduction=\"none\")\n    attr_bce = (attr_bce * mask).sum() / mask.sum().clamp_min(1.0)\n\n    weak_sum = weak_targets.sum(dim=1, keepdim=True)\n    weak_norm = weak_targets / weak_sum.clamp_min(1e-8)\n\n    valid_rank = (weak_sum.squeeze(1) > 0).float()\n    kl = weak_norm * (torch.log(weak_norm.clamp_min(1e-8)) - torch.log(attn.clamp_min(1e-8)))\n    kl = kl.sum(dim=1)\n    rank_consistency = (kl * valid_rank).sum() / valid_rank.sum().clamp_min(1.0)\n\n    attr_loss = attr_bce + rank_consistency\n\n    entropy = -(attn * torch.log(attn.clamp_min(1e-8)) * mask).sum(dim=1) / mask.sum(dim=1).clamp_min(1.0)\n    sparse_loss = entropy.mean()\n\n    total = cls_loss + lambda_attr * attr_loss + lambda_sparse * sparse_loss\n\n    return {\n        \"total\": total,\n        \"cls\": cls_loss.detach().item(),\n        \"attr\": attr_loss.detach().item(),\n        \"sparse\": sparse_loss.detach().item(),\n    }\n\n\ndef build_bag_records_from_meta(\n    meta_df: pd.DataFrame,\n    table_cache: dict,\n    bundle: dict,\n    bag_cfg: dict,\n    random_state: int = 42,\n):\n    rows = []\n    for i, (_, row) in enumerate(meta_df.iterrows()):\n        rep_id = row[\"repertoire_id\"]\n        agg_df = table_cache.get(rep_id, pd.DataFrame(columns=[\"junction_aa\", \"templates\", \"v_call\", \"j_call\"]))\n        bag = build_bag_from_table(\n            agg_df=agg_df,\n            bundle=bundle,\n            top_k_abundant=bag_cfg[\"top_k_abundant\"],\n            random_tail=bag_cfg[\"random_tail\"],\n            force_public_top=bag_cfg[\"force_public_top\"],\n            max_seq_len=bag_cfg[\"max_seq_len\"],\n            random_state=random_state + i,\n        )\n        rows.append({\n            \"ID\": rep_id,\n            \"dataset\": row[\"dataset_name\"],\n            \"label_positive\": int(row[\"label_positive\"]) if pd.notna(row[\"label_positive\"]) else 0,\n            **bag,\n        })\n    return rows\n\"\"\")\n\ntrain_branch_d_py = dedent(\"\"\"\nfrom __future__ import annotations\n\nimport sys\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import DataLoader\nimport yaml\n\nPROJECT_ROOT = Path(__file__).resolve().parent\nif str(PROJECT_ROOT) not in sys.path:\n    sys.path.insert(0, str(PROJECT_ROOT))\n\nfrom src.branch_b_public import build_dataset_cache, build_fold_signal_catalog\nfrom src.branch_d_mil import (\n    DeepMILModel,\n    RepertoireBagDataset,\n    build_bag_records_from_meta,\n    build_repertoire_table_cache,\n    compute_losses,\n    seed_torch,\n)\nfrom src.branch_a_models import safe_auc\nfrom src.utils import ensure_dir, seed_everything, dataset_id_from_name\n\nSCALE_POS_WEIGHT = {\n    1: 1.0, 2: 1.0, 3: 1.0, 4: 1.0, 5: 1.0, 6: 1.0, 7: 5.04, 8: 2.05\n}\n\n\ndef collate_fn(batch):\n    return {\n        \"ID\": [x[\"ID\"] for x in batch],\n        \"dataset\": [x[\"dataset\"] for x in batch],\n        \"label\": torch.stack([x[\"label\"] for x in batch]),\n        \"tokens\": torch.stack([x[\"tokens\"] for x in batch]),\n        \"v_ids\": torch.stack([x[\"v_ids\"] for x in batch]),\n        \"j_ids\": torch.stack([x[\"j_ids\"] for x in batch]),\n        \"abundance\": torch.stack([x[\"abundance\"] for x in batch]),\n        \"lengths\": torch.stack([x[\"lengths\"] for x in batch]),\n        \"weak_targets\": torch.stack([x[\"weak_targets\"] for x in batch]),\n        \"mask\": torch.stack([x[\"mask\"] for x in batch]),\n        \"seqs\": [x[\"seqs\"] for x in batch],\n    }\n\n\ndef move_to_device(batch, device):\n    out = {}\n    for k, v in batch.items():\n        if torch.is_tensor(v):\n            out[k] = v.to(device)\n        else:\n            out[k] = v\n    return out\n\n\ndef evaluate_model(model, loader, device):\n    model.eval()\n    bag_probs = []\n    bag_labels = []\n\n    seq_rows = []\n    attn_rows = []\n\n    with torch.no_grad():\n        for batch in loader:\n            batch_dev = move_to_device(batch, device)\n            outputs = model(\n                tokens=batch_dev[\"tokens\"],\n                v_ids=batch_dev[\"v_ids\"],\n                j_ids=batch_dev[\"j_ids\"],\n                abundance=batch_dev[\"abundance\"],\n                lengths=batch_dev[\"lengths\"],\n                mask=batch_dev[\"mask\"],\n            )\n\n            probs = torch.sigmoid(outputs[\"bag_logit\"]).detach().cpu().numpy()\n            attn = outputs[\"attn_weights\"].detach().cpu().numpy()\n            attr = outputs[\"attr_probs\"].detach().cpu().numpy()\n            final_seq = outputs[\"final_seq_score\"].detach().cpu().numpy()\n            weak = batch[\"weak_targets\"].numpy()\n            mask = batch[\"mask\"].numpy()\n\n            bag_probs.extend(probs.tolist())\n            bag_labels.extend(batch[\"label\"].numpy().tolist())\n\n            for i, rep_id in enumerate(batch[\"ID\"]):\n                attn_rows.append({\n                    \"ID\": rep_id,\n                    \"attn\": attn[i].astype(np.float32),\n                    \"attr\": attr[i].astype(np.float32),\n                    \"final_seq\": final_seq[i].astype(np.float32),\n                    \"mask\": mask[i].astype(np.float32),\n                })\n\n                valid_idx = np.where(mask[i] > 0)[0].tolist()\n                seq_entries = []\n                for j in valid_idx:\n                    seq_entries.append({\n                        \"ID\": rep_id,\n                        \"dataset\": batch[\"dataset\"][i],\n                        \"rank\": int(j + 1),\n                        \"sequence\": batch[\"seqs\"][i][j],\n                        \"attention_weight\": float(attn[i][j]),\n                        \"attr_prob\": float(attr[i][j]),\n                        \"final_seq_score\": float(final_seq[i][j]),\n                        \"weak_target\": float(weak[i][j]),\n                    })\n                seq_rows.extend(seq_entries)\n\n    auc = safe_auc(np.array(bag_labels), np.array(bag_probs))\n    return auc, np.array(bag_probs), seq_rows, attn_rows\n\n\ndef train_one_fold(train_records, val_records, cfg_model, cfg_train, device, seed):\n    seed_torch(seed)\n\n    train_ds = RepertoireBagDataset(train_records)\n    val_ds = RepertoireBagDataset(val_records)\n\n    train_loader = DataLoader(train_ds, batch_size=cfg_train[\"batch_size\"], shuffle=True, collate_fn=collate_fn)\n    val_loader = DataLoader(val_ds, batch_size=cfg_train[\"batch_size\"], shuffle=False, collate_fn=collate_fn)\n\n    model = DeepMILModel(\n        aa_emb_dim=cfg_model[\"aa_emb_dim\"],\n        gene_emb_dim=cfg_model[\"gene_emb_dim\"],\n        conv_channels=cfg_model[\"conv_channels\"],\n        hidden_dim=cfg_model[\"hidden_dim\"],\n        dropout=cfg_model[\"dropout\"],\n    ).to(device)\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=cfg_train[\"lr\"],\n        weight_decay=cfg_train[\"weight_decay\"],\n    )\n\n    best_state = None\n    best_auc = -np.inf\n    best_epoch = -1\n\n    for epoch in range(cfg_train[\"epochs\"]):\n        model.train()\n        for batch in train_loader:\n            batch_dev = move_to_device(batch, device)\n            outputs = model(\n                tokens=batch_dev[\"tokens\"],\n                v_ids=batch_dev[\"v_ids\"],\n                j_ids=batch_dev[\"j_ids\"],\n                abundance=batch_dev[\"abundance\"],\n                lengths=batch_dev[\"lengths\"],\n                mask=batch_dev[\"mask\"],\n            )\n            losses = compute_losses(\n                outputs=outputs,\n                labels=batch_dev[\"label\"],\n                weak_targets=batch_dev[\"weak_targets\"],\n                mask=batch_dev[\"mask\"],\n                lambda_attr=cfg_train[\"lambda_attr\"],\n                lambda_sparse=cfg_train[\"lambda_sparse\"],\n            )\n\n            optimizer.zero_grad()\n            losses[\"total\"].backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0)\n            optimizer.step()\n\n        val_auc, _, _, _ = evaluate_model(model, val_loader, device)\n        if np.isnan(val_auc):\n            val_auc = -1.0\n\n        if val_auc > best_auc:\n            best_auc = val_auc\n            best_epoch = epoch\n            best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}\n\n    if best_state is not None:\n        model.load_state_dict(best_state)\n\n    val_auc, val_probs, seq_rows, attn_rows = evaluate_model(model, val_loader, device)\n    return model, val_auc, val_probs, seq_rows, attn_rows, best_epoch\n\n\ndef main():\n    cfg_base = yaml.safe_load((PROJECT_ROOT / \"configs\" / \"base.yaml\").read_text())\n    cfg_branch = yaml.safe_load((PROJECT_ROOT / \"configs\" / \"branch_d.yaml\").read_text())\n\n    runtime = cfg_base[\"runtime\"]\n    phase1_paths = cfg_base[\"paths\"]\n    out_dir = ensure_dir(cfg_branch[\"paths\"][\"output_root\"])\n\n    seed_everything(runtime[\"random_state\"])\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n    canonical = pd.read_csv(Path(phase1_paths[\"output_root\"]) / \"canonical_metadata.csv\")\n    all_folds = pd.read_csv(Path(phase1_paths[\"output_root\"]) / \"all_folds.csv\")\n    train_root = Path(phase1_paths[\"train_root\"])\n\n    train_meta_all = canonical[canonical[\"source\"] == \"train\"].copy()\n    train_meta_all = train_meta_all[pd.notna(train_meta_all[\"label_positive\"])].copy()\n\n    dataset_names = sorted(train_meta_all[\"dataset_name\"].unique().tolist())\n    seeds = list(cfg_branch[\"training\"][\"seeds\"])\n\n    cfg_dict = cfg_branch[\"dictionary\"]\n    cfg_bag = cfg_branch[\"bag\"]\n    cfg_model = cfg_branch[\"model\"]\n    cfg_train = cfg_branch[\"training\"]\n\n    all_oof = []\n\n    print(\"=\" * 80)\n    print(\"PHASE 5: Branch D - Deep MIL repertoire model\")\n    print(\"=\" * 80)\n    print(\"Device:\", device)\n    print(\"Datasets:\", dataset_names)\n\n    for ds_name in dataset_names:\n        ds_id = dataset_id_from_name(ds_name)\n        ds_path = train_root / ds_name\n\n        print(f\"\\\\n{'-' * 80}\")\n        print(f\"Branch D training: {ds_name} (id={ds_id})\")\n        print(f\"{'-' * 80}\")\n\n        meta_ds = train_meta_all[train_meta_all[\"dataset_name\"] == ds_name].copy()\n        fold_ds = all_folds[all_folds[\"dataset_name\"] == ds_name].copy()\n\n        meta_ds = meta_ds.merge(\n            fold_ds[[\"dataset_name\", \"repertoire_id\", \"fold\", \"splitter\", \"n_splits_used\"]],\n            on=[\"dataset_name\", \"repertoire_id\"],\n            how=\"inner\",\n        ).reset_index(drop=True)\n\n        signal_cache = build_dataset_cache(\n            meta_ds=meta_ds,\n            ds_path=ds_path,\n            max_sequences_per_file=cfg_dict[\"max_sequences_per_file\"],\n            random_state=runtime[\"random_state\"],\n            n_jobs=runtime[\"n_jobs\"],\n        )\n        table_cache = build_repertoire_table_cache(\n            meta_ds=meta_ds,\n            ds_path=ds_path,\n            max_sequences_per_file=cfg_dict[\"max_sequences_per_file\"],\n            random_state=runtime[\"random_state\"],\n        )\n\n        unique_folds = sorted(meta_ds[\"fold\"].dropna().unique().tolist())\n        ds_seed_frames = []\n        ds_seq_rows = []\n        ds_attn_maps = []\n        ds_summary_rows = []\n\n        for seed in seeds:\n            pred_store = {\n                \"ID\": meta_ds[\"repertoire_id\"].astype(str).tolist(),\n                \"dataset\": meta_ds[\"dataset_name\"].astype(str).tolist(),\n                \"label_positive\": meta_ds[\"label_positive\"].astype(int).tolist(),\n                \"fold\": meta_ds[\"fold\"].tolist(),\n                \"seed\": [seed] * len(meta_ds),\n                \"MIL_PROB\": [np.nan] * len(meta_ds),\n            }\n\n            for fold in unique_folds:\n                tr_meta = meta_ds[meta_ds[\"fold\"] != fold].copy().reset_index(drop=True)\n                va_meta = meta_ds[meta_ds[\"fold\"] == fold].copy().reset_index(drop=True)\n                va_global_idx = meta_ds.index[meta_ds[\"fold\"] == fold].to_numpy()\n\n                bundle, exact_df, cluster_df = build_fold_signal_catalog(\n                    fold_train_meta=tr_meta,\n                    cache=signal_cache,\n                    min_freq=cfg_dict[\"min_freq\"],\n                    enrichment=cfg_dict[\"enrichment\"],\n                    top_exact=cfg_dict[\"top_exact\"],\n                    top_clusters=cfg_dict[\"top_clusters\"],\n                    max_repertoires_per_class=cfg_dict[\"max_repertoires_per_class\"],\n                    random_state=seed,\n                )\n\n                train_records = build_bag_records_from_meta(\n                    meta_df=tr_meta,\n                    table_cache=table_cache,\n                    bundle=bundle,\n                    bag_cfg=cfg_bag,\n                    random_state=seed,\n                )\n                val_records = build_bag_records_from_meta(\n                    meta_df=va_meta,\n                    table_cache=table_cache,\n                    bundle=bundle,\n                    bag_cfg=cfg_bag,\n                    random_state=seed + 1000,\n                )\n\n                model, val_auc, val_probs, seq_rows, attn_rows, best_epoch = train_one_fold(\n                    train_records=train_records,\n                    val_records=val_records,\n                    cfg_model=cfg_model,\n                    cfg_train=cfg_train,\n                    device=device,\n                    seed=seed,\n                )\n\n                for idx, p in zip(va_global_idx, val_probs):\n                    pred_store[\"MIL_PROB\"][idx] = float(p)\n\n                ds_summary_rows.append({\n                    \"dataset\": ds_name,\n                    \"seed\": seed,\n                    \"fold\": fold,\n                    \"model\": \"MIL\",\n                    \"auc\": val_auc,\n                    \"best_epoch\": best_epoch,\n                })\n\n                for row in seq_rows:\n                    row[\"seed\"] = seed\n                    row[\"fold\"] = fold\n                ds_seq_rows.extend(seq_rows)\n\n                for row in attn_rows:\n                    row[\"seed\"] = seed\n                    row[\"fold\"] = fold\n                ds_attn_maps.extend(attn_rows)\n\n            seed_df = pd.DataFrame(pred_store)\n            ds_seed_frames.append(seed_df)\n\n            ds_summary_rows.append({\n                \"dataset\": ds_name,\n                \"seed\": seed,\n                \"fold\": -1,\n                \"model\": \"MIL\",\n                \"auc\": safe_auc(seed_df[\"label_positive\"].values, seed_df[\"MIL_PROB\"].values),\n                \"best_epoch\": -1,\n            })\n\n        ds_oof = pd.concat(ds_seed_frames, ignore_index=True)\n        ds_oof.to_csv(out_dir / f\"branch_d_oof_{ds_name}.csv\", index=False)\n        all_oof.append(ds_oof)\n\n        # save sequence-level scores\n        seq_df = pd.DataFrame(ds_seq_rows)\n        if len(seq_df):\n            seq_path = out_dir / f\"mil_sequence_scores_{ds_name}.parquet\"\n            try:\n                seq_df.to_parquet(seq_path, index=False)\n            except Exception:\n                seq_df.to_csv(out_dir / f\"mil_sequence_scores_{ds_name}_fallback.csv\", index=False)\n\n        # save attention maps\n        if len(ds_attn_maps):\n            attn_ids = [x[\"ID\"] for x in ds_attn_maps]\n            attn_seed = [x[\"seed\"] for x in ds_attn_maps]\n            attn_fold = [x[\"fold\"] for x in ds_attn_maps]\n            attn_arr = np.vstack([x[\"attn\"] for x in ds_attn_maps]).astype(np.float32)\n            attr_arr = np.vstack([x[\"attr\"] for x in ds_attn_maps]).astype(np.float32)\n            final_arr = np.vstack([x[\"final_seq\"] for x in ds_attn_maps]).astype(np.float32)\n            mask_arr = np.vstack([x[\"mask\"] for x in ds_attn_maps]).astype(np.float32)\n\n            np.savez_compressed(\n                out_dir / f\"mil_attention_maps_{ds_name}.npz\",\n                repertoire_ids=np.array(attn_ids, dtype=object),\n                seeds=np.array(attn_seed, dtype=np.int64),\n                folds=np.array(attn_fold, dtype=np.int64),\n                attention=attn_arr,\n                attr_probs=attr_arr,\n                final_seq_scores=final_arr,\n                mask=mask_arr,\n            )\n\n        summary_df = pd.DataFrame(ds_summary_rows)\n        if len(summary_df):\n            summary_df.to_csv(out_dir / f\"branch_d_summary_{ds_name}.csv\", index=False)\n\n            overall = (\n                summary_df[summary_df[\"fold\"] >= 0]\n                .groupby([\"dataset\", \"seed\", \"model\"], as_index=False)\n                .agg(mean_auc=(\"auc\", \"mean\"), std_auc=(\"auc\", \"std\"), max_auc=(\"auc\", \"max\"))\n                .sort_values([\"dataset\", \"seed\", \"mean_auc\"], ascending=[True, True, False])\n            )\n            print(\"\\\\nSummary:\")\n            print(overall.to_string(index=False))\n\n        print(f\"Saved OOF -> {out_dir / f'branch_d_oof_{ds_name}.csv'}\")\n        print(f\"Saved seq parquet -> {out_dir / f'mil_sequence_scores_{ds_name}.parquet'}\")\n        print(f\"Saved attn npz -> {out_dir / f'mil_attention_maps_{ds_name}.npz'}\")\n\n    final_oof = pd.concat(all_oof, ignore_index=True) if len(all_oof) else pd.DataFrame()\n    final_oof.to_csv(out_dir / \"branch_d_oof.csv\", index=False)\n\n    print(\"\\\\n\" + \"=\" * 80)\n    print(\"PHASE 5 DONE\")\n    print(\"Saved:\")\n    print(out_dir / \"branch_d_oof.csv\")\n    print(\"=\" * 80)\n\n\nif __name__ == \"__main__\":\n    main()\n\"\"\")\n\n(PROJECT_ROOT / \"src\" / \"branch_d_mil.py\").write_text(branch_d_mil_py, encoding=\"utf-8\")\n(PROJECT_ROOT / \"train_branch_d.py\").write_text(train_branch_d_py, encoding=\"utf-8\")\n\nprint(\"Written:\")\nprint(\"-\", PROJECT_ROOT / \"src\" / \"branch_d_mil.py\")\nprint(\"-\", PROJECT_ROOT / \"train_branch_d.py\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T11:16:46.527536Z","iopub.execute_input":"2026-07-24T11:16:46.528157Z","iopub.status.idle":"2026-07-24T11:16:46.548313Z","shell.execute_reply.started":"2026-07-24T11:16:46.528129Z","shell.execute_reply":"2026-07-24T11:16:46.547690Z"}},"outputs":[],"execution_count":null},{"id":"9843c5d1","cell_type":"code","source":"#Cell 3 — imports verify করো\nfrom pathlib import Path\np = Path(\"/kaggle/working/project/src\")\nprint(\"src/ exists:\", p.exists())\nprint(\"contents:\", sorted(f.name for f in p.glob(\"*\")) if p.exists() else \"N/A\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T11:16:46.549943Z","iopub.execute_input":"2026-07-24T11:16:46.550229Z","iopub.status.idle":"2026-07-24T11:16:46.566865Z","shell.execute_reply.started":"2026-07-24T11:16:46.550196Z","shell.execute_reply":"2026-07-24T11:16:46.566115Z"}},"outputs":[],"execution_count":null},{"id":"d64cf62a","cell_type":"code","source":"#Cell 4 — Phase 5 run করো\n!python /kaggle/working/project/train_branch_d.py","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T11:16:46.567690Z","iopub.execute_input":"2026-07-24T11:16:46.568004Z","iopub.status.idle":"2026-07-24T11:16:51.057886Z","shell.execute_reply.started":"2026-07-24T11:16:46.567983Z","shell.execute_reply":"2026-07-24T11:16:51.057156Z"}},"outputs":[],"execution_count":null},{"id":"d2a19b35","cell_type":"code","source":"#Cell 5 — outputs inspect করো\nfrom pathlib import Path\nimport pandas as pd\n\nOUT_DIR = Path(\"/kaggle/working/project/artifacts/phase3_branch_b\")\n\nprint(\"OUT_DIR exists:\", OUT_DIR.exists())\nprint(\"Files:\")\nfound = sorted(OUT_DIR.glob(\"*\")) if OUT_DIR.exists() else []\nfor p in found:\n    print(\"-\", p.name)\nif not found:\n    print(\"(none — Branch B training/inference did not write any outputs here.)\")\n\ndef show(name, path):\n    print(f\"\\n{name}\")\n    if path.exists():\n        display(pd.read_csv(path).head())\n    else:\n        print(f\"  -> missing: {path}\")\n\nshow(\"branch_b_oof.csv\", OUT_DIR / \"branch_b_oof.csv\")\n\nsummary_path = OUT_DIR / \"branch_b_summary.csv\"\nif summary_path.exists():\n    print(\"\\nbranch_b_summary.csv\")\n    display(pd.read_csv(summary_path))\nelse:\n    print(\"\\nbranch_b_summary.csv -> missing\")\n\nexact_files = sorted(OUT_DIR.glob(\"disease_enriched_sequences_*.csv\"))\nif exact_files:\n    print(f\"\\n{exact_files[0].name}\")\n    display(pd.read_csv(exact_files[0]).head())\nelse:\n    print(\"\\ndisease_enriched_sequences_*.csv -> none found\")\n\ncluster_files = sorted(OUT_DIR.glob(\"cluster_catalog_*.csv\"))\nif cluster_files:\n    print(f\"\\n{cluster_files[0].name}\")\n    display(pd.read_csv(cluster_files[0]).head())\nelse:\n    print(\"\\ncluster_catalog_*.csv -> none found\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-24T11:16:51.059145Z","iopub.execute_input":"2026-07-24T11:16:51.059514Z","iopub.status.idle":"2026-07-24T11:16:51.068584Z","shell.execute_reply.started":"2026-07-24T11:16:51.059460Z","shell.execute_reply":"2026-07-24T11:16:51.068001Z"}},"outputs":[],"execution_count":null}]}