# %% [code]
"""
Stanford RNA 3D Folding Part 2 - Baseline Submission
=====================================================

This script generates 3D RNA structure predictions using a simple
helical model with random rotations for ensemble diversity.
"""

import numpy as np
import pandas as pd
from pathlib import Path

# Constants
HELIX_RISE = 2.81
HELIX_TWIST_DEG = 32.7
HELIX_RADIUS = 10.0

print("=" * 60)
print("Stanford RNA 3D Folding Part 2 - Baseline")
print("=" * 60)


def simple_fold(sequence: str, seed: int = 42) -> np.ndarray:
    """Simple helical folding."""
    n = len(sequence)
    rng = np.random.RandomState(seed)

    coords = np.zeros((n, 3), dtype=np.float32)
    twist_rad = np.radians(HELIX_TWIST_DEG)

    for i in range(n):
        angle = i * twist_rad
        coords[i] = [
            HELIX_RADIUS * np.cos(angle),
            HELIX_RADIUS * np.sin(angle),
            i * HELIX_RISE
        ]

    coords += rng.randn(n, 3).astype(np.float32) * 0.5
    coords -= coords.mean(axis=0)

    return coords


def random_rotation(coords: np.ndarray, seed: int) -> np.ndarray:
    """Apply random rotation."""
    rng = np.random.RandomState(seed)
    alpha, beta, gamma = rng.uniform(0, 2*np.pi, 3)

    ca, sa = np.cos(alpha), np.sin(alpha)
    cb, sb = np.cos(beta), np.sin(beta)
    cg, sg = np.cos(gamma), np.sin(gamma)

    R = np.array([
        [cg*cb*ca - sg*sa, -cg*cb*sa - sg*ca, cg*sb],
        [sg*cb*ca + cg*sa, -sg*cb*sa + cg*ca, sg*sb],
        [-sb*ca, sb*sa, cb]
    ], dtype=np.float32)

    return (coords - coords.mean(axis=0)) @ R.T


def predict_ensemble(sequence: str, num_models: int = 5) -> list:
    """Generate ensemble of diverse predictions."""
    sequence = sequence.upper().replace('T', 'U')

    models = []
    for i in range(num_models):
        coords = simple_fold(sequence, seed=i * 7919)
        if i > 0:
            coords = random_rotation(coords, seed=i * 17)
        models.append(coords)

    return models


# Load test data
INPUT_DIR = Path("/kaggle/input/stanford-rna-3d-folding-2")
test_path = INPUT_DIR / "test_sequences.csv"

print(f"Looking for test data at: {test_path}")

if test_path.exists():
    test_df = pd.read_csv(test_path)
    print(f"Loaded {len(test_df)} test sequences")
else:
    print(f"Error: Test file not found at {test_path}")
    # List available files
    print(f"Available in input dir: {list(INPUT_DIR.glob('*'))}")
    raise FileNotFoundError(f"Test file not found at {test_path}")

# Generate predictions
rows = []

for idx, row in test_df.iterrows():
    target_id = row['target_id']
    sequence = str(row['sequence'])

    # Clean sequence
    if '\n' in sequence:
        sequence = sequence.split('\n')[0]
    sequence = ''.join(c for c in sequence.upper() if c in 'AUGC')

    if len(sequence) == 0:
        print(f"Warning: Empty sequence for {target_id}")
        continue

    if idx % 5 == 0:
        print(f"[{idx+1}/{len(test_df)}] {target_id}: {len(sequence)} nt")

    models = predict_ensemble(sequence, num_models=5)

    for resid, nucleotide in enumerate(sequence, start=1):
        row_data = {
            'ID': f"{target_id}_{resid}",
            'resname': nucleotide,
            'resid': resid
        }

        for m_idx, coords in enumerate(models):
            i = resid - 1
            if i < len(coords):
                row_data[f'x_{m_idx+1}'] = round(float(coords[i, 0]), 3)
                row_data[f'y_{m_idx+1}'] = round(float(coords[i, 1]), 3)
                row_data[f'z_{m_idx+1}'] = round(float(coords[i, 2]), 3)
            else:
                row_data[f'x_{m_idx+1}'] = 0.0
                row_data[f'y_{m_idx+1}'] = 0.0
                row_data[f'z_{m_idx+1}'] = 0.0

        rows.append(row_data)

# Create submission
cols = ['ID', 'resname', 'resid']
for i in range(1, 6):
    cols.extend([f'x_{i}', f'y_{i}', f'z_{i}'])

submission_df = pd.DataFrame(rows)[cols]

# Save
output_path = Path("/kaggle/working/submission.csv")
submission_df.to_csv(output_path, index=False)

print(f"\nSubmission saved to {output_path}")
print(f"Shape: {submission_df.shape}")
print("\nFirst 5 rows:")
print(submission_df.head())
print("\nDone!")
