"""
Training script with better error handling and GPU optimization.
"""

import torch
import argparse
import sys
import os
from datetime import datetime
import traceback

# Force GPU usage and optimize
os.environ['CUDA_LAUNCH_BLOCKING'] = '0'  # Async operations
torch.backends.cudnn.benchmark = True  # Optimize for consistent input sizes

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 with better error handling
        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 cache
            # Set memory fraction to avoid issues
            torch.cuda.set_per_process_memory_fraction(0.9)
        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...")
        try:
            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  # Keep at 0 for Windows
            )
        except Exception as e:
            print(f"ERROR creating data loaders: {e}")
            traceback.print_exc()
            return
        
        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...")
        try:
            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,
            )
            
            # 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:,}")
            
        except Exception as e:
            print(f"ERROR creating model: {e}")
            traceback.print_exc()
            return
        
        # Create trainer
        print("Creating trainer...")
        try:
            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,
            )
        except Exception as e:
            print(f"ERROR creating trainer: {e}")
            traceback.print_exc()
            return
        
        # Save config
        os.makedirs(config.training.save_dir, exist_ok=True)
        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}")
                try:
                    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")
                except Exception as e:
                    print(f"ERROR loading test dataset {test_name}: {e}")
                    traceback.print_exc()
        
        # Train with better error handling
        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
                    train_metrics = trainer.train_epoch()
                    
                    # 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)
                        
                        try:
                            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"ERROR during benchmark evaluation: {e}")
                            traceback.print_exc()
                            
                except KeyboardInterrupt:
                    print("\nTraining interrupted by user")
                    break
                except Exception as e:
                    print(f"\nERROR during training epoch {epoch + 1}: {e}")
                    traceback.print_exc()
                    print("\nContinuing to next epoch...")
                    continue
            
        except Exception as e:
            print(f"\nFATAL ERROR during training: {e}")
            traceback.print_exc()
            return
        
        # Final evaluation
        if test_loaders:
            print("\n" + "="*70)
            print("Final Benchmark Evaluation")
            print("="*70)
            
            try:
                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}")
            except Exception as e:
                print(f"ERROR during final evaluation: {e}")
                traceback.print_exc()
        
        print("\nTraining and evaluation complete!")
        
    except Exception as e:
        print(f"\nFATAL ERROR: {e}")
        traceback.print_exc()
        sys.exit(1)


if __name__ == '__main__':
    main()



