"""
Run comprehensive benchmark evaluations and generate comparison reports.
"""

import torch
import json
import argparse
import sys
import os
from datetime import datetime

sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from config.config import Config
from model.architecture import NovelAIModel
from utils.tokenizer import SimpleTokenizer
from utils.data_loader import create_data_loaders, load_text_file
from evaluation.benchmarks import BenchmarkSuite


# Baseline results for common open-source models (for comparison)
# These are approximate values for reference
BASELINE_RESULTS = {
    'GPT-2 Small (124M)': {
        'perplexity': 35.0,
        'accuracy': 0.45,
    },
    'GPT-2 Medium (355M)': {
        'perplexity': 26.0,
        'accuracy': 0.50,
    },
    'GPT-2 Large (774M)': {
        'perplexity': 22.0,
        'accuracy': 0.52,
    },
    'Transformer Base': {
        'perplexity': 40.0,
        'accuracy': 0.42,
    },
}


def load_model(checkpoint_path, config_path=None):
    """Load a trained model from checkpoint."""
    print(f"Loading model from {checkpoint_path}...")
    
    # Load config
    if config_path is None:
        config_path = os.path.join(os.path.dirname(checkpoint_path), 'config.json')
    
    if os.path.exists(config_path):
        config = Config.load(config_path)
    else:
        config = Config.default()
        print("Warning: Config file not found, using defaults")
    
    # Create tokenizer (in production, save/load tokenizer separately)
    tokenizer = SimpleTokenizer(is_char_level=False)
    
    # Create model
    formula_config = {
        'w1': config.model.formula_w1,
        'w2': config.model.formula_w2,
        'w3': config.model.formula_w3,
        'lambda_decay': config.model.formula_lambda,
        'k_fatigue': config.model.formula_k,
        'learnable_weights': config.model.formula_learnable_weights,
        'learnable_decay': config.model.formula_learnable_decay,
    }
    
    model = NovelAIModel(
        vocab_size=config.model.vocab_size,
        embedding_dim=config.model.embedding_dim,
        num_layers=config.model.num_layers,
        num_heads=config.model.num_heads,
        feedforward_dim=config.model.feedforward_dim,
        max_seq_length=config.model.max_seq_length,
        dropout=config.model.dropout,
        formula_config=formula_config,
        use_memory_buffer=config.model.use_memory_buffer,
        memory_buffer_size=config.model.memory_buffer_size,
        use_formula_attention=config.model.use_formula_attention,
    )
    
    # Load weights
    device = 'cuda' if torch.cuda.is_available() else 'cpu'
    checkpoint = torch.load(checkpoint_path, map_location=device)
    
    if 'model_state_dict' in checkpoint:
        model.load_state_dict(checkpoint['model_state_dict'])
    else:
        model.load_state_dict(checkpoint)
    
    model = model.to(device)
    model.eval()
    
    print(f"Model loaded successfully on {device}")
    print(f"Parameters: {sum(p.numel() for p in model.parameters()):,}")
    
    return model, tokenizer, config, device


def evaluate_on_datasets(model, tokenizer, datasets, device, output_dir):
    """Evaluate model on multiple datasets."""
    benchmark_suite = BenchmarkSuite(model, tokenizer, device=device)
    
    # Prepare data loaders
    test_loaders = {}
    for dataset_name, dataset_path in datasets.items():
        print(f"\nLoading dataset: {dataset_name} from {dataset_path}")
        texts = load_text_file(dataset_path)
        
        # Create data loader
        _, test_loader, _ = create_data_loaders(
            train_texts=texts,
            val_texts=None,
            tokenizer=tokenizer,
            max_length=512,  # Use model's max length
            batch_size=8,
            num_workers=0,
            build_vocab=False
        )
        
        test_loaders[dataset_name] = test_loader
        print(f"  Loaded {len(texts)} texts")
    
    # Run evaluations
    results = benchmark_suite.evaluate_multiple_datasets(test_loaders)
    
    # Save results
    results_path = os.path.join(output_dir, 'benchmark_results.json')
    with open(results_path, 'w') as f:
        json.dump(results, f, indent=2)
    print(f"\nSaved results to {results_path}")
    
    return results


def generate_comparison_report(results, model_name, output_dir):
    """Generate a detailed comparison report."""
    
    # Calculate model size
    total_params = sum(p.numel() for p in results.get('model_info', {}).get('parameters', [0]))
    
    report_lines = []
    report_lines.append("\n" + "="*80)
    report_lines.append("BENCHMARK COMPARISON REPORT")
    report_lines.append("="*80)
    report_lines.append(f"Model: {model_name}")
    report_lines.append(f"Parameters: {total_params:,}")
    report_lines.append(f"Date: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
    report_lines.append("="*80)
    
    # Perplexity comparison
    report_lines.append("\n## Perplexity Results (lower is better)")
    report_lines.append("-"*80)
    report_lines.append(f"{'Dataset':<30} {model_name:<20} {'Baseline (GPT-2 Small)':<25} {'Improvement':<15}")
    report_lines.append("-"*80)
    
    baseline_ppl = BASELINE_RESULTS.get('GPT-2 Small (124M)', {}).get('perplexity', None)
    
    for dataset_name, metrics in results.items():
        if dataset_name == 'model_info':
            continue
            
        our_ppl = metrics.get('perplexity', float('inf'))
        
        if baseline_ppl:
            improvement = ((baseline_ppl - our_ppl) / baseline_ppl) * 100
            report_lines.append(f"{dataset_name:<30} {our_ppl:<20.2f} {baseline_ppl:<25.2f} {improvement:+.1f}%")
        else:
            report_lines.append(f"{dataset_name:<30} {our_ppl:<20.2f} {'N/A':<25}")
    
    # Accuracy comparison
    report_lines.append("\n## Accuracy Results (higher is better)")
    report_lines.append("-"*80)
    report_lines.append(f"{'Dataset':<30} {model_name:<20} {'Baseline (GPT-2 Small)':<25} {'Improvement':<15}")
    report_lines.append("-"*80)
    
    baseline_acc = BASELINE_RESULTS.get('GPT-2 Small (124M)', {}).get('accuracy', None)
    
    for dataset_name, metrics in results.items():
        if dataset_name == 'model_info':
            continue
            
        our_acc = metrics.get('accuracy', 0.0) * 100
        
        if baseline_acc:
            improvement = ((our_acc - baseline_acc * 100) / (baseline_acc * 100)) * 100
            report_lines.append(f"{dataset_name:<30} {our_acc:<20.2f} {baseline_acc*100:<25.2f} {improvement:+.1f}%")
        else:
            report_lines.append(f"{dataset_name:<30} {our_acc:<20.2f} {'N/A':<25}")
    
    # Summary statistics
    report_lines.append("\n## Summary Statistics")
    report_lines.append("-"*80)
    
    perplexities = [m.get('perplexity', float('inf')) for k, m in results.items() if k != 'model_info' and m.get('perplexity')]
    accuracies = [m.get('accuracy', 0.0) for k, m in results.items() if k != 'model_info' and m.get('accuracy')]
    
    if perplexities:
        report_lines.append(f"Average Perplexity: {sum(perplexities)/len(perplexities):.2f}")
        report_lines.append(f"Best Perplexity: {min(perplexities):.2f}")
        report_lines.append(f"Worst Perplexity: {max(perplexities):.2f}")
    
    if accuracies:
        report_lines.append(f"Average Accuracy: {sum(accuracies)/len(accuracies):.4f}")
        report_lines.append(f"Best Accuracy: {max(accuracies):.4f}")
        report_lines.append(f"Worst Accuracy: {min(accuracies):.4f}")
    
    report_lines.append("\n" + "="*80)
    
    report_text = "\n".join(report_lines)
    
    # Save report
    report_path = os.path.join(output_dir, 'benchmark_comparison_report.txt')
    with open(report_path, 'w') as f:
        f.write(report_text)
    
    print(report_text)
    print(f"\nSaved report to {report_path}")
    
    return report_text


def main():
    parser = argparse.ArgumentParser(description='Run benchmark evaluations')
    parser.add_argument('--checkpoint', type=str, required=True, help='Path to model checkpoint')
    parser.add_argument('--config', type=str, default=None, help='Path to config file')
    parser.add_argument('--datasets', type=str, nargs='+', required=True, help='Dataset paths')
    parser.add_argument('--dataset_names', type=str, nargs='+', help='Dataset names')
    parser.add_argument('--output_dir', type=str, default='./benchmark_results', help='Output directory')
    parser.add_argument('--model_name', type=str, default='Novel AI Model', help='Model name for report')
    
    args = parser.parse_args()
    
    os.makedirs(args.output_dir, exist_ok=True)
    
    # Load model
    model, tokenizer, config, device = load_model(args.checkpoint, args.config)
    
    # Prepare datasets
    dataset_names = args.dataset_names if args.dataset_names else [f"dataset_{i}" for i in range(len(args.datasets))]
    datasets = dict(zip(dataset_names, args.datasets))
    
    # Run evaluations
    results = evaluate_on_datasets(model, tokenizer, datasets, device, args.output_dir)
    
    # Add model info
    results['model_info'] = {
        'parameters': list(model.parameters()),
        'total_params': sum(p.numel() for p in model.parameters()),
        'embedding_dim': config.model.embedding_dim,
        'num_layers': config.model.num_layers,
    }
    
    # Generate comparison report
    generate_comparison_report(results, args.model_name, args.output_dir)
    
    print("\nBenchmark evaluation complete!")


if __name__ == '__main__':
    main()



