"""
Benchmark evaluation module for comparing model performance.
"""

import torch
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
from typing import Dict, List, Tuple
import math
import numpy as np
from tqdm import tqdm


class PerplexityEvaluator:
    """
    Evaluates model perplexity on test datasets.
    Perplexity is the standard metric for language models.
    """
    
    def __init__(self, model, tokenizer, device='cuda'):
        self.model = model
        self.tokenizer = tokenizer
        self.device = device
        self.model.eval()
    
    def evaluate(self, data_loader: DataLoader) -> Dict[str, float]:
        """
        Evaluate perplexity on a dataset.
        
        Returns:
            Dictionary with:
            - perplexity: Overall perplexity
            - avg_loss: Average cross-entropy loss
            - num_tokens: Total number of tokens evaluated
        """
        total_loss = 0.0
        total_tokens = 0
        
        import sys
        use_tqdm = sys.platform != 'win32'
        
        eval_iter = tqdm(data_loader, desc="Evaluating perplexity") if use_tqdm else data_loader
        if not use_tqdm:
            print(f"Evaluating perplexity on {len(data_loader)} batches...")
        
        with torch.no_grad():
            for batch in eval_iter:
                input_ids = batch['input_ids'].to(self.device)
                target_ids = batch['target_ids'].to(self.device)
                attention_mask = batch.get('attention_mask', None)
                if attention_mask is not None:
                    attention_mask = attention_mask.to(self.device)
                
                # Forward pass
                output = self.model(input_ids=input_ids, attention_mask=attention_mask)
                logits = output['logits']  # [batch, seq_len, vocab_size]
                
                # Compute loss
                logits_flat = logits.view(-1, logits.shape[-1])
                targets_flat = target_ids.view(-1)
                
                # Mask padding tokens
                mask = (targets_flat != self.tokenizer.pad_token_id) & (targets_flat >= 0)
                if mask.sum() == 0:
                    continue
                
                loss = F.cross_entropy(logits_flat, targets_flat, reduction='none')
                loss = loss[mask]
                
                total_loss += loss.sum().item()
                total_tokens += mask.sum().item()
        
        if total_tokens == 0:
            return {'perplexity': float('inf'), 'avg_loss': float('inf'), 'num_tokens': 0}
        
        avg_loss = total_loss / total_tokens
        perplexity = math.exp(avg_loss)
        
        return {
            'perplexity': perplexity,
            'avg_loss': avg_loss,
            'num_tokens': total_tokens
        }


class AccuracyEvaluator:
    """
    Evaluates next token prediction accuracy.
    """
    
    def __init__(self, model, tokenizer, device='cuda'):
        self.model = model
        self.tokenizer = tokenizer
        self.device = device
        self.model.eval()
    
    def evaluate(self, data_loader: DataLoader) -> Dict[str, float]:
        """
        Evaluate accuracy on next token prediction.
        
        Returns:
            Dictionary with:
            - accuracy: Overall accuracy
            - top_k_accuracy: Top-k accuracy (for k=5)
            - num_predictions: Total number of predictions
        """
        correct = 0
        top_k_correct = 0
        total = 0
        
        import sys
        use_tqdm = sys.platform != 'win32'
        
        eval_iter = tqdm(data_loader, desc="Evaluating accuracy") if use_tqdm else data_loader
        if not use_tqdm:
            print(f"Evaluating accuracy on {len(data_loader)} batches...")
        
        with torch.no_grad():
            for batch in eval_iter:
                input_ids = batch['input_ids'].to(self.device)
                target_ids = batch['target_ids'].to(self.device)
                attention_mask = batch.get('attention_mask', None)
                if attention_mask is not None:
                    attention_mask = attention_mask.to(self.device)
                
                # Forward pass
                output = self.model(input_ids=input_ids, attention_mask=attention_mask)
                logits = output['logits']  # [batch, seq_len, vocab_size]
                
                # Get predictions
                predictions = logits.argmax(dim=-1)  # [batch, seq_len]
                
                # Get top-k predictions
                top_k_predictions = logits.topk(k=5, dim=-1)[1]  # [batch, seq_len, 5]
                
                # Mask padding tokens
                mask = (target_ids != self.tokenizer.pad_token_id) & (target_ids >= 0)
                
                # Compute accuracy
                correct_predictions = (predictions == target_ids) & mask
                correct += correct_predictions.sum().item()
                
                # Compute top-k accuracy
                target_expanded = target_ids.unsqueeze(-1).expand_as(top_k_predictions)
                top_k_correct_predictions = (top_k_predictions == target_expanded).any(dim=-1) & mask
                top_k_correct += top_k_correct_predictions.sum().item()
                
                total += mask.sum().item()
        
        if total == 0:
            return {'accuracy': 0.0, 'top_k_accuracy': 0.0, 'num_predictions': 0}
        
        accuracy = correct / total
        top_k_accuracy = top_k_correct / total
        
        return {
            'accuracy': accuracy,
            'top_k_accuracy': top_k_accuracy,
            'num_predictions': total
        }


class BenchmarkSuite:
    """
    Comprehensive benchmark evaluation suite.
    """
    
    def __init__(self, model, tokenizer, device='cuda'):
        self.model = model
        self.tokenizer = tokenizer
        self.device = device
        self.perplexity_evaluator = PerplexityEvaluator(model, tokenizer, device)
        self.accuracy_evaluator = AccuracyEvaluator(model, tokenizer, device)
    
    def evaluate_dataset(
        self,
        data_loader: DataLoader,
        dataset_name: str = "unknown"
    ) -> Dict[str, float]:
        """
        Evaluate model on a dataset.
        
        Returns:
            Dictionary with all metrics
        """
        print(f"\n{'='*70}")
        print(f"Evaluating on dataset: {dataset_name}")
        print(f"{'='*70}")
        
        # Evaluate perplexity
        perplexity_metrics = self.perplexity_evaluator.evaluate(data_loader)
        
        # Recreate data loader for accuracy (or reuse if possible)
        accuracy_metrics = self.accuracy_evaluator.evaluate(data_loader)
        
        # Combine metrics
        metrics = {
            **perplexity_metrics,
            **accuracy_metrics,
            'dataset_name': dataset_name
        }
        
        # Print results
        print(f"\nResults for {dataset_name}:")
        print(f"  Perplexity: {metrics['perplexity']:.2f}")
        print(f"  Average Loss: {metrics['avg_loss']:.4f}")
        print(f"  Accuracy: {metrics['accuracy']:.4f} ({metrics['accuracy']*100:.2f}%)")
        print(f"  Top-5 Accuracy: {metrics['top_k_accuracy']:.4f} ({metrics['top_k_accuracy']*100:.2f}%)")
        print(f"  Tokens Evaluated: {metrics['num_tokens']:,}")
        
        return metrics
    
    def evaluate_multiple_datasets(
        self,
        datasets: Dict[str, DataLoader]
    ) -> Dict[str, Dict[str, float]]:
        """
        Evaluate on multiple datasets.
        
        Args:
            datasets: Dictionary mapping dataset names to DataLoaders
            
        Returns:
            Dictionary mapping dataset names to their metrics
        """
        results = {}
        
        for dataset_name, data_loader in datasets.items():
            metrics = self.evaluate_dataset(data_loader, dataset_name)
            results[dataset_name] = metrics
        
        return results
    
    def generate_comparison_report(
        self,
        results: Dict[str, Dict[str, float]],
        baseline_results: Dict[str, Dict[str, float]] = None
    ) -> str:
        """
        Generate a comparison report.
        
        Args:
            results: Our model's results
            baseline_results: Baseline model results for comparison
            
        Returns:
            Formatted report string
        """
        report = []
        report.append("\n" + "="*70)
        report.append("BENCHMARK COMPARISON REPORT")
        report.append("="*70)
        
        # Perplexity comparison
        report.append("\n## Perplexity Results (lower is better)")
        report.append("-" * 70)
        report.append(f"{'Dataset':<30} {'Our Model':<15} {'Baseline':<15} {'Difference':<15}")
        report.append("-" * 70)
        
        for dataset_name, metrics in results.items():
            our_ppl = metrics['perplexity']
            baseline_ppl = baseline_results.get(dataset_name, {}).get('perplexity', None) if baseline_results else None
            
            if baseline_ppl:
                diff = our_ppl - baseline_ppl
                diff_pct = (diff / baseline_ppl) * 100
                report.append(f"{dataset_name:<30} {our_ppl:<15.2f} {baseline_ppl:<15.2f} {diff:+.2f} ({diff_pct:+.1f}%)")
            else:
                report.append(f"{dataset_name:<30} {our_ppl:<15.2f} {'N/A':<15} {'N/A':<15}")
        
        # Accuracy comparison
        report.append("\n## Accuracy Results (higher is better)")
        report.append("-" * 70)
        report.append(f"{'Dataset':<30} {'Our Model':<15} {'Baseline':<15} {'Difference':<15}")
        report.append("-" * 70)
        
        for dataset_name, metrics in results.items():
            our_acc = metrics['accuracy'] * 100
            baseline_acc = baseline_results.get(dataset_name, {}).get('accuracy', None) * 100 if baseline_results else None
            
            if baseline_acc:
                diff = our_acc - baseline_acc
                report.append(f"{dataset_name:<30} {our_acc:<15.2f} {baseline_acc:<15.2f} {diff:+.2f}%")
            else:
                report.append(f"{dataset_name:<30} {our_acc:<15.2f} {'N/A':<15} {'N/A':<15}")
        
        report.append("\n" + "="*70)
        
        return "\n".join(report)

