"""
Stanford RNA 3D Folding Part 2 - Template-Based Baseline
Adapted from public notebook by parthenos (nihilisticneuralnet)

This approach:
1. Finds similar sequences in the training data
2. Uses their 3D coordinates as templates
3. Adapts templates to the query sequence via alignment
4. Applies RNA geometric constraints
"""

import pandas as pd
import numpy as np
from scipy.spatial.transform import Rotation as R
import random
from Bio import pairwise2
from Bio.Seq import Seq
import time
from scipy.spatial import distance_matrix
import warnings
warnings.filterwarnings('ignore')

print("Loading data...")
train_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/train_sequences.csv')
valid_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/validation_sequences.csv')
test_seqs = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv')
train_labels = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/train_labels.csv')
valid_labels = pd.read_csv('/kaggle/input/stanford-rna-3d-folding-2/validation_labels.csv')

print(f"Train sequences: {len(train_seqs)}")
print(f"Validation sequences: {len(valid_seqs)}")
print(f"Test sequences: {len(test_seqs)}")

def process_labels(labels_df):
    """Process labels dataframe to create a dictionary mapping target_id to coordinates"""
    coords_dict = {}
    for id_prefix, group in labels_df.groupby(lambda x: labels_df['ID'][x].rsplit('_', 1)[0]):
        coords = []
        for _, row in group.sort_values('resid').iterrows():
            coords.append([row['x_1'], row['y_1'], row['z_1']])
        coords_dict[id_prefix] = np.array(coords)
    return coords_dict

print("Processing training labels...")
train_coords_dict = process_labels(train_labels)
valid_coords_dict = process_labels(valid_labels)
print(f"Processed {len(train_coords_dict)} training and {len(valid_coords_dict)} validation structures")

def find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, temporal_cutoff=None, top_n=5):
    """Find sequences in training data similar to query sequence"""
    similar_seqs = []
    query_seq_obj = Seq(query_seq)
    
    if temporal_cutoff:
        filtered_train_seqs = train_seqs_df[train_seqs_df['temporal_cutoff'] < temporal_cutoff]
    else:
        filtered_train_seqs = train_seqs_df
    
    for _, row in filtered_train_seqs.iterrows():
        target_id = row['target_id']
        train_seq = row['sequence']
        
        if target_id not in train_coords_dict:
            continue
        if abs(len(train_seq) - len(query_seq)) / max(len(train_seq), len(query_seq)) > 0.5:
            continue
        
        alignments = pairwise2.align.globalms(query_seq_obj, train_seq, 2, -1, -10, -0.5, one_alignment_only=True)
        
        if alignments:
            alignment = alignments[0]
            similarity_score = alignment.score / (2 * min(len(query_seq), len(train_seq)))
            similar_seqs.append((target_id, train_seq, similarity_score, train_coords_dict[target_id]))
    
    similar_seqs.sort(key=lambda x: x[2], reverse=True)
    return similar_seqs[:top_n]

def adaptive_rna_constraints(coordinates, sequence, confidence=1.0):
    """Apply RNA geometric constraints with adaptive strength"""
    refined_coords = coordinates.copy()
    n_residues = len(sequence)
    constraint_strength = 0.8 * (1.0 - min(confidence, 0.8))
    
    seq_min_dist = 5.5
    seq_max_dist = 6.5
    
    for i in range(n_residues - 1):
        current_pos = refined_coords[i]
        next_pos = refined_coords[i+1]
        current_dist = np.linalg.norm(next_pos - current_pos)
        
        if current_dist < seq_min_dist or current_dist > seq_max_dist:
            target_dist = (seq_min_dist + seq_max_dist) / 2
            direction = next_pos - current_pos
            direction = direction / (np.linalg.norm(direction) + 1e-10)
            adjustment = (target_dist - current_dist) * constraint_strength
            refined_coords[i+1] = current_pos + direction * (current_dist + adjustment)
    
    return refined_coords

def adapt_template_to_query(query_seq, template_seq, template_coords, alignment=None):
    """Adapt template coordinates to fit query sequence"""
    if alignment is None:
        query_seq_obj = Seq(query_seq)
        template_seq_obj = Seq(template_seq)
        alignments = pairwise2.align.globalms(query_seq_obj, template_seq_obj, 2, -1, -10, -0.5, one_alignment_only=True)
        
        if not alignments:
            return generate_basic_structure(query_seq)
        alignment = alignments[0]
    
    aligned_query = alignment.seqA
    aligned_template = alignment.seqB
    
    query_coords = np.zeros((len(query_seq), 3))
    query_coords.fill(np.nan)
    
    query_idx = 0
    template_idx = 0
    
    for i in range(len(aligned_query)):
        query_char = aligned_query[i]
        template_char = aligned_template[i]
        
        if query_char != '-' and template_char != '-':
            if template_idx < len(template_coords):
                query_coords[query_idx] = template_coords[template_idx]
            template_idx += 1
            query_idx += 1
        elif query_char != '-' and template_char == '-':
            query_idx += 1
        elif query_char == '-' and template_char != '-':
            template_idx += 1
    
    # Fill NaN positions
    typical_step = 4.0
    for i in range(len(query_coords)):
        if np.isnan(query_coords[i, 0]):
            prev_valid = -1
            for j in range(i-1, -1, -1):
                if not np.isnan(query_coords[j, 0]):
                    prev_valid = j
                    break
            
            next_valid = -1
            for j in range(i+1, len(query_coords)):
                if not np.isnan(query_coords[j, 0]):
                    next_valid = j
                    break
            
            if prev_valid >= 0 and next_valid >= 0:
                weight = (i - prev_valid) / (next_valid - prev_valid)
                query_coords[i] = (1 - weight) * query_coords[prev_valid] + weight * query_coords[next_valid]
            elif prev_valid >= 0:
                direction = np.random.normal(0, 1, 3)
                direction = direction / (np.linalg.norm(direction) + 1e-10) * typical_step
                query_coords[i] = query_coords[prev_valid] + direction
            elif i == 0 and next_valid >= 0:
                for j in range(next_valid-1, -1, -1):
                    direction = np.random.normal(0, 1, 3)
                    direction = direction / (np.linalg.norm(direction) + 1e-10) * typical_step
                    query_coords[j] = query_coords[j+1] - direction
            else:
                angle = i * 0.6
                query_coords[i] = [10.0 * np.cos(angle), 10.0 * np.sin(angle), i * 2.5]
    
    if np.isnan(query_coords).any():
        query_coords = np.nan_to_num(query_coords)
    
    return query_coords

def generate_basic_structure(sequence):
    """Generate a simple helical structure"""
    n_residues = len(sequence)
    coordinates = np.zeros((n_residues, 3))
    radius = 10.0
    rise_per_residue = 2.5
    angle_per_residue = 0.5
    
    for i in range(n_residues):
        angle = i * angle_per_residue
        coordinates[i] = [radius * np.cos(angle), radius * np.sin(angle), i * rise_per_residue]
    
    return coordinates

def generate_rna_structure(sequence, seed=None):
    """Generate a more realistic RNA structure"""
    if seed is not None:
        np.random.seed(seed)
        random.seed(seed)
    
    n_residues = len(sequence)
    coordinates = np.zeros((n_residues, 3))
    
    for i in range(min(3, n_residues)):
        angle = i * 0.6
        coordinates[i] = [10.0 * np.cos(angle), 10.0 * np.sin(angle), i * 2.5]
    
    current_direction = np.array([0.0, 0.0, 1.0])
    
    for i in range(3, n_residues):
        has_pair = False
        pair_idx = -1
        complementary = {'G': 'C', 'C': 'G', 'A': 'U', 'U': 'A'}
        current_base = sequence[i]
        
        window_size = min(i, 15)
        for j in range(i-window_size, i):
            if j >= 0 and sequence[j] == complementary.get(current_base, 'X'):
                has_pair = True
                pair_idx = j
                break
        
        if has_pair and i - pair_idx <= 10 and random.random() < 0.7:
            pair_pos = coordinates[pair_idx]
            random_offset = np.random.normal(0, 1, 3) * 2.0
            base_pair_distance = 10.0 + random.uniform(-1.0, 1.0)
            center = np.mean(coordinates[:i], axis=0)
            direction = center - pair_pos
            direction = direction / (np.linalg.norm(direction) + 1e-10)
            coordinates[i] = pair_pos + direction * base_pair_distance + random_offset
            current_direction = np.random.normal(0, 0.3, 3)
            current_direction = current_direction / (np.linalg.norm(current_direction) + 1e-10)
        else:
            if random.random() < 0.3:
                angle = random.uniform(0.2, 0.6)
                axis = np.random.normal(0, 1, 3)
                axis = axis / (np.linalg.norm(axis) + 1e-10)
                rotation = R.from_rotvec(angle * axis)
                current_direction = rotation.apply(current_direction)
            else:
                current_direction += np.random.normal(0, 0.15, 3)
                current_direction = current_direction / (np.linalg.norm(current_direction) + 1e-10)
            
            step_size = random.uniform(3.5, 4.5)
            coordinates[i] = coordinates[i-1] + step_size * current_direction
    
    return coordinates

def predict_rna_structures(sequence, target_id, train_seqs_df, train_coords_dict, n_predictions=5, temporal_cutoff=None):
    """Generate n predictions for a sequence"""
    predictions = []
    similar_seqs = find_similar_sequences(sequence, train_seqs_df, train_coords_dict, 
                                         temporal_cutoff=temporal_cutoff, top_n=n_predictions)
    
    if similar_seqs:
        for i, (template_id, template_seq, similarity, template_coords) in enumerate(similar_seqs):
            adapted_coords = adapt_template_to_query(sequence, template_seq, template_coords)
            if adapted_coords is not None:
                refined_coords = adaptive_rna_constraints(adapted_coords, sequence, confidence=similarity)
                random_scale = max(0.05, 0.8 - similarity)
                randomized_coords = refined_coords.copy()
                randomized_coords += np.random.normal(0, random_scale, randomized_coords.shape)
                predictions.append(randomized_coords)
                
                if len(predictions) >= n_predictions:
                    break
    
    while len(predictions) < n_predictions:
        seed_value = hash(target_id) % 10000 + len(predictions) * 1000
        de_novo_coords = generate_rna_structure(sequence, seed=seed_value)
        refined_de_novo = adaptive_rna_constraints(de_novo_coords, sequence, confidence=0.2)
        predictions.append(refined_de_novo)
    
    return predictions[:n_predictions]

# Generate predictions
print("Generating predictions for test sequences...")
all_predictions = []
start_time = time.time()
total_targets = len(test_seqs)

for idx, row in test_seqs.iterrows():
    target_id = row['target_id']
    sequence = row['sequence']
    temporal_cutoff = row['temporal_cutoff'] if 'temporal_cutoff' in row else None
    
    if idx % 5 == 0:
        elapsed = time.time() - start_time
        targets_processed = idx + 1
        if targets_processed > 0:
            avg_time_per_target = elapsed / targets_processed
            est_time_remaining = avg_time_per_target * (total_targets - targets_processed)
            print(f"Processing {targets_processed}/{total_targets}: {target_id} ({len(sequence)} nt), "
                  f"elapsed: {elapsed:.1f}s, est. remaining: {est_time_remaining:.1f}s")
    
    predictions = predict_rna_structures(sequence, target_id, train_seqs, train_coords_dict, 
                                        n_predictions=5, temporal_cutoff=temporal_cutoff)
    
    for j in range(len(sequence)):
        pred_row = {
            'ID': f"{target_id}_{j+1}",
            'resname': sequence[j],
            'resid': j + 1
        }
        for i in range(5):
            pred_row[f'x_{i+1}'] = predictions[i][j][0]
            pred_row[f'y_{i+1}'] = predictions[i][j][1]
            pred_row[f'z_{i+1}'] = predictions[i][j][2]
        
        all_predictions.append(pred_row)

# Create and save submission
submission_df = pd.DataFrame(all_predictions)
column_order = ['ID', 'resname', 'resid']
for i in range(1, 6):
    for coord in ['x', 'y', 'z']:
        column_order.append(f'{coord}_{i}')
submission_df = submission_df[column_order]

submission_df.to_csv('submission.csv', index=False)
print(f"Generated predictions for {len(test_seqs)} RNA sequences")
print(f"Total runtime: {time.time() - start_time:.1f} seconds")
print(submission_df.head())
