#!/usr/bin/env python3
"""
RNA 3D Structure Prediction V7 - RibonanzaNet 3D
Based on top-scoring Kaggle notebook approach
Uses pre-trained RibonanzaNet 3D weights with sliding window inference
"""
import sys
import random
import yaml
import numpy as np
import pandas as pd
import torch
import torch.nn as nn
from torch.utils.data import Dataset
from pathlib import Path

# Configuration
class Config:
    test_seq = "/kaggle/input/stanford-rna-3d-folding-2/test_sequences.csv"
    model_config_path = "/kaggle/input/ribonanzanet2d-final/configs/pairwise.yaml"
    finetuned_weights_path = "/kaggle/input/ribonanzanet-3d-v6/RibonanzaNet-3D.pt"
    
    max_len = 384
    batch_size = 1
    seed = 42

config = Config()

# Set seed
def set_seed(seed: int):
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    torch.backends.cudnn.deterministic = True
    torch.backends.cudnn.benchmark = False

set_seed(config.seed)

# Dataset
class RNADataset(Dataset):
    def __init__(self, data, max_len=384):
        self.data = data
        self.max_len = max_len
        self.tokens = {nt: i for i, nt in enumerate("ACGU")}

    def __len__(self):
        return len(self.data)

    def __getitem__(self, idx):
        sequence = [self.tokens[nt] for nt in self.data.loc[idx, "sequence"]]
        sequence = torch.tensor(np.array(sequence), dtype=torch.long)
        
        return {
            "sequence": sequence,
            "target_id": self.data.loc[idx, "target_id"]
        }

# Model
sys.path.append("/kaggle/input/ribonanzanet2d-final")
from Network import RibonanzaNet

class Config2:
    def __init__(self, **entries):
        self.__dict__.update(entries)

def load_config_from_yaml(file_path):
    with open(file_path, "r") as file:
        cfg = yaml.safe_load(file)
    return Config2(**cfg)

class FinetunedRibonanzaNet(RibonanzaNet):
    def __init__(self, config_obj, dropout=0.2):
        config_obj.dropout = dropout
        super(FinetunedRibonanzaNet, self).__init__(config_obj)
        
        self.dropout = nn.Dropout(p=0.0)
        self.xyz_predictor = nn.Linear(256, 3)

    def forward(self, src):
        sequence_features, _ = self.get_embeddings(
            src, torch.ones_like(src).long().to(src.device)
        )
        xyz_pred = self.xyz_predictor(sequence_features)
        return xyz_pred

# Inference with sliding window
def predict_sequence(model, sequence, max_len=384):
    """Predict XYZ coordinates with sliding window for long sequences"""
    device = next(model.parameters()).device
    seq_len = len(sequence)

    if seq_len <= max_len:
        src = sequence.unsqueeze(0).to(device)
        model.eval()
        with torch.no_grad():
            xyz = model(src).squeeze(0)
        return xyz.cpu().numpy()
    else:
        # Sliding window with 50% overlap
        step_size = max_len // 2
        predictions_sum = np.zeros((seq_len, 3))
        counts = np.zeros(seq_len)

        for start in range(0, seq_len - max_len + 1, step_size):
            end = start + max_len
            window = sequence[start:end].unsqueeze(0).to(device)
            
            model.eval()
            with torch.no_grad():
                xyz = model(window).squeeze(0)
            
            window_pred = xyz.cpu().numpy()
            predictions_sum[start:end] += window_pred
            counts[start:end] += 1

        # Handle last segment
        if (seq_len - max_len) % step_size != 0:
            start = seq_len - max_len
            window = sequence[start:].unsqueeze(0).to(device)
            
            model.eval()
            with torch.no_grad():
                xyz = model(window).squeeze(0)
            
            window_pred = xyz.cpu().numpy()
            predictions_sum[start:] += window_pred
            counts[start:] += 1

        # Average overlapping predictions
        final_pred = predictions_sum / counts[:, np.newaxis]
        return final_pred

def main():
    # Load test data
    test_data = pd.read_csv(config.test_seq)
    print(f"Test sequences: {len(test_data)}")
    
    test_dataset = RNADataset(test_data, max_len=config.max_len)
    
    # Load model
    model_cfg = load_config_from_yaml(config.model_config_path)
    model = FinetunedRibonanzaNet(model_cfg, dropout=0.2).cuda()
    
    # Load fine-tuned weights
    model.load_state_dict(torch.load(config.finetuned_weights_path, map_location="cuda"))
    print("Model loaded successfully!")
    
    # Run inference
    print("Running inference...")
    all_predictions = []
    
    for i in range(len(test_dataset)):
        sample = test_dataset[i]
        sequence = sample["sequence"]
        
        pred = predict_sequence(model, sequence, max_len=config.max_len)
        all_predictions.append(pred)
        
        if (i + 1) % 10 == 0:
            print(f"Processed {i + 1}/{len(test_dataset)} sequences")
    
    print("Inference complete!")
    
    # Create submission
    data = []
    
    for i in range(len(test_data)):
        target_id = test_data.loc[i, "target_id"]
        sequence = test_data.loc[i, "sequence"]
        seq_length = len(sequence)
        
        preds = all_predictions[i]
        
        for j in range(seq_length):
            row = [f"{target_id}_{j+1}", sequence[j], j + 1]
            
            # Add same prediction 5 times
            for k in range(5):
                row.extend([preds[j, 0], preds[j, 1], preds[j, 2]])
            
            data.append(row)
    
    # Create DataFrame
    columns = ["ID", "resname", "resid"]
    for i in range(1, 6):
        columns += [f"x_{i}", f"y_{i}", f"z_{i}"]
    
    submission = pd.DataFrame(data, columns=columns)
    
    print(f"\nSubmission shape: {submission.shape}")
    submission.to_csv("submission.csv", index=False)
    print("Submission saved to submission.csv")

if __name__ == "__main__":
    main()
