#!/usr/bin/env python3
"""
TBM-Only Notebook for Stanford RNA 3D Folding Part 2
Based on 1st Place Solution from Part 1 by g john rao
"""

# %% [markdown]
# # Stanford RNA 3D Folding Part 2 - TBM Only Approach
# 
# This notebook implements a Template-Based Modeling (TBM) approach for RNA 3D structure prediction.
# Based on the winning solution from Part 1, TBM alone achieved 0.593 TM-score, outperforming hybrid approaches.
# 
# ## Pipeline Overview:
# 1. **SEARCH**: Find similar structures via sequence alignment
# 2. **ALIGN**: Global sequence alignment with RNA-optimized gap penalties  
# 3. **TRANSFER**: Copy 3D coordinates for matched positions
# 4. **GAP FILL**: Geometric backbone reconstruction
# 5. **REFINE**: Confidence-based adaptive optimization
# 6. **ENSEMBLE**: Generate 5 diverse predictions per target

# %% [code]
import os
import sys
import numpy as np
import pandas as pd
from Bio import Align, SeqIO
from Bio.Align import substitution_matrices
import warnings
warnings.filterwarnings('ignore')

# Set paths
DATA_PATH = "/kaggle/input/stanford-rna-3d-folding-2/"
OUTPUT_PATH = "/kaggle/working/"

# %% [markdown]
# ## 1. Data Loading and Preprocessing

# %% [code]
def load_sequences(path):
    """Load RNA sequences from CSV."""
    df = pd.read_csv(path)
    print(f"Loaded {len(df)} sequences from {path}")
    return df

def load_template_library():
    """
    Load template library from training data and PDB_RNA.
    For Kaggle submission, we use all available training structures as templates.
    """
    # Load training sequences with coordinates
    train_seqs = pd.read_csv(os.path.join(DATA_PATH, "train_sequences.csv"))
    train_labels = pd.read_csv(os.path.join(DATA_PATH, "train_labels.csv"))
    
    print(f"Training sequences: {len(train_seqs)}")
    print(f"Training labels: {len(train_labels)}")
    
    return train_seqs, train_labels

# Load test sequences
test_seqs = load_sequences(os.path.join(DATA_PATH, "test_sequences.csv"))
train_seqs, train_labels = load_template_library()

# Display sample
print("Test sequences sample:")
print(test_seqs.head())

# %% [markdown]
# ## 2. Sequence Similarity Scoring
# 
# Composite scoring function combining multiple similarity metrics:
# - Global alignment score (40%)
# - Local alignment score (30%)
# - Feature similarity - dinucleotide composition (20%)
# - K-mer similarity (10%)

# %% [code]
class RNASimilarityScorer:
    """Calculate composite similarity scores between RNA sequences."""
    
    # Important dinucleotides for RNA structure
    IMPORTANT_DINUCS = ['AU', 'UA', 'GC', 'CG', 'GU', 'UG', 'AA', 'UU', 'GG', 'CC']
    
    def __init__(self):
        # RNA-specific substitution matrix (NUC44)
        self.matrix = substitution_matrices.load("NUC.4.4")
        self.aligner = Align.PairwiseAligner()
        self.aligner.substitution_matrix = self.matrix
        self.aligner.open_gap_score = -10
        self.aligner.extend_gap_score = -1
    
    def global_alignment_score(self, seq1, seq2):
        """Compute global alignment score."""
        try:
            alignments = self.aligner.align(seq1, seq2)
            return alignments[0].score / max(len(seq1), len(seq2))
        except:
            return 0.0
    
    def local_alignment_score(self, seq1, seq2):
        """Compute local alignment score."""
        try:
            aligner_local = Align.PairwiseAligner()
            aligner_local.substitution_matrix = self.matrix
            aligner_local.mode = 'local'
            alignments = aligner_local.align(seq1, seq2)
            return alignments[0].score / min(len(seq1), len(seq2))
        except:
            return 0.0
    
    def extract_features(self, seq):
        """Extract sequence composition features."""
        seq = seq.upper().replace('T', 'U')
        features = []
        
        # 1. Single nucleotide frequencies
        for nuc in ['A', 'U', 'G', 'C']:
            freq = seq.count(nuc) / len(seq) if len(seq) > 0 else 0
            features.append(freq)
        
        # 2. GC content
        gc_content = (seq.count('G') + seq.count('C')) / len(seq) if len(seq) > 0 else 0
        features.append(gc_content)
        
        # 3. Dinucleotide frequencies (reduced set)
        for dinuc in self.IMPORTANT_DINUCS:
            count = sum(1 for i in range(len(seq)-1) if seq[i:i+2] == dinuc)
            freq = count / (len(seq) - 1) if len(seq) > 1 else 0
            features.append(freq)
        
        return np.array(features)
    
    def feature_similarity(self, seq1, seq2):
        """Compute feature-based similarity."""
        f1 = self.extract_features(seq1)
        f2 = self.extract_features(seq2)
        
        # Cosine similarity
        norm1 = np.linalg.norm(f1)
        norm2 = np.linalg.norm(f2)
        
        if norm1 == 0 or norm2 == 0:
            return 0.0
        
        return np.dot(f1, f2) / (norm1 * norm2)
    
    def kmer_similarity(self, seq1, seq2, k=3):
        """Compute k-mer similarity."""
        def get_kmers(seq, k):
            return set(seq[i:i+k] for i in range(len(seq)-k+1))
        
        kmers1 = get_kmers(seq1, k)
        kmers2 = get_kmers(seq2, k)
        
        if not kmers1 or not kmers2:
            return 0.0
        
        intersection = len(kmers1 & kmers2)
        union = len(kmers1 | kmers2)
        
        return intersection / union if union > 0 else 0.0
    
    def composite_score(self, seq1, seq2):
        """
        Calculate composite similarity score.
        Weights tuned based on Part 1 winning solution.
        """
        global_score = self.global_alignment_score(seq1, seq2)
        local_score = self.local_alignment_score(seq1, seq2)
        feature_score = self.feature_similarity(seq1, seq2)
        kmer_score = self.kmer_similarity(seq1, seq2)
        
        # Composite formula from winning solution
        composite = (
            0.4 * global_score +
            0.3 * local_score +
            0.2 * feature_score +
            0.1 * kmer_score
        )
        
        return composite, {
            'global': global_score,
            'local': local_score,
            'feature': feature_score,
            'kmer': kmer_score
        }

# Initialize scorer
scorer = RNASimilarityScorer()
print("Similarity scorer initialized")

# %% [markdown]
# ## 3. Template Database Construction

# %% [code]
class TemplateDatabase:
    """Database of RNA structural templates."""
    
    def __init__(self, sequences_df, labels_df):
        self.sequences = sequences_df.copy()
        self.labels = labels_df.copy()
        self._build_index()
    
    def _build_index(self):
        """Index templates by sequence for fast lookup."""
        self.seq_to_template = {}
        
        # Group labels by target_id
        grouped = self.labels.groupby('target_id')
        
        for target_id, group in grouped:
            # Get coordinates for this target
            coords = {}
            for _, row in group.iterrows():
                resid = int(row['resid'])
                coords[resid] = {
                    'resname': row['resname'],
                    'x': row.get('x_1', np.nan),
                    'y': row.get('y_1', np.nan),
                    'z': row.get('z_1', np.nan)
                }
            
            self.seq_to_template[target_id] = {
                'coords': coords,
                'length': len(group)
            }
    
    def get_template(self, target_id):
        """Get template by target_id."""
        return self.seq_to_template.get(target_id)
    
    def get_all_template_ids(self):
        """Get all available template IDs."""
        return list(self.seq_to_template.keys())

# Build template database
print("Building template database...")
template_db = TemplateDatabase(train_seqs, train_labels)
print(f"Templates available: {len(template_db.get_all_template_ids())}")

# %% [markdown]
# ## 4. Template Search

# %% [code]
class TemplateSearcher:
    """Search for similar templates given a query sequence."""
    
    def __init__(self, template_db, scorer, top_k=14):
        self.db = template_db
        self.scorer = scorer
        self.top_k = top_k
    
    def search(self, query_seq, query_id=None):
        """
        Find top-K similar templates for query sequence.
        Returns list of (template_id, score, details).
        """
        scores = []
        template_ids = self.db.get_all_template_ids()
        
        # Limit templates for speed
        for tmpl_id in template_ids[:min(500, len(template_ids))]:
            tmpl_data = self.db.get_template(tmpl_id)
            if tmpl_data is None:
                continue
            
            # Get template sequence from training data
            tmpl_seq_row = train_seqs[train_seqs['target_id'] == tmpl_id]
            if len(tmpl_seq_row) == 0:
                continue
            
            tmpl_seq = tmpl_seq_row.iloc[0]['sequence']
            
            # Calculate composite score
            score, details = self.scorer.composite_score(query_seq, tmpl_seq)
            
            if score > 0.3:  # Threshold for similarity
                scores.append((tmpl_id, score, details, tmpl_seq))
        
        # Sort by score descending
        scores.sort(key=lambda x: x[1], reverse=True)
        
        return scores[:self.top_k]

# Initialize searcher
searcher = TemplateSearcher(template_db, scorer, top_k=14)
print("Template searcher initialized")

# %% [markdown]
# ## 5. Coordinate Transfer and Gap Filling

# %% [code]
class StructureBuilder:
    """Build 3D structure from template alignment."""
    
    def __init__(self):
        self.aligner = Align.PairwiseAligner()
        matrix = substitution_matrices.load("NUC.4.4")
        self.aligner.substitution_matrix = matrix
        self.aligner.open_gap_score = -10
        self.aligner.extend_gap_score = -1
    
    def align_sequences(self, query_seq, tmpl_seq):
        """Align query to template sequence."""
        alignments = self.aligner.align(tmpl_seq, query_seq)
        alignment = alignments[0]
        
        # Extract aligned positions
        aligned_tmpl = alignment[0]
        aligned_query = alignment[1]
        
        # Build mapping from query position to template position
        mapping = {}
        tmpl_pos = 0
        query_pos = 0
        
        for i in range(len(aligned_tmpl)):
            tmpl_char = aligned_tmpl[i]
            query_char = aligned_query[i]
            
            if tmpl_char != '-':
                tmpl_pos += 1
            if query_char != '-':
                query_pos += 1
                if tmpl_char != '-':
                    mapping[query_pos] = tmpl_pos
        
        return mapping
    
    def transfer_coordinates(self, query_seq, template_id, template_db):
        """
        Transfer coordinates from template to query based on alignment.
        Returns array of (x, y, z) coordinates.
        """
        # Get template data
        tmpl_data = template_db.get_template(template_id)
        if tmpl_data is None:
            return None
        
        # Get template sequence
        tmpl_seq_row = train_seqs[train_seqs['target_id'] == template_id]
        if len(tmpl_seq_row) == 0:
            return None
        
        tmpl_seq = tmpl_seq_row.iloc[0]['sequence']
        
        # Align sequences
        mapping = self.align_sequences(query_seq, tmpl_seq)
        
        # Build coordinate array
        coords = np.zeros((len(query_seq), 3))
        coords.fill(np.nan)
        
        tmpl_coords = tmpl_data['coords']
        
        # Transfer coordinates
        for query_pos, tmpl_pos in mapping.items():
            if 1 <= tmpl_pos <= len(tmpl_coords):
                coord_data = tmpl_coords.get(tmpl_pos)
                if coord_data and not np.isnan(coord_data.get('x', np.nan)):
                    coords[query_pos - 1] = [coord_data['x'], coord_data['y'], coord_data['z']]
        
        return coords, mapping
    
    def fill_gaps(self, coords, query_seq):
        """
        Fill missing coordinates using geometric interpolation.
        Uses sinusoidal perturbations for compressed gaps and linear for normal gaps.
        """
        n = len(coords)
        filled_coords = coords.copy()
        
        # Find gaps (consecutive NaN positions)
        i = 0
        while i < n:
            if np.isnan(filled_coords[i, 0]):
                # Find gap end
                gap_start = i
                while i < n and np.isnan(filled_coords[i, 0]):
                    i += 1
                gap_end = i
                gap_len = gap_end - gap_start
                
                # Get anchor points
                if gap_start > 0 and gap_end < n:
                    # Gap in middle - interpolate between anchors
                    start_coord = filled_coords[gap_start - 1]
                    end_coord = filled_coords[gap_end]
                    
                    for j in range(gap_len):
                        t = (j + 1) / (gap_len + 1)
                        
                        # Direction vector
                        direction = end_coord - start_coord
                        distance = np.linalg.norm(direction)
                        
                        if distance > 0:
                            unit_direction = direction / distance
                            
                            # Base position
                            base_pos = start_coord + t * direction
                            
                            # Add sinusoidal perturbation for compressed gaps
                            if gap_len > 3:
                                # Perpendicular perturbation
                                perp = np.array([-unit_direction[1], unit_direction[0], unit_direction[2]])
                                if np.linalg.norm(perp) == 0:
                                    perp = np.array([0, 0, 1])
                                perp = perp / np.linalg.norm(perp)
                                
                                # Sinusoidal amplitude
                                amplitude = 2.0 * np.sin(np.pi * t)
                                base_pos += amplitude * perp
                            
                            filled_coords[gap_start + j] = base_pos
                        else:
                            # Fallback: place along line
                            filled_coords[gap_start + j] = start_coord + t * direction
                
                elif gap_start == 0 and gap_end < n:
                    # Gap at start - extend backward from first known
                    known_coord = filled_coords[gap_end]
                    next_coord = filled_coords[gap_end + 1] if gap_end + 1 < n else known_coord
                    direction = next_coord - known_coord
                    
                    for j in range(gap_len - 1, -1, -1):
                        offset = gap_end - j
                        filled_coords[j] = known_coord - offset * direction
                
                elif gap_end == n and gap_start > 0:
                    # Gap at end - extend forward from last known
                    known_coord = filled_coords[gap_start - 1]
                    prev_coord = filled_coords[gap_start - 2] if gap_start > 1 else known_coord
                    direction = known_coord - prev_coord
                    
                    for j in range(gap_len):
                        offset = j + 1
                        filled_coords[gap_start + j] = known_coord + offset * direction
            else:
                i += 1
        
        return filled_coords

# Initialize builder
builder = StructureBuilder()
print("Structure builder initialized")

# %% [markdown]
# ## 6. Structure Prediction Pipeline

# %% [code]
class TBMPredictor:
    """Template-Based Modeling predictor."""
    
    def __init__(self, searcher, builder, template_db):
        self.searcher = searcher
        self.builder = builder
        self.db = template_db
    
    def predict(self, target_id, query_seq, n_models=5):
        """
        Generate n_models predictions for a target sequence.
        Returns list of coordinate arrays.
        """
        # Search for templates
        templates = self.searcher.search(query_seq, target_id)
        
        if not templates:
            # No good templates - generate random structure as fallback
            return self._generate_fallback(query_seq, n_models)
        
        # Build structures from top templates
        structures = []
        
        for tmpl_id, score, details, tmpl_seq in templates[:n_models]:
            # Transfer coordinates
            result = self.builder.transfer_coordinates(query_seq, tmpl_id, self.db)
            
            if result is None:
                continue
            
            coords, mapping = result
            
            # Fill gaps
            filled_coords = self.builder.fill_gaps(coords, query_seq)
            
            # Validate structure
            if self._validate_structure(filled_coords):
                structures.append({
                    'coords': filled_coords,
                    'template_id': tmpl_id,
                    'score': score,
                    'mapping_ratio': len(mapping) / len(query_seq)
                })
        
        # If we don't have enough structures, duplicate best ones
        while len(structures) < n_models and structures:
            structures.append(structures[len(structures) % len(structures)].copy())
        
        if not structures:
            return self._generate_fallback(query_seq, n_models)
        
        # Return top n_models by score
        structures.sort(key=lambda x: x['score'], reverse=True)
        return [s['coords'] for s in structures[:n_models]]
    
    def _validate_structure(self, coords):
        """Check if structure is valid (no NaN values)."""
        return not np.any(np.isnan(coords))
    
    def _generate_fallback(self, query_seq, n_models):
        """Generate fallback random structures when no templates found."""
        structures = []
        n = len(query_seq)
        
        for _ in range(n_models):
            # Generate random walk with realistic bond lengths
            coords = np.zeros((n, 3))
            coords[0] = [0, 0, 0]
            
            for i in range(1, n):
                # Random direction with ~5.5 Angstrom step (typical C1'-C1' distance)
                direction = np.random.randn(3)
                direction = direction / np.linalg.norm(direction)
                coords[i] = coords[i-1] + direction * 5.5
            
            structures.append(coords)
        
        return structures

# Initialize predictor
predictor = TBMPredictor(searcher, builder, template_db)
print("TBM Predictor initialized")

# %% [markdown]
# ## 7. Generate Predictions

# %% [code]
def create_submission_row(target_id, resnames, resids, structures):
    """
    Create a submission row for a target.
    structures: list of 5 coordinate arrays
    """
    rows = []
    
    for i, (resname, resid) in enumerate(zip(resnames, resids)):
        row_id = f"{target_id}_{resid}"
        
        # Collect coordinates from all 5 models
        coords = []
        for model_idx in range(5):
            if model_idx < len(structures):
                coord = structures[model_idx][i]
                coords.extend([coord[0], coord[1], coord[2]])
            else:
                coords.extend([0.0, 0.0, 0.0])
        
        row = {
            'ID': row_id,
            'resname': resname,
            'resid': resid
        }
        
        # Add coordinates for each model
        for model_idx in range(5):
            idx = model_idx * 3
            row[f'x_{model_idx+1}'] = coords[idx]
            row[f'y_{model_idx+1}'] = coords[idx+1]
            row[f'z_{model_idx+1}'] = coords[idx+2]
        
        rows.append(row)
    
    return rows

# %% [code]
# Generate predictions for all test sequences
print("Generating predictions...")

all_predictions = []

for idx, row in test_seqs.iterrows():
    target_id = row['target_id']
    sequence = row['sequence']
    
    print(f"Processing {target_id} (length: {len(sequence)})")
    
    # Generate 5 predictions
    structures = predictor.predict(target_id, sequence, n_models=5)
    
    # Create submission rows
    resnames = list(sequence)
    resids = list(range(1, len(sequence) + 1))
    
    rows = create_submission_row(target_id, resnames, resids, structures)
    all_predictions.extend(rows)
    
    print(f"  Generated {len(structures)} models")

print(f"\nTotal prediction rows: {len(all_predictions)}")

# %% [markdown]
# ## 8. Save Submission

# %% [code]
# Create submission DataFrame
submission_df = pd.DataFrame(all_predictions)

# Reorder columns to match expected format
cols = ['ID', 'resname', 'resid']
for i in range(1, 6):
    cols.extend([f'x_{i}', f'y_{i}', f'z_{i}'])

submission_df = submission_df[cols]

# Save to CSV
output_file = os.path.join(OUTPUT_PATH, "submission.csv")
submission_df.to_csv(output_file, index=False)

print(f"Submission saved to: {output_file}")
print(f"Total rows: {len(submission_df)}")
print("\nSubmission preview:")
print(submission_df.head(10))

# %% [markdown]
# ## 9. Validation and Summary

# %% [code]
# Check submission format
print("\n" + "="*50)
print("Submission Validation:")
print("="*50)
print(f"  Rows: {len(submission_df)}")
print(f"  Columns: {list(submission_df.columns)}")
print(f"  Missing values: {submission_df.isnull().sum().sum()}")
print(f"  Unique targets: {submission_df['ID'].str.split('_').str[0].nunique()}")

# Check coordinate ranges
print("\nCoordinate ranges per model:")
for i in range(1, 6):
    x_col, y_col, z_col = f'x_{i}', f'y_{i}', f'z_{i}'
    print(f"  Model {i}:")
    print(f"    X: [{submission_df[x_col].min():.2f}, {submission_df[x_col].max():.2f}]")
    print(f"    Y: [{submission_df[y_col].min():.2f}, {submission_df[y_col].max():.2f}]")
    print(f"    Z: [{submission_df[z_col].min():.2f}, {submission_df[z_col].max():.2f}]")

print("\n" + "="*50)
print("Submission ready for upload!")
print("="*50)

# %% [code]
# Show sample predictions for first target
sample_target = test_seqs.iloc[0]['target_id']
sample_df = submission_df[submission_df['ID'].str.startswith(sample_target)]
print(f"\nSample predictions for {sample_target}:")
print(sample_df.head())
