{"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},{"sourceType":"kernelVersion","sourceId":314566256}],"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 numpy pandas biopython","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T10:57:45.979114Z","iopub.execute_input":"2026-04-25T10:57:45.979378Z","iopub.status.idle":"2026-04-25T10:57:52.964839Z","shell.execute_reply.started":"2026-04-25T10:57:45.979355Z","shell.execute_reply":"2026-04-25T10:57:52.964037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport json\nimport random\nimport numpy as np\nimport pandas as pd\nfrom collections import defaultdict, deque\nfrom Bio import SeqIO\n \n# Paths\nCAFA_DIR   = \"/kaggle/input/competitions/cafa-5-protein-function-prediction\"\nFASTA_FILE = f\"{CAFA_DIR}/Train/train_sequences.fasta\"\nTERMS_FILE = f\"{CAFA_DIR}/Train/train_terms.tsv\"\nOBO_FILE   = f\"{CAFA_DIR}/Train/go-basic.obo\"\nTEST_FASTA = f\"{CAFA_DIR}/Test (Targets)/testsuperset.fasta\"\nIA_FILE    = f\"{CAFA_DIR}/IA.txt\"\nOUT_DIR    = \"/kaggle/working/preprocessed\"\nos.makedirs(OUT_DIR, exist_ok=True)\n \n# Config\nMIN_SEQ_LEN           = 30\nMIN_PROTEINS_PER_TERM = 50\nREMOVE_ROOT_TERM      = True\nREMOVE_IA_ZERO        = True\nVAL_FRACTION          = 0.2\nRANDOM_SEED           = 42\n \nrandom.seed(RANDOM_SEED)\nnp.random.seed(RANDOM_SEED)\n \n# تأكد الملفات موجودة\nfor name, path in [\n    (\"FASTA\",  FASTA_FILE),\n    (\"TERMS\",  TERMS_FILE),\n    (\"OBO\",    OBO_FILE),\n    (\"TEST\",   TEST_FASTA),\n    (\"IA\",     IA_FILE),\n]:\n    status = \"✓\" if os.path.exists(path) else \"✗ NOT FOUND\"\n    print(f\"  {status}  {name}: {path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T10:57:57.185036Z","iopub.execute_input":"2026-04-25T10:57:57.185715Z","iopub.status.idle":"2026-04-25T10:57:57.532447Z","shell.execute_reply.started":"2026-04-25T10:57:57.185681Z","shell.execute_reply":"2026-04-25T10:57:57.531734Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Loading train sequences...\")\ntrain_sequences = {}\nfor record in SeqIO.parse(FASTA_FILE, \"fasta\"):\n    pid = record.id.split(\"|\")[1] if \"|\" in record.id else record.id\n    train_sequences[pid] = str(record.seq).upper()\nprint(f\"  Train sequences: {len(train_sequences):,}\")\n \nprint(\"Loading test sequences...\")\ntest_sequences = {}\nfor record in SeqIO.parse(TEST_FASTA, \"fasta\"):\n    pid = record.id.split(\"|\")[1] if \"|\" in record.id else record.id\n    test_sequences[pid] = str(record.seq).upper()\nprint(f\"  Test  sequences: {len(test_sequences):,}\")\n \n# ── 2b. Load BP annotations ────────────────────────────────────────\nprint(\"Loading BPO annotations...\")\nprotein_to_terms_raw = defaultdict(set)\nterm_to_proteins_raw = defaultdict(set)\ntotal_bpo_rows = 0\n \nwith open(TERMS_FILE, encoding=\"utf-8\") as f:\n    next(f)  # skip header\n    for line in f:\n        parts = line.strip().split(\"\\t\")\n        if len(parts) == 3 and parts[2] == \"BPO\":\n            pid, go_term = parts[0], parts[1]\n            protein_to_terms_raw[pid].add(go_term)\n            term_to_proteins_raw[go_term].add(pid)\n            total_bpo_rows += 1\n \nprint(f\"  BPO rows        : {total_bpo_rows:,}\")\nprint(f\"  Unique proteins : {len(protein_to_terms_raw):,}\")\nprint(f\"  Unique GO terms : {len(term_to_proteins_raw):,}\")\n \n# ── 2c. Load IA weights ────────────────────────────────────────────\nia_weights = {}\nif os.path.exists(IA_FILE):\n    with open(IA_FILE) as f:\n        for line in f:\n            parts = line.strip().split(\"\\t\")\n            if len(parts) == 2:\n                try:\n                    ia_weights[parts[0]] = float(parts[1])\n                except ValueError:\n                    pass\nprint(f\"  IA weights loaded: {len(ia_weights):,} terms\")\n \n# ── 2d. Parse GO OBO (parents map) ────────────────────────────────\nprint(\"Parsing GO OBO file...\")\ngo_parents = defaultdict(set)\ngo_names   = {}\n \ncurrent_id  = None\nis_obsolete = False\nwith open(OBO_FILE, encoding=\"utf-8\") as f:\n    for line in f:\n        line = line.strip()\n        if line == \"[Term]\":\n            current_id  = None\n            is_obsolete = False\n        elif line.startswith(\"id: GO:\"):\n            current_id = line[4:].strip()\n        elif line.startswith(\"name: \") and current_id:\n            go_names[current_id] = line[6:].strip()\n        elif line == \"is_obsolete: true\":\n            is_obsolete = True\n        elif line.startswith(\"is_a: \") and current_id and not is_obsolete:\n            parent = line[6:].split(\"!\")[0].strip()\n            if parent.startswith(\"GO:\"):\n                go_parents[current_id].add(parent)\n        elif line.startswith(\"relationship: \") and current_id and not is_obsolete:\n            parts = line[14:].split()\n            if len(parts) >= 2 and parts[1].startswith(\"GO:\"):\n                go_parents[current_id].add(parts[1])\n \nprint(f\"  GO terms in OBO: {len(go_names):,}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T10:58:00.40796Z","iopub.execute_input":"2026-04-25T10:58:00.408561Z","iopub.status.idle":"2026-04-25T10:58:11.039661Z","shell.execute_reply.started":"2026-04-25T10:58:00.408528Z","shell.execute_reply":"2026-04-25T10:58:11.03903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"STEP 1: Sequence Cleaning\")\nprint(\"=\"*60)\n \nALLOWED_CHARS = set(\"ACDEFGHIKLMNPQRSTVWYXBZU\")\n \ndef clean_sequences(sequences, bp_proteins=None):\n    \"\"\"\n    - شيل البروتينات الأقصر من MIN_SEQ_LEN\n    - حوّل U → C (Selenocysteine → Cysteine)\n    - سيب X كما هو (ESM-2 بيتعامل معاه)\n    - شيل أي حروف غير معروفة\n    - لو bp_proteins محدد، شيل البروتينات اللي مالهاش BPO annotation\n    \"\"\"\n    cleaned = {}\n    stats = defaultdict(int)\n    stats[\"total\"] = len(sequences)\n \n    for pid, seq in sequences.items():\n        # لو محتاجين BP annotation بس\n        if bp_proteins is not None and pid not in bp_proteins:\n            stats[\"no_annotation\"] += 1\n            continue\n \n        # شيل قصيرة\n        if len(seq) < MIN_SEQ_LEN:\n            stats[\"too_short\"] += 1\n            continue\n \n        # U → C\n        if \"U\" in seq:\n            seq = seq.replace(\"U\", \"C\")\n            stats[\"replaced_U\"] += 1\n \n        # X — نسجّل بس مش نشيل\n        if \"X\" in seq:\n            stats[\"has_X\"] += 1\n \n        # شيل أي حروف غريبة تماماً\n        if not set(seq).issubset(ALLOWED_CHARS):\n            seq = \"\".join(c if c in ALLOWED_CHARS else \"X\" for c in seq)\n            stats[\"fixed_chars\"] += 1\n \n        cleaned[pid] = seq\n \n    stats[\"total_after\"] = len(cleaned)\n    return cleaned, stats\n \n \n# تنظيف الـ train\nbp_proteins = set(protein_to_terms_raw.keys())\ntrain_clean, stats_train = clean_sequences(train_sequences, bp_proteins)\n \nprint(f\"  Total train proteins    : {stats_train['total']:,}\")\nprint(f\"  Removed (no BP annot.)  : {stats_train['no_annotation']:,}\")\nprint(f\"  Removed (too short)     : {stats_train['too_short']:,}\")\nprint(f\"  U→C replacements        : {stats_train['replaced_U']:,}\")\nprint(f\"  Contains X (kept)       : {stats_train['has_X']:,}\")\nprint(f\"  Fixed unusual chars     : {stats_train['fixed_chars']:,}\")\nprint(f\"  After cleaning          : {stats_train['total_after']:,}\")\n \n# تنظيف الـ test (بدون فلترة BP)\ntest_clean, stats_test = clean_sequences(test_sequences)\nprint(f\"\\n  Test proteins before    : {stats_test['total']:,}\")\nprint(f\"  Test removed (short)    : {stats_test['too_short']:,}\")\nprint(f\"  Test after cleaning     : {stats_test['total_after']:,}\")\n \n# إحصاء الطول\ntrain_lens = [len(s) for s in train_clean.values()]\nprint(f\"\\n  Train length stats:\")\nprint(f\"    Min    : {min(train_lens)}\")\nprint(f\"    Max    : {max(train_lens):,}\")\nprint(f\"    Mean   : {np.mean(train_lens):.0f}\")\nprint(f\"    Median : {np.median(train_lens):.0f}\")\nprint(f\"    > 1022 : {sum(1 for l in train_lens if l > 1022):,}  \"\n      f\"({100*sum(1 for l in train_lens if l > 1022)/len(train_lens):.1f}%)\")\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T10:58:22.637827Z","iopub.execute_input":"2026-04-25T10:58:22.638146Z","iopub.status.idle":"2026-04-25T10:58:24.272368Z","shell.execute_reply.started":"2026-04-25T10:58:22.638104Z","shell.execute_reply":"2026-04-25T10:58:24.271715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"STEP 2: GO Term Filtering\")\nprint(\"=\"*60)\n \nterms_to_keep = set(term_to_proteins_raw.keys())\nprint(f\"  Starting GO terms: {len(terms_to_keep):,}\")\n \n# 2a. شيل الـ root term\nif REMOVE_ROOT_TERM:\n    root = \"GO:0008150\"\n    if root in terms_to_keep:\n        terms_to_keep.discard(root)\n        print(f\"  Removed root term {root} (biological_process)\")\n        print(f\"    → Remaining: {len(terms_to_keep):,}\")\n \n# 2b. شيل الـ terms بـ IA = 0\nif REMOVE_IA_ZERO and ia_weights:\n    ia_zero = {t for t in terms_to_keep if ia_weights.get(t, -1) == 0.0}\n    terms_to_keep -= ia_zero\n    print(f\"  Removed {len(ia_zero):,} terms with IA = 0\")\n    print(f\"    → Remaining: {len(terms_to_keep):,}\")\n \n# 2c. شيل الـ rare terms (< MIN_PROTEINS_PER_TERM)\nterm_counts = {}\nfor term in terms_to_keep:\n    count = len(term_to_proteins_raw[term] & set(train_clean.keys()))\n    term_counts[term] = count\n \nrare_terms = {t for t, c in term_counts.items() if c < MIN_PROTEINS_PER_TERM}\nterms_to_keep -= rare_terms\nprint(f\"  Removed {len(rare_terms):,} rare terms (< {MIN_PROTEINS_PER_TERM} proteins)\")\nprint(f\"    → Remaining: {len(terms_to_keep):,}\")\n \n# الترتيب النهائي\nfinal_terms = sorted(terms_to_keep)\nterm_to_idx = {t: i for i, t in enumerate(final_terms)}\nM = len(final_terms)\n \n# إحصاء الـ annotations المتبقية\nannotations_kept = sum(\n    1 for p in train_clean\n    for t in protein_to_terms_raw.get(p, set())\n    if t in terms_to_keep\n)\nprint(f\"\\n  Final GO terms          : {M:,}\")\nprint(f\"  Annotations kept        : {annotations_kept:,} / {total_bpo_rows:,} \"\n      f\"({100*annotations_kept/total_bpo_rows:.1f}%)\")\n \n# عرض sample من الـ terms\nprint(f\"\\n  Sample terms (first 5):\")\nfor t in final_terms[:5]:\n    print(f\"    {t} | {go_names.get(t,'?')[:45]:45s} | \"\n          f\"n={term_counts.get(t,0):,} | IA={ia_weights.get(t,0):.3f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T10:58:28.317914Z","iopub.execute_input":"2026-04-25T10:58:28.318595Z","iopub.status.idle":"2026-04-25T10:59:46.966285Z","shell.execute_reply.started":"2026-04-25T10:58:28.318565Z","shell.execute_reply":"2026-04-25T10:59:46.965622Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"STEP 3: Build Final Protein Set\")\nprint(\"=\"*60)\n \n# بس البروتينات اللي عندها على الأقل 1 term بعد الفلترة\nfinal_proteins        = []\nprotein_to_terms_clean = {}\n \nfor pid in sorted(train_clean.keys()):\n    remaining = protein_to_terms_raw[pid] & terms_to_keep\n    if remaining:\n        final_proteins.append(pid)\n        protein_to_terms_clean[pid] = remaining\n \nprotein_to_idx = {p: i for i, p in enumerate(final_proteins)}\nN = len(final_proteins)\n \nprint(f\"  Proteins with ≥1 remaining term: {N:,}\")\n \n# إحصاء terms لكل بروتين\nterms_per_prot = [len(protein_to_terms_clean[p]) for p in final_proteins]\nprint(f\"  Terms per protein:\")\nprint(f\"    Min    : {min(terms_per_prot)}\")\nprint(f\"    Max    : {max(terms_per_prot)}\")\nprint(f\"    Mean   : {np.mean(terms_per_prot):.1f}\")\nprint(f\"    Median : {np.median(terms_per_prot):.0f}\")\n \n# final sequences (بدون truncation — هيتعمل في الـ embedding step)\ntrain_sequences_final = {pid: train_clean[pid] for pid in final_proteins}\ntest_sequences_final  = {pid: test_clean[pid] for pid in test_clean}\n \nprint(f\"\\n  Train sequences (no truncation): {len(train_sequences_final):,}\")\nprint(f\"  Test  sequences (no truncation): {len(test_sequences_final):,}\")\n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T10:59:55.372949Z","iopub.execute_input":"2026-04-25T10:59:55.373437Z","iopub.status.idle":"2026-04-25T10:59:56.803594Z","shell.execute_reply.started":"2026-04-25T10:59:55.373409Z","shell.execute_reply":"2026-04-25T10:59:56.802926Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"STEP 4: GO Label Propagation\")\nprint(\"=\"*60)\n \ndef get_all_ancestors(term_id, parents_map):\n    \"\"\"BFS للأعلى في الـ DAG — بيرجع كل الـ ancestor terms\"\"\"\n    visited = set()\n    queue   = deque([term_id])\n    while queue:\n        current = queue.popleft()\n        for parent in parents_map.get(current, set()):\n            if parent not in visited:\n                visited.add(parent)\n                queue.append(parent)\n    return visited\n \n# نبني ancestor cache للـ final_terms بس\nprint(\"  Building ancestor cache...\")\nfinal_terms_set       = set(final_terms)\nterm_ancestors_cache  = {}\nfor term in final_terms:\n    anc = get_all_ancestors(term, go_parents)\n    term_ancestors_cache[term] = anc & final_terms_set\n \nanc_counts = [len(v) for v in term_ancestors_cache.values()]\nprint(f\"  Mean ancestors per term : {np.mean(anc_counts):.1f}\")\nprint(f\"  Max ancestors           : {max(anc_counts)}\")\nprint(f\"  Terms with 0 ancestors  : {sum(1 for x in anc_counts if x == 0)}\")\n \n# Propagate\nprint(\"  Propagating labels...\")\nprotein_to_terms_propagated = {}\nadded_total = 0\n \nfor pid in final_proteins:\n    original = protein_to_terms_clean[pid]\n    propagated = set(original)\n    for term in original:\n        propagated |= term_ancestors_cache.get(term, set())\n    protein_to_terms_propagated[pid] = propagated\n    added_total += len(propagated) - len(original)\n \norig_total = sum(len(v) for v in protein_to_terms_clean.values())\nprop_total = sum(len(v) for v in protein_to_terms_propagated.values())\n \nprint(f\"\\n  Before propagation: {orig_total:,} labels \"\n      f\"(avg {orig_total/N:.1f}/protein)\")\nprint(f\"  After propagation : {prop_total:,} labels \"\n      f\"(avg {prop_total/N:.1f}/protein)\")\nprint(f\"  Added             : {prop_total - orig_total:,} labels\")\n \n# Sanity check\nif prop_total/N >= M * 0.95:\n    print(\"  ⚠ Propagation suspicious — using original labels\")\n    protein_to_terms_propagated = protein_to_terms_clean\nelse:\n    print(\"  ✓ Propagation OK\")\n \n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T11:00:00.01582Z","iopub.execute_input":"2026-04-25T11:00:00.016547Z","iopub.status.idle":"2026-04-25T11:00:03.054815Z","shell.execute_reply.started":"2026-04-25T11:00:00.016508Z","shell.execute_reply":"2026-04-25T11:00:03.054182Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"STEP 5: Label Matrix + Stratified Train/Val Split\")\nprint(\"=\"*60)\n \n# ── Build label matrix ─────────────────────────────────────────────\nprint(f\"  Building label matrix: {N:,} × {M:,}...\")\nlabel_matrix = np.zeros((N, M), dtype=np.float32)\nfor i, pid in enumerate(final_proteins):\n    for term in protein_to_terms_propagated[pid]:\n        j = term_to_idx.get(term)\n        if j is not None:\n            label_matrix[i, j] = 1.0\n \ntotal_pos = int(label_matrix.sum())\nprint(f\"  Total positive labels   : {total_pos:,}\")\nprint(f\"  Label density           : {100*total_pos/(N*M):.4f}%\")\nprint(f\"  Avg labels per protein  : {total_pos/N:.1f}\")\nprint(f\"  Avg proteins per term   : {total_pos/M:.1f}\")\n \n# ── Stratified split ────────────────────────────────────────────────\n# نقسم على أساس عدد الـ labels عشان كل split يكون فيه توزيع متوازن\nprint(f\"\\n  Building stratified train/val split...\")\n \n# نقسم البروتينات لـ buckets حسب عدد labels\nlabel_counts = label_matrix.sum(axis=1)\nbuckets      = defaultdict(list)\nfor i, count in enumerate(label_counts):\n    bucket = min(int(count) // 5, 20)  # bucket size = 5 labels\n    buckets[bucket].append(i)\n \ntrain_indices = []\nval_indices   = []\n \nfor bucket, indices in buckets.items():\n    random.shuffle(indices)\n    val_size = max(1, int(len(indices) * VAL_FRACTION))\n    val_indices.extend(indices[:val_size])\n    train_indices.extend(indices[val_size:])\n \ntrain_indices = sorted(train_indices)\nval_indices   = sorted(val_indices)\n \n# تحقق من التوزيع\ntrain_labels      = label_matrix[train_indices]\nval_labels        = label_matrix[val_indices]\nzero_val_terms    = int((val_labels.sum(0) == 0).sum())\n \nprint(f\"  Train: {len(train_indices):,} proteins ({100*len(train_indices)/N:.1f}%)\")\nprint(f\"  Val  : {len(val_indices):,} proteins ({100*len(val_indices)/N:.1f}%)\")\nprint(f\"  Terms with 0 positives in val: {zero_val_terms} \"\n      f\"(stratified — should be ~0)\")\n \n# ── IA weight vector ───────────────────────────────────────────────\nia_weight_vector = np.array(\n    [ia_weights.get(t, 0.0) for t in final_terms], dtype=np.float32)\nprint(f\"\\n  IA weight vector: {ia_weight_vector.shape}\")\nprint(f\"  Non-zero IA     : {int((ia_weight_vector > 0).sum()):,} / {M:,}\")\nprint(f\"  Max IA          : {ia_weight_vector.max():.4f}\")\nprint(f\"  Mean IA (>0)    : {ia_weight_vector[ia_weight_vector > 0].mean():.4f}\")\n \n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T11:00:05.385214Z","iopub.execute_input":"2026-04-25T11:00:05.385541Z","iopub.status.idle":"2026-04-25T11:00:08.434699Z","shell.execute_reply.started":"2026-04-25T11:00:05.385516Z","shell.execute_reply":"2026-04-25T11:00:08.433945Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"STEP 6: Saving Outputs\")\nprint(\"=\"*60)\n \n# 1. Protein IDs\nwith open(f\"{OUT_DIR}/train_protein_ids.txt\", \"w\") as f:\n    f.write(\"\\n\".join(final_proteins))\nprint(f\"  ✓ train_protein_ids.txt       ({N:,} proteins)\")\n \nwith open(f\"{OUT_DIR}/test_protein_ids.txt\", \"w\") as f:\n    f.write(\"\\n\".join(sorted(test_sequences_final.keys())))\nprint(f\"  ✓ test_protein_ids.txt        ({len(test_sequences_final):,} proteins)\")\n \n# 2. GO terms list\nwith open(f\"{OUT_DIR}/go_terms_list.txt\", \"w\") as f:\n    for term in final_terms:\n        f.write(f\"{term}\\t{go_names.get(term,'')}\\t{ia_weights.get(term,0):.6f}\\n\")\nprint(f\"  ✓ go_terms_list.txt           ({M:,} terms)\")\n \n# 3. Label matrix\nnp.save(f\"{OUT_DIR}/label_matrix.npy\", label_matrix)\nprint(f\"  ✓ label_matrix.npy            {label_matrix.shape}  \"\n      f\"{label_matrix.nbytes/1e6:.0f} MB\")\n \n# 4. Train/Val indices\nnp.save(f\"{OUT_DIR}/train_indices.npy\", np.array(train_indices))\nnp.save(f\"{OUT_DIR}/val_indices.npy\",   np.array(val_indices))\nprint(f\"  ✓ train_indices.npy           ({len(train_indices):,})\")\nprint(f\"  ✓ val_indices.npy             ({len(val_indices):,})\")\n \n# 5. IA weights\nnp.save(f\"{OUT_DIR}/ia_weights.npy\", ia_weight_vector)\nprint(f\"  ✓ ia_weights.npy              {ia_weight_vector.shape}\")\n \n# 6. Train sequences FASTA (بدون truncation)\nfasta_path = f\"{OUT_DIR}/train_sequences_clean.fasta\"\nwith open(fasta_path, \"w\") as f:\n    for pid in final_proteins:\n        seq = train_sequences_final[pid]\n        f.write(f\">{pid}\\n\")\n        for i in range(0, len(seq), 60):\n            f.write(seq[i:i+60] + \"\\n\")\nprint(f\"  ✓ train_sequences_clean.fasta ({len(final_proteins):,} sequences)\")\n \n# 7. Test sequences FASTA (بدون truncation)\nfasta_path_test = f\"{OUT_DIR}/test_sequences_clean.fasta\"\nwith open(fasta_path_test, \"w\") as f:\n    for pid in sorted(test_sequences_final.keys()):\n        seq = test_sequences_final[pid]\n        f.write(f\">{pid}\\n\")\n        for i in range(0, len(seq), 60):\n            f.write(seq[i:i+60] + \"\\n\")\nprint(f\"  ✓ test_sequences_clean.fasta  ({len(test_sequences_final):,} sequences)\")\n \n# 8. Mappings JSON\nwith open(f\"{OUT_DIR}/term_to_idx.json\", \"w\") as f:\n    json.dump(term_to_idx, f)\nwith open(f\"{OUT_DIR}/protein_to_idx.json\", \"w\") as f:\n    json.dump(protein_to_idx, f)\nprint(f\"  ✓ term_to_idx.json\")\nprint(f\"  ✓ protein_to_idx.json\")\n \n# 9. Protein → terms mapping (للـ debugging)\nwith open(f\"{OUT_DIR}/protein_to_terms.json\", \"w\") as f:\n    json.dump(\n        {p: sorted(list(terms))\n         for p, terms in protein_to_terms_propagated.items()},\n        f\n    )\nprint(f\"  ✓ protein_to_terms.json\")\n \n# 10. Config\nconfig = {\n    \"min_seq_len\"           : MIN_SEQ_LEN,\n    \"min_proteins_per_term\" : MIN_PROTEINS_PER_TERM,\n    \"remove_root_term\"      : REMOVE_ROOT_TERM,\n    \"remove_ia_zero\"        : REMOVE_IA_ZERO,\n    \"val_fraction\"          : VAL_FRACTION,\n    \"random_seed\"           : RANDOM_SEED,\n    \"go_propagation\"        : True,\n    \"truncation\"            : False,\n    \"pooling\"               : \"per_residue\",\n    \"n_train_proteins\"      : N,\n    \"n_test_proteins\"       : len(test_sequences_final),\n    \"n_go_terms\"            : M,\n    \"n_train_split\"         : len(train_indices),\n    \"n_val_split\"           : len(val_indices),\n    \"total_positive_labels\" : total_pos,\n    \"label_density_pct\"     : round(100 * total_pos / (N * M), 4),\n}\nwith open(f\"{OUT_DIR}/config.json\", \"w\") as f:\n    json.dump(config, f, indent=2)\nprint(f\"  ✓ config.json\")\n \n ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T11:00:13.534993Z","iopub.execute_input":"2026-04-25T11:00:13.535292Z","iopub.status.idle":"2026-04-25T11:00:18.131609Z","shell.execute_reply.started":"2026-04-25T11:00:13.535264Z","shell.execute_reply":"2026-04-25T11:00:18.130996Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\" + \"=\"*60)\nprint(\"PREPROCESSING COMPLETE — SUMMARY\")\nprint(\"=\"*60)\n \ntrain_lens_final = [len(train_sequences_final[p]) for p in final_proteins]\ntest_lens_final  = [len(test_sequences_final[p])  for p in test_sequences_final]\n \nlong_train = sum(1 for l in train_lens_final if l > 1022)\nlong_test  = sum(1 for l in test_lens_final  if l > 1022)\n \nprint(f\"\"\"\nINPUT:\n  Train proteins raw   : {len(train_sequences):>10,}\n  Test  proteins raw   : {len(test_sequences):>10,}\n  BPO annotation rows  : {total_bpo_rows:>10,}\n  GO terms raw         : {len(term_to_proteins_raw):>10,}\n \nAFTER PREPROCESSING:\n  Train proteins       : {N:>10,}\n  Test  proteins       : {len(test_sequences_final):>10,}\n  GO terms             : {M:>10,}\n  Label matrix         : {N:,} × {M:,}\n  Positive labels      : {total_pos:>10,}\n  Label density        : {100*total_pos/(N*M):>10.4f}%\n  Avg labels/protein   : {total_pos/N:>10.1f}\n \nSEQUENCE LENGTHS (no truncation):\n  Train — min={min(train_lens_final)}, max={max(train_lens_final):,}, mean={np.mean(train_lens_final):.0f}\n  Test  — min={min(test_lens_final)},  max={max(test_lens_final):,},  mean={np.mean(test_lens_final):.0f}\n  Train proteins > 1022 AA : {long_train:,} ({100*long_train/N:.1f}%) ← هيتعمل sliding window في الـ embedding\n  Test  proteins > 1022 AA : {long_test:,}  ({100*long_test/len(test_sequences_final):.1f}%)\n \nTRAIN/VAL SPLIT (stratified):\n  Train : {len(train_indices):>10,} proteins\n  Val   : {len(val_indices):>10,} proteins\n \nOUTPUT FILES → {OUT_DIR}/\n  train_protein_ids.txt\n  test_protein_ids.txt\n  go_terms_list.txt\n  label_matrix.npy            ← {N:,} × {M:,}\n  train_indices.npy\n  val_indices.npy\n  ia_weights.npy              ← للـ weighted loss function\n  train_sequences_clean.fasta ← بدون truncation (للـ sliding window)\n  test_sequences_clean.fasta  ← بدون truncation\n  term_to_idx.json\n  protein_to_idx.json\n  protein_to_terms.json\n  config.json\n\"\"\")\nprint(\"=\"*60)\nprint(\"✓ Ready for per-residue embedding extraction!\")\nprint(\"=\"*60)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-25T11:00:26.24006Z","iopub.execute_input":"2026-04-25T11:00:26.240375Z","iopub.status.idle":"2026-04-25T11:00:26.351431Z","shell.execute_reply.started":"2026-04-25T11:00:26.240349Z","shell.execute_reply":"2026-04-25T11:00:26.350729Z"}},"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-26T14:18:33.557855Z","iopub.execute_input":"2026-04-26T14:18:33.558586Z","iopub.status.idle":"2026-04-26T14:18:37.122898Z","shell.execute_reply.started":"2026-04-26T14:18:33.558546Z","shell.execute_reply":"2026-04-26T14:18:37.121968Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport os\nimport json\nimport numpy as np\nimport pandas as pd\nimport gc\n\n# ── Paths ──────────────────────────────────────────────────────────\nPREP_DIR   = \"/kaggle/working/preprocessed\"   # output من الـ preprocessing\nOUTPUT_DIR = \"/kaggle/working\"\n\nTRAIN_IDS_FILE   = f\"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed/train_protein_ids.txt\"\nTEST_IDS_FILE    = f\"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed/test_protein_ids.txt\"\nTRAIN_FASTA      = f\"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed/train_sequences_clean.fasta\"\nTEST_FASTA       = f\"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed/test_sequences_clean.fasta\"\nLABEL_MATRIX     = f\"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed/label_matrix.npy\"\nTRAIN_IDX_FILE   = f\"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed/train_indices.npy\"\nVAL_IDX_FILE     = f\"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed/val_indices.npy\"\nIA_WEIGHTS_FILE  = f\"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed/ia_weights.npy\"\nTERMS_FILE       = f\"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed/go_terms_list.txt\"\nTERM_TO_IDX_FILE = f\"/kaggle/input/datasets/mazroa/cafa-5-2/preprocessed/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!'}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T14:27:03.071802Z","iopub.execute_input":"2026-04-26T14:27:03.072412Z","iopub.status.idle":"2026-04-26T14:27:03.115425Z","shell.execute_reply.started":"2026-04-26T14:27:03.072374Z","shell.execute_reply":"2026-04-26T14:27:03.114775Z"}},"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\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T14:27:06.85023Z","iopub.execute_input":"2026-04-26T14:27:06.85086Z","iopub.status.idle":"2026-04-26T14:27:17.965688Z","shell.execute_reply.started":"2026-04-26T14:27:06.850826Z","shell.execute_reply":"2026-04-26T14:27:17.964748Z"}},"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\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T14:27:27.067884Z","iopub.execute_input":"2026-04-26T14:27:27.068218Z","iopub.status.idle":"2026-04-26T14:27:28.798313Z","shell.execute_reply.started":"2026-04-26T14:27:27.068184Z","shell.execute_reply":"2026-04-26T14:27:28.797529Z"}},"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        \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-26T14:28:19.600731Z","iopub.execute_input":"2026-04-26T14:28:19.601525Z","iopub.status.idle":"2026-04-26T14:28:19.619963Z","shell.execute_reply.started":"2026-04-26T14:28:19.601493Z","shell.execute_reply":"2026-04-26T14:28:19.619306Z"}},"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-26T14:18:51.091586Z","iopub.execute_input":"2026-04-26T14:18:51.09217Z","iopub.status.idle":"2026-04-26T14:19:01.128886Z","shell.execute_reply.started":"2026-04-26T14:18:51.092128Z","shell.execute_reply":"2026-04-26T14:19:01.127939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\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-26T14:31:51.540368Z","iopub.execute_input":"2026-04-26T14:31:51.5409Z","iopub.status.idle":"2026-04-26T14:31:51.55252Z","shell.execute_reply.started":"2026-04-26T14:31:51.540867Z","shell.execute_reply":"2026-04-26T14:31:51.551631Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport faiss\nfrom sklearn.preprocessing import normalize as sk_normalize\nfrom tqdm import tqdm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device: {device}\")\n\n# ── Load ESM-2 + Attention Pooling ────────────────────────────────\nprint(\"\\nLoading ESM-2...\")\nencoder = ESM2Encoder()\nencoder = encoder.to(device)\nencoder.eval()\n\n# ── Compute per-protein embeddings ────────────────────────────────\n# نحسب embedding لكل بروتين (ESM-2 frozen + Attention Pooling)\n# النتيجة: [N, 1280] — ده اللي هيبقى الـ node features في الـ graph\n\nEMBED_BATCH = 8   # عدد البروتينات في كل batch\nnode_features = np.zeros((N_TRAIN, ESM2_DIM), dtype=np.float32)\n\nprint(f\"\\nComputing node features for {N_TRAIN:,} proteins...\")\nprint(f\"  Batch size : {EMBED_BATCH}\")\nprint(f\"  Long seqs  : \"\n      f\"{sum(1 for p in train_proteins if len(train_seqs[p]) > MAX_SEQ_LEN):,}\"\n      f\" → sliding window\")\n\nfor i in tqdm(range(0, N_TRAIN, EMBED_BATCH), desc=\"Encoding\"):\n    batch_pids = train_proteins[i:i+EMBED_BATCH]\n    batch_seqs = [train_seqs[p] for p in batch_pids]\n\n    with torch.no_grad():\n        embs = encoder(batch_seqs, device)   # [B, 1280]\n\n    node_features[i:i+len(batch_pids)] = embs.cpu().float().numpy()\n\n    torch.cuda.empty_cache()\n    gc.collect()\n\nprint(f\"✓ Node features: {node_features.shape}\")\n\n# ── نشيل ESM-2 من الـ GPU عشان نوفر VRAM للـ GNN ──────────────────\n# الـ Attention Pooling weights هنحتفظ بيهم\nattn_pool_state = encoder.attn_pool.state_dict()\ndel encoder\ntorch.cuda.empty_cache()\ngc.collect()\nprint(\"✓ ESM-2 unloaded from GPU\")\n\n# ── Normalize + Build Graph (FAISS) ───────────────────────────────\nprint(\"\\nBuilding similarity graph...\")\nnode_features_norm = sk_normalize(node_features, norm='l2').astype(np.float32)\nnp.save(f\"{OUTPUT_DIR}/node_features.npy\", node_features)\n\nd     = node_features_norm.shape[1]\nindex = faiss.IndexFlatIP(d)\nindex.add(node_features_norm)\nprint(f\"  FAISS index: {index.ntotal:,} vectors\")\n\nFAISS_BATCH = 1000\nK           = 100\nedges_src, edges_dst = [], []\n\nfor i in tqdm(range(0, N_TRAIN, FAISS_BATCH), desc=\"Building edges\"):\n    batch  = node_features_norm[i:i+FAISS_BATCH]\n    scores, indices = index.search(batch, K)\n    for bi in range(len(batch)):\n        actual_i = i + bi\n        for k in range(K):\n            j   = int(indices[bi, k])\n            sim = float(scores[bi, k])\n            if j > actual_i and sim >= SIM_THRESHOLD:\n                edges_src.append(actual_i)\n                edges_dst.append(j)\n\ndel index\ngc.collect()\n\nprint(f\"  Undirected edges: {len(edges_src):,}\")\n\n# ── Build PyG Data object ─────────────────────────────────────────\nfrom torch_geometric.data import Data\nfrom sklearn.model_selection import train_test_split\n\nedge_index = torch.tensor(\n    [edges_src + edges_dst,\n     edges_dst + edges_src], dtype=torch.long)\ndel edges_src, edges_dst\ngc.collect()\n\nx = torch.tensor(node_features_norm, dtype=torch.float32)\ny = torch.tensor(label_matrix,       dtype=torch.float32)\n\ntrain_mask = torch.zeros(N_TRAIN, dtype=torch.bool)\nval_mask   = torch.zeros(N_TRAIN, dtype=torch.bool)\ntrain_mask[train_indices] = True\nval_mask[val_indices]     = True\n\ndata = Data(x=x, edge_index=edge_index, y=y,\n            train_mask=train_mask,\n            val_mask=val_mask)\n\ndegrees  = torch.zeros(N_TRAIN)\ndegrees.scatter_add_(0, edge_index[0],\n                     torch.ones(edge_index.shape[1]))\nisolated = (degrees == 0).sum().item()\n\nprint(f\"\\n{'='*45}\")\nprint(f\"Nodes      : {data.num_nodes:,}\")\nprint(f\"Edges      : {data.num_edges:,}\")\nprint(f\"Avg degree : {data.num_edges/data.num_nodes:.1f}\")\nprint(f\"Train      : {train_mask.sum():,}\")\nprint(f\"Val        : {val_mask.sum():,}\")\nprint(f\"Isolated   : {isolated:,} ({100*isolated/N_TRAIN:.1f}%)\")\n\ntorch.save({'data': data, 'attn_pool_state': attn_pool_state,\n            'train_proteins': train_proteins,\n            'test_proteins' : test_proteins,\n            'go_terms'      : go_terms},\n           GRAPH_PATH)\nprint(f\"\\n✓ Graph saved → {GRAPH_PATH}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-26T14:21:14.801893Z","iopub.execute_input":"2026-04-26T14:21:14.80253Z","iopub.status.idle":"2026-04-26T14:21:15.648953Z","shell.execute_reply.started":"2026-04-26T14:21:14.802495Z","shell.execute_reply":"2026-04-26T14:21:15.648016Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from 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# ── 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},"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}]}