"""
Training script with built-in benchmark evaluation - fixed version.
"""

import torch
import argparse
import sys
import os
from datetime import datetime

# Better error reporting on Windows
os.environ['CUDA_LAUNCH_BLOCKING'] = '0'
torch.backends.cudnn.benchmark = True

# Fix Windows output buffering
if sys.platform == 'win32':
    if hasattr(sys.stdout, 'reconfigure'):
        sys.stdout.reconfigure(encoding='utf-8', errors='replace')

sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from config.config import Config
from model.architecture import NovelAIModel
from training.trainer import Trainer
from utils.data_loader import create_data_loaders, load_text_file
from utils.tokenizer import SimpleTokenizer
from evaluation.benchmarks import BenchmarkSuite


def main():
    parser = argparse.ArgumentParser(description='Train and evaluate the novel AI model')
    parser.add_argument('--config', type=str, default=None, help='Path to config file')
    parser.add_argument('--train_data', type=str, required=True, help='Path to training data file')
    parser.add_argument('--val_data', type=str, default=None, help='Path to validation data file')
    parser.add_argument('--test_data', type=str, nargs='+', help='Path(s) to test data file(s)')
    parser.add_argument('--test_names', type=str, nargs='+', help='Names for test datasets')
    parser.add_argument('--output_dir', type=str, default='./checkpoints', help='Output directory')
    parser.add_argument('--epochs', type=int, default=None, help='Number of epochs')
    parser.add_argument('--batch_size', type=int, default=None, help='Batch size')
    parser.add_argument('--learning_rate', type=float, default=None, help='Learning rate')
    parser.add_argument('--eval_interval', type=int, default=1, help='Evaluate every N epochs')
    parser.add_argument('--save_best', action='store_true', help='Save best model based on validation')
    
    args = parser.parse_args()
    
    try:
        # Load config
        if args.config and os.path.exists(args.config):
            config = Config.load(args.config)
        else:
            config = Config.default()
        
        # Override config
        if args.epochs is not None:
            config.training.num_epochs = args.epochs
        if args.batch_size is not None:
            config.training.batch_size = args.batch_size
        if args.learning_rate is not None:
            config.training.learning_rate = args.learning_rate
        
        config.training.train_data_path = args.train_data
        config.training.val_data_path = args.val_data
        config.training.save_dir = args.output_dir
        
        # Set device
        device = 'cuda' if torch.cuda.is_available() else 'cpu'
        config.training.device = device
        print(f"Using device: {device}")
        if device == 'cuda':
            print(f"GPU: {torch.cuda.get_device_name(0)}")
            torch.cuda.empty_cache()
        
        # Load data
        print("Loading data...")
        train_texts = load_text_file(args.train_data)
        print(f"Loaded {len(train_texts)} training texts")
        
        val_texts = None
        if args.val_data:
            val_texts = load_text_file(args.val_data)
            print(f"Loaded {len(val_texts)} validation texts")
        
        # Create data loaders
        print("Creating data loaders...")
        train_loader, val_loader, tokenizer = create_data_loaders(
            train_texts=train_texts,
            val_texts=val_texts,
            max_length=config.model.max_seq_length,
            batch_size=config.training.batch_size,
            num_workers=0
        )
        
        print(f"Vocabulary size: {tokenizer.vocab_size}")
        print(f"Training batches: {len(train_loader)}")
        if val_loader:
            print(f"Validation batches: {len(val_loader)}")
        
        # Update model vocab size
        config.model.vocab_size = tokenizer.vocab_size
        
        # Create model
        print("Creating 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=False,  # Disable to avoid CUDA memory issues
            memory_buffer_size=config.model.memory_buffer_size,
            use_formula_attention=config.model.use_formula_attention,
        )
        
        model = model.to(device)
        print(f"Model moved to {device}")
        
        total_params = sum(p.numel() for p in model.parameters())
        print(f"Model parameters: {total_params:,}")
        
        # Create trainer
        trainer = Trainer(
            model=model,
            train_loader=train_loader,
            val_loader=val_loader,
            learning_rate=config.training.learning_rate,
            weight_decay=config.training.weight_decay,
            warmup_steps=config.training.warmup_steps,
            max_grad_norm=config.training.max_grad_norm,
            device=device,
            save_dir=config.training.save_dir,
            log_interval=config.training.log_interval,
        )
        
        # Save config
        os.makedirs(config.training.save_dir, exist_ok=True)
        config.save(os.path.join(config.training.save_dir, 'config.json'))
        
        # Load test datasets
        test_loaders = {}
        if args.test_data:
            test_names = args.test_names if args.test_names else [f"test_{i}" for i in range(len(args.test_data))]
            for test_path, test_name in zip(args.test_data, test_names):
                print(f"Loading test dataset: {test_name}")
                try:
                    test_texts = load_text_file(test_path)
                    if len(test_texts) == 0:
                        print(f"  Warning: {test_name} is empty, skipping")
                        continue
                    train_loader_test, val_loader_test, _ = create_data_loaders(
                        train_texts=test_texts,
                        val_texts=None,
                        tokenizer=tokenizer,
                        max_length=config.model.max_seq_length,
                        batch_size=config.training.batch_size,
                        num_workers=0,
                        build_vocab=False
                    )
                    # Use train_loader_test as the test loader (since val_texts was None)
                    test_loader_final = train_loader_test if train_loader_test else val_loader_test
                    if test_loader_final is None or len(test_loader_final) == 0:
                        print(f"  Warning: {test_name} created 0 batches, skipping")
                        continue
                    test_loaders[test_name] = test_loader_final
                    print(f"  Loaded {len(test_texts)} texts, {len(test_loader_final)} batches")
                except Exception as e:
                    print(f"  Error loading {test_name}: {e}")
                    import traceback
                    traceback.print_exc()
                    continue
        
        # Train
        print("\n" + "="*70)
        print("Starting Training")
        print("="*70)
        
        best_val_loss = float('inf')
        
        for epoch in range(config.training.num_epochs):
            print(f"\n{'='*70}")
            print(f"Epoch {epoch + 1}/{config.training.num_epochs}")
            print(f"{'='*70}")
            
            try:
                train_metrics = trainer.train_epoch()
                print(f"\nEpoch {epoch + 1} Train Loss: {train_metrics['train_loss']:.4f}")
                
                val_metrics = trainer.validate()
                if val_metrics:
                    print(f"Epoch {epoch + 1} Val Loss: {val_metrics.get('val_loss', 'N/A')}")
                    
                    if val_metrics.get('val_loss', float('inf')) < best_val_loss:
                        best_val_loss = val_metrics['val_loss']
                        if args.save_best:
                            checkpoint_path = os.path.join(config.training.save_dir, 'best_model.pt')
                            torch.save({
                                'epoch': epoch,
                                'model_state_dict': model.state_dict(),
                                'optimizer_state_dict': trainer.optimizer.state_dict(),
                                'val_loss': best_val_loss,
                            }, checkpoint_path)
                            print(f"Saved best model (val_loss={best_val_loss:.4f})")
                
                # Evaluate on test sets
                if args.eval_interval > 0 and (epoch + 1) % args.eval_interval == 0 and test_loaders:
                    print("\n" + "="*70)
                    print(f"Running Benchmarks (Epoch {epoch + 1})")
                    print("="*70)
                    
                    benchmark_suite = BenchmarkSuite(model, tokenizer, device=device)
                    test_results = benchmark_suite.evaluate_multiple_datasets(test_loaders)
                    
                    results_path = os.path.join(config.training.save_dir, f'eval_results_epoch_{epoch+1}.json')
                    import json
                    with open(results_path, 'w') as f:
                        json.dump(test_results, f, indent=2)
                    print(f"Saved evaluation results")
                    
            except Exception as e:
                print(f"\nERROR during epoch {epoch + 1}: {e}")
                import traceback
                traceback.print_exc()
                continue
        
        # Final evaluation
        if test_loaders:
            print("\n" + "="*70)
            print("Final Benchmark Evaluation")
            print("="*70)
            
            benchmark_suite = BenchmarkSuite(model, tokenizer, device=device)
            final_results = benchmark_suite.evaluate_multiple_datasets(test_loaders)
            
            report = benchmark_suite.generate_comparison_report(final_results)
            print(report)
            
            final_results_path = os.path.join(config.training.save_dir, 'final_benchmark_results.json')
            import json
            with open(final_results_path, 'w') as f:
                json.dump(final_results, f, indent=2)
        
        print("\nTraining and evaluation complete!")
        
    except Exception as e:
        print(f"\nFATAL ERROR: {e}")
        import traceback
        traceback.print_exc()
        sys.exit(1)


if __name__ == '__main__':
    main()

