#!/usr/bin/env python3
"""exp029: exp011 base + N_SAMPLE=10 (best-of-10 submission).

Based on exp011 (0.434 LB) with key change:
- N_SAMPLE=10: generate 10 predictions, submit all 10 (x_1..x_10)
- Kaggle scoring evaluates best-of-N, so more predictions = higher score
- Inspired by asimandia notebook (#1 public notebook)
"""

import subprocess, os, sys, glob
if os.path.exists("/kaggle/input"):
    print("Mounted inputs:", os.listdir("/kaggle/input"))

    def _pip_install_wheels(search_dirs, label=""):
        for d in search_dirs:
            if os.path.isdir(d):
                wheels = sorted(glob.glob(os.path.join(d, "*.whl")))
                if wheels:
                    print(f"  Installing {label} from {d}: {[os.path.basename(w) for w in wheels]}")
                    r = subprocess.run(
                        [sys.executable, "-m", "pip", "install", "--no-deps"] + wheels,
                        capture_output=True, text=True,
                    )
                    print(f"  exit={r.returncode}: {r.stdout.strip()}")
                    if r.returncode != 0:
                        print(f"  stderr: {r.stderr.strip()}")
                    return True
        return False

    _pip_install_wheels([
        "/kaggle/input/biopython-cp312",
        "/kaggle/input/datasets/yuto0712/biopython-cp312",
        "/kaggle/input/datasets/ogurtsov/biopython",
    ], "BioPython")
    _pip_install_wheels([
        "/kaggle/input/protenix-deps-cp312",
        "/kaggle/input/datasets/yuto0712/protenix-deps-cp312",
    ], "Protenix deps")

import gc
import json
import time
import warnings
from pathlib import Path

import numpy as np
import pandas as pd

warnings.filterwarnings("ignore", category=DeprecationWarning, module="Bio")
from Bio.Align import PairwiseAligner

# ─────────────── Configuration ──────────────────────────────────────────────
os.environ["LAYERNORM_TYPE"] = "torch"
os.environ.setdefault("RNA_MSA_DEPTH_LIMIT", "512")

IS_KAGGLE = os.path.exists("/kaggle/input")
IS_COMPETITION_RERUN = bool(os.environ.get("KAGGLE_IS_COMPETITION_RERUN", ""))

if IS_KAGGLE:
    _candidates = [
        "/kaggle/input/stanford-rna-3d-folding-2",
        "/kaggle/input/competitions/stanford-rna-3d-folding-2",
    ]
    DATA_BASE = next((p for p in _candidates if os.path.isdir(p)), _candidates[0])
    OUTPUT_PATH = "/kaggle/working/submission.csv"
    PDB_RNA_DIR = os.path.join(DATA_BASE, "PDB_RNA")
    _code_candidates = [
        "/kaggle/input/protenix-v1-adjusted/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1",
        "/kaggle/input/datasets/qiweiyin/protenix-v1-adjusted/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1",
    ]
    CODE_DIR = next((p for p in _code_candidates if os.path.isdir(p)), _code_candidates[0])
else:
    DATA_BASE = os.environ.get("DATA_DIR", "data/raw")
    OUTPUT_PATH = os.environ.get("OUTPUT_PATH", "experiments/exp011/output/submission.csv")
    PDB_RNA_DIR = os.path.join(DATA_BASE, "PDB_RNA")
    CODE_DIR = os.environ.get("PROTENIX_CODE_DIR", "")

LOCAL_N_SAMPLES = 2

# Protenix config (matching Artem 0.413)
MODEL_NAME = "protenix_base_20250630_v1.0.0"
N_SAMPLE = 10
SEED = 42
MAX_SEQ_LEN = 512
CHUNK_OVERLAP = 128
USE_MSA = "false"
USE_TEMPLATE = "false"
USE_RNA_MSA = "true"

# TBM thresholds (matching Artem 0.413)
MIN_SIMILARITY = 0.0
MIN_PERCENT_IDENTITY = 50.0
LENGTH_DIFF_CUTOFF = 0.3
TOP_N_TEMPLATES = 30

# RNA structural constraint parameters
BOND_DISTANCE = 5.95
BOND_I2_DISTANCE = 10.2
SELF_AVOIDANCE_MIN = 3.2

# CIF parsing
RESNAME_MAP = {
    "A": "A", "C": "C", "G": "G", "U": "U",
    "ADE": "A", "CYT": "C", "GUA": "G", "URA": "U", "URI": "U",
    "PSU": "U", "H2U": "U", "5MU": "U", "5MC": "C",
    "OMC": "C", "OMG": "G", "7MG": "G", "2MG": "G",
    "1MA": "A", "M2G": "G", "I": "G", "YYG": "G",
}


# ─────────────── Utilities ──────────────────────────────────────────────────
def seed_everything(seed: int) -> None:
    import torch
    os.environ["CUBLAS_WORKSPACE_CONFIG"] = ":4096:8"
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    np.random.seed(seed)
    torch.backends.cudnn.benchmark = False
    torch.backends.cudnn.deterministic = True
    torch.use_deterministic_algorithms(True)


def parse_fasta(fasta_content: str) -> dict:
    out, cur, parts = {}, None, []
    for line in str(fasta_content).splitlines():
        line = line.strip()
        if not line:
            continue
        if line.startswith(">"):
            if cur is not None:
                out[cur] = "".join(parts)
            cur = line[1:].split()[0]
            parts = []
        else:
            parts.append(line.replace(" ", ""))
    if cur is not None:
        out[cur] = "".join(parts)
    return out


def parse_stoichiometry(stoich: str) -> list:
    if pd.isna(stoich) or str(stoich).strip() == "":
        return []
    return [(ch.strip(), int(cnt)) for part in str(stoich).split(";")
            for ch, cnt in [part.split(":")]]


def get_chain_segments(row) -> list:
    seq = row["sequence"]
    stoich = row.get("stoichiometry", "")
    all_sq = row.get("all_sequences", "")
    if pd.isna(stoich) or pd.isna(all_sq) or str(stoich).strip() == "" or str(all_sq).strip() == "":
        return [(0, len(seq))]
    try:
        chain_dict = parse_fasta(all_sq)
        order = parse_stoichiometry(stoich)
        segs, pos = [], 0
        for ch, cnt in order:
            base = chain_dict.get(ch)
            if base is None:
                return [(0, len(seq))]
            for _ in range(cnt):
                segs.append((pos, pos + len(base)))
                pos += len(base)
        return segs if pos == len(seq) else [(0, len(seq))]
    except Exception:
        return [(0, len(seq))]


def build_segments_map(df: pd.DataFrame) -> dict:
    seg_map = {}
    for _, r in df.iterrows():
        seg_map[r["target_id"]] = get_chain_segments(r)
    return seg_map


def process_labels(labels_df: pd.DataFrame) -> dict:
    coords = {}
    prefixes = labels_df["ID"].str.rsplit("_", n=1).str[0]
    for prefix, grp in labels_df.groupby(prefixes):
        coords[prefix] = grp.sort_values("resid")[["x_1", "y_1", "z_1"]].values.astype(np.float32)
    return coords


# ─────────────── Chunk + Stitch (from OpenRNAFold 0.434) ───────────────────
def split_into_chunks(seq_len: int, max_len: int, overlap: int) -> list:
    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
    return chunks


def kabsch_align(P: np.ndarray, Q: np.ndarray):
    centroid_P = P.mean(axis=0)
    centroid_Q = Q.mean(axis=0)
    Pc = P - centroid_P
    Qc = Q - centroid_Q
    H = Pc.T @ Qc
    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 = centroid_Q - R @ centroid_P
    return R, t


def stitch_chunk_coords(chunk_coords_list, chunk_ranges, seq_len):
    if len(chunk_coords_list) == 1:
        coords = chunk_coords_list[0]
        if coords.shape[0] >= seq_len:
            return coords[:seq_len]
        out = np.zeros((seq_len, 3), dtype=coords.dtype)
        out[:coords.shape[0]] = coords
        return out

    aligned = [chunk_coords_list[0].copy()]
    for i in range(1, len(chunk_coords_list)):
        prev_start, prev_end = chunk_ranges[i - 1]
        cur_start, cur_end = chunk_ranges[i]
        ov_start = cur_start
        ov_end = min(prev_end, cur_end)
        ov_len = ov_end - ov_start

        if ov_len < 3:
            aligned.append(chunk_coords_list[i].copy())
            continue

        prev_ov = aligned[i - 1][ov_start - prev_start: ov_end - prev_start]
        cur_ov = chunk_coords_list[i][ov_start - cur_start: ov_end - cur_start]
        valid = ~(np.isnan(prev_ov).any(axis=1) | np.isnan(cur_ov).any(axis=1))
        if valid.sum() < 3:
            aligned.append(chunk_coords_list[i].copy())
            continue

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

    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 = coords.shape[0]
        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_start = s
            ov_end = min(chunk_ranges[i - 1][1], e)
            ramp_len = ov_end - ov_start
            if ramp_len > 0:
                w[:ramp_len] = np.linspace(0.0, 1.0, ramp_len)

        if i < len(chunk_ranges) - 1:
            next_s = chunk_ranges[i + 1][0]
            ramp_start = next_s - s
            ramp_len = actual_end - next_s
            if ramp_len > 0 and ramp_start < used_len:
                w[ramp_start:used_len] = np.linspace(1.0, 0.0, ramp_len)

        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)


# ─────────────── PDB_RNA Template Loading ───────────────────────────────────
def parse_cif_c1_coords_by_chain(cif_path: str) -> dict:
    from collections import defaultdict
    chains = defaultdict(dict)
    in_atom_site = False
    columns = []

    with open(cif_path, "r", errors="replace") as f:
        for line in f:
            line = line.rstrip("\n")
            if line.startswith("_atom_site."):
                if not in_atom_site:
                    in_atom_site = True
                    columns = []
                columns.append(line.split(".")[1].strip())
                continue
            if in_atom_site and line.startswith("ATOM") or (in_atom_site and line.startswith("HETATM")):
                parts = line.split()
                if len(parts) < len(columns):
                    continue
                col_dict = {columns[i]: parts[i] for i in range(len(columns))}

                model = col_dict.get("pdbx_PDB_model_num", "1")
                if model != "1":
                    continue
                atom_name = col_dict.get("label_atom_id", col_dict.get("auth_atom_id", ""))
                if atom_name.startswith('"') and atom_name.endswith('"'):
                    atom_name = atom_name[1:-1]
                if atom_name != "C1'":
                    continue
                comp_id = col_dict.get("label_comp_id", col_dict.get("auth_comp_id", ""))
                canonical = RESNAME_MAP.get(comp_id)
                if canonical is None:
                    continue
                chain_id = col_dict.get("auth_asym_id", col_dict.get("label_asym_id", ""))
                seq_id = col_dict.get("label_seq_id", col_dict.get("auth_seq_id", "0"))
                try:
                    x = float(col_dict.get("Cartn_x", 0))
                    y = float(col_dict.get("Cartn_y", 0))
                    z = float(col_dict.get("Cartn_z", 0))
                    seq_num = int(seq_id) if seq_id != "." else 0
                except (ValueError, TypeError):
                    continue
                if seq_num not in chains[chain_id]:
                    chains[chain_id][seq_num] = (canonical, [x, y, z])
            elif in_atom_site and (line.startswith("#") or line == ""):
                in_atom_site = False

    result = {}
    for chain_id, residues in chains.items():
        if len(residues) < 10:
            continue
        sorted_ids = sorted(residues.keys())
        seq = "".join(residues[sid][0] for sid in sorted_ids)
        coords = np.array([residues[sid][1] for sid in sorted_ids], dtype=np.float32)
        result[chain_id] = {"sequence": seq, "coords": coords}
    return result


def load_pdb_rna_templates(pdb_rna_dir: str, existing_pdb_ids: set) -> tuple:
    t0 = time.time()
    cif_dir = Path(pdb_rna_dir)
    if not cif_dir.exists():
        print("  PDB_RNA directory not found, skipping")
        return pd.DataFrame(columns=["target_id", "sequence"]), {}

    rows, coords_dict, n_parsed = [], {}, 0
    for cif_path in sorted(cif_dir.glob("*.cif")):
        pdb_id = cif_path.stem.upper()
        if pdb_id in existing_pdb_ids:
            continue
        try:
            chains = parse_cif_c1_coords_by_chain(str(cif_path))
            for chain_id, data in chains.items():
                target_id = f"{pdb_id}_{chain_id}"
                rows.append({"target_id": target_id, "sequence": data["sequence"]})
                coords_dict[target_id] = data["coords"]
            n_parsed += 1
        except Exception:
            pass

    seqs_df = pd.DataFrame(rows) if rows else pd.DataFrame(columns=["target_id", "sequence"])
    print(f"  PDB_RNA: {n_parsed} CIF -> {len(rows)} chains in {time.time()-t0:.1f}s")
    return seqs_df, coords_dict


def load_new_pdb_templates(existing_ids: set) -> tuple:
    for d in ["/kaggle/input/rna3d-new-pdb-templates",
              "/kaggle/input/datasets/billbafare/rna3d-new-pdb-templates"]:
        if os.path.isdir(d):
            csv_path = os.path.join(d, "template_sequences.csv")
            npz_path = os.path.join(d, "template_coords.npz")
            if os.path.exists(csv_path) and os.path.exists(npz_path):
                df = pd.read_csv(csv_path)
                npz = np.load(npz_path, allow_pickle=True)
                rows, coords_dict = [], {}
                for _, row in df.iterrows():
                    cid = row["chain_id"]
                    if cid in existing_ids:
                        continue
                    rows.append({"target_id": cid, "sequence": row["sequence"]})
                    coords_dict[cid] = npz[cid].astype(np.float32)
                print(f"  New PDB templates: {len(rows)} chains")
                return pd.DataFrame(rows) if rows else pd.DataFrame(columns=["target_id", "sequence"]), coords_dict
    print("  New PDB templates: dataset not found, skipping")
    return pd.DataFrame(columns=["target_id", "sequence"]), {}


# ─────────────── TBM Core Functions ──────────────────────────────────────────
_aligner = PairwiseAligner()
_aligner.mode = "global"
_aligner.match_score = 2
_aligner.mismatch_score = -1.5
_aligner.open_gap_score = -8
_aligner.extend_gap_score = -0.4
_aligner.query_left_open_gap_score = -8
_aligner.query_left_extend_gap_score = -0.4
_aligner.query_right_open_gap_score = -8
_aligner.query_right_extend_gap_score = -0.4
_aligner.target_left_open_gap_score = -8
_aligner.target_left_extend_gap_score = -0.4
_aligner.target_right_open_gap_score = -8
_aligner.target_right_extend_gap_score = -0.4


def find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, top_n=30):
    results = []
    for _, row in train_seqs_df.iterrows():
        tid, tseq = row["target_id"], row["sequence"]
        if tid not in train_coords_dict:
            continue
        if abs(len(tseq) - len(query_seq)) / max(len(tseq), len(query_seq)) > LENGTH_DIFF_CUTOFF:
            continue
        aln = next(iter(_aligner.align(query_seq, tseq)))
        norm_s = aln.score / (2 * min(len(query_seq), len(tseq)))
        identical = sum(
            1 for (qs, qe), (ts, te) in zip(*aln.aligned)
            for qp, tp in zip(range(qs, qe), range(ts, te))
            if query_seq[qp] == tseq[tp]
        )
        pct_id = 100 * identical / len(query_seq)
        results.append((tid, tseq, norm_s, train_coords_dict[tid], pct_id))
    results.sort(key=lambda x: x[2], reverse=True)
    return results[:top_n]


def adapt_template_to_query(query_seq, template_seq, template_coords) -> np.ndarray:
    aln = next(iter(_aligner.align(query_seq, template_seq)))
    new_coords = np.full((len(query_seq), 3), np.nan)
    for (qs, qe), (ts, te) in zip(*aln.aligned):
        chunk = template_coords[ts:te]
        if len(chunk) == (qe - qs):
            new_coords[qs:qe] = chunk
    for i in range(len(new_coords)):
        if np.isnan(new_coords[i, 0]):
            pv = next((j for j in range(i - 1, -1, -1) if not np.isnan(new_coords[j, 0])), -1)
            nv = next((j for j in range(i + 1, len(new_coords)) if not np.isnan(new_coords[j, 0])), -1)
            if pv >= 0 and nv >= 0:
                w = (i - pv) / (nv - pv)
                new_coords[i] = (1 - w) * new_coords[pv] + w * new_coords[nv]
            elif pv >= 0:
                new_coords[i] = new_coords[pv] + [3, 0, 0]
            elif nv >= 0:
                new_coords[i] = new_coords[nv] + [3, 0, 0]
            else:
                new_coords[i] = [i * 3, 0, 0]
    return np.nan_to_num(new_coords)


def adaptive_rna_constraints(coords, target_id, segments_map, confidence=1.0, passes=2) -> np.ndarray:
    X = coords.copy()
    segments = segments_map.get(target_id, [(0, len(X))])
    strength = max(0.75 * (1.0 - min(confidence, 0.97)), 0.02)
    for _ in range(passes):
        for s, e in segments:
            C = X[s:e]
            L = e - s
            if L < 3:
                continue
            d = C[1:] - C[:-1]
            dist = np.linalg.norm(d, axis=1) + 1e-6
            adj = d * ((BOND_DISTANCE - dist) / dist)[:, None] * (0.22 * strength)
            C[:-1] -= adj
            C[1:] += adj
            d2 = C[2:] - C[:-2]
            d2n = np.linalg.norm(d2, axis=1) + 1e-6
            adj2 = d2 * ((BOND_I2_DISTANCE - d2n) / d2n)[:, None] * (0.10 * strength)
            C[:-2] -= adj2
            C[2:] += adj2
            C[1:-1] += (0.06 * strength) * (0.5 * (C[:-2] + C[2:]) - C[1:-1])
            if L >= 25:
                idx = np.linspace(0, L - 1, min(L, 160)).astype(int) if L > 220 else np.arange(L)
                P = C[idx]
                diff = P[:, None, :] - P[None, :, :]
                dm = np.linalg.norm(diff, axis=2) + 1e-6
                sep = np.abs(idx[:, None] - idx[None, :])
                mask = (sep > 2) & (dm < SELF_AVOIDANCE_MIN)
                if np.any(mask):
                    vec = (diff * ((SELF_AVOIDANCE_MIN - dm) / dm)[:, :, None] * mask[:, :, None]).sum(axis=1)
                    C[idx] += (0.015 * strength) * vec
            X[s:e] = C
    return X


# ─────────────── Diversity Transforms (confidence-scaled, from OpenRNAFold) ─
def _rotmat(axis, ang):
    a = np.asarray(axis, float)
    a /= np.linalg.norm(a) + 1e-12
    x, y, z = a
    c, s = np.cos(ang), np.sin(ang)
    CC = 1 - c
    return np.array([[c+x*x*CC, x*y*CC-z*s, x*z*CC+y*s],
                     [y*x*CC+z*s, c+y*y*CC, y*z*CC-x*s],
                     [z*x*CC-y*s, z*y*CC+x*s, c+z*z*CC]])


def apply_hinge(coords, seg, rng, deg=22, confidence=1.0):
    s, e = seg
    L = e - s
    if L < 30:
        return coords
    deg_scaled = deg * max(0.2, min(2.0, 1.5 - confidence))
    pivot = s + int(rng.integers(10, L - 10))
    R = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg_scaled, deg_scaled))))
    X = coords.copy()
    p0 = X[pivot].copy()
    X[pivot+1:e] = (X[pivot+1:e] - p0) @ R.T + p0
    return X


def jitter_chains(coords, segs, rng, deg=12, trans=1.5, confidence=1.0):
    X = coords.copy()
    gc_ = X.mean(0, keepdims=True)
    scale = max(0.2, min(2.0, 1.5 - confidence))
    deg_scaled = deg * scale
    trans_scaled = trans * scale
    for s, e in segs:
        R = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg_scaled, deg_scaled))))
        shift = rng.normal(size=3)
        shift = shift / (np.linalg.norm(shift) + 1e-12) * float(rng.uniform(0, trans_scaled))
        c = X[s:e].mean(0, keepdims=True)
        X[s:e] = (X[s:e] - c) @ R.T + c + shift
    X -= X.mean(0, keepdims=True) - gc_
    return X


def smooth_wiggle(coords, segs, rng, amp=0.8, confidence=1.0):
    X = coords.copy()
    scale = max(0.2, min(2.0, 1.5 - confidence))
    amp_scaled = amp * scale
    for s, e in segs:
        L = e - s
        if L < 20:
            continue
        ctrl = np.linspace(0, L - 1, 6)
        disp = rng.normal(0, amp_scaled, (6, 3))
        t = np.arange(L)
        X[s:e] += np.vstack([np.interp(t, ctrl, disp[:, k]) for k in range(3)]).T
    return X


def generate_rna_structure(sequence: str, seed=None) -> np.ndarray:
    if seed is not None:
        np.random.seed(seed)
    n = len(sequence)
    coords = np.zeros((n, 3))
    for i in range(n):
        ang = i * 0.6
        coords[i] = [10.0 * np.cos(ang), 10.0 * np.sin(ang), i * 2.5]
    return coords


# ─────────────── Protenix Helpers ───────────────────────────────────────────
def get_c1_mask(data, atom_array):
    import torch
    if atom_array is not None:
        try:
            if hasattr(atom_array, "centre_atom_mask"):
                m = atom_array.centre_atom_mask == 1
                if hasattr(atom_array, "is_rna"):
                    m = m & atom_array.is_rna
                return torch.from_numpy(m).bool()
            if hasattr(atom_array, "atom_name"):
                base = atom_array.atom_name == "C1'"
                if hasattr(atom_array, "is_rna"):
                    base = base & atom_array.is_rna
                return torch.from_numpy(base).bool()
        except Exception:
            pass
    f = data["input_feature_dict"]
    if "centre_atom_mask" in f:
        return (f["centre_atom_mask"] == 1).bool()
    if "center_atom_mask" in f:
        return (f["center_atom_mask"] == 1).bool()
    n_tokens = data.get("N_token", torch.tensor(0)).item()
    mask11 = (f["atom_to_tokatom_idx"] == 11).bool()
    mask12 = (f["atom_to_tokatom_idx"] == 12).bool()
    c11, c12 = mask11.sum().item(), mask12.sum().item()
    if abs(c11 - n_tokens) < abs(c12 - n_tokens):
        return mask11
    return mask12


def build_input_json(df, json_path):
    data = [
        {
            "name": row["target_id"],
            "covalent_bonds": [],
            "sequences": [{"rnaSequence": {"sequence": row["sequence"], "count": 1}}],
        }
        for _, row in df.iterrows()
    ]
    with open(json_path, "w") as f:
        json.dump(data, f)


def build_configs(input_json_path, dump_dir, model_name):
    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

    base = {**configs_base, **{"data": data_configs}, **inference_configs}

    def deep_update(t, p):
        for k, v in p.items():
            if isinstance(v, dict) and k in t and isinstance(t[k], dict):
                deep_update(t[k], v)
            else:
                t[k] = v

    deep_update(base, model_configs[model_name])
    arg_str = " ".join([
        f"--model_name {model_name}",
        f"--input_json_path {input_json_path}",
        f"--dump_dir {dump_dir}",
        f"--use_msa {USE_MSA}",
        f"--use_template {USE_TEMPLATE}",
        f"--use_rna_msa {USE_RNA_MSA}",
        f"--sample_diffusion.N_sample {N_SAMPLE}",
        f"--seeds {SEED}",
    ])
    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)


def coords_to_rows(target_id, seq, coords):
    rows = []
    for i in range(len(seq)):
        row = {"ID": f"{target_id}_{i + 1}", "resname": seq[i], "resid": i + 1}
        for s in range(N_SAMPLE):
            if s < coords.shape[0] and i < coords.shape[1]:
                x, y, z = coords[s, i]
            else:
                x, y, z = 0.0, 0.0, 0.0
            row[f"x_{s + 1}"] = np.clip(float(x), -999.999, 9999.999)
            row[f"y_{s + 1}"] = np.clip(float(y), -999.999, 9999.999)
            row[f"z_{s + 1}"] = np.clip(float(z), -999.999, 9999.999)
        rows.append(row)
    return rows


# ─────────────── TBM Phase ──────────────────────────────────────────────────
def tbm_phase(test_df, train_seqs_df, train_coords_dict, segments_map):
    print(f"\n{'='*60}")
    print(f"PHASE 1: Template-Based Modeling")
    print(f"  MIN_SIMILARITY={MIN_SIMILARITY}  MIN_PID={MIN_PERCENT_IDENTITY}")
    print(f"{'='*60}")
    t0 = time.time()

    template_predictions = {}
    protenix_queue = {}

    for _, row in test_df.iterrows():
        tid = row["target_id"]
        seq = row["sequence"]
        segs = segments_map.get(tid, [(0, len(seq))])

        similar = find_similar_sequences(seq, train_seqs_df, train_coords_dict, top_n=TOP_N_TEMPLATES)
        preds = []
        used = set()

        for i, (tmpl_id, tmpl_seq, sim, tmpl_coords, pct_id) in enumerate(similar):
            if len(preds) >= N_SAMPLE:
                break
            if sim < MIN_SIMILARITY or pct_id < MIN_PERCENT_IDENTITY:
                break
            if tmpl_id in used:
                continue

            rng = np.random.default_rng((row.name * 10000000000 + i * 10007) % (2**32))
            adapted = adapt_template_to_query(seq, tmpl_seq, tmpl_coords)

            # Confidence-scaled diversity transforms (OpenRNAFold style)
            slot = len(preds)
            if slot == 0:
                X = adapted
            elif slot == 1:
                X = adapted + rng.normal(0, max(0.01, (0.40 - sim) * 0.06), adapted.shape)
            elif slot == 2:
                longest = max(segs, key=lambda se: se[1] - se[0])
                X = apply_hinge(adapted, longest, rng, confidence=sim)
            elif slot == 3:
                X = jitter_chains(adapted, segs, rng, confidence=sim)
            else:
                X = smooth_wiggle(adapted, segs, rng, confidence=sim)

            refined = adaptive_rna_constraints(X, tid, segments_map, confidence=sim)
            preds.append(refined)
            used.add(tmpl_id)

        template_predictions[tid] = preds
        n_needed = N_SAMPLE - len(preds)
        if n_needed > 0:
            protenix_queue[tid] = (n_needed, seq)
            print(f"  {tid} ({len(seq)}nt): {len(preds)} TBM -> need {n_needed} Protenix")
        else:
            print(f"  {tid} ({len(seq)}nt): all {N_SAMPLE} from TBM")

    print(f"\nPhase 1 done in {time.time()-t0:.1f}s")
    print(f"  TBM-covered: {len(test_df) - len(protenix_queue)}, Need Protenix: {len(protenix_queue)}")
    return template_predictions, protenix_queue


# ─────────────── Protenix Phase (with chunking) ───────────────────────────
def protenix_phase(protenix_queue, test_df, segments_map):
    import torch
    from tqdm import tqdm

    print(f"\n{'='*60}")
    print(f"PHASE 2: Protenix for {len(protenix_queue)} targets (with chunking)")
    print(f"{'='*60}")

    if not os.path.isdir(CODE_DIR):
        print(f"  CODE_DIR not found: {CODE_DIR}")
        return {}

    root_dir = os.environ["PROTENIX_ROOT_DIR"]
    for p, name in [
        (Path(root_dir) / "checkpoint" / f"{MODEL_NAME}.pt", "checkpoint"),
        (Path(root_dir) / "common" / "components.cif", "CCD file"),
        (Path(root_dir) / "common" / "components.cif.rdkit_mol.pkl", "CCD cache"),
    ]:
        if not p.exists():
            print(f"  Missing {name}: {p}")
            return {}

    from protenix.data.inference.infer_dataloader import InferenceDataset
    from runner.inference import InferenceRunner, update_gpu_compatible_configs, update_inference_configs

    work_dir = Path("/kaggle/working") if IS_KAGGLE else Path("experiments/exp011/output")
    work_dir.mkdir(parents=True, exist_ok=True)

    # Build tasks: split long sequences into chunks
    tasks = []
    chunk_info = {}  # target_id -> [{"name": ..., "range": (s, e)}, ...]

    for target_id, (n_needed, full_seq) in protenix_queue.items():
        seq_len = len(full_seq)
        if seq_len <= MAX_SEQ_LEN:
            tasks.append({"target_id": target_id, "sequence": full_seq})
            chunk_info[target_id] = [{"name": target_id, "range": (0, seq_len)}]
            print(f"  {target_id} ({seq_len}nt): single pass")
        else:
            chunks = split_into_chunks(seq_len, MAX_SEQ_LEN, CHUNK_OVERLAP)
            chunk_info[target_id] = []
            for ci, (cs, ce) in enumerate(chunks):
                chunk_name = f"{target_id}_chunk{ci}"
                tasks.append({"target_id": chunk_name, "sequence": full_seq[cs:ce]})
                chunk_info[target_id].append({"name": chunk_name, "range": (cs, ce)})
            print(f"  {target_id} ({seq_len}nt): {len(chunks)} chunks {[(s,e) for s,e in chunks]}")

    tasks_df = pd.DataFrame(tasks)
    input_json_path = str(work_dir / "protenix_queue_input.json")
    build_input_json(tasks_df, input_json_path)

    configs = build_configs(input_json_path, str(work_dir / "outputs"), MODEL_NAME)
    configs = update_gpu_compatible_configs(configs)
    runner = InferenceRunner(configs)
    dataset = InferenceDataset(configs)

    raw_predictions = {}

    for i in tqdm(range(len(dataset)), desc="Protenix"):
        data, atom_array, error_message = dataset[i]
        sample_name = data.get("sample_name", f"sample_{i}")

        # Determine parent target
        parent_id = sample_name.split("_chunk")[0] if "_chunk" in sample_name else sample_name
        if parent_id not in protenix_queue:
            del data, atom_array, error_message
            gc.collect()
            torch.cuda.empty_cache()
            continue

        n_needed = protenix_queue[parent_id][0]

        if error_message:
            print(f"  {sample_name}: data error - {error_message}")
            raw_predictions[sample_name] = None
            del data, atom_array, error_message
            gc.collect()
            torch.cuda.empty_cache()
            continue

        try:
            new_cfg = update_inference_configs(configs, data["N_token"].item())
            new_cfg.sample_diffusion.N_sample = n_needed
            runner.update_model_configs(new_cfg)

            prediction = runner.predict(data)
            raw_coords = prediction["coordinate"]

            mask = get_c1_mask(data, atom_array).to(raw_coords.device)
            coords = raw_coords[:, mask, :].detach().cpu().numpy()

            # Collapse check
            if coords.shape[1] > 1:
                diffs = np.linalg.norm(coords[0, 1:] - coords[0, :-1], axis=-1)
                if np.all(diffs < 1e-4):
                    print(f"  WARNING: {sample_name} collapsed")
                    raw_predictions[sample_name] = None
                    continue

            sub_seq_len = data["N_token"].item()
            if coords.shape[1] != sub_seq_len:
                padded = np.zeros((coords.shape[0], sub_seq_len, 3), dtype=np.float32)
                ml = min(coords.shape[1], sub_seq_len)
                padded[:, :ml, :] = coords[:, :ml, :]
                coords = padded

            raw_predictions[sample_name] = coords
            print(f"  {sample_name}: {coords.shape[0]} preds, {coords.shape[1]}nt")

        except Exception as exc:
            print(f"  {sample_name}: FAILED - {exc}")
            import traceback
            traceback.print_exc()
            raw_predictions[sample_name] = None
        finally:
            del data, atom_array
            gc.collect()
            torch.cuda.empty_cache()

    # Post-process: stitch chunks back together
    protenix_preds = {}
    for target_id, (n_needed, full_seq) in protenix_queue.items():
        seq_len = len(full_seq)
        chunks = chunk_info.get(target_id, [])
        if not chunks:
            continue

        if len(chunks) == 1:
            coords = raw_predictions.get(target_id)
            if coords is not None:
                # Pad to full seq len if needed
                if coords.shape[1] != seq_len:
                    padded = np.zeros((coords.shape[0], seq_len, 3), dtype=np.float32)
                    ml = min(coords.shape[1], seq_len)
                    padded[:, :ml, :] = coords[:, :ml, :]
                    coords = padded
                protenix_preds[target_id] = coords
                print(f"  {target_id}: {coords.shape[0]} predictions")
            else:
                print(f"  {target_id}: FAILED")
                protenix_preds[target_id] = None
        else:
            # Stitch chunks
            all_ok = True
            chunk_results_per_sample = {s: [] for s in range(n_needed)}

            for cinfo in chunks:
                cname = cinfo["name"]
                crange = cinfo["range"]
                ccoords = raw_predictions.get(cname)

                if ccoords is None:
                    all_ok = False
                    break

                for s_idx in range(n_needed):
                    if s_idx < ccoords.shape[0]:
                        chunk_results_per_sample[s_idx].append((ccoords[s_idx], crange))
                    else:
                        chunk_results_per_sample[s_idx].append((ccoords[-1], crange))

            if not all_ok:
                print(f"  {target_id}: chunked inference incomplete")
                protenix_preds[target_id] = None
                continue

            stitched_samples = []
            for s_idx in range(n_needed):
                items = chunk_results_per_sample[s_idx]
                coords_list = [c for c, _ in items]
                ranges_list = [r for _, r in items]
                full_coords = stitch_chunk_coords(coords_list, ranges_list, seq_len)
                stitched_samples.append(full_coords)

            result = np.stack(stitched_samples, axis=0)
            protenix_preds[target_id] = result
            print(f"  {target_id}: {result.shape[0]} stitched predictions ({len(chunks)} chunks)")

    return protenix_preds


# ─────────────── Main ────────────────────────────────────────────────────────
def main():
    if os.path.isdir(CODE_DIR):
        os.environ["PROTENIX_ROOT_DIR"] = CODE_DIR
        sys.path.append(CODE_DIR)

        ckpt_dir = Path(CODE_DIR) / "checkpoint"
        expected_ckpt = ckpt_dir / f"{MODEL_NAME}.pt"
        if not expected_ckpt.exists() and ckpt_dir.is_dir():
            existing = list(ckpt_dir.glob("*.pt"))
            if existing:
                try:
                    os.symlink(str(existing[0]), str(expected_ckpt))
                    print(f"  Checkpoint alias: {existing[0].name} -> {expected_ckpt.name}")
                except OSError:
                    import shutil
                    writable_ckpt = Path("/kaggle/working" if IS_KAGGLE else "experiments/exp011/output") / "checkpoint"
                    writable_ckpt.mkdir(parents=True, exist_ok=True)
                    dst = writable_ckpt / f"{MODEL_NAME}.pt"
                    shutil.copy2(str(existing[0]), str(dst))
                    writable_root = writable_ckpt.parent
                    common_link = writable_root / "common"
                    if not common_link.exists():
                        os.symlink(str(Path(CODE_DIR) / "common"), str(common_link))
                    os.environ["PROTENIX_ROOT_DIR"] = str(writable_root)
                    print(f"  Checkpoint copied to writable dir: {dst}")
            else:
                print(f"  WARNING: No .pt files in {ckpt_dir}")
        elif expected_ckpt.exists():
            print(f"  Checkpoint OK: {expected_ckpt.name}")

    test_csv = os.path.join(DATA_BASE, "test_sequences.csv")
    train_csv = os.path.join(DATA_BASE, "train_sequences.csv")
    train_lbls = os.path.join(DATA_BASE, "train_labels.csv")
    val_csv = os.path.join(DATA_BASE, "validation_sequences.csv")
    val_lbls = os.path.join(DATA_BASE, "validation_labels.csv")

    test_df_full = pd.read_csv(test_csv)
    test_df = test_df_full.reset_index(drop=True)
    print(f"Test targets: {len(test_df)}")

    segments_map = build_segments_map(test_df)

    # Load template pool
    print("\nLoading templates...")
    train_seqs = pd.read_csv(train_csv)
    val_seqs = pd.read_csv(val_csv)
    train_labels = pd.read_csv(train_lbls, usecols=["ID", "resid", "x_1", "y_1", "z_1"],
                               dtype={"x_1": np.float32, "y_1": np.float32, "z_1": np.float32})
    val_labels = pd.read_csv(val_lbls, usecols=["ID", "resid", "x_1", "y_1", "z_1"],
                             dtype={"x_1": np.float32, "y_1": np.float32, "z_1": np.float32})

    combined_seqs = pd.concat([train_seqs, val_seqs], ignore_index=True)
    combined_labels = pd.concat([train_labels, val_labels], ignore_index=True)
    train_coords = process_labels(combined_labels)
    print(f"  Train+Val: {len(combined_seqs)} sequences, {len(train_coords)} structures")

    existing_pdb_ids = set(combined_seqs["target_id"].str[:4].str.upper())
    pdb_seqs, pdb_coords = load_pdb_rna_templates(PDB_RNA_DIR, existing_pdb_ids)
    if len(pdb_seqs) > 0:
        combined_seqs = pd.concat([combined_seqs, pdb_seqs], ignore_index=True)
        train_coords.update(pdb_coords)

    existing_ids = set(combined_seqs["target_id"])
    extra_seqs, extra_coords = load_new_pdb_templates(existing_ids)
    if len(extra_seqs) > 0:
        combined_seqs = pd.concat([combined_seqs, extra_seqs], ignore_index=True)
        train_coords.update(extra_coords)

    print(f"  Total template pool: {len(combined_seqs)} sequences, {len(train_coords)} structures")

    # Phase 1: TBM
    seed_everything(SEED)
    template_preds, protenix_queue = tbm_phase(test_df, combined_seqs, train_coords, segments_map)

    # Phase 2: Protenix (with chunking for long sequences)
    protenix_preds = {}
    if protenix_queue and IS_KAGGLE:
        protenix_preds = protenix_phase(protenix_queue, test_df, segments_map)
    elif protenix_queue:
        print(f"\nSkipping Protenix (not on Kaggle). {len(protenix_queue)} targets will use de-novo.")

    # Phase 3: Combine (TBM first, then Protenix, then de-novo)
    print(f"\n{'='*60}")
    print("PHASE 3: Combine")
    print(f"{'='*60}")

    all_rows = []
    for _, row in test_df.iterrows():
        tid = row["target_id"]
        seq = row["sequence"]

        combined = list(template_preds.get(tid, []))

        ptx = protenix_preds.get(tid)
        if ptx is not None and ptx.ndim == 3:
            for j in range(ptx.shape[0]):
                if len(combined) >= N_SAMPLE:
                    break
                combined.append(ptx[j])

        n_denovo = 0
        while len(combined) < N_SAMPLE:
            seed_val = row.name * 1000000 + len(combined) * 1000
            dn = generate_rna_structure(seq, seed=seed_val)
            combined.append(adaptive_rna_constraints(dn, tid, segments_map, confidence=0.2))
            n_denovo += 1

        if n_denovo:
            print(f"  {tid}: {n_denovo} de-novo fallback(s)")

        stacked = np.stack(combined[:N_SAMPLE], axis=0)
        all_rows.extend(coords_to_rows(tid, seq, stacked))

    sub = pd.DataFrame(all_rows)
    cols = ["ID", "resname", "resid"] + [
        f"{c}_{i}" for i in range(1, N_SAMPLE + 1) for c in ["x", "y", "z"]
    ]
    sub[cols].to_csv(OUTPUT_PATH, index=False)
    print(f"\nSaved submission to {OUTPUT_PATH} ({len(sub):,} rows)")


if __name__ == "__main__":
    main()
