#!/usr/bin/env python3
"""
Protenix Fine-Tuned Inference for Stanford RNA 3D Folding Part 2
================================================================
Internet OFF, GPU ON (T4 16GB or P100 16GB)
Loads fine-tuned Protenix checkpoint from Kaggle Dataset.

Based on 7th place Part 1 solution inference approach:
- Diffusion steps: 200
- Trunk recycling: 10 (N_cycle)  
- Generate 10 predictions per sequence (5 from each model or 10 from 1 model)
- K-medoids structure selection (select 5 diverse from 10)
- Output submission.csv

This is designed as a Kaggle notebook. Key assumptions:
- Fine-tuned model checkpoint uploaded as Kaggle Dataset
- Competition data available at /kaggle/input/stanford-rna-3d-folding-2/
- MSA data available at /kaggle/input/stanford-rna-3d-folding-2/MSA/
"""

import os
import sys
import json
import time
import subprocess
import warnings
import gc
warnings.filterwarnings('ignore')

# ============================================================
# CONFIGURATION
# ============================================================

# Kaggle paths
COMPETITION_DATA = "/kaggle/input/stanford-rna-3d-folding-2"
FINETUNED_MODEL = "/kaggle/input/protenix-rna-finetuned-v2"  # Kaggle Dataset with checkpoint
OUTPUT_DIR = "/kaggle/working"

# Inference params (from 7th place)
N_DIFFUSION_STEPS = 200  # More steps = better quality
N_CYCLE = 10  # Trunk recycling iterations
N_SAMPLE = 5  # Predictions per run (can do 2 runs with different seeds = 10 total)
SEEDS = [101, 102]  # Two seeds for diversity
DTYPE = "bf16"

# Sequence length handling
MAX_SEQ_LEN_GPU = 850  # Max tokens for T4/P100 GPU
LONG_SEQ_THRESHOLD = 850

# ============================================================
# INSTALLATION
# ============================================================

def install_protenix():
    """Install Protenix from the fine-tuned model package or PyPI."""
    print("Installing Protenix...")
    
    # Check if Protenix code is bundled with the dataset
    protenix_src = os.path.join(FINETUNED_MODEL, "Protenix-RNA-Kaggle")
    if os.path.exists(protenix_src):
        subprocess.run([sys.executable, "-m", "pip", "install", "-e", protenix_src], 
                      check=True, capture_output=True)
    else:
        subprocess.run([sys.executable, "-m", "pip", "install", "protenix"],
                      check=True, capture_output=True)
    
    # Install additional deps
    subprocess.run([
        sys.executable, "-m", "pip", "install", 
        "scikit-learn-extra", "biopython==1.83", "biotite==1.0.1"
    ], check=True, capture_output=True)
    
    print("  Protenix installed successfully")


# ============================================================
# DATA PREPARATION
# ============================================================

def prepare_input_json(test_csv_path, msa_dir=None):
    """
    Convert test_sequences.csv to Protenix input JSON format.
    
    Returns list of dicts, each containing:
    {
        "sequences": [{"rnaSequence": {"sequence": "...", "count": 1}}],
        "name": "target_id"
    }
    """
    import pandas as pd
    
    df = pd.read_csv(test_csv_path)
    inputs = []
    
    for _, row in df.iterrows():
        target_id = row['target_id']
        sequence = row['sequence']
        seq_len = len(sequence)
        
        entry = {
            "sequences": [{
                "rnaSequence": {
                    "sequence": sequence,
                    "count": 1,
                }
            }],
            "name": target_id
        }
        
        # Add MSA if available
        if msa_dir and os.path.exists(os.path.join(msa_dir, f"{target_id}.MSA.fasta")):
            entry["sequences"][0]["rnaSequence"]["msa"] = {
                "precomputed_msa_dir": msa_dir,
                "pairing_db": ""
            }
        
        inputs.append(entry)
    
    return inputs, df


# ============================================================
# INFERENCE
# ============================================================

def run_protenix_inference(input_json, checkpoint_path, output_dir, 
                           seed=101, use_msa=False, n_sample=5, n_step=200):
    """
    Run Protenix inference on a single sequence.
    """
    # Write input JSON
    json_path = os.path.join(output_dir, "input.json")
    with open(json_path, 'w') as f:
        json.dump([input_json], f)
    
    cmd = [
        "protenix", "predict",
        "--input", json_path,
        "--out_dir", output_dir,
        "--seeds", str(seed),
        "--use_msa", str(use_msa).lower(),
        "--load_checkpoint_path", checkpoint_path,
        "--dtype", DTYPE,
        "--sample_diffusion.N_step", str(n_step),
        "--sample_diffusion.N_sample", str(n_sample),
        "--model.N_cycle", str(N_CYCLE),
    ]
    
    result = subprocess.run(cmd, capture_output=True, text=True, timeout=7200)
    return result.returncode == 0


def extract_c1_prime_coords(cif_path):
    """Extract C1' atom coordinates from a CIF file."""
    from Bio.PDB import MMCIFParser
    
    parser = MMCIFParser(QUIET=True)
    structure = parser.get_structure('pred', cif_path)
    
    coords = []
    for model in structure:
        for chain in model:
            for residue in chain:
                if residue.get_resname() in ['A', 'U', 'G', 'C']:
                    for atom in residue:
                        if atom.get_name() == "C1'":
                            coords.append(atom.get_coord().tolist())
    
    return coords


# ============================================================
# K-MEDOIDS STRUCTURE SELECTION
# ============================================================

def compute_tm_score_matrix(all_coords):
    """
    Compute pairwise TM-scores between predictions.
    Uses simplified TM-score based on C1' atoms.
    """
    import numpy as np
    
    n = len(all_coords)
    tm_matrix = np.zeros((n, n))
    
    for i in range(n):
        for j in range(i, n):
            if i == j:
                tm_matrix[i, j] = 1.0
            else:
                tm = simplified_tm_score(
                    np.array(all_coords[i]), 
                    np.array(all_coords[j])
                )
                tm_matrix[i, j] = tm
                tm_matrix[j, i] = tm
    
    return tm_matrix


def simplified_tm_score(coords1, coords2):
    """
    Simplified TM-score computation using C1' atoms.
    Based on Zhang & Skolnick (2004).
    """
    import numpy as np
    from scipy.spatial.transform import Rotation
    
    if len(coords1) != len(coords2):
        return 0.0
    
    L = len(coords1)
    if L == 0:
        return 0.0
    
    d0 = 1.24 * (L - 15) ** (1.0/3.0) - 1.8
    d0 = max(d0, 0.5)
    
    # Superpose using Kabsch algorithm
    c1 = coords1 - coords1.mean(axis=0)
    c2 = coords2 - coords2.mean(axis=0)
    
    H = c1.T @ c2
    U, S, Vt = np.linalg.svd(H)
    d = np.sign(np.linalg.det(Vt.T @ U.T))
    D = np.diag([1, 1, d])
    R = Vt.T @ D @ U.T
    
    c2_aligned = (R @ c2.T).T
    
    distances = np.sqrt(np.sum((c1 - c2_aligned) ** 2, axis=1))
    tm = np.sum(1.0 / (1.0 + (distances / d0) ** 2)) / L
    
    return tm


def select_diverse_structures(all_coords, k=5):
    """
    Select k diverse structures using k-medoids clustering.
    Distance metric: 1 - TM-score
    
    Returns indices of selected structures.
    """
    import numpy as np
    
    n = len(all_coords)
    if n <= k:
        return list(range(n))
    
    # Compute TM-score matrix
    tm_matrix = compute_tm_score_matrix(all_coords)
    distance_matrix = 1 - tm_matrix
    
    # K-medoids clustering
    try:
        from sklearn_extra.cluster import KMedoids
        kmedoids = KMedoids(n_clusters=k, metric='precomputed', random_state=42)
        kmedoids.fit(distance_matrix)
        return kmedoids.medoid_indices_.tolist()
    except ImportError:
        # Fallback: greedy farthest-point selection
        print("  sklearn_extra not available, using greedy selection")
        selected = [0]
        for _ in range(k - 1):
            min_dists = np.min(distance_matrix[selected], axis=0)
            min_dists[selected] = -1  # Exclude already selected
            next_idx = np.argmax(min_dists)
            selected.append(next_idx)
        return selected


# ============================================================
# SUBMISSION GENERATION
# ============================================================

def generate_submission(predictions, test_df, output_path):
    """
    Generate submission.csv in the required format.
    
    predictions: dict of {target_id: [list of 5 coordinate arrays]}
    Each coordinate array has shape (seq_len, 3) for C1' atoms.
    """
    import pandas as pd
    import numpy as np
    
    rows = []
    
    for _, row in test_df.iterrows():
        target_id = row['target_id']
        sequence = row['sequence']
        seq_len = len(sequence)
        
        if target_id not in predictions:
            print(f"  WARNING: No predictions for {target_id}, using zeros")
            for resid in range(1, seq_len + 1):
                resname = sequence[resid - 1]
                row_id = f"{target_id}_{resid}"
                row_data = {"ID": row_id, "resname": resname, "resid": resid}
                for i in range(1, 6):  # 5 predictions
                    row_data[f"x_{i}"] = 0.0
                    row_data[f"y_{i}"] = 0.0
                    row_data[f"z_{i}"] = 0.0
                rows.append(row_data)
            continue
        
        pred_coords = predictions[target_id]  # List of 5 coordinate arrays
        
        for resid in range(1, seq_len + 1):
            resname = sequence[resid - 1]
            row_id = f"{target_id}_{resid}"
            row_data = {"ID": row_id, "resname": resname, "resid": resid}
            
            for i in range(len(pred_coords)):
                idx = resid - 1
                if idx < len(pred_coords[i]):
                    row_data[f"x_{i+1}"] = pred_coords[i][idx][0]
                    row_data[f"y_{i+1}"] = pred_coords[i][idx][1]
                    row_data[f"z_{i+1}"] = pred_coords[i][idx][2]
                else:
                    # End-same padding (repeat last coordinate)
                    last_coord = pred_coords[i][-1]
                    row_data[f"x_{i+1}"] = last_coord[0]
                    row_data[f"y_{i+1}"] = last_coord[1]
                    row_data[f"z_{i+1}"] = last_coord[2]
            
            # If fewer than 5 predictions, duplicate last prediction
            for i in range(len(pred_coords), 5):
                idx = resid - 1
                src = pred_coords[-1]
                if idx < len(src):
                    row_data[f"x_{i+1}"] = src[idx][0]
                    row_data[f"y_{i+1}"] = src[idx][1]
                    row_data[f"z_{i+1}"] = src[idx][2]
                else:
                    last_coord = src[-1]
                    row_data[f"x_{i+1}"] = last_coord[0]
                    row_data[f"y_{i+1}"] = last_coord[1]
                    row_data[f"z_{i+1}"] = last_coord[2]
            
            rows.append(row_data)
    
    submission = pd.DataFrame(rows)
    submission.to_csv(output_path, index=False)
    print(f"Submission saved to {output_path}")
    print(f"  Shape: {submission.shape}")
    
    return submission


# ============================================================
# MAIN PIPELINE
# ============================================================

def main():
    import pandas as pd
    import numpy as np
    
    print("=" * 60)
    print("Protenix Fine-Tuned Inference")
    print("Stanford RNA 3D Folding Part 2")
    print("=" * 60)
    
    # Find checkpoint
    checkpoint_path = None
    for f in os.listdir(FINETUNED_MODEL):
        if f.endswith('_ema_0.995.pt') or f.endswith('.pt'):
            checkpoint_path = os.path.join(FINETUNED_MODEL, f)
            break
    
    # Also check subdirectories
    if checkpoint_path is None:
        for root, dirs, files in os.walk(FINETUNED_MODEL):
            for f in files:
                if f.endswith('_ema_0.995.pt'):
                    checkpoint_path = os.path.join(root, f)
                    break
            if checkpoint_path:
                break
    
    if checkpoint_path is None:
        # Fall back to any .pt file
        for root, dirs, files in os.walk(FINETUNED_MODEL):
            for f in files:
                if f.endswith('.pt'):
                    checkpoint_path = os.path.join(root, f)
                    break
            if checkpoint_path:
                break
    
    print(f"Checkpoint: {checkpoint_path}")
    assert checkpoint_path is not None, "No checkpoint found!"
    
    # Load test data
    test_csv = os.path.join(COMPETITION_DATA, "test_sequences.csv")
    test_df = pd.read_csv(test_csv)
    msa_dir = os.path.join(COMPETITION_DATA, "MSA")
    
    print(f"Test sequences: {len(test_df)}")
    
    # Prepare inputs
    inputs, _ = prepare_input_json(test_csv, msa_dir)
    
    # Run inference for each sequence
    all_predictions = {}
    
    for idx, (inp, (_, row)) in enumerate(zip(inputs, test_df.iterrows())):
        target_id = row['target_id']
        seq_len = len(row['sequence'])
        
        print(f"\n[{idx+1}/{len(inputs)}] {target_id} (len={seq_len})")
        
        if seq_len > MAX_SEQ_LEN_GPU:
            print(f"  Skipping (too long for GPU, will use RNAPro or baseline)")
            continue
        
        target_dir = os.path.join(OUTPUT_DIR, "predictions", target_id)
        os.makedirs(target_dir, exist_ok=True)
        
        all_coords = []
        
        for seed in SEEDS:
            print(f"  Running seed={seed}...")
            t0 = time.time()
            
            success = run_protenix_inference(
                input_json=inp,
                checkpoint_path=checkpoint_path,
                output_dir=target_dir,
                seed=seed,
                use_msa=os.path.exists(os.path.join(msa_dir, f"{target_id}.MSA.fasta")),
                n_sample=N_SAMPLE,
                n_step=N_DIFFUSION_STEPS,
            )
            
            elapsed = time.time() - t0
            print(f"  Seed {seed}: {'OK' if success else 'FAILED'} ({elapsed:.0f}s)")
            
            # Collect output CIF files
            for cif_file in sorted(os.listdir(target_dir)):
                if cif_file.endswith('.cif') and f"seed{seed}" in cif_file:
                    cif_path = os.path.join(target_dir, cif_file)
                    try:
                        coords = extract_c1_prime_coords(cif_path)
                        if len(coords) > 0:
                            all_coords.append(coords)
                    except Exception as e:
                        print(f"    Error reading {cif_file}: {e}")
            
            # Clean up to save memory
            gc.collect()
        
        print(f"  Total predictions: {len(all_coords)}")
        
        if len(all_coords) > 0:
            # Select 5 diverse structures using k-medoids
            if len(all_coords) > 5:
                selected_idx = select_diverse_structures(all_coords, k=5)
                print(f"  K-medoids selected: {selected_idx}")
                selected_coords = [all_coords[i] for i in selected_idx]
            else:
                selected_coords = all_coords
            
            all_predictions[target_id] = selected_coords
        else:
            print(f"  WARNING: No valid predictions for {target_id}")
    
    # Generate submission
    print("\n" + "=" * 60)
    print("Generating submission.csv")
    print("=" * 60)
    
    submission_path = os.path.join(OUTPUT_DIR, "submission.csv")
    generate_submission(all_predictions, test_df, submission_path)
    
    print("\nDone!")


if __name__ == "__main__":
    install_protenix()
    main()
