{"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},{"sourceType":"datasetVersion","sourceId":15950910,"datasetId":10229225,"databundleVersionId":16909808}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q transformers accelerate biopython\n \nimport torch, transformers\nprint(f\"torch        : {torch.__version__}\")\nprint(f\"transformers : {transformers.__version__}\")\nprint(f\"CUDA         : {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"GPU          : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM         : \"\n          f\"{torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\nprint(\"✓ Ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-27T13:20:31.321095Z","iopub.execute_input":"2026-04-27T13:20:31.321808Z","iopub.status.idle":"2026-04-27T13:20:34.937031Z","shell.execute_reply.started":"2026-04-27T13:20:31.321755Z","shell.execute_reply":"2026-04-27T13:20:34.936064Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport json\nimport numpy as np\nimport pandas as pd\nimport gc\n\n# ── Paths ──────────────────────────────────────────────────────────\nPREP_DIR   = \"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed\"   # output من الـ preprocessing\nOUTPUT_DIR = \"/kaggle/working\"\n\nTRAIN_IDS_FILE   = f\"{PREP_DIR}/train_protein_ids.txt\"\nTEST_IDS_FILE    = f\"{PREP_DIR}/test_protein_ids.txt\"\nTRAIN_FASTA      = f\"{PREP_DIR}/train_sequences_clean.fasta\"\nTEST_FASTA       = f\"{PREP_DIR}/test_sequences_clean.fasta\"\nLABEL_MATRIX     = f\"{PREP_DIR}/label_matrix.npy\"\nTRAIN_IDX_FILE   = f\"{PREP_DIR}/train_indices.npy\"\nVAL_IDX_FILE     = f\"{PREP_DIR}/val_indices.npy\"\nIA_WEIGHTS_FILE  = f\"{PREP_DIR}/ia_weights.npy\"\nTERMS_FILE       = f\"{PREP_DIR}/go_terms_list.txt\"\nTERM_TO_IDX_FILE = f\"{PREP_DIR}/term_to_idx.json\"\n\nGRAPH_PATH = f\"{OUTPUT_DIR}/protein_graph.pt\"\nCKPT_PATH  = f\"{OUTPUT_DIR}/best_model.pt\"\nSUB_PATH   = f\"{OUTPUT_DIR}/submission.tsv\"\n\n# ── Model Config ───────────────────────────────────────────────────\nESM2_MODEL    = \"facebook/esm2_t33_650M_UR50D\"\nESM2_DIM      = 1280       # output dim من ESM-2\nHIDDEN_DIM    = 512        # hidden dim في الـ GNN\nNUM_GNN_LAYERS = 3\nDROPOUT       = 0.3\nMAX_SEQ_LEN   = 1022       # ESM-2 limit\nWINDOW_OVERLAP = 256       # overlap للـ sliding window\n\n# ── Training Config ────────────────────────────────────────────────\nBATCH_SIZE    = 16         # عدد البروتينات في كل batch\nEPOCHS        = 80\nLR            = 3e-4\nWEIGHT_DECAY  = 1e-4\nPATIENCE      = 12         # early stopping\nSIM_THRESHOLD = 0.80       # لبناء الـ graph\n\n# تحقق من الملفات\nprint(\"Checking files...\")\nall_ok = True\nfor name, path in [\n    (\"train_ids\",    TRAIN_IDS_FILE),\n    (\"test_ids\",     TEST_IDS_FILE),\n    (\"train_fasta\",  TRAIN_FASTA),\n    (\"test_fasta\",   TEST_FASTA),\n    (\"labels\",       LABEL_MATRIX),\n    (\"train_idx\",    TRAIN_IDX_FILE),\n    (\"val_idx\",      VAL_IDX_FILE),\n    (\"ia_weights\",   IA_WEIGHTS_FILE),\n    (\"terms\",        TERMS_FILE),\n]:\n    exists = os.path.exists(path)\n    print(f\"  {'✓' if exists else '✗'} {name}\")\n    if not exists:\n        all_ok = False\nprint(f\"\\n{'✓ All files OK' if all_ok else '✗ Run preprocessing first!'}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-27T13:21:29.665915Z","iopub.execute_input":"2026-04-27T13:21:29.666661Z","iopub.status.idle":"2026-04-27T13:21:29.700894Z","shell.execute_reply.started":"2026-04-27T13:21:29.666623Z","shell.execute_reply":"2026-04-27T13:21:29.700282Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from Bio import SeqIO\n\n# ── Protein IDs ────────────────────────────────────────────────────\nprint(\"Loading protein IDs...\")\nwith open(TRAIN_IDS_FILE) as f:\n    train_proteins = [l.strip() for l in f if l.strip()]\nwith open(TEST_IDS_FILE) as f:\n    test_proteins  = [l.strip() for l in f if l.strip()]\n\nN_TRAIN = len(train_proteins)\nN_TEST  = len(test_proteins)\nprint(f\"  Train: {N_TRAIN:,}\")\nprint(f\"  Test : {N_TEST:,}\")\n\n# ── Sequences ──────────────────────────────────────────────────────\nprint(\"Loading sequences...\")\ntrain_seqs = {}\nfor record in SeqIO.parse(TRAIN_FASTA, \"fasta\"):\n    pid = record.id.split(\"|\")[1] if \"|\" in record.id else record.id\n    train_seqs[pid] = str(record.seq)\n\ntest_seqs = {}\nfor record in SeqIO.parse(TEST_FASTA, \"fasta\"):\n    pid = record.id.split(\"|\")[1] if \"|\" in record.id else record.id\n    test_seqs[pid] = str(record.seq)\n\nprint(f\"  Train seqs: {len(train_seqs):,}\")\nprint(f\"  Test  seqs: {len(test_seqs):,}\")\n\n# ── Labels ─────────────────────────────────────────────────────────\nprint(\"Loading labels...\")\nlabel_matrix  = np.load(LABEL_MATRIX)\ntrain_indices = np.load(TRAIN_IDX_FILE)\nval_indices   = np.load(VAL_IDX_FILE)\nia_weights    = np.load(IA_WEIGHTS_FILE)\n\nN, M = label_matrix.shape\nprint(f\"  Label matrix : {N:,} × {M:,}\")\nprint(f\"  Train split  : {len(train_indices):,}\")\nprint(f\"  Val split    : {len(val_indices):,}\")\nprint(f\"  IA weights   : {M:,} terms\")\n\n# ── GO terms ───────────────────────────────────────────────────────\ngo_terms = []\nwith open(TERMS_FILE) as f:\n    for line in f:\n        parts = line.strip().split(\"\\t\")\n        if parts:\n            go_terms.append(parts[0])\n\nwith open(TERM_TO_IDX_FILE) as f:\n    term_to_idx = json.load(f)\n\nprint(f\"  GO terms     : {len(go_terms):,}\")\n\n# ── Protein → index ────────────────────────────────────────────────\nprot_to_idx = {p: i for i, p in enumerate(train_proteins)}\n\nprint(\"\\n✓ Data loaded\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-27T13:21:32.749649Z","iopub.execute_input":"2026-04-27T13:21:32.750069Z","iopub.status.idle":"2026-04-27T13:21:34.530677Z","shell.execute_reply.started":"2026-04-27T13:21:32.750036Z","shell.execute_reply":"2026-04-27T13:21:34.529844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom transformers import EsmTokenizer, EsmModel\n\nclass AttentionPooling(nn.Module):\n    \"\"\"\n    Attention Pooling:\n    بدل mean pooling، بنعلّم الموديل يركز على الـ residues المهمة.\n\n    input : [seq_len, 1280]  ← per-residue embeddings من ESM-2\n    output: [1280]           ← protein-level embedding\n\n    الميكانيزم:\n      score_i = tanh(W * h_i) · v     ← attention score لكل residue\n      α = softmax(scores)              ← normalize\n      output = Σ α_i * h_i            ← weighted sum\n    \"\"\"\n    def __init__(self, input_dim):\n        super().__init__()\n        self.attention = nn.Sequential(\n            nn.Linear(input_dim, input_dim // 2),\n            nn.Tanh(),\n            nn.Linear(input_dim // 2, 1)\n        )\n\n    def forward(self, hidden_states, attention_mask=None):\n        \"\"\"\n        hidden_states : [B, L, D]  أو [L, D]\n        attention_mask: [B, L]     (1=real token, 0=padding/CLS/EOS)\n        \"\"\"\n        if hidden_states.dim() == 2:\n            hidden_states = hidden_states.unsqueeze(0)\n\n        # Attention scores\n        scores = self.attention(hidden_states).squeeze(-1)  # [B, L]\n\n        # Mask لو موجود\n        if attention_mask is not None:\n            scores = scores.masked_fill(attention_mask == 0, -1e9)\n\n        weights = F.softmax(scores, dim=-1)                 # [B, L]\n\n        # Weighted sum\n        pooled = (weights.unsqueeze(-1) * hidden_states).sum(1)  # [B, D]\n        return pooled.squeeze(0) if pooled.shape[0] == 1 else pooled\n\n\nclass ESM2Encoder(nn.Module):\n    \"\"\"\n    ESM-2 encoder مع Sliding Window للبروتينات الطويلة.\n    ESM-2 frozen (مش هنعمل fine-tune) عشان نوفر VRAM.\n\n    الـ forward pass:\n      1. لو seq_len ≤ 1022: pass مباشر لـ ESM-2\n      2. لو seq_len > 1022: نقسم لـ chunks متداخلة\n         كل chunk → ESM-2 → per-residue embeddings\n         نجمع الـ chunks بـ overlap averaging\n      3. Attention Pooling → protein-level vector [1280]\n    \"\"\"\n    def __init__(self, model_name=ESM2_MODEL,\n                 max_len=MAX_SEQ_LEN,\n                 overlap=WINDOW_OVERLAP):\n        super().__init__()\n        self.tokenizer  = EsmTokenizer.from_pretrained(model_name)\n        self.esm2       = EsmModel.from_pretrained(\n            model_name, torch_dtype=torch.float16)\n        self.attn_pool  = AttentionPooling(ESM2_DIM)\n        self.max_len    = max_len\n        self.overlap    = overlap\n\n        # Freeze كل الـ ESM-2 parameters\n        for param in self.esm2.parameters():\n            param.requires_grad = False\n\n        print(f\"ESM-2 loaded & frozen\")\n        trainable = sum(p.numel() for p in self.parameters()\n                        if p.requires_grad)\n        total     = sum(p.numel() for p in self.parameters())\n        print(f\"  Trainable params: {trainable:,} / {total:,}\")\n\n    def _encode_chunk(self, sequence, device):\n        \"\"\"يعمل per-residue embeddings لـ sequence واحدة ≤ 1022\"\"\"\n        inputs = self.tokenizer(\n            sequence,\n            return_tensors=\"pt\",\n            add_special_tokens=True,\n            max_length=self.max_len + 2,\n            truncation=True\n        )\n        inputs = {k: v.to(device) for k, v in inputs.items()}\n\n        with torch.no_grad():\n            out = self.esm2(**inputs)\n\n        # نشيل CLS (0) و EOS (آخر real token)\n        hidden = out.last_hidden_state.float()  # [1, L, 1280]\n        mask   = inputs[\"attention_mask\"]       # [1, L]\n\n        # mask للـ residues بس (بدون CLS وEOS)\n        residue_mask = mask.clone()\n        residue_mask[0, 0] = 0\n        last_real = mask[0].sum().item() - 1\n        if last_real > 0:\n            residue_mask[0, int(last_real)] = 0\n\n        return hidden[0], residue_mask[0]  # [L, 1280], [L]\n\n    def _sliding_window(self, sequence, device):\n        \"\"\"\n        Sliding Window للبروتينات الطويلة.\n        بنعمل overlap averaging في مناطق التداخل.\n        \"\"\"\n        step   = self.max_len - self.overlap\n        L_full = len(sequence)\n\n        # accumulate: مجموع الـ embeddings + عداد\n        accumulated = torch.zeros(L_full, ESM2_DIM, device=device)\n        counts      = torch.zeros(L_full, device=device)\n\n        start = 0\n        while start < L_full:\n            end   = min(start + self.max_len, L_full)\n            chunk = sequence[start:end]\n\n            hidden, residue_mask = self._encode_chunk(chunk, device)\n            # hidden: [chunk_len+2, 1280] مع CLS وEOS\n            # نأخذ بس الـ residue embeddings (بدون CLS وEOS)\n            valid = residue_mask.bool()\n            residue_embs = hidden[valid]  # [chunk_len, 1280]\n            chunk_len    = end - start\n\n            accumulated[start:end] += residue_embs[:chunk_len]\n            counts[start:end]      += 1.0\n\n            if end == L_full:\n                break\n            start += step\n\n        # متوسط في مناطق الـ overlap\n        full_residues = accumulated / counts.unsqueeze(-1).clamp(min=1)\n        return full_residues  # [L_full, 1280]\n\n    def forward(self, sequences, device):\n        \"\"\"\n        sequences: list of strings\n        returns  : [B, 1280]  protein-level embeddings\n        \"\"\"\n        batch_embeddings = []\n\n        for seq in sequences:\n            if len(seq) <= self.max_len:\n                # Short: مباشر\n                hidden, residue_mask = self._encode_chunk(seq, device)\n                valid       = residue_mask.bool()\n                residues    = hidden[valid].unsqueeze(0)   # [1, L, 1280]\n                attn_input  = residue_mask[valid].unsqueeze(0)\n            else:\n                # Long: sliding window\n                residues   = self._sliding_window(seq, device).unsqueeze(0)\n                attn_input = None\n\n            # Attention Pooling → [1280]\n            pooled = self.attn_pool(residues, attn_input)  # [1, 1280]\n            batch_embeddings.append(pooled.squeeze(0))\n\n        return torch.stack(batch_embeddings)  # [B, 1280]\nprint(\"✓ Model classes defined\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-27T13:21:38.390159Z","iopub.execute_input":"2026-04-27T13:21:38.390845Z","iopub.status.idle":"2026-04-27T13:21:38.41005Z","shell.execute_reply.started":"2026-04-27T13:21:38.390807Z","shell.execute_reply":"2026-04-27T13:21:38.409133Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n!pip install -q torch-geometric\n!pip install -q torch-scatter torch-sparse \\\n    -f https://data.pyg.org/whl/torch-2.0.0+cu118.html\n!pip install -q transformers accelerate faiss-cpu\n\nimport torch, torch_geometric, transformers\nprint(f\"torch          : {torch.__version__}\")\nprint(f\"torch_geometric: {torch_geometric.__version__}\")\nprint(f\"transformers   : {transformers.__version__}\")\nprint(f\"CUDA           : {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"GPU            : {torch.cuda.get_device_name(0)}\")\n    print(f\"VRAM           : \"\n          f\"{torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB\")\nprint(\"✓ Ready\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-27T13:21:54.426987Z","iopub.execute_input":"2026-04-27T13:21:54.428016Z","iopub.status.idle":"2026-04-27T13:22:05.250151Z","shell.execute_reply.started":"2026-04-27T13:21:54.427965Z","shell.execute_reply":"2026-04-27T13:22:05.249174Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nfrom torch_geometric.nn import SAGEConv, JumpingKnowledge\n\nclass ProteinGNN(nn.Module):\n    \"\"\"\n    GraphSAGE + Jumping Knowledge\n    input : [N, 1280]  ← protein embeddings من ESM-2 + Attention Pool\n    output: [N, M]     ← probability لكل GO term\n    \"\"\"\n    def __init__(self, in_dim=ESM2_DIM,\n                 hidden_dim=HIDDEN_DIM,\n                 out_dim=M,\n                 num_layers=NUM_GNN_LAYERS,\n                 dropout=DROPOUT):\n        super().__init__()\n        self.convs   = nn.ModuleList()\n        self.norms   = nn.ModuleList()\n        self.dropout = dropout\n\n        for i in range(num_layers):\n            in_c = in_dim if i == 0 else hidden_dim\n            self.convs.append(SAGEConv(in_c, hidden_dim))\n            self.norms.append(nn.LayerNorm(hidden_dim))\n\n        self.jk          = JumpingKnowledge(\"cat\")\n        jk_dim           = hidden_dim * num_layers\n\n        self.classifier  = nn.Sequential(\n            nn.Linear(jk_dim, hidden_dim),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(hidden_dim, out_dim)\n        )\n\n    def forward(self, x, edge_index):\n        layer_outs = []\n        for conv, norm in zip(self.convs, self.norms):\n            x = conv(x, edge_index)\n            x = norm(x)\n            x = F.relu(x)\n            x = F.dropout(x, p=self.dropout, training=self.training)\n            layer_outs.append(x)\n        x = self.jk(layer_outs)\n        return self.classifier(x)  # logits (بدون sigmoid)\n\n\nclass FullModel(nn.Module):\n    \"\"\"\n    الموديل الكامل:\n      ESM-2 (frozen) → Attention Pooling → GNN → Classifier\n    \"\"\"\n    def __init__(self, gnn):\n        super().__init__()\n        self.gnn = gnn\n\n    def forward(self, node_features, edge_index):\n        \"\"\"\n        node_features: [N, 1280]  ← محسوبة مسبقاً من ESM-2\n        edge_index   : [2, E]\n        \"\"\"\n        return self.gnn(node_features, edge_index)\n\n\nprint(\"✓ GNN model defined\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-27T13:22:11.497778Z","iopub.execute_input":"2026-04-27T13:22:11.498731Z","iopub.status.idle":"2026-04-27T13:22:11.508125Z","shell.execute_reply.started":"2026-04-27T13:22:11.498687Z","shell.execute_reply":"2026-04-27T13:22:11.50733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport numpy as np\n\n# ── حط مساراتك هنا ───────────────────────────────────────────────\nGRAPH_PATH = \"/kaggle/input/datasets/mazroa/cafa-5-2/protein_graph.pt\"   # protein_graph.pt\nPREP_DIR   = \"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed\"   # باقي الملفات\n\nfiles = {\n    \"protein_graph.pt\"  : GRAPH_PATH,\n    \"label_matrix.npy\"  : f\"{PREP_DIR}/label_matrix.npy\",\n    \"train_indices.npy\" : f\"{PREP_DIR}/train_indices.npy\",\n    \"val_indices.npy\"   : f\"{PREP_DIR}/val_indices.npy\",\n    \"ia_weights.npy\"    : f\"{PREP_DIR}/ia_weights.npy\",\n    \"go_terms_list.txt\" : f\"{PREP_DIR}/go_terms_list.txt\",\n}\n\nprint(\"Checking files...\")\nall_ok = True\nfor name, path in files.items():\n    exists = os.path.exists(path)\n    print(f\"  {'✓' if exists else '✗ NOT FOUND'} {name}\")\n    if not exists:\n        all_ok = False\n\nif all_ok:\n    # تحقق من الأبعاد\n    checkpoint   = torch.load(GRAPH_PATH, weights_only=False)\n    data         = checkpoint['data']\n    label_matrix = np.load(f\"{PREP_DIR}/label_matrix.npy\")\n    train_idx    = np.load(f\"{PREP_DIR}/train_indices.npy\")\n    val_idx      = np.load(f\"{PREP_DIR}/val_indices.npy\")\n    ia_weights   = np.load(f\"{PREP_DIR}/ia_weights.npy\")\n\n    go_terms = []\n    with open(f\"{PREP_DIR}/go_terms_list.txt\") as f:\n        for line in f:\n            parts = line.strip().split(\"\\t\")\n            if parts:\n                go_terms.append(parts[0])\n\n    N = data.num_nodes\n    M = len(go_terms)\n\n    print(f\"\\nDimension Check:\")\n    print(f\"  Graph nodes      : {N:,}\")\n    print(f\"  Graph edges      : {data.num_edges:,}\")\n    print(f\"  Node features    : {data.x.shape}\")\n    print(f\"  Label matrix     : {label_matrix.shape}\")\n    print(f\"  Train indices    : {len(train_idx):,}\")\n    print(f\"  Val indices      : {len(val_idx):,}\")\n    print(f\"  IA weights       : {ia_weights.shape}\")\n    print(f\"  GO terms         : {M:,}\")\n\n    # تحقق إن الأبعاد متوافقة\n    assert label_matrix.shape[0] == N, \\\n        f\"✗ label_matrix rows ({label_matrix.shape[0]}) != nodes ({N})\"\n    assert label_matrix.shape[1] == M, \\\n        f\"✗ label_matrix cols ({label_matrix.shape[1]}) != GO terms ({M})\"\n    assert ia_weights.shape[0] == M, \\\n        f\"✗ ia_weights ({ia_weights.shape[0]}) != GO terms ({M})\"\n    assert train_idx.max() < N, \\\n        f\"✗ train_idx out of range\"\n    assert val_idx.max() < N, \\\n        f\"✗ val_idx out of range\"\n\n    print(f\"\\n✓ All dimensions match — ready to train!\")\n    print(f\"\\nConfig للـ GNN:\")\n    print(f\"  in_dim  = {data.x.shape[1]}  (ESM-2 dim)\")\n    print(f\"  out_dim = {M}  (GO terms)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-27T13:20:03.976225Z","iopub.execute_input":"2026-04-27T13:20:03.977037Z","iopub.status.idle":"2026-04-27T13:20:05.91027Z","shell.execute_reply.started":"2026-04-27T13:20:03.976998Z","shell.execute_reply":"2026-04-27T13:20:05.909302Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nfrom torch_geometric.loader import NeighborLoader\nimport torch\nimport torch.nn.functional as F\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nimport numpy as np\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n\n# ── NeighborLoader ─────────────────────────────────────────────────\n# كل batch: 512 node × أقصى 10 جيران لكل layer\ntrain_loader = NeighborLoader(\n    data,\n    num_neighbors=[10, 10, 10],   # جيران لكل layer\n    batch_size=512,\n    input_nodes=data.train_mask,\n    shuffle=True\n)\n\nval_loader = NeighborLoader(\n    data,\n    num_neighbors=[10, 10, 10],\n    batch_size=512,\n    input_nodes=data.val_mask,\n    shuffle=False\n)\n\nprint(f\"Train batches: {len(train_loader):,}\")\nprint(f\"Val batches  : {len(val_loader):,}\")\n\n# ── Model ──────────────────────────────────────────────────────────\ngnn   = ProteinGNN(\n    in_dim     = ESM2_DIM,\n    hidden_dim = HIDDEN_DIM,\n    out_dim    = M,\n    num_layers = NUM_GNN_LAYERS,\n    dropout    = DROPOUT\n).to(device)\nmodel = FullModel(gnn).to(device)\n\ntotal_params = sum(p.numel() for p in model.parameters())\nprint(f\"Model params : {total_params:,}\")\n\n# ── Loss ───────────────────────────────────────────────────────────\nia_tensor  = torch.tensor(ia_weights,  dtype=torch.float32).to(device)\npos_weight = torch.tensor(\n    [(label_matrix[:, i] == 0).sum() /\n     max(label_matrix[:, i].sum(), 1)\n     for i in range(M)],\n    dtype=torch.float32).to(device)\n\ndef ia_weighted_bce(logits, targets, pos_weight, ia_tensor):\n    bce = F.binary_cross_entropy_with_logits(\n        logits, targets,\n        pos_weight=pos_weight,\n        reduction='none'\n    )\n    weighted = bce * (1.0 + ia_tensor.unsqueeze(0))\n    return weighted.mean()\n\ndef compute_fmax(y_true_np, y_pred_np):\n    best_f1 = 0.0\n    for thr in np.arange(0.05, 0.95, 0.05):\n        pred = (y_pred_np > thr).astype(float)\n        tp   = (y_true_np * pred).sum(1)\n        prec = (tp / (pred.sum(1) + 1e-8)).mean()\n        rec  = (tp / (y_true_np.sum(1) + 1e-8)).mean()\n        f1   = 2 * prec * rec / (prec + rec + 1e-8)\n        best_f1 = max(best_f1, float(f1))\n    return best_f1\n\n# ── Optimizer ──────────────────────────────────────────────────────\noptimizer = AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\nscheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS)\n\n# ── Training Loop (Mini-Batch) ─────────────────────────────────────\nprint(f\"\\nTraining — {EPOCHS} epochs | patience={PATIENCE}\\n\")\n\nbest_fmax  = 0.0\nno_improve = 0\n\nfor epoch in range(1, EPOCHS + 1):\n\n    # ── Train ───────────────────────────────────────────────────────\n    model.train()\n    total_loss = 0\n    n_batches  = 0\n\n    for batch in train_loader:\n        batch = batch.to(device)\n        optimizer.zero_grad()\n\n        # batch.num_sampled_nodes = عدد الـ nodes في الـ batch\n        # بناخد بس الـ output بتاع الـ seed nodes (مش الجيران)\n        logits = model(batch.x, batch.edge_index)\n        # الـ seed nodes هي أول batch_size node في الـ batch\n        seed_nodes = batch.batch_size\n\n        loss = ia_weighted_bce(\n            logits[:seed_nodes],\n            batch.y[:seed_nodes],\n            pos_weight,\n            ia_tensor\n        )\n\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n\n        total_loss += loss.item()\n        n_batches  += 1\n\n    avg_train_loss = total_loss / n_batches\n    scheduler.step()\n\n    # ── Validate ────────────────────────────────────────────────────\n    model.eval()\n    val_preds  = []\n    val_labels = []\n    val_losses = []\n\n    with torch.no_grad():\n        for batch in val_loader:\n            batch = batch.to(device)\n            logits = model(batch.x, batch.edge_index)\n            seed_nodes = batch.batch_size\n\n            v_loss = ia_weighted_bce(\n                logits[:seed_nodes],\n                batch.y[:seed_nodes],\n                pos_weight,\n                ia_tensor\n            )\n            val_losses.append(v_loss.item())\n\n            probs = torch.sigmoid(logits[:seed_nodes]).cpu().numpy()\n            labs  = batch.y[:seed_nodes].cpu().numpy()\n            val_preds.append(probs)\n            val_labels.append(labs)\n\n    val_preds  = np.vstack(val_preds)\n    val_labels = np.vstack(val_labels)\n    avg_val_loss = np.mean(val_losses)\n    val_fmax     = compute_fmax(val_labels, val_preds)\n\n    # ── Checkpoint ──────────────────────────────────────────────────\n    is_best = val_fmax > best_fmax\n    if is_best:\n        best_fmax  = val_fmax\n        no_improve = 0\n        torch.save({\n            'epoch'      : epoch,\n            'model_state': model.state_dict(),\n            'val_fmax'   : best_fmax,\n            'go_terms'   : go_terms,\n        }, CKPT_PATH)\n    else:\n        no_improve += 1\n\n    if epoch % 5 == 0 or epoch == 1:\n        print(f\"Epoch {epoch:3d} | \"\n              f\"train_loss={avg_train_loss:.4f} | \"\n              f\"val_loss={avg_val_loss:.4f} | \"\n              f\"val_Fmax={val_fmax:.4f}\"\n              f\"{'  ★ best' if is_best else ''}\")\n\n    if no_improve >= PATIENCE:\n        print(f\"\\nEarly stop at epoch {epoch}\")\n        break\n\nprint(f\"\\n✓ Best val Fmax : {best_fmax:.4f}\")\nprint(f\"✓ Checkpoint    : {CKPT_PATH}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-27T13:22:55.374799Z","iopub.execute_input":"2026-04-27T13:22:55.375653Z","iopub.status.idle":"2026-04-27T13:23:08.651023Z","shell.execute_reply.started":"2026-04-27T13:22:55.375616Z","shell.execute_reply":"2026-04-27T13:23:08.649849Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Load best model ───────────────────────────────────────────────\nckpt  = torch.load(CKPT_PATH)\nmodel.load_state_dict(ckpt['model_state'])\nmodel.eval()\nprint(f\"Best model loaded (epoch {ckpt['epoch']}, \"\n      f\"Fmax={ckpt['val_fmax']:.4f})\")\n\n# ── Fmax analysis ─────────────────────────────────────────────────\nwith torch.no_grad():\n    all_logits  = model(data.x, data.edge_index)\n    val_probs   = torch.sigmoid(all_logits[data.val_mask]).cpu().numpy()\n    val_labels  = data.y[data.val_mask].cpu().numpy()\n\nprint(\"\\nFmax @ different thresholds:\")\nprint(f\"{'Threshold':>10} | {'Precision':>10} | {'Recall':>10} | {'F1':>10}\")\nprint(\"-\" * 46)\n\nbest_thr  = 0.3\nbest_fmax = 0.0\nresults   = []\n\nfor thr in np.arange(0.05, 0.80, 0.05):\n    pred  = (val_probs > thr).astype(float)\n    tp    = (val_labels * pred).sum(1)\n    prec  = (tp / (pred.sum(1) + 1e-8)).mean()\n    rec   = (tp / (val_labels.sum(1) + 1e-8)).mean()\n    f1    = 2 * prec * rec / (prec + rec + 1e-8)\n    results.append((float(thr), float(prec), float(rec), float(f1)))\n    if float(f1) > best_fmax:\n        best_fmax = float(f1)\n        best_thr  = float(thr)\n    if round(thr, 2) in [0.1, 0.2, 0.3, 0.4, 0.5]:\n        print(f\"{thr:>10.2f} | {float(prec):>10.4f} | \"\n              f\"{float(rec):>10.4f} | {float(f1):>10.4f}\")\n\nprint(f\"\\n✓ Best threshold : {best_thr:.2f}\")\nprint(f\"✓ Best Fmax      : {best_fmax:.4f}\")\n\n# ── Per-ontology breakdown ────────────────────────────────────────\nprint(f\"\\nPredictions @ threshold={best_thr:.2f}:\")\npred_final   = (val_probs > best_thr).astype(float)\navg_pred     = pred_final.sum(1).mean()\navg_true     = val_labels.sum(1).mean()\nprint(f\"  Avg predicted labels/protein : {avg_pred:.1f}\")\nprint(f\"  Avg true labels/protein      : {avg_true:.1f}\")\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import csv\n\n# ── Compute test embeddings ───────────────────────────────────────\nprint(\"Computing test protein embeddings...\")\n\n# نعيد تحميل ESM-2 للـ test inference\nencoder = ESM2Encoder().to(device)\n# نرجّع الـ attention pooling weights المحفوظة\nencoder.attn_pool.load_state_dict(\n    torch.load(GRAPH_PATH)['attn_pool_state'])\nencoder.eval()\n\ntest_features = np.zeros((N_TEST, ESM2_DIM), dtype=np.float32)\n\nfor i in tqdm(range(0, N_TEST, EMBED_BATCH), desc=\"Test encoding\"):\n    batch_pids = test_proteins[i:i+EMBED_BATCH]\n    batch_seqs = [test_seqs.get(p, \"A\") for p in batch_pids]\n\n    with torch.no_grad():\n        embs = encoder(batch_seqs, device)\n\n    test_features[i:i+len(batch_pids)] = embs.cpu().float().numpy()\n    torch.cuda.empty_cache()\n\ndel encoder\ntorch.cuda.empty_cache()\ngc.collect()\n\nprint(f\"✓ Test features: {test_features.shape}\")\n\n# ── Normalize ─────────────────────────────────────────────────────\ntest_norm = sk_normalize(test_features, norm='l2').astype(np.float32)\ntest_x    = torch.tensor(test_norm, dtype=torch.float32).to(device)\n\n# ── Inference (بدون graph edges للـ test) ─────────────────────────\n# Test proteins مش موجودين في الـ graph\n# بنمررهم عبر الـ GNN classifier بدون message passing\ndummy_edge = torch.zeros((2, 0), dtype=torch.long).to(device)\n\nmodel.eval()\nwith torch.no_grad():\n    test_logits = model(test_x, dummy_edge)\n    test_probs  = torch.sigmoid(test_logits).cpu().numpy()\n\nprint(f\"✓ Test predictions: {test_probs.shape}\")\n\n# ── Write submission ──────────────────────────────────────────────\nprint(f\"\\nWriting submission (threshold={best_thr:.2f})...\")\nn_lines = 0\n\nwith open(SUB_PATH, \"w\", newline=\"\") as f:\n    writer = csv.writer(f, delimiter=\"\\t\")\n    for i, pid in enumerate(test_proteins):\n        for j, term in enumerate(go_terms):\n            score = float(test_probs[i, j])\n            if score >= best_thr:\n                writer.writerow([pid, term, f\"{score:.4f}\"])\n                n_lines += 1\n\nprint(f\"✓ Submission saved  → {SUB_PATH}\")\nprint(f\"  Lines written     : {n_lines:,}\")\nprint(f\"  Threshold used    : {best_thr:.2f}\")\nprint(f\"  Best val Fmax     : {best_fmax:.4f}\")\nprint(f\"  GO terms covered  : {M:,}\")\n\n# ── Final summary ─────────────────────────────────────────────────\nprint(f\"\"\"\n══════════════════════════════════════════════\nDONE!\n  Train proteins : {N_TRAIN:,}\n  Test  proteins : {N_TEST:,}\n  GO terms       : {M:,}\n  Best val Fmax  : {best_fmax:.4f}\n  Threshold      : {best_thr:.2f}\n  Submission     : {SUB_PATH}\n══════════════════════════════════════════════\n\"\"\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}