"""
Final working training script - all fixes applied.
Runs training with benchmark evaluation.
"""

import torch
import sys
import os
import json
from datetime import datetime

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

os.environ['CUDA_LAUNCH_BLOCKING'] = '0'
torch.backends.cudnn.benchmark = True

sys.path.insert(0, '.')

from config.config import Config
from model.architecture import NovelAIModel
from training.trainer import Trainer
from training.validation import validate_training_setup
from training.logger import setup_logging
from utils.data_loader import create_data_loaders, load_text_file, TextDataset
from utils.tokenizer import SimpleTokenizer
from torch.utils.data import DataLoader

def main():
    print("="*70)
    print("NOVEL AI MODEL - BENCHMARK TRAINING")
    print("="*70)
    
    # Setup logger
    logger = setup_logging(log_dir='./logs')
    logger.info("Starting benchmark training")
    
    try:
        # Device
        device = 'cuda' if torch.cuda.is_available() else 'cpu'
        logger.info(f"Device: {device}")
        if device == 'cuda':
            try:
                logger.info(f"GPU: {torch.cuda.get_device_name(0)}")
                torch.cuda.empty_cache()
            except Exception as e:
                logger.warning(f"GPU detection error: {e}")
        
        # Load data
        logger.info("Loading data files...")
        train_path = './data/sample_train.txt'
        val_path = './data/sample_valid.txt'
        test_path = './data/sample_test.txt'
        
        if not os.path.exists(train_path):
            error_msg = (
                f"Training data file not found: {train_path}\n"
                f"  → Create sample data: python scripts/create_sample_data.py\n"
                f"  → Or download data: python scripts/download_wikitext.py"
            )
            logger.critical(error_msg)
            raise FileNotFoundError(error_msg)
        
        train_texts = load_text_file(train_path)
        val_texts = load_text_file(val_path) if os.path.exists(val_path) else []
        test_texts = load_text_file(test_path) if os.path.exists(test_path) else []
        
        logger.info(f"Loaded data: Train={len(train_texts)}, Val={len(val_texts)}, Test={len(test_texts)}")
        
        if len(train_texts) == 0:
            error_msg = "Training data is empty. Cannot train without data."
            logger.critical(error_msg)
            raise ValueError(error_msg)
        
        # Create tokenizer
        logger.info("Creating tokenizer...")
        tokenizer = SimpleTokenizer(is_char_level=False)
        tokenizer.build_vocab(train_texts + val_texts + test_texts)
        logger.info(f"Vocab size: {tokenizer.vocab_size}")
        
        # Create data loaders
        logger.info("Creating data loaders...")
        try:
            train_loader, val_loader, _ = create_data_loaders(
                train_texts=train_texts,
                val_texts=val_texts,
                tokenizer=tokenizer,
                max_length=256,  # Reduced for memory
                batch_size=2,  # Reduced for memory
                num_workers=0,
                build_vocab=False
            )
            
            # Test loader
            if test_texts:
                test_dataset = TextDataset(test_texts, tokenizer, max_length=256)
                test_loader = DataLoader(test_dataset, batch_size=2, shuffle=False, num_workers=0)
            else:
                test_loader = None
            
            logger.info(f"Train batches: {len(train_loader)}, Val: {len(val_loader) if val_loader else 0}, Test: {len(test_loader) if test_loader else 0}")
        except Exception as e:
            error_msg = f"Failed to create data loaders: {e}"
            logger.critical(error_msg, exc_info=True)
            raise
        
        # Validate setup before training
        model_config = {
            'vocab_size': tokenizer.vocab_size,
            'embedding_dim': 256,
            'num_layers': 4,
            'num_heads': 8,
            'max_seq_length': 256,  # Reduced for memory
        }
        training_config = {
            'batch_size': 2,  # Reduced for memory
            'learning_rate': 1e-3,
            'num_epochs': 5,
        }
        data_paths = {
            'train': train_path,
            'val': val_path if os.path.exists(val_path) else None,
            'test': test_path if os.path.exists(test_path) else None,
        }
        
        logger.info("Running pre-flight validation...")
        is_valid = validate_training_setup(
            data_paths=data_paths,
            model_config=model_config,
            training_config=training_config,
            checkpoint_dir='./checkpoints/benchmark_run'
        )
        
        if not is_valid:
            logger.critical("Pre-flight validation failed. Please fix errors before training.")
            raise RuntimeError("Validation failed. Check output above for details.")
        
        # Create model
        logger.info("Creating model...")
        try:
            model = NovelAIModel(
                vocab_size=tokenizer.vocab_size,
                embedding_dim=256,  # Smaller for faster training
                num_layers=4,
                num_heads=8,
                max_seq_length=256,  # Reduced for memory
                use_memory_buffer=True,  # Re-enabled with CUDA-compatible implementation
                formula_config={}
            )
            
            model = model.to(device)
            param_count = sum(p.numel() for p in model.parameters())
            logger.info(f"Model created: {param_count:,} parameters")
        except Exception as e:
            error_msg = f"Failed to create model: {e}"
            logger.critical(error_msg, exc_info=True)
            raise
        
        # Trainer
        logger.info("Initializing trainer...")
        trainer = Trainer(
            model=model,
            train_loader=train_loader,
            val_loader=val_loader,
            learning_rate=1e-3,
            weight_decay=0.01,
            device=device,
            save_dir='./checkpoints/benchmark_run',
            log_interval=10,
            logger=logger  # Use shared logger
        )
        
        # Train
        logger.info("="*70)
        logger.info("STARTING TRAINING")
        logger.info("="*70)
        
        num_epochs = 5
        
        for epoch in range(num_epochs):
            logger.info(f"\nEpoch {epoch + 1}/{num_epochs}")
            
            try:
                train_metrics = trainer.train_epoch()
                logger.info(f"Train Loss: {train_metrics['train_loss']:.4f}")
                
                val_metrics = trainer.validate()
                if val_metrics:
                    val_loss = val_metrics.get('val_loss', float('inf'))
                    logger.info(f"Val Loss: {val_loss:.4f}")
                
                # Evaluate on test set every 2 epochs
                if test_loader and (epoch + 1) % 2 == 0:
                    logger.info("Running benchmark evaluation...")
                    try:
                        from evaluation.benchmarks import BenchmarkSuite
                        suite = BenchmarkSuite(model, tokenizer, device=device)
                        
                        results = suite.evaluate_dataset(test_loader, 'sample_test')
                        
                        results_path = f'./checkpoints/benchmark_run/eval_epoch_{epoch+1}.json'
                        with open(results_path, 'w') as f:
                            json.dump(results, f, indent=2)
                        logger.info(f"Saved results to {results_path}")
                    except Exception as e:
                        logger.error(f"Evaluation failed: {e}", exc_info=True)
                
            except KeyboardInterrupt:
                logger.info("Training interrupted by user")
                raise
            except Exception as e:
                logger.error(f"Error in epoch {epoch + 1}: {e}", exc_info=True)
                logger.warning(f"Continuing to next epoch...")
                continue
        
        # Final evaluation
        if test_loader:
            logger.info("="*70)
            logger.info("FINAL BENCHMARK EVALUATION")
            logger.info("="*70)
            
            try:
                from evaluation.benchmarks import BenchmarkSuite
                suite = BenchmarkSuite(model, tokenizer, device=device)
                final_results = suite.evaluate_dataset(test_loader, 'sample_test')
                
                report = suite.generate_comparison_report({'sample_test': final_results})
                logger.info(report)
                
                with open('./checkpoints/benchmark_run/final_results.json', 'w') as f:
                    json.dump(final_results, f, indent=2)
                
                logger.info("Final results saved")
            except Exception as e:
                logger.error(f"Final evaluation failed: {e}", exc_info=True)
        
        logger.info("\nTraining complete!")
        logger.info(f"Results saved in: ./checkpoints/benchmark_run/")
        logger.info(f"Logs saved in: {logger.get_log_file()}")
        
    except KeyboardInterrupt:
        logger.info("\nTraining interrupted by user")
        sys.exit(0)
    except Exception as e:
        logger.critical(f"\nFATAL ERROR: {e}", exc_info=True)
        logger.critical(f"Check log file for details: {logger.get_log_file()}")
        sys.exit(1)

if __name__ == '__main__':
    main()

