"""
Main training script for the novel AI model.
"""

import torch
import argparse
import sys
import os

# Add project root to path
sys.path.append(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():
    parser = argparse.ArgumentParser(description='Train 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('--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')
    
    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 with command line arguments
    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}")
    
    # 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=config.model.use_memory_buffer,
        memory_buffer_size=config.model.memory_buffer_size,
        use_formula_attention=config.model.use_formula_attention,
    )
    
    # 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')}")
    
    # Train
    trainer.train(num_epochs=config.training.num_epochs)
    
    print("\nTraining complete!")


if __name__ == '__main__':
    main()



