{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":41875,"databundleVersionId":5521661}],"dockerImageVersionId":31328,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 1 — List input files                                  ║\n# ╚══════════════════════════════════════════════════════════════╝\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-03T15:21:00.180712Z","iopub.execute_input":"2026-05-03T15:21:00.181169Z","iopub.status.idle":"2026-05-03T15:21:01.095486Z","shell.execute_reply.started":"2026-05-03T15:21:00.181133Z","shell.execute_reply":"2026-05-03T15:21:01.094673Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 2 — Install packages                                  ║\n# ╚══════════════════════════════════════════════════════════════╝\nimport subprocess, sys\n\ndef pip(pkg):\n    subprocess.run([sys.executable, '-m', 'pip', 'install', '-q', pkg], check=True)\n\npip('fair-esm')\npip('biopython')\npip('goatools')\npip('faiss-cpu')       # faiss-cpu works fine on T4 for 5k proteins\npip('networkx')\npip('plotly')\n\nimport torch\nTORCH = torch.__version__.split('+')[0]\nCUDA  = 'cu' + torch.version.cuda.replace('.', '') if torch.cuda.is_available() else 'cpu'\nprint(f'torch={TORCH}  cuda={CUDA}')\n\n# Install PyTorch Geometric\nsubprocess.run([\n    sys.executable, '-m', 'pip', 'install', '-q',\n    'torch-geometric',\n    '-f', f'https://data.pyg.org/whl/torch-{TORCH}+{CUDA}.html'\n], check=True)\n\nprint('✅ All packages installed')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T16:08:11.185011Z","iopub.execute_input":"2026-05-02T16:08:11.185408Z","iopub.status.idle":"2026-05-02T16:38:50.705143Z","shell.execute_reply.started":"2026-05-02T16:08:11.185377Z","shell.execute_reply":"2026-05-02T16:38:50.704273Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 3 — Imports & Config                                  ║\n# ╚══════════════════════════════════════════════════════════════╝\nimport os, warnings\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch_geometric.data import Data\nfrom torch_geometric.loader import NeighborLoader\nfrom torch_geometric.nn import SAGEConv, JumpingKnowledge\nfrom sklearn.preprocessing import MultiLabelBinarizer\nimport faiss\nfrom Bio import SeqIO\nfrom goatools.obo_parser import GODag\nimport networkx as nx\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nimport seaborn as sns\nimport plotly.graph_objects as go\nimport plotly.express as px\nfrom collections import Counter\nwarnings.filterwarnings('ignore')\n\n# ── Config ────────────────────────────────────────────────────────\nCFG = dict(\n    train_fasta    = '/kaggle/input/competitions/cafa-5-protein-function-prediction/Train/train_sequences.fasta',\n    train_terms    = '/kaggle/input/competitions/cafa-5-protein-function-prediction/Train/train_terms.tsv',\n    test_fasta     = '/kaggle/input/competitions/cafa-5-protein-function-prediction/Test (Targets)/testsuperset.fasta',\n    obo_path       = '/kaggle/input/competitions/cafa-5-protein-function-prediction/Train/go-basic.obo',\n    ia_path        = '/kaggle/input/competitions/cafa-5-protein-function-prediction/IA.txt',\n    emb_cache      = '/kaggle/working/node_features.npy',\n    test_emb_cache = '/kaggle/working/test_features.npy',\n    ontology       = 'BPO',\n    hidden_dim     = 512,\n    num_layers     = 3,\n    dropout        = 0.3,\n    lr             = 3e-4,\n    weight_decay   = 1e-4,\n    batch_size     = 512,\n    epochs         = 60,\n    patience       = 10,\n    max_grad_norm  = 1.0,\n    hier_lambda    = 0.5,\n    graph_k        = 20,\n    graph_thresh   = 0.80,\n    val_frac       = 0.1,\n    seed           = 42,\n    device         = 'cuda' if torch.cuda.is_available() else 'cpu',\n    demo_limit     = 5000,   # set None for full run\n    use_real_esm   = True,   # False = random embeddings for fast testing\n)\n\ntorch.manual_seed(CFG['seed'])\nnp.random.seed(CFG['seed'])\nDEVICE = torch.device(CFG['device'])\nprint(f'Device : {DEVICE}')\nif torch.cuda.is_available():\n    print(f'GPU    : {torch.cuda.get_device_name(0)}')\n    print(f'VRAM   : {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-03T15:21:25.936765Z","iopub.execute_input":"2026-05-03T15:21:25.937261Z","iopub.status.idle":"2026-05-03T15:21:29.623353Z","shell.execute_reply.started":"2026-05-03T15:21:25.937227Z","shell.execute_reply":"2026-05-03T15:21:29.622336Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 4 — Read all data files                               ║\n# ╚══════════════════════════════════════════════════════════════╝\n\n# ── 4a. FASTA sequences ──────────────────────────────────────────\ndef load_fasta(path):\n    return {rec.id: str(rec.seq) for rec in SeqIO.parse(path, 'fasta')}\n\nprint('Reading train sequences...')\ntrain_seqs_full = load_fasta(CFG['train_fasta'])\nprint(f'  → {len(train_seqs_full):,} train proteins')\n\nprint('Reading test sequences...')\ntest_seqs = load_fasta(CFG['test_fasta'])\nprint(f'  → {len(test_seqs):,} test proteins')\n\n# Demo mode subsample\nif CFG['demo_limit']:\n    prot_ids  = list(train_seqs_full.keys())[:CFG['demo_limit']]\n    train_seqs = {k: train_seqs_full[k] for k in prot_ids}\n    print(f'  [Demo] Using {len(train_seqs):,} proteins')\nelse:\n    prot_ids   = list(train_seqs_full.keys())\n    train_seqs = train_seqs_full\nN = len(prot_ids)\n\n# ── 4b. GO term labels ───────────────────────────────────────────\nprint('\\nReading train_terms.tsv...')\ndf_terms = pd.read_csv(CFG['train_terms'], sep='\\t', header=None,\n                       names=['protein_id', 'term', 'aspect'])\nprint(f'  → {len(df_terms):,} rows')\nprint(f'  Aspects: {df_terms[\"aspect\"].value_counts().to_dict()}')\n\ndf_bp        = df_terms[df_terms['aspect'] == CFG['ontology']]\ngo_label_map = df_bp.groupby('protein_id')['term'].apply(list).to_dict()\nlabel_counts = [len(v) for v in go_label_map.values()]\nprint(f'  → {len(go_label_map):,} proteins with BP annotations')\nprint(f'  Labels per protein: mean={np.mean(label_counts):.1f}, '\n      f'median={np.median(label_counts):.0f}, max={max(label_counts)}')\n\n# ── 4c. GO OBO hierarchy ─────────────────────────────────────────\nprint('\\nParsing go-basic.obo...')\ngo_dag   = GODag(CFG['obo_path'], optional_attrs=['relationship'])\nbp_terms = [t for t, node in go_dag.items()\n            if hasattr(node, 'namespace') and node.namespace == 'biological_process']\nprint(f'  → {len(bp_terms):,} BP terms in GO hierarchy')\n\n# ── 4d. IA weights ───────────────────────────────────────────────\nprint('\\nReading IA.txt...')\nia_map = {}\nwith open(CFG['ia_path']) as f:\n    for line in f:\n        parts = line.strip().split('\\t')\n        if len(parts) == 2:\n            try:    ia_map[parts[0]] = float(parts[1])\n            except: pass\nprint(f'  → {len(ia_map):,} terms with IA weights')\nprint(f'  IA range: [{min(ia_map.values()):.3f}, {max(ia_map.values()):.3f}]')\nprint('\\n✅ All files loaded')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T16:48:20.200036Z","iopub.execute_input":"2026-05-02T16:48:20.20035Z","iopub.status.idle":"2026-05-02T16:48:27.879646Z","shell.execute_reply.started":"2026-05-02T16:48:20.200324Z","shell.execute_reply":"2026-05-02T16:48:27.878761Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 5 — EDA Plots                                         ║\n# ╚══════════════════════════════════════════════════════════════╝\nterm_freq = Counter(t for terms in go_label_map.values() for t in terms)\n\nfig, axes = plt.subplots(2, 3, figsize=(18, 10))\nfig.suptitle('CAFA 5 — Data Exploration', fontsize=16, fontweight='bold')\n\n# Plot 1: Sequence length\nseq_lens = [len(s) for s in train_seqs.values()]\naxes[0,0].hist(seq_lens, bins=60, color='#4A90D9', edgecolor='white', lw=0.4)\naxes[0,0].axvline(np.median(seq_lens), color='#E55', lw=2,\n                  label=f'Median={int(np.median(seq_lens))}')\naxes[0,0].set(xlabel='Sequence length (aa)', ylabel='Count',\n              title='Protein sequence lengths')\naxes[0,0].legend()\n\n# Plot 2: GO annotations per protein\nann_counts = [len(go_label_map.get(p, [])) for p in prot_ids]\npos_counts = [c for c in ann_counts if c > 0]\naxes[0,1].hist(pos_counts, bins=50, color='#6BCB77', edgecolor='white', lw=0.4)\naxes[0,1].axvline(np.median(pos_counts), color='#E55', lw=2,\n                  label=f'Median={int(np.median(pos_counts))}')\naxes[0,1].set(xlabel='GO terms per protein', ylabel='Count',\n              title='GO annotations per protein (BP)')\naxes[0,1].legend()\n\n# Plot 3: Aspect distribution\naspect_counts = df_terms[df_terms['aspect'] != 'aspect']['aspect'].value_counts()\naxes[0,2].bar(aspect_counts.index, aspect_counts.values,\n              color=['#4A90D9','#6BCB77','#FF6B6B'])\nfor i, (asp, cnt) in enumerate(aspect_counts.items()):\n    axes[0,2].text(i, cnt + 5000, f'{cnt:,}', ha='center', fontsize=9)\naxes[0,2].set(xlabel='Ontology aspect', ylabel='Count',\n              title='Annotations by ontology')\n\n# Plot 4: Top 20 GO terms\ntop20 = term_freq.most_common(20)\nt_labels, t_vals = zip(*top20)\naxes[1,0].barh(range(20), t_vals, color='#9B59B6')\naxes[1,0].set_yticks(range(20))\naxes[1,0].set_yticklabels(list(t_labels), fontsize=7)\naxes[1,0].invert_yaxis()\naxes[1,0].set(xlabel='Frequency', title='Top 20 GO terms (BP)')\n\n# Plot 5: IA weights\nia_vals = list(ia_map.values())\naxes[1,1].hist(ia_vals, bins=60, color='#FF6B6B', edgecolor='white', lw=0.4)\naxes[1,1].axvline(np.mean(ia_vals), color='#333', lw=2,\n                  label=f'Mean={np.mean(ia_vals):.2f}')\naxes[1,1].set(xlabel='IA weight', ylabel='Count', title='IA weight distribution')\naxes[1,1].legend()\n\n# Plot 6: GO depth\ndepths = [go_dag[t].depth for t in bp_terms if t in go_dag]\naxes[1,2].hist(depths, bins=30, color='#F39C12', edgecolor='white', lw=0.4)\naxes[1,2].set(xlabel='Depth in GO hierarchy', ylabel='Count',\n              title='GO term depth distribution')\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/eda_plots.png', dpi=150, bbox_inches='tight')\nplt.show()\nprint('✅ EDA saved')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T16:48:54.295449Z","iopub.execute_input":"2026-05-02T16:48:54.296105Z","iopub.status.idle":"2026-05-02T16:48:57.465687Z","shell.execute_reply.started":"2026-05-02T16:48:54.296076Z","shell.execute_reply":"2026-05-02T16:48:57.464789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 6 — Build label matrix & GO hierarchy index           ║\n# ╚══════════════════════════════════════════════════════════════╝\nprint('Building label matrix...')\nmlb        = MultiLabelBinarizer()\nlabel_list = [go_label_map.get(pid, []) for pid in prot_ids]\nY          = mlb.fit_transform(label_list)   # [N, C]\ngo_terms   = list(mlb.classes_)\nn_go       = len(go_terms)\nprint(f'  Shape: {Y.shape}  |  Positive rate: {Y.mean():.4f}')\nprint(f'  GO terms in matrix: {n_go:,}')\n\nprint('\\nBuilding GO hierarchy index for loss...')\nterm2idx          = {t: i for i, t in enumerate(go_terms)}\nparent_idx_list   = []\nchild_idx_list    = []\n\nfor term in go_terms:\n    if term not in go_dag:\n        continue\n    node    = go_dag[term]\n    parents = set(node.parents)\n    if hasattr(node, 'relationship'):\n        parents |= set(node.relationship.get('part_of', []))\n    for parent in parents:\n        pid = parent.item_id\n        if pid in term2idx:\n            child_idx_list.append(term2idx[term])\n            parent_idx_list.append(term2idx[pid])\n\nprint(f'  Parent-child pairs: {len(parent_idx_list):,}')\n\nia_weights = torch.tensor(\n    [ia_map.get(t, 1.0) for t in go_terms], dtype=torch.float32\n)\nprint(f'  IA weights shape  : {ia_weights.shape}')\nprint('✅ Labels & hierarchy ready')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T17:05:02.63743Z","iopub.execute_input":"2026-05-02T17:05:02.637767Z","iopub.status.idle":"2026-05-02T17:05:02.973111Z","shell.execute_reply.started":"2026-05-02T17:05:02.637741Z","shell.execute_reply":"2026-05-02T17:05:02.972268Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 7 — Define ESM-2 Embedder                             ║\n# ╚══════════════════════════════════════════════════════════════╝\nimport esm\n\nclass AttentionPooling(nn.Module):\n    \"\"\"\n    Replaces mean-pooling with a learned weighted sum.\n    score_i = w^T tanh(W h_i)  →  alpha = softmax  →  out = Σ alpha_i h_i\n    \"\"\"\n    def __init__(self, embed_dim=1280):\n        super().__init__()\n        self.proj  = nn.Linear(embed_dim, 256, bias=False)\n        self.score = nn.Linear(256, 1,    bias=False)\n\n    def forward(self, x):            # x: [B, L, D]\n        a = torch.tanh(self.proj(x)) # [B, L, 256]\n        a = self.score(a).squeeze(-1)# [B, L]\n        a = torch.softmax(a, dim=-1) # [B, L]\n        return (a.unsqueeze(-1) * x).sum(dim=1)  # [B, D]\n\n\nclass ESM2Embedder:\n    \"\"\"\n    Wraps ESM-2 (650M).\n    - AttentionPooling over residue tokens\n    - Sliding window for sequences > 1022 aa\n    - Output: 1280-d vector per protein\n    \"\"\"\n    MAX_LEN = 1022\n    OVERLAP  = 64\n\n    def __init__(self, device):\n        self.device = device\n        print('Loading ESM-2 (650M) — may take ~1 min...')\n        self.model, self.alphabet = esm.pretrained.esm2_t33_650M_UR50D()\n        self.model = self.model.eval().to(device)\n        self.batch_converter = self.alphabet.get_batch_converter()\n        self.attn_pool = AttentionPooling(1280).to(device)\n        print('  ESM-2 loaded ✅')\n\n    @torch.no_grad()\n    def embed(self, sequences: dict, batch_size=4):\n        ids, seqs = list(sequences.keys()), list(sequences.values())\n        all_embs  = []\n        for i in range(0, len(ids), batch_size):\n            batch_embs = [self._embed_one(s) for s in seqs[i:i+batch_size]]\n            all_embs.append(torch.stack(batch_embs))\n            if (i // batch_size) % 50 == 0:\n                pct = 100 * i / len(ids)\n                print(f'  [{i:>5}/{len(ids)}]  {pct:4.1f}%', end='\\r')\n        print(f'  [{len(ids)}/{len(ids)}]  100.0% — done    ')\n        return torch.cat(all_embs, dim=0).cpu().numpy()\n\n    def _embed_one(self, seq):\n        if len(seq) <= self.MAX_LEN:\n            return self._forward_chunk(seq)\n        step   = self.MAX_LEN - self.OVERLAP\n        chunks = []\n        start  = 0\n        while start < len(seq):\n            end = min(start + self.MAX_LEN, len(seq))\n            chunks.append(self._forward_chunk(seq[start:end]))\n            if end == len(seq): break\n            start += step\n        return torch.stack(chunks).mean(0)\n\n    def _forward_chunk(self, seq):\n        _, _, tokens = self.batch_converter([('p', seq)])\n        tokens = tokens.to(self.device)\n        out    = self.model(tokens, repr_layers=[33], return_contacts=False)\n        reps   = out['representations'][33][0]  # [L+2, 1280]\n        cls_e  = reps[0]\n        eos_e  = reps[-1]\n        mid    = reps[1:-1]                     # [L, 1280]\n        pooled = self.attn_pool(mid.unsqueeze(0)).squeeze(0)\n        return ((cls_e + eos_e + pooled) / 3.0).float()\n\nprint('ESM-2 embedder class defined ✅')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T17:05:35.190655Z","iopub.execute_input":"2026-05-02T17:05:35.1914Z","iopub.status.idle":"2026-05-02T17:05:35.21056Z","shell.execute_reply.started":"2026-05-02T17:05:35.191371Z","shell.execute_reply":"2026-05-02T17:05:35.209735Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 8 — Compute / Load ESM-2 Embeddings  ← FIXED CELL    ║\n# ╚══════════════════════════════════════════════════════════════╝\n#\n# This cell was missing before — it creates the `embeddings`\n# variable that build_protein_graph() needs.\n#\n# ── Option A: load cached ──────────────────────────────────────\nif os.path.exists(CFG['emb_cache']):\n    print(f'Loading cached embeddings → {CFG[\"emb_cache\"]}')\n    embeddings = np.load(CFG['emb_cache'])\n    print(f'  Shape: {embeddings.shape}')\n\n# ── Option B: compute with real ESM-2 ─────────────────────────\nelif CFG['use_real_esm']:\n    print('Computing ESM-2 embeddings...')\n    print('Estimated time: ~25 min for 5 000 proteins on T4 GPU')\n    embedder   = ESM2Embedder(DEVICE)\n    embeddings = embedder.embed(train_seqs, batch_size=4)\n    np.save(CFG['emb_cache'], embeddings)\n    print(f'  Saved → {CFG[\"emb_cache\"]}')\n    del embedder\n    torch.cuda.empty_cache()\n\n# ── Option C: random embeddings (fast pipeline test) ──────────\nelse:\n    print('⚠️  Using RANDOM embeddings — pipeline test only!')\n    print('   Set CFG[\"use_real_esm\"] = True for real results.')\n    rng        = np.random.RandomState(CFG['seed'])\n    embeddings = rng.randn(N, 1280).astype(np.float32)\n\n# ── Sanity check ──────────────────────────────────────────────\nassert embeddings.shape == (N, 1280), \\\n    f'Expected ({N}, 1280), got {embeddings.shape}'\nprint(f'✅ embeddings ready — shape: {embeddings.shape}')\n\n# ── Quick PCA visualisation ────────────────────────────────────\nfrom sklearn.decomposition import PCA\npca    = PCA(n_components=2, random_state=42)\nemb_2d = pca.fit_transform(embeddings[:min(2000, N)])\ncol    = [len(go_label_map.get(p, [])) for p in prot_ids[:min(2000, N)]]\n\nplt.figure(figsize=(8, 6))\nsc = plt.scatter(emb_2d[:,0], emb_2d[:,1], c=col,\n                 cmap='viridis', s=6, alpha=0.7)\nplt.colorbar(sc, label='BP GO annotation count')\nplt.xlabel(f'PC1 ({pca.explained_variance_ratio_[0]*100:.1f}%)')\nplt.ylabel(f'PC2 ({pca.explained_variance_ratio_[1]*100:.1f}%)')\nplt.title('PCA of protein embeddings')\nplt.tight_layout()\nplt.savefig('/kaggle/working/pca_embeddings.png', dpi=150)\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 9 — Build Protein Similarity Graph (FAISS)            ║\n# ╚══════════════════════════════════════════════════════════════╝\ndef build_protein_graph(embeddings, k=20, threshold=0.80):\n    \"\"\"\n    1. L2-normalise embeddings  →  inner product == cosine similarity\n    2. FAISS IndexFlatIP exact kNN search\n    3. Keep edges with cosine ≥ threshold\n    Returns: edge_index [2, E], edge_weight [E]\n    \"\"\"\n    n, d   = embeddings.shape\n    norms  = np.linalg.norm(embeddings, axis=1, keepdims=True) + 1e-12\n    normed = (embeddings / norms).astype(np.float32)\n\n    index = faiss.IndexFlatIP(d)\n    # GPU FAISS — comment out if faiss-cpu only\n    # res   = faiss.StandardGpuResources()\n    # index = faiss.index_cpu_to_gpu(res, 0, index)\n    index.add(normed)\n\n    print(f'  Searching top-{k} neighbours for {n:,} proteins...')\n    sims, nbrs = index.search(normed, k + 1)   # +1 because idx 0 = self\n\n    src, dst, wts = [], [], []\n    for i in range(n):\n        for j, sim in zip(nbrs[i], sims[i]):\n            if j == i or j < 0 or sim < threshold:\n                continue\n            src.append(i); dst.append(int(j)); wts.append(float(sim))\n\n    edge_index  = torch.tensor([src, dst], dtype=torch.long)\n    edge_weight = torch.tensor(wts,        dtype=torch.float32)\n    return edge_index, edge_weight\n\n\nprint('Building protein similarity graph...')\nedge_index, edge_weight = build_protein_graph(\n    embeddings,\n    k         = CFG['graph_k'],\n    threshold = CFG['graph_thresh'],\n)\n\nn_nodes = N\nn_edges = edge_index.shape[1]\navg_deg = n_edges / n_nodes\n\nprint(f'\\n📊 Graph Statistics')\nprint(f'  Nodes      : {n_nodes:,}')\nprint(f'  Edges      : {n_edges:,}')\nprint(f'  Avg degree : {avg_deg:.2f}')\nif n_edges > 0:\n    print(f'  Sim range  : [{edge_weight.min():.4f}, {edge_weight.max():.4f}]')\nelse:\n    print('  ⚠️  0 edges — lower graph_thresh or check embeddings')\nprint('✅ Graph built')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-02T17:06:44.61649Z","iopub.execute_input":"2026-05-02T17:06:44.616905Z","iopub.status.idle":"2026-05-02T17:06:44.629118Z","shell.execute_reply.started":"2026-05-02T17:06:44.616876Z","shell.execute_reply":"2026-05-02T17:06:44.628261Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 10 — Graph Visualization (static matplotlib)          ║\n# ╚══════════════════════════════════════════════════════════════╝\nVIZ_N = 300    # subgraph size for clarity\n\n# ── Build NetworkX subgraph ───────────────────────────────────────\nmask   = (edge_index[0] < VIZ_N) & (edge_index[1] < VIZ_N)\nei_sub = edge_index[:, mask]\new_sub = edge_weight[mask] if n_edges > 0 else torch.tensor([])\n\nG = nx.Graph()\nG.add_nodes_from(range(VIZ_N))\nfor s, d, w in zip(ei_sub[0].tolist(), ei_sub[1].tolist(),\n                   ew_sub.tolist() if len(ew_sub) else []):\n    G.add_edge(s, d, weight=w)\n\nnode_ann = [len(go_label_map.get(prot_ids[i], [])) for i in range(VIZ_N)]\nnode_deg = [G.degree(i) for i in range(VIZ_N)]\nprint(f'Subgraph: {G.number_of_nodes()} nodes | {G.number_of_edges()} edges')\n\nprint('Computing spring layout...')\npos = nx.spring_layout(G, seed=CFG['seed'], k=0.4, iterations=60)\n\nfig, axes = plt.subplots(2, 2, figsize=(18, 14))\nfig.suptitle(f'Protein Similarity Graph  (n={VIZ_N} subgraph)',\n             fontsize=16, fontweight='bold')\n\n# ── Plot 1: Colored by degree ──────────────────────────────────\nax = axes[0, 0]\nnx.draw_networkx_edges(G, pos, ax=ax, alpha=0.18, width=0.5, edge_color='#aaa')\nsc1 = nx.draw_networkx_nodes(G, pos, ax=ax,\n    node_color=node_deg, cmap='plasma',\n    node_size=[8 + d*14 for d in node_deg], alpha=0.85)\nplt.colorbar(sc1, ax=ax, label='Node degree')\nax.set_title('Colored by node degree'); ax.axis('off')\n\n# ── Plot 2: Colored by GO annotation count ─────────────────────\nax = axes[0, 1]\nnx.draw_networkx_edges(G, pos, ax=ax, alpha=0.15, width=0.4, edge_color='#aaa')\nsc2 = nx.draw_networkx_nodes(G, pos, ax=ax,\n    node_color=node_ann, cmap='YlOrRd',\n    node_size=30, alpha=0.85)\nplt.colorbar(sc2, ax=ax, label='BP GO annotations')\nax.set_title('Colored by GO annotation count'); ax.axis('off')\n\n# ── Plot 3: Edge weight distribution ──────────────────────────\nax = axes[1, 0]\nif n_edges > 0:\n    ax.hist(edge_weight.numpy(), bins=60, color='#4A90D9',\n            edgecolor='white', linewidth=0.4)\n    ax.axvline(CFG['graph_thresh'], color='red', lw=2, ls='--',\n               label=f'threshold={CFG[\"graph_thresh\"]}')\n    ax.legend()\nax.set(xlabel='Cosine similarity', ylabel='Edge count',\n       title='Edge weight distribution')\n\n# ── Plot 4: Degree distribution log-log ───────────────────────\nax = axes[1, 1]\nall_deg    = [d for _, d in G.degree() if d > 0]\ndeg_cnt    = Counter(all_deg)\ndegs, cnts = zip(*sorted(deg_cnt.items())) if deg_cnt else ([0],[0])\nax.loglog(degs, cnts, 'o-', color='#6BCB77', ms=4, lw=1.5)\nax.set(xlabel='Degree (log)', ylabel='Count (log)',\n       title='Degree distribution (log-log)')\nax.grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/graph_viz_static.png', dpi=150, bbox_inches='tight')\nplt.show()\nprint('✅ Static graph visualization saved')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 11 — Interactive Plotly Graph                         ║\n# ╚══════════════════════════════════════════════════════════════╝\nINT_N = 200\n\nmask_i = (edge_index[0] < INT_N) & (edge_index[1] < INT_N)\nei_i   = edge_index[:, mask_i]\n\nG_i = nx.Graph()\nG_i.add_nodes_from(range(INT_N))\nfor s, d in zip(ei_i[0].tolist(), ei_i[1].tolist()):\n    G_i.add_edge(s, d)\n\npos_i = nx.spring_layout(G_i, seed=CFG['seed'], k=0.5, iterations=80)\n\n# Edge traces\nex, ey = [], []\nfor s, d in G_i.edges():\n    x0,y0 = pos_i[s]; x1,y1 = pos_i[d]\n    ex += [x0, x1, None]; ey += [y0, y1, None]\n\nedge_tr = go.Scatter(x=ex, y=ey, mode='lines',\n    line=dict(width=0.6, color='#bbb'), hoverinfo='none')\n\n# Node traces\nnx_ = [pos_i[n][0] for n in range(INT_N)]\nny_ = [pos_i[n][1] for n in range(INT_N)]\nnc  = [G_i.degree(n) for n in range(INT_N)]\nann_i = [len(go_label_map.get(prot_ids[n], [])) for n in range(INT_N)]\n\nhover = [\n    f'<b>{prot_ids[n]}</b><br>Degree: {G_i.degree(n)}<br>GO annotations: {ann_i[n]}'\n    for n in range(INT_N)\n]\nnode_tr = go.Scatter(x=nx_, y=ny_, mode='markers',\n    hoverinfo='text', text=hover,\n    marker=dict(\n        size=[6 + c*2 for c in nc], color=nc, colorscale='Plasma',\n        colorbar=dict(title='Degree', thickness=14),\n        line=dict(width=0.5, color='white'),\n    ))\n\nfig_int = go.Figure(\n    data=[edge_tr, node_tr],\n    layout=go.Layout(\n        title='<b>Protein Similarity Graph</b> — hover for details',\n        titlefont_size=14, showlegend=False, hovermode='closest',\n        margin=dict(b=20, l=5, r=5, t=50), height=600,\n        xaxis=dict(showgrid=False, zeroline=False, showticklabels=False),\n        yaxis=dict(showgrid=False, zeroline=False, showticklabels=False),\n        template='plotly_white',\n    )\n)\nfig_int.write_html('/kaggle/working/graph_interactive.html')\nfig_int.show()\nprint('✅ Interactive graph saved → graph_interactive.html')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 12 — GO Hierarchy Subgraph Visualization              ║\n# ╚══════════════════════════════════════════════════════════════╝\ntop_terms = [t for t, _ in term_freq.most_common(20) if t in go_dag]\n\ngo_viz_nodes = set(top_terms)\nfor term in top_terms:\n    for p in go_dag[term].parents:\n        go_viz_nodes.add(p.item_id)\n        for pp in go_dag[p.item_id].parents:\n            go_viz_nodes.add(pp.item_id)\n\nG_go = nx.DiGraph()\nfor t in go_viz_nodes:\n    if t not in go_dag: continue\n    G_go.add_node(t, depth=go_dag[t].depth)\n    for p in go_dag[t].parents:\n        pid = p.item_id\n        if pid in go_viz_nodes:\n            G_go.add_edge(pid, t)\n\nprint(f'GO subgraph: {G_go.number_of_nodes()} nodes, {G_go.number_of_edges()} edges')\ntry:\n    pos_go = nx.nx_agraph.graphviz_layout(G_go, prog='dot')\nexcept Exception:\n    pos_go = nx.spring_layout(G_go, seed=42)\n\ndepths_go = [G_go.nodes[n].get('depth', 0) for n in G_go.nodes()]\nis_top    = {n: n in set(top_terms) for n in G_go.nodes()}\nsizes     = [200 if is_top[n] else 60 for n in G_go.nodes()]\n\nfig, ax = plt.subplots(figsize=(16, 10))\nnx.draw_networkx_edges(G_go, pos_go, ax=ax, arrows=True,\n    arrowsize=10, edge_color='#888', width=0.8, alpha=0.6,\n    connectionstyle='arc3,rad=0.1')\nsc = nx.draw_networkx_nodes(G_go, pos_go, ax=ax,\n    node_color=depths_go, cmap='RdYlGn_r', node_size=sizes, alpha=0.9)\ntop_labels = {n: go_dag[n].name[:22] for n in G_go.nodes() if is_top[n] and n in go_dag}\nnx.draw_networkx_labels(G_go, pos_go, labels=top_labels, ax=ax, font_size=6)\nplt.colorbar(sc, ax=ax, label='GO term depth')\nax.set_title('GO Hierarchy (top-20 BP terms + 2 ancestor levels)', fontsize=13)\nax.axis('off')\nplt.tight_layout()\nplt.savefig('/kaggle/working/go_hierarchy.png', dpi=150, bbox_inches='tight')\nplt.show()\nprint('✅ GO hierarchy visualization saved')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 13 — Model Architecture                               ║\n# ╚══════════════════════════════════════════════════════════════╝\nclass ProteinGNN(nn.Module):\n    \"\"\"\n    GraphSAGE × num_layers  +  Jumping Knowledge (concat)  +  MLP head\n\n    GraphSAGE update rule:\n        h_v^{l+1} = σ( W · CONCAT[ h_v^l,  MEAN_{u∈N(v)} h_u^l ] )\n\n    JK concat:\n        h_v^JK = CONCAT[ h_v^1 || h_v^2 || h_v^3 ]\n        → all layers contribute; each protein uses its best receptive field\n    \"\"\"\n    def __init__(self, in_ch, hidden, out_ch, n_layers=3, dropout=0.3):\n        super().__init__()\n        self.n_layers = n_layers\n        self.drop     = dropout\n\n        # Input projection: 1280 → hidden\n        self.proj = nn.Sequential(\n            nn.Linear(in_ch, hidden),\n            nn.BatchNorm1d(hidden),\n            nn.GELU(),\n            nn.Dropout(dropout),\n        )\n\n        # GraphSAGE layers\n        self.convs = nn.ModuleList(\n            [SAGEConv(hidden, hidden, aggr='mean') for _ in range(n_layers)]\n        )\n\n        # Jumping Knowledge\n        self.jk    = JumpingKnowledge(mode='cat', channels=hidden, num_layers=n_layers)\n        jk_dim     = hidden * n_layers\n\n        # MLP classifier\n        self.head = nn.Sequential(\n            nn.Linear(jk_dim, hidden),\n            nn.BatchNorm1d(hidden),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden, hidden // 2),\n            nn.BatchNorm1d(hidden // 2),\n            nn.GELU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden // 2, out_ch),\n        )\n        for m in self.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.xavier_uniform_(m.weight)\n                if m.bias is not None: nn.init.zeros_(m.bias)\n\n    def forward(self, x, edge_index):\n        h    = self.proj(x)\n        outs = []\n        for conv in self.convs:\n            h = F.gelu(conv(h, edge_index))\n            h = F.dropout(h, p=self.drop, training=self.training)\n            outs.append(h)\n        h = self.jk(outs)      # [N, hidden*n_layers]\n        return self.head(h)    # [N, n_go]\n\n\nmodel = ProteinGNN(\n    in_ch   = 1280,\n    hidden  = CFG['hidden_dim'],\n    out_ch  = n_go,\n    n_layers= CFG['num_layers'],\n    dropout = CFG['dropout'],\n).to(DEVICE)\n\nn_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'✅ ProteinGNN')\nprint(f'   Params    : {n_params:,}')\nprint(f'   GO terms  : {n_go:,}')\nprint(f'   Hidden    : {CFG[\"hidden_dim\"]}')\nprint(model)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 14 — Hybrid Loss                                      ║\n# ╚══════════════════════════════════════════════════════════════╝\nclass HybridLoss(nn.Module):\n    \"\"\"\n    L = BCE  +  IA-weighted BCE  +  λ * Hierarchical Consistency\n\n    Hierarchical loss:\n        For every (parent, child) pair from GO DAG:\n            penalty = max(0, P(child) - P(parent))^2\n\n    GO hierarchy is used ONLY here → ZERO data leakage.\n    \"\"\"\n    def __init__(self, ia_weights, parent_idx, child_idx, lam=0.5):\n        super().__init__()\n        self.register_buffer('ia_w',    ia_weights)\n        self.register_buffer('par_idx', torch.tensor(parent_idx, dtype=torch.long))\n        self.register_buffer('chi_idx', torch.tensor(child_idx,  dtype=torch.long))\n        self.lam = lam\n\n    def forward(self, logits, targets):\n        probs = torch.sigmoid(logits)\n\n        # 1. Standard BCE\n        bce = F.binary_cross_entropy_with_logits(\n            logits, targets, reduction='mean'\n        )\n        # 2. IA-weighted BCE (upweight rare/specific GO terms)\n        ia_bce = F.binary_cross_entropy_with_logits(\n            logits, targets, pos_weight=self.ia_w, reduction='mean'\n        )\n        # 3. Hierarchical consistency: P(child) <= P(parent)\n        if len(self.par_idx) > 0:\n            viol = torch.clamp(probs[:, self.chi_idx] - probs[:, self.par_idx], min=0.0)\n            hier = viol.pow(2).mean()\n        else:\n            hier = torch.zeros(1, device=logits.device).squeeze()\n\n        total = bce + ia_bce + self.lam * hier\n        return {'loss': total, 'bce': bce, 'ia_bce': ia_bce, 'hier': hier}\n\n\ncriterion = HybridLoss(\n    ia_weights = ia_weights,\n    parent_idx = parent_idx_list,\n    child_idx  = child_idx_list,\n    lam        = CFG['hier_lambda'],\n)\nprint(f'✅ HybridLoss  λ={CFG[\"hier_lambda\"]}  pairs={len(parent_idx_list):,}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 15 — Build PyG Data & NeighborLoaders                 ║\n# ╚══════════════════════════════════════════════════════════════╝\nx     = torch.tensor(embeddings, dtype=torch.float32)\ny     = torch.tensor(Y,          dtype=torch.float32)\ngraph = Data(x=x, edge_index=edge_index, edge_attr=edge_weight, y=y)\n\n# Train / val split\nperm       = torch.randperm(N)\nval_size   = int(N * CFG['val_frac'])\nval_mask   = torch.zeros(N, dtype=torch.bool)\ntrain_mask = torch.zeros(N, dtype=torch.bool)\nval_mask[perm[:val_size]]   = True\ntrain_mask[perm[val_size:]] = True\ngraph.train_mask = train_mask\ngraph.val_mask   = val_mask\n\ngraph = graph.to(DEVICE)\n\ntrain_loader = NeighborLoader(\n    graph,\n    num_neighbors = [10, 10, 10],\n    batch_size    = CFG['batch_size'],\n    input_nodes   = graph.train_mask,\n    shuffle       = True,\n)\nval_loader = NeighborLoader(\n    graph,\n    num_neighbors = [10, 10, 10],\n    batch_size    = CFG['batch_size'],\n    input_nodes   = graph.val_mask,\n    shuffle       = False,\n)\n\nprint(f'✅ Graph data ready')\nprint(f'   Train nodes : {train_mask.sum().item():,}')\nprint(f'   Val nodes   : {val_mask.sum().item():,}')\nprint(f'   Feature dim : {graph.x.shape[1]}')\nprint(f'   Label dim   : {graph.y.shape[1]}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 16 — Training Loop                                    ║\n# ╚══════════════════════════════════════════════════════════════╝\noptimizer = AdamW(model.parameters(),\n                  lr=CFG['lr'], weight_decay=CFG['weight_decay'])\nscheduler = CosineAnnealingLR(optimizer,\n                               T_max=CFG['epochs'], eta_min=CFG['lr']*0.01)\n\nhistory = {k: [] for k in\n           ['train_loss','val_loss','train_bce','val_bce','train_hier','val_hier']}\n\nbest_val   = float('inf')\npatience_c = 0\nbest_state = None\n\nprint(f'Training for up to {CFG[\"epochs\"]} epochs (patience={CFG[\"patience\"]})\\n')\nprint(f'{\"Epoch\":>6}  {\"TrLoss\":>8}  {\"BCE\":>8}  {\"Hier\":>8}  {\"ValLoss\":>8}')\nprint('-' * 52)\n\nfor epoch in range(1, CFG['epochs'] + 1):\n\n    # ── Train ──────────────────────────────────────────────────\n    model.train()\n    tr = {k: 0.0 for k in ['loss','bce','ia_bce','hier']}\n    nb = 0\n    for batch in train_loader:\n        batch   = batch.to(DEVICE)\n        logits  = model(batch.x, batch.edge_index)[:batch.batch_size]\n        targets = batch.y[:batch.batch_size].float()\n        ld = criterion(logits, targets)\n        optimizer.zero_grad()\n        ld['loss'].backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), CFG['max_grad_norm'])\n        optimizer.step()\n        for k in tr: tr[k] += ld[k].item()\n        nb += 1\n    tr = {k: v/nb for k, v in tr.items()}\n\n    # ── Validate ───────────────────────────────────────────────\n    model.eval()\n    vl = {k: 0.0 for k in ['loss','bce','ia_bce','hier']}\n    nb = 0\n    with torch.no_grad():\n        for batch in val_loader:\n            batch   = batch.to(DEVICE)\n            logits  = model(batch.x, batch.edge_index)[:batch.batch_size]\n            targets = batch.y[:batch.batch_size].float()\n            ld = criterion(logits, targets)\n            for k in vl: vl[k] += ld[k].item()\n            nb += 1\n    vl = {k: v/nb for k, v in vl.items()}\n    scheduler.step()\n\n    for k in ['loss','bce','hier']:\n        history[f'train_{k}'].append(tr[k])\n        history[f'val_{k}'].append(vl[k])\n\n    if epoch % 5 == 0 or epoch == 1:\n        print(f'{epoch:>6}  {tr[\"loss\"]:>8.4f}  {tr[\"bce\"]:>8.4f}  '\n              f'{tr[\"hier\"]:>8.4f}  {vl[\"loss\"]:>8.4f}')\n\n    # Early stopping\n    if vl['loss'] < best_val - 1e-4:\n        best_val   = vl['loss']\n        patience_c = 0\n        best_state = {k: v.clone() for k, v in model.state_dict().items()}\n    else:\n        patience_c += 1\n        if patience_c >= CFG['patience']:\n            print(f'\\n⏹  Early stopping at epoch {epoch}')\n            break\n\nmodel.load_state_dict(best_state)\ntorch.save(best_state, '/kaggle/working/best_model.pt')\nprint(f'\\n✅ Training done  |  Best val loss: {best_val:.4f}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 17 — Training Curves                                  ║\n# ╚══════════════════════════════════════════════════════════════╝\nep  = range(1, len(history['train_loss']) + 1)\nfig, axes = plt.subplots(1, 3, figsize=(18, 5))\nfig.suptitle('Training Curves', fontsize=14, fontweight='bold')\n\nfor ax, key, title in zip(\n    axes,\n    ['loss', 'bce', 'hier'],\n    ['Total loss', 'BCE loss', 'Hierarchical loss']\n):\n    ax.plot(ep, history[f'train_{key}'], label='Train', color='#4A90D9', lw=2)\n    ax.plot(ep, history[f'val_{key}'],   label='Val',   color='#E55',   lw=2, ls='--')\n    ax.set(xlabel='Epoch', ylabel='Loss', title=title)\n    ax.legend(); ax.grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/training_curves.png', dpi=150)\nplt.show()\nprint('✅ Training curves saved')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 18 — Threshold Tuning (Fmax)                          ║\n# ╚══════════════════════════════════════════════════════════════╝\nmodel.eval()\nall_probs, all_labels = [], []\n\nwith torch.no_grad():\n    for batch in val_loader:\n        batch = batch.to(DEVICE)\n        logits  = model(batch.x, batch.edge_index)[:batch.batch_size]\n        targets = batch.y[:batch.batch_size]\n        all_probs.append(torch.sigmoid(logits).cpu().numpy())\n        all_labels.append(targets.cpu().numpy())\n\nprobs  = np.vstack(all_probs)\nlabels = np.vstack(all_labels)\n\nthresholds = np.arange(0.05, 0.80, 0.025)\nfmax_sc, prec_sc, rec_sc = [], [], []\n\nfor t in thresholds:\n    preds = (probs >= t).astype(float)\n    tp = (preds * labels).sum()\n    fp = (preds * (1 - labels)).sum()\n    fn = ((1 - preds) * labels).sum()\n    p  = tp / (tp + fp + 1e-12)\n    r  = tp / (tp + fn + 1e-12)\n    f  = 2*p*r / (p + r + 1e-12)\n    fmax_sc.append(f); prec_sc.append(p); rec_sc.append(r)\n\nbest_idx  = int(np.argmax(fmax_sc))\nbest_t    = float(thresholds[best_idx])\nbest_fmax = fmax_sc[best_idx]\n\nprint(f'Best threshold : {best_t:.3f}')\nprint(f'Fmax           : {best_fmax:.4f}')\nprint(f'Precision      : {prec_sc[best_idx]:.4f}')\nprint(f'Recall         : {rec_sc[best_idx]:.4f}')\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\naxes[0].plot(thresholds, fmax_sc, 'b-o', ms=4, label='Fmax')\naxes[0].plot(thresholds, prec_sc, 'g--', lw=1.5, label='Precision')\naxes[0].plot(thresholds, rec_sc,  'r--', lw=1.5, label='Recall')\naxes[0].axvline(best_t, color='orange', lw=2, ls=':', label=f'Best t={best_t:.2f}')\naxes[0].set(xlabel='Threshold', ylabel='Score', title='Metrics vs Threshold')\naxes[0].legend(); axes[0].grid(True, alpha=0.3)\n\naxes[1].plot(rec_sc, prec_sc, 'b-o', ms=4)\naxes[1].scatter([rec_sc[best_idx]], [prec_sc[best_idx]],\n                color='red', s=100, zorder=5, label=f'Fmax={best_fmax:.3f}')\naxes[1].set(xlabel='Recall', ylabel='Precision', title='Precision-Recall Curve')\naxes[1].legend(); axes[1].grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/fmax_curves.png', dpi=150)\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 19 — Inference on Test Set                            ║\n# ╚══════════════════════════════════════════════════════════════╝\n\n# ── Embed test proteins ──────────────────────────────────────────\nif os.path.exists(CFG['test_emb_cache']):\n    print('Loading cached test embeddings...')\n    test_embs = np.load(CFG['test_emb_cache'])\nelif CFG['use_real_esm']:\n    print('Computing test embeddings with ESM-2...')\n    embedder  = ESM2Embedder(DEVICE)\n    test_embs = embedder.embed(test_seqs, batch_size=4)\n    np.save(CFG['test_emb_cache'], test_embs)\n    del embedder; torch.cuda.empty_cache()\nelse:\n    print('⚠️  Random test embeddings (pipeline test)')\n    test_embs = np.random.randn(len(test_seqs), 1280).astype(np.float32)\n\ntest_ids      = list(test_seqs.keys())\nn_train_nodes = N\nprint(f'Test embeddings: {test_embs.shape}')\n\n# ── Build full graph: train + test ───────────────────────────────\nall_embs    = np.vstack([embeddings, test_embs])\nn_all       = len(all_embs)\nei_full, ew_full = build_protein_graph(\n    all_embs, k=CFG['graph_k'], threshold=CFG['graph_thresh']\n)\n\ntest_mask_inf = torch.zeros(n_all, dtype=torch.bool)\ntest_mask_inf[n_train_nodes:] = True\n\ng_full = Data(\n    x          = torch.tensor(all_embs, dtype=torch.float32),\n    edge_index = ei_full,\n    edge_attr  = ew_full,\n).to(DEVICE)\n\ninfer_loader = NeighborLoader(\n    g_full, num_neighbors=[10, 10, 10],\n    batch_size=512, input_nodes=test_mask_inf, shuffle=False,\n)\n\n# ── Forward pass ─────────────────────────────────────────────────\nmodel.eval()\nall_test_probs, all_test_nids = [], []\nwith torch.no_grad():\n    for batch in infer_loader:\n        batch  = batch.to(DEVICE)\n        logits = model(batch.x, batch.edge_index)[:batch.batch_size]\n        all_test_probs.append(torch.sigmoid(logits).cpu().numpy())\n        all_test_nids.append(batch.n_id[:batch.batch_size].cpu().numpy())\n\ntest_probs = np.vstack(all_test_probs)\ntest_nids  = np.concatenate(all_test_nids)\n\n# ── Build prediction DataFrame ────────────────────────────────────\nrows = []\nfor nid, pvec in zip(test_nids, test_probs):\n    pid = test_ids[int(nid) - n_train_nodes]\n    for j, p in enumerate(pvec):\n        if p >= best_t:\n            rows.append({'protein_id': pid, 'go_term': go_terms[j],\n                         'score': round(float(p), 4)})\n\ndf_pred = pd.DataFrame(rows)\ndf_pred.to_csv('/kaggle/working/cafa5_predictions.tsv', sep='\\t', index=False)\nprint(f'\\n✅ Predictions saved')\nprint(f'   Total rows      : {len(df_pred):,}')\nprint(f'   Unique proteins : {df_pred[\"protein_id\"].nunique():,}')\nprint(f'   Unique GO terms : {df_pred[\"go_term\"].nunique():,}')\ndf_pred.head(10)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ╔══════════════════════════════════════════════════════════════╗\n# ║  CELL 20 — Final Prediction Analysis Plots                  ║\n# ╚══════════════════════════════════════════════════════════════╝\nfig, axes = plt.subplots(2, 2, figsize=(16, 11))\nfig.suptitle('Prediction Analysis', fontsize=15, fontweight='bold')\n\n# 1. Predictions per protein\nppp = df_pred.groupby('protein_id').size()\naxes[0,0].hist(ppp, bins=50, color='#4A90D9', edgecolor='white')\naxes[0,0].axvline(ppp.median(), color='red', lw=2,\n                  label=f'Median={ppp.median():.0f}')\naxes[0,0].set(xlabel='Predicted GO terms per protein', ylabel='Count',\n              title='Predictions per protein')\naxes[0,0].legend()\n\n# 2. Score distribution\naxes[0,1].hist(df_pred['score'], bins=60, color='#6BCB77', edgecolor='white')\naxes[0,1].axvline(best_t, color='red', lw=2, ls='--',\n                  label=f'threshold={best_t:.2f}')\naxes[0,1].set(xlabel='Prediction score', ylabel='Count',\n              title='Score distribution')\naxes[0,1].legend()\n\n# 3. Top predicted GO terms\ntop_p = df_pred['go_term'].value_counts().head(15)\naxes[1,0].barh(range(15), top_p.values, color='#9B59B6')\naxes[1,0].set_yticks(range(15))\naxes[1,0].set_yticklabels(top_p.index.tolist(), fontsize=7)\naxes[1,0].invert_yaxis()\naxes[1,0].set(xlabel='Times predicted', title='Top 15 predicted GO terms')\n\n# 4. TP / FP / FN on val sample (first 30 terms)\ns_idx  = np.random.choice(len(probs), min(200, len(probs)), replace=False)\nsp     = (probs[s_idx, :30] >= best_t).astype(int)\nsl     = labels[s_idx, :30].astype(int)\ntp_m   = (sp * sl).sum(0)\nfp_m   = (sp * (1 - sl)).sum(0)\nfn_m   = ((1-sp) * sl).sum(0)\nx_pts  = np.arange(30)\nw      = 0.28\naxes[1,1].bar(x_pts-w, tp_m, w, label='TP', color='#27AE60')\naxes[1,1].bar(x_pts,   fp_m, w, label='FP', color='#E74C3C')\naxes[1,1].bar(x_pts+w, fn_m, w, label='FN', color='#F39C12')\naxes[1,1].set(xlabel='GO term index (first 30)', ylabel='Count',\n              title='TP/FP/FN per term (val sample)')\naxes[1,1].legend()\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/prediction_analysis.png', dpi=150, bbox_inches='tight')\nplt.show()\n\n# ── Summary of all saved files ────────────────────────────────────\nprint('\\n🎉  All done!  Output files:')\nfiles = [\n    'eda_plots.png', 'pca_embeddings.png',\n    'graph_viz_static.png', 'graph_interactive.html', 'go_hierarchy.png',\n    'training_curves.png', 'fmax_curves.png',\n    'prediction_analysis.png',\n    'best_model.pt', 'cafa5_predictions.tsv',\n    'node_features.npy',\n]\nfor f in files:\n    path = f'/kaggle/working/{f}'\n    if os.path.exists(path):\n        kb = os.path.getsize(path) / 1024\n        print(f'  ✅  {f:<40}  {kb:>8.0f} KB')\n    else:\n        print(f'  ⬜  {f:<40}  (not generated yet)')","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}