"""
Main training script for MK3.

Usage:
    python train.py --train_data data/train.txt --val_data data/val.txt

For testing with sample data:
    python train.py --test_mode
"""

import argparse
import torch
import sys
import os

from training.config import ModelConfig, TrainingConfig
from training.trainer import CompleteMK3Trainer
from utils.tokenizer import SimpleTokenizer
from utils.data_loader import TextDataset, TextDataLoader, load_text_file, create_train_val_split
from utils.reproducibility import print_system_info, set_seed


def parse_args():
    parser = argparse.ArgumentParser(description='Train MK3 Continuous Autoregressive Model')

    # Data arguments
    parser.add_argument('--train_data', type=str, help='Path to training data')
    parser.add_argument('--val_data', type=str, help='Path to validation data')
    parser.add_argument('--test_mode', action='store_true', help='Run in test mode with synthetic data')

    # Model arguments
    parser.add_argument('--vocab_size', type=int, default=10000, help='Vocabulary size')
    parser.add_argument('--embedding_dim', type=int, default=768, help='Embedding dimension')
    parser.add_argument('--vector_dim', type=int, default=1024, help='Continuous vector dimension')
    parser.add_argument('--chunk_size', type=int, default=8, help='Tokens per continuous vector')
    parser.add_argument('--num_layers', type=int, default=12, help='Number of transformer layers')
    parser.add_argument('--num_heads', type=int, default=8, help='Number of attention heads')

    # Training arguments
    parser.add_argument('--batch_size', type=int, default=32, help='Batch size')
    parser.add_argument('--learning_rate', type=float, default=1e-4, help='Learning rate')
    parser.add_argument('--num_epochs', type=int, default=10, help='Number of epochs')
    parser.add_argument('--max_seq_length', type=int, default=512, help='Maximum sequence length')

    # Other arguments
    parser.add_argument('--checkpoint_dir', type=str, default='./checkpoints', help='Checkpoint directory')
    parser.add_argument('--seed', type=int, default=42, help='Random seed')
    parser.add_argument('--device', type=str, default=None, help='Device (cuda/cpu)')

    return parser.parse_args()


def create_synthetic_data():
    """Create synthetic data for testing."""
    print("Creating synthetic test data...")

    # Simple repeating patterns for easy learning
    train_texts = [
        "The quick brown fox jumps over the lazy dog. " * 10,
        "Hello world, this is a test. " * 10,
        "Machine learning is fascinating and powerful. " * 10,
        "Natural language processing enables computers to understand text. " * 10,
        "Deep learning models can learn complex patterns. " * 10,
    ] * 20  # Repeat for more data

    val_texts = [
        "The quick brown fox jumps over the lazy dog. " * 5,
        "Hello world, this is a test. " * 5,
    ] * 5

    return train_texts, val_texts


def main():
    args = parse_args()

    # Print system info
    print_system_info()

    # Set seed
    set_seed(args.seed)

    # Load or create data
    if args.test_mode:
        print("\n=== RUNNING IN TEST MODE ===\n")
        train_texts, val_texts = create_synthetic_data()
    else:
        if not args.train_data:
            raise ValueError("Must provide --train_data or use --test_mode")

        print(f"Loading training data from {args.train_data}")
        train_texts = load_text_file(args.train_data)

        if args.val_data:
            print(f"Loading validation data from {args.val_data}")
            val_texts = load_text_file(args.val_data)
        else:
            print("Splitting training data for validation")
            train_texts, val_texts = create_train_val_split(train_texts, val_ratio=0.1, seed=args.seed)

    print(f"Training examples: {len(train_texts)}")
    print(f"Validation examples: {len(val_texts)}")

    # Build tokenizer
    print("\nBuilding tokenizer...")
    tokenizer = SimpleTokenizer(vocab_size=args.vocab_size, min_frequency=2)
    tokenizer.train(train_texts)
    print(f"Vocabulary size: {tokenizer.vocab_size}")

    # Save tokenizer
    os.makedirs(args.checkpoint_dir, exist_ok=True)
    tokenizer.save(os.path.join(args.checkpoint_dir, 'tokenizer.json'))

    # Create datasets
    print("\nCreating datasets...")
    train_dataset = TextDataset(
        train_texts,
        tokenizer,
        max_length=args.max_seq_length
    )
    val_dataset = TextDataset(
        val_texts,
        tokenizer,
        max_length=args.max_seq_length
    )

    # Create data loaders
    train_loader = TextDataLoader(
        train_dataset,
        batch_size=args.batch_size,
        shuffle=True
    )
    val_loader = TextDataLoader(
        val_dataset,
        batch_size=args.batch_size,
        shuffle=False
    )

    print(f"Training batches: {len(train_loader)}")
    print(f"Validation batches: {len(val_loader)}")

    # Create configs
    model_config = ModelConfig(
        vocab_size=len(tokenizer.token_to_id),
        embedding_dim=args.embedding_dim,
        vector_dim=args.vector_dim,
        chunk_size=args.chunk_size,
        num_layers=args.num_layers,
        num_heads=args.num_heads,
        max_seq_length=args.max_seq_length
    )

    training_config = TrainingConfig(
        batch_size=args.batch_size,
        learning_rate=args.learning_rate,
        num_epochs=args.num_epochs,
        checkpoint_dir=args.checkpoint_dir,
        seed=args.seed
    )

    # Special settings for test mode
    if args.test_mode:
        training_config.num_epochs = 3
        training_config.autoencoder_steps = 500
        training_config.save_every = 100
        training_config.eval_every = 50

    print("\n" + "=" * 60)
    print("Model Configuration:")
    print("=" * 60)
    for key, value in model_config.to_dict().items():
        print(f"{key}: {value}")

    print("\n" + "=" * 60)
    print("Training Configuration:")
    print("=" * 60)
    print(training_config)

    # Get device
    if args.device:
        device = torch.device(args.device)
    else:
        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

    # Create trainer
    print(f"\n=== Initializing Trainer on {device} ===\n")
    trainer = CompleteMK3Trainer(
        model_config=model_config,
        training_config=training_config,
        train_loader=train_loader,
        val_loader=val_loader,
        device=device
    )

    # Train
    print("\n=== Starting Training ===\n")
    try:
        trainer.train()
    except KeyboardInterrupt:
        print("\n\nTraining interrupted by user")
        print("Saving checkpoint...")
        trainer.save_checkpoint('interrupted_checkpoint.pt')
        print("Checkpoint saved")

    print("\n=== Training Complete ===\n")
    print(f"Checkpoints saved to: {args.checkpoint_dir}")
    print(f"Best validation loss: {trainer.best_val_loss:.4f}")


if __name__ == '__main__':
    main()
