{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":41875,"databundleVersionId":5521661,"sourceType":"competition"},{"sourceId":5538488,"sourceType":"datasetVersion","datasetId":3192134},{"sourceId":5549164,"sourceType":"datasetVersion","datasetId":3197305},{"sourceId":6057158,"sourceType":"datasetVersion","datasetId":3465890},{"sourceId":6080228,"sourceType":"datasetVersion","datasetId":3480882},{"sourceId":6194252,"sourceType":"datasetVersion","datasetId":3555818},{"sourceId":6240759,"sourceType":"datasetVersion","datasetId":3585518}],"dockerImageVersionId":31153,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!/usr/bin/env python3\n# ============================================================\n# CAFA-5 — ESM2-650M Embedder (Streaming + Shards + Hotfix)\n# Outputs:\n#   /kaggle/working/emb_cache/train_<id>.npz\n#   /kaggle/working/emb_cache/test_<id>.npz\n#\n# Run multiple times with:\n#   MODE=train / test\n#   SHARDS=<N>, SHARD_ID=0..N-1\n#\n# Then bundle emb_cache as dataset for Notebook B.\n# ============================================================\n\nimport os\nos.environ.setdefault(\"TRANSFORMERS_SKIP_CHAT_TEMPLATE_LOAD\", \"1\")\nos.environ.setdefault(\"HF_HUB_ENABLE_HF_TRANSFER\", \"1\")\nos.environ.setdefault(\"TOKENIZERS_PARALLELISM\", \"true\")\n\n# ----------------------\n# Chat-template hotfix\n# ----------------------\ndef _install_chat_template_hotfix():\n    try:\n        import transformers\n        from transformers.utils import hub as _tf_hub\n        orig = getattr(_tf_hub, \"list_repo_templates\", None)\n        if callable(orig):\n            def _safe(*a, **k):\n                try:\n                    return orig(*a, **k)\n                except Exception:\n                    return []\n            _tf_hub.list_repo_templates = _safe\n            print(\"[Hotfix] transformers.utils.hub.list_repo_templates patched.\")\n    except Exception as e:\n        print(f\"[Hotfix] hub patch skipped: {e}\")\n    try:\n        from transformers import tokenization_utils_base as _tk\n        orig2 = getattr(_tk, \"list_repo_templates\", None)\n        if callable(orig2):\n            def _safe2(*a, **k):\n                try:\n                    return orig2(*a, **k)\n                except Exception:\n                    return []\n            _tk.list_repo_templates = _safe2\n            print(\"[Hotfix] tokenization_utils_base.list_repo_templates patched.\")\n    except Exception as e:\n        print(f\"[Hotfix] tokenization patch skipped: {e}\")\n\n_install_chat_template_hotfix()\nprint(\"Chat-template hotfix installed.\\n\")\n\nimport gzip\nfrom pathlib import Path\nfrom glob import glob\nfrom contextlib import nullcontext\n\nimport numpy as np\nfrom tqdm import tqdm\nimport torch\nfrom transformers import AutoTokenizer, AutoModel\n\n# ----------------------\n# Config\n# ----------------------\nINPUT_DIR = \"/kaggle/input/cafa-5-protein-function-prediction\"\nWORK_DIR  = \"/kaggle/working\"\nCACHE_DIR = f\"{WORK_DIR}/emb_cache\"\nos.makedirs(CACHE_DIR, exist_ok=True)\n\n# Which part to embed this run:\nMODE = os.environ.get(\"MODE\", \"train\")  # \"train\", \"test\", or \"both\"\n\n# Sharding: strongly recommend SHARDS>=16 for 650M on P100\nSHARDS   = int(os.environ.get(\"SHARDS\", \"32\"))\nSHARD_ID = int(os.environ.get(\"SHARD_ID\", \"0\"))\n\nCFG = dict(\n    seed=42,\n    device=\"cuda\" if torch.cuda.is_available() else \"cpu\",\n    esm2_name=\"facebook/esm2_t33_650M_UR50D\",\n    seq_stride=1022,\n    seq_pool=\"mean\",\n    batch_size_embed=1,      # keep 1 for safety with 650M\n    cache_fp16=True,\n)\n\ndevice = torch.device(CFG[\"device\"])\nUSE_CUDA = torch.cuda.is_available()\n\ntorch.backends.cuda.matmul.allow_tf32 = True\ntry:\n    torch.set_float32_matmul_precision(\"high\")\nexcept Exception:\n    pass\n\ndef seed_everything(seed=42):\n    import random\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\nseed_everything(CFG[\"seed\"])\n\ndef _glob_one(patterns):\n    for pat in patterns:\n        hits = sorted(glob(pat, recursive=True))\n        if hits:\n            return hits[0]\n    return None\n\nTRAIN_FASTA = _glob_one([\n    f\"{INPUT_DIR}/Train/train_sequences.fasta\",\n    f\"{INPUT_DIR}/Train/train_sequences.fa\",\n    f\"{INPUT_DIR}/**/*train*sequence*.fa*\",\n])\nTEST_FASTA  = _glob_one([\n    f\"{INPUT_DIR}/Test/testsuperset.fasta\",\n    f\"{INPUT_DIR}/Test/test_superset.fasta\",\n    f\"{INPUT_DIR}/Test/targets.fasta\",\n    f\"{INPUT_DIR}/**/*test*super*set*.fa*\",\n    f\"{INPUT_DIR}/**/*target*.fa*\",\n])\n\nprint(f\"Device: {device} | MODE={MODE} | SHARD {SHARD_ID}/{SHARDS-1}\")\nprint(\"TRAIN_FASTA:\", TRAIN_FASTA)\nprint(\"TEST_FASTA :\", TEST_FASTA)\n\nassert TRAIN_FASTA or TEST_FASTA, \"No FASTA files found.\"\n\n# ----------------------\n# Dtype + AMP\n# ----------------------\ndef _choose_dtype():\n    if not USE_CUDA:\n        return torch.float32\n    if torch.cuda.is_bf16_supported():\n        return torch.bfloat16\n    return torch.float16\n\nWEIGHT_DTYPE = _choose_dtype()\nAMP_ENABLED = USE_CUDA\n\ndef amp_ctx():\n    if AMP_ENABLED:\n        return torch.amp.autocast(device_type=\"cuda\", dtype=WEIGHT_DTYPE)\n    return nullcontext()\n\n# ----------------------\n# Load ESM2-650M\n# ----------------------\nprint(f\"Loading backbone: {CFG['esm2_name']} with dtype={WEIGHT_DTYPE}\")\nesm_tokenizer = AutoTokenizer.from_pretrained(\n    CFG[\"esm2_name\"],\n    use_fast=True\n)\nesm_model = AutoModel.from_pretrained(\n    CFG[\"esm2_name\"],\n    torch_dtype=WEIGHT_DTYPE,\n    low_cpu_mem_usage=True\n).to(device)\nesm_model.eval()\n\nesm_dim = getattr(getattr(esm_model, \"config\", None), \"hidden_size\", 1280)\nprint(\"esm_dim:\", esm_dim)\n\n# ----------------------\n# FASTA streaming (no big dict)\n# ----------------------\ndef fasta_records(path):\n    \"\"\"\n    Stream (id, seq) from FASTA without keeping all in RAM.\n    \"\"\"\n    if not path:\n        return\n    is_gz = str(path).endswith(\".gz\")\n    fh = gzip.open(path, \"rt\") if is_gz else open(path, \"r\")\n    ident = None\n    buf = []\n    with fh as f:\n        for line in f:\n            if not line:\n                continue\n            if line.startswith(\">\"):\n                # flush previous\n                if ident is not None and buf:\n                    seq = \"\".join(buf).strip()\n                    if seq:\n                        yield ident, seq\n                # new id\n                raw = line[1:].strip().split()[0]\n                if \"|\" in raw:\n                    parts = raw.split(\"|\")\n                    if len(parts) >= 2:\n                        raw = parts[1]\n                ident = raw\n                buf = []\n            else:\n                buf.append(line.strip())\n        # last record\n        if ident is not None and buf:\n            seq = \"\".join(buf).strip()\n            if seq:\n                yield ident, seq\n\ndef in_shard(tid: str) -> bool:\n    if SHARDS <= 1:\n        return True\n    h = (hash(tid) % SHARDS + SHARDS) % SHARDS\n    return h == SHARD_ID\n\n# ----------------------\n# Embedding helpers\n# ----------------------\ndef _chunk_ids(ids, max_len=1022, stride=1022):\n    s = 0\n    out = []\n    while s < len(ids):\n        e = min(s + max_len, len(ids))\n        out.append(ids[s:e])\n        if e == len(ids):\n            break\n        s = max(0, e - stride)\n    return out\n\n_EMBED_BS = [max(1, CFG[\"batch_size_embed\"])]\n\n@torch.inference_mode()\ndef embed_sequence(seq: str):\n    toks = esm_tokenizer(seq, add_special_tokens=False)[\"input_ids\"]\n    if len(toks) == 0:\n        return np.zeros((esm_dim,), np.float32)\n\n    chunks = _chunk_ids(toks, max_len=1022, stride=CFG[\"seq_stride\"])\n    outs = []\n    i = 0\n\n    while i < len(chunks):\n        bs = _EMBED_BS[0]\n        batch = chunks[i : i + bs]\n        built = [esm_tokenizer.build_inputs_with_special_tokens(x) for x in batch]\n        L = max(len(x) for x in built)\n\n        input_ids = torch.zeros((len(built), L), dtype=torch.long, device=device)\n        mask      = torch.zeros((len(built), L), dtype=torch.long, device=device)\n        for bi, x in enumerate(built):\n            l = len(x)\n            input_ids[bi, :l] = torch.as_tensor(x, device=device)\n            mask[bi, :l] = 1\n\n        try:\n            with amp_ctx():\n                out = esm_model(input_ids=input_ids, attention_mask=mask).last_hidden_state\n                pooled = (out * mask.unsqueeze(-1)).sum(1) / mask.sum(1, keepdim=True).clamp_min(1)\n            outs.append(pooled.detach().float().cpu().numpy())\n            del out, pooled, input_ids, mask\n            if USE_CUDA:\n                torch.cuda.empty_cache()\n            i += bs\n        except torch.cuda.OutOfMemoryError:\n            if USE_CUDA:\n                torch.cuda.empty_cache()\n            _EMBED_BS[0] = max(1, bs // 2)\n            print(f\"[embed_sequence] OOM → reducing batch_size_embed to {_EMBED_BS[0]}\")\n            if _EMBED_BS[0] == bs:\n                # stuck; return safe zero to avoid killing job\n                return np.zeros((esm_dim,), np.float32)\n\n    if not outs:\n        return np.zeros((esm_dim,), np.float32)\n    return np.vstack(outs).mean(0)\n\ndef _save_npz(path, v_seq):\n    if v_seq is None:\n        return\n    if CFG[\"cache_fp16\"]:\n        v_seq = v_seq.astype(np.float16)\n    np.savez_compressed(path, v_seq=v_seq)\n\n# ----------------------\n# Main embedding loops\n# ----------------------\ndef run_split(name, fasta_path):\n    if not fasta_path:\n        return\n    print(f\"\\n=== Embedding {name} split on shard {SHARD_ID}/{SHARDS-1} ===\")\n    # We don't know count without scanning; just stream with tqdm without total.\n    for tid, seq in tqdm(fasta_records(fasta_path), desc=f\"{name} (stream)\"):\n        if not tid or not seq:\n            continue\n        if not in_shard(tid):\n            continue\n        path = Path(CACHE_DIR) / f\"{name}_{tid}.npz\"\n        if path.exists():\n            continue\n        try:\n            v_seq = embed_sequence(seq)\n        except Exception as e:\n            print(f\"[WARN] {name} {tid} failed: {e}\")\n            continue\n        _save_npz(path, v_seq=v_seq)\n\nif MODE in (\"train\", \"both\"):\n    run_split(\"train\", TRAIN_FASTA)\n\nif MODE in (\"test\", \"both\"):\n    run_split(\"test\", TEST_FASTA)\n\nprint(\"\\n✅ Done for this shard/mode. Collect emb_cache from /kaggle/working.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}