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

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

# Better error reporting on Windows
os.environ['CUDA_LAUNCH_BLOCKING'] = '0'  # Set to '1' for debugging if needed
torch.backends.cudnn.benchmark = True

# Fix Windows output buffering
if sys.platform == 'win32':
    import sys
    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()
    
    # 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 - FORCE GPU if available
    if torch.cuda.is_available():
        device = 'cuda'
        print(f"[GPU] Available: {torch.cuda.get_device_name(0)}")
        print(f"[GPU] CUDA Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB")
        torch.cuda.empty_cache()  # Clear any previous allocations
    else:
        device = 'cpu'
        print("[WARNING] No GPU available, using CPU (will be very slow)")
    
    config.training.device = device
    print(f"Using device: {device}")
    
    # 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 and tokenizer
    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,
            )
        
        # Move to GPU immediately
        model = model.to(device)
        print(f"Model moved to {device}")
        
        # Print model info
    total_params = sum(p.numel() for p in model.parameters())
    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
    print(f"Model parameters: {total_params:,}")
    print(f"Trainable parameters: {trainable_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
    config.save(os.path.join(config.training.save_dir, 'config.json'))
    print(f"Saved config to {os.path.join(config.training.save_dir, 'config.json')}")
    
    # Load test datasets if provided
    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} from {test_path}")
            test_texts = load_text_file(test_path)
            _, test_loader, _ = 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
            )
            test_loaders[test_name] = test_loader
            print(f"  Loaded {len(test_texts)} test texts")
    
        # Train
        print("\n" + "="*70)
        print("Starting Training")
        print("="*70)
        
        best_val_loss = float('inf')
        
        try:
            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
                    print("Starting training epoch...")
                    train_metrics = trainer.train_epoch()
                    print("Training epoch completed successfully")
                    
                    # Validate
                    val_metrics = trainer.validate()
                    
                    # Update best model
                    if val_metrics and val_metrics.get('val_loss', float('inf')) < best_val_loss:
                        best_val_loss = val_metrics['val_loss']
                        if args.save_best:
                            # Save best model
                            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"\nSaved best model to {checkpoint_path}")
                    
                    # 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)
                        
                        # Save evaluation results
                        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"\nSaved evaluation results to {results_path}")
                
                except Exception as e:
                    print(f"\nERROR during epoch {epoch + 1}: {e}")
                    import traceback
                    traceback.print_exc()
                    print("Continuing to next epoch...")
                    continue
        except KeyboardInterrupt:
            print("\nTraining interrupted by user")
        except Exception as e:
            print(f"\nFATAL ERROR: {e}")
            import traceback
            traceback.print_exc()
            raise
    
    # 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)
        
        # Generate comparison report
        report = benchmark_suite.generate_comparison_report(final_results)
        print(report)
        
        # Save final results
        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(f"\nSaved final benchmark results to {final_results_path}")
    
    print("\nTraining and evaluation complete!")


if __name__ == '__main__':
    main()

