"""
Training infrastructure for the novel AI model.
"""

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, Dataset
from typing import Dict, Optional, List, Tuple
import os
import json
import time
from tqdm import tqdm
import sys
import traceback

sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from model.architecture import NovelAIModel
from training.logger import setup_logging, TrainingLogger


class LanguageModelDataset(Dataset):
    """
    Dataset for language modeling tasks.
    """
    
    def __init__(self, texts: List[str], tokenizer, max_length: int = 512):
        self.texts = texts
        self.tokenizer = tokenizer
        self.max_length = max_length
        
    def __len__(self):
        return len(self.texts)
    
    def __getitem__(self, idx):
        text = self.texts[idx]
        
        # Tokenize
        tokens = self.tokenizer.encode(text)
        
        # Truncate or pad
        if len(tokens) > self.max_length:
            tokens = tokens[:self.max_length]
        else:
            tokens = tokens + [self.tokenizer.pad_token_id] * (self.max_length - len(tokens))
        
        # Input and target (shifted by one)
        input_ids = torch.tensor(tokens[:-1], dtype=torch.long)
        target_ids = torch.tensor(tokens[1:], dtype=torch.long)
        
        return {
            'input_ids': input_ids,
            'target_ids': target_ids,
            'attention_mask': (input_ids != self.tokenizer.pad_token_id).long()
        }


class Trainer:
    """
    Trainer for the novel AI model.
    """
    
    def __init__(
        self,
        model: NovelAIModel,
        train_loader: DataLoader,
        val_loader: Optional[DataLoader] = None,
        learning_rate: float = 1e-4,
        weight_decay: float = 0.01,
        warmup_steps: int = 1000,
        max_grad_norm: float = 1.0,
        device: str = 'cuda' if torch.cuda.is_available() else 'cpu',
        save_dir: str = './checkpoints',
        log_interval: int = 100,
        logger: Optional[TrainingLogger] = None,
    ):
        # Initialize logger
        if logger is None:
            logger = setup_logging(log_dir=os.path.join(save_dir, 'logs'))
        self.logger = logger
        
        # Force GPU if available
        if device == 'cpu' and torch.cuda.is_available():
            self.logger.warning("CPU specified but GPU available. Using GPU instead.")
            device = 'cuda'
        
        self.device = device
        self.logger.info(f"Using device: {device}")
        if device == 'cuda':
            try:
                gpu_name = torch.cuda.get_device_name(0)
                gpu_memory = torch.cuda.get_device_properties(0).total_memory / 1e9
                self.logger.info(f"GPU: {gpu_name}")
                self.logger.info(f"CUDA Memory: {gpu_memory:.2f} GB")
            except Exception as e:
                self.logger.error(f"Error detecting GPU: {e}")
                self.logger.error("Falling back to CPU")
                device = 'cpu'
                self.device = device
        
        try:
            self.model = model.to(device)
            self.logger.debug(f"Model moved to {device}")
        except Exception as e:
            self.logger.critical(f"Failed to move model to {device}: {e}")
            raise
        
        self.train_loader = train_loader
        self.val_loader = val_loader
        self.device = device
        self.save_dir = save_dir
        self.log_interval = log_interval
        self.max_grad_norm = max_grad_norm
        
        # Validate data loaders
        if train_loader is None or len(train_loader) == 0:
            error_msg = "Train loader is empty or None. Cannot train without data."
            self.logger.error(error_msg)
            raise ValueError(error_msg)
        
        self.logger.info(f"Train batches: {len(train_loader)}")
        if val_loader:
            self.logger.info(f"Validation batches: {len(val_loader)}")
        else:
            self.logger.warning("No validation loader provided. Cannot track validation loss.")
        
        # Create save directory
        try:
            os.makedirs(save_dir, exist_ok=True)
            self.logger.info(f"Checkpoint directory: {save_dir}")
        except Exception as e:
            error_msg = f"Cannot create checkpoint directory {save_dir}: {e}"
            self.logger.critical(error_msg)
            raise RuntimeError(error_msg)
        
        # Setup optimizer
        self.optimizer = optim.AdamW(
            model.parameters(),
            lr=learning_rate,
            weight_decay=weight_decay,
            betas=(0.9, 0.95)
        )
        
        # Learning rate scheduler with warmup
        self.scheduler = optim.lr_scheduler.LambdaLR(
            self.optimizer,
            lr_lambda=lambda step: min(1.0, step / warmup_steps) if step < warmup_steps else 1.0
        )
        
        # Loss function
        self.criterion = nn.CrossEntropyLoss(ignore_index=-100)
        
        # Training state
        self.global_step = 0
        self.best_val_loss = float('inf')
        self.train_losses = []
        self.val_losses = []
    
    def train_epoch(self) -> Dict[str, float]:
        """Train for one epoch."""
        self.model.train()
        total_loss = 0.0
        num_batches = 0
        
        try:
            # Disable tqdm on Windows to avoid PowerShell issues
            import sys
            use_tqdm = sys.platform != 'win32'
            
            if use_tqdm:
                progress_bar = tqdm(self.train_loader, desc='Training')
                batch_iter = progress_bar
            else:
                batch_iter = self.train_loader
                print(f"Training {len(self.train_loader)} batches...")
            
            for batch_idx, batch in enumerate(batch_iter):
                try:
                    # Move to device
                    input_ids = batch['input_ids'].to(self.device, non_blocking=True)
                    target_ids = batch['target_ids'].to(self.device, non_blocking=True)
                    attention_mask = batch['attention_mask'].to(self.device, non_blocking=True)
                    
                    # Forward pass
                    output = self.model(input_ids=input_ids, attention_mask=attention_mask, use_cache=True)
                    logits = output['logits']  # [batch, seq_len, vocab_size]
                    
                    # Compute loss
                    # Reshape for cross entropy: [batch * seq_len, vocab_size] vs [batch * seq_len]
                    logits_flat = logits.view(-1, logits.shape[-1])
                    targets_flat = target_ids.view(-1)
                    
                    # Mask out padding tokens
                    mask = (targets_flat != -100) & (targets_flat != 0) & (targets_flat != self.model.token_embedding.weight.shape[0] - 1)
                    if mask.sum() == 0:
                        continue
                    
                    loss = self.criterion(logits_flat, targets_flat)
            
                    # Backward pass
                    self.optimizer.zero_grad()
                    
                    # Check for NaN/inf
                    if torch.isnan(loss) or torch.isinf(loss):
                        error_msg = (
                            f"Invalid loss value {loss.item():.6f} at batch {batch_idx}, step {self.global_step}.\n"
                            f"  → This may indicate:\n"
                            f"     1. Learning rate too high (current: {self.scheduler.get_last_lr()[0]:.2e})\n"
                            f"     2. Numerical instability in model\n"
                            f"     3. Corrupted data in batch"
                        )
                        self.logger.error(error_msg)
                        self.logger.warning(f"Skipping batch {batch_idx}")
                        continue
                    
                    loss.backward()
                    
                    # Gradient clipping
                    grad_norm = torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.max_grad_norm)
                    
                    if torch.isnan(grad_norm) or torch.isinf(grad_norm):
                        error_msg = (
                            f"Invalid gradient norm {grad_norm.item():.6f} at batch {batch_idx}.\n"
                            f"  → This indicates exploding gradients.\n"
                            f"  → Try: reducing learning rate, enabling gradient clipping, or reducing batch size"
                        )
                        self.logger.error(error_msg)
                        self.logger.warning(f"Skipping batch {batch_idx}")
                        continue
                    
                    # Log very large gradient norms as warning
                    if grad_norm.item() > 10.0:
                        self.logger.warning(
                            f"Large gradient norm {grad_norm.item():.4f} at batch {batch_idx}. "
                            f"Consider reducing learning rate."
                        )
                    
                    self.optimizer.step()
                    self.scheduler.step()
                    
                    # Update statistics
                    total_loss += loss.item()
                    num_batches += 1
                    self.global_step += 1
                    
                    # Update progress bar
                    current_lr = self.scheduler.get_last_lr()[0]
                    if use_tqdm:
                        progress_bar.set_postfix({
                            'loss': f"{loss.item():.4f}",
                            'avg_loss': f"{total_loss / num_batches:.4f}",
                            'lr': f"{current_lr:.2e}",
                            'step': self.global_step
                        })
                    elif batch_idx % 10 == 0:  # Print every 10 batches on Windows
                        print(f"Batch {batch_idx}/{len(self.train_loader)}: Loss={loss.item():.4f}, Avg={total_loss/num_batches:.4f}")
                    
                    # Logging
                    if self.global_step % self.log_interval == 0:
                        self.logger.info(
                            f"Step {self.global_step}: Loss={loss.item():.4f}, "
                            f"Avg Loss={total_loss/num_batches:.4f}, LR={current_lr:.2e}, "
                            f"GradNorm={grad_norm.item():.4f}"
                        )
                    
                    # Clear cache periodically
                    if batch_idx % 10 == 0 and self.device == 'cuda':
                        torch.cuda.empty_cache()
                        
                except RuntimeError as e:
                    if "out of memory" in str(e).lower():
                        error_msg = (
                            f"GPU out of memory at batch {batch_idx}, step {self.global_step}.\n"
                            f"  → Solutions:\n"
                            f"     1. Reduce batch_size (current: {self.train_loader.batch_size})\n"
                            f"     2. Reduce max_seq_length in model config\n"
                            f"     3. Use gradient accumulation (process smaller batches)\n"
                            f"     4. Close other GPU applications\n"
                            f"     5. Enable mixed precision training"
                        )
                        self.logger.critical(error_msg)
                        torch.cuda.empty_cache()
                        self.logger.info("GPU cache cleared. Attempting to continue...")
                        # Try to continue with next batch
                        continue
                    else:
                        error_msg = f"RuntimeError at batch {batch_idx}, step {self.global_step}: {e}"
                        self.logger.error(error_msg, exc_info=True)
                        raise
                except Exception as e:
                    error_msg = (
                        f"Unexpected error at batch {batch_idx}, step {self.global_step}.\n"
                        f"  Error: {type(e).__name__}: {e}\n"
                        f"  → Check data format, model configuration, or system resources"
                    )
                    self.logger.critical(error_msg, exc_info=True)
                    raise
        
        except KeyboardInterrupt:
            self.logger.info("Training interrupted by user")
            raise
        except Exception as e:
            error_msg = (
                f"FATAL ERROR in train_epoch: {type(e).__name__}: {e}\n"
                f"  → Check logs for detailed error information"
            )
            self.logger.critical(error_msg, exc_info=True)
            raise
        
        avg_loss = total_loss / num_batches if num_batches > 0 else 0.0
        self.train_losses.append(avg_loss)
        
        return {
            'train_loss': avg_loss,
            'learning_rate': self.scheduler.get_last_lr()[0]
        }
    
    def validate(self) -> Dict[str, float]:
        """Validate the model."""
        if self.val_loader is None:
            return {}
        
        self.model.eval()
        total_loss = 0.0
        num_batches = 0
        
        import sys
        use_tqdm = sys.platform != 'win32'
        
        val_iter = tqdm(self.val_loader, desc='Validating') if use_tqdm else self.val_loader
        if not use_tqdm:
            print(f"Validating {len(self.val_loader)} batches...")
        
        with torch.no_grad():
            for batch in val_iter:
                input_ids = batch['input_ids'].to(self.device)
                target_ids = batch['target_ids'].to(self.device)
                attention_mask = batch['attention_mask'].to(self.device)
                
                output = self.model(input_ids=input_ids, attention_mask=attention_mask, use_cache=False)
                logits = output['logits']
                
                logits_flat = logits.view(-1, logits.shape[-1])
                targets_flat = target_ids.view(-1)
                
                mask = (targets_flat != -100) & (targets_flat != 0)
                if mask.sum() == 0:
                    continue
                
                loss = self.criterion(logits_flat, targets_flat)
                
                total_loss += loss.item()
                num_batches += 1
        
        avg_loss = total_loss / num_batches if num_batches > 0 else float('inf')
        self.val_losses.append(avg_loss)
        
        return {'val_loss': avg_loss}
    
    def train(self, num_epochs: int):
        """Train the model for multiple epochs."""
        print(f"Starting training for {num_epochs} epochs")
        print(f"Device: {self.device}")
        print(f"Model parameters: {sum(p.numel() for p in self.model.parameters()):,}")
        print(f"Trainable parameters: {sum(p.numel() for p in self.model.parameters() if p.requires_grad):,}")
        
        for epoch in range(num_epochs):
            print(f"\n{'='*50}")
            print(f"Epoch {epoch + 1}/{num_epochs}")
            print(f"{'='*50}")
            
            # Train
            train_metrics = self.train_epoch()
            
            # Validate
            val_metrics = self.validate()
            
            # Print epoch summary
            print(f"\nEpoch {epoch + 1} Summary:")
            print(f"  Train Loss: {train_metrics['train_loss']:.4f}")
            if val_metrics:
                print(f"  Val Loss: {val_metrics['val_loss']:.4f}")
            
            # Save checkpoint
            is_best = False
            if val_metrics and val_metrics['val_loss'] < self.best_val_loss:
                self.best_val_loss = val_metrics['val_loss']
                is_best = True
            
            self.save_checkpoint(epoch, train_metrics, val_metrics, is_best)
            
            # Print formula weights
            if hasattr(self.model.layers[0].attention, 'scoring_formula'):
                formula = self.model.layers[0].attention.scoring_formula
                print(f"\nFormula Weights:")
                print(f"  w1 (novelty): {formula.w1.item():.4f}")
                print(f"  w2 (retention): {formula.w2.item():.4f}")
                print(f"  w3 (payoff): {formula.w3.item():.4f}")
                print(f"  λ (decay): {formula.lambda_decay.item():.4f}")
                print(f"  k (fatigue): {formula.k_fatigue.item():.4f}")
    
    def save_checkpoint(
        self,
        epoch: int,
        train_metrics: Dict[str, float],
        val_metrics: Dict[str, float],
        is_best: bool = False
    ):
        """Save a training checkpoint."""
        try:
            checkpoint = {
                'epoch': epoch,
                'global_step': self.global_step,
                'model_state_dict': self.model.state_dict(),
                'optimizer_state_dict': self.optimizer.state_dict(),
                'scheduler_state_dict': self.scheduler.state_dict(),
                'train_metrics': train_metrics,
                'val_metrics': val_metrics,
                'best_val_loss': self.best_val_loss,
            }
            
            # Save regular checkpoint
            checkpoint_path = os.path.join(self.save_dir, f'checkpoint_epoch_{epoch}.pt')
            torch.save(checkpoint, checkpoint_path)
            self.logger.debug(f"Saved checkpoint: {checkpoint_path}")
            
            # Save best model
            if is_best:
                best_path = os.path.join(self.save_dir, 'best_model.pt')
                torch.save(checkpoint, best_path)
                self.logger.info(f"Saved best model (val_loss={self.best_val_loss:.6f}) to {best_path}")
            
            # Save latest checkpoint
            latest_path = os.path.join(self.save_dir, 'latest_checkpoint.pt')
            torch.save(checkpoint, latest_path)
            
        except Exception as e:
            error_msg = (
                f"Failed to save checkpoint at epoch {epoch}: {e}\n"
                f"  → Check disk space and write permissions\n"
                f"  → Checkpoint directory: {self.save_dir}"
            )
            self.logger.critical(error_msg, exc_info=True)
            raise
    
    def load_checkpoint(self, checkpoint_path: str):
        """Load a training checkpoint with error handling."""
        if not os.path.exists(checkpoint_path):
            error_msg = f"Checkpoint file not found: {checkpoint_path}"
            self.logger.error(error_msg)
            raise FileNotFoundError(error_msg)
        
        try:
            self.logger.info(f"Loading checkpoint from {checkpoint_path}")
            checkpoint = torch.load(checkpoint_path, map_location=self.device)
            
            # Validate checkpoint structure
            required_keys = ['model_state_dict', 'optimizer_state_dict', 'scheduler_state_dict']
            missing_keys = [k for k in required_keys if k not in checkpoint]
            if missing_keys:
                error_msg = (
                    f"Checkpoint is missing required keys: {missing_keys}\n"
                    f"  → Checkpoint may be corrupted or from an incompatible version"
                )
                self.logger.error(error_msg)
                raise ValueError(error_msg)
            
            # Load state dicts with error handling
            try:
                self.model.load_state_dict(checkpoint['model_state_dict'], strict=False)
            except Exception as e:
                error_msg = (
                    f"Failed to load model state dict: {e}\n"
                    f"  → Model architecture may have changed\n"
                    f"  → Try retraining from scratch or check model config"
                )
                self.logger.error(error_msg)
                raise
            
            try:
                self.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
            except Exception as e:
                self.logger.warning(f"Could not load optimizer state: {e}. Using fresh optimizer.")
            
            try:
                self.scheduler.load_state_dict(checkpoint['scheduler_state_dict'])
            except Exception as e:
                self.logger.warning(f"Could not load scheduler state: {e}. Using fresh scheduler.")
            
            self.global_step = checkpoint.get('global_step', 0)
            self.best_val_loss = checkpoint.get('best_val_loss', float('inf'))
            
            epoch = checkpoint.get('epoch', 0)
            self.logger.info(f"Loaded checkpoint. Resuming from epoch {epoch}, step {self.global_step}")
            self.logger.info(f"Best validation loss so far: {self.best_val_loss:.6f}")
            
        except Exception as e:
            error_msg = f"Failed to load checkpoint {checkpoint_path}: {e}"
            self.logger.critical(error_msg, exc_info=True)
            raise

