import subprocess
import sys
subprocess.check_call([sys.executable, "-m", "pip", "install", "--no-index", "/kaggle/input/biopython-cp312/biopython-1.86-cp312-cp312-manylinux2014_x86_64.manylinux_2_17_x86_64.manylinux_2_28_x86_64.whl"])

import pandas as pd
import numpy as np
import random
import time
import warnings
import os, sys
from scipy.spatial.distance import cdist
from scipy.optimize import minimize
warnings.filterwarnings('ignore')

DATA_PATH = '/kaggle/input/stanford-rna-3d-folding-2/'
train_seqs = pd.read_csv(DATA_PATH + 'train_sequences.csv')
test_seqs = pd.read_csv(DATA_PATH + 'test_sequences.csv')
train_labels = pd.read_csv(DATA_PATH + 'train_labels.csv')

sys.path.append(os.path.join(DATA_PATH, "extra"))

# --- Robust import for Kaggle's extra/parse_fasta_py.py ---
try:
    import typing as _typing
    import builtins as _builtins
    _builtins.Dict  = getattr(_typing, "Dict")
    _builtins.Tuple = getattr(_typing, "Tuple")
    _builtins.List  = getattr(_typing, "List")
    from parse_fasta_py import parse_fasta as _parse_fasta_raw

    def parse_fasta(fasta_content: str):
        d = _parse_fasta_raw(fasta_content)
        out = {}
        for k, v in d.items():
            out[k] = v[0] if isinstance(v, tuple) else v
        return out
except Exception:
    def parse_fasta(fasta_content: str):
        out = {}
        cur = None
        seq_parts = []
        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(seq_parts)
                header = line[1:]
                cur = header.split()[0]
                seq_parts = []
            else:
                seq_parts.append(line.replace(" ", ""))
        if cur is not None:
            out[cur] = "".join(seq_parts)
        return out

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

def get_chain_segments(row):
    seq = row['sequence']
    stoich = row.get('stoichiometry', '')
    all_seq = row.get('all_sequences', '')

    if pd.isna(stoich) or pd.isna(all_seq) or str(stoich).strip()=="" or str(all_seq).strip()=="":
        return [(0, len(seq))]

    try:
        chain_dict = parse_fasta(all_seq)
        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):
                L = len(base)
                segs.append((pos, pos + L))
                pos += L
        if pos != len(seq):
            return [(0, len(seq))]
        return segs
    except Exception:
        return [(0, len(seq))]

def build_segments_map(df):
    seg_map = {}
    stoich_map = {}
    for _, r in df.iterrows():
        tid = r['target_id']
        seg_map[tid] = get_chain_segments(r)
        stoich_map[tid] = str(r.get('stoichiometry', '') if not pd.isna(r.get('stoichiometry', '')) else '')
    return seg_map, stoich_map

train_segs_map, train_stoich_map = build_segments_map(train_seqs)
test_segs_map,  test_stoich_map  = build_segments_map(test_seqs)

def process_labels(labels_df):
    coords_dict = {}
    prefixes = labels_df['ID'].str.rsplit('_', n=1).str[0]
    for id_prefix, group in labels_df.groupby(prefixes):
        coords_dict[id_prefix] = group.sort_values('resid')[['x_1', 'y_1', 'z_1']].values
    return coords_dict

train_coords_dict = process_labels(train_labels)

from Bio.Align import PairwiseAligner

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=5):
    """Enhanced template search with better similarity metrics"""
    similar_seqs = []
    
    for _, row in train_seqs_df.iterrows():
        target_id, train_seq = row['target_id'], row['sequence']
        if target_id not in train_coords_dict: continue
        
        # More lenient length filter for better template coverage
        len_ratio = abs(len(train_seq) - len(query_seq)) / max(len(train_seq), len(query_seq))
        if len_ratio > 0.4: continue
        
        raw_score = aligner.score(query_seq, train_seq)
        normalized_score = raw_score / (2 * min(len(query_seq), len(train_seq)))
        
        # Bonus for exact length match
        if len(train_seq) == len(query_seq):
            normalized_score *= 1.05
        
        similar_seqs.append((target_id, train_seq, normalized_score, train_coords_dict[target_id]))
    
    similar_seqs.sort(key=lambda x: x[2], reverse=True)
    return similar_seqs[:top_n]

def adapt_template_to_query(query_seq, template_seq, template_coords):
    """Enhanced adaptation with better gap handling"""
    alignment = next(iter(aligner.align(query_seq, template_seq)))
    new_coords = np.full((len(query_seq), 3), np.nan)
    
    for (q_start, q_end), (t_start, t_end) in zip(*alignment.aligned):
        t_chunk = template_coords[t_start:t_end]
        if len(t_chunk) == (q_end - q_start):
            new_coords[q_start:q_end] = t_chunk

    # Enhanced interpolation with cubic spline-like behavior
    for i in range(len(new_coords)):
        if np.isnan(new_coords[i, 0]):
            prev_v = next((j for j in range(i-1, -1, -1) if not np.isnan(new_coords[j, 0])), -1)
            next_v = next((j for j in range(i+1, len(new_coords)) if not np.isnan(new_coords[j, 0])), -1)
            
            if prev_v >= 0 and next_v >= 0:
                # Smooth interpolation
                gap_size = next_v - prev_v
                w = (i - prev_v) / gap_size
                # Add slight curvature for more realistic gaps
                curve = 0.2 * np.sin(w * np.pi)
                new_coords[i] = (1-w)*new_coords[prev_v] + w*new_coords[next_v]
                # Slight perpendicular offset for natural curve
                direction = new_coords[next_v] - new_coords[prev_v]
                perp = np.array([-direction[1], direction[0], direction[2]]) if abs(direction[0]) > abs(direction[1]) else np.array([direction[2], direction[1], -direction[0]])
                perp = perp / (np.linalg.norm(perp) + 1e-12)
                new_coords[i] += curve * perp
            elif prev_v >= 0: 
                new_coords[i] = new_coords[prev_v] + [5.95, 0, 0]
            elif next_v >= 0: 
                new_coords[i] = new_coords[next_v] - [5.95, 0, 0]
            else: 
                new_coords[i] = [i*5.95, 0, 0]
            
    return np.nan_to_num(new_coords)

def adaptive_rna_constraints(coordinates, target_id, confidence=1.0, passes=3):
    """Enhanced constraint refinement with better geometry"""
    coords = coordinates.copy()
    segments = test_segs_map.get(target_id, [(0, len(coords))])

    # Adaptive strength based on confidence
    strength = 0.85 * (1.0 - min(confidence, 0.98))
    strength = max(strength, 0.03)

    for pass_num in range(passes):
        # Reduce strength in later passes for convergence
        current_strength = strength * (1.0 - 0.15 * pass_num / passes)
        
        for (s, e) in segments:
            X = coords[s:e]
            L = e - s
            if L < 3:
                coords[s:e] = X
                continue

            # (1) Bond length constraint: i,i+1 to ~5.95Å
            d = X[1:] - X[:-1]
            dist = np.linalg.norm(d, axis=1) + 1e-6
            target = 5.95
            scale = (target - dist) / dist
            adj = (d * scale[:, None]) * (0.25 * current_strength)
            X[:-1] -= adj
            X[1:]  += adj

            # (2) i,i+2 constraint to ~10.2Å
            if L >= 3:
                d2 = X[2:] - X[:-2]
                dist2 = np.linalg.norm(d2, axis=1) + 1e-6
                target2 = 10.2
                scale2 = (target2 - dist2) / dist2
                adj2 = (d2 * scale2[:, None]) * (0.12 * current_strength)
                X[:-2] -= adj2
                X[2:]  += adj2

            # (3) i,i+3 constraint for better stacking geometry (~14.5Å)
            if L >= 4:
                d3 = X[3:] - X[:-3]
                dist3 = np.linalg.norm(d3, axis=1) + 1e-6
                target3 = 14.5
                scale3 = (target3 - dist3) / dist3
                adj3 = (d3 * scale3[:, None]) * (0.08 * current_strength)
                X[:-3] -= adj3
                X[3:]  += adj3

            # (4) Laplacian smoothing
            if L >= 3:
                lap = 0.5 * (X[:-2] + X[2:]) - X[1:-1]
                X[1:-1] += (0.07 * current_strength) * lap

            # (5) Self-avoidance with better collision handling
            if L >= 20:
                k = min(L, 180) if L > 250 else L
                if k < L:
                    idx = np.linspace(0, L - 1, k).astype(int)
                else:
                    idx = np.arange(L)

                P = X[idx]
                diff = P[:, None, :] - P[None, :, :]
                distm = np.linalg.norm(diff, axis=2) + 1e-6
                sep = np.abs(idx[:, None] - idx[None, :])

                # More aggressive clash resolution
                mask = (sep > 2) & (distm < 3.5)
                if np.any(mask):
                    force = (3.5 - distm) / distm
                    vec = (diff * force[:, :, None] * mask[:, :, None]).sum(axis=1)
                    X[idx] += (0.020 * current_strength) * vec

            coords[s:e] = X

    return coords

def _rotmat(axis, ang):
    axis = np.asarray(axis, float)
    axis = axis / (np.linalg.norm(axis) + 1e-12)
    x, y, z = axis
    c, s = np.cos(ang), np.sin(ang)
    C = 1.0 - c
    return np.array([
        [c + x*x*C,     x*y*C - z*s, x*z*C + y*s],
        [y*x*C + z*s,   c + y*y*C,   y*z*C - x*s],
        [z*x*C - y*s,   z*y*C + x*s, c + z*z*C]
    ], dtype=float)

def apply_hinge(coords, seg, rng, max_angle_deg=25):
    s, e = seg
    L = e - s
    if L < 30:
        return coords
    pivot = s + int(rng.integers(10, L - 10))
    axis = rng.normal(size=3)
    ang = np.deg2rad(float(rng.uniform(-max_angle_deg, max_angle_deg)))
    R = _rotmat(axis, ang)
    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, segments, rng, max_angle_deg=12, max_trans=1.5):
    X = coords.copy()
    global_center = X.mean(axis=0, keepdims=True)
    for (s, e) in segments:
        axis = rng.normal(size=3)
        ang = np.deg2rad(float(rng.uniform(-max_angle_deg, max_angle_deg)))
        R = _rotmat(axis, ang)
        shift = rng.normal(size=3)
        shift = shift / (np.linalg.norm(shift) + 1e-12) * float(rng.uniform(0.0, max_trans))
        c = X[s:e].mean(axis=0, keepdims=True)
        X[s:e] = (X[s:e] - c) @ R.T + c + shift
    X -= X.mean(axis=0, keepdims=True) - global_center
    return X

def smooth_wiggle(coords, segments, rng, amp=0.8):
    X = coords.copy()
    for (s, e) in segments:
        L = e - s
        if L < 20:
            continue
        n_ctrl = 6
        ctrl_x = np.linspace(0, L - 1, n_ctrl)
        ctrl_disp = rng.normal(0, amp, size=(n_ctrl, 3))
        t = np.arange(L)
        disp = np.vstack([np.interp(t, ctrl_x, ctrl_disp[:, k]) for k in range(3)]).T
        X[s:e] += disp
    return X

def local_twist(coords, segments, rng, max_twist_deg=18):
    """New: Add local twist variations"""
    X = coords.copy()
    for (s, e) in segments:
        L = e - s
        if L < 25:
            continue
        # Random twist regions
        n_twists = int(rng.integers(1, 4))
        for _ in range(n_twists):
            center = s + int(rng.integers(8, L - 8))
            radius = int(rng.integers(5, 12))
            start = max(s, center - radius)
            end = min(e, center + radius)
            
            axis = X[end-1] - X[start] if end > start+1 else rng.normal(size=3)
            ang = np.deg2rad(float(rng.uniform(-max_twist_deg, max_twist_deg)))
            R = _rotmat(axis, ang)
            pivot = X[center].copy()
            
            for i in range(start, end):
                w = 1.0 - abs(i - center) / radius
                R_blend = w * R + (1-w) * np.eye(3)
                X[i] = (X[i] - pivot) @ R_blend.T + pivot
    return X

def ensemble_blend(predictions, weights=None):
    """Blend multiple predictions with optional weights"""
    if weights is None:
        weights = np.ones(len(predictions)) / len(predictions)
    else:
        weights = np.array(weights) / np.sum(weights)
    
    blended = np.zeros_like(predictions[0])
    for pred, w in zip(predictions, weights):
        blended += w * pred
    return blended

def predict_rna_structures(row, train_seqs_df, train_coords_dict, n_predictions=5):
    tid = row['target_id']
    seq = row['sequence']
    assert set(seq).issubset(set("ACGU")), f"Non-ACGU in {tid}"

    segments = test_segs_map.get(tid, [(0, len(seq))])

    # Larger candidate pool for better diversity
    cands = find_similar_sequences(query_seq=seq, train_seqs_df=train_seqs_df, 
                                   train_coords_dict=train_coords_dict, top_n=40)
    assert all(len(c[3]) == len(c[1]) for c in cands), "Template coords/seq length mismatch"
    predictions = []
    used = set()

    for i in range(n_predictions):
        seed = (abs(hash(tid)) + i * 10007) % (2**32)
        rng = np.random.default_rng(seed)

        if not cands:
            coords = np.zeros((len(seq), 3), dtype=float)
            for (s, e) in segments:
                for j in range(s+1, e):
                    coords[j] = coords[j-1] + [5.95, 0, 0]
            predictions.append(coords)
            continue

        # Enhanced template selection
        if i == 0:
            # Best template
            t_id, t_seq, sim, t_coords = cands[0]
        else:
            # Sample from top candidates with diversity
            K = min(15, len(cands))
            sims = np.array([cands[k][2] for k in range(K)], float)
            w = np.exp((sims - sims.max()) / 0.10)  # Temperature for diversity
            
            for k in range(K):
                if cands[k][0] in used:
                    w[k] *= 0.05  # Strong penalty for reuse
            
            w = w / (w.sum() + 1e-12)
            k = int(rng.choice(np.arange(K), p=w))
            t_id, t_seq, sim, t_coords = cands[k]

        used.add(t_id)

        adapted = adapt_template_to_query(query_seq=seq, template_seq=t_seq, template_coords=t_coords)

        # Diversification strategies
        if i == 0:
            # Keep best template pure
            X = adapted
        elif i == 1:
            # Mild Gaussian noise
            noise_scale = max(0.02, (0.50 - sim) * 0.08)
            X = adapted + rng.normal(0, noise_scale, adapted.shape)
        elif i == 2:
            # Hinge motion in longest chain
            longest = max(segments, key=lambda se: se[1] - se[0])
            X = apply_hinge(adapted, longest, rng, max_angle_deg=24)
        elif i == 3:
            # Inter-chain jitter for multi-chain structures
            if len(segments) > 1:
                X = jitter_chains(adapted, segments, rng, max_angle_deg=12, max_trans=1.2)
            else:
                X = smooth_wiggle(adapted, segments, rng, amp=0.9)
        else:
            # Local twist + smooth deformation
            X = local_twist(adapted, segments, rng, max_twist_deg=16)
            X = smooth_wiggle(X, segments, rng, amp=0.6)

        # Refinement with more passes for better geometry
        refined = adaptive_rna_constraints(X, tid, confidence=sim, passes=3)
        predictions.append(refined)

    return predictions

# Main prediction loop
all_predictions = []
start_time = time.time()

for idx, row in test_seqs.iterrows():
    if idx % 10 == 0: 
        print(f"Processing {idx}/{len(test_seqs)} | {time.time()-start_time:.1f}s")
    
    tid, seq = row['target_id'], row['sequence']
    preds = predict_rna_structures(row, train_seqs, train_coords_dict, n_predictions=5)
    
    for j in range(len(seq)):
        res = {'ID': f"{tid}_{j+1}", 'resname': seq[j], 'resid': j+1}
        for i in range(5):
            res[f'x_{i+1}'], res[f'y_{i+1}'], res[f'z_{i+1}'] = preds[i][j]
        all_predictions.append(res)

sub = pd.DataFrame(all_predictions)
cols = ['ID', 'resname', 'resid'] + [f'{c}_{i}' for i in range(1,6) for c in ['x','y','z']]

# Clip coordinates
coord_cols = [c for c in cols if c.startswith(('x_','y_','z_'))]
sub[coord_cols] = sub[coord_cols].clip(-999.999, 9999.999)

sub[cols].to_csv('submission.csv', index=False)
print(f"✓ submission.csv saved! Total time: {time.time()-start_time:.1f}s")

