import os
import numpy as np
import matplotlib.pyplot as plt
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, TensorDataset
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import StandardScaler
from sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score
import torch.nn.functional as F
import math
import joblib
import pickle


# Positional Encoding for Transformer
class PositionalEncoding(nn.Module):
    """Adds positional information to the input embeddings."""
    def __init__(self, d_model, max_len=200):
        super(PositionalEncoding, self).__init__()
        
        pe = torch.zeros(max_len, d_model)
        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
        
        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)
        
        pe = pe.unsqueeze(0)
        self.register_buffer('pe', pe)
    
    def forward(self, x):
        return x + self.pe[:, :x.size(1), :]


# Transformer-based Ray Predictor
class TransformerRayPredictor(nn.Module):
    """Transformer model for predicting realistic ray data from ideal spatial data."""
    def __init__(self, input_dim=676, output_dim=676, d_model=128, nhead=8, 
                 num_encoder_layers=4, dim_feedforward=512, dropout=0.1):
        super(TransformerRayPredictor, self).__init__()
        
        self.input_dim = input_dim
        self.output_dim = output_dim
        self.d_model = d_model
        self.num_positions = 169
        self.features_per_position = 4
        
        self.input_embedding = nn.Linear(self.features_per_position, d_model)
        self.pos_encoder = PositionalEncoding(d_model, max_len=self.num_positions)
        
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=d_model,
            nhead=nhead,
            dim_feedforward=dim_feedforward,
            dropout=dropout,
            batch_first=True
        )
        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_encoder_layers)
        self.output_projection = nn.Linear(d_model, self.features_per_position)
        
        self._initialize_weights()
    
    def _initialize_weights(self):
        for module in self.modules():
            if isinstance(module, nn.Linear):
                nn.init.xavier_uniform_(module.weight)
                if module.bias is not None:
                    nn.init.constant_(module.bias, 0)
    
    def forward(self, x):
        batch_size = x.size(0)
        x = x.reshape(batch_size, self.num_positions, self.features_per_position)
        x = self.input_embedding(x)
        x = self.pos_encoder(x)
        x = self.transformer_encoder(x)
        x = self.output_projection(x)
        output = x.reshape(batch_size, -1)
        return output


# Enhanced Transformer with Pre-Layer Normalization
class EnhancedTransformerRayPredictor(nn.Module):
    """Enhanced transformer with modern architecture improvements."""
    def __init__(self, input_dim=676, output_dim=676, d_model=256, nhead=8, 
                 num_encoder_layers=6, dim_feedforward=1024, dropout=0.1):
        super(EnhancedTransformerRayPredictor, self).__init__()
        
        self.input_dim = input_dim
        self.output_dim = output_dim
        self.d_model = d_model
        self.num_positions = 169
        self.features_per_position = 4
        
        self.input_embedding = nn.Sequential(
            nn.Linear(self.features_per_position, d_model),
            nn.LayerNorm(d_model),
            nn.Dropout(dropout)
        )
        
        self.pos_embedding = nn.Parameter(torch.randn(1, self.num_positions, d_model))
        
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=d_model,
            nhead=nhead,
            dim_feedforward=dim_feedforward,
            dropout=dropout,
            batch_first=True,
            norm_first=True
        )
        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_encoder_layers)
        
        self.output_projection = nn.Sequential(
            nn.LayerNorm(d_model),
            nn.Linear(d_model, dim_feedforward // 2),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(dim_feedforward // 2, self.features_per_position)
        )
        
        self._initialize_weights()
    
    def _initialize_weights(self):
        for module in self.modules():
            if isinstance(module, nn.Linear):
                nn.init.xavier_uniform_(module.weight)
                if module.bias is not None:
                    nn.init.constant_(module.bias, 0)
    
    def forward(self, x):
        batch_size = x.size(0)
        x = x.reshape(batch_size, self.num_positions, self.features_per_position)
        x = self.input_embedding(x)
        x = x + self.pos_embedding
        x_transformed = self.transformer_encoder(x)
        output = self.output_projection(x_transformed)
        return output.reshape(batch_size, -1)


class TransformerRayDataProcessor:
    """Data processor with transformer model support."""
    def __init__(self, ideal_dir, realistic_dir):
        self.ideal_dir = ideal_dir
        self.realistic_dir = realistic_dir
        
        #change _z_ to _x_ or _y_ depending on which axis the code is running for
        self.test_files = [
            "2015-09-27T150500_z_stats.txt",
            "a2_2016-04-21T114700_z_stats.txt", 
            "2014-07-17T054500_z_stats.txt",
            "a3_2014-10-01T112000_z_stats.txt",
            "a1_2014-09-03T160908_z_stats.txt"
        ]
        
        #change the first value in paranthesis for range looping, 
        # in this case '2' which represents z-axis to '0' which represents x-axis or '1' which represents y-axis
        all_files = os.listdir(ideal_dir)
        all_y_files = [all_files[i] for i in range(2, len(all_files), 3)] 
        available_files = [f for f in all_y_files if f not in self.test_files]
        
        train_files, val_files = train_test_split(
            available_files, test_size=0.2, random_state=42
        )
        
        self.train_files = train_files
        self.val_files = val_files
        self.model_type = None  # Track which model type was built
        
        print(f"Total y-direction files found: {len(all_y_files)}")
        print(f"Training files: {len(self.train_files)}")
        print(f"Validation files: {len(self.val_files)}")
        print(f"Test files: {len(self.test_files)}")
        
        self.X_train = None
        self.X_val = None
        self.X_test = None
        self.y_train = None
        self.y_val = None
        self.y_test = None
        self.model = None
        self.scaler_X = StandardScaler()
        self.scaler_y = StandardScaler()
        
        self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
        print(f"Using device: {self.device}")
    
    def read_data_file(self, file_path):
        with open(file_path, 'r') as f:
            lines = f.readlines()
        
        data = {'alpha': [], 'beta': [], 'location': [], 'scale': []}
        
        for i, line in enumerate(lines):
            mod = i % 4
            value = float(line.strip())
            
            if mod == 0:
                data['alpha'].append(value)
            elif mod == 1:
                data['beta'].append(value)
            elif mod == 2:
                data['location'].append(value)
            elif mod == 3:
                data['scale'].append(value)
        
        return data
    
    def prepare_spatial_data(self, data_dict):
        features = []
        for key in ['alpha', 'beta', 'location', 'scale']:
            features.extend(data_dict[key])
        return np.array(features)
    
    def load_data(self):
        X_train_data = []
        y_train_data = []
        
        print("Loading training data...")
        for filename in self.train_files:
            ideal_file_path = os.path.join(self.ideal_dir, filename)
            realistic_file_path = os.path.join(self.realistic_dir, filename)
            
            if os.path.exists(ideal_file_path) and os.path.exists(realistic_file_path):
                ideal_data = self.read_data_file(ideal_file_path)
                realistic_data = self.read_data_file(realistic_file_path)
                
                ideal_features = self.prepare_spatial_data(ideal_data)
                realistic_features = self.prepare_spatial_data(realistic_data)
                
                X_train_data.append(ideal_features)
                y_train_data.append(realistic_features)
        
        X_val_data = []
        y_val_data = []
        
        print("Loading validation data...")
        for filename in self.val_files:
            ideal_file_path = os.path.join(self.ideal_dir, filename)
            realistic_file_path = os.path.join(self.realistic_dir, filename)
            
            if os.path.exists(ideal_file_path) and os.path.exists(realistic_file_path):
                ideal_data = self.read_data_file(ideal_file_path)
                realistic_data = self.read_data_file(realistic_file_path)
                
                ideal_features = self.prepare_spatial_data(ideal_data)
                realistic_features = self.prepare_spatial_data(realistic_data)
                
                X_val_data.append(ideal_features)
                y_val_data.append(realistic_features)
        
        self.X_train = np.array(X_train_data)
        self.y_train = np.array(y_train_data)
        self.X_val = np.array(X_val_data)
        self.y_val = np.array(y_val_data)
        
        X_test_data = []
        y_test_data = []
        
        print("Loading test data...")
        for filename in self.test_files:
            ideal_file_path = os.path.join(self.ideal_dir, filename)
            realistic_file_path = os.path.join(self.realistic_dir, filename)
            
            if os.path.exists(ideal_file_path) and os.path.exists(realistic_file_path):
                ideal_data = self.read_data_file(ideal_file_path)
                realistic_data = self.read_data_file(realistic_file_path)
                
                ideal_features = self.prepare_spatial_data(ideal_data)
                realistic_features = self.prepare_spatial_data(realistic_data)
                
                X_test_data.append(ideal_features)
                y_test_data.append(realistic_features)
        
        if X_test_data:
            self.X_test = np.array(X_test_data)
            self.y_test = np.array(y_test_data)
        else:
            raise ValueError("No test files found!")
        
        print("Scaling data...")
        self.X_train = self.scaler_X.fit_transform(self.X_train)
        self.X_val = self.scaler_X.transform(self.X_val)
        self.X_test = self.scaler_X.transform(self.X_test)
        
        self.y_train = self.scaler_y.fit_transform(self.y_train)
        self.y_val = self.scaler_y.transform(self.y_val)
        self.y_test = self.scaler_y.transform(self.y_test)
        
        print(f"Data loaded:")
        print(f"  Training samples: {len(self.X_train)}")
        print(f"  Validation samples: {len(self.X_val)}")
        print(f"  Test samples: {len(self.X_test)}")
    
    def build_model(self, model_type='transformer', **kwargs):
        """Build model based on specified type."""
        input_dim = self.X_train.shape[1]
        output_dim = self.y_train.shape[1]
        
        self.model_type = model_type  # Store model type
        
        if model_type == 'transformer':
            d_model = kwargs.get('d_model', 128)
            nhead = kwargs.get('nhead', 8)
            num_layers = kwargs.get('num_encoder_layers', 4)
            dim_feedforward = kwargs.get('dim_feedforward', 512)
            
            self.model = TransformerRayPredictor(
                input_dim=input_dim,
                output_dim=output_dim,
                d_model=d_model,
                nhead=nhead,
                num_encoder_layers=num_layers,
                dim_feedforward=dim_feedforward
            ).to(self.device)
            print(f"Built Transformer model (d_model={d_model}, heads={nhead}, layers={num_layers})")
            
        elif model_type == 'enhanced_transformer':
            d_model = kwargs.get('d_model', 256)
            nhead = kwargs.get('nhead', 8)
            num_layers = kwargs.get('num_encoder_layers', 6)
            dim_feedforward = kwargs.get('dim_feedforward', 1024)
            
            self.model = EnhancedTransformerRayPredictor(
                input_dim=input_dim,
                output_dim=output_dim,
                d_model=d_model,
                nhead=nhead,
                num_encoder_layers=num_layers,
                dim_feedforward=dim_feedforward
            ).to(self.device)
            print(f"Built Enhanced Transformer model (d_model={d_model}, heads={nhead}, layers={num_layers})")
        
        else:
            raise ValueError(f"Unknown model_type: {model_type}")
        
        print(f"Total parameters: {sum(p.numel() for p in self.model.parameters()):,}")
    
    def train_model(self, epochs=200, batch_size=32, learning_rate=0.001):
        if self.model is None:
            self.build_model()
        
        X_train_tensor = torch.FloatTensor(self.X_train).to(self.device)
        y_train_tensor = torch.FloatTensor(self.y_train).to(self.device)
        X_val_tensor = torch.FloatTensor(self.X_val).to(self.device)
        y_val_tensor = torch.FloatTensor(self.y_val).to(self.device)
        
        train_dataset = TensorDataset(X_train_tensor, y_train_tensor)
        train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
        
        val_dataset = TensorDataset(X_val_tensor, y_val_tensor)
        val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)
        
        optimizer = optim.AdamW(self.model.parameters(), lr=learning_rate, weight_decay=1e-5)
        criterion = nn.MSELoss()
        scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.7, patience=10)
        
        history = {'loss': [], 'val_loss': [], 'mae': [], 'val_mae': []}
        
        print(f"Starting training for {epochs} epochs...")
        
        best_val_loss = float('inf')
        patience_counter = 0
        patience = 20
        
        for epoch in range(epochs):
            self.model.train()
            train_loss = 0
            train_mae = 0
            num_batches = 0
            
            for batch_X, batch_y in train_loader:
                optimizer.zero_grad()
                outputs = self.model(batch_X)
                loss = criterion(outputs, batch_y)
                loss.backward()
                torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=1.0)
                optimizer.step()
                
                train_loss += loss.item()
                train_mae += F.l1_loss(outputs, batch_y).item()
                num_batches += 1
            
            self.model.eval()
            val_loss = 0
            val_mae = 0
            val_batches = 0
            
            with torch.no_grad():
                for batch_X, batch_y in val_loader:
                    outputs = self.model(batch_X)
                    loss = criterion(outputs, batch_y)
                    
                    val_loss += loss.item()
                    val_mae += F.l1_loss(outputs, batch_y).item()
                    val_batches += 1
            
            avg_train_loss = train_loss / num_batches
            avg_val_loss = val_loss / val_batches
            avg_train_mae = train_mae / num_batches
            avg_val_mae = val_mae / val_batches
            
            history['loss'].append(avg_train_loss)
            history['val_loss'].append(avg_val_loss)
            history['mae'].append(avg_train_mae)
            history['val_mae'].append(avg_val_mae)
            
            scheduler.step(avg_val_loss)
            
            if avg_val_loss < best_val_loss:
                best_val_loss = avg_val_loss
                patience_counter = 0
                torch.save(self.model.state_dict(), 'best_model_temp.pth')
            else:
                patience_counter += 1
            
            if (epoch + 1) % 10 == 0:
                print(f"Epoch [{epoch+1}/{epochs}] - "
                      f"Train Loss: {avg_train_loss:.4f}, Train MAE: {avg_train_mae:.4f}, "
                      f"Val Loss: {avg_val_loss:.4f}, Val MAE: {avg_val_mae:.4f}")
            
            if patience_counter >= patience:
                print(f"Early stopping at epoch {epoch+1}")
                break
        
        if os.path.exists('best_model_temp.pth'):
            self.model.load_state_dict(torch.load('best_model_temp.pth'))
            os.remove('best_model_temp.pth')
        
        print(f"Training complete! Best validation loss: {best_val_loss:.4f}")
        return history
    
    def evaluate_model(self):
        if self.model is None:
            raise ValueError("Model hasn't been trained yet")
        
        self.model.eval()
        X_test_tensor = torch.FloatTensor(self.X_test).to(self.device)
        
        with torch.no_grad():
            y_pred_tensor = self.model(X_test_tensor)
            y_pred = y_pred_tensor.cpu().numpy()
        
        mse = np.mean((y_pred - self.y_test) ** 2)
        mae = np.mean(np.abs(y_pred - self.y_test))
        
        y_pred_orig = self.scaler_y.inverse_transform(y_pred)
        y_test_orig = self.scaler_y.inverse_transform(self.y_test)
        
        ss_res = np.sum((y_test_orig - y_pred_orig) ** 2)
        ss_tot = np.sum((y_test_orig - np.mean(y_test_orig, axis=0)) ** 2)
        r2 = 1 - (ss_res / ss_tot)
        
        mape = np.mean(np.abs((y_test_orig - y_pred_orig) / (y_test_orig + 1e-10))) * 100
        
        return {'mse': mse, 'mae': mae, 'r2': r2, 'mape': mape}
    
    def save_model(self, model_path):
        """Save the trained model and scalers."""
        if self.model is None:
            raise ValueError("Model hasn't been trained yet")
        
        # Create directory if it doesn't exist
        os.makedirs(os.path.dirname(model_path), exist_ok=True)
        
        # Save model state
        torch.save(self.model.state_dict(), model_path)
        print(f"✅ Model weights saved to: {model_path}")
        
        # Save scalers
        scaler_X_path = model_path.replace('.pth', '_scaler_X.joblib')
        scaler_y_path = model_path.replace('.pth', '_scaler_y.joblib')
        
        joblib.dump(self.scaler_X, scaler_X_path)
        joblib.dump(self.scaler_y, scaler_y_path)
        print(f"✅ Scalers saved to:")
        print(f"   {scaler_X_path}")
        print(f"   {scaler_y_path}")
        
        # Save model metadata
        model_meta = {
            'model_type': self.model_type,
            'model_class': type(self.model).__name__,
            'input_dim': self.model.input_dim,
            'output_dim': self.model.output_dim,
            'num_positions': self.model.num_positions,
            'features_per_position': self.model.features_per_position,
        }
        
        # Add model-specific parameters
        if hasattr(self.model, 'd_model'):
            model_meta['d_model'] = self.model.d_model
        
        meta_path = model_path.replace('.pth', '_meta.pkl')
        with open(meta_path, 'wb') as f:
            pickle.dump(model_meta, f)
        print(f"✅ Metadata saved to: {meta_path}")
        
        print(f"\n{'='*60}")
        print("MODEL SAVED SUCCESSFULLY!")
        print(f"{'='*60}")
        print(f"Model type: {self.model_type}")
        print(f"Total parameters: {sum(p.numel() for p in self.model.parameters()):,}")
        print(f"\nFiles created:")
        print(f"  1. {os.path.basename(model_path)} - Model weights")
        print(f"  2. {os.path.basename(scaler_X_path)} - Input scaler")
        print(f"  3. {os.path.basename(scaler_y_path)} - Output scaler")
        print(f"  4. {os.path.basename(meta_path)} - Model metadata")
    
    def visualize_results(self, history=None):
        if history:
            plt.figure(figsize=(15, 10))
            
            plt.subplot(2, 2, 1)
            plt.plot(history['loss'], label='Train Loss')
            plt.plot(history['val_loss'], label='Validation Loss')
            plt.title('Model Loss')
            plt.ylabel('Loss')
            plt.xlabel('Epoch')
            plt.legend()
            plt.grid(True)
            
            plt.subplot(2, 2, 2)
            plt.plot(history['mae'], label='Train MAE')
            plt.plot(history['val_mae'], label='Validation MAE')
            plt.title('Model MAE')
            plt.ylabel('MAE')
            plt.xlabel('Epoch')
            plt.legend()
            plt.grid(True)
            
            plt.subplot(2, 2, 3)
            loss_diff = [abs(train - val) for train, val in zip(history['loss'], history['val_loss'])]
            plt.plot(loss_diff, label='|Train Loss - Val Loss|', color='red')
            plt.title('Training vs Validation Loss Difference')
            plt.ylabel('Loss Difference')
            plt.xlabel('Epoch')
            plt.legend()
            plt.grid(True)
            
            plt.subplot(2, 2, 4)
            plt.plot(history['loss'], label='Train Loss', alpha=0.7)
            plt.plot(history['val_loss'], label='Val Loss', alpha=0.7)
            plt.yscale('log')
            plt.title('Learning Curve (Log Scale)')
            plt.ylabel('Loss (Log Scale)')
            plt.xlabel('Epoch')
            plt.legend()
            plt.grid(True)
            
            plt.tight_layout()
            plt.show()


# Main execution
if __name__ == "__main__":
    # put the path for ideal trained files in ideal_dir, and same for realistic_dir for realistic trained files, in this case trained with smoothing method 13
    # model_save_path is the path where the model of transformer for the corresponding axis is saved to, in this case z for z-axis and 13 for the smoothing method
    ideal_dir = r"D:\allen\DLR\Original_data\ray_dataset_txt\smoothened data\13\ideal_method_13_training"
    realistic_dir = r"D:\allen\DLR\Original_data\ray_dataset_txt\smoothened data\13\realistic_method_13_training"
    model_save_path = r"D:\allen\DLR\models\cnn_transformer\z_axis\transformer_ray_predictor_z_axis_13.pth"
    
    # Create processor
    processor = TransformerRayDataProcessor(ideal_dir, realistic_dir)
    
    # Load data
    processor.load_data()
    
    print("\n" + "="*70)
    print("TRANSFORMER MODEL OPTIONS")
    print("="*70)
    print("\n1. 'transformer' - Standard Transformer")
    print("   • Good starting point")
    print("   • Fewer parameters, faster training")
    print("   • Recommended: d_model=128, nhead=8, layers=4")
    print("")
    print("2. 'enhanced_transformer' - Enhanced Transformer ⭐ RECOMMENDED")
    print("   • Pre-layer normalization (more stable)")
    print("   • Learnable positional embeddings")
    print("   • Better for complex patterns")
    print("   • Recommended: d_model=256, nhead=8, layers=6")
    print("="*70 + "\n")
    
    # Build transformer model
    processor.build_model(
        model_type='enhanced_transformer',
        d_model=256,
        nhead=8,
        num_encoder_layers=6,
        dim_feedforward=1024
    )
    
    # Train model
    print("\nStarting transformer training...")
    history = processor.train_model(epochs=200, batch_size=32, learning_rate=0.0001)
    
    # Evaluate
    print("\nEvaluating transformer...")
    metrics = processor.evaluate_model()
    print("\nTransformer Evaluation:")
    for key, value in metrics.items():
        print(f"  {key}: {value:.6f}")
    
    # Visualize
    processor.visualize_results(history)
    
    # ⭐ SAVE THE MODEL ⭐
    print("\n" + "="*70)
    print("SAVING TRANSFORMER MODEL")
    print("="*70)
    processor.save_model(model_save_path)
    
    print("\n" + "="*70)
    print("TRAINING COMPLETE!")
    print("="*70)
    print(f"\n✅ Model ready for inference")
    print(f"✅ Location: {model_save_path}")
    print(f"\n📌 Next step: Use the Transformer Inference Pipeline")