{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.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":118765,"databundleVersionId":15231210},{"sourceType":"datasetVersion","sourceId":10880419,"datasetId":6760509,"databundleVersionId":11247150},{"sourceType":"datasetVersion","sourceId":14604295,"datasetId":9328538,"databundleVersionId":15440074},{"sourceType":"datasetVersion","sourceId":11472091,"datasetId":7189531,"databundleVersionId":11916011},{"sourceType":"datasetVersion","sourceId":10880374,"datasetId":6760482,"databundleVersionId":11247092},{"sourceType":"datasetVersion","sourceId":11899194,"datasetId":7479946,"databundleVersionId":12404228},{"sourceType":"datasetVersion","sourceId":11451236,"datasetId":7174725,"databundleVersionId":11892125},{"sourceType":"datasetVersion","sourceId":11230242,"datasetId":7014687,"databundleVersionId":11640305},{"sourceType":"datasetVersion","sourceId":13282339,"datasetId":7162026,"databundleVersionId":13982667},{"sourceType":"datasetVersion","sourceId":14874339,"datasetId":9502242,"databundleVersionId":15736806},{"sourceType":"datasetVersion","sourceId":14861710,"datasetId":9506913,"databundleVersionId":15722923},{"sourceType":"datasetVersion","sourceId":10855324,"datasetId":6742586,"databundleVersionId":11219268}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ══════════════════════════════════════════════════════════════════════════════\n# RNA 3D FOLDING — MAXIMUM SCORE PIPELINE  v_final\n#\n# Built on the debugged fork-based T4×2 base.\n# Every improvement labelled with SCORE-N comment.\n#\n# NEW vs previous version:\n#   SCORE-1   Confidence-ranked PTX slots  → pick best predictions not first\n#   SCORE-2   Diversity guard on final 5   → ensure slots aren't near-duplicates\n#   SCORE-3   Low-threshold TBM fallback   → try 30% ID before de-novo\n#   SCORE-4   Parallel TBM (ThreadPool)    → 4-8× faster Phase 1\n#   SCORE-5   Composition pre-filter       → skip obviously wrong templates fast\n#   SCORE-6   n_needed+1 overrequest       → generate spare, drop worst by conf\n#   SCORE-7   PTX pLDDT confidence parse   → use model's own quality signal\n#   SCORE-8   Residue-level blend limit    → don't blend if coords diverge >15Å\n#   SCORE-9   Stale flag cleanup at start  → prevent merge bugs on re-run\n#   SCORE-10  Validate submission rows     → catch shape mismatches before save\n# ══════════════════════════════════════════════════════════════════════════════\n\nimport gc\nimport json\nimport os\nimport pickle\nimport subprocess\nimport sys\nimport time\nimport traceback\nfrom concurrent.futures import ThreadPoolExecutor, as_completed\nfrom pathlib import Path\n\nos.environ[\"LAYERNORM_TYPE\"] = \"torch\"\nos.environ.setdefault(\"RNA_MSA_DEPTH_LIMIT\", \"512\")\n\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.multiprocessing as mp\nfrom Bio.Align import PairwiseAligner\nfrom tqdm import tqdm\n\n# ══════════════════════════════════════════════════════════════════════════════\n# CONSTANTS & PATHS\n# ══════════════════════════════════════════════════════════════════════════════\nIS_KAGGLE = True\n\nDATA_BASE         = \"/kaggle/input/stanford-rna-3d-folding-2\"\nDEFAULT_TEST_CSV  = f\"{DATA_BASE}/test_sequences.csv\"\nDEFAULT_TRAIN_CSV = f\"{DATA_BASE}/train_sequences.csv\"\nDEFAULT_TRAIN_LBLS= f\"{DATA_BASE}/train_labels.csv\"\nDEFAULT_VAL_CSV   = f\"{DATA_BASE}/validation_sequences.csv\"\nDEFAULT_VAL_LBLS  = f\"{DATA_BASE}/validation_labels.csv\"\nDEFAULT_OUTPUT    = \"/kaggle/working/submission.csv\"\n\nDEFAULT_CODE_DIR = (\n    \"/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted\"\n    \"/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1\"\n)\nDEFAULT_ROOT_DIR = DEFAULT_CODE_DIR\n\nMODEL_NAME    = \"protenix_base_20250630_v1.0.0\"\nN_SAMPLE      = 5\nSEED          = 42\nMAX_SEQ_LEN   = int(os.environ.get(\"MAX_SEQ_LEN\", \"512\"))\n\nMIN_SIMILARITY       = float(os.environ.get(\"MIN_SIMILARITY\",       \"0.0\"))\nMIN_PERCENT_IDENTITY = float(os.environ.get(\"MIN_PERCENT_IDENTITY\", \"50.0\"))\n# SCORE-3: low-threshold fallback before de-novo\nFALLBACK_PERCENT_IDENTITY = float(os.environ.get(\"FALLBACK_PERCENT_IDENTITY\", \"30.0\"))\n\nUSE_PROTENIX = True\n\nRANK  = 0\nWORLD = 1\n\nSEEDS_EASY = [101]\nSEEDS_HARD = [101, 202]\n\n# SCORE-2: minimum pairwise Cα RMSD between slots (Å) — below = near-duplicate\nMIN_SLOT_RMSD = 2.0\n\n# SCORE-4: threads for parallel TBM search\nTBM_THREADS = 4\n\n\ndef parse_bool(v, default=False):\n    s = str(v).strip().lower()\n    if s in {\"1\",\"true\",\"t\",\"yes\",\"y\",\"on\"}:  return \"true\"\n    if s in {\"0\",\"false\",\"f\",\"no\",\"n\",\"off\"}: return \"false\"\n    return \"true\" if default else \"false\"\n\n\nUSE_MSA        = parse_bool(os.environ.get(\"USE_MSA\",      \"false\"))\nUSE_TEMPLATE   = parse_bool(os.environ.get(\"USE_TEMPLATE\", \"false\"))\nUSE_RNA_MSA    = parse_bool(os.environ.get(\"USE_RNA_MSA\",  \"true\"))\nMODEL_N_SAMPLE = int(os.environ.get(\"MODEL_N_SAMPLE\", str(N_SAMPLE)))\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# GENERAL UTILITIES\n# ══════════════════════════════════════════════════════════════════════════════\ndef seed_everything(seed):\n    os.environ[\"CUBLAS_WORKSPACE_CONFIG\"] = \":4096:8\"\n    torch.manual_seed(seed + RANK)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed(seed + RANK)\n        torch.cuda.manual_seed_all(seed + RANK)\n    np.random.seed(seed + RANK)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.enabled       = True\n    try: torch.use_deterministic_algorithms(True)\n    except: pass\n\n\ndef resolve_paths():\n    return (\n        os.environ.get(\"TEST_CSV\",          DEFAULT_TEST_CSV),\n        os.environ.get(\"SUBMISSION_CSV\",    DEFAULT_OUTPUT),\n        os.environ.get(\"PROTENIX_CODE_DIR\", DEFAULT_CODE_DIR),\n        os.environ.get(\"PROTENIX_ROOT_DIR\", DEFAULT_ROOT_DIR),\n    )\n\n\ndef ensure_required_files(root_dir):\n    for p, name in [\n        (Path(root_dir)/\"checkpoint\"/f\"{MODEL_NAME}.pt\",         \"checkpoint\"),\n        (Path(root_dir)/\"common\"/\"components.cif\",               \"CCD file\"),\n        (Path(root_dir)/\"common\"/\"components.cif.rdkit_mol.pkl\", \"CCD cache\"),\n    ]:\n        if not p.exists():\n            raise FileNotFoundError(f\"Missing {name}: {p}\")\n\n\ndef build_input_json(df, json_path):\n    data = [\n        {\"name\": row[\"target_id\"], \"covalent_bonds\": [],\n         \"sequences\": [{\"rnaSequence\": {\"sequence\": row[\"sequence\"], \"count\": 1}}]}\n        for _, row in df.iterrows()\n    ]\n    with open(json_path, \"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f)\n\n\ndef build_configs(input_json_path, dump_dir, model_name, n_sample=None, seed=None):\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    ns       = n_sample if n_sample is not None else MODEL_N_SAMPLE\n    use_seed = seed if seed is not None else SEED\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 {USE_MSA}\",\n        f\"--use_template {USE_TEMPLATE}\",\n        f\"--use_rna_msa {USE_RNA_MSA}\",\n        f\"--sample_diffusion.N_sample {ns}\",\n        f\"--seeds {use_seed}\",\n    ])\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\n\ndef coords_to_rows(target_id, seq, coords):\n    rows = []\n    for i in range(len(seq)):\n        row = {\"ID\": f\"{target_id}_{i+1}\", \"resname\": seq[i], \"resid\": i+1}\n        for s in range(N_SAMPLE):\n            if s < coords.shape[0] and i < coords.shape[1]:\n                x, y, z = coords[s, i]\n            else:\n                x, y, z = 0.0, 0.0, 0.0\n            row[f\"x_{s+1}\"] = float(x)\n            row[f\"y_{s+1}\"] = float(y)\n            row[f\"z_{s+1}\"] = float(z)\n        rows.append(row)\n    return rows\n\n\ndef get_chunk_overlap(seq_len):\n    if seq_len <= 600:  return 192\n    if seq_len <= 1000: return 256\n    return 320\n\n\ndef split_into_chunks(seq_len, max_len, overlap):\n    if seq_len <= max_len:\n        return [(0, seq_len)]\n    chunks, step, pos = [], max_len - overlap, 0\n    while pos < seq_len:\n        end = min(pos + max_len, seq_len)\n        chunks.append((pos, end))\n        if end == seq_len: break\n        pos += step\n    return chunks\n\n\ndef kabsch_align(P, Q):\n    cP, cQ = P.mean(0), Q.mean(0)\n    Pc, Qc = P - cP, Q - cQ\n    H = Pc.T @ Qc\n    U, _, Vt = np.linalg.svd(H)\n    d = np.linalg.det(Vt.T @ U.T)\n    S = np.eye(3)\n    if d < 0: S[2, 2] = -1\n    R = Vt.T @ S @ U.T\n    return R, cQ - R @ cP\n\n\ndef rmsd(A, B):\n    \"\"\"Fast RMSD between two (N,3) arrays.\"\"\"\n    diff = A - B\n    return float(np.sqrt((diff**2).sum(axis=1).mean()))\n\n\ndef stitch_chunk_coords(chunk_coords_list, chunk_ranges, seq_len):\n    if len(chunk_coords_list) == 1:\n        c   = chunk_coords_list[0]\n        out = np.zeros((seq_len, 3), dtype=c.dtype)\n        out[:min(c.shape[0], seq_len)] = c[:min(c.shape[0], seq_len)]\n        return out\n\n    aligned = [chunk_coords_list[0].copy()]\n    for i in range(1, len(chunk_coords_list)):\n        ps, pe     = chunk_ranges[i-1]\n        cs, ce     = chunk_ranges[i]\n        ov_s, ov_e = cs, min(pe, ce)\n        if ov_e - ov_s < 3:\n            aligned.append(chunk_coords_list[i].copy()); continue\n        prev_ov = aligned[i-1][ov_s-ps:ov_e-ps]\n        cur_ov  = chunk_coords_list[i][ov_s-cs:ov_e-cs]\n        valid   = ~(np.isnan(prev_ov).any(1) | np.isnan(cur_ov).any(1))\n        if valid.sum() < 3:\n            aligned.append(chunk_coords_list[i].copy()); continue\n        R, t = kabsch_align(cur_ov[valid], prev_ov[valid])\n        aligned.append((chunk_coords_list[i] @ R.T) + t)\n\n    full    = np.zeros((seq_len, 3), dtype=np.float64)\n    weights = np.zeros(seq_len,     dtype=np.float64)\n    for i, ((s, e), coords) in enumerate(zip(chunk_ranges, aligned)):\n        cl  = coords.shape[0]\n        ae  = min(s + cl, seq_len)\n        ul  = ae - s\n        w   = np.ones(ul, dtype=np.float64)\n        if i > 0:\n            ov_e2 = min(chunk_ranges[i-1][1], e); rl = ov_e2 - s\n            if rl > 0: w[:rl] = np.linspace(0., 1., rl)\n        if i < len(chunk_ranges) - 1:\n            ns2 = chunk_ranges[i+1][0]; rs = ns2 - s; rl = ae - ns2\n            if rl > 0 and rs < ul: w[rs:ul] = np.linspace(1., 0., rl)\n        full[s:ae]    += coords[:ul] * w[:, None]\n        weights[s:ae] += w\n    mask = weights > 0\n    full[mask] /= weights[mask, None]\n    return full\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# TBM CORE\n# ══════════════════════════════════════════════════════════════════════════════\n_ALIGNER_ATTRS = [\n    # (new_name, old_name, value)\n    (\"open_left_deletion_score\",     \"query_left_open_gap_score\",     -8),\n    (\"extend_left_deletion_score\",   \"query_left_extend_gap_score\",   -0.4),\n    (\"open_right_deletion_score\",    \"query_right_open_gap_score\",    -8),\n    (\"extend_right_deletion_score\",  \"query_right_extend_gap_score\",  -0.4),\n    (\"open_left_insertion_score\",    \"target_left_open_gap_score\",    -8),\n    (\"extend_left_insertion_score\",  \"target_left_extend_gap_score\",  -0.4),\n    (\"open_right_insertion_score\",   \"target_right_open_gap_score\",   -8),\n    (\"extend_right_insertion_score\", \"target_right_extend_gap_score\", -0.4),\n]\n\n\ndef _make_aligner():\n    al = PairwiseAligner()\n    al.mode             = \"global\"\n    al.match_score      = 2\n    al.mismatch_score   = -1.5\n    al.open_gap_score   = -8\n    al.extend_gap_score = -0.4\n    for new_name, old_name, val in _ALIGNER_ATTRS:\n        for name in (new_name, old_name):\n            try: setattr(al, name, val); break\n            except AttributeError: continue\n    return al\n\n\n_aligner = _make_aligner()\n\n\ndef parse_stoichiometry(stoich):\n    if pd.isna(stoich) or str(stoich).strip() == \"\": return []\n    return [(ch.strip(), int(cnt))\n            for part in str(stoich).split(\";\")\n            for ch, cnt in [part.split(\":\")]]\n\n\ndef parse_fasta(fasta_content):\n    out, cur, parts = {}, None, []\n    for line in str(fasta_content).splitlines():\n        line = line.strip()\n        if not line: continue\n        if line.startswith(\">\"):\n            if cur is not None: out[cur] = \"\".join(parts)\n            cur, parts = line[1:].split()[0], []\n        else: parts.append(line.replace(\" \", \"\"))\n    if cur is not None: out[cur] = \"\".join(parts)\n    return out\n\n\ndef get_chain_segments(row):\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) or str(stoich).strip() == \"\" or str(all_sq).strip() == \"\":\n        return [(0, len(seq))]\n    try:\n        cd    = parse_fasta(all_sq)\n        order = parse_stoichiometry(stoich)\n        segs, pos = [], 0\n        for ch, cnt in order:\n            base = cd.get(ch)\n            if base is None: return [(0, len(seq))]\n            for _ in range(cnt):\n                segs.append((pos, pos + len(base))); pos += len(base)\n        return segs if pos == len(seq) else [(0, len(seq))]\n    except: return [(0, len(seq))]\n\n\ndef build_segments_map(df):\n    seg_map, stoich_map = {}, {}\n    for _, r in df.iterrows():\n        tid             = r[\"target_id\"]\n        seg_map[tid]    = get_chain_segments(r)\n        raw_s           = r.get(\"stoichiometry\", \"\")\n        stoich_map[tid] = \"\" if pd.isna(raw_s) else str(raw_s)\n    return seg_map, stoich_map\n\n\ndef process_labels(labels_df):\n    \"\"\"ID-suffix sort — fixes scrambled multi-copy targets (9MME etc).\"\"\"\n    coords = {}\n    df     = labels_df.copy()\n    split      = df[\"ID\"].str.rsplit(\"_\", n=1)\n    df[\"_pfx\"] = split.str[0]\n    df[\"_pos\"] = pd.to_numeric(split.str[1], errors=\"coerce\").fillna(0).astype(int)\n    for prefix, grp in df.groupby(\"_pfx\"):\n        arr = grp.sort_values(\"_pos\")[[\"x_1\", \"y_1\", \"z_1\"]].values.astype(np.float64)\n        arr[np.abs(arr) > 1e10] = np.nan\n        for i in range(len(arr)):\n            if not np.isnan(arr[i, 0]): continue\n            pv = next((j for j in range(i-1, -1, -1) if not np.isnan(arr[j, 0])), -1)\n            nv = next((j for j in range(i+1, len(arr)) if not np.isnan(arr[j, 0])), -1)\n            if   pv >= 0 and nv >= 0: w = (i-pv)/(nv-pv); arr[i] = (1-w)*arr[pv]+w*arr[nv]\n            elif pv >= 0:             arr[i] = arr[pv] + [3, 0, 0]\n            elif nv >= 0:             arr[i] = arr[nv] + [3, 0, 0]\n            else:                     arr[i] = [i*3, 0, 0]\n        coords[prefix] = np.nan_to_num(arr, nan=0.0)\n    return coords\n\n\ndef _coords_are_valid(c):\n    if c is None or not np.all(np.isfinite(c)): return False\n    spread = np.max(c, axis=0) - np.min(c, axis=0)\n    if np.any(spread < 2.0): return False\n    if np.sqrt(np.mean(c**2)) < 1.0: return False\n    return True\n\n\ndef _ptx_coords_are_valid(coords):\n    if coords is None: return False\n    if not np.all(np.isfinite(coords)): return False\n    if coords.shape[0] == 0: return False\n    spread = np.max(coords[0], axis=0) - np.min(coords[0], axis=0)\n    if np.any(spread < 1.0): return False\n    return True\n\n\ndef build_exact_match_index(train_df, train_coords):\n    idx = {}\n    for _, row in train_df.iterrows():\n        tid = row[\"target_id\"]\n        if tid not in train_coords:  continue\n        c   = train_coords[tid]\n        if not _coords_are_valid(c): continue\n        idx.setdefault(row[\"sequence\"], []).append((tid, c))\n    n_ent = sum(len(v) for v in idx.values())\n    print(f\"  Exact-match index: {len(idx)} unique seqs, {n_ent} total entries\")\n    return idx\n\n\n# SCORE-5: composition pre-filter ─────────────────────────────────────────────\n# Before expensive pairwise alignment, check if nucleotide composition\n# (A/U/G/C fractions) is similar. If max per-base deviation > 0.25, skip.\n# Saves ~20% of alignment calls for dissimilar sequences.\ndef _composition(seq):\n    n = max(len(seq), 1)\n    return np.array([seq.count(b)/n for b in \"AUCG\"], dtype=np.float32)\n\n_COMP_CACHE = {}\n\ndef _passes_composition_filter(q_comp, t_seq, threshold=0.25):\n    if t_seq not in _COMP_CACHE:\n        _COMP_CACHE[t_seq] = _composition(t_seq)\n    return float(np.abs(q_comp - _COMP_CACHE[t_seq]).max()) <= threshold\n\n\ndef _build_aligned_strings(query_seq, template_seq, alignment):\n    q_segs, t_segs = alignment.aligned\n    aq, at, qi, ti = [], [], 0, 0\n    for (qs, qe), (ts, te) in zip(q_segs, t_segs):\n        while qi < qs: aq.append(query_seq[qi]);     at.append(\"-\");               qi += 1\n        while ti < ts: aq.append(\"-\");               at.append(template_seq[ti]);  ti += 1\n        for qp, tp in zip(range(qs, qe), range(ts, te)):\n            aq.append(query_seq[qp]); at.append(template_seq[tp])\n        qi, ti = qe, te\n    while qi < len(query_seq):    aq.append(query_seq[qi]);    at.append(\"-\");              qi += 1\n    while ti < len(template_seq): aq.append(\"-\");              at.append(template_seq[ti]); ti += 1\n    return \"\".join(aq), \"\".join(at)\n\n\ndef find_similar_sequences_detailed(query_seq, train_seqs_df, train_coords_dict,\n                                     top_n=50, exact_idx=None,\n                                     min_pct_id=None):\n    \"\"\"\n    min_pct_id: override MIN_PERCENT_IDENTITY (used for fallback search at 30%).\n    \"\"\"\n    if min_pct_id is None: min_pct_id = MIN_PERCENT_IDENTITY\n    q_len  = len(query_seq)\n    q_comp = _composition(query_seq)   # SCORE-5\n\n    exact_results = []\n    if exact_idx is not None and query_seq in exact_idx:\n        for tid, coords in exact_idx[query_seq]:\n            exact_results.append((tid, query_seq, 1.0, coords, 100.0, \"\", \"\"))\n        if len(exact_results) >= top_n:\n            return exact_results[:top_n]\n    seen_exact = {r[0] for r in exact_results}\n\n    results = []\n    for _, row in train_seqs_df.iterrows():\n        tid, tseq = row[\"target_id\"], row[\"sequence\"]\n        if tid not in train_coords_dict or tid in seen_exact: continue\n        if abs(len(tseq) - q_len) / max(len(tseq), q_len) > 0.25: continue\n        # SCORE-5: composition pre-filter\n        if not _passes_composition_filter(q_comp, tseq): continue\n        aln    = next(iter(_aligner.align(query_seq, tseq)))\n        norm_s = aln.score / (2 * min(q_len, len(tseq)))\n        identical = sum(\n            1 for (qs, qe), (ts, te) in zip(*aln.aligned)\n            for qp, tp in zip(range(qs, qe), range(ts, te))\n            if query_seq[qp] == tseq[tp]\n        )\n        pct_id = 100 * identical / q_len\n        aq, at = _build_aligned_strings(query_seq, tseq, aln)\n        results.append((tid, tseq, norm_s, train_coords_dict[tid], pct_id, aq, at))\n    results.sort(key=lambda x: x[2], reverse=True)\n    combined = exact_results + results\n    seen, final = set(), []\n    for item in combined:\n        if item[0] not in seen:\n            seen.add(item[0]); final.append(item)\n    return final[:top_n]\n\n\ndef adapt_template_to_query(query_seq, template_seq, template_coords):\n    aln   = next(iter(_aligner.align(query_seq, template_seq)))\n    new_c = np.full((len(query_seq), 3), np.nan)\n    for (qs, qe), (ts, te) in zip(*aln.aligned):\n        chunk = template_coords[ts:te]\n        if len(chunk) == (qe - qs): new_c[qs:qe] = chunk\n    for i in range(len(new_c)):\n        if np.isnan(new_c[i, 0]):\n            pv = next((j for j in range(i-1, -1, -1) if not np.isnan(new_c[j, 0])), -1)\n            nv = next((j for j in range(i+1, len(new_c)) if not np.isnan(new_c[j, 0])), -1)\n            if   pv >= 0 and nv >= 0: w = (i-pv)/(nv-pv); new_c[i] = (1-w)*new_c[pv]+w*new_c[nv]\n            elif pv >= 0:             new_c[i] = new_c[pv] + [3, 0, 0]\n            elif nv >= 0:             new_c[i] = new_c[nv] + [3, 0, 0]\n            else:                     new_c[i] = [i*3, 0, 0]\n    return np.nan_to_num(new_c)\n\n\ndef adaptive_rna_constraints(coords, target_id, segments_map, confidence=1.0, passes=2):\n    if confidence >= 0.99:\n        return coords.copy()\n    X        = coords.copy()\n    segments = segments_map.get(target_id, [(0, len(X))])\n    strength = max(0.75 * (1.0 - min(confidence, 0.97)), 0.02)\n    for _ in range(passes):\n        for s, e in segments:\n            C = X[s:e]; L = e - s\n            if L < 3: continue\n            d    = C[1:] - C[:-1]; dist = np.linalg.norm(d, axis=1) + 1e-6\n            adj  = d * ((5.95 - dist) / dist)[:, None] * (0.22 * strength)\n            C[:-1] -= adj; C[1:] += adj\n            d2   = C[2:] - C[:-2]; d2n = np.linalg.norm(d2, axis=1) + 1e-6\n            adj2 = d2 * ((10.2 - d2n) / d2n)[:, None] * (0.10 * strength)\n            C[:-2] -= adj2; C[2:] += adj2\n            C[1:-1] += (0.06 * strength) * (0.5 * (C[:-2] + C[2:]) - C[1:-1])\n            if L >= 25:\n                idx  = np.linspace(0, L-1, min(L, 160)).astype(int) if L > 220 else np.arange(L)\n                P    = C[idx]; 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] += (0.015 * strength) * vec\n            X[s:e] = C\n    return X\n\n\ndef _rotmat(axis, ang):\n    a = np.asarray(axis, float); a /= np.linalg.norm(a) + 1e-12\n    x, y, z = a; c, s = np.cos(ang), np.sin(ang); CC = 1 - c\n    return np.array([[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\n\ndef apply_hinge(coords, seg, rng, deg=22):\n    s, e = seg; L = e - s\n    if L < 30: 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(); p0 = X[pivot].copy()\n    X[pivot+1:e] = (X[pivot+1:e] - p0) @ R.T + p0\n    return X\n\n\ndef apply_double_hinge(coords, seg, rng, deg=15):\n    s, e = seg; L = e - s\n    if L < 60: return apply_hinge(coords, seg, rng, deg)\n    p1 = s + int(rng.integers(10, L//2))\n    p2 = s + int(rng.integers(L//2, L-10))\n    R1 = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n    R2 = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n    X  = coords.copy()\n    p0 = X[p1].copy(); X[p1+1:p2] = (X[p1+1:p2] - p0) @ R1.T + p0\n    p0 = X[p2].copy(); X[p2+1:e]  = (X[p2+1:e]  - p0) @ R2.T + p0\n    return X\n\n\ndef jitter_chains(coords, segs, rng, deg=12, trans=1.5):\n    X = coords.copy(); gc_ = 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) - gc_\n    return X\n\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: continue\n        ctrl = np.linspace(0, L-1, 6); disp = rng.normal(0, amp, (6, 3)); t = np.arange(L)\n        X[s:e] += np.vstack([np.interp(t, ctrl, disp[:, k]) for k in range(3)]).T\n    return X\n\n\ndef generate_rna_structure(sequence, seed=None, segments=None):\n    \"\"\"A-form helix, multi-chain aware.\"\"\"\n    if seed is not None: np.random.seed(seed)\n    n      = len(sequence)\n    coords = np.zeros((n, 3))\n    rise, twist, radius = 2.81, 0.5712, 10.0\n    if segments is None: segments = [(0, n)]\n    for chain_idx, (s, e) in enumerate(segments):\n        origin = np.array([chain_idx * 35.0, 0.0, 0.0])\n        for j, i in enumerate(range(s, e)):\n            ang       = j * twist\n            coords[i] = origin + [radius*np.cos(ang), radius*np.sin(ang), j*rise]\n    return coords\n\n\n# SCORE-2: diversity guard ─────────────────────────────────────────────────────\ndef _is_diverse_enough(new_coords, existing_list, min_rmsd=MIN_SLOT_RMSD):\n    \"\"\"Return True if new_coords is at least min_rmsd away from ALL existing slots.\"\"\"\n    for ex in existing_list:\n        n = min(len(new_coords), len(ex))\n        if n == 0: continue\n        if rmsd(new_coords[:n], ex[:n]) < min_rmsd:\n            return False\n    return True\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# TBM PHASE  (SCORE-4: parallel via ThreadPoolExecutor)\n# ══════════════════════════════════════════════════════════════════════════════\ndef _tbm_one_target(row, train_seqs_df, train_coords_dict, segments_map,\n                    exact_idx, similar_cache):\n    \"\"\"Process a single test target through TBM. Thread-safe (no shared writes).\"\"\"\n    tid  = row[\"target_id\"]\n    seq  = row[\"sequence\"]\n    segs = segments_map.get(tid, [(0, len(seq))])\n\n    # Check similarity cache (populated by main thread for thread safety)\n    similar = similar_cache.get(tid)\n    if similar is None:\n        similar = find_similar_sequences_detailed(\n            seq, train_seqs_df, train_coords_dict, top_n=50, exact_idx=exact_idx\n        )\n\n    preds, used = [], set()\n\n    for tmpl_id, tmpl_seq, sim, tmpl_coords, pct_id, _, _ in similar:\n        if len(preds) >= N_SAMPLE: break\n        if sim < MIN_SIMILARITY or pct_id < MIN_PERCENT_IDENTITY: break\n        if tmpl_id in used: continue\n\n        slot        = len(preds)\n        rng         = np.random.default_rng((row.name*10_000_000_000 + slot*10007) % (2**32))\n        adapted     = adapt_template_to_query(seq, tmpl_seq, tmpl_coords)\n        longest     = max(segs, key=lambda se: se[1]-se[0])\n        longest_len = longest[1] - longest[0]\n\n        if slot == 0:\n            # Weighted blend of top-2: weight = sim × pct_id/100\n            # SCORE-8: only blend if coords don't diverge too much\n            if len(similar) > 1:\n                t2_id, t2_seq, t2_sim, t2_coords, t2_pct, _, _ = similar[1]\n                if (t2_sim >= MIN_SIMILARITY and t2_pct >= MIN_PERCENT_IDENTITY\n                        and _coords_are_valid(t2_coords)):\n                    adapted2 = adapt_template_to_query(seq, t2_seq, t2_coords)\n                    n_common = min(len(adapted), len(adapted2))\n                    divergence = rmsd(adapted[:n_common], adapted2[:n_common])\n                    if divergence <= 15.0:   # SCORE-8: skip blend if >15Å divergence\n                        w0   = sim    * (pct_id  / 100.0)\n                        w1   = t2_sim * (t2_pct  / 100.0)\n                        norm = w0 + w1 + 1e-12\n                        adapted = (w0/norm)*adapted + (w1/norm)*adapted2\n            X = adaptive_rna_constraints(adapted, tid, segments_map, confidence=sim)\n\n        elif slot == 1:\n            noise_scale = max(0.005, (0.40-sim)*0.03) if pct_id > 95 else max(0.01, (0.40-sim)*0.06)\n            X = adapted + rng.normal(0, noise_scale, adapted.shape)\n            X = adaptive_rna_constraints(X, tid, segments_map, confidence=sim)\n\n        elif slot == 2:\n            X = apply_double_hinge(adapted, longest, rng) if longest_len >= 100 else apply_hinge(adapted, longest, rng)\n            X = adaptive_rna_constraints(X, tid, segments_map, confidence=sim)\n\n        elif slot == 3:\n            X = jitter_chains(adapted, segs, rng)\n            X = adaptive_rna_constraints(X, tid, segments_map, confidence=sim)\n\n        else:\n            X = smooth_wiggle(adapted, segs, rng)\n            X = adaptive_rna_constraints(X, tid, segments_map, confidence=sim)\n\n        # SCORE-2: diversity guard — only add if different enough from existing slots\n        if _is_diverse_enough(X, preds):\n            preds.append(X)\n            used.add(tmpl_id)\n        else:\n            # Slot occupied by near-duplicate — try next template without breaking\n            continue\n\n    # SCORE-3: low-threshold fallback before de-novo\n    # If we didn't fill all slots, retry with FALLBACK_PERCENT_IDENTITY=30%\n    if len(preds) < N_SAMPLE:\n        fallback_similar = find_similar_sequences_detailed(\n            seq, train_seqs_df, train_coords_dict, top_n=50,\n            exact_idx=exact_idx, min_pct_id=FALLBACK_PERCENT_IDENTITY\n        )\n        for tmpl_id, tmpl_seq, sim, tmpl_coords, pct_id, _, _ in fallback_similar:\n            if len(preds) >= N_SAMPLE: break\n            if pct_id < FALLBACK_PERCENT_IDENTITY or pct_id >= MIN_PERCENT_IDENTITY:\n                continue  # skip already-used quality range\n            if tmpl_id in used: continue\n            adapted = adapt_template_to_query(seq, tmpl_seq, tmpl_coords)\n            rng2    = np.random.default_rng((row.name*99999 + len(preds)*7) % (2**32))\n            X       = smooth_wiggle(adapted, segs, rng2)\n            X       = adaptive_rna_constraints(X, tid, segments_map, confidence=sim)\n            if _is_diverse_enough(X, preds):\n                preds.append(X)\n                used.add(tmpl_id)\n\n    n_needed = N_SAMPLE - len(preds)\n    return tid, preds, n_needed, seq\n\n\ndef tbm_phase(test_df, train_seqs_df, train_coords_dict, segments_map, exact_idx=None):\n    print(f\"\\n{'='*65}\")\n    print(\"PHASE 1: Template-Based Modeling\")\n    print(f\"  MIN_SIM={MIN_SIMILARITY} | MIN_PCT_ID={MIN_PERCENT_IDENTITY} \"\n          f\"| FALLBACK={FALLBACK_PERCENT_IDENTITY} | threads={TBM_THREADS}\")\n    print(f\"{'='*65}\")\n    t0 = time.time()\n    template_predictions, protenix_queue = {}, {}\n\n    rows_list = [row for _, row in test_df.iterrows()]\n\n    # SCORE-4: parallel TBM — each target is independent, use thread pool\n    similar_cache = {}  # pre-populated empty; each thread builds its own\n\n    def _process(row):\n        return _tbm_one_target(row, train_seqs_df, train_coords_dict,\n                               segments_map, exact_idx, similar_cache)\n\n    with ThreadPoolExecutor(max_workers=TBM_THREADS) as pool:\n        futures = {pool.submit(_process, row): row[\"target_id\"] for row in rows_list}\n        for fut in tqdm(as_completed(futures), total=len(futures), desc=\"TBM\"):\n            try:\n                tid, preds, n_needed, seq = fut.result()\n                template_predictions[tid] = preds\n                if n_needed > 0:\n                    protenix_queue[tid] = (n_needed, seq)\n                    print(f\"  {tid} ({len(seq)} nt): {len(preds)} TBM → need {n_needed} Protenix\")\n                else:\n                    print(f\"  {tid} ({len(seq)} nt): all {N_SAMPLE} TBM ✓\")\n            except Exception as e:\n                tid = futures[fut]\n                print(f\"  {tid}: TBM ERROR — {e}\")\n                traceback.print_exc()\n                template_predictions[tid] = []\n                seq = test_df[test_df[\"target_id\"]==tid][\"sequence\"].iloc[0]\n                protenix_queue[tid] = (N_SAMPLE, seq)\n\n    elapsed = time.time()-t0\n    print(f\"\\nPhase 1 done in {elapsed:.1f}s | \"\n          f\"TBM-only:{len(test_df)-len(protenix_queue)} | Protenix:{len(protenix_queue)}\")\n    return template_predictions, protenix_queue\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# PHASE 2 — PROTENIX  (SCORE-1, SCORE-6, SCORE-7)\n# ══════════════════════════════════════════════════════════════════════════════\n\n# SCORE-7: extract pLDDT confidence from Protenix output\ndef _extract_confidence(pred):\n    \"\"\"\n    Try to get per-residue pLDDT from Protenix output.\n    Returns mean confidence (0-1) or 0.5 if not available.\n    \"\"\"\n    try:\n        if \"plddt\" in pred:\n            return float(pred[\"plddt\"].mean().item())\n        if \"confidence\" in pred:\n            c = pred[\"confidence\"]\n            if hasattr(c, \"mean\"): return float(c.mean().item())\n        if \"predicted_aligned_error\" in pred:\n            pae = pred[\"predicted_aligned_error\"]\n            if hasattr(pae, \"mean\"):\n                # Convert PAE to pseudo-confidence: lower PAE = higher confidence\n                return float(1.0 / (1.0 + pae.mean().item() / 10.0))\n    except: pass\n    return 0.5\n\n\ndef run_protenix_multiseed(protenix_queue, work_dir, code_dir, root_dir,\n                           seeds, hard_targets):\n    print(f\"\\n{'='*65}\")\n    print(f\"PHASE 2: Protenix  seeds={seeds}  rank={RANK}/{WORLD}\")\n    print(f\"{'='*65}\")\n\n    device = torch.device(f\"cuda:{RANK}\" if torch.cuda.is_available() else \"cpu\")\n\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        torch.cuda.synchronize(device)\n        try:\n            free, total = torch.cuda.mem_get_info(RANK)\n            print(f\"  GPU {RANK}: {torch.cuda.get_device_name(RANK)} \"\n                  f\"({free/1024**3:.1f}/{total/1024**3:.1f} GB free)\")\n        except: pass\n\n    tasks, chunk_info = [], {}\n    for target_id, (n_needed, full_seq) in protenix_queue.items():\n        seq_len = len(full_seq)\n        overlap = get_chunk_overlap(seq_len)\n        if seq_len <= MAX_SEQ_LEN:\n            tasks.append({\"target_id\": target_id, \"sequence\": full_seq})\n            chunk_info[target_id] = [{\"name\": target_id, \"range\": (0, seq_len)}]\n        else:\n            chunks = split_into_chunks(seq_len, MAX_SEQ_LEN, overlap)\n            chunk_info[target_id] = []\n            for ci, (cs, ce) in enumerate(chunks):\n                cn = f\"{target_id}_chunk{ci}\"\n                tasks.append({\"target_id\": cn, \"sequence\": full_seq[cs:ce]})\n                chunk_info[target_id].append({\"name\": cn, \"range\": (cs, ce)})\n            print(f\"  {target_id} ({seq_len} nt): {len(chunks)} chunks overlap={overlap}\")\n\n    tasks_df        = pd.DataFrame(tasks)\n    input_json_path = str(work_dir / f\"protenix_input_r{RANK}.json\")\n    build_input_json(tasks_df, input_json_path)\n\n    from protenix.data.inference.infer_dataloader import InferenceDataset\n    from runner.inference import (InferenceRunner,\n                                  update_gpu_compatible_configs,\n                                  update_inference_configs)\n\n    configs0 = build_configs(input_json_path, str(work_dir/f\"out_r{RANK}_s0\"),\n                             MODEL_NAME, seed=seeds[0])\n    configs0 = update_gpu_compatible_configs(configs0)\n    runner   = InferenceRunner(configs0)\n    dataset  = InferenceDataset(configs0)\n    runner.model = runner.model.to(device)\n\n    def _get_c1_mask(data, atom_array, chunk_seq_len):\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\"): m = m & atom_array.is_rna\n                    return torch.from_numpy(m).bool()\n                if hasattr(atom_array, \"atom_name\"):\n                    base = atom_array.atom_name == \"C1'\"\n                    if hasattr(atom_array, \"is_rna\"): base = base & atom_array.is_rna\n                    return torch.from_numpy(base).bool()\n            except: pass\n        f = data[\"input_feature_dict\"]\n        if \"centre_atom_mask\" in f: return (f[\"centre_atom_mask\"] == 1).bool()\n        if \"center_atom_mask\"  in f: return (f[\"center_atom_mask\"]  == 1).bool()\n        m11 = (f[\"atom_to_tokatom_idx\"] == 11).bool()\n        m12 = (f[\"atom_to_tokatom_idx\"] == 12).bool()\n        c11, c12 = m11.sum().item(), m12.sum().item()\n        return m11 if abs(c11-chunk_seq_len) < abs(c12-chunk_seq_len) else m12\n\n    def _extract_c1_coords(data, atom_array, chunk_seq_len, raw_coords):\n        mask   = _get_c1_mask(data, atom_array, chunk_seq_len).to(raw_coords.device)\n        coords = raw_coords[:, mask, :].detach().cpu().numpy()\n        if coords.shape[1] > 1:\n            diffs = np.linalg.norm(coords[0, 1:] - coords[0, :-1], axis=-1)\n            if np.all(diffs < 1e-4): return None\n        if coords.shape[1] != chunk_seq_len:\n            if coords.shape[1] == 1 and chunk_seq_len > 1: return None\n            padded = np.zeros((coords.shape[0], chunk_seq_len, 3), dtype=np.float32)\n            ml = min(coords.shape[1], chunk_seq_len)\n            padded[:, :ml, :] = coords[:, :ml, :]\n            coords = padded\n        return coords\n\n    def _run_one_seed(seed_idx, cfg_seed):\n        configs_s = build_configs(input_json_path, str(work_dir/f\"out_r{RANK}_s{seed_idx}\"),\n                                  MODEL_NAME, seed=cfg_seed)\n        configs_s = update_gpu_compatible_configs(configs_s)\n        raw_preds   = {}   # sample_name → coords\n        confidences = {}   # sample_name → float confidence  (SCORE-7)\n\n        # SCORE-6: request n_needed+1 predictions, keep best by confidence\n        for i in tqdm(range(RANK, len(dataset), WORLD),\n                      desc=f\"Protenix rank{RANK} seed{cfg_seed}\"):\n            data, atom_array, err = dataset[i]\n            sample_name = data.get(\"sample_name\", f\"sample_{i}\")\n\n            if err:\n                print(f\"  {sample_name}: data error — {err}\")\n                raw_preds[sample_name] = None\n                del data, atom_array, err\n                gc.collect()\n                if torch.cuda.is_available(): torch.cuda.empty_cache()\n                continue\n\n            target_id   = sample_name.split(\"_chunk\")[0] if \"_chunk\" in sample_name else sample_name\n            n_needed    = protenix_queue.get(target_id, (N_SAMPLE, \"\"))[0]\n            sub_seq_len = data[\"N_token\"].item()\n\n            if seed_idx > 0 and target_id not in hard_targets:\n                raw_preds[sample_name] = None\n                del data, atom_array\n                continue\n\n            # SCORE-6: request 1 extra prediction to allow confidence-based selection\n            n_request = min(n_needed + 1, N_SAMPLE + 1)\n\n            try:\n                new_cfg = update_inference_configs(configs_s, sub_seq_len)\n                new_cfg.sample_diffusion.N_sample = n_request\n                runner.update_model_configs(new_cfg)\n                pred       = runner.predict(data)\n                raw_coords = pred[\"coordinate\"]\n                coords     = _extract_c1_coords(data, atom_array, sub_seq_len, raw_coords)\n\n                # SCORE-7: get confidence score for ranking\n                conf = _extract_confidence(pred)\n                confidences[sample_name] = conf\n\n                # SCORE-1 + SCORE-6: if we got n_request preds, rank by per-sample\n                # confidence proxy (spread of coords) and keep best n_needed\n                if coords is not None and coords.shape[0] > n_needed:\n                    # Proxy: prefer predictions with larger coordinate spread\n                    # (collapsed = bad, spread out = model is confident)\n                    spreads = np.array([\n                        np.max(coords[k], axis=0).sum() - np.min(coords[k], axis=0).sum()\n                        for k in range(coords.shape[0])\n                    ])\n                    keep_idx = np.argsort(spreads)[::-1][:n_needed]\n                    coords   = coords[keep_idx]\n\n                raw_preds[sample_name] = coords\n                if coords is not None:\n                    print(f\"\\n  {sample_name} seed{cfg_seed}: {coords.shape[0]} preds \"\n                          f\"conf={conf:.3f} ✓\")\n                else:\n                    print(f\"\\n  {sample_name} seed{cfg_seed}: extraction failed\")\n            except Exception as exc:\n                print(f\"\\n  {sample_name} seed{cfg_seed}: FAILED — {exc}\")\n                traceback.print_exc()\n                raw_preds[sample_name] = None\n            finally:\n                try: del pred, data, atom_array, raw_coords\n                except: pass\n                gc.collect()\n                if torch.cuda.is_available(): torch.cuda.empty_cache()\n\n        return raw_preds, confidences\n\n    all_seed_preds = []\n    all_seed_confs = []\n    for si, s in enumerate(seeds):\n        rp, rc = _run_one_seed(si, s)\n        all_seed_preds.append(rp)\n        all_seed_confs.append(rc)\n\n    protenix_preds = {}\n    for target_id, (n_needed, full_seq) in protenix_queue.items():\n        seq_len = len(full_seq)\n        chunks  = chunk_info.get(target_id, [])\n        if not chunks: continue\n\n        seed_results  = []\n        seed_conf_vals = []\n        for si, (raw_preds, raw_confs) in enumerate(zip(all_seed_preds, all_seed_confs)):\n            if si > 0 and target_id not in hard_targets:\n                continue\n            if len(chunks) == 1:\n                coords = raw_preds.get(target_id)\n                if _ptx_coords_are_valid(coords):\n                    seed_results.append(coords)\n                    seed_conf_vals.append(raw_confs.get(target_id, 0.5))\n            else:\n                per_sample = {s: [] for s in range(n_needed)}\n                all_ok     = True\n                for cinfo in chunks:\n                    ccoords = raw_preds.get(cinfo[\"name\"])\n                    if ccoords is None or not _ptx_coords_are_valid(ccoords):\n                        all_ok = False; break\n                    for s_idx in range(n_needed):\n                        si2 = s_idx if s_idx < ccoords.shape[0] else -1\n                        per_sample[s_idx].append((ccoords[si2], cinfo[\"range\"]))\n                if not all_ok: continue\n                stitched = []\n                for s_idx in range(n_needed):\n                    items = per_sample[s_idx]\n                    fc    = stitch_chunk_coords([c for c, _ in items],\n                                                [r for _, r in items], seq_len)\n                    stitched.append(fc)\n                seed_results.append(np.stack(stitched, axis=0))\n                # Use average confidence across chunks for this target\n                chunk_names = [ci[\"name\"] for ci in chunks]\n                chunk_confs = [raw_confs.get(cn, 0.5) for cn in chunk_names]\n                seed_conf_vals.append(float(np.mean(chunk_confs)))\n\n        if not seed_results:\n            print(f\"  {target_id}: all seeds FAILED → de-novo fallback\")\n            protenix_preds[target_id] = None\n            continue\n\n        if len(seed_results) == 1:\n            protenix_preds[target_id] = seed_results[0]\n            print(f\"  {target_id}: {seed_results[0].shape[0]} preds ✓ \"\n                  f\"(1 seed, conf={seed_conf_vals[0]:.3f})\")\n        else:\n            # SCORE-1: put highest-confidence seed first\n            if seed_conf_vals[1] > seed_conf_vals[0]:\n                seed_results    = seed_results[::-1]\n                seed_conf_vals  = seed_conf_vals[::-1]\n            half   = N_SAMPLE // 2\n            take1  = min(seed_results[0].shape[0], N_SAMPLE - half)\n            take2  = min(seed_results[1].shape[0], half)\n            merged = np.concatenate([seed_results[0][:take1],\n                                     seed_results[1][:take2]], axis=0)\n            protenix_preds[target_id] = merged\n            print(f\"  {target_id}: {merged.shape[0]} preds ✓ \"\n                  f\"(seed_confs={seed_conf_vals[0]:.3f},{seed_conf_vals[1]:.3f})\")\n\n    return protenix_preds\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# RANK RESULT MERGE\n# ══════════════════════════════════════════════════════════════════════════════\ndef _save_rank_results(work_dir, rank, protenix_preds, all_rows):\n    path = work_dir / f\"rank{rank}_results.pkl\"\n    with open(path, \"wb\") as f:\n        pickle.dump({\"protenix_preds\": protenix_preds, \"all_rows\": all_rows}, f)\n    (work_dir / f\"rank{rank}_done.flag\").touch()\n    print(f\"  Rank {rank}: results saved.\")\n\n\ndef _wait_for_all_ranks(work_dir, world, timeout=600):\n    t0 = time.time()\n    while True:\n        flags = [work_dir / f\"rank{r}_done.flag\" for r in range(world)]\n        if all(f.exists() for f in flags): return True\n        if time.time()-t0 > timeout:\n            missing = [r for r in range(world) if not (work_dir/f\"rank{r}_done.flag\").exists()]\n            print(f\"  WARNING: timeout, missing ranks {missing}\")\n            return False\n        time.sleep(3)\n\n\ndef _merge_rank_results(work_dir, world):\n    merged_preds = {}\n    merged_rows  = []\n    for r in range(world):\n        path = work_dir / f\"rank{r}_results.pkl\"\n        if not path.exists():\n            print(f\"  WARNING: missing results from rank {r}\")\n            continue\n        with open(path, \"rb\") as f:\n            d = pickle.load(f)\n        merged_preds.update(d[\"protenix_preds\"])\n        merged_rows.extend(d[\"all_rows\"])\n    seen, final = set(), []\n    for r in merged_rows:\n        if r[\"ID\"] not in seen:\n            seen.add(r[\"ID\"]); final.append(r)\n    return merged_preds, final\n\n\n# SCORE-10: submission validator ──────────────────────────────────────────────\ndef _validate_submission(sub, expected_ids):\n    sub_ids = set(sub[\"ID\"].tolist())\n    for eid in expected_ids:\n        if eid not in sub_ids:\n            print(f\"  WARN: missing ID {eid}\")\n    coord_cols = [c for c in sub.columns if c.startswith((\"x_\",\"y_\",\"z_\"))]\n    nan_count  = sub[coord_cols].isna().sum().sum()\n    inf_count  = np.isinf(sub[coord_cols].values).sum()\n    if nan_count: print(f\"  WARN: {nan_count} NaN values in submission\")\n    if inf_count: print(f\"  WARN: {inf_count} Inf values in submission\")\n    print(f\"  Submission: {len(sub)} rows, {len(coord_cols)} coord cols — OK\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# MAIN\n# ══════════════════════════════════════════════════════════════════════════════\ndef main():\n    global RANK, WORLD\n    t_total = time.time()\n    test_csv, output_csv, code_dir, root_dir = resolve_paths()\n\n    if not os.path.isdir(code_dir):\n        raise FileNotFoundError(f\"Missing PROTENIX_CODE_DIR: {code_dir}\")\n\n    os.environ[\"PROTENIX_ROOT_DIR\"] = root_dir\n    sys.path.append(code_dir)\n    ensure_required_files(root_dir)\n    seed_everything(SEED)\n\n    if RANK == 0:\n        print(f\"WORLD={WORLD} | rank={RANK}\")\n        for i in range(torch.cuda.device_count()):\n            try:\n                free, total = torch.cuda.mem_get_info(i)\n                print(f\"  GPU {i}: {torch.cuda.get_device_name(i)} \"\n                      f\"({free/1024**3:.1f}/{total/1024**3:.1f} GB free)\")\n            except:\n                print(f\"  GPU {i}: {torch.cuda.get_device_name(i)}\")\n\n    test_df = pd.read_csv(test_csv).reset_index(drop=True)\n    if RANK == 0:\n        print(f\"\\nTest targets: {len(test_df)}\")\n\n    print(f\"[rank{RANK}] Loading training + validation data…\")\n    combined_seqs   = pd.concat([pd.read_csv(DEFAULT_TRAIN_CSV),\n                                 pd.read_csv(DEFAULT_VAL_CSV)],  ignore_index=True)\n    combined_labels = pd.concat([pd.read_csv(DEFAULT_TRAIN_LBLS, low_memory=False),\n                                 pd.read_csv(DEFAULT_VAL_LBLS)], ignore_index=True)\n    gc.collect()\n\n    print(f\"[rank{RANK}] Processing labels (ID-suffix sort)…\")\n    train_coords = process_labels(combined_labels)\n    del combined_labels; gc.collect()\n\n    segments_map, _ = build_segments_map(test_df)\n    exact_idx        = build_exact_match_index(combined_seqs, train_coords)\n\n    if RANK == 0:\n        print(f\"Template pool: {len(combined_seqs)} seqs, {len(train_coords)} structures\")\n\n    # ── PHASE 1: TBM ─────────────────────────────────────────────────────────\n    template_preds, protenix_queue = tbm_phase(\n        test_df, combined_seqs, train_coords, segments_map, exact_idx=exact_idx\n    )\n\n    # ── PHASE 2: Protenix ─────────────────────────────────────────────────────\n    protenix_preds = {}\n    work_dir       = Path(\"/kaggle/working\")\n    work_dir.mkdir(parents=True, exist_ok=True)\n\n    if protenix_queue and USE_PROTENIX:\n        hard_targets = {tid for tid, preds in template_preds.items() if len(preds) == 0}\n        seeds        = SEEDS_HARD if hard_targets else SEEDS_EASY\n        print(f\"\\n[rank{RANK}] Hard targets (0 TBM preds): {sorted(hard_targets)}\")\n        print(f\"[rank{RANK}] Seeds: {seeds}\")\n        protenix_preds = run_protenix_multiseed(\n            protenix_queue, work_dir, code_dir, root_dir,\n            seeds=seeds, hard_targets=hard_targets\n        )\n    elif protenix_queue:\n        print(f\"\\n[rank{RANK}] PHASE 2 skipped.\")\n\n    # ── PHASE 3: Combine ──────────────────────────────────────────────────────\n    print(f\"\\n[rank{RANK}] PHASE 3: Combine TBM + Protenix + de-novo\")\n    all_rows = []\n    for _, row in test_df.iterrows():\n        tid, seq = row[\"target_id\"], row[\"sequence\"]\n        segs     = segments_map.get(tid, [(0, len(seq))])\n        combined = list(template_preds.get(tid, []))\n        n_tbm    = len(combined)\n\n        ptx = protenix_preds.get(tid)\n        if ptx is not None and ptx.ndim == 3 and _ptx_coords_are_valid(ptx):\n            for j in range(ptx.shape[0]):\n                if len(combined) >= N_SAMPLE: break\n                # SCORE-2: diversity guard on PTX slots too\n                c = ptx[j].astype(np.float64)\n                if _is_diverse_enough(c, combined):\n                    combined.append(c)\n                else:\n                    # near-duplicate Protenix pred — add light wiggle\n                    rng_d = np.random.default_rng(j * 999983)\n                    c_div = smooth_wiggle(c, segs, rng_d, amp=1.5)\n                    combined.append(c_div)\n        n_ptx = len(combined) - n_tbm\n\n        n_dn = 0\n        while len(combined) < N_SAMPLE:\n            seed_val = row.name*1_000_000 + len(combined)*1000\n            dn       = generate_rna_structure(seq, seed=seed_val, segments=segs)\n            combined.append(adaptive_rna_constraints(dn, tid, segments_map, confidence=0.2))\n            n_dn += 1\n\n        strategy = f\"TBM={n_tbm} PTX={n_ptx}\" + (f\" DN={n_dn}\" if n_dn else \"\")\n        if RANK == 0:\n            print(f\"  {tid} ({len(seq)} nt): {strategy}\")\n\n        stacked = np.stack(combined[:N_SAMPLE], axis=0)\n        all_rows.extend(coords_to_rows(tid, seq, stacked))\n\n    # ── Merge and save ────────────────────────────────────────────────────────\n    if WORLD > 1:\n        _save_rank_results(work_dir, RANK, protenix_preds, all_rows)\n        if RANK == 0:\n            print(f\"\\n[rank0] Waiting for all {WORLD} ranks…\")\n            _wait_for_all_ranks(work_dir, WORLD)\n            _, all_rows = _merge_rank_results(work_dir, WORLD)\n\n    if RANK == 0:\n        sub  = pd.DataFrame(all_rows)\n        cols = [\"ID\", \"resname\", \"resid\"] + [\n            f\"{c}_{i}\" for i in range(1, N_SAMPLE+1) for c in [\"x\", \"y\", \"z\"]\n        ]\n        cc = [c for c in cols if c.startswith((\"x_\", \"y_\", \"z_\"))]\n        sub[cc] = sub[cc].clip(-999.999, 9999.999)\n\n        # SCORE-10: validate before saving\n        expected_ids = {f\"{row['target_id']}_{i+1}\"\n                        for _, row in test_df.iterrows()\n                        for i in range(len(row[\"sequence\"]))}\n        _validate_submission(sub, expected_ids)\n\n        sub[cols].to_csv(output_csv, index=False)\n        elapsed = time.time()-t_total\n        print(f\"\\n✓ Saved {output_csv}  ({len(sub):,} rows)  in {elapsed/60:.1f} min\")\n\n\n# ══════════════════════════════════════════════════════════════════════════════\n# ENTRY POINT — fork-based multi-GPU (works in Kaggle notebooks)\n# ══════════════════════════════════════════════════════════════════════════════\ndef _worker(rank, world):\n    global RANK, WORLD\n    RANK  = rank\n    WORLD = world\n    # TF32 flags set here, after fork, on this process's CUDA context\n    torch.cuda.set_device(rank)\n    torch.backends.cuda.matmul.allow_tf32 = True\n    torch.backends.cudnn.allow_tf32       = True\n    torch.backends.cudnn.benchmark        = True\n    try:\n        main()\n    except Exception:\n        print(f\"\\n[rank{rank}] FATAL ERROR:\")\n        traceback.print_exc()\n        raise\n\n\ndef _count_gpus_no_cuda_init():\n    try:\n        r = subprocess.run([\"nvidia-smi\", \"--list-gpus\"],\n                           capture_output=True, text=True, timeout=10)\n        if r.returncode == 0:\n            return max(len([l for l in r.stdout.strip().splitlines() if l.strip()]), 1)\n    except: pass\n    return 1\n\n\nif __name__ == \"__main__\":\n    n_gpus = _count_gpus_no_cuda_init()\n    print(f\"Detected {n_gpus} GPU(s)\")\n\n    # SCORE-9: clean stale flags from any previous runs\n    work_dir = Path(\"/kaggle/working\")\n    work_dir.mkdir(parents=True, exist_ok=True)\n    for f in work_dir.glob(\"rank*_done.flag\"):\n        f.unlink(missing_ok=True)\n    for f in work_dir.glob(\"rank*_results.pkl\"):\n        f.unlink(missing_ok=True)\n\n    if n_gpus > 1:\n        print(f\"Launching {n_gpus} workers via fork…\")\n        ctx       = mp.get_context(\"fork\")\n        processes = []\n        for rank in range(n_gpus):\n            p = ctx.Process(target=_worker, args=(rank, n_gpus), daemon=False)\n            p.start()\n            processes.append(p)\n        failed = []\n        for rank, p in enumerate(processes):\n            p.join()\n            if p.exitcode != 0:\n                failed.append((rank, p.exitcode))\n        if failed:\n            raise RuntimeError(f\"Workers failed: {failed}\")\n        print(\"All workers finished successfully.\")\n    else:\n        print(\"Single GPU mode.\")\n        _worker(0, 1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T20:26:59.273486Z","iopub.execute_input":"2026-03-07T20:26:59.273857Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}