#!/usr/bin/env python3
"""
RNA 3D Folding Part 2 — 4th place solution
Pipeline: Cross-Attention template reranker + Protenix v1 + RNAPro

Slots (5 predictions per target):
  1: TBM bio top-1 (BioPython global alignment)
  2: TBM model top-1 (Cross-Attn reranked, mmseqs_0.300 cluster-diverse from S1)
  3: TBM combined top-1 (sim * model_score, cluster-diverse from S1/S2)
  4: Protenix v1 (co-fold with protein/DNA/ligands if they fit)
  5: RNAPro

For rna_len > 512 nt the Cross-Attn reranker is skipped (training-time length
limit), and slots 2/3 fall back to bio ranks 2/3 with cluster diversity.

Kaggle datasets:
  - stanford-rna-3d-folding-2 (competition)
  - gapchenko/rna3d2-pip-wheels (pip wheels — offline install)
  - gapchenko/rna3d2-bundle (checkpoints, source, template DB, MSA, etc.)
  - gapchenko/rna3d2-tbm-hgb-v1 (Cross-Attn checkpoint, template DB metadata)
  - gapchenko/rna3d2-prot-msa-v2 (protein MSA a3m files)
  - gapchenko/rna3d2-rnapro-public-best-checkpoint (RNAPro public-best checkpoint)
"""
import os, sys, time, json, gc, re, traceback, warnings, glob, subprocess, shutil
from collections import OrderedDict
warnings.filterwarnings("ignore")
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"

import numpy as np
import pandas as pd

# ============================================================
# CONFIG PROFILES — change only CONFIG_PROFILE to switch
# ============================================================
CONFIG_PROFILE = "A"  # ← CHANGE THIS: A, B, C, D

PROFILES = {
    # A: Balanced — no time limits (private test has more time)
    "A": dict(
        PTX_N_CYCLE=4, PTX_N_STEP=40, PTX_N_SAMPLE=3,
        RNAPRO_N_CYCLE=4, RNAPRO_N_STEP=40, RNAPRO_N_SAMPLE=3,
        PTX_MAX_RNA_LEN=950, PTX_MAX_COFOLD_TOKENS=450,
        RNAPRO_MAX_LEN=500, RNAPRO_SKIP_LEN=99999,
        PTX_TIME_LIMIT=90000, RNAPRO_TIME_LIMIT=60000,
        PHASE_A_BUDGET=2020000, PHASE_B_BUDGET=700000,
        GLOBAL_DEADLINE=2820000,  # ~32 days — effectively no limit
        ADAPTIVE_TIME=True,
    ),
    # B: Fast — fewer cycles, more targets covered
    "B": dict(
        PTX_N_CYCLE=4, PTX_N_STEP=20, PTX_N_SAMPLE=3,
        RNAPRO_N_CYCLE=4, RNAPRO_N_STEP=20, RNAPRO_N_SAMPLE=2,
        PTX_MAX_RNA_LEN=600, PTX_MAX_COFOLD_TOKENS=1500,
        RNAPRO_MAX_LEN=400, RNAPRO_SKIP_LEN=99999,
        PTX_TIME_LIMIT=600, RNAPRO_TIME_LIMIT=300,
        PHASE_A_BUDGET=12000, PHASE_B_BUDGET=4800,
        GLOBAL_DEADLINE=28200,
        ADAPTIVE_TIME=True,
    ),
    # C: Quality — more cycles for PTX co-fold targets
    "C": dict(
        PTX_N_CYCLE=10, PTX_N_STEP=50, PTX_N_SAMPLE=3,
        RNAPRO_N_CYCLE=4, RNAPRO_N_STEP=40, RNAPRO_N_SAMPLE=3,
        PTX_MAX_RNA_LEN=900, PTX_MAX_COFOLD_TOKENS=1800,
        RNAPRO_MAX_LEN=400, RNAPRO_SKIP_LEN=99999,
        PTX_TIME_LIMIT=1200, RNAPRO_TIME_LIMIT=600,
        PHASE_A_BUDGET=18000, PHASE_B_BUDGET=5400,
        GLOBAL_DEADLINE=28200,
        ADAPTIVE_TIME=True,
    ),
    # D: PTX-heavy — skip RNAPro, give all time to Protenix
    "D": dict(
        PTX_N_CYCLE=8, PTX_N_STEP=40, PTX_N_SAMPLE=5,
        RNAPRO_N_CYCLE=4, RNAPRO_N_STEP=40, RNAPRO_N_SAMPLE=1,
        PTX_MAX_RNA_LEN=900, PTX_MAX_COFOLD_TOKENS=1800,
        RNAPRO_MAX_LEN=400, RNAPRO_SKIP_LEN=300,
        PTX_TIME_LIMIT=1200, RNAPRO_TIME_LIMIT=300,
        PHASE_A_BUDGET=21600, PHASE_B_BUDGET=2400,
        GLOBAL_DEADLINE=28200,
        ADAPTIVE_TIME=True,
    ),
}

P = PROFILES[CONFIG_PROFILE]
# TEST_MODE: process only TEST_TIDS (a small fixed list) instead of the full test
# set; lets you sanity-check the pipeline end-to-end on a Kaggle session without
# burning the 8 h submission budget. Turn OFF before a real submission.
TEST_MODE = True
TEST_TIDS = [
    "2RSK",  # ~214 nt, cofold (protein partner) — exercises the co-fold path
    "9JGM",  # ~210 nt — pure RNA, exercises the RNA-only path
]

# Protenix v1
PTX_N_CYCLE = P["PTX_N_CYCLE"]
PTX_N_STEP = P["PTX_N_STEP"]
PTX_N_SAMPLE = P["PTX_N_SAMPLE"]
PTX_SEED = 42
PTX_MAX_RNA_LEN = P["PTX_MAX_RNA_LEN"]
PTX_MAX_COFOLD_TOKENS = P["PTX_MAX_COFOLD_TOKENS"]
PTX_MODEL_NAME = "protenix_base_20250630_v1.0.0"
PTX_TIME_LIMIT = P["PTX_TIME_LIMIT"]       # max seconds per target for PTX

# RNAPro
RNAPRO_N_CYCLE = P["RNAPRO_N_CYCLE"]
RNAPRO_N_STEP = P["RNAPRO_N_STEP"]
RNAPRO_N_SAMPLE = P["RNAPRO_N_SAMPLE"]
RNAPRO_MAX_LEN = P["RNAPRO_MAX_LEN"]
RNAPRO_SKIP_LEN = P["RNAPRO_SKIP_LEN"]     # skip RNAPro if RNA > this length
RNAPRO_TIME_LIMIT = P["RNAPRO_TIME_LIMIT"]  # max seconds per target for RNAPro

# Phase budgets
PHASE_A_BUDGET = P["PHASE_A_BUDGET"]  # max seconds for all PTX targets
PHASE_B_BUDGET = P["PHASE_B_BUDGET"]  # max seconds for all RNAPro targets
GLOBAL_DEADLINE = P["GLOBAL_DEADLINE"]  # hard wall-clock deadline from start (Kaggle=8h)
ADAPTIVE_TIME = P.get("ADAPTIVE_TIME", True)  # per-target PTX time based on complexity+MSA

# Sequence chunking
PTX_CHUNK_OVERLAP = 12
RNAPRO_CHUNK_OVERLAP = 6
PTX_COFOLD_PROTEIN_BUDGET = 450   # max protein tokens for chain-aware cofold

IS_COMP_RERUN = bool(os.environ.get("KAGGLE_IS_COMPETITION_RERUN", ""))

# RNAPro checkpoint selection: "private" (competition winner) or "public" (public LB best)
RNAPRO_CKPT_VARIANT = "private"  # uses rnapro-public-best-inference.ckpt override

# Ligand injection — blacklist crystallization/buffer artifacts
LIGAND_BLACKLIST = {
    # Buffers and crystallization agents
    "GOL", "EDO", "PEG", "PG4", "P6G", "1PE", "PE4", "EPE",  # polyethylene glycols
    "SO4", "PO4", "CIT", "ACT", "FMT", "ACY", "TRS", "MES",  # buffer salts
    "MPD", "IPA", "DMS", "BME",  # solvents
    # Common crystallization additives (not biologically relevant to RNA fold)
    "SPM", "SPD",  # polyamines (stabilize crystal packing, not fold-critical)
    "HEZ",  # hexanediol (precipitant — hurt 5TPY by -0.09)
    "NCO",  # cobalt hexammine (crystallographic probe)
    "IRI",  # iridium hexammine (crystallographic heavy atom)
    "RHD",  # rhodium ammine (heavy atom derivative)
}


def parse_ligands_for_ptx(ligand_ids_str):
    """Parse ligand_ids column → list of Protenix ligand entities, filtering junk.

    Returns list of dicts like {"ligand": "CCD_SAM", "count": 1}.
    """
    if not ligand_ids_str or str(ligand_ids_str) == 'nan':
        return []
    ids = [x.strip() for x in str(ligand_ids_str).split(";") if x.strip()]
    ligands = []
    for lid in ids:
        if lid in LIGAND_BLACKLIST:
            continue
        ligands.append({"ligand": f"CCD_{lid}", "count": 1})
    return ligands


# ============================================================
# 1. PATHS — auto-discover Kaggle datasets
# ============================================================
print("=" * 70)
print(f"RNA 3D Folding — v8 Triple Model (profile {CONFIG_PROFILE})")
print("=" * 70)
print(f"PTX: cycles={PTX_N_CYCLE} steps={PTX_N_STEP} samples={PTX_N_SAMPLE} "
      f"max_rna={PTX_MAX_RNA_LEN} max_cofold={PTX_MAX_COFOLD_TOKENS} tlimit={PTX_TIME_LIMIT}s")
print(f"RNAPro: cycles={RNAPRO_N_CYCLE} steps={RNAPRO_N_STEP} samples={RNAPRO_N_SAMPLE} "
      f"max_len={RNAPRO_MAX_LEN} skip>{RNAPRO_SKIP_LEN} tlimit={RNAPRO_TIME_LIMIT}s")
print(f"Budgets: phaseA={PHASE_A_BUDGET}s phaseB={PHASE_B_BUDGET}s global={GLOBAL_DEADLINE}s adaptive={ADAPTIVE_TIME}")
print(f"Test mode: {TEST_MODE}")

OUTPUT_PATH = '/kaggle/working' if os.path.exists('/kaggle/working') else '/workspace/kaggle_test_output'
os.makedirs(OUTPUT_PATH, exist_ok=True)

# Find competition data
DATA_PATH = None
for candidate in [
    '/kaggle/input/stanford-rna-3d-folding-2',
    '/kaggle/input/competitions/stanford-rna-3d-folding-2',
    '/workspace/lora_ws/competition',  # pod fallback
]:
    if os.path.exists(candidate) and os.path.exists(os.path.join(candidate, 'test_sequences.csv')):
        DATA_PATH = candidate
        break
if DATA_PATH is None:
    # Deep search
    for root, dirs, files in os.walk('/kaggle/input'):
        if 'test_sequences.csv' in files:
            DATA_PATH = root
            break
if DATA_PATH is None:
    raise ValueError("Competition data not found!")
print(f"Competition data: {DATA_PATH}")

# Find our datasets
WHEELS_DIR = None
BUNDLE_DIR = None

search_dirs = []
if os.path.exists('/kaggle/input'):
    for d in os.listdir('/kaggle/input'):
        search_dirs.append(os.path.join('/kaggle/input', d))
    # Nested: /kaggle/input/datasets/<user>/<slug>
    ds_root = '/kaggle/input/datasets'
    if os.path.exists(ds_root):
        for user in os.listdir(ds_root):
            user_path = os.path.join(ds_root, user)
            if os.path.isdir(user_path):
                for ds in os.listdir(user_path):
                    search_dirs.append(os.path.join(user_path, ds))

for d in search_dirs:
    if not os.path.isdir(d):
        continue
    contents = os.listdir(d)
    if any(f.endswith('.whl') for f in contents) or any(f.endswith('.tar.gz') for f in contents):
        WHEELS_DIR = d
    if 'checkpoints' in contents and 'template_db_v4' in contents:
        BUNDLE_DIR = d

# Pod fallback
if WHEELS_DIR is None and os.path.exists('/workspace/kaggle_datasets/wheels'):
    WHEELS_DIR = '/workspace/kaggle_datasets/wheels'
if BUNDLE_DIR is None and os.path.exists('/workspace/kaggle_datasets/bundle'):
    BUNDLE_DIR = '/workspace/kaggle_datasets/bundle'

if WHEELS_DIR is None:
    raise ValueError("Wheels dataset not found!")
if BUNDLE_DIR is None:
    raise ValueError("Bundle dataset not found!")

# Find Cross-Attn data dataset (holds the CA checkpoint + template DB metadata)
CA_DATA_DIR = None
for d in search_dirs:
    if not os.path.isdir(d):
        continue
    if os.path.exists(os.path.join(d, 'tm_crossattn_z1.pt')):
        CA_DATA_DIR = d
        break
if CA_DATA_DIR is None and os.path.exists('/workspace/kaggle_datasets/tbm_hgb_v1'):
    CA_DATA_DIR = '/workspace/kaggle_datasets/tbm_hgb_v1'
if CA_DATA_DIR is None:
    raise ValueError("Cross-Attn data dataset (rna3d2-tbm-hgb-v1) not found!")

print(f"Wheels: {WHEELS_DIR}")
print(f"Bundle: {BUNDLE_DIR}")
print(f"CA data: {CA_DATA_DIR}")

# Debug: show bundle contents
if BUNDLE_DIR:
    print(f"  Bundle contents: {sorted(os.listdir(BUNDLE_DIR))}")

# Derived paths — use _find to handle Kaggle flat-extracting archives
def _find(base, name):
    """Find file/dir in base, searching recursively if needed."""
    direct = os.path.join(base, name)
    if os.path.exists(direct):
        return direct
    # Try without .gz (Kaggle decompresses)
    if name.endswith('.gz'):
        nogz = os.path.join(base, name[:-3])
        if os.path.exists(nogz):
            return nogz
    # Search one level deep
    for d in os.listdir(base):
        candidate = os.path.join(base, d, name)
        if os.path.exists(candidate):
            return candidate
    return direct  # fallback to direct path (will fail with clear error)

_RNAPRO_CKPT_PRIVATE = _find(BUNDLE_DIR, os.path.join('checkpoints', 'rnapro-inference-only.ckpt'))

# Public-best: search in dedicated dataset, then bundle fallback
_RNAPRO_CKPT_PUBLIC = None
for d in search_dirs:
    if not os.path.isdir(d): continue
    candidate = os.path.join(d, 'rnapro-public-best-inference.ckpt')
    if os.path.exists(candidate):
        _RNAPRO_CKPT_PUBLIC = candidate
        break
if _RNAPRO_CKPT_PUBLIC is None:
    _RNAPRO_CKPT_PUBLIC = _find(BUNDLE_DIR, os.path.join('checkpoints', 'rnapro-public-best-inference.ckpt'))

RNAPRO_CHECKPOINT = _RNAPRO_CKPT_PUBLIC if RNAPRO_CKPT_VARIANT == "public" else _RNAPRO_CKPT_PRIVATE
print(f"RNAPro ckpt: {RNAPRO_CKPT_VARIANT} → {os.path.basename(RNAPRO_CHECKPOINT)}")
PTX_CHECKPOINT = _find(BUNDLE_DIR, os.path.join('checkpoints', PTX_MODEL_NAME + '.pt'))
RIBONANZA_PATH = _find(BUNDLE_DIR, os.path.join('checkpoints', 'ribonanzanet2_checkpoint'))
CCD_CACHE = _find(BUNDLE_DIR, 'ccd_cache')
TEMPLATE_DB_DIR = _find(BUNDLE_DIR, 'template_db_v4')
CA_TEMPLATE_META_PATH = os.path.join(CA_DATA_DIR, 'template_db_v4_unified', 'template_db.json')
RNA_METADATA_CSV = os.path.join(DATA_PATH, 'extra', 'rna_metadata.csv')
if not os.path.exists(RNA_METADATA_CSV):
    RNA_METADATA_CSV = '/workspace/lora_ws/competition/extra/rna_metadata.csv'
COMP_MSA_DIR = os.path.join(DATA_PATH, 'MSA')
EXT_MSA_DIR = _find(BUNDLE_DIR, os.path.join('ext_msa', 'msas'))
PDB_TO_URS = _find(BUNDLE_DIR, os.path.join('ext_msa', 'pdb_to_primary_urs.json'))

# Protein MSA (a3m files from uniref30 search)
PROTEIN_A3M_DIR = None
_protein_search_dirs = search_dirs + ([BUNDLE_DIR] if BUNDLE_DIR else [])
for d in _protein_search_dirs:
    if not os.path.isdir(d):
        continue
    # Check for extracted dir
    if os.path.isdir(os.path.join(d, 'protein_a3m')):
        PROTEIN_A3M_DIR = os.path.join(d, 'protein_a3m')
        break
    if os.path.isdir(os.path.join(d, 'a3m_named')):
        PROTEIN_A3M_DIR = os.path.join(d, 'a3m_named')
        break
    if os.path.isdir(os.path.join(d, 'protein_a3m_kaggle')):
        PROTEIN_A3M_DIR = os.path.join(d, 'protein_a3m_kaggle')
        break
    # Check for tar.gz to extract
    tgz = os.path.join(d, 'protein_a3m.tar.gz')
    if os.path.exists(tgz):
        PROTEIN_A3M_DIR = os.path.join(OUTPUT_PATH, 'protein_a3m')
        os.makedirs(PROTEIN_A3M_DIR, exist_ok=True)
        import tarfile
        with tarfile.open(tgz, 'r:gz') as tf:
            tf.extractall(PROTEIN_A3M_DIR)
        print(f"Extracted protein a3m: {len(os.listdir(PROTEIN_A3M_DIR))} files")
        break
    # Check for loose a3m files (Kaggle auto-extracts tar.gz)
    a3ms = [f for f in os.listdir(d) if f.endswith('.a3m')]
    if len(a3ms) > 100:
        PROTEIN_A3M_DIR = d
        break
# Pod fallback
if PROTEIN_A3M_DIR is None and os.path.isdir('/workspace/protein_msa_full/a3m_named'):
    PROTEIN_A3M_DIR = '/workspace/protein_msa_full/a3m_named'
if PROTEIN_A3M_DIR:
    n_a3m = len([f for f in os.listdir(PROTEIN_A3M_DIR) if f.endswith('.a3m')])
    print(f"Protein A3M dir: {PROTEIN_A3M_DIR} ({n_a3m} files)")
else:
    print("Protein A3M: not found (co-fold will run without protein MSA)")
USALIGN_BIN = _find(BUNDLE_DIR, os.path.join('scripts', 'USalign'))
PTX_SOURCE = _find(BUNDLE_DIR, 'Protenix_v1')
RNAPRO_SOURCE = _find(BUNDLE_DIR, 'RNAPro')
TBM_SCRIPT = _find(BUNDLE_DIR, os.path.join('scripts', 'tbm_phase1.py'))


# Verify critical files
for label, path in [
    ("RNAPro ckpt", RNAPRO_CHECKPOINT),
    ("PTX ckpt", PTX_CHECKPOINT),
    ("RibonanzaNet2", RIBONANZA_PATH),
    ("CCD cache", CCD_CACHE),
    ("Template DB", TEMPLATE_DB_DIR),
    ("RNA metadata", RNA_METADATA_CSV),
    ("Template meta", CA_TEMPLATE_META_PATH),
    ("TBM script", TBM_SCRIPT),
    ("PTX source", PTX_SOURCE),
    ("RNAPro source", RNAPRO_SOURCE),
]:
    exists = os.path.exists(path)
    size = ""
    if exists and os.path.isfile(path):
        size = f" ({os.path.getsize(path)/1e6:.1f} MB)"
    print(f"  {'OK' if exists else 'MISSING'}: {label}{size}")

# ============================================================
# 2. INSTALL PACKAGES
# ============================================================
print("\n=== Installing wheels ===")
import platform, tarfile
py_ver = f"cp{sys.version_info.major}{sys.version_info.minor}"
print(f"Python: {sys.version_info.major}.{sys.version_info.minor}, ABI: {py_ver}")

# Unpack tar.gz if wheels dir contains it
tar_files = glob.glob(os.path.join(WHEELS_DIR, '*.tar.gz'))
if tar_files and not glob.glob(os.path.join(WHEELS_DIR, '*.whl')):
    print(f"  Extracting {len(tar_files)} tar.gz archive(s)...")
    for tf in tar_files:
        with tarfile.open(tf, 'r:gz') as t:
            t.extractall(WHEELS_DIR)

# Install all wheels with --no-deps (don't touch Kaggle's torch/numpy/etc)
whl_files = sorted(glob.glob(os.path.join(WHEELS_DIR, '*.whl')))
# Filter: matching cpXY or pure python
compatible = []
for w in whl_files:
    bn = os.path.basename(w)
    if f'-{py_ver}-' in bn or 'py3-none' in bn or 'py2.py3-none' in bn or f'-cp3{sys.version_info.minor}-' in bn:
        compatible.append(w)
    # Also accept abi3 wheels
    elif '-cp3' in bn and '-abi3-' in bn:
        compatible.append(w)

print(f"  Total wheels: {len(whl_files)}, compatible: {len(compatible)}")
if compatible:
    r = subprocess.run(
        [sys.executable, '-m', 'pip', 'install', '--no-deps', '--quiet'] + compatible,
        capture_output=True, text=True, timeout=300
    )
    print(f"  Install: rc={r.returncode}")
    if r.returncode != 0:
        print(f"  STDERR: {r.stderr[:500]}")

# Install Protenix v1 from source (--no-deps to avoid pulling torch etc)
print("\n=== Installing Protenix v1 from source ===")
# Try pip install first; if it fails (no build deps), add to sys.path
try:
    import protenix
    print(f"  protenix already installed: {protenix.__file__}")
except ImportError:
    r = subprocess.run(
        [sys.executable, '-m', 'pip', 'install', '--no-deps', '--quiet', PTX_SOURCE],
        capture_output=True, text=True, timeout=120
    )
    if r.returncode != 0:
        print(f"  pip install failed (rc={r.returncode}): {r.stderr[:300] if r.stderr else ''}")
        print(f"  adding PTX_SOURCE to sys.path")
        sys.path.insert(0, PTX_SOURCE)
    else:
        print(f"  Installed OK")

# Verify critical imports
print("\n=== Verifying imports ===")
for pkg, mod in [
    ('biopython', 'Bio'), ('biotite', 'biotite'), ('rdkit', 'rdkit'),
    ('ml_collections', 'ml_collections'), ('optree', 'optree'),
    ('einops', 'einops'), ('gemmi', 'gemmi'), ('protenix', 'protenix'),
]:
    try:
        __import__(mod)
        print(f"  {pkg}: OK")
    except ImportError as e:
        print(f"  {pkg}: FAILED ({e})")

# Set environment
os.environ['LAYERNORM_TYPE'] = 'torch'
os.environ['TRIANGLE_ATTENTION'] = 'torch'
os.environ['TRIANGLE_MULTIPLICATIVE'] = 'torch'
#os.environ['NVIDIA_TF32_OVERRIDE'] = '1'
os.environ['PROTENIX_CCD_CACHE'] = CCD_CACHE

# Copy USalign to writable dir and make executable (Kaggle input is read-only)
if os.path.exists(USALIGN_BIN):
    usalign_local = os.path.join(OUTPUT_PATH, 'USalign')
    shutil.copy2(USALIGN_BIN, usalign_local)
    os.chmod(usalign_local, 0o755)
    USALIGN_BIN = usalign_local

# Protenix needs ~/common/ with CCD + cluster + obsolete files
# (otherwise download_inference_cache tries to download from internet)
ptx_common_dir = os.path.expanduser('~/common')
if not os.path.exists(ptx_common_dir):
    os.symlink(CCD_CACHE, ptx_common_dir)
    print(f"  PTX common symlink: {ptx_common_dir} → {CCD_CACHE}")

# Protenix checkpoint symlink: protenix CLI looks in ~/checkpoint/{model_name}.pt
ptx_ckpt_dir = os.path.expanduser('~/checkpoint')
os.makedirs(ptx_ckpt_dir, exist_ok=True)
ptx_link = os.path.join(ptx_ckpt_dir, PTX_MODEL_NAME + '.pt')
if not os.path.exists(ptx_link):
    os.symlink(PTX_CHECKPOINT, ptx_link)
    print(f"  PTX checkpoint symlink: {ptx_link} → {PTX_CHECKPOINT}")

# ============================================================
# 3. LOAD DATA
# ============================================================
print("\n=== Loading data ===")
test_seqs = pd.read_csv(os.path.join(DATA_PATH, 'test_sequences.csv'))
print(f"Test targets: {len(test_seqs)}")

if TEST_MODE:
    # Filter to multiple targets — check test first, then train
    found_rows = []
    train_df = None
    for tid in TEST_TIDS:
        if tid in test_seqs['target_id'].values:
            found_rows.append(test_seqs[test_seqs['target_id'] == tid].iloc[0])
        else:
            if train_df is None:
                train_path = os.path.join(DATA_PATH, 'train_sequences.csv')
                if os.path.exists(train_path):
                    train_df = pd.read_csv(train_path)
                else:
                    train_df = pd.DataFrame()
            if len(train_df) and tid in train_df['target_id'].values:
                found_rows.append(train_df[train_df['target_id'] == tid].iloc[0])
                print(f"TEST MODE: {tid} loaded from train_sequences.csv")
    if found_rows:
        test_seqs = pd.DataFrame(found_rows).reset_index(drop=True)
    else:
        test_seqs = test_seqs.iloc[:1].copy()
    TEST_TID = test_seqs.iloc[0]['target_id']  # for backwards compat
    print(f"TEST MODE: {len(test_seqs)} targets: {list(test_seqs['target_id'])}")

# ============================================================
# 4. TBM pipeline + Cross-Attention template reranker
# ============================================================
print("\n=== Loading TBM + Cross-Attention reranker ===")
import importlib.util

def _load_module(name, path):
    spec = importlib.util.spec_from_file_location(name, path)
    mod = importlib.util.module_from_spec(spec)
    sys.modules[name] = mod
    spec.loader.exec_module(mod)
    return mod

# Add scripts dir to path so inter-module imports work
scripts_dir = os.path.join(BUNDLE_DIR, 'scripts')
if scripts_dir not in sys.path:
    sys.path.insert(0, scripts_dir)
# Also add CA data dir (for any helper imports)
if CA_DATA_DIR not in sys.path:
    sys.path.insert(0, CA_DATA_DIR)

tbm = _load_module("tbm_phase1", TBM_SCRIPT)

# Load template DB
templates_full = tbm.load_template_db(TEMPLATE_DB_DIR)
aligner = tbm.make_aligner()
print(f"Template DB: {len(templates_full)} entries")

# Template meta (optional — not used by the diversity selector)
try:
    with open(CA_TEMPLATE_META_PATH) as f:
        template_meta = json.load(f)
    print(f"Template meta: {len(template_meta)} entries")
except FileNotFoundError:
    template_meta = {}
    print(f"Template meta: SKIPPED (not needed for diversity selector)")

# --- Cross-Attention TM predictor (diversity template selector) ---
import torch
import torch.nn as nn
import torch.nn.functional as F
import pickle as _pickle
import math as _math
import editdistance as _editdistance

_CA_PAD = 4; _CA_VOCAB = 6
_CA_NUC = {'A': 0, 'U': 1, 'G': 2, 'C': 3}

class _CASelfAttn(nn.Module):
    def __init__(self, d, h, do=0.1):
        super().__init__()
        self.h = h; self.dh = d // h; self.scale = self.dh ** -0.5
        self.qkv = nn.Linear(d, 3*d); self.out = nn.Linear(d, d); self.drop = nn.Dropout(do)
    def forward(self, x, mask=None):
        B, L, D = x.shape
        qkv = self.qkv(x).view(B, L, 3, self.h, self.dh)
        q, k, v = qkv.permute(2, 0, 3, 1, 4)
        a = (q @ k.transpose(-2, -1)) * self.scale
        if mask is not None: a = a.masked_fill(mask.unsqueeze(1).unsqueeze(2), float('-inf'))
        return self.out((self.drop(F.softmax(a, -1)) @ v).transpose(1, 2).contiguous().view(B, L, D))

class _CABlock(nn.Module):
    def __init__(self, d, h, do=0.1):
        super().__init__()
        self.n1 = nn.LayerNorm(d); self.attn = _CASelfAttn(d, h, do); self.n2 = nn.LayerNorm(d)
        self.ffn = nn.Sequential(nn.Linear(d, d*2), nn.GELU(), nn.Dropout(do), nn.Linear(d*2, d), nn.Dropout(do))
    def forward(self, x, m=None): x = x + self.attn(self.n1(x), m); return x + self.ffn(self.n2(x))

class _CrossAttnTM(nn.Module):
    """Pure sequence: embed → self-attn → bidir cross-attn → pool → predict TM."""
    def __init__(self, d=128, h=4, L=5, ml=512, do=0.1):
        super().__init__()
        self.emb = nn.Embedding(_CA_VOCAB, d, padding_idx=_CA_PAD)
        self.pos = nn.Embedding(ml, d)
        self.drop = nn.Dropout(do)
        self.layers = nn.ModuleList([_CABlock(d, h, do) for _ in range(L)])
        self.cq2t = nn.MultiheadAttention(d, h, dropout=do, batch_first=True)
        self.ct2q = nn.MultiheadAttention(d, h, dropout=do, batch_first=True)
        self.cnq = nn.LayerNorm(d); self.cnt = nn.LayerNorm(d)
        self.cffn = nn.Sequential(nn.Linear(d, d*2), nn.GELU(), nn.Dropout(do), nn.Linear(d*2, d), nn.Dropout(do))
        self.cn2 = nn.LayerNorm(d); self.fn = nn.LayerNorm(d)
        self.ctx_proj = nn.Linear(10, d)
        self.head = nn.Sequential(nn.Linear(d*2, d), nn.GELU(), nn.Dropout(do), nn.Linear(d, 1))
    def enc(self, ids, mask):
        B, L = ids.shape; p = torch.arange(L, device=ids.device).unsqueeze(0).expand(B, L)
        x = self.drop(self.emb(ids) + self.pos(p))
        for l in self.layers: x = l(x, mask)
        return x
    def forward(self, qi, qm, ti, tm, ctx=None):
        B = qi.size(0); dev = qi.device
        q = self.enc(qi, qm); t = self.enc(ti, tm)
        qn = self.cnq(q); tn = self.cnt(t)
        q = q + self.cq2t(qn, tn, tn, key_padding_mask=tm)[0]
        t = t + self.ct2q(tn, qn, qn, key_padding_mask=qm)[0]
        q = q + self.cffn(self.cn2(q)); q = self.fn(q)
        qv = (~qm).unsqueeze(-1).float(); pooled = (q * qv).sum(1) / qv.sum(1).clamp(min=1)
        ql = (~qm).sum(1).float(); tl = (~tm).sum(1).float()
        lf = torch.stack([torch.log(ql+1)/6, torch.log(tl+1)/6, ql/(tl+1)], -1)
        if ctx is not None: af = torch.cat([lf, ctx], -1)
        else: af = torch.cat([lf, torch.zeros(B, 7, device=dev)], -1)
        return self.head(torch.cat([pooled, self.ctx_proj(af)], -1)).squeeze(-1)

# Load cross-attention model
_CA_CKPT_PATH = os.path.join(CA_DATA_DIR, 'tm_crossattn_z1.pt')
if not os.path.exists(_CA_CKPT_PATH):
    _CA_CKPT_PATH = '/workspace/overnight_runs/Z1_d128_4h_5L_best.pt'  # fallback
_ca_model = None
if os.path.exists(_CA_CKPT_PATH):
    _ca_ckpt = torch.load(_CA_CKPT_PATH, map_location='cpu', weights_only=False)
    _ca_cfg = _ca_ckpt['config']
    _ca_model = _CrossAttnTM(_ca_cfg['d_model'], _ca_cfg.get('n_heads', 4),
                              _ca_cfg['n_layers'], _ca_cfg['max_len'])
    _ca_model.load_state_dict(_ca_ckpt['model'], strict=False)
    _ca_model.eval()
    print(f"Cross-Attention TM: loaded (ep{_ca_ckpt['epoch']}, d={_ca_cfg['d_model']}, "
          f"L={_ca_cfg['n_layers']}, AUC={_ca_ckpt['metrics'].get('auc', 0):.3f})")
else:
    print(f"WARNING: Cross-Attention TM model not found at {_CA_CKPT_PATH}")

def _ca_seq_to_ids(seq, max_len=512):
    return [_CA_NUC.get(c, _CA_PAD) for c in seq[:max_len]]

def _ca_predict_scores(q_seq, t_seqs):
    """Score template sequences against query using cross-attention model."""
    if _ca_model is None:
        return [0.0] * len(t_seqs)
    qi = _ca_seq_to_ids(q_seq); ql = len(qi)
    results = []
    BS = 64
    for i in range(0, len(t_seqs), BS):
        bt = t_seqs[i:i+BS]
        tl = [_ca_seq_to_ids(s) for s in bt]; mt = max(len(t) for t in tl)
        qt = torch.tensor([qi]*len(bt)); qm = torch.zeros(len(bt), ql, dtype=torch.bool)
        tt = torch.full((len(bt), mt), _CA_PAD); tm = torch.ones(len(bt), mt, dtype=torch.bool)
        ctx = torch.zeros(len(bt), 7)
        for j in range(len(bt)):
            tt[j,:len(tl[j])] = torch.tensor(tl[j]); tm[j,:len(tl[j])] = False
            lev = 1.0 - _editdistance.eval(q_seq[:500], bt[j][:500]) / max(len(q_seq[:500]), len(bt[j][:500]), 1)
            ctx[j, 0] = lev
        with torch.no_grad():
            logits = _ca_model(qt, qm, tt, tm, ctx)
            results.extend(torch.sigmoid(logits.float()).numpy().tolist())
    return results

# Load RNA metadata (used by the diversity selector for mmseqs_0.300 cluster lookup).
_rna_meta_df = pd.read_csv(RNA_METADATA_CSV)
print(f"RNA metadata: {len(_rna_meta_df)} rows (for mmseqs_0.300 cluster lookup)")

# ============================================================
# 5. Diversity template selector (bio / model / combined)
# ============================================================
def predict_diversity_for_target(tid, seq, segments, row, n_predictions=3):
    """Diversity template selector: Bio + Model + Combined, cluster-diverse.
    S1: Bio top-1 (safe baseline)
    S2: Model top-1 (different mmseqs cluster)
    S3: Bio×Model combined top-1 (different cluster from S1,S2)
    """
    pool = templates_full
    exclude = set()
    cands = tbm.find_templates(seq, pool, aligner, top_n=500, exclude_ids=exclude)

    if not cands:
        coords = generate_aform_helix(seq)
        return [(coords, 0.0)] * n_predictions

    rna_len = len(seq)
    K = len(cands)
    t_ids = [c[0] for c in cands]
    t_seqs = [c[1] for c in cands]
    t_sims = [c[2] for c in cands]

    # Cluster lookup for diversity
    def _get_cl(t):
        try:
            p = t.split('_')[0]
            rows = _rna_meta_df[_rna_meta_df['pdb_id'].str.upper() == p.upper()]
            if len(rows) > 0 and pd.notna(rows.iloc[0].get('mmseqs_0.300')):
                return str(rows.iloc[0]['mmseqs_0.300'])
        except: pass
        return f'_unk_{t}'

    def _adapt(idx):
        t_id, t_seq, sim, t_coords = cands[idx]
        coords, _ = tbm.adapt_template_to_query(seq, t_seq, t_coords, aligner)
        coords = tbm.adaptive_rna_constraints(coords, segments, confidence=max(sim, 0.3), passes=2)
        return coords

    # S1: Bio top-1 (always)
    bio_order = sorted(range(K), key=lambda i: -t_sims[i])
    s1_idx = bio_order[0]
    s1_cl = _get_cl(t_ids[s1_idx])
    used_cls = {s1_cl}; used_ids = {t_ids[s1_idx]}

    # For long sequences (>512nt): model unreliable → bio_sim fallback with cluster diversity
    use_model = rna_len <= 512 and _ca_model is not None

    if use_model:
        pred_scores = _ca_predict_scores(seq, t_seqs)
        combined = [max(t_sims[i], 0) * pred_scores[i] for i in range(K)]

        # S2: Model top-1 (different cluster)
        mod_order = sorted(range(K), key=lambda i: -pred_scores[i])
        s2_idx = None
        for mi in mod_order:
            cl = _get_cl(t_ids[mi])
            if cl not in used_cls and t_ids[mi] not in used_ids:
                s2_idx = mi; used_cls.add(cl); used_ids.add(t_ids[mi]); break
        if s2_idx is None: s2_idx = mod_order[0]

        # S3: Combined top-1 (different cluster)
        comb_order = sorted(range(K), key=lambda i: -combined[i])
        s3_idx = None
        for ci in comb_order:
            cl = _get_cl(t_ids[ci])
            if cl not in used_cls and t_ids[ci] not in used_ids:
                s3_idx = ci; break
        if s3_idx is None:
            for ci in comb_order:
                if t_ids[ci] not in used_ids: s3_idx = ci; break
        if s3_idx is None: s3_idx = comb_order[0]

        mode_label = "model"
    else:
        # Fallback: bio_sim ranks 2,3 with cluster diversity
        pred_scores = [0.0] * K
        combined = [0.0] * K
        s2_idx = None
        for bi in bio_order[1:]:
            cl = _get_cl(t_ids[bi])
            if cl not in used_cls and t_ids[bi] not in used_ids:
                s2_idx = bi; used_cls.add(cl); used_ids.add(t_ids[bi]); break
        if s2_idx is None: s2_idx = bio_order[1] if len(bio_order) > 1 else bio_order[0]

        s3_idx = None
        for bi in bio_order[1:]:
            cl = _get_cl(t_ids[bi])
            if cl not in used_cls and t_ids[bi] not in used_ids:
                s3_idx = bi; break
        if s3_idx is None:
            for bi in bio_order[1:]:
                if t_ids[bi] not in used_ids: s3_idx = bi; break
        if s3_idx is None: s3_idx = bio_order[min(2, len(bio_order)-1)]

        mode_label = f"bio_fallback(>{512}nt)" if rna_len > 512 else "bio_fallback(no_model)"

    s1_coords = _adapt(s1_idx)
    s2_coords = _adapt(s2_idx)
    s3_coords = _adapt(s3_idx)

    all_cls = {_get_cl(t_ids[s1_idx]), _get_cl(t_ids[s2_idx]), _get_cl(t_ids[s3_idx])}
    print(f"  [{tid}] TBM {mode_label} ({rna_len}nt, {K} cands, {len(all_cls)} clusters):")
    if use_model:
        print(f"    S1(bio):  {t_ids[s1_idx]:>12} sim={t_sims[s1_idx]:.3f} mod={pred_scores[s1_idx]:.3f}")
        print(f"    S2(mod):  {t_ids[s2_idx]:>12} sim={t_sims[s2_idx]:.3f} mod={pred_scores[s2_idx]:.3f}")
        print(f"    S3(comb): {t_ids[s3_idx]:>12} sim={t_sims[s3_idx]:.3f} mod={pred_scores[s3_idx]:.3f} "
              f"comb={combined[s3_idx]:.4f}")
    else:
        print(f"    S1(bio):  {t_ids[s1_idx]:>12} sim={t_sims[s1_idx]:.3f}")
        print(f"    S2(bio2): {t_ids[s2_idx]:>12} sim={t_sims[s2_idx]:.3f}")
        print(f"    S3(bio3): {t_ids[s3_idx]:>12} sim={t_sims[s3_idx]:.3f}")

    result = [(s1_coords, t_sims[s1_idx]), (s2_coords, t_sims[s2_idx]),
              (s3_coords, t_sims[s3_idx])]
    return result[:n_predictions]


def generate_aform_helix(sequence):
    """A-form RNA helix fallback."""
    n = len(sequence)
    coords = np.zeros((n, 3))
    rise, twist, radius = 2.81, np.radians(32.7), 9.4
    for i in range(n):
        angle = i * twist
        coords[i] = [radius * np.cos(angle), radius * np.sin(angle), i * rise]
    return coords


# ============================================================
# SEQUENCE CHUNKING UTILITIES
# ============================================================
def split_into_chunks(seq_len, max_len, overlap):
    """Split a sequence into overlapping (start, end) chunks.
    Last chunk is always full-size (anchored to end) to avoid short tails
    with poor edge quality. Overlap with penultimate chunk may be larger
    than `overlap` — that's fine, more Kabsch points = better alignment.
    """
    if seq_len <= max_len:
        return [(0, seq_len)]
    chunks = []
    step = max_len - overlap
    pos = 0
    while pos < seq_len:
        end = min(pos + max_len, seq_len)
        chunks.append((pos, end))
        if end == seq_len:
            break
        pos += step
    # Fix: if last chunk is shorter than max_len * 0.6, replace with
    # a full-size chunk anchored at the end (overlap with previous grows)
    if len(chunks) >= 2:
        last_s, last_e = chunks[-1]
        last_len = last_e - last_s
        if last_len < max_len * 0.6:
            new_start = max(0, seq_len - max_len)
            # Only replace if new chunk has at least `overlap` overlap with previous
            prev_s, prev_e = chunks[-2]
            if prev_e - new_start >= overlap:
                chunks[-1] = (new_start, seq_len)
    return chunks


def kabsch_align(P, Q):
    """Optimal rotation R and translation t so that R@P + t ≈ Q."""
    cp = P.mean(axis=0)
    cq = Q.mean(axis=0)
    H = (P - cp).T @ (Q - cq)
    U, _, Vt = np.linalg.svd(H)
    d = np.linalg.det(Vt.T @ U.T)
    S = np.eye(3)
    if d < 0:
        S[2, 2] = -1
    R = Vt.T @ S @ U.T
    t = cq - R @ cp
    return R, t


def stitch_chunk_coords(chunk_coords_list, chunk_ranges, seq_len):
    """Merge overlapping chunk coordinates using Kabsch alignment + narrow blend.

    Kabsch uses the FULL overlap for best rotation estimation.
    Blending is restricted to a narrow zone (BLEND_HALF residues each side of
    the midpoint). Outside the blend zone, each residue belongs exclusively to
    whichever chunk has it closer to center.  This avoids the quality loss from
    averaging mismatched edge predictions over large overlap regions.
    """
    BLEND_HALF = 6  # blend only ±6 residues around the handover point

    if len(chunk_coords_list) == 1:
        c = chunk_coords_list[0]
        if len(c) >= seq_len:
            return c[:seq_len].astype(np.float32)
        out = np.zeros((seq_len, 3), dtype=np.float32)
        out[:len(c)] = c
        return out

    # Align each chunk to previous using FULL overlap region (best Kabsch)
    aligned = [chunk_coords_list[0].copy().astype(np.float64)]
    for i in range(1, len(chunk_coords_list)):
        prev_s, prev_e = chunk_ranges[i - 1]
        cur_s, cur_e = chunk_ranges[i]
        ov_start = cur_s
        ov_end = min(prev_e, cur_e)
        ov_len = ov_end - ov_start

        cur_coords = chunk_coords_list[i].copy().astype(np.float64)

        if ov_len < 3:
            aligned.append(cur_coords)
            continue

        prev_ov = aligned[i - 1][ov_start - prev_s: ov_end - prev_s]
        cur_ov = cur_coords[ov_start - cur_s: ov_end - cur_s]

        valid = ~(np.isnan(prev_ov).any(axis=1) | np.isnan(cur_ov).any(axis=1))
        if valid.sum() < 3:
            aligned.append(cur_coords)
            continue

        R, t = kabsch_align(cur_ov[valid], prev_ov[valid])
        transformed = (cur_coords @ R.T) + t
        aligned.append(transformed)

    # Narrow-blend merge: hard switch with small crossfade around midpoint
    full = np.zeros((seq_len, 3), dtype=np.float64)
    weights = np.zeros(seq_len, dtype=np.float64)

    for i, ((s, e), coords) in enumerate(zip(chunk_ranges, aligned)):
        chunk_len = len(coords)
        actual_end = min(s + chunk_len, seq_len)
        used_len = actual_end - s
        w = np.ones(used_len, dtype=np.float64)

        if i > 0:
            ov_end = min(chunk_ranges[i - 1][1], e)
            ov_len = ov_end - s
            if ov_len > 0:
                # Midpoint of overlap
                mid = ov_len // 2
                # Before midpoint - BLEND_HALF: weight = 0 (previous chunk owns)
                zero_end = max(0, mid - BLEND_HALF)
                w[:zero_end] = 0.0
                # Narrow blend zone around midpoint
                blend_start = zero_end
                blend_end = min(ov_len, mid + BLEND_HALF)
                blend_len = blend_end - blend_start
                if blend_len > 0:
                    w[blend_start:blend_end] = np.linspace(0.0, 1.0, blend_len)
                # After midpoint + BLEND_HALF: weight = 1 (this chunk owns)

        if i < len(chunk_ranges) - 1:
            next_s = chunk_ranges[i + 1][0]
            ramp_start = next_s - s
            ov_len = actual_end - next_s
            if ov_len > 0 and ramp_start < used_len:
                mid = ramp_start + ov_len // 2
                # Before midpoint - BLEND_HALF: weight = 1 (this chunk owns)
                blend_start = max(ramp_start, mid - BLEND_HALF)
                blend_end = min(used_len, mid + BLEND_HALF)
                blend_len = blend_end - blend_start
                if blend_len > 0:
                    w[blend_start:blend_end] = np.linspace(1.0, 0.0, blend_len)
                # After midpoint + BLEND_HALF: weight = 0 (next chunk owns)
                zero_start = min(used_len, mid + BLEND_HALF)
                w[zero_start:used_len] = 0.0

        full[s:actual_end] += coords[:used_len] * w[:, None]
        weights[s:actual_end] += w

    mask = weights > 0
    full[mask] /= weights[mask, None]
    return full.astype(np.float32)


# ============================================================
# 6. PROTENIX V1 CO-FOLD FUNCTIONS
# ============================================================
def parse_chains(all_sequences_str):
    """Parse all_sequences FASTA into typed chains.

    Returns list of dicts: {header, sequence, type} where type is 'rna', 'protein', or 'dna'.
    RNA chains are returned separately (not concatenated) to support separate entities.
    """
    if pd.isna(all_sequences_str):
        return []
    seqs = [s for s in str(all_sequences_str).split(">") if s.strip()]
    chains = []
    for s in seqs:
        lines = s.strip().split("\n")
        header = lines[0]
        seq = "".join(lines[1:]).replace(" ", "")
        chars = set(seq.upper()) - set("- \n")
        if not chars:
            continue
        if chars <= set("ACGT") and "T" in chars:
            ctype = "dna"
        elif len(chars - set("AUGC")) > 2:
            ctype = "protein"
        else:
            ctype = "rna"
        chains.append({"header": header, "sequence": seq, "type": ctype})
    return chains


def slice_msa_fasta(msa_path, cs, ce, out_path):
    """Slice MSA FASTA columns [cs:ce] for a sequence chunk.

    Maps query (ungapped) positions cs:ce to alignment columns,
    then extracts those columns from all sequences, dropping gap-only rows.
    Returns out_path on success, None if source MSA missing/empty.
    """
    if not os.path.exists(msa_path):
        return None

    headers, sequences = [], []
    with open(msa_path) as f:
        cur_h, cur_s = None, []
        for line in f:
            line = line.rstrip()
            if line.startswith(">"):
                if cur_h is not None:
                    headers.append(cur_h)
                    sequences.append("".join(cur_s))
                cur_h = line
                cur_s = []
            else:
                cur_s.append(line)
        if cur_h is not None:
            headers.append(cur_h)
            sequences.append("".join(cur_s))

    if not sequences:
        return None

    # Map ungapped query positions to alignment columns
    query_seq = sequences[0]
    query_pos = 0
    col_start = col_end = None
    for col_idx, char in enumerate(query_seq):
        if char not in "-. ":
            if query_pos == cs:
                col_start = col_idx
            if query_pos == ce - 1:
                col_end = col_idx + 1
                break
            query_pos += 1
    if col_start is None:
        col_start = 0
    if col_end is None:
        col_end = len(query_seq)

    # Slice and filter gap-only sequences
    out_h, out_s = [], []
    for h, s in zip(headers, sequences):
        chunk = s[col_start:col_end]
        if chunk.replace("-", "").replace(".", "").strip():
            out_h.append(h)
            out_s.append(chunk)

    os.makedirs(os.path.dirname(out_path), exist_ok=True)
    with open(out_path, "w") as f:
        for h, s in zip(out_h, out_s):
            f.write(f"{h}\n{s}\n")

    return out_path


def _build_kmer_index(seq_to_path, k=3):
    """Build inverted k-mer index: kmer -> set of sequences containing it."""
    kmer_idx = {}
    for seq in seq_to_path:
        for i in range(len(seq) - k + 1):
            kmer = seq[i:i+k]
            if kmer not in kmer_idx:
                kmer_idx[kmer] = []
            kmer_idx[kmer].append(seq)
    return kmer_idx


def _find_homolog_a3m(query_seq, seq_to_path, kmer_idx, tid, k=3, min_identity=0.30):
    """Find best homologous protein MSA via k-mer filter + global alignment.

    Returns (adapted_a3m_path, identity) or (None, 0).
    """
    from collections import Counter
    # K-mer overlap scoring: count how many query k-mers each db seq shares
    query_kmers = set()
    for i in range(len(query_seq) - k + 1):
        query_kmers.add(query_seq[i:i+k])
    if not query_kmers:
        return None, 0

    hit_counts = Counter()
    for kmer in query_kmers:
        for seq in kmer_idx.get(kmer, []):
            hit_counts[seq] += 1

    if not hit_counts:
        return None, 0

    # Top candidates by Jaccard-like score
    candidates = []
    for seq, shared in hit_counts.most_common(20):
        seq_kmers = max(len(seq) - k + 1, 1)
        jaccard = shared / (len(query_kmers) + seq_kmers - shared)
        candidates.append((jaccard, seq))

    # Precise alignment on top-10
    try:
        from Bio.Align import PairwiseAligner
        aligner = PairwiseAligner()
        aligner.mode = 'global'
        aligner.match_score = 2.0
        aligner.mismatch_score = -1.5
        aligner.open_gap_score = -8.0
        aligner.extend_gap_score = -0.4
    except ImportError:
        # No BioPython — use simple k-mer identity estimate
        best_jacc, best_seq = candidates[0]
        if best_jacc >= min_identity:
            return _adapt_a3m(query_seq, seq_to_path[best_seq], tid), best_jacc
        return None, 0

    best_identity = 0
    best_seq = None
    for _, cand_seq in candidates[:10]:
        aln = aligner.align(query_seq, cand_seq)
        raw_score = aln.score if hasattr(aln, 'score') else aln[0].score
        max_possible = 2.0 * min(len(query_seq), len(cand_seq))
        identity = raw_score / max_possible if max_possible > 0 else 0
        if identity > best_identity:
            best_identity = identity
            best_seq = cand_seq

    if best_identity >= min_identity and best_seq is not None:
        adapted = _adapt_a3m(query_seq, seq_to_path[best_seq], tid)
        return adapted, best_identity

    return None, 0


def _adapt_a3m(query_seq, donor_a3m_path, tid):
    """Create adapted a3m: replace query (first) sequence, keep MSA hits.

    The donor a3m may have a different-length query. We keep all hit sequences
    as-is (they align to the donor, not query), but replace the first seq line.
    This is an approximation — works well for close homologs (>50% identity)
    where the MSA profile is transferable.
    """
    adapted_dir = os.path.join(OUTPUT_PATH, 'ptx_msa_homolog')
    os.makedirs(adapted_dir, exist_ok=True)
    import hashlib
    h = hashlib.md5(query_seq.encode()).hexdigest()[:12]
    adapted_path = os.path.join(adapted_dir, f"{tid}_{h}.a3m")

    with open(donor_a3m_path) as f:
        lines = f.readlines()

    # Replace first header + sequence
    out_lines = [f">query_{tid}\n", query_seq + "\n"]
    # Skip donor's first header + sequence, keep rest
    i = 0
    while i < len(lines) and lines[i].startswith(">"):
        i += 1  # skip first header
    if i < len(lines):
        i += 1  # skip first sequence
    # Append remaining MSA hits as-is
    out_lines.extend(lines[i:])

    with open(adapted_path, 'w') as f:
        f.writelines(out_lines)
    return adapted_path


def _find_protein_a3m(protein_seq, tid):
    """Find pre-computed protein MSA a3m for a given target+chain.

    Searches PROTEIN_A3M_DIR for {PDB}_{chain}.a3m files matching by sequence.
    Falls back to homologous MSA (>=30% identity) if exact match not found.
    Returns path to a3m file or None.
    """
    if PROTEIN_A3M_DIR is None:
        return None
    # Build index on first call: sequence -> filepath
    if not hasattr(_find_protein_a3m, '_cache'):
        _find_protein_a3m._cache = {}
        for fn in os.listdir(PROTEIN_A3M_DIR):
            if not fn.endswith('.a3m'):
                continue
            a3m_path = os.path.join(PROTEIN_A3M_DIR, fn)
            try:
                with open(a3m_path) as f:
                    hdr = f.readline()  # >header
                    seq = f.readline().strip()  # first sequence
                _find_protein_a3m._cache[seq] = a3m_path
            except Exception:
                continue
        _find_protein_a3m._kmer_idx = None  # built lazily on first miss
        print(f"  Protein A3M index: {len(_find_protein_a3m._cache)} sequences")

    # Exact match
    exact = _find_protein_a3m._cache.get(protein_seq)
    if exact:
        return exact

    # Homolog fallback: build k-mer index lazily on first miss
    if _find_protein_a3m._kmer_idx is None:
        _find_protein_a3m._kmer_idx = _build_kmer_index(_find_protein_a3m._cache, k=3)
        print(f"  Protein A3M k-mer index built ({len(_find_protein_a3m._kmer_idx)} 3-mers)")

    adapted, identity = _find_homolog_a3m(
        protein_seq, _find_protein_a3m._cache, _find_protein_a3m._kmer_idx, tid)
    if adapted:
        print(f"    [MSA] Homolog fallback for {tid}: identity={identity:.1%}, "
              f"len={len(protein_seq)}")
        return adapted
    return None


def _build_protein_entry(protein_seq, tid):
    """Build proteinChain JSON entry, with MSA if available."""
    entry = {"sequence": protein_seq, "count": 1}
    a3m = _find_protein_a3m(protein_seq, tid)
    if a3m:
        entry["unpairedMsaPath"] = a3m
    else:
        print(f"    [MSA] WARNING: no protein MSA for {tid} (len={len(protein_seq)})")
    return entry


def _trim_protein_a3m(a3m_path, trim_left, trim_right, out_path):
    """Trim a3m MSA symmetrically: remove trim_left cols from start, trim_right from end.

    a3m format: lowercase = insertions (not alignment columns), uppercase + '-' = columns.
    We trim by counting alignment columns (uppercase + '-'), skipping insertions.
    """
    os.makedirs(os.path.dirname(out_path), exist_ok=True)
    with open(a3m_path) as f:
        lines = f.readlines()
    out_lines = []
    for line in lines:
        if line.startswith(">"):
            out_lines.append(line)
            continue
        seq = line.rstrip("\n")
        trimmed = []
        col_idx = 0
        total_cols = sum(1 for ch in seq if ch == '-' or ch.isupper())
        keep_end = total_cols - trim_right
        for ch in seq:
            if ch == '-' or ch.isupper():
                if trim_left <= col_idx < keep_end:
                    trimmed.append(ch)
                col_idx += 1
            else:
                # Insertion (lowercase) — belongs to preceding column
                if trim_left <= (col_idx - 1) < keep_end and col_idx > 0:
                    trimmed.append(ch)
        out_lines.append("".join(trimmed) + "\n")
    with open(out_path, 'w') as f:
        f.writelines(out_lines)
    return out_path


def _build_trimmed_protein_entry(orig_seq, trimmed_seq, trim_left, trim_right, tid):
    """Build proteinChain entry with trimmed sequence and trimmed MSA."""
    entry = {"sequence": trimmed_seq, "count": 1}
    a3m = _find_protein_a3m(orig_seq, tid)
    if a3m:
        trimmed_dir = os.path.join(OUTPUT_PATH, 'ptx_msa_trimmed')
        trimmed_path = os.path.join(trimmed_dir, f"{tid}_trimmed.a3m")
        _trim_protein_a3m(a3m, trim_left, trim_right, trimmed_path)
        entry["unpairedMsaPath"] = trimmed_path
        print(f"    [PTX] Trimmed protein MSA: {len(orig_seq)} → {len(trimmed_seq)} "
              f"(cut {trim_left}+{trim_right})")
    else:
        print(f"    [MSA] WARNING: no protein MSA for trimmed {tid} (len={len(trimmed_seq)})")
    return entry


def _ensure_msa_file(rna_seq, tid):
    """Ensure an MSA file exists for the RNA sequence.

    Tries competition MSA, then external MSA, then creates a minimal
    single-sequence FASTA so that Protenix never invokes nhmmer.
    """
    # 1. Competition MSA
    msa_file = os.path.join(COMP_MSA_DIR, f"{tid}.MSA.fasta")
    if os.path.exists(msa_file):
        return msa_file

    # 2. External MSA via PDB→URS mapping
    if os.path.exists(PDB_TO_URS):
        try:
            with open(PDB_TO_URS) as f:
                pdb_to_urs = json.load(f)
            pdb_code = tid.split("_")[0].lower() if "_" in tid else tid.lower()
            urs_id = pdb_to_urs.get(pdb_code)
            if urs_id:
                ext_msa = os.path.join(EXT_MSA_DIR, urs_id, f"{urs_id}_all.a3m")
                if os.path.exists(ext_msa):
                    return ext_msa
        except Exception:
            pass

    # 3. Create minimal single-sequence FASTA (prevents nhmmer invocation)
    print(f"    [MSA] WARNING: no RNA MSA for {tid} — using dummy single-seq")
    dummy_dir = os.path.join(OUTPUT_PATH, 'ptx_msa_dummy')
    os.makedirs(dummy_dir, exist_ok=True)
    dummy_path = os.path.join(dummy_dir, f"{tid}.fasta")
    with open(dummy_path, 'w') as f:
        f.write(f">{tid}\n{rna_seq}\n")
    return dummy_path


def build_ptx_json(tid, rna_seq, chains, use_cofold, truncated):
    """Build Protenix v1 input JSON."""
    sequences = []
    rna_entry = {"sequence": rna_seq, "count": 1}

    # Always provide unpairedMsaPath to prevent nhmmer search (no internet on Kaggle)
    rna_entry["unpairedMsaPath"] = _ensure_msa_file(rna_seq, tid)

    sequences.append({"rnaSequence": rna_entry})

    if use_cofold:
        for c in chains:
            if c["type"] == "protein":
                sequences.append({"proteinChain": _build_protein_entry(c["sequence"], tid)})
            elif c["type"] == "dna":
                sequences.append({"dnaSequence": {"sequence": c["sequence"], "count": 1}})

    return [{"name": tid, "sequences": sequences, "modelSeeds": [PTX_SEED]}]


def extract_rna_c1(cif_path):
    """Extract RNA-only C1' coords from CIF, preserving entity order.

    Sorts by chain appearance order (not lexicographic) then res_id.
    This is critical for multi-entity Protenix output where chain IDs
    may be multi-letter (AA, AB, ...) and lexsort would break ordering.
    """
    from biotite.structure.io import pdbx

    with open(cif_path) as f:
        cif_data = pdbx.CIFFile.read(f)
    struct = pdbx.get_structure(cif_data, model=1)
    atom_names = np.char.strip(struct.atom_name.astype(str))
    c1_mask = atom_names == "C1'"
    rna_mask = np.isin(struct.res_name, ["A", "G", "C", "U"])
    rna_c1 = struct[c1_mask & rna_mask]

    # Build chain order map: first-appearance index for each chain_id
    chain_ids = rna_c1.chain_id.astype(str)
    seen = {}
    for cid in chain_ids:
        if cid not in seen:
            seen[cid] = len(seen)
    chain_order = np.array([seen[cid] for cid in chain_ids])

    sort_idx = np.lexsort((rna_c1.res_id, chain_order))
    return rna_c1[sort_idx].coord.astype(np.float32)


def setup_protenix_inprocess():
    """Load Protenix v1 model ONCE for in-process inference.

    Returns (runner, configs, infer_predict_fn, saved_cwd).
    """
    from typing import Mapping

    saved_cwd = os.getcwd()

    if PTX_SOURCE not in sys.path:
        sys.path.insert(0, PTX_SOURCE)
    os.chdir(PTX_SOURCE)

    # Dummy JSON for initial config parsing (overridden per target)
    dummy_json = os.path.join(OUTPUT_PATH, '_ptx_init.json')
    with open(dummy_json, 'w') as f:
        json.dump([{"name": "init", "sequences": [{"rnaSequence": {"sequence": "AAAA", "count": 1}}],
                     "modelSeeds": [PTX_SEED]}], f)

    saved_argv = sys.argv
    sys.argv = [
        "runner/inference.py",
        "--seeds", str(PTX_SEED),
        "--dump_dir", os.path.join(OUTPUT_PATH, 'ptx_preds'),
        "--input_json_path", dummy_json,
        "--model_name", PTX_MODEL_NAME,
        "--model.N_cycle", str(PTX_N_CYCLE),
        "--sample_diffusion.N_sample", str(PTX_N_SAMPLE),
        "--sample_diffusion.N_step", str(PTX_N_STEP),
        "--use_msa", "true",
        "--use_rna_msa", "true",
        "--use_template", "false",
        "--need_atom_confidence", "true",
        "--dtype", "fp32",
        "--triangle_multiplicative", "torch",
        "--triangle_attention", "torch",
    ]

    from runner.inference import (
        InferenceRunner, infer_predict,
        update_gpu_compatible_configs,
        download_inference_cache
    )
    from configs.configs_base import configs as configs_base
    from configs.configs_data import data_configs
    from configs.configs_inference import inference_configs
    from configs.configs_model_type import model_configs
    from protenix.config.config import parse_configs, parse_sys_args

    arg_str = parse_sys_args()
    configs = {**configs_base, **{"data": data_configs}, **inference_configs}
    configs = parse_configs(configs=configs, arg_str=arg_str, fill_required_with_null=True)

    base_configs = {**configs_base, **{"data": data_configs}, **inference_configs}
    model_specifics = model_configs.get(configs.model_name, {})

    def _deep_update(d, u):
        for k, v in u.items():
            if isinstance(v, Mapping) and k in d and isinstance(d[k], Mapping):
                _deep_update(d[k], v)
            else:
                d[k] = v
        return d

    if model_specifics:
        _deep_update(base_configs, model_specifics)
    configs = parse_configs(configs=base_configs, arg_str=arg_str, fill_required_with_null=True)
    configs = update_gpu_compatible_configs(configs)

    sys.argv = saved_argv

    t0 = time.time()
    download_inference_cache(configs)
    print(f"  [PTX] Cache check: {time.time()-t0:.1f}s")

    t0 = time.time()
    runner = InferenceRunner(configs)
    print(f"  [PTX] Model loaded in {time.time()-t0:.0f}s")

    return runner, configs, infer_predict, saved_cwd


# Global dict to pass chunk ranges to _run_ptx_chunk for MSA slicing
_ptx_chunk_ranges = {}  # chunk_tid -> (start, end) in original sequence


def _run_ptx_chunk(runner, configs, infer_predict_fn, tid, rna_seq, chains, ligands=None):
    """Run Protenix v1 inference for a single RNA chunk (or full sequence).

    Returns list of (rna_len, 3) coord arrays, or empty list on failure.
    """
    import torch
    rna_len = len(rna_seq)

    # Co-fold decision: include non-RNA chains if they fit in VRAM budget
    extra_tokens = sum(len(c["sequence"]) for c in chains if c["type"] in ("protein", "dna"))
    ligand_tokens = len(ligands) if ligands else 0
    total_tokens = rna_len + extra_tokens + ligand_tokens
    use_cofold = extra_tokens > 0 and total_tokens <= PTX_MAX_COFOLD_TOKENS

    # Protein trimming: if cofold overflows by ≤100 tokens, trim largest protein (≥500aa)
    PTX_COFOLD_TRIM_MARGIN = 100
    trimmed_protein_idx = None
    trimmed_protein_seq = None
    trim_left = trim_right = 0
    if (not use_cofold and extra_tokens > 0
            and total_tokens <= PTX_MAX_COFOLD_TOKENS + PTX_COFOLD_TRIM_MARGIN):
        overflow = total_tokens - PTX_MAX_COFOLD_TOKENS
        prot_chains = [(i, c) for i, c in enumerate(chains) if c["type"] == "protein"]
        if prot_chains:
            largest_idx, largest_chain = max(prot_chains, key=lambda x: len(x[1]["sequence"]))
            if len(largest_chain["sequence"]) >= 500:
                trim_left = overflow // 2
                trim_right = overflow - trim_left
                orig_seq = largest_chain["sequence"]
                trimmed_protein_seq = orig_seq[trim_left:len(orig_seq) - trim_right]
                trimmed_protein_idx = largest_idx
                total_tokens -= overflow
                use_cofold = True
                print(f"    [PTX] PROTEIN TRIM: {len(orig_seq)} → {len(trimmed_protein_seq)} aa "
                      f"(cut {trim_left}+{trim_right}, overflow was {overflow})")

    # For chunks, use original tid for MSA lookup but chunk tid for output
    orig_tid = tid.split("_c")[0] if "_c" in tid else tid
    is_chunk = (tid != orig_tid)

    if is_chunk:
        # Try slicing competition MSA for this chunk range
        full_msa = os.path.join(COMP_MSA_DIR, f"{orig_tid}.MSA.fasta")
        chunk_range = _ptx_chunk_ranges.get(tid, (0, len(rna_seq)))
        sliced_dir = os.path.join(OUTPUT_PATH, 'ptx_msa_sliced')
        sliced_path = os.path.join(sliced_dir, f"{tid}.fasta")
        msa_path = slice_msa_fasta(full_msa, chunk_range[0], chunk_range[1], sliced_path)
        if msa_path is not None:
            n_sliced = sum(1 for l in open(msa_path) if l.startswith(">"))
            print(f"    [PTX] Chunk MSA: sliced cols [{chunk_range[0]}:{chunk_range[1]}] → {n_sliced} seqs")
        if msa_path is None:
            # Fallback to dummy single-seq
            dummy_dir = os.path.join(OUTPUT_PATH, 'ptx_msa_dummy')
            os.makedirs(dummy_dir, exist_ok=True)
            msa_path = os.path.join(dummy_dir, f"{tid}.fasta")
            with open(msa_path, 'w') as f:
                f.write(f">{tid}\n{rna_seq}\n")
    else:
        msa_path = _ensure_msa_file(rna_seq, orig_tid)

    mode = "cofold" if use_cofold else "rna_only"
    if trimmed_protein_idx is not None:
        mode = "cofold+trim"
    chain_summary = ", ".join(f"{c['type']}({len(c['sequence'])})" for c in chains if c["type"] != "rna")
    msa_src = "sliced" if "msa_sliced" in msa_path else (
        "comp" if COMP_MSA_DIR in msa_path else ("ext" if "ext_msa" in msa_path else "dummy"))
    print(f"    [PTX] {tid}: rna={rna_len}, extra=[{chain_summary}], "
          f"tokens={total_tokens}, mode={mode}, msa={msa_src}")

    # Build JSON with separate RNA entities for multi-chain heterodimers
    sequences = []

    # Check if we have multiple RNA chains (heterodimer — different sequences)
    rna_chains = [c for c in chains if c["type"] == "rna"]
    if len(rna_chains) > 1 and not is_chunk:
        # Separate RNA entities: each gets its own rnaSequence entry
        # This gives +0.32 TM on heterodimers (EXP-B result)
        rna_seqs_distinct = []
        for rc in rna_chains:
            found = False
            for existing_seq, count_ref in rna_seqs_distinct:
                if rc["sequence"] == existing_seq:
                    count_ref[0] += 1
                    found = True
                    break
            if not found:
                rna_seqs_distinct.append((rc["sequence"], [1]))
        # Slice competition MSA per chain: MSA columns are concat of all RNA chains
        # in order. We know chain boundaries from cumulative lengths.
        full_msa = os.path.join(COMP_MSA_DIR, f"{orig_tid}.MSA.fasta")
        chain_offset = 0
        for idx, (rna_s, count_ref) in enumerate(rna_seqs_distinct):
            chain_len = len(rna_s) * count_ref[0]  # total nt for this entity (all copies)
            # Slice MSA columns for this chain's range
            sliced_dir = os.path.join(OUTPUT_PATH, 'ptx_msa_perchain')
            os.makedirs(sliced_dir, exist_ok=True)
            chain_msa = os.path.join(sliced_dir, f"{orig_tid}_chain{idx}.fasta")
            # For homodimer (count>1), MSA has count copies concatenated;
            # slice just the first copy's columns (they're identical)
            slice_start = chain_offset
            slice_end = chain_offset + len(rna_s)
            sliced = slice_msa_fasta(full_msa, slice_start, slice_end, chain_msa)
            if sliced is None:
                # Fallback to dummy if MSA not available
                dummy_dir = os.path.join(OUTPUT_PATH, 'ptx_msa_dummy')
                os.makedirs(dummy_dir, exist_ok=True)
                chain_msa = os.path.join(dummy_dir, f"{orig_tid}_chain{idx}.fasta")
                with open(chain_msa, 'w') as fmsa:
                    fmsa.write(f">{orig_tid}_chain{idx}\n{rna_s}\n")
                print(f"    [PTX] Chain {idx}: no MSA, using dummy")
            else:
                print(f"    [PTX] Chain {idx}: sliced MSA cols [{slice_start}:{slice_end}]")
            chain_offset += chain_len
            rna_entry = {"sequence": rna_s, "count": count_ref[0],
                         "unpairedMsaPath": chain_msa}
            sequences.append({"rnaSequence": rna_entry})
        print(f"    [PTX] Multi-RNA: {len(rna_seqs_distinct)} distinct entities "
              f"(from {len(rna_chains)} chains)")
    else:
        # Single RNA entity (original behavior, also for chunks)
        rna_entry = {"sequence": rna_seq, "count": 1, "unpairedMsaPath": msa_path}
        sequences.append({"rnaSequence": rna_entry})

    if use_cofold:
        for ci_chain, c in enumerate(chains):
            if c["type"] == "protein":
                if ci_chain == trimmed_protein_idx and trimmed_protein_seq is not None:
                    pentry = _build_trimmed_protein_entry(
                        c["sequence"], trimmed_protein_seq, trim_left, trim_right, orig_tid)
                else:
                    pentry = _build_protein_entry(c["sequence"], orig_tid)
                if "unpairedMsaPath" in pentry:
                    print(f"    [PTX] Protein MSA found for {orig_tid}")
                sequences.append({"proteinChain": pentry})
            elif c["type"] == "dna":
                sequences.append({"dnaSequence": {"sequence": c["sequence"], "count": 1}})

    # Ligand injection: add biologically relevant ligands as CCD entities
    if ligands:
        for lig in ligands:
            sequences.append({"ligand": lig})
        lig_names = [l["ligand"] for l in ligands]
        print(f"    [PTX] Ligands injected: {lig_names} (+{len(ligands)} tokens)")

    json_data = [{"name": tid, "sequences": sequences, "modelSeeds": [PTX_SEED]}]

    json_dir = os.path.join(OUTPUT_PATH, 'ptx_inputs')
    os.makedirs(json_dir, exist_ok=True)
    json_path = os.path.join(json_dir, f"{tid}.json")
    with open(json_path, "w") as f:
        json.dump(json_data, f)

    pred_dir = os.path.join(OUTPUT_PATH, 'ptx_preds', tid)
    if os.path.exists(pred_dir):
        shutil.rmtree(pred_dir)
    os.makedirs(pred_dir, exist_ok=True)

    # Update configs
    configs.input_json_path = json_path
    configs.dump_dir = pred_dir
    # CRITICAL: runner caches dump_dir→dumper.base_dir in __init__, must update all
    runner.dump_dir = pred_dir
    runner.dumper.base_dir = pred_dir
    runner.error_dir = os.path.join(pred_dir, "ERR")
    os.makedirs(runner.error_dir, exist_ok=True)

    # Dynamic VRAM settings — tuned for P100/T4 16GB (Kaggle GPU)
    # Key insight: OOM root cause was pair_transition [N,N,128]→[N,N,512] in fp32.
    # With AMP enabled (skip_amp=False), pair_transition runs in fp16 → half VRAM.
    # P100 16GB note: chunk=256 OOMs on 1136 tokens (softmax in tri_att_start), chunk=64 OK.
    if total_tokens > 800:
        configs.infer_setting["sample_diffusion_chunk_size"] = 1
        configs.infer_setting["chunk_size"] = 64
        configs.infer_setting["dynamic_chunk_size"] = False
        configs.skip_amp.confidence_head = False
        configs.skip_amp.sample_diffusion = False
        print(f"    [PTX] VRAM: high chunk=64, AMP=on ({total_tokens} tokens)")
    elif total_tokens > 600:
        configs.infer_setting["sample_diffusion_chunk_size"] = 1
        configs.infer_setting["chunk_size"] = 128
        configs.infer_setting["dynamic_chunk_size"] = False
        configs.skip_amp.confidence_head = False
        print(f"    [PTX] VRAM: mid chunk=128, conf_AMP=on ({total_tokens} tokens)")
    elif total_tokens > 400:
        configs.infer_setting["sample_diffusion_chunk_size"] = 1
        configs.infer_setting["chunk_size"] = 256
        configs.infer_setting["dynamic_chunk_size"] = False
        configs.skip_amp.confidence_head = False
        print(f"    [PTX] VRAM: low chunk=256, AMP=on ({total_tokens} tokens)")
    elif total_tokens > 200:
        configs.infer_setting["sample_diffusion_chunk_size"] = PTX_N_SAMPLE
        configs.infer_setting["chunk_size"] = 256
        configs.infer_setting["dynamic_chunk_size"] = False
        print(f"    [PTX] VRAM: chunk=256 ({total_tokens} tokens)")
    else:
        configs.infer_setting["sample_diffusion_chunk_size"] = PTX_N_SAMPLE
        configs.infer_setting["chunk_size"] = None
        configs.infer_setting["dynamic_chunk_size"] = False
        print(f"    [PTX] VRAM: no chunking ({total_tokens} tokens)")

    # OOM fallback: Protenix catches OOM internally (runner/inference.py:500)
    # so our try/except won't see it. Instead: detect via missing CIF output.
    oom_fallback_chain = [
        (64,  True,  "chunk=64,AMP=full"),
    ]

    t0 = time.time()
    succeeded = False
    try:
        infer_predict_fn(runner, configs)
    except Exception as e:
        print(f"    [PTX] ERROR (outer): {e}")

    # Check if CIF was produced (Protenix swallows OOM internally)
    cif_files = sorted(glob.glob(os.path.join(pred_dir, "**/*.cif"), recursive=True))
    if cif_files:
        succeeded = True
        # Partial OOM: some samples may have silently failed
        if len(cif_files) < PTX_N_SAMPLE:
            print(f"    [PTX] PARTIAL OOM: {len(cif_files)}/{PTX_N_SAMPLE} samples "
                  f"at {total_tokens} tokens (chunk={configs.infer_setting.get('chunk_size')})")
    else:
        print(f"    [PTX] No CIF — likely OOM at {total_tokens} tokens "
              f"(chunk={configs.infer_setting.get('chunk_size')})")
        import torch
        torch.cuda.empty_cache()

        for fb_chunk, fb_amp, fb_label in oom_fallback_chain:
            current = configs.infer_setting.get("chunk_size")
            if current is not None and current <= fb_chunk:
                continue
            print(f"    [PTX] OOM retry: {fb_label}")
            configs.infer_setting["sample_diffusion_chunk_size"] = 1
            configs.infer_setting["chunk_size"] = fb_chunk
            configs.infer_setting["dynamic_chunk_size"] = False
            if fb_amp:
                configs.skip_amp.confidence_head = False
                configs.skip_amp.sample_diffusion = False
            # Clean pred dir for fresh output
            if os.path.exists(pred_dir):
                shutil.rmtree(pred_dir)
            os.makedirs(pred_dir, exist_ok=True)
            configs.dump_dir = pred_dir
            # CRITICAL: runner caches dump_dir→dumper.base_dir in __init__, must update all
            runner.dump_dir = pred_dir
            runner.dumper.base_dir = pred_dir
            runner.error_dir = os.path.join(pred_dir, "ERR")
            os.makedirs(runner.error_dir, exist_ok=True)
            torch.cuda.empty_cache()
            try:
                infer_predict_fn(runner, configs)
            except Exception as e2:
                print(f"    [PTX] RETRY ERROR: {e2}")
            cif_files = sorted(glob.glob(os.path.join(pred_dir, "**/*.cif"), recursive=True))
            if cif_files:
                succeeded = True
                print(f"    [PTX] OOM recovered with {fb_label}")
                break
            torch.cuda.empty_cache()

        if not succeeded:
            # Cofold OOM fallback: retry RNA-only (drop all proteins/DNA)
            if use_cofold:
                print(f"    [PTX] COFOLD OOM — retrying RNA-only (drop proteins)")
                rna_only_seqs = [s for s in json.load(open(json_path))[0]["sequences"]
                                 if "rnaSequence" in s]
                json_rna = [{"name": tid, "sequences": rna_only_seqs, "modelSeeds": [PTX_SEED]}]
                with open(json_path, "w") as fj:
                    json.dump(json_rna, fj)
                rna_only_tokens = rna_len
                configs.infer_setting["sample_diffusion_chunk_size"] = 1
                configs.infer_setting["chunk_size"] = 64 if rna_only_tokens > 600 else 256
                configs.infer_setting["dynamic_chunk_size"] = False
                configs.skip_amp.confidence_head = False
                configs.skip_amp.sample_diffusion = False
                if os.path.exists(pred_dir):
                    shutil.rmtree(pred_dir)
                os.makedirs(pred_dir, exist_ok=True)
                configs.dump_dir = pred_dir
                runner.dump_dir = pred_dir
                runner.dumper.base_dir = pred_dir
                runner.error_dir = os.path.join(pred_dir, "ERR")
                os.makedirs(runner.error_dir, exist_ok=True)
                torch.cuda.empty_cache()
                try:
                    infer_predict_fn(runner, configs)
                except Exception as e3:
                    print(f"    [PTX] RNA-only retry ERROR: {e3}")
                cif_files = sorted(glob.glob(os.path.join(pred_dir, "**/*.cif"), recursive=True))
                if cif_files:
                    succeeded = True
                    print(f"    [PTX] RNA-only fallback OK: {len(cif_files)} samples")
                else:
                    print(f"    [PTX] RNA-only fallback also failed")

            if not succeeded:
                print(f"    [PTX] ALL RETRIES EXHAUSTED for {total_tokens} tokens")
                return []
    elapsed = time.time() - t0

    # cif_files already populated above (from fallback logic)
    if not cif_files:
        print(f"    [PTX] No CIF output")
        return []

    all_coords = []
    for cif_path in cif_files:
        try:
            coords_i = extract_rna_c1(cif_path)
            if len(coords_i) < rna_len:
                coords_i = np.concatenate([coords_i,
                    np.zeros((rna_len - len(coords_i), 3), dtype=np.float32)])
            elif len(coords_i) > rna_len:
                coords_i = coords_i[:rna_len]
            all_coords.append(coords_i)
        except Exception as e:
            print(f"    [PTX] CIF parse error: {e}")

    print(f"    [PTX] {tid}: {len(all_coords)} samples in {elapsed:.0f}s")
    return all_coords


# --- Chain-aware Protenix packing for multi-chain RNA targets ---

def _parse_target_rna_chains(rna_seq_full, all_sequences, stoichiometry):
    """Parse target RNA chains from stoichiometry + all_sequences.

    Returns [(start, end, chain_seq, auth_chain_id, count), ...] sorted by start,
    or [] if stoichiometry is empty/unparseable (triggers flat chunking fallback).
    """
    if pd.isna(stoichiometry) or not str(stoichiometry).strip():
        return []

    # Parse stoichiometry → {auth_chain_id: copy_count}
    stoi_map = {}
    for part in str(stoichiometry).split(";"):
        part = part.strip()
        if ":" not in part:
            continue
        chain_id, count_str = part.split(":", 1)
        stoi_map[chain_id.strip()] = int(count_str.strip())

    if not stoi_map:
        return []

    # Parse all_sequences → RNA entries with auth chain IDs
    rna_entries = []  # (auth_id, seq)
    for entry in str(all_sequences).split(">"):
        if not entry.strip():
            continue
        lines = entry.strip().split("\n")
        header = lines[0]
        seq = "".join(lines[1:]).replace(" ", "")
        # Only RNA (AUGC only, no non-canonical)
        if not seq or not set(seq.upper()) <= {"A", "U", "G", "C"}:
            continue
        # Extract auth chain ID: "Chain XX[auth YY]"
        auth_match = re.search(r"auth ([^\]]+)", header)
        if not auth_match:
            continue
        auth_id = auth_match.group(1).strip()
        if auth_id in stoi_map:
            rna_entries.append((auth_id, seq))

    if not rna_entries:
        return []

    # Map chains to positions in concatenated rna_seq_full
    # Two-pass: first find ALL possible positions, then greedy non-overlapping assignment
    candidates = []  # (auth_id, seq, count, all_positions)
    for auth_id, seq in rna_entries:
        count = stoi_map.get(auth_id, 1)
        # Find all occurrences of this seq in rna_seq_full
        positions = []
        start = 0
        while True:
            idx = rna_seq_full.find(seq, start)
            if idx < 0:
                break
            positions.append(idx)
            start = idx + 1  # allow overlapping finds
        candidates.append((auth_id, seq, count, positions))

    # Greedy assignment: process chains by position of their first occurrence
    # to maintain original sequence order
    result = []
    used = set()  # positions already assigned

    # Flatten: for each chain × count, pick earliest unused position
    assign_list = []
    for auth_id, seq, count, positions in candidates:
        for _ in range(count):
            assign_list.append((auth_id, seq, positions))

    # Sort by earliest available position (to assign in sequence order)
    # Multiple passes to handle dependencies
    for _ in range(3):  # max 3 passes for convergence
        changed = False
        for i, (auth_id, seq, positions) in enumerate(assign_list):
            if id((auth_id, seq, positions)) in used:
                continue
            for pos in sorted(positions):
                # Check no overlap with already assigned
                overlaps = False
                for rs, re_end, _, _, _ in result:
                    if pos < re_end and pos + len(seq) > rs:
                        overlaps = True
                        break
                if not overlaps:
                    result.append((pos, pos + len(seq), seq, auth_id, 1))
                    used.add(id((auth_id, seq, positions)))
                    changed = True
                    break
        if not changed:
            break

    result.sort(key=lambda x: x[0])
    return result


def _select_proteins_for_cofold(chains, budget):
    """Select small protein chains that fit within token budget.

    Args:
        chains: output of parse_chains(all_sequences)
        budget: max total protein tokens

    Returns:
        (selected_chains, total_tokens)
    """
    proteins = [c for c in chains if c["type"] == "protein"]
    proteins.sort(key=lambda c: len(c["sequence"]))

    selected = []
    total = 0
    for p in proteins:
        plen = len(p["sequence"])
        if total + plen <= budget:
            selected.append(p)
            total += plen
    return selected, total


def _pack_rna_groups(rna_chain_list, rna_budget):
    """Pack RNA chains into groups that fit within rna_budget tokens.

    Each group will be one Protenix inference call.
    Tries to keep whole chains together; only splits if a single chain exceeds budget.

    Returns: list of groups, each group = list of (start, end, seq, chain_id, count)
    """
    groups = []
    current_group = []
    current_tokens = 0

    for chain_info in rna_chain_list:
        start, end, seq, chain_id, count = chain_info
        chain_tokens = len(seq) * count

        if chain_tokens > rna_budget:
            # Single chain too large — flush current group, then split this chain
            if current_group:
                groups.append(current_group)
                current_group = []
                current_tokens = 0
            if count == 1:
                # Split the chain into sub-chunks
                sub_chunks = split_into_chunks(len(seq), rna_budget, PTX_CHUNK_OVERLAP)
                for cs, ce in sub_chunks:
                    groups.append([(start + cs, start + ce, seq[cs:ce], chain_id, 1)])
            else:
                # Multi-copy chain too large — put as single group, may OOM
                groups.append([chain_info])
            continue

        if current_tokens + chain_tokens > rna_budget:
            # Start new group
            if current_group:
                groups.append(current_group)
            current_group = [chain_info]
            current_tokens = chain_tokens
        else:
            current_group.append(chain_info)
            current_tokens += chain_tokens

    if current_group:
        groups.append(current_group)

    return groups


def _run_ptx_chainpack(runner, configs, infer_predict_fn,
                       tid, rna_seq_full, rna_chains, all_chains, ligands):
    """Chain-aware Protenix inference for multi-chain RNA targets.

    Instead of flat residue chunking, packs whole RNA chains into groups
    with a fixed protein subset. Each group = one Protenix call.
    """
    import torch

    full_len = len(rna_seq_full)

    # Select proteins
    selected_prots, prot_tokens = _select_proteins_for_cofold(
        all_chains, PTX_COFOLD_PROTEIN_BUDGET)
    rna_budget = PTX_MAX_COFOLD_TOKENS - prot_tokens

    # Pack RNA chains into groups
    groups = _pack_rna_groups(rna_chains, rna_budget)

    prot_summary = ", ".join(f"{len(p['sequence'])}aa" for p in selected_prots)
    print(f"    [PTX] CHAINPACK: {len(groups)} groups, "
          f"rna_budget={rna_budget}, prots=[{prot_summary}] ({prot_tokens}tok)")
    for gi, grp in enumerate(groups):
        chain_desc = ", ".join(f"{ci}:{e-s}nt" for s, e, _, ci, _ in grp)
        grp_tokens = sum(len(seq) * cnt for _, _, seq, _, cnt in grp)
        print(f"      group {gi}: {chain_desc} ({grp_tokens}tok)")

    # Run each group
    group_results = []  # [(global_ranges, [sample_coords, ...]), ...]

    for gi, grp in enumerate(groups):
        grp_tid = f"{tid}_g{gi}"
        grp_rna_total = sum(len(seq) * cnt for _, _, seq, _, cnt in grp)

        # Build protein-only chains list for _run_ptx_chunk
        prot_chains_for_chunk = selected_prots  # same for all groups

        # Build custom JSON for this group: separate rnaSequence per chain
        # We need to bypass _run_ptx_chunk's JSON building and do it ourselves
        msa_dir = os.path.join(OUTPUT_PATH, 'ptx_msa_chainpack')
        os.makedirs(msa_dir, exist_ok=True)

        sequences = []
        chain_order = []  # track (global_start, global_end, n_residues) for coord mapping

        # Group identical RNA sequences for Protenix count>1
        seq_groups = OrderedDict()  # seq → [(start, end, chain_id, count), ...]
        for start, end, seq, chain_id, count in grp:
            seq_groups.setdefault(seq, []).append((start, end, chain_id, count))

        for seq, entries in seq_groups.items():
            total_count = sum(cnt for _, _, _, cnt in entries)
            # MSA: try slicing from competition MSA
            first_start = entries[0][0]
            first_end = entries[0][1]
            chain_msa = os.path.join(msa_dir, f"{grp_tid}_{first_start}.fasta")
            orig_tid = tid.split("_c")[0] if "_c" in tid else tid
            full_msa = os.path.join(COMP_MSA_DIR, f"{orig_tid}.MSA.fasta")
            sliced = slice_msa_fasta(full_msa, first_start, first_end, chain_msa)
            if sliced is None:
                # Dummy single-seq MSA
                with open(chain_msa, 'w') as f:
                    f.write(f">{grp_tid}\n{seq}\n")
            rna_entry = {"sequence": seq, "count": total_count,
                         "unpairedMsaPath": chain_msa}
            sequences.append({"rnaSequence": rna_entry})
            # Track for coord mapping: each copy contributes len(seq) residues
            for s, e, ci, cnt in entries:
                for _ in range(cnt):
                    chain_order.append((s, e))

        # Add proteins
        for p in selected_prots:
            pentry = _build_protein_entry(p["sequence"], tid)
            sequences.append({"proteinChain": pentry})

        # Add ligands
        if ligands:
            for lig in ligands:
                sequences.append({"ligand": lig})

        json_data = [{"name": grp_tid, "sequences": sequences, "modelSeeds": [PTX_SEED]}]
        json_dir = os.path.join(OUTPUT_PATH, 'ptx_inputs')
        os.makedirs(json_dir, exist_ok=True)
        json_path = os.path.join(json_dir, f"{grp_tid}.json")
        with open(json_path, "w") as f:
            json.dump(json_data, f)

        total_tokens = grp_rna_total + prot_tokens + (len(ligands) if ligands else 0)
        pred_dir = os.path.join(OUTPUT_PATH, 'ptx_preds', grp_tid)
        if os.path.exists(pred_dir):
            shutil.rmtree(pred_dir)
        os.makedirs(pred_dir, exist_ok=True)

        # Configure VRAM
        if total_tokens > 800:
            configs.infer_setting["sample_diffusion_chunk_size"] = 1
            configs.infer_setting["chunk_size"] = 64
            configs.infer_setting["dynamic_chunk_size"] = False
            configs.skip_amp.confidence_head = False
            configs.skip_amp.sample_diffusion = False
        elif total_tokens > 400:
            configs.infer_setting["sample_diffusion_chunk_size"] = 1
            configs.infer_setting["chunk_size"] = 256
            configs.infer_setting["dynamic_chunk_size"] = False
            configs.skip_amp.confidence_head = False

        configs.input_json_path = json_path
        configs.dump_dir = pred_dir
        runner.dump_dir = pred_dir
        runner.dumper.base_dir = pred_dir
        runner.error_dir = os.path.join(pred_dir, "ERR")
        os.makedirs(runner.error_dir, exist_ok=True)

        torch.cuda.empty_cache()
        gc.collect()

        print(f"    [PTX] group {gi}: {total_tokens} tokens, inferring...")
        t0 = time.time()
        succeeded = False
        try:
            infer_predict_fn(runner, configs)
        except Exception as e:
            print(f"    [PTX] group {gi} ERROR: {e}")

        cif_files = sorted(glob.glob(os.path.join(pred_dir, "**/*.cif"), recursive=True))
        if cif_files:
            succeeded = True
            if len(cif_files) < PTX_N_SAMPLE:
                print(f"    [PTX] group {gi}: PARTIAL OOM {len(cif_files)}/{PTX_N_SAMPLE}")
        else:
            # Cofold OOM → retry RNA-only for this group (drop proteins)
            torch.cuda.empty_cache()
            if selected_prots:
                print(f"    [PTX] group {gi} COFOLD OOM — retrying RNA-only")
                rna_only_seqs = [s for s in json_data[0]["sequences"] if "rnaSequence" in s]
                json_rna = [{"name": grp_tid, "sequences": rna_only_seqs,
                             "modelSeeds": [PTX_SEED]}]
                with open(json_path, "w") as fj:
                    json.dump(json_rna, fj)
                rna_only_tok = grp_rna_total
                configs.infer_setting["chunk_size"] = 64 if rna_only_tok > 600 else 256
                configs.skip_amp.confidence_head = False
                configs.skip_amp.sample_diffusion = False
                if os.path.exists(pred_dir):
                    shutil.rmtree(pred_dir)
                os.makedirs(pred_dir, exist_ok=True)
                configs.dump_dir = pred_dir
                runner.dump_dir = pred_dir
                runner.dumper.base_dir = pred_dir
                runner.error_dir = os.path.join(pred_dir, "ERR")
                os.makedirs(runner.error_dir, exist_ok=True)
                torch.cuda.empty_cache()
                try:
                    infer_predict_fn(runner, configs)
                except Exception as e3:
                    print(f"    [PTX] group {gi} RNA-only ERROR: {e3}")
                cif_files = sorted(glob.glob(os.path.join(pred_dir, "**/*.cif"), recursive=True))
                if cif_files:
                    succeeded = True
                    print(f"    [PTX] group {gi} RNA-only fallback OK: {len(cif_files)} samples")

        elapsed = time.time() - t0

        if not cif_files:
            print(f"    [PTX] group {gi}: ALL RETRIES FAILED ({elapsed:.0f}s)")
            group_results.append((chain_order, []))
            continue

        # Extract RNA C1' from each CIF sample
        sample_coords = []
        for cif_path in cif_files:
            try:
                coords = extract_rna_c1(cif_path)
                sample_coords.append(coords)
            except Exception as e:
                print(f"    [PTX] group {gi} CIF parse error: {e}")

        print(f"    [PTX] group {gi}: {len(sample_coords)} samples, "
              f"{sample_coords[0].shape[0] if sample_coords else 0} C1' atoms ({elapsed:.0f}s)")
        group_results.append((chain_order, sample_coords))

    # Assemble: map group coords back to full sequence positions
    if not group_results or all(not sc for _, sc in group_results):
        return []

    n_samples = min(len(sc) for _, sc in group_results if sc) if any(sc for _, sc in group_results) else 0
    if n_samples == 0:
        return []

    all_stitched = []
    for s_idx in range(n_samples):
        full_coords = np.full((full_len, 3), np.nan, dtype=np.float32)

        for chain_order, sample_coords_list in group_results:
            if not sample_coords_list or s_idx >= len(sample_coords_list):
                continue
            coords = sample_coords_list[s_idx]
            # Map coords back to global positions using chain_order
            pos = 0
            for global_start, global_end in chain_order:
                chain_len = global_end - global_start
                if pos + chain_len <= len(coords):
                    full_coords[global_start:global_end] = coords[pos:pos + chain_len]
                pos += chain_len

        # Fill any remaining NaN with linear interpolation
        for d in range(3):
            valid = np.isfinite(full_coords[:, d])
            if valid.sum() >= 2 and (~valid).any():
                full_coords[:, d] = np.interp(
                    np.arange(full_len), np.where(valid)[0], full_coords[valid, d])

        all_stitched.append(full_coords)

    print(f"    [PTX] CHAINPACK done: {len(all_stitched)} samples × {full_len} residues")
    return all_stitched


def run_protenix_target_inprocess(runner, configs, infer_predict_fn, tid, rna_seq_full, all_sequences, ligands=None, stoichiometry=""):
    """Run Protenix v1 in-process for one target. Single pass only.

    Logic:
      - If target has protein/DNA AND fits in token budget → cofold ONLY.
        RNA-only runs ONLY as OOM fallback.
      - If target is RNA-only OR doesn't fit in token budget → rna_only mode.

    Returns list of (L, 3) coord arrays, or empty list on failure.
    """
    import torch
    chains = parse_chains(all_sequences)
    full_len = len(rna_seq_full)
    rna_chains = _parse_target_rna_chains(rna_seq_full, all_sequences, stoichiometry)

    # Compute total tokens for cofold
    total_tokens = sum(len(c["sequence"]) for c in chains)
    has_nonrna = any(c["type"] in ("protein", "dna") for c in chains)
    fits_cofold = total_tokens <= PTX_MAX_COFOLD_TOKENS
    rna_only_chains = [c for c in chains if c["type"] == "rna"]

    # Decide mode
    if has_nonrna and fits_cofold:
        mode = "cofold"
        use_chains = chains
    else:
        mode = "rna_only"
        use_chains = rna_only_chains
    print(f"    [PTX] {tid}: mode={mode}, rna={full_len}, total_tok={total_tokens}, "
          f"has_nonrna={has_nonrna}, fits={fits_cofold}")

    # ── Single inference pass ──
    primary_result = []

    def _do_inference(run_chains, run_ligands, label=""):
        """Run inference with given chains. Returns list of coord arrays."""
        result = []
        # Chain-aware packing for multi-chain RNA
        if len(rna_chains) > 1 and full_len > PTX_MAX_RNA_LEN:
            try:
                r = _run_ptx_chainpack(runner, configs, infer_predict_fn,
                                       tid + label, rna_seq_full, rna_chains,
                                       run_chains, run_ligands)
                if r:
                    return r
            except Exception as e:
                print(f"    [PTX] chainpack error: {e}, falling back to flat")
                traceback.print_exc()

        # Split into chunks if needed
        chunks = split_into_chunks(full_len, PTX_MAX_RNA_LEN, PTX_CHUNK_OVERLAP)

        if len(chunks) == 1:
            coords = _run_ptx_chunk(runner, configs, infer_predict_fn,
                                    tid + label, rna_seq_full, run_chains,
                                    ligands=run_ligands)
            for c in coords:
                if len(c) < full_len:
                    c = np.concatenate([c, np.zeros((full_len - len(c), 3), dtype=np.float32)])
                elif len(c) > full_len:
                    c = c[:full_len]
                result.append(c)
            if result:
                print(f"    [PTX] OK{label}: {len(result)} samples, {full_len} residues")
        else:
            print(f"    [PTX] CHUNKING{label}: {len(chunks)} chunks for {full_len} nt")
            chunk_all_coords = []
            for ci, (cs, ce) in enumerate(chunks):
                chunk_seq = rna_seq_full[cs:ce]
                chunk_tid = f"{tid}{label}_c{ci}"
                chunk_len = ce - cs
                _ptx_chunk_ranges[chunk_tid] = (cs, ce)
                torch.cuda.empty_cache(); gc.collect()
                coords = _run_ptx_chunk(runner, configs, infer_predict_fn,
                                        chunk_tid, chunk_seq, run_chains,
                                        ligands=run_ligands)
                if not coords:
                    print(f"    [PTX] Chunk {ci} FAILED, using zeros")
                    coords = [np.zeros((chunk_len, 3), dtype=np.float32)] * PTX_N_SAMPLE
                fixed = []
                for c in coords:
                    if len(c) != chunk_len:
                        c_fixed = np.zeros((chunk_len, 3), dtype=np.float32)
                        n = min(len(c), chunk_len)
                        c_fixed[:n] = c[:n]
                        c = c_fixed
                    fixed.append(c)
                chunk_all_coords.append(((cs, ce), fixed))
            n_samples = min(len(cc) for _, cc in chunk_all_coords)
            if n_samples > 0:
                for s_idx in range(n_samples):
                    sample_chunks = [cc[min(s_idx, len(cc) - 1)] for _, cc in chunk_all_coords]
                    ranges = [r for r, _ in chunk_all_coords]
                    stitched = stitch_chunk_coords(sample_chunks, ranges, full_len)
                    result.append(stitched)
                print(f"    [PTX] Stitched{label}: {len(result)} samples × {full_len} residues")
        return result

    # Run primary mode
    try:
        primary_result = _do_inference(use_chains, ligands if mode == "cofold" else None)
    except RuntimeError as e:
        if "out of memory" in str(e).lower() and mode == "cofold":
            # OOM on cofold → fallback to RNA-only
            print(f"    [PTX] OOM on cofold, falling back to rna_only")
            torch.cuda.empty_cache(); gc.collect()
            try:
                primary_result = _do_inference(rna_only_chains, None, "_rnaonly")
            except Exception as e2:
                print(f"    [PTX] RNA-only fallback also failed: {e2}")
        else:
            print(f"    [PTX] Error: {e}")
            traceback.print_exc()

    return primary_result


# ============================================================
# 7. RNAPRO IN-MEMORY INFERENCE (v7: load once, predict per target)
# ============================================================
def setup_rnapro_inprocess():
    """Load RNAPro model ONCE for in-process inference.

    Returns (runner, configs, infer_predict_fn, rnapro_dir, saved_cwd).
    Mirrors the Protenix setup_protenix_inprocess() pattern.
    """
    import torch

    saved_cwd = os.getcwd()

    # Copy RNAPro source to writable dir
    rnapro_dir = os.path.join(OUTPUT_PATH, 'RNAPro')
    if os.path.exists(rnapro_dir):
        shutil.rmtree(rnapro_dir)
    shutil.copytree(RNAPRO_SOURCE, rnapro_dir)

    # Symlink CCD cache
    rnapro_ccd = os.path.join(rnapro_dir, 'release_data', 'ccd_cache')
    os.makedirs(os.path.dirname(rnapro_ccd), exist_ok=True)
    if not os.path.exists(rnapro_ccd):
        os.symlink(CCD_CACHE, rnapro_ccd)

    # Symlink RibonanzaNet2
    rnapro_ribo = os.path.join(rnapro_dir, 'release_data', 'ribonanzanet2_checkpoint')
    if not os.path.exists(rnapro_ribo):
        os.symlink(RIBONANZA_PATH, rnapro_ribo)

    sys.path.insert(0, rnapro_dir)
    os.chdir(rnapro_dir)

    # Parse configs via CLI args (same as run() does, but we keep the runner)
    saved_argv = sys.argv
    sys.argv = [
        "runner/inference.py",
        "--model_name=rnapro_base",
        "--seeds=42",
        f"--dump_dir={os.path.join(OUTPUT_PATH, 'rnapro_output')}",
        f"--sequences_csv={os.path.join(OUTPUT_PATH, '_rnapro_dummy.csv')}",
        "--dtype=fp32",
        f"--model.N_cycle={RNAPRO_N_CYCLE}",
        f"--sample_diffusion.N_sample={RNAPRO_N_SAMPLE}",
        f"--sample_diffusion.N_step={RNAPRO_N_STEP}",
        "--use_msa=True",
        f"--rna_msa_dir={COMP_MSA_DIR}",
        f"--load_checkpoint_path={RNAPRO_CHECKPOINT}",
        "--model.use_RibonanzaNet2=True",
        f"--model.ribonanza_net_path={RIBONANZA_PATH}",
        "--use_template=masked_templates",
        "--model.use_template=masked_templates",
        "--model.template_embedder.n_blocks=2",
        "--triangle_attention=torch",
        "--triangle_multiplicative=torch",
        "--load_strict=False",
        f"--max_len={RNAPRO_MAX_LEN}",
        "--num_workers=0",
        "--logger=logging",
        "--n_templates_inf=5",
        f"--template_data={os.path.join(rnapro_dir, 'test_templates.pt')}",
    ]

    # Write dummy CSV so config parsing doesn't fail
    pd.DataFrame([{"target_id": "dummy", "sequence": "AAAA"}]).to_csv(
        os.path.join(OUTPUT_PATH, '_rnapro_dummy.csv'), index=False)

    from configs.configs_base import configs as configs_base_rn
    from configs.configs_data import data_configs as data_configs_rn
    from configs.configs_inference import inference_configs as inference_configs_rn
    from rnapro.config import parse_configs as rn_parse_configs, parse_sys_args as rn_parse_sys_args
    from runner.inference import InferenceRunner as RNAProRunner, infer_predict as rn_infer_predict

    configs_base_rn["use_deepspeed_evo_attention"] = (
        os.environ.get("USE_DEEPSPEED_EVO_ATTENTION", False) == "true")

    rn_configs = {**configs_base_rn, **{"data": data_configs_rn}, **inference_configs_rn}
    rn_configs = rn_parse_configs(
        configs=rn_configs,
        arg_str=rn_parse_sys_args(),
        fill_required_with_null=True,
    )

    sys.argv = saved_argv

    # Build model + load checkpoint ONCE
    t0 = time.time()
    runner = RNAProRunner(rn_configs)
    print(f"  [RNAPro] Model loaded in {time.time()-t0:.0f}s")

    return runner, rn_configs, rn_infer_predict, rnapro_dir, saved_cwd


def run_rnapro_target_inprocess(runner, configs, infer_predict_fn, rnapro_dir,
                                 tid, seq, rnapro_msa_dir=None):
    """Run RNAPro inference for a single target (in-memory model reuse).

    Handles chunking for long sequences. Returns list of (L, 3) coord arrays.
    """
    import torch
    from rnapro.utils.inference import process_sequence, extract_c1_coordinates

    full_len = len(seq)

    # Determine MSA dir to use
    msa_dir = rnapro_msa_dir or COMP_MSA_DIR

    # Split into chunks if needed
    chunks = split_into_chunks(full_len, RNAPRO_MAX_LEN, RNAPRO_CHUNK_OVERLAP)
    is_chunked = len(chunks) > 1

    if is_chunked:
        print(f"    [RNAPro] CHUNKING: {len(chunks)} chunks for {full_len} nt: {chunks}")
        # Need merged MSA dir for sliced MSA files
        msa_merged = os.path.join(OUTPUT_PATH, 'rnapro_msa_merged')
        os.makedirs(msa_merged, exist_ok=True)
        # Symlink original MSA files
        if os.path.isdir(COMP_MSA_DIR):
            for fn in os.listdir(COMP_MSA_DIR):
                src = os.path.join(COMP_MSA_DIR, fn)
                dst = os.path.join(msa_merged, fn)
                if not os.path.exists(dst):
                    try:
                        os.symlink(src, dst)
                    except OSError:
                        pass
        msa_dir = msa_merged

    all_chunk_preds = []  # list of (range, [sample_coords])

    items = []  # (chunk_tid, chunk_seq, chunk_start, chunk_end)
    if is_chunked:
        full_msa = os.path.join(COMP_MSA_DIR, f"{tid}.MSA.fasta")
        for ci, (cs, ce) in enumerate(chunks):
            chunk_tid = f"{tid}_rnc{ci}"
            chunk_seq = seq[cs:ce]
            items.append((chunk_tid, chunk_seq, cs, ce))
            # Slice MSA
            sliced_path = os.path.join(msa_dir, f"{chunk_tid}.MSA.fasta")
            result = slice_msa_fasta(full_msa, cs, ce, sliced_path)
            if result:
                print(f"      [RNAPro] Sliced MSA for {chunk_tid}: cols [{cs}:{ce}]")
    else:
        items.append((tid, seq, 0, full_len))

    for chunk_tid, chunk_seq, cs, ce in items:
        chunk_len = len(chunk_seq)

        # Create dummy template for this target
        template_pt = os.path.join(rnapro_dir, 'test_templates.pt')
        template_dict = {chunk_tid: {"xyz": np.zeros((chunk_len, 5, 3), dtype=np.float32)}}
        torch.save(template_dict, template_pt)

        # Setup input files
        temp_dir = os.path.join(configs.dump_dir, 'input')
        os.makedirs(temp_dir, exist_ok=True)
        process_sequence(sequence=chunk_seq, target_id=chunk_tid, temp_dir=temp_dir)

        # Update configs for this target
        configs.input_json_path = os.path.join(temp_dir, f"{chunk_tid}_input.json")
        configs.rna_msa_dir = msa_dir
        configs.template_data = template_pt
        configs.sequences_csv = os.path.join(OUTPUT_PATH, '_rnapro_dummy.csv')

        # Write temp CSV for this target
        pd.DataFrame([{"target_id": chunk_tid, "sequence": chunk_seq}]).to_csv(
            configs.sequences_csv, index=False)

        t0 = time.time()
        chunk_coords = []

        # Single pass: N_sample diffusion samples give diversity (no template_idx loop)
        configs.template_idx = 0
        try:
            infer_predict_fn(runner, configs)
        except Exception as e:
            print(f"      [RNAPro] {chunk_tid} ERROR: {e}")

        # Extract coordinates from CIF output
        cif_pattern = os.path.join(
            configs.dump_dir, chunk_tid, "seed_42", "predictions",
            f"{chunk_tid}_sample_*.cif")
        cif_files = sorted(glob.glob(cif_pattern))
        for cif_path in cif_files:
            try:
                coord = extract_c1_coordinates(cif_path)
                if coord is None:
                    coord = np.zeros((chunk_len, 3), dtype=np.float32)
                elif coord.shape[0] < chunk_len:
                    pad = np.zeros((chunk_len - coord.shape[0], 3), dtype=np.float32)
                    coord = np.concatenate([coord, pad], axis=0)
                elif coord.shape[0] > chunk_len:
                    coord = coord[:chunk_len]
                chunk_coords.append(coord)
            except Exception as e:
                print(f"      [RNAPro] CIF parse error: {e}")

        torch.cuda.empty_cache()

        elapsed = time.time() - t0
        print(f"    [RNAPro] {chunk_tid}: {len(chunk_coords)} samples in {elapsed:.0f}s")
        all_chunk_preds.append(((cs, ce), chunk_coords))

    # Stitch chunks if needed
    if not is_chunked:
        _, coords_list = all_chunk_preds[0]
        # Trim/pad to full_len
        result = []
        for c in coords_list:
            if len(c) < full_len:
                c = np.concatenate([c, np.zeros((full_len - len(c), 3), dtype=np.float32)])
            elif len(c) > full_len:
                c = c[:full_len]
            result.append(c)
        return result

    # Stitch each sample independently
    n_samples = min(len(cc) for _, cc in all_chunk_preds) if all_chunk_preds else 0
    if n_samples == 0:
        return []

    all_stitched = []
    for s_idx in range(n_samples):
        sample_chunks = [cc[min(s_idx, len(cc) - 1)] for _, cc in all_chunk_preds]
        ranges = [r for r, _ in all_chunk_preds]
        stitched = stitch_chunk_coords(sample_chunks, ranges, full_len)
        all_stitched.append(stitched)

    print(f"    [RNAPro] Stitched: {len(all_stitched)} samples x {full_len} residues")
    return all_stitched


# ============================================================
# 8. MAIN HYBRID PIPELINE
# ============================================================
def generate_predictions(sequences_df):
    """
    Slot 1: HGB ranker top-1 (always)
    Slot 2: RNAPro (try) or HGB top-2 (fallback)
    Slots 3-5: Protenix v1 co-fold (try) or HGB top-3..5 (fallback)

    Targets sorted by complexity: short RNA-only first → co-fold last.
    Time budgets enforced per-target and per-phase.
    Adaptive time: co-fold with protein MSA gets more time (high ROI).
    """
    import torch
    t_global_start = time.time()

    # --- Classify each target for smart time allocation ---
    def _classify_target(row):
        """Classify target → (type, extra_tokens, has_protein_msa, ptx_time_limit).

        Types:
          'cofold_msa'  — has protein chains + pre-computed protein MSA → highest PTX ROI
          'cofold_nomsa'— has protein chains but no MSA → moderate ROI
          'rna_long'    — RNA-only but needs chunking → slow
          'rna_short'   — RNA-only, fits in one pass → fast
        """
        rna_len = len(row['sequence'])
        all_seq = row.get('all_sequences', '')

        # Parse non-RNA chains
        protein_seqs = []
        extra_tokens = 0
        if pd.notna(all_seq):
            for entry in all_seq.split('>'):
                if not entry.strip():
                    continue
                lines = entry.strip().split('\n')
                seq = ''.join(lines[1:]).strip()
                if not seq:
                    continue
                unique = set(seq.upper())
                if unique <= {'A', 'U', 'G', 'C'}:
                    continue  # RNA
                extra_tokens += len(seq)
                if not (unique <= {'A', 'T', 'G', 'C'}):
                    protein_seqs.append(seq)  # protein (not DNA)

        total_tokens = rna_len + extra_tokens
        has_cofold = extra_tokens > 0

        # Check protein MSA availability
        has_prot_msa = False
        if protein_seqs:
            for ps in protein_seqs:
                if _find_protein_a3m(ps, row['target_id']):
                    has_prot_msa = True
                    break

        # Classify
        if has_cofold and has_prot_msa:
            ttype = 'cofold_msa'
        elif has_cofold:
            ttype = 'cofold_nomsa'
        elif rna_len > PTX_MAX_RNA_LEN:
            ttype = 'rna_long'
        else:
            ttype = 'rna_short'

        # Adaptive PTX time limit based on type + complexity
        if ADAPTIVE_TIME:
            if ttype == 'cofold_msa':
                # High ROI: co-fold + MSA gives biggest TM gain (e.g. 9J09: +0.386)
                # Scale by token count: bigger complex = more time needed
                ptx_tlimit = min(240000, max(1200, total_tokens * 2))
            elif ttype == 'cofold_nomsa':
                # Moderate ROI: co-fold without MSA still helps but less
                ptx_tlimit = min(120000, max(600, total_tokens))
            elif ttype == 'rna_long':
                # Chunking is slow but needed, scale by chunks
                n_chunks = max(1, (rna_len - 1) // PTX_MAX_RNA_LEN + 1)
                ptx_tlimit = min(180000, 600 * n_chunks)
            else:
                # RNA-only short: fast, minimal time
                ptx_tlimit = min(60000, max(180, rna_len * 2))
        else:
            ptx_tlimit = PTX_TIME_LIMIT

        return {
            'type': ttype,
            'rna_len': rna_len,
            'extra_tokens': extra_tokens,
            'total_tokens': total_tokens,
            'has_prot_msa': has_prot_msa,
            'n_proteins': len(protein_seqs),
            'ptx_tlimit': ptx_tlimit,
        }

    # Classify all targets
    target_info = {}
    for _, row in sequences_df.iterrows():
        target_info[row['target_id']] = _classify_target(row)

    # --- Sort targets by expected speed (fast first, slow last) ---
    # Priority: rna_short (fast) → rna_long → cofold_msa (high ROI, do early) → cofold_nomsa
    _type_order = {'rna_short': 0, 'rna_long': 1, 'cofold_msa': 2, 'cofold_nomsa': 3}

    def _sort_key(row):
        info = target_info[row['target_id']]
        return (_type_order.get(info['type'], 9), info['total_tokens'])

    sort_keys = [_sort_key(row) for _, row in sequences_df.iterrows()]
    sort_order = sorted(range(len(sort_keys)), key=lambda i: sort_keys[i])
    sequences_df = sequences_df.iloc[sort_order].reset_index(drop=True)

    print(f"\n{'='*70}")
    print(f"HYBRID PREDICTION: {len(sequences_df)} targets (sorted by type+complexity)")
    print(f"PyTorch: {torch.__version__}, CUDA: {torch.cuda.is_available()}")
    if torch.cuda.is_available():
        print(f"GPU: {torch.cuda.get_device_name(0)}")
    print(f"{'='*70}")

    # Preview sort order with classification
    total_ptx_budget_est = 0
    for i, (_, row) in enumerate(sequences_df.iterrows()):
        info = target_info[row['target_id']]
        msa_tag = "+MSA" if info['has_prot_msa'] else ""
        print(f"  [{i+1}] {row['target_id']}: {info['rna_len']}nt, "
              f"tokens={info['total_tokens']}, {info['type']}{msa_tag}, "
              f"ptx_tlimit={info['ptx_tlimit']}s")
        total_ptx_budget_est += info['ptx_tlimit']
    print(f"  Total estimated PTX budget: {total_ptx_budget_est}s "
          f"(phase_a_budget={PHASE_A_BUDGET}s)")

    total = len(sequences_df)

    def _time_left():
        return GLOBAL_DEADLINE - (time.time() - t_global_start)

    # --- Phase A: Protenix v1 co-fold ---
    ptx_preds = {}
    print(f"\n--- Phase A: Protenix v1 co-fold (budget {PHASE_A_BUDGET}s) ---")
    phase_a_start = time.time()
    ptx_runner = None
    ptx_saved_cwd = None
    try:
        ptx_runner, ptx_configs, ptx_infer_fn, ptx_saved_cwd = setup_protenix_inprocess()

        for i, (idx, row) in enumerate(sequences_df.iterrows()):
            tid = row['target_id']
            seq = row['sequence']
            all_seq = row.get('all_sequences', '')
            info = target_info[tid]

            # Time checks
            phase_elapsed = time.time() - phase_a_start
            remaining = _time_left()
            if phase_elapsed > PHASE_A_BUDGET:
                print(f"  Phase A budget exhausted ({phase_elapsed:.0f}s > {PHASE_A_BUDGET}s), "
                      f"skipping remaining {total - i} targets")
                break
            # Reserve time for RNAPro + RF: min 20min for phase B+C
            min_reserve = 1200
            if remaining < min_reserve:
                print(f"  Global deadline approaching ({remaining:.0f}s left < {min_reserve}s reserve), stopping PTX")
                break

            print(f"  [{i+1}/{total}] {tid} ({len(seq)} nt, {info['type']}) "
                  f"[phaseA: {phase_elapsed:.0f}s/{PHASE_A_BUDGET}s, "
                  f"global: {remaining:.0f}s left, ptx_tlimit={info['ptx_tlimit']}s]")
            t_target = time.time()
            target_ligands = parse_ligands_for_ptx(row.get('ligand_ids', ''))
            try:
                coords_list = run_protenix_target_inprocess(
                    ptx_runner, ptx_configs, ptx_infer_fn, tid, seq, all_seq,
                    ligands=target_ligands, stoichiometry=row.get('stoichiometry', ''))
                if coords_list:
                    ptx_preds[tid] = coords_list
                target_time = time.time() - t_target
                print(f"    [PTX] {tid}: {'OK' if coords_list else 'EMPTY'} in {target_time:.0f}s")
                if target_time > info['ptx_tlimit']:
                    print(f"    [PTX] WARNING: {tid} took {target_time:.0f}s > adaptive limit {info['ptx_tlimit']}s")
            except Exception as e:
                print(f"    [PTX] EXCEPTION: {e}")
                traceback.print_exc()
    except Exception as e:
        print(f"  [PTX] SETUP FAILED: {e}")
        traceback.print_exc()
    finally:
        if ptx_runner is not None:
            del ptx_runner
        if ptx_saved_cwd is not None:
            os.chdir(ptx_saved_cwd)
        ptx_mods = [m for m in sys.modules if m.startswith(('runner.', 'configs.'))]
        for m in ptx_mods:
            del sys.modules[m]
        if 'runner' in sys.modules:
            del sys.modules['runner']
        if 'configs' in sys.modules:
            del sys.modules['configs']
        if PTX_SOURCE in sys.path:
            sys.path.remove(PTX_SOURCE)
        if torch.cuda.is_available():
            torch.cuda.empty_cache()
        gc.collect()

    phase_a_time = time.time() - phase_a_start
    print(f"\n  Protenix: {len(ptx_preds)}/{total} targets OK in {phase_a_time:.0f}s")

    # --- Phase B: RNAPro in-memory ---
    rnapro_preds = {}
    rnapro_dir = None
    rnapro_saved_cwd = None
    phase_b_start = time.time()
    try:
        print(f"\n--- Phase B: RNAPro in-memory (budget {PHASE_B_BUDGET}s) ---")
        rn_runner, rn_configs, rn_infer_fn, rnapro_dir, rnapro_saved_cwd = setup_rnapro_inprocess()

        # Sort by RNA length (not total tokens) — RNAPro sees only RNA.
        # Short RNA first → more targets covered before budget runs out.
        rnapro_order = sorted(
            range(len(sequences_df)),
            key=lambda i: len(sequences_df.iloc[i]['sequence'])
        )

        for i, seq_idx in enumerate(rnapro_order):
            row = sequences_df.iloc[seq_idx]
            tid = row['target_id']
            seq = row['sequence']

            # Time checks
            phase_elapsed = time.time() - phase_b_start
            remaining = _time_left()
            if phase_elapsed > PHASE_B_BUDGET:
                print(f"  Phase B budget exhausted ({phase_elapsed:.0f}s), "
                      f"skipping remaining {total - i} targets")
                break
            if remaining < 600:  # need 10min for RF assembly
                print(f"  Global deadline approaching ({remaining:.0f}s left), stopping RNAPro")
                break

            # Skip very long sequences
            if len(seq) > RNAPRO_SKIP_LEN:
                print(f"  [{i+1}/{total}] {tid} ({len(seq)} nt) — SKIP (>{RNAPRO_SKIP_LEN}nt)")
                continue

            print(f"  [{i+1}/{total}] {tid} ({len(seq)} nt) "
                  f"[phaseB: {phase_elapsed:.0f}s/{PHASE_B_BUDGET}s]")
            t_target = time.time()
            try:
                coords_list = run_rnapro_target_inprocess(
                    rn_runner, rn_configs, rn_infer_fn, rnapro_dir, tid, seq)
                if coords_list:
                    rnapro_preds[tid] = coords_list
                    print(f"    [RNAPro] OK: {len(coords_list)} samples in {time.time()-t_target:.0f}s")
                target_time = time.time() - t_target
                if target_time > RNAPRO_TIME_LIMIT:
                    print(f"    [RNAPro] WARNING: {tid} took {target_time:.0f}s > limit {RNAPRO_TIME_LIMIT}s")
            except Exception as e:
                print(f"    [RNAPro] EXCEPTION: {e}")
                traceback.print_exc()
    except Exception as e:
        print(f"  [RNAPro] SETUP FAILED: {e}")
        traceback.print_exc()
    finally:
        if 'rn_runner' in dir():
            del rn_runner
        if rnapro_saved_cwd is not None:
            os.chdir(rnapro_saved_cwd)
        rn_mods = [m for m in sys.modules if m.startswith(('runner.', 'configs.', 'rnapro.'))]
        for m in rn_mods:
            del sys.modules[m]
        for m in ['runner', 'configs', 'rnapro']:
            if m in sys.modules:
                del sys.modules[m]
        if rnapro_dir and rnapro_dir in sys.path:
            sys.path.remove(rnapro_dir)
        if torch.cuda.is_available():
            torch.cuda.empty_cache()
        gc.collect()

    phase_b_time = time.time() - phase_b_start
    print(f"\n  RNAPro: {len(rnapro_preds)}/{total} targets OK in {phase_b_time:.0f}s")

    # Free GPU
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
    gc.collect()

    # Phase C
    print(f"\n--- Phase C: HGB + Smart Routing Assembly ---")
    all_preds = {}
    route_stats = {"ptx_heavy": 0, "rnapro_heavy": 0, "balanced": 0, "hgb_heavy": 0}

    for idx, row in sequences_df.iterrows():
        tid = row['target_id']
        seq = row['sequence']
        rna_len = len(seq)
        segments = tbm.get_chain_segments(row)

        try:
            tbm_preds = predict_diversity_for_target(tid, seq, segments, row, n_predictions=3)
        except Exception as e:
            print(f"  [{tid}] HGB ERROR: {e}")
            tbm_preds = [(generate_aform_helix(seq), 0.0)] * 5

        rn = rnapro_preds.get(tid, [])
        pt = ptx_preds.get(tid, [])
        info = target_info.get(tid, {})
        total_tok = info.get('total_tokens', rna_len)
        has_protein = info.get('n_proteins', 0) > 0
        has_dna = any(c['type'] == 'dna' for c in parse_chains(row.get('all_sequences', ''))) if pd.notna(row.get('all_sequences')) else False
        desc = str(row.get('description', '')).lower()

        # --- Determine routing: how many PTX vs RNAPro in slots 2-5 ---
        # Default: balanced 2+2
        n_ptx_slots = 1
        n_rnapro_slots = 1
        route_label = "balanced"

        if total_tok > 3000:
            # Huge targets: mostly HGB, 1 PTX + 1 RNAPro for diversity
            n_ptx_slots = 1; n_rnapro_slots = 1
            route_label = "hgb_heavy"
        elif has_protein and has_dna:
            # CRISPR / RNA+Prot+DNA: RNAPro dominates massively (+0.256)
            n_ptx_slots = 1; n_rnapro_slots = 1
            route_label = "rnapro_heavy"
        elif rna_len < 50:
            # Very short: PTX better (+0.062)
            n_ptx_slots = 1; n_rnapro_slots = 1
            route_label = "ptx_heavy"
        elif rna_len >= 1000:
            # Very long: PTX better on ribosomes (+0.143)
            n_ptx_slots = 1; n_rnapro_slots = 1
            route_label = "ptx_heavy"
        elif has_protein and not has_dna:
            # RNA+Protein (no DNA): roughly tied, slight PTX edge
            n_ptx_slots = 1; n_rnapro_slots = 1
            route_label = "balanced"
        elif 50 <= rna_len < 1000:
            # Medium RNA: RNAPro dominates
            n_ptx_slots = 1; n_rnapro_slots = 1
            route_label = "rnapro_heavy"

        route_stats[route_label] = route_stats.get(route_label, 0) + 1

        # --- Fill 5 slots: S1-S3 = TBM diversity, S4 = PTX, S5 = RNaPro ---
        slots = []

        # Slots 1-3: Diversity TBM (Bio, Model, Combined)
        for i in range(min(3, len(tbm_preds))):
            slots.append(tbm_preds[i][0])

        # Slot 4: PTX preferred, else RNaPro, else TBM S1 duplicate
        s4_src = "FB"
        if pt:
            slots.append(pt[0]); s4_src = "PTX"
        elif rn:
            slots.append(rn[0]); s4_src = "RN"
        else:
            slots.append(tbm_preds[0][0]); s4_src = "TBM_dup"  # better than aform

        # Slot 5: RNaPro preferred (avoid duplicate with S4)
        s5_src = "FB"
        if rn and s4_src != "RN":
            slots.append(rn[0]); s5_src = "RN"
        elif rn and len(rn) > 1:
            slots.append(rn[1]); s5_src = "RN_2"
        elif pt and s4_src != "PTX":
            slots.append(pt[0]); s5_src = "PTX"
        elif pt and len(pt) > 1:
            slots.append(pt[1]); s5_src = "PTX_2"
        else:
            # Fallback: duplicate best TBM slot (S1)
            slots.append(tbm_preds[0][0]); s5_src = "TBM_dup"

        all_preds[tid] = slots[:5]

        print(f"  [{tid}] {route_label} (tok={total_tok}) "
              f"S1-3=TBM_div S4={s4_src} S5={s5_src}")

    print(f"\nRouting stats: {route_stats}")
    return all_preds


# ============================================================
# 9. FORMAT SUBMISSION
# ============================================================
def format_submission(predictions, sequences_df):
    rows = []
    for _, row in sequences_df.iterrows():
        tid = row['target_id']
        seq = row['sequence']
        preds = predictions.get(tid)
        if preds is None:
            preds = [generate_aform_helix(seq)] * 5

        for j in range(len(seq)):
            r = {'ID': f"{tid}_{j+1}", 'resname': seq[j], 'resid': j + 1}
            for i in range(5):
                x, y, z = float(preds[i][j][0]), float(preds[i][j][1]), float(preds[i][j][2])
                r[f'x_{i+1}'] = min(max(x, -999.999), 9999.999)
                r[f'y_{i+1}'] = min(max(y, -999.999), 9999.999)
                r[f'z_{i+1}'] = min(max(z, -999.999), 9999.999)
            rows.append(r)

    df = pd.DataFrame(rows)
    cols = ['ID', 'resname', 'resid']
    for i in range(1, 6):
        for c in ['x', 'y', 'z']:
            cols.append(f'{c}_{i}')
    return df[cols]


# ============================================================
# 10. RUN
# ============================================================
t_start = time.time()
predictions = generate_predictions(test_seqs)
submission = format_submission(predictions, test_seqs)
sub_path = os.path.join(OUTPUT_PATH, 'submission.csv')
submission.to_csv(sub_path, index=False)
elapsed = time.time() - t_start

print(f"\n{'='*70}")
print(f"DONE in {elapsed:.0f}s ({elapsed/60:.1f}m) — profile {CONFIG_PROFILE}")
print(f"Submission: {sub_path} ({len(submission)} rows)")
print(submission.head())
print(f"{'='*70}")