{"cells":[{"cell_type":"markdown","id":"intro","metadata":{},"source":"# RNA 3D: Protenix + Template-Based Modeling (TBM) Hybrid\n\n**Strategy**: Two-phase approach for RNA 3D structure prediction.\n\n### Phase 1 — Template-Based Modeling (TBM)\n- Align each test sequence against all training sequences using Biopython `PairwiseAligner`\n- If best match has ≥50% identity, morph the template's known 3D C1' coordinates onto the query\n- Generate up to 5 diverse predictions per target via hinge rotation, jitter, and wiggle transforms\n- Apply physics-inspired RNA constraints (bond length, angle, clash)\n\n### Phase 2 — Protenix (ByteDance AlphaFold3 adaptation for RNA)\n- For targets without good TBM matches, run Protenix GPU inference\n- Sequences >512 nt are split into overlapping chunks, predicted separately, then stitched with Kabsch alignment\n- Only generates as many predictions as needed (fills remaining slots)\n\n**Datasets used**:\n- `qiweiyin/protenix-v1-adjusted` — Protenix v1 model + weights\n- `kami1976/biopython-cp312` — Biopython wheel (offline install)\n- `amirrezaaleyasin/biotite` — Biotite wheel\n- `amirrezaaleyasin/rdkit-2025-9-5` — RDKit wheel (Protenix dependency)"},{"cell_type":"code","execution_count":null,"id":"cell-install","metadata":{},"outputs":[],"source":"# Install offline wheels (no internet needed)\nimport subprocess, sys, os\n\ndef install_wheel(path, quiet=True):\n    if not os.path.exists(path):\n        print(f'  SKIP (not found): {path}')\n        return\n    cmd = [sys.executable, '-m', 'pip', 'install', '--no-index', '--no-deps', path]\n    if quiet: cmd += ['-q']\n    subprocess.check_call(cmd)\n    print(f'  Installed: {os.path.basename(path)}')\n\ntry:\n    from Bio.Align import PairwiseAligner\n    print('biopython already available')\nexcept ImportError:\n    bp_dir = '/kaggle/input/biopython-cp312'\n    if os.path.isdir(bp_dir):\n        for root, dirs, files in os.walk(bp_dir):\n            for f in sorted(files):\n                if f.endswith('.whl') and 'biopython' in f:\n                    install_wheel(os.path.join(root, f)); break\n\nfor whl_dir in ['/kaggle/input/biotite', '/kaggle/input/rdkit-2025-9-5']:\n    if os.path.isdir(whl_dir):\n        for root, dirs, files in os.walk(whl_dir):\n            for f in sorted(files):\n                if f.endswith('.whl'):\n                    install_wheel(os.path.join(root, f)); break\nprint('Done.')"},{"cell_type":"code","execution_count":null,"id":"cell-main","metadata":{},"outputs":[],"source":"# RNA 3D: Protenix + TBM Hybrid\nimport gc, json, os, sys, time\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom tqdm import tqdm\nfrom Bio.Align import PairwiseAligner\n\nos.environ['LAYERNORM_TYPE'] = 'torch'\nos.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'\nos.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'expandable_segments:True'\n\n# ── Competition data path ─────────────────────────────────────────────────\nfor _cand in ['/kaggle/input/stanford-rna-3d-folding-2',\n               '/kaggle/input/competitions/stanford-rna-3d-folding-2']:\n    if os.path.isdir(_cand):\n        DATA_BASE = _cand; break\nelse:\n    raise FileNotFoundError('Competition data not found')\nprint(f'DATA_BASE: {DATA_BASE}')\n\n# ── Protenix model path ───────────────────────────────────────────────────\n_PTX_BASE = '/kaggle/input/protenix-v1-adjusted/Protenix-v1-adjust-v2/Protenix-v1-adjust-v2/Protenix-v1'\nDEFAULT_CODE_DIR = _PTX_BASE\nDEFAULT_ROOT_DIR = _PTX_BASE\n\n# ── Config ────────────────────────────────────────────────────────────────\nMODEL_NAME    = 'protenix_base_20250630_v1.0.0'\nN_SAMPLE      = 5      # predictions per target\nSEED          = 42\nMAX_SEQ_LEN   = 512    # split longer sequences into chunks\nCHUNK_OVERLAP = 128    # overlap between chunks for Kabsch stitching\nMIN_PERCENT_IDENTITY = 50.0  # minimum identity to use TBM\nMIN_SIMILARITY       = 0.0\nUSE_PROTENIX = True\nOUTPUT_CSV = '/kaggle/working/submission.csv'\n\n\n# ════════════════════════════════════════════════════════════════════════\n# UTILITIES\n# ════════════════════════════════════════════════════════════════════════\n\ndef seed_everything(seed):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    np.random.seed(seed)\n    torch.backends.cudnn.benchmark = False\n    torch.backends.cudnn.deterministic = True\n    torch.use_deterministic_algorithms(True)\n\n\n# ── Chunking for long sequences ───────────────────────────────────────────\ndef get_chunks(seq_len, max_len=MAX_SEQ_LEN, overlap=CHUNK_OVERLAP):\n    \"\"\"Split a long sequence into overlapping chunks.\"\"\"\n    if seq_len <= max_len:\n        return [(0, seq_len)]\n    chunks = []\n    start = 0\n    while start < seq_len:\n        end = min(start + max_len, seq_len)\n        chunks.append((start, end))\n        if end == seq_len: break\n        start = end - overlap\n    return chunks\n\n\n# ── Kabsch alignment (SVD-based rigid alignment) ──────────────────────────\ndef kabsch_align(moving, fixed):\n    \"\"\"Align 'moving' onto 'fixed' using Kabsch/SVD algorithm.\"\"\"\n    mu_m, mu_f = moving.mean(0), fixed.mean(0)\n    Mc, Fc = moving - mu_m, fixed - mu_f\n    H = Mc.T @ Fc\n    U, S, Vt = np.linalg.svd(H)\n    d = np.sign(np.linalg.det(Vt.T @ U.T))\n    R = Vt.T @ np.diag([1.0, 1.0, d]) @ U.T\n    return Mc @ R.T + mu_f, R, mu_m, mu_f\n\n\ndef stitch_chunk_coords(chunk_coords_list, chunk_ranges, full_len):\n    \"\"\"\n    Stitch per-chunk predictions into a full-length structure.\n    Overlapping regions are aligned with Kabsch and blended linearly.\n    \"\"\"\n    n_sample = chunk_coords_list[0].shape[0]\n    result = np.zeros((n_sample, full_len, 3), dtype=np.float32)\n    for s in range(n_sample):\n        placed = np.zeros((full_len, 3), dtype=np.float32)\n        s0, e0 = chunk_ranges[0]\n        placed[s0:e0] = chunk_coords_list[0][s]\n        for i in range(1, len(chunk_coords_list)):\n            cs, ce = chunk_ranges[i]\n            prev_e = chunk_ranges[i-1][1]\n            cp = chunk_coords_list[i][s]\n            ov_s, ov_e = cs, prev_e\n            ov_len = ov_e - ov_s\n            if ov_len > 3:\n                moving = cp[:ov_len]\n                fixed  = placed[ov_s:ov_e]\n                cp_aligned, R, mu_m, mu_f = kabsch_align(moving, fixed)\n                rest = cp[ov_len:]\n                rest_aligned = (rest - mu_m) @ R.T + mu_f\n                # Linear blend in overlap region\n                t = np.linspace(0, 1, ov_len)[:, None]\n                placed[ov_s:ov_e] = (1-t)*placed[ov_s:ov_e] + t*cp_aligned\n                if ce > ov_e:\n                    placed[ov_e:ce] = rest_aligned\n            else:\n                placed[cs:ce] = cp\n        result[s] = placed\n    return result\n\n\n# ── Protenix config builder ───────────────────────────────────────────────\ndef build_input_json(rows, json_path):\n    data = [\n        {'name': r['target_id'], 'covalent_bonds': [],\n         'sequences': [{'rnaSequence': {'sequence': r['sequence'], 'count': 1}}]}\n        for r in rows\n    ]\n    with open(json_path, 'w') as f:\n        json.dump(data, f)\n\n\ndef build_configs(input_json_path, dump_dir, model_name, n_sample):\n    from configs.configs_base import configs as configs_base\n    from configs.configs_data import data_configs\n    from configs.configs_inference import inference_configs\n    from configs.configs_model_type import model_configs\n    from protenix.config.config import parse_configs\n    base = {**configs_base, **{'data': data_configs}, **inference_configs}\n    def deep_update(t, p):\n        for k, v in p.items():\n            if isinstance(v, dict) and k in t and isinstance(t[k], dict):\n                deep_update(t[k], v)\n            else:\n                t[k] = v\n    deep_update(base, model_configs[model_name])\n    arg_str = ' '.join([\n        f'--model_name {model_name}',\n        f'--input_json_path {input_json_path}',\n        f'--dump_dir {dump_dir}',\n        '--use_msa false',\n        '--use_template false',\n        '--use_rna_msa false',\n        f'--sample_diffusion.N_sample {n_sample}',\n        f'--seeds {SEED}',\n    ])\n    return parse_configs(configs=base, arg_str=arg_str, fill_required_with_null=True)\n\n\ndef get_c1_mask(data, atom_array, seq_len):\n    \"\"\"Extract the C1' atom mask from Protenix output features.\"\"\"\n    if atom_array is not None:\n        try:\n            if hasattr(atom_array, 'centre_atom_mask'):\n                m = atom_array.centre_atom_mask == 1\n                if hasattr(atom_array, 'is_rna'): m = m & atom_array.is_rna\n                return torch.from_numpy(m).bool()\n            if hasattr(atom_array, 'atom_name'):\n                base = atom_array.atom_name == \"C1'\"\n                if hasattr(atom_array, 'is_rna'): base = base & atom_array.is_rna\n                return torch.from_numpy(base).bool()\n        except Exception:\n            pass\n    f = data['input_feature_dict']\n    if 'centre_atom_mask' in f: return (f['centre_atom_mask'] == 1).bool()\n    if 'center_atom_mask'  in f: return (f['center_atom_mask']  == 1).bool()\n    m11 = (f['atom_to_tokatom_idx'] == 11).bool()\n    m12 = (f['atom_to_tokatom_idx'] == 12).bool()\n    return m11 if abs(m11.sum().item()-seq_len) < abs(m12.sum().item()-seq_len) else m12\n\n\ndef coords_to_rows(target_id, seq, coords):\n    rows = []\n    for i in range(len(seq)):\n        row = {'ID': f'{target_id}_{i+1}', 'resname': seq[i], 'resid': i+1}\n        for s in range(N_SAMPLE):\n            if s < coords.shape[0] and i < coords.shape[1]:\n                x, y, z = float(coords[s,i,0]), float(coords[s,i,1]), float(coords[s,i,2])\n            else:\n                x, y, z = 0.0, 0.0, 0.0\n            row[f'x_{s+1}'] = x; row[f'y_{s+1}'] = y; row[f'z_{s+1}'] = z\n        rows.append(row)\n    return rows\n\n\n# ════════════════════════════════════════════════════════════════════════\n# DATA LOADING\n# ════════════════════════════════════════════════════════════════════════\n\ndef process_labels(labels_df):\n    \"\"\"Convert labels DataFrame to {target_id: np.array (L,3)} of C1' coordinates.\"\"\"\n    coords = {}\n    prefixes = labels_df['ID'].str.rsplit('_', n=1).str[0]\n    for prefix, grp in labels_df.groupby(prefixes):\n        coords[prefix] = grp.sort_values('resid')[['x_1','y_1','z_1']].values\n    return coords\n\n\ndef parse_stoichiometry(stoich):\n    if pd.isna(stoich) or str(stoich).strip() == '': return []\n    return [(ch.strip(), int(cnt)) for part in str(stoich).split(';')\n            for ch, cnt in [part.split(':')]]\n\n\ndef parse_fasta(fasta_content):\n    out, cur, parts = {}, None, []\n    for line in str(fasta_content).splitlines():\n        line = line.strip()\n        if not line: continue\n        if line.startswith('>'):\n            if cur is not None: out[cur] = ''.join(parts)\n            cur = line[1:].split()[0]; parts = []\n        else:\n            parts.append(line.replace(' ', ''))\n    if cur is not None: out[cur] = ''.join(parts)\n    return out\n\n\ndef get_chain_segments(row):\n    \"\"\"Return list of (start, end) index ranges for each RNA chain in the target.\"\"\"\n    seq = row['sequence']\n    stoich = row.get('stoichiometry', '')\n    all_sq  = row.get('all_sequences', '')\n    if pd.isna(stoich) or pd.isna(all_sq) or str(stoich).strip() == '' or str(all_sq).strip() == '':\n        return [(0, len(seq))]\n    try:\n        chain_dict = parse_fasta(all_sq)\n        order = parse_stoichiometry(stoich)\n        segs, pos = [], 0\n        for ch, cnt in order:\n            base = chain_dict.get(ch)\n            if base is None: return [(0, len(seq))]\n            for _ in range(cnt):\n                segs.append((pos, pos+len(base))); pos += len(base)\n        return segs if pos == len(seq) else [(0, len(seq))]\n    except Exception:\n        return [(0, len(seq))]\n\n\ndef build_segments_map(df):\n    return {r['target_id']: get_chain_segments(r) for _, r in df.iterrows()}\n\n\n# ════════════════════════════════════════════════════════════════════════\n# TEMPLATE-BASED MODELING (TBM)\n# ════════════════════════════════════════════════════════════════════════\n\ndef _make_aligner():\n    al = PairwiseAligner()\n    al.mode = 'global'\n    al.match_score = 2.0\n    al.mismatch_score = -1.0\n    al.open_gap_score = -10\n    al.extend_gap_score = -0.5\n    al.query_left_open_gap_score    = al.query_right_open_gap_score    = -10\n    al.query_left_extend_gap_score  = al.query_right_extend_gap_score  = -0.5\n    al.target_left_open_gap_score   = al.target_right_open_gap_score   = -10\n    al.target_left_extend_gap_score = al.target_right_extend_gap_score = -0.5\n    return al\n\n_aligner = _make_aligner()\n\n\ndef find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, top_n=50):\n    \"\"\"Find the top_n most similar training sequences to query_seq.\"\"\"\n    results = []\n    for _, row in train_seqs_df.iterrows():\n        tid, tseq = row['target_id'], row['sequence']\n        if tid not in train_coords_dict: continue\n        # Quick length filter: skip if >30% length difference\n        if abs(len(tseq)-len(query_seq)) / max(len(tseq), len(query_seq)) > 0.3: continue\n        aln = next(iter(_aligner.align(query_seq, tseq)))\n        norm_s = aln.score / (2.0 * min(len(query_seq), len(tseq)))\n        identical = sum(\n            1 for (qs,qe),(ts,te) in zip(*aln.aligned)\n            for qp, tp in zip(range(qs,qe), range(ts,te))\n            if query_seq[qp] == tseq[tp]\n        )\n        pct_id = 100 * identical / len(query_seq)\n        results.append((tid, tseq, norm_s, train_coords_dict[tid], pct_id, aln))\n    results.sort(key=lambda x: x[2], reverse=True)\n    return results[:top_n]\n\n\ndef adapt_template_to_query(query_seq, template_coords, aln):\n    \"\"\"Map template C1' coordinates onto query sequence positions via alignment.\"\"\"\n    new_coords = np.full((len(query_seq), 3), np.nan)\n    for (qs, qe), (ts, te) in zip(*aln.aligned):\n        chunk = template_coords[ts:te]\n        if len(chunk) == (qe - qs):\n            new_coords[qs:qe] = chunk\n    # Fill gaps by linear interpolation\n    for i in range(len(new_coords)):\n        if np.isnan(new_coords[i, 0]):\n            pv = next((j for j in range(i-1, -1, -1) if not np.isnan(new_coords[j,0])), -1)\n            nv = next((j for j in range(i+1, len(new_coords)) if not np.isnan(new_coords[j,0])), -1)\n            if pv >= 0 and nv >= 0:\n                w = (i-pv)/(nv-pv)\n                new_coords[i] = (1-w)*new_coords[pv] + w*new_coords[nv]\n            elif pv >= 0: new_coords[i] = new_coords[pv] + [3, 0, 0]\n            elif nv >= 0: new_coords[i] = new_coords[nv] + [3, 0, 0]\n            else: new_coords[i] = [i*3, 0, 0]\n    return np.nan_to_num(new_coords)\n\n\n# ── RNA physics constraints ───────────────────────────────────────────────\n\ndef adaptive_rna_constraints(coords, target_id, segments_map, confidence=1.0, passes=2):\n    \"\"\"Soft physics constraints: bond length ~5.9Å, 2-bond dist ~10.2Å, clash >3.2Å.\"\"\"\n    if confidence >= 0.99:\n        return coords.copy()\n    X = coords.copy()\n    segments = segments_map.get(target_id, [(0, len(X))])\n    strength = max(0.75 * (1.0 - min(confidence, 0.97)), 0.02)\n    for _ in range(passes):\n        for s, e in segments:\n            C = X[s:e]; L = e - s\n            if L < 3: continue\n            d = C[1:]-C[:-1]; dist = np.linalg.norm(d, axis=1)+1e-6\n            adj = d * ((5.95-dist)/dist)[:,None] * (0.22*strength)\n            C[:-1] -= adj; C[1:] += adj\n            d2 = C[2:]-C[:-2]; d2n = np.linalg.norm(d2, axis=1)+1e-6\n            adj2 = d2 * ((10.2-d2n)/d2n)[:,None] * (0.10*strength)\n            C[:-2] -= adj2; C[2:] += adj2\n            C[1:-1] += (0.06*strength) * (0.5*(C[:-2]+C[2:])-C[1:-1])\n            if L >= 25:\n                idx = np.linspace(0, L-1, min(L,160)).astype(int) if L > 220 else np.arange(L)\n                P = C[idx]; diff = P[:,None,:] - P[None,:,:]\n                dm = np.linalg.norm(diff, axis=2)+1e-6\n                sep = np.abs(idx[:,None] - idx[None,:])\n                mask = (sep > 2) & (dm < 3.2)\n                if np.any(mask):\n                    vec = (diff*((3.2-dm)/dm)[:,:,None]*mask[:,:,None]).sum(axis=1)\n                    C[idx] += (0.015*strength) * vec\n            X[s:e] = C\n    return X\n\n\ndef _rotmat(axis, ang):\n    a = np.asarray(axis, float); a /= np.linalg.norm(a)+1e-12\n    x,y,z = a; c,s = np.cos(ang),np.sin(ang); CC = 1-c\n    return np.array([[c+x*x*CC, x*y*CC-z*s, x*z*CC+y*s],\n                     [y*x*CC+z*s, c+y*y*CC, y*z*CC-x*s],\n                     [z*x*CC-y*s, z*y*CC+x*s, c+z*z*CC]])\n\n\ndef apply_hinge(coords, seg, rng, deg=22):\n    \"\"\"Rotate one half of a chain segment around a random pivot.\"\"\"\n    s, e = seg; L = e - s\n    if L < 30: return coords\n    pivot = s + int(rng.integers(10, L-10))\n    R = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n    X = coords.copy(); p0 = X[pivot].copy()\n    X[pivot+1:e] = (X[pivot+1:e] - p0) @ R.T + p0\n    return X\n\n\ndef jitter_chains(coords, segs, rng, deg=12, trans=1.5):\n    \"\"\"Apply random rotation+translation to each chain segment.\"\"\"\n    X = coords.copy(); gc_ = X.mean(0, keepdims=True)\n    for s, e in segs:\n        R = _rotmat(rng.normal(size=3), np.deg2rad(float(rng.uniform(-deg, deg))))\n        shift = rng.normal(size=3)\n        shift = shift / (np.linalg.norm(shift)+1e-12) * float(rng.uniform(0, trans))\n        c = X[s:e].mean(0, keepdims=True)\n        X[s:e] = (X[s:e]-c) @ R.T + c + shift\n    X -= X.mean(0, keepdims=True) - gc_\n    return X\n\n\ndef smooth_wiggle(coords, segs, rng, amp=0.8):\n    \"\"\"Add smooth low-frequency displacement noise along the chain.\"\"\"\n    X = coords.copy()\n    for s, e in segs:\n        L = e - s\n        if L < 20: continue\n        ctrl = np.linspace(0, L-1, 6)\n        disp = rng.normal(0, amp, (6,3))\n        t = np.arange(L)\n        X[s:e] += np.vstack([np.interp(t, ctrl, disp[:,k]) for k in range(3)]).T\n    return X\n\n\ndef generate_rna_structure(sequence, seed=None):\n    \"\"\"Simple de-novo A-form helix as fallback.\"\"\"\n    if seed is not None: np.random.seed(seed)\n    n = len(sequence)\n    coords = np.zeros((n, 3))\n    for i in range(n):\n        ang = i * 0.6\n        coords[i] = [10.0*np.cos(ang), 10.0*np.sin(ang), i*2.5]\n    return coords\n\n\n# ════════════════════════════════════════════════════════════════════════\n# TBM PHASE\n# ════════════════════════════════════════════════════════════════════════\n\ndef tbm_phase(test_df, train_seqs_df, train_coords_dict, segments_map):\n    print(f\"\\n{'='*60}\")\n    print('PHASE 1: Template-Based Modeling')\n    print(f\"{'='*60}\")\n    t0 = time.time()\n    template_predictions = {}\n    protenix_queue = {}\n\n    for _, row in test_df.iterrows():\n        tid  = row['target_id']\n        seq  = row['sequence']\n        segs = segments_map.get(tid, [(0, len(seq))])\n\n        similar = find_similar_sequences(seq, train_seqs_df, train_coords_dict, top_n=50)\n        preds = []; used = set()\n\n        for i, (tmpl_id, tmpl_seq, sim, tmpl_coords, pct_id, aln) in enumerate(similar):\n            if len(preds) >= N_SAMPLE: break\n            if sim < MIN_SIMILARITY or pct_id < MIN_PERCENT_IDENTITY: break\n            if tmpl_id in used: continue\n\n            rng = np.random.default_rng((row.name * 10000000000 + i * 10007) % (2**32))\n            adapted = adapt_template_to_query(seq, tmpl_coords, aln)\n            slot = len(preds)\n            if slot == 0:   X = adapted\n            elif slot == 1: X = adapted + rng.normal(0, max(0.01,(0.40-sim)*0.06), adapted.shape)\n            elif slot == 2:\n                longest = max(segs, key=lambda se: se[1]-se[0])\n                X = apply_hinge(adapted, longest, rng)\n            elif slot == 3: X = jitter_chains(adapted, segs, rng)\n            else:           X = smooth_wiggle(adapted, segs, rng)\n\n            refined = adaptive_rna_constraints(X, tid, segments_map, confidence=sim)\n            preds.append(refined); used.add(tmpl_id)\n\n        template_predictions[tid] = preds\n        n_needed = N_SAMPLE - len(preds)\n        if n_needed > 0:\n            protenix_queue[tid] = (n_needed, seq)\n            print(f'  {tid} ({len(seq)}nt): {len(preds)} TBM + {n_needed} Protenix  chunks={len(get_chunks(len(seq)))}')\n        else:\n            print(f'  {tid} ({len(seq)}nt): {N_SAMPLE} from TBM')\n\n    print(f'Phase 1: {time.time()-t0:.1f}s | '\n          f'TBM-only: {len(test_df)-len(protenix_queue)} | '\n          f'Protenix queue: {len(protenix_queue)}')\n    return template_predictions, protenix_queue\n\n\n# ════════════════════════════════════════════════════════════════════════\n# MAIN\n# ════════════════════════════════════════════════════════════════════════\n\ndef main():\n    # Locate Protenix code\n    code_dir = DEFAULT_CODE_DIR\n    if not os.path.isdir(code_dir):\n        for root, dirs, files in os.walk('/kaggle/input/protenix-v1-adjusted'):\n            for d in dirs:\n                if d == 'Protenix-v1':\n                    code_dir = os.path.join(root, d); break\n            if os.path.isdir(code_dir): break\n    if not os.path.isdir(code_dir):\n        raise FileNotFoundError('Protenix not found')\n    os.environ['PROTENIX_ROOT_DIR'] = code_dir\n    sys.path.insert(0, code_dir)\n    seed_everything(SEED)\n\n    print('Loading competition data...')\n    test_df      = pd.read_csv(f'{DATA_BASE}/test_sequences.csv')\n    train_seqs   = pd.read_csv(f'{DATA_BASE}/train_sequences.csv')\n    train_labels = pd.read_csv(f'{DATA_BASE}/train_labels.csv')\n    train_coords  = process_labels(train_labels)\n    segments_map  = build_segments_map(test_df)\n    print(f'Test targets: {len(test_df)} | Template pool: {len(train_seqs)}')\n\n    template_preds, protenix_queue = tbm_phase(\n        test_df.reset_index(drop=True), train_seqs, train_coords, segments_map)\n\n    # ── Phase 2: Protenix ──────────────────────────────────────────────\n    protenix_preds = {}\n    if protenix_queue and USE_PROTENIX:\n        print(f\"\\n{'='*60}\")\n        print(f'PHASE 2: Protenix for {len(protenix_queue)} targets')\n        print(f\"{'='*60}\")\n        work_dir = Path('/kaggle/working')\n\n        chunk_meta    = {}\n        json_rows     = []\n        chunk_seq_len = {}\n        chunk_n_needed = {}\n\n        for _, row in test_df.iterrows():\n            tid = row['target_id']\n            if tid not in protenix_queue: continue\n            n_needed, full_seq = protenix_queue[tid]\n            chunks = get_chunks(len(full_seq))\n            if len(chunks) == 1:\n                chunk_meta[tid] = None\n                chunk_seq = full_seq[:MAX_SEQ_LEN]\n                json_rows.append({'target_id': tid, 'sequence': chunk_seq})\n                chunk_seq_len[tid] = len(chunk_seq)\n                chunk_n_needed[tid] = n_needed\n            else:\n                ci_list = []\n                for ci, (cs, ce) in enumerate(chunks):\n                    chunk_id  = f'{tid}__chunk{ci}'\n                    chunk_seq = full_seq[cs:ce]\n                    json_rows.append({'target_id': chunk_id, 'sequence': chunk_seq})\n                    chunk_seq_len[chunk_id] = len(chunk_seq)\n                    chunk_n_needed[chunk_id] = n_needed\n                    ci_list.append((chunk_id, cs, ce))\n                chunk_meta[tid] = ci_list\n                print(f'  {tid} ({len(full_seq)}nt) -> {len(chunks)} chunks')\n\n        input_json_path = str(work_dir / 'protenix_input.json')\n        build_input_json(json_rows, input_json_path)\n\n        from protenix.data.inference.infer_dataloader import InferenceDataset\n        from runner.inference import (InferenceRunner,\n                                      update_gpu_compatible_configs,\n                                      update_inference_configs)\n\n        configs = build_configs(input_json_path, str(work_dir / 'ptx_out'), MODEL_NAME, n_sample=1)\n        configs = update_gpu_compatible_configs(configs)\n        runner  = InferenceRunner(configs)\n        dataset = InferenceDataset(configs)\n        raw_preds = {}\n\n        for i in tqdm(range(len(dataset)), desc='Protenix'):\n            data, atom_array, error_message = dataset[i]\n            entry_id = data.get('sample_name', f'sample_{i}')\n            n_needed = chunk_n_needed.get(entry_id, N_SAMPLE)\n            seq_len  = chunk_seq_len.get(entry_id, data['N_token'].item())\n            if error_message:\n                print(f'  {entry_id}: data error — {error_message}')\n                raw_preds[entry_id] = None\n                del data, atom_array; gc.collect(); torch.cuda.empty_cache(); continue\n            try:\n                new_cfg = update_inference_configs(configs, data['N_token'].item())\n                new_cfg.sample_diffusion.N_sample = n_needed\n                runner.update_model_configs(new_cfg)\n                prediction = runner.predict(data)\n                raw_coords = prediction['coordinate']\n                mask   = get_c1_mask(data, atom_array, seq_len)\n                mask   = mask.to(raw_coords.device)\n                coords = raw_coords[:, mask, :].detach().cpu().numpy()\n                print(f'  {entry_id}: coords {coords.shape}')\n                if coords.shape[1] > 1:\n                    diffs = np.linalg.norm(coords[0,1:]-coords[0,:-1], axis=-1)\n                    if np.all(diffs < 1e-4):\n                        print(f'    WARNING: collapsed'); raw_preds[entry_id] = None; continue\n                if coords.shape[1] != seq_len:\n                    padded = np.zeros((coords.shape[0], seq_len, 3), dtype=np.float32)\n                    ml = min(coords.shape[1], seq_len)\n                    padded[:, :ml, :] = coords[:, :ml, :]\n                    coords = padded\n                raw_preds[entry_id] = coords\n            except Exception as exc:\n                import traceback\n                print(f'  {entry_id}: FAILED — {exc}'); traceback.print_exc()\n                raw_preds[entry_id] = None\n            finally:\n                try: del data, atom_array\n                except: pass\n                gc.collect(); torch.cuda.empty_cache()\n\n        # Stitch chunks\n        for tid, ci_list in chunk_meta.items():\n            full_seq = protenix_queue[tid][1]\n            if ci_list is None:\n                protenix_preds[tid] = raw_preds.get(tid)\n            else:\n                coords_list, ranges_list, ok = [], [], True\n                for chunk_id, cs, ce in ci_list:\n                    cc = raw_preds.get(chunk_id)\n                    if cc is None: ok = False; break\n                    coords_list.append(cc); ranges_list.append((cs, ce))\n                if ok and coords_list:\n                    stitched = stitch_chunk_coords(coords_list, ranges_list, len(full_seq))\n                    print(f'  {tid}: stitched {len(coords_list)} chunks -> {stitched.shape}')\n                    protenix_preds[tid] = stitched\n                else:\n                    protenix_preds[tid] = None\n\n    # ── Phase 3: Combine ───────────────────────────────────────────────\n    print(f\"\\n{'='*60}\")\n    print('PHASE 3: Combine TBM + Protenix + de-novo fallback')\n    print(f\"{'='*60}\")\n    all_rows = []\n    for _, row in test_df.iterrows():\n        tid  = row['target_id']\n        seq  = row['sequence']\n        combined = list(template_preds.get(tid, []))\n        ptx = protenix_preds.get(tid)\n        if ptx is not None and ptx.ndim == 3:\n            for j in range(ptx.shape[0]):\n                if len(combined) >= N_SAMPLE: break\n                combined.append(ptx[j])\n        # De-novo fallback for any remaining slots\n        n_denovo = 0\n        while len(combined) < N_SAMPLE:\n            seed_val = (hash(tid) + len(combined) * 997) % (2**32)\n            dn = generate_rna_structure(seq, seed=seed_val)\n            combined.append(adaptive_rna_constraints(dn, tid, segments_map, confidence=0.2))\n            n_denovo += 1\n        if n_denovo:\n            print(f'  {tid}: {n_denovo} de-novo slot(s)')\n        stacked = np.stack(combined[:N_SAMPLE], axis=0)\n        all_rows.extend(coords_to_rows(tid, seq, stacked))\n\n    sub = pd.DataFrame(all_rows)\n    cols = ['ID','resname','resid'] + [f'{c}_{i}' for i in range(1, N_SAMPLE+1) for c in ['x','y','z']]\n    coord_cols = [c for c in cols if c.startswith(('x_','y_','z_'))]\n    sub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)\n    sub[cols].to_csv(OUTPUT_CSV, index=False)\n    print(f'Saved {len(sub):,} rows to {OUTPUT_CSV}')\n\n\nmain()"},{"cell_type":"code","execution_count":null,"id":"cell-verify","metadata":{},"outputs":[],"source":"# Verify submission\nimport pandas as pd\nsub = pd.read_csv('/kaggle/working/submission.csv')\nprint(f'Shape: {sub.shape}')\nprint(sub.head(3))\nprint(f'\\nNaN count: {sub.isna().sum().sum()}')\nprint(f'x_1 range: [{sub.x_1.min():.2f}, {sub.x_1.max():.2f}]')"}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.0"}},"nbformat":4,"nbformat_minor":5}