{"metadata":{"kernelspec":{"display_name":"Python 3 (ipykernel)","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":118765,"databundleVersionId":15231210},{"sourceType":"datasetVersion","sourceId":14604295,"datasetId":9328538,"databundleVersionId":15440074},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268},{"sourceType":"kernelVersion","sourceId":299296890}],"dockerImageVersionId":31260,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Reference\n\nhttps://www.kaggle.com/code/llkh0a/stanford-rna-3d-folding-part-2-protenix-tbm","metadata":{"execution":{"iopub.execute_input":"2026-02-16T19:38:36.011993Z","iopub.status.busy":"2026-02-16T19:38:36.011656Z","iopub.status.idle":"2026-02-16T19:38:36.016065Z","shell.execute_reply":"2026-02-16T19:38:36.015433Z","shell.execute_reply.started":"2026-02-16T19:38:36.011965Z"}}},{"cell_type":"code","source":"# !pip install /kaggle/input/datasets/ogurtsov/biopython/biopython-1.85-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# SUBMISSION (SINGLE CELL) — BEST GENERALIZED TBM (A+B+C packs)\n# + meta weights + 2-stage consensus + optional Protenix fallback\n# + format-safe writer (3 decimals, no NaN/Inf)\n# ============================================================\n\nimport os, sys, gc, json, pickle, csv, math, time\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nfrom Bio.Align import PairwiseAligner\nfrom tqdm import tqdm\n\n# -------------------------\n# Paths / knobs\n# -------------------------\nDATA_ROOT = os.environ.get(\"RNA_DATA_ROOT\", \"/kaggle/input/stanford-rna-3d-folding-2\")\nTEST_CSV  = f\"{DATA_ROOT}/test_sequences.csv\"\n\nPACK_DIR = os.environ.get(\"PACK_DIR\", \"/kaggle/input/notebooks/kumarandatascientist/3-pretrained-rna-3d-models-metaensemble\")\nPACK_A_PATH = Path(PACK_DIR) / \"rna3d_tbm_pack_A.pkl\"\nPACK_B_PATH = Path(PACK_DIR) / \"rna3d_tbm_pack_B.pkl\"\nPACK_C_PATH = Path(PACK_DIR) / \"rna3d_tbm_pack_C.pkl\"\nMETA_PATH   = Path(PACK_DIR) / \"rna3d_meta_ensemble.json\"\n\nOUT_CSV = Path(os.environ.get(\"SUBMISSION_CSV\", \"/kaggle/working/submission.csv\"))\nOUT_CSV.parent.mkdir(parents=True, exist_ok=True)\n\nN_SAMPLE = int(os.environ.get(\"N_SAMPLE\", \"2\"))\nSEED     = int(os.environ.get(\"SEED\", \"42\"))\n\n# Fast candidate preselect\nKMER_K         = int(os.environ.get(\"KMER_K\", \"5\"))      # 5-mer = 1024 dims\nPRESEL_TOPN    = int(os.environ.get(\"PRESEL_TOPN\", \"250\"))  # candidates to align per target\nLEN_RATIO_MAX  = float(os.environ.get(\"LEN_RATIO_MAX\", \"0.35\"))  # global length window\n\n# Consensus\nCONSENSUS_K    = int(os.environ.get(\"CONSENSUS_K\", \"14\"))   # top templates per pack used in consensus\nCONS_ITERS     = int(os.environ.get(\"CONS_ITERS\", \"2\"))\n\n# TBM gating → Protenix usage\nUSE_PROTENIX = str(os.environ.get(\"USE_PROTENIX\", \"false\")).strip().lower() in {\"1\",\"true\",\"t\",\"yes\",\"y\",\"on\"}\nTBM_CONF_TH  = float(os.environ.get(\"TBM_CONF_TH\", \"0.14\"))  # lower => use Protenix more often\nPROTENIX_SAMPLES_FOR_HARD = int(os.environ.get(\"PROTENIX_SAMPLES_FOR_HARD\", \"2\"))\n\n# Protenix config (optional; only used if USE_PROTENIX=true)\nDEFAULT_CODE_DIR = os.environ.get(\n    \"PROTENIX_CODE_DIR\",\n    \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"\n)\nDEFAULT_ROOT_DIR = os.environ.get(\"PROTENIX_ROOT_DIR\", DEFAULT_CODE_DIR)\nMODEL_NAME    = os.environ.get(\"MODEL_NAME\", \"protenix_base_20250630_v1.0.0\")\nMAX_SEQ_LEN   = int(os.environ.get(\"MAX_SEQ_LEN\", \"512\"))\nCHUNK_OVERLAP = int(os.environ.get(\"CHUNK_OVERLAP\", \"128\"))\nUSE_MSA      = str(os.environ.get(\"USE_MSA\", \"false\")).strip().lower() in {\"1\",\"true\",\"t\",\"yes\",\"y\",\"on\"}\nUSE_TEMPLATE = str(os.environ.get(\"USE_TEMPLATE\", \"false\")).strip().lower() in {\"1\",\"true\",\"t\",\"yes\",\"y\",\"on\"}\nUSE_RNA_MSA  = str(os.environ.get(\"USE_RNA_MSA\", \"true\")).strip().lower() in {\"1\",\"true\",\"t\",\"yes\",\"y\",\"on\"}\n\nnp.random.seed(SEED)\n\nprint(\"============================================================\")\nprint(\"SUBMISSION: Packs A+B+C + meta + consensus + optional Protenix\")\nprint(\"DATA_ROOT :\", DATA_ROOT)\nprint(\"TEST_CSV  :\", TEST_CSV)\nprint(\"PACK_DIR  :\", PACK_DIR)\nprint(\"OUT_CSV   :\", str(OUT_CSV))\nprint(\"N_SAMPLE  :\", N_SAMPLE)\nprint(\"PRESEL_TOPN:\", PRESEL_TOPN, \"| CONSENSUS_K:\", CONSENSUS_K, \"| USE_PROTENIX:\", USE_PROTENIX)\nprint(\"============================================================\\n\")\n\n# -------------------------\n# Load packs + meta\n# -------------------------\ndef load_pack(p: Path):\n    if not p.exists():\n        raise FileNotFoundError(f\"Missing pack: {p}\")\n    with open(p, \"rb\") as f:\n        obj = pickle.load(f)\n    return obj\n\npackA = load_pack(PACK_A_PATH)\npackB = load_pack(PACK_B_PATH)\npackC = load_pack(PACK_C_PATH)\n\n# Use one shared template pool (they should match)\ntemplate_df = packA[\"template_df\"]\ntemplate_coord_map = packA[\"template_coord_map\"]\n\npacks = [\n    (\"A\", packA[\"pack_cfg\"]),\n    (\"B\", packB[\"pack_cfg\"]),\n    (\"C\", packC[\"pack_cfg\"]),\n]\n\nmeta_weights = {\"A\": 0.34, \"B\": 0.35, \"C\": 0.31}\nif META_PATH.exists():\n    with open(META_PATH, \"r\", encoding=\"utf-8\") as f:\n        meta = json.load(f)\n    w = meta.get(\"ensemble_weights\", {})\n    # keys are pack_name strings; map by prefix A/B/C\n    # If meta stores by pack_name, fallback to equal\n    # We support both forms:\n    if \"A\" in w and \"B\" in w and \"C\" in w:\n        meta_weights = {k: float(w[k]) for k in [\"A\",\"B\",\"C\"]}\n    else:\n        # try mapping by pack_cfg[\"pack_name\"]\n        name_to_key = {cfg[\"pack_name\"]: key for key, cfg in packs}\n        mapped = {}\n        for name, val in w.items():\n            if name in name_to_key:\n                mapped[name_to_key[name]] = float(val)\n        if set(mapped.keys()) == {\"A\",\"B\",\"C\"}:\n            meta_weights = mapped\n\n# normalize\ns = sum(meta_weights.values())\nmeta_weights = {k: v/(s+1e-12) for k,v in meta_weights.items()}\n\nprint(\"[TEMPLATE POOL]\", template_df.shape, \"| coord_map:\", len(template_coord_map))\nprint(\"[META WEIGHTS]\", meta_weights, \"\\n\")\n\n# -------------------------\n# Build aligners per pack\n# -------------------------\ndef build_aligner(params: dict) -> PairwiseAligner:\n    al = PairwiseAligner()\n    al.mode = \"global\"\n    al.match_score = float(params.get(\"match_score\", 2.0))\n    al.mismatch_score = float(params.get(\"mismatch_score\", -1.5))\n    al.open_gap_score = float(params.get(\"open_gap_score\", -8.0))\n    al.extend_gap_score = float(params.get(\"extend_gap_score\", -0.4))\n\n    al.query_left_open_gap_score  = float(params.get(\"query_left_open_gap_score\",  al.open_gap_score))\n    al.query_left_extend_gap_score= float(params.get(\"query_left_extend_gap_score\", al.extend_gap_score))\n    al.query_right_open_gap_score = float(params.get(\"query_right_open_gap_score\", al.open_gap_score))\n    al.query_right_extend_gap_score= float(params.get(\"query_right_extend_gap_score\", al.extend_gap_score))\n\n    al.target_left_open_gap_score  = float(params.get(\"target_left_open_gap_score\",  al.open_gap_score))\n    al.target_left_extend_gap_score= float(params.get(\"target_left_extend_gap_score\", al.extend_gap_score))\n    al.target_right_open_gap_score = float(params.get(\"target_right_open_gap_score\", al.open_gap_score))\n    al.target_right_extend_gap_score= float(params.get(\"target_right_extend_gap_score\", al.extend_gap_score))\n    return al\n\nALIGNERS = {key: build_aligner(cfg[\"aligner\"]) for key, cfg in packs}\n\n# -------------------------\n# Stoichiometry segmentation (keeps constraints within chains)\n# -------------------------\ndef parse_stoichiometry(stoich: str) -> list:\n    if pd.isna(stoich) or str(stoich).strip() == \"\":\n        return []\n    return [(ch.strip(), int(cnt)) for part in str(stoich).split(\";\")\n            for ch, cnt in [part.split(\":\")]]\n\ndef parse_fasta(fasta_content: str) -> dict:\n    out, cur, parts = {}, None, []\n    for line in str(fasta_content).splitlines():\n        line = line.strip()\n        if not line:\n            continue\n        if line.startswith(\">\"):\n            if cur is not None:\n                out[cur] = \"\".join(parts)\n            cur = line[1:].split()[0]\n            parts = []\n        else:\n            parts.append(line.replace(\" \", \"\"))\n    if cur is not None:\n        out[cur] = \"\".join(parts)\n    return out\n\ndef get_chain_segments(row) -> list:\n    seq    = row[\"sequence\"]\n    stoich = row.get(\"stoichiometry\", \"\")\n    all_sq = row.get(\"all_sequences\", \"\")\n    if (pd.isna(stoich) or pd.isna(all_sq)\n            or str(stoich).strip() == \"\" or str(all_sq).strip() == \"\"):\n        return [(0, len(seq))]\n    try:\n        chain_dict = parse_fasta(all_sq)\n        order = parse_stoichiometry(stoich)\n        segs, pos = [], 0\n        for ch, cnt in order:\n            base = chain_dict.get(ch)\n            if base is None:\n                return [(0, len(seq))]\n            for _ in range(cnt):\n                segs.append((pos, pos + len(base)))\n                pos += len(base)\n        return segs if pos == len(seq) else [(0, len(seq))]\n    except Exception:\n        return [(0, len(seq))]\n\n# -------------------------\n# Kabsch + apply\n# -------------------------\ndef kabsch_align(P: np.ndarray, Q: np.ndarray):\n    P = np.asarray(P, dtype=np.float64)\n    Q = np.asarray(Q, dtype=np.float64)\n    if P.shape[0] < 3:\n        t = Q.mean(0) - P.mean(0)\n        return np.eye(3), t\n    Pc = P - P.mean(0, keepdims=True)\n    Qc = Q - Q.mean(0, keepdims=True)\n    C = Pc.T @ Qc\n    V, S, Wt = np.linalg.svd(C)\n    d = np.sign(np.linalg.det(V @ Wt))\n    D = np.diag([1.0, 1.0, d])\n    R = V @ D @ Wt\n    t = Q.mean(0) - (P.mean(0) @ R)\n    return R, t\n\ndef apply_rt(X: np.ndarray, R: np.ndarray, t: np.ndarray) -> np.ndarray:\n    return (X @ R) + t\n\n# -------------------------\n# Lightweight constraints (chain-wise)\n# -------------------------\ndef constraint_refine(coords: np.ndarray, segs: list, passes: int,\n                      bond_strength: float, d2_strength: float, lap_strength: float, avoid_strength: float,\n                      confidence: float = 1.0) -> np.ndarray:\n    X = coords.astype(np.float32).copy()\n    strength = max(0.78 * (1.0 - min(float(confidence), 0.97)), 0.02)\n    for _ in range(int(passes)):\n        for s, e in segs:\n            C = X[s:e]\n            L = e - s\n            if L < 3:\n                continue\n\n            d = C[1:] - C[:-1]\n            dist = np.linalg.norm(d, axis=1) + 1e-6\n            adj  = d * ((5.95 - dist) / dist)[:, None] * (bond_strength * strength)\n            C[:-1] -= adj\n            C[1:]  += adj\n\n            if L >= 5:\n                d2 = C[2:] - C[:-2]\n                d2n = np.linalg.norm(d2, axis=1) + 1e-6\n                adj2 = d2 * ((10.2 - d2n) / d2n)[:, None] * (d2_strength * strength)\n                C[:-2] -= adj2\n                C[2:]  += adj2\n\n            C[1:-1] += (lap_strength * strength) * (0.5 * (C[:-2] + C[2:]) - C[1:-1])\n\n            if L >= 25 and avoid_strength > 0:\n                idx = np.linspace(0, L - 1, min(L, 160)).astype(int) if L > 220 else np.arange(L)\n                P = C[idx]\n                diff = P[:, None, :] - P[None, :, :]\n                dm = np.linalg.norm(diff, axis=2) + 1e-6\n                sep = np.abs(idx[:, None] - idx[None, :])\n                mask = (sep > 2) & (dm < 3.2)\n                if np.any(mask):\n                    vec = (diff * ((3.2 - dm) / dm)[:, :, None] * mask[:, :, None]).sum(axis=1)\n                    C[idx] += (avoid_strength * strength) * vec\n\n            X[s:e] = C\n\n    return np.nan_to_num(X, nan=0.0, posinf=0.0, neginf=0.0).astype(np.float32)\n\n# -------------------------\n# Diversity transforms\n# -------------------------\ndef rotmat(axis, ang):\n    a = np.asarray(axis, float)\n    a /= np.linalg.norm(a) + 1e-12\n    x, y, z = a\n    c, s = np.cos(ang), np.sin(ang)\n    CC = 1 - c\n    return np.array([\n        [c + x*x*CC, x*y*CC - z*s, x*z*CC + y*s],\n        [y*x*CC + z*s, c + y*y*CC, y*z*CC - x*s],\n        [z*x*CC - y*s, z*y*CC + x*s, c + z*z*CC],\n    ], dtype=np.float32)\n\ndef apply_hinge(coords, seg, rng, deg=20.0):\n    s, e = seg\n    L = e - s\n    if L < 30:\n        return coords\n    pivot = s + int(rng.integers(10, L - 10))\n    R = rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n    X = coords.copy()\n    p0 = X[pivot].copy()\n    X[pivot+1:e] = (X[pivot+1:e] - p0) @ R.T + p0\n    return X\n\ndef jitter_chains(coords, segs, rng, deg=12.0, trans=1.5):\n    X = coords.copy()\n    g0 = X.mean(0, keepdims=True)\n    for s, e in segs:\n        R = rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n        shift = rng.normal(size=3)\n        shift = shift / (np.linalg.norm(shift) + 1e-12) * float(rng.uniform(0, trans))\n        c = X[s:e].mean(0, keepdims=True)\n        X[s:e] = (X[s:e] - c) @ R.T + c + shift\n    X -= X.mean(0, keepdims=True) - g0\n    return X\n\ndef smooth_wiggle(coords, segs, rng, amp=0.8):\n    X = coords.copy()\n    for s, e in segs:\n        L = e - s\n        if L < 20:\n            continue\n        ctrl = np.linspace(0, L - 1, 6)\n        disp = rng.normal(0, amp, (6, 3))\n        t = np.arange(L)\n        X[s:e] += np.vstack([np.interp(t, ctrl, disp[:, k]) for k in range(3)]).T.astype(np.float32)\n    return X\n\ndef helix(seq: str, seed: int):\n    rng = np.random.default_rng(seed)\n    n = len(seq)\n    X = np.zeros((n,3), dtype=np.float32)\n    for i in range(n):\n        ang = i * 0.6\n        X[i] = [10*np.cos(ang), 10*np.sin(ang), i*2.5]\n    X += rng.normal(0, 0.05, X.shape).astype(np.float32)\n    return X\n\n# -------------------------\n# K-mer preselect index (5-mer cosine)\n# -------------------------\nVOCAB = 4 ** KMER_K\nMAP = {ord(\"A\"):0, ord(\"C\"):1, ord(\"G\"):2, ord(\"U\"):3, ord(\"T\"):3, ord(\"N\"):0}\n\ndef kmer_vec(seq: str, k: int):\n    b = np.frombuffer(seq.encode(\"ascii\", \"ignore\"), dtype=np.uint8)\n    x = np.array([MAP.get(int(ch), 0) for ch in b], dtype=np.int32)\n    if len(x) < k:\n        v = np.zeros((4**k,), dtype=np.float32)\n        return v, 0.0\n    # rolling base-4\n    base = 4\n    powk = base ** (k-1)\n    idx = 0\n    v = np.zeros((4**k,), dtype=np.float32)\n    for i in range(k):\n        idx = idx*base + x[i]\n    v[idx] += 1.0\n    for i in range(k, len(x)):\n        idx = (idx - x[i-k]*powk) * base + x[i]\n        v[idx] += 1.0\n    nrm = float(np.linalg.norm(v) + 1e-12)\n    return v, nrm\n\n# Build template arrays\ntmpl_ids  = template_df[\"target_id\"].tolist()\ntmpl_seqs = template_df[\"sequence\"].tolist()\ntmpl_lens = np.array([len(s) for s in tmpl_seqs], dtype=np.int32)\n\n# build vectors (can take ~1-2 min depending on template count)\nprint(\"[INDEX] building k-mer matrix ...\")\nt0 = time.time()\ntmpl_mat = np.zeros((len(tmpl_seqs), VOCAB), dtype=np.float32)\ntmpl_norm = np.zeros((len(tmpl_seqs),), dtype=np.float32)\nfor i, s in enumerate(tmpl_seqs):\n    v, n = kmer_vec(s, KMER_K)\n    tmpl_mat[i] = v\n    tmpl_norm[i] = n\nprint(f\"[INDEX] done. N={len(tmpl_seqs)} | vocab={VOCAB} | time={time.time()-t0:.1f}s\\n\")\n\n# -------------------------\n# Alignment stats + adapt (pack-specific aligner)\n# -------------------------\ndef score_alignment(aligner: PairwiseAligner, qseq: str, tseq: str):\n    aln = next(iter(aligner.align(qseq, tseq)))\n    # normalized score (rough)\n    norm_s = float(aln.score / (2.0 * max(1, min(len(qseq), len(tseq)))))\n    identical = 0\n    for (qs, qe), (ts, te) in zip(*aln.aligned):\n        for qp, tp in zip(range(qs, qe), range(ts, te)):\n            identical += (qseq[qp] == tseq[tp])\n    pct_id_q = float(100.0 * identical / max(1, len(qseq)))\n    return norm_s, pct_id_q, aln\n\ndef adapt_template(aligner: PairwiseAligner, qseq: str, tseq: str, tcoords: np.ndarray):\n    aln = next(iter(aligner.align(qseq, tseq)))\n    out = np.full((len(qseq), 3), np.nan, dtype=np.float32)\n    hard = np.zeros((len(qseq),), dtype=bool)\n    for (qs, qe), (ts, te) in zip(*aln.aligned):\n        chunk = tcoords[ts:te]\n        if len(chunk) == (qe - qs):\n            out[qs:qe] = chunk.astype(np.float32)\n            hard[qs:qe] = True\n\n    for i in range(len(out)):\n        if np.isnan(out[i, 0]):\n            pv = next((j for j in range(i - 1, -1, -1) if not np.isnan(out[j, 0])), -1)\n            nv = next((j for j in range(i + 1, len(out)) if not np.isnan(out[j, 0])), -1)\n            if pv >= 0 and nv >= 0:\n                w = (i - pv) / (nv - pv)\n                out[i] = (1 - w) * out[pv] + w * out[nv]\n            elif pv >= 0:\n                out[i] = out[pv] + np.array([3.0, 0.0, 0.0], dtype=np.float32)\n            elif nv >= 0:\n                out[i] = out[nv] + np.array([3.0, 0.0, 0.0], dtype=np.float32)\n            else:\n                out[i] = np.array([i*3.0, 0.0, 0.0], dtype=np.float32)\n\n    return np.nan_to_num(out, nan=0.0, posinf=0.0, neginf=0.0).astype(np.float32), hard\n\n# -------------------------\n# 2-stage consensus within a pack\n# -------------------------\ndef consensus_merge(qseq: str, segs: list, aligner: PairwiseAligner, candidates, refine_cfg, iters=2):\n    \"\"\"\n    candidates: list of tuples (tid, tseq, sim, pct_id_q) sorted desc by sim\n    returns: proto coords (L,3), confidence scalar\n    \"\"\"\n    K = min(CONSENSUS_K, len(candidates))\n    if K <= 0:\n        return None, 0.0\n\n    adapted = []\n    masks   = []\n    weights = []\n\n    for j in range(K):\n        tid, tseq, sim, pid = candidates[j]\n        if tid not in template_coord_map:\n            continue\n        cq, hard = adapt_template(aligner, qseq, tseq, template_coord_map[tid])\n        # initial weight: emphasize similarity and identity\n        w = (max(sim, 0.0) ** 2.0) * ((max(pid, 1e-3)/100.0) ** 1.3)\n        adapted.append(cq.astype(np.float64))\n        masks.append(hard)\n        weights.append(w)\n\n    if len(adapted) == 0:\n        return None, 0.0\n\n    weights = np.asarray(weights, dtype=np.float64)\n    weights = weights / (weights.sum() + 1e-12)\n\n    proto = adapted[0].copy()\n\n    for _ in range(int(iters)):\n        aligned = []\n        for X, m in zip(adapted, masks):\n            if m.sum() >= 8:\n                R, t = kabsch_align(X[m], proto[m])\n                Xa = apply_rt(X, R, t)\n            else:\n                Xa = X\n            aligned.append(Xa)\n\n        proto = np.zeros_like(aligned[0])\n        for w, Xa in zip(weights, aligned):\n            proto += w * Xa\n\n        # reweight by RMSD-to-proto (consistency)\n        rmsd = []\n        for Xa, m in zip(aligned, masks):\n            use = m if m.sum() >= 8 else slice(None)\n            d = Xa[use] - proto[use]\n            rmsd.append(float(np.sqrt((d*d).sum() / max(1, d.shape[0]))))\n        rmsd = np.asarray(rmsd, dtype=np.float64)\n        sigma = max(3.0, float(np.median(rmsd)) + 1e-6)\n        cons_w = np.exp(-(rmsd / sigma) ** 2)\n        weights = weights * cons_w\n        weights = weights / (weights.sum() + 1e-12)\n\n    proto = proto.astype(np.float32)\n    # light refine\n    proto = constraint_refine(\n        proto, segs=segs, passes=int(refine_cfg[\"passes\"]),\n        bond_strength=float(refine_cfg[\"bond_strength\"]),\n        d2_strength=float(refine_cfg[\"d2_strength\"]),\n        lap_strength=float(refine_cfg[\"lap_strength\"]),\n        avoid_strength=float(refine_cfg[\"avoid_strength\"]),\n        confidence=float(max(c[2] for c in candidates[:K])) if K else 1.0\n    )\n    conf = float(max(c[2] for c in candidates[:K])) if K else 0.0\n    return proto, conf\n\n# -------------------------\n# Protenix (optional) — chunked inference for only \"hard\" targets\n# -------------------------\ndef ensure_required_files(root_dir: str) -> None:\n    ckpt = Path(root_dir) / \"checkpoint\" / f\"{MODEL_NAME}.pt\"\n    ccd  = Path(root_dir) / \"common\" / \"components.cif\"\n    ccdp = Path(root_dir) / \"common\" / \"components.cif.rdkit_mol.pkl\"\n    missing = [p for p in [ckpt, ccd, ccdp] if not p.exists()]\n    if missing:\n        raise FileNotFoundError(f\"Missing Protenix files: {missing}\")\n\ndef build_input_json(df: pd.DataFrame, json_path: str) -> None:\n    data = [\n        {\n            \"name\": row[\"target_id\"],\n            \"covalent_bonds\": [],\n            \"sequences\": [{\"rnaSequence\": {\"sequence\": row[\"sequence\"], \"count\": 1}}],\n        }\n        for _, row in df.iterrows()\n    ]\n    with open(json_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n\ndef build_configs(input_json_path: str, dump_dir: str, model_name: str):\n    from configs.configs_base import configs as configs_base\n    from configs.configs_data import data_configs\n    from configs.configs_inference import inference_configs\n    from configs.configs_model_type import model_configs\n    from protenix.config.config import parse_configs\n\n    base = {**configs_base, **{\"data\": data_configs}, **inference_configs}\n\n    def deep_update(t, p):\n        for k, v in p.items():\n            if isinstance(v, dict) and k in t and isinstance(t[k], dict):\n                deep_update(t[k], v)\n            else:\n                t[k] = v\n\n    deep_update(base, model_configs[model_name])\n    arg_str = \" \".join([\n        f\"--model_name {model_name}\",\n        f\"--input_json_path {input_json_path}\",\n        f\"--dump_dir {dump_dir}\",\n        f\"--use_msa {str(USE_MSA).lower()}\",\n        f\"--use_template {str(USE_TEMPLATE).lower()}\",\n        f\"--use_rna_msa {str(USE_RNA_MSA).lower()}\",\n        f\"--sample_diffusion.N_sample {N_SAMPLE}\",\n        f\"--seeds {SEED}\",\n    ])\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\ndef get_c1_mask(data: dict, atom_array):\n    # robust C1' selection\n    if atom_array is not None:\n        try:\n            if hasattr(atom_array, \"centre_atom_mask\"):\n                m = atom_array.centre_atom_mask == 1\n                if hasattr(atom_array, \"is_rna\"):\n                    m = m & atom_array.is_rna\n                return m\n            if hasattr(atom_array, \"atom_name\"):\n                base = (atom_array.atom_name == \"C1'\")\n                if hasattr(atom_array, \"is_rna\"):\n                    base = base & atom_array.is_rna\n                return base\n        except Exception:\n            pass\n\n    f = data[\"input_feature_dict\"]\n    if \"centre_atom_mask\" in f:\n        return (f[\"centre_atom_mask\"] == 1).detach().cpu().numpy()\n    if \"center_atom_mask\" in f:\n        return (f[\"center_atom_mask\"] == 1).detach().cpu().numpy()\n\n    n_tokens = int(data.get(\"N_token\", 0))\n    m11 = (f.get(\"atom_to_tokatom_idx\") == 11).detach().cpu().numpy()\n    m12 = (f.get(\"atom_to_tokatom_idx\") == 12).detach().cpu().numpy()\n    if abs(int(m11.sum()) - n_tokens) < abs(int(m12.sum()) - n_tokens):\n        return m11\n    return m12\n\ndef chunk_plan(L: int, max_len: int, overlap: int):\n    if L <= max_len:\n        return [(0, L)]\n    step = max_len - overlap\n    starts = list(range(0, L, step))\n    out = []\n    for s in starts:\n        e = min(L, s + max_len)\n        out.append((s, e))\n        if e == L:\n            break\n    return out\n\ndef stitch_chunks(chunks, starts, ends, overlap):\n    full_L = ends[-1]\n    out = np.zeros((full_L,3), dtype=np.float32)\n    out[starts[0]:ends[0]] = chunks[0].astype(np.float32)\n    cur_end = ends[0]\n    for j in range(1, len(chunks)):\n        s, e = starts[j], ends[j]\n        X = chunks[j].astype(np.float64)\n\n        ov_s = max(s, cur_end - overlap)\n        ov_e = min(cur_end, e)\n        if ov_e - ov_s >= 6:\n            P = X[(ov_s - s):(ov_e - s)]\n            Q = out[ov_s:ov_e].astype(np.float64)\n            R, t = kabsch_align(P, Q)\n            X = apply_rt(X, R, t)\n        else:\n            shift = out[cur_end-1].astype(np.float64) - X[max(0, cur_end-1 - s)]\n            X = X + shift\n\n        if ov_e > ov_s:\n            a = ov_s - s\n            b = ov_e - s\n            w = np.linspace(0, 1, b-a, dtype=np.float64)[:, None]\n            out[ov_s:ov_e] = (1-w) * out[ov_s:ov_e] + w * X[a:b]\n\n        tail_s = ov_e\n        if tail_s < e:\n            out[tail_s:e] = X[(tail_s - s):(e - s)]\n        cur_end = max(cur_end, e)\n\n    return out\n\ndef run_protenix_for_targets(targets: dict, segs_map: dict):\n    \"\"\"\n    targets: tid -> (n_needed, full_seq)\n    returns: tid -> np.ndarray (n_needed, L, 3)\n    \"\"\"\n    if not targets:\n        return {}\n\n    if not os.path.isdir(DEFAULT_CODE_DIR) or not os.path.isdir(DEFAULT_ROOT_DIR):\n        print(\"[WARN] Protenix directories not found. Skipping Protenix.\")\n        return {}\n\n    sys.path.append(DEFAULT_CODE_DIR)\n    os.environ[\"PROTENIX_ROOT_DIR\"] = DEFAULT_ROOT_DIR\n    ensure_required_files(DEFAULT_ROOT_DIR)\n\n    from protenix.data.inference.infer_dataloader import InferenceDataset\n    from runner.inference import InferenceRunner, update_gpu_compatible_configs, update_inference_configs\n\n    chunk_rows = []\n    chunk_meta = {}  # cname -> (tid, s, e, L, n_needed)\n    for tid, (n_needed, seq) in targets.items():\n        L = len(seq)\n        plan = chunk_plan(L, MAX_SEQ_LEN, CHUNK_OVERLAP)\n        for ci, (s,e) in enumerate(plan):\n            cname = f\"{tid}__chunk{ci:02d}__{s}_{e}\"\n            chunk_rows.append({\"target_id\": cname, \"sequence\": seq[s:e]})\n            chunk_meta[cname] = (tid, s, e, L, int(n_needed))\n\n    chunk_df = pd.DataFrame(chunk_rows)\n    work_dir = Path(\"/kaggle/working\")\n    input_json = str(work_dir / \"protenix_chunks_input.json\")\n    build_input_json(chunk_df, input_json)\n\n    configs = build_configs(input_json, str(work_dir / \"outputs\"), MODEL_NAME)\n    configs = update_gpu_compatible_configs(configs)\n    runner = InferenceRunner(configs)\n    dataset = InferenceDataset(configs)\n\n    chunk_preds = {}\n    print(f\"\\n[PROTENIX] targets={len(targets)} chunks={len(chunk_df)}\")\n    for i in tqdm(range(len(dataset)), desc=\"Protenix-chunks\"):\n        data, atom_array, error = dataset[i]\n        cname = data.get(\"sample_name\", f\"sample_{i}\")\n        if cname not in chunk_meta:\n            continue\n        tid, s, e, L, n_needed = chunk_meta[cname]\n        if error:\n            chunk_preds[cname] = None\n            continue\n        try:\n            new_cfg = update_inference_configs(configs, int(data[\"N_token\"].item()))\n            new_cfg.sample_diffusion.N_sample = int(n_needed)\n            runner.update_model_configs(new_cfg)\n\n            pred = runner.predict(data)\n            raw = pred[\"coordinate\"]  # (n_needed, all_atoms, 3)\n            mask = get_c1_mask(data, atom_array)\n            coords = raw[:, mask, :].detach().cpu().numpy().astype(np.float32)\n\n            chunk_len = e - s\n            if coords.shape[1] != chunk_len:\n                fixed = np.zeros((coords.shape[0], chunk_len, 3), dtype=np.float32)\n                m = min(coords.shape[1], chunk_len)\n                fixed[:, :m, :] = coords[:, :m, :]\n                coords = fixed\n\n            chunk_preds[cname] = np.nan_to_num(coords, nan=0.0, posinf=0.0, neginf=0.0).astype(np.float32)\n\n        except Exception as exc:\n            print(\"  chunk fail:\", cname, exc)\n            chunk_preds[cname] = None\n        finally:\n            try:\n                del pred, raw, coords, mask, data, atom_array\n            except Exception:\n                pass\n            gc.collect()\n\n    # assemble per target\n    by_target = {tid: None for tid in targets}\n    chunks_by_tid = {}\n    for cname, meta in chunk_meta.items():\n        tid, s, e, L, n_needed = meta\n        chunks_by_tid.setdefault(tid, []).append((s,e,cname))\n    for tid in chunks_by_tid:\n        chunks_by_tid[tid].sort(key=lambda x: x[0])\n\n    for tid, (n_needed, seq) in targets.items():\n        L = len(seq)\n        plans = chunks_by_tid.get(tid, [])\n        if not plans:\n            continue\n        segs = segs_map.get(tid, [(0, L)])\n        assembled = []\n        for sidx in range(int(n_needed)):\n            ch_list, starts, ends = [], [], []\n            for (cs, ce, cname) in plans:\n                pred = chunk_preds.get(cname)\n                if pred is None:\n                    ch = helix(seq[cs:ce], seed=hash((tid,cs,sidx))%(2**32))\n                else:\n                    ch = pred[min(sidx, pred.shape[0]-1)]\n                ch_list.append(ch)\n                starts.append(cs); ends.append(ce)\n            merged = stitch_chunks(ch_list, starts, ends, CHUNK_OVERLAP)\n            merged = constraint_refine(merged, segs, passes=2,\n                                      bond_strength=0.22, d2_strength=0.10, lap_strength=0.06, avoid_strength=0.015,\n                                      confidence=0.85)\n            assembled.append(merged.astype(np.float32))\n        by_target[tid] = np.stack(assembled, axis=0).astype(np.float32)\n    return by_target\n\n# -------------------------\n# Submission writer (hard locked)\n# -------------------------\ndef expected_columns(n_sample: int):\n    cols = [\"ID\",\"resname\",\"resid\"]\n    for s in range(1, n_sample+1):\n        cols += [f\"x_{s}\", f\"y_{s}\", f\"z_{s}\"]\n    return cols\n\ndef write_submission(test_df: pd.DataFrame, pred_map: dict, out_csv: Path, n_sample: int):\n    cols = expected_columns(n_sample)\n    coord_cols = [c for c in cols if c.startswith((\"x_\",\"y_\",\"z_\"))]\n    total_rows = int(test_df[\"sequence\"].str.len().sum())\n\n    ids = [None]*total_rows\n    resname = [None]*total_rows\n    resid = np.zeros((total_rows,), dtype=np.int32)\n    coord_mat = np.zeros((total_rows, 3*n_sample), dtype=np.float32)\n\n    r = 0\n    for _, row in test_df.iterrows():\n        tid = row[\"target_id\"]\n        seq = row[\"sequence\"]\n        L = len(seq)\n\n        X = pred_map.get(tid)\n        if X is None:\n            X = np.stack([helix(seq, seed=hash((tid,k))%(2**32)) for k in range(n_sample)], axis=0).astype(np.float32)\n\n        X = np.asarray(X, dtype=np.float32)\n        if X.ndim == 2:\n            X = np.stack([X]*n_sample, axis=0)\n        if X.shape[0] < n_sample:\n            X = np.concatenate([X, np.stack([X[-1]]*(n_sample-X.shape[0]), axis=0)], axis=0)\n        X = X[:n_sample]\n\n        if X.shape[1] != L:\n            fixed = np.zeros((n_sample, L, 3), dtype=np.float32)\n            m = min(X.shape[1], L)\n            fixed[:, :m, :] = X[:, :m, :]\n            X = fixed\n\n        X = np.nan_to_num(X, nan=0.0, posinf=0.0, neginf=0.0).astype(np.float32)\n\n        for i, b in enumerate(seq):\n            ids[r+i] = f\"{tid}_{i+1}\"\n            resname[r+i] = str(b)\n            resid[r+i] = i+1\n\n        coord_mat[r:r+L, :] = X.transpose(1,0,2).reshape(L, 3*n_sample)\n        r += L\n\n    coord_mat = np.nan_to_num(coord_mat, nan=0.0, posinf=0.0, neginf=0.0).astype(np.float32)\n    coord_mat = np.clip(coord_mat, -999.999, 9999.999).astype(np.float32)\n    coord_mat = np.round(coord_mat, 3).astype(np.float32)\n\n    df = pd.DataFrame({\"ID\": ids, \"resname\": pd.Series(resname, dtype=\"string\").str.slice(0,1), \"resid\": resid})\n    for j in range(n_sample):\n        df[f\"x_{j+1}\"] = coord_mat[:, 3*j+0]\n        df[f\"y_{j+1}\"] = coord_mat[:, 3*j+1]\n        df[f\"z_{j+1}\"] = coord_mat[:, 3*j+2]\n\n    df = df[cols]\n    vals = df[coord_cols].to_numpy()\n    if not np.isfinite(vals).all():\n        bad = np.where(~np.isfinite(vals))\n        raise ValueError(f\"Non-finite coords at row={bad[0][0]} col={coord_cols[bad[1][0]]}\")\n\n    df.to_csv(out_csv, index=False, float_format=\"%.3f\", quoting=csv.QUOTE_MINIMAL)\n    print(f\"\\n✅ Saved submission: {out_csv} | rows={len(df):,} | size(MB)={out_csv.stat().st_size/(1024*1024):.3f}\")\n    print(df.head(3))\n\n# -------------------------\n# Main prediction loop\n# -------------------------\ntest_df = pd.read_csv(TEST_CSV)\nsegs_map = {r[\"target_id\"]: get_chain_segments(r) for _, r in test_df.iterrows()}\n\ndef presel_candidates(qseq: str):\n    \"\"\"Return indices of templates in length window, ranked by k-mer cosine similarity.\"\"\"\n    Lq = len(qseq)\n    lo = int((1.0 - LEN_RATIO_MAX) * Lq)\n    hi = int((1.0 + LEN_RATIO_MAX) * Lq)\n    idx = np.where((tmpl_lens >= lo) & (tmpl_lens <= hi))[0]\n    if idx.size == 0:\n        return np.arange(min(len(tmpl_seqs), PRESEL_TOPN))\n\n    qv, qn = kmer_vec(qseq, KMER_K)\n    if qn <= 0:\n        return idx[:min(idx.size, PRESEL_TOPN)]\n\n    sub = tmpl_mat[idx]\n    dot = sub @ qv\n    denom = (tmpl_norm[idx] * qn) + 1e-12\n    sim = dot / denom\n    k = min(PRESEL_TOPN, idx.size)\n    top = np.argpartition(-sim, k-1)[:k]\n    ranked = top[np.argsort(-sim[top])]\n    return idx[ranked]\n\ndef per_pack_rank(qseq: str, cand_idx, pack_key: str, pack_cfg: dict):\n    \"\"\"Compute accurate align scores for candidates for one pack.\"\"\"\n    aligner = ALIGNERS[pack_key]\n    tbm_cfg = pack_cfg[\"tbm\"]\n    # evaluate only cand_idx\n    scored = []\n    for ii in cand_idx:\n        tid = tmpl_ids[ii]\n        tseq = tmpl_seqs[ii]\n        if tid not in template_coord_map:\n            continue\n        # pack-specific length ratio limit (tighter)\n        lr = float(tbm_cfg[\"len_ratio_limit\"])\n        if abs(len(tseq) - len(qseq)) / max(len(tseq), len(qseq)) > lr:\n            continue\n        sim, pid, _ = score_alignment(aligner, qseq, tseq)\n        scored.append((tid, tseq, sim, pid))\n    scored.sort(key=lambda x: x[2], reverse=True)\n    return scored[:int(tbm_cfg[\"top_n\"])]\n\npred_map = {}\nprotenix_queue = {}\n\nprint(\"[RUN] predicting TBM across packs ...\")\nfor _, row in tqdm(test_df.iterrows(), total=len(test_df), desc=\"Targets\"):\n    tid = row[\"target_id\"]\n    qseq = row[\"sequence\"]\n    segs = segs_map.get(tid, [(0, len(qseq))])\n\n    cand_idx = presel_candidates(qseq)\n\n    pack_outputs = {}   # pack_key -> dict(proto, conf, best_single, best_conf)\n    pack_conf_score = []  # (pack_key, confscore)\n\n    for key, cfg in packs:\n        ranked = per_pack_rank(qseq, cand_idx, key, cfg)\n        tbm_cfg = cfg[\"tbm\"]\n        ref_cfg = cfg[\"refine\"]\n        div_cfg = cfg[\"diversity\"]\n\n        # filter by thresholds\n        good = []\n        for (ttid, tseq, sim, pid) in ranked:\n            if sim < float(tbm_cfg[\"min_similarity\"]) or pid < float(tbm_cfg[\"min_pct_identity\"]):\n                continue\n            good.append((ttid, tseq, sim, pid))\n        if len(good) == 0:\n            # no template for this pack\n            pack_outputs[key] = {\"proto\": None, \"conf\": 0.0, \"best\": None, \"best_conf\": 0.0, \"good\": []}\n            pack_conf_score.append((key, 0.0))\n            continue\n\n        # best single template refined\n        ttid, tseq, sim, pid = good[0]\n        base, _hard = adapt_template(ALIGNERS[key], qseq, tseq, template_coord_map[ttid])\n        best_single = constraint_refine(\n            base, segs=segs, passes=int(ref_cfg[\"passes\"]),\n            bond_strength=float(ref_cfg[\"bond_strength\"]),\n            d2_strength=float(ref_cfg[\"d2_strength\"]),\n            lap_strength=float(ref_cfg[\"lap_strength\"]),\n            avoid_strength=float(ref_cfg[\"avoid_strength\"]),\n            confidence=float(sim)\n        )\n\n        # consensus proto\n        proto, conf = consensus_merge(qseq, segs, ALIGNERS[key], good, ref_cfg, iters=CONS_ITERS)\n        if proto is None:\n            proto = best_single\n            conf = float(sim)\n\n        # pack confidence (mix sim + identity)\n        confscore = float(0.75*conf + 0.25*(pid/100.0))\n\n        pack_outputs[key] = {\"proto\": proto, \"conf\": float(conf), \"best\": best_single, \"best_conf\": float(sim), \"good\": good}\n        pack_conf_score.append((key, confscore))\n\n    pack_conf_score.sort(key=lambda x: x[1], reverse=True)\n    best_pack = pack_conf_score[0][0]\n    second_pack = pack_conf_score[1][0] if len(pack_conf_score) > 1 else best_pack\n\n    # overall TBM confidence = best pack score\n    tbm_conf = float(pack_conf_score[0][1])\n\n    # Build 5 final candidates across packs\n    cands = []\n\n    # 1) Best pack consensus proto\n    P0 = pack_outputs[best_pack][\"proto\"]\n    if P0 is None:\n        P0 = helix(qseq, seed=hash((tid,\"helix0\"))%(2**32))\n    cands.append(P0)\n\n    # 2) Second pack consensus proto (diversity across aligners)\n    P1 = pack_outputs[second_pack][\"proto\"]\n    if P1 is None:\n        P1 = helix(qseq, seed=hash((tid,\"helix1\"))%(2**32))\n    cands.append(P1)\n\n    # 3) Meta-weighted average of (A,B,C) protos AFTER aligning to best pack proto\n    ref = P0.astype(np.float64)\n    blend = np.zeros_like(ref, dtype=np.float64)\n    for key in [\"A\",\"B\",\"C\"]:\n        X = pack_outputs[key][\"proto\"]\n        if X is None:\n            continue\n        X = X.astype(np.float64)\n        R, t = kabsch_align(X, ref)\n        Xa = apply_rt(X, R, t)\n        blend += float(meta_weights.get(key, 0.0)) * Xa\n    if np.abs(blend).sum() < 1e-6:\n        blend = ref.copy()\n    blend = blend.astype(np.float32)\n    # light refine (use best pack refine params)\n    ref_cfg = dict([p for p in packs if p[0]==best_pack][0][1][\"refine\"]) if True else packB[\"pack_cfg\"][\"refine\"]\n    blend = constraint_refine(\n        blend, segs=segs, passes=int(ref_cfg[\"passes\"]),\n        bond_strength=float(ref_cfg[\"bond_strength\"]),\n        d2_strength=float(ref_cfg[\"d2_strength\"]),\n        lap_strength=float(ref_cfg[\"lap_strength\"]),\n        avoid_strength=float(ref_cfg[\"avoid_strength\"]),\n        confidence=float(pack_outputs[best_pack][\"conf\"])\n    )\n    cands.append(blend)\n\n    # 4) Best pack best-single-template refined (often different from proto)\n    Pbest = pack_outputs[best_pack][\"best\"]\n    if Pbest is None:\n        Pbest = helix(qseq, seed=hash((tid,\"helix2\"))%(2**32))\n    cands.append(Pbest)\n\n    # 5) Structured diversity from best pack proto (hinge/jitter/wiggle)\n    rng = np.random.default_rng(hash((tid, \"div\")) % (2**32))\n    div_cfg = dict([p for p in packs if p[0]==best_pack][0][1][\"diversity\"])\n    longest = max(segs, key=lambda se: se[1]-se[0])\n    Pdiv = P0.copy()\n    # choose a deterministic transform depending on id hash\n    mode = (hash(tid) % 3)\n    if mode == 0:\n        Pdiv = apply_hinge(Pdiv, longest, rng, deg=float(div_cfg.get(\"hinge_deg\", 22.0)))\n    elif mode == 1:\n        Pdiv = jitter_chains(Pdiv, segs, rng, deg=float(div_cfg.get(\"jitter_deg\", 12.0)), trans=float(div_cfg.get(\"jitter_trans\", 1.5)))\n    else:\n        Pdiv = smooth_wiggle(Pdiv, segs, rng, amp=float(div_cfg.get(\"wiggle_amp\", 0.8)))\n    ref_cfg = dict([p for p in packs if p[0]==best_pack][0][1][\"refine\"])\n    Pdiv = constraint_refine(\n        Pdiv, segs=segs, passes=int(ref_cfg[\"passes\"]),\n        bond_strength=float(ref_cfg[\"bond_strength\"]),\n        d2_strength=float(ref_cfg[\"d2_strength\"]),\n        lap_strength=float(ref_cfg[\"lap_strength\"]),\n        avoid_strength=float(ref_cfg[\"avoid_strength\"]),\n        confidence=float(pack_outputs[best_pack][\"conf\"])\n    )\n    cands.append(Pdiv)\n\n    # Gate Protenix: if TBM confidence low, replace last 1–2 candidates with Protenix later\n    need_ptx = USE_PROTENIX and (tbm_conf < TBM_CONF_TH)\n    if need_ptx:\n        n_replace = min(PROTENIX_SAMPLES_FOR_HARD, 2)  # replace up to 2\n        # mark queue\n        protenix_queue[tid] = (n_replace, qseq)\n        # temporarily keep placeholders; we will overwrite after Protenix run\n        # (keep list length 5)\n        # no-op here\n\n    pred_map[tid] = np.stack(cands[:N_SAMPLE], axis=0).astype(np.float32)\n\n# Run Protenix only for queued targets and patch into pred_map\nif USE_PROTENIX and len(protenix_queue) > 0:\n    print(f\"\\n[GATE] Protenix queued targets = {len(protenix_queue)} (tbm_conf < {TBM_CONF_TH})\")\n    ptx = run_protenix_for_targets(protenix_queue, segs_map)\n    # Replace last slots with Protenix samples\n    for tid, arr in ptx.items():\n        if arr is None:\n            continue\n        # arr: (n_replace, L,3)\n        X = pred_map[tid]\n        n_replace = min(arr.shape[0], N_SAMPLE)\n        # Replace from the end\n        for j in range(n_replace):\n            X[-(j+1)] = arr[j]\n        pred_map[tid] = np.nan_to_num(X, nan=0.0, posinf=0.0, neginf=0.0).astype(np.float32)\n\n# Write submission\nwrite_submission(test_df, pred_map, OUT_CSV, N_SAMPLE)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}