"""
Simplified training script with clear error handling - will use GPU automatically.
"""

import torch
import sys
import os

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


def main():
    # Setup
    device = 'cuda' if torch.cuda.is_available() else 'cpu'
    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('./data/sample_train.txt')
    val_texts = load_text_file('./data/sample_valid.txt')
    print(f"Train: {len(train_texts)}, Val: {len(val_texts)}")
    
    # Create data loaders
    train_loader, val_loader, tokenizer = create_data_loaders(
        train_texts=train_texts,
        val_texts=val_texts,
        max_length=512,
        batch_size=4,
        num_workers=0
    )
    
    # Create model
    config = Config.default()
    config.model.vocab_size = tokenizer.vocab_size
    
    model = NovelAIModel(
        vocab_size=config.model.vocab_size,
        embedding_dim=256,  # Smaller for testing
        num_layers=4,
        max_seq_length=512,
        formula_config={}
    )
    
    model = model.to(device)
    print(f"Model on {device}, params: {sum(p.numel() for p in model.parameters()):,}")
    
    # Create trainer
    trainer = Trainer(
        model=model,
        train_loader=train_loader,
        val_loader=val_loader,
        learning_rate=1e-3,
        device=device,
        save_dir='./checkpoints/simple_run'
    )
    
    # Train one epoch
    print("\n" + "="*70)
    print("Starting Training (1 epoch)")
    print("="*70)
    
    try:
        train_metrics = trainer.train_epoch()
        print(f"\nTraining complete! Loss: {train_metrics['train_loss']:.4f}")
        
        val_metrics = trainer.validate()
        print(f"Validation complete! Loss: {val_metrics.get('val_loss', 'N/A')}")
        
    except Exception as e:
        print(f"\nERROR: {e}")
        import traceback
        traceback.print_exc()


if __name__ == '__main__':
    main()



