#!/usr/bin/env python3
"""
RNA TaBM v73 - TRUE Multi-Method Ensemble

BREAKTHROUGH APPROACH - Fixing the 0.371 Ceiling:

Previous attempts all hit 0.371:
- v70: TBM + diverse DRfold2 → 0.371
- v71: TBM + physics/structure DRfold2 → 0.371
- v72: Pure Protenix → 0.371

ROOT CAUSE: All previous approaches created 5 models that differed by only 0.5-1Å
(essentially one prediction + noise). For effective best-of-5 ensemble, models
should differ by 5-10Å+.

THE SOLUTION - TRUE Multi-Method Ensemble:
- Model 1: Protenix (deep learning with MSA + templates)
- Model 2: DRfold2 (different deep learning architecture)
- Model 3: TBM template #1 (pure template-based, highest similarity)
- Model 4: TBM template #2 (second-best template)
- Model 5: TBM template #3 (third-best template)

These use fundamentally different information sources and will produce
genuinely different structures (5-10Å+ differences).

Expected improvement: 0.371 -> 0.42-0.48

If each method is right 40% of the time for different residues, best-of-5
could capture 70-80% correct residues.
"""

import subprocess
import sys
import os
import glob

# Install biopython
BIOPYTHON_PATTERNS = [
    '/kaggle/input/biopython-cp312/*.whl',
    '/kaggle/input/datasets/kami1976/biopython-cp312/*.whl',
]

for pattern in BIOPYTHON_PATTERNS:
    wheels = glob.glob(pattern)
    if wheels:
        path = wheels[0]
        print(f"Installing biopython from: {path}")
        subprocess.check_call([sys.executable, '-m', 'pip', 'install', '--no-index', path, '-q'])
        break

import pandas as pd
import numpy as np
import warnings

warnings.filterwarnings('ignore')

print("=" * 60)
print("RNA TaBM v73 - TRUE Multi-Method Ensemble")
print("=" * 60)
print("Fixing the 0.371 ceiling with genuinely diverse models!")
print()

# Load data
DATA_PATHS = [
    '/kaggle/input/stanford-rna-3d-folding-2/',
    '/kaggle/input/competitions/stanford-rna-3d-folding-2/',
]

DATA_PATH = None
for path in DATA_PATHS:
    if os.path.exists(path + 'train_sequences.csv'):
        DATA_PATH = path
        print(f"Using data path: {DATA_PATH}")
        break

if DATA_PATH is None:
    raise RuntimeError("Could not find competition data!")

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')

print(f"Train sequences: {len(train_seqs)}")
print(f"Test sequences: {len(test_seqs)}")
print(f"Train labels: {len(train_labels)}")

# Load Protenix predictions (Model 1)
PROTENIX_PATHS = [
    '/kaggle/input/protenix-predictions/protenix_predictions.csv',
    '/kaggle/input/datasets/pawanmali/protenix-predictions/protenix_predictions.csv',
]

protenix_predictions = None
for path in PROTENIX_PATHS:
    if os.path.exists(path):
        print(f"\nLoading Protenix predictions from: {path}")
        protenix_predictions = pd.read_csv(path)
        print(f"Loaded {len(protenix_predictions)} Protenix predictions")
        print("Model 1: Protenix (deep learning with MSA)")
        break

# Load DRfold2 best model (Model 2)
DRFOLD2_PATHS = [
    '/kaggle/input/drfold2-smart-selection/drfold2_smart_selection.csv',
    '/kaggle/input/datasets/pawanmali/drfold2-smart-selection/drfold2_smart_selection.csv',
]

drfold2_predictions = None
for path in DRFOLD2_PATHS:
    if os.path.exists(path):
        print(f"\nLoading DRfold2 predictions from: {path}")
        drfold2_predictions = pd.read_csv(path)
        print(f"Loaded {len(drfold2_predictions)} DRfold2 predictions")
        print("Model 2: DRfold2 (physics-based quality selection)")
        break

# Process labels for TBM (Models 3-5)
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 = group.sort_values('resid')[['x_1', 'y_1', 'z_1']].values
        coords_dict[id_prefix] = coords

    return coords_dict

train_coords_dict = process_labels(train_labels)
print(f"\nProcessed {len(train_coords_dict)} training structures for TBM")
print("Models 3-5: TBM templates (top 3 by sequence similarity)")

# Setup aligner for TBM
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

print("Aligner configured")

def find_similar_sequences(query_seq, train_seqs_df, train_coords_dict, top_n=5):
    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

        if abs(len(train_seq) - len(query_seq)) / max(len(train_seq), len(query_seq)) > 0.3:
            continue

        raw_score = aligner.score(query_seq, train_seq)
        normalized_score = raw_score / (2 * min(len(query_seq), len(train_seq)))

        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):
    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

    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:
                w = (i - prev_v) / (next_v - prev_v)
                new_coords[i] = (1-w)*new_coords[prev_v] + w*new_coords[next_v]
            elif prev_v >= 0:
                new_coords[i] = new_coords[prev_v] + [3, 0, 0]
            elif next_v >= 0:
                new_coords[i] = new_coords[next_v] + [3, 0, 0]
            else:
                new_coords[i] = [i*3, 0, 0]

    return np.nan_to_num(new_coords)

def get_tbm_templates(sequence, train_seqs_df, train_coords_dict, num_templates=3):
    """Get top N TBM templates (different templates, not noisy copies)"""
    similar_seqs = find_similar_sequences(sequence, train_seqs_df, train_coords_dict, top_n=num_templates)

    templates = []
    for t_id, t_seq, score, t_coords in similar_seqs:
        adapted = adapt_template_to_query(sequence, t_seq, t_coords)
        templates.append(adapted)

    # Fill remaining slots if needed
    while len(templates) < num_templates:
        n = len(sequence)
        coords = np.zeros((n, 3))
        for j in range(1, n):
            coords[j] = coords[j-1] + [5.9, 0, 0]
        templates.append(coords)

    return templates[:num_templates]

# Main prediction loop
print("\nGenerating TRUE multi-method ensemble...")
all_predictions = []
method_counts = {'protenix': 0, 'drfold2': 0, 'tbm_only': 0}

for idx, row in test_seqs.iterrows():
    tid, seq = row['target_id'], row['sequence']
    n = len(seq)

    if idx % 10 == 0:
        print(f"Processing {idx}/{len(test_seqs)}: {tid}")

    # Collect predictions from different methods
    models = [None] * 5  # 5 model slots

    # Model 1: Protenix (if available)
    has_protenix = False
    if protenix_predictions is not None:
        protenix_data = protenix_predictions[protenix_predictions['ID'].str.startswith(tid + '_')]
        if len(protenix_data) == n:
            has_protenix = True
            # Extract only the FIRST model from Protenix (x_1, y_1, z_1)
            protenix_coords = protenix_data[['x_1', 'y_1', 'z_1']].values
            models[0] = protenix_coords

    # Model 2: DRfold2 (if available)
    has_drfold2 = False
    if drfold2_predictions is not None:
        drfold2_data = drfold2_predictions[drfold2_predictions['ID'].str.startswith(tid + '_')]
        if len(drfold2_data) == n:
            has_drfold2 = True
            # Extract only the FIRST model from DRfold2 (x_1, y_1, z_1)
            drfold2_coords = drfold2_data[['x_1', 'y_1', 'z_1']].values
            models[1] = drfold2_coords

    # Models 3-5: TBM templates
    if has_protenix or has_drfold2:
        # Use 3 TBM templates
        tbm_templates = get_tbm_templates(seq, train_seqs, train_coords_dict, num_templates=3)
        models[2] = tbm_templates[0]
        models[3] = tbm_templates[1]
        models[4] = tbm_templates[2]

        if has_protenix and has_drfold2:
            method_counts['protenix'] += 1
        elif has_protenix:
            # No DRfold2, use 4 TBM templates
            models[1] = get_tbm_templates(seq, train_seqs, train_coords_dict, num_templates=4)[3]
            method_counts['protenix'] += 1
        else:
            # No Protenix, use 4 TBM templates
            models[0] = get_tbm_templates(seq, train_seqs, train_coords_dict, num_templates=4)[3]
            method_counts['drfold2'] += 1
    else:
        # No deep learning predictions, use 5 TBM templates
        tbm_templates = get_tbm_templates(seq, train_seqs, train_coords_dict, num_templates=5)
        models = tbm_templates
        method_counts['tbm_only'] += 1

    # Create prediction rows
    for j in range(n):
        res = {'ID': f"{tid}_{j+1}", 'resname': seq[j], 'resid': j+1}
        for i in range(5):
            res[f'x_{i+1}'] = models[i][j, 0]
            res[f'y_{i+1}'] = models[i][j, 1]
            res[f'z_{i+1}'] = models[i][j, 2]
        all_predictions.append(res)

print(f"\nPrediction summary:")
print(f"  Protenix + DRfold2 + TBM: {method_counts['protenix']}/28")
print(f"  DRfold2 + TBM only: {method_counts['drfold2']}/28")
print(f"  TBM only: {method_counts['tbm_only']}/28")
print(f"  Total residues: {len(all_predictions)}")

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

# Validate
print(f"\nSubmission shape: {sub.shape}")
coord_cols = [c for c in sub.columns if c.startswith(('x_', 'y_', 'z_'))]
print(f"Coordinate range: [{sub[coord_cols].min().min():.2f}, {sub[coord_cols].max().max():.2f}]")
print(f"Coordinate mean: {sub[coord_cols].mean().mean():.2f}")

nan_count = sub[coord_cols].isna().sum().sum()
if nan_count > 0:
    print(f"WARNING: {nan_count} NaN values found!")
    sub[coord_cols] = sub[coord_cols].fillna(0)
else:
    print("No NaN values - Good!")

sub.to_csv('submission.csv', index=False)
print("\nSubmission saved!")
print("=" * 60)
print("This is the CRITICAL test of true multi-method ensemble!")
print("Expected: 0.371 -> 0.42-0.48")
print("=" * 60)
print("DONE!")
