"""
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')
    parser.add_argument('--warmup_steps', type=int, default=None, help='Number of LR warmup steps')
    parser.add_argument('--gradient_accumulation_steps', type=int, default=None, help='Gradient accumulation steps')
    parser.add_argument('--autoencoder_steps', type=int, default=None, help='Override autoencoder pretraining steps')
    parser.add_argument('--autoencoder_target_accuracy', type=float, default=None,
                        help='Override autoencoder target accuracy')
    parser.add_argument('--skip_autoencoder_pretrain', action='store_true',
                        help='Skip autoencoder pretraining stage')
    parser.add_argument('--freeze_autoencoder_after_pretrain', action='store_true',
                        help='Freeze autoencoder parameters after pretraining')
    parser.add_argument('--autoencoder_freeze_steps', type=int, default=None,
                        help='Number of Stage 2 steps to wait before freezing autoencoder')
    parser.add_argument('--save_every', type=int, default=None, help='Checkpoint save frequency (steps)')
    parser.add_argument('--eval_every', type=int, default=None, help='Evaluation frequency (steps)')
    parser.add_argument('--keep_n_checkpoints', type=int, default=None, help='Number of checkpoints to retain')
    parser.add_argument('--disable_curriculum', action='store_true', help='Disable curriculum learning')
    parser.add_argument('--initial_seq_length', type=int, default=None,
                        help='Initial sequence length for curriculum learning')
    parser.add_argument('--no_amp', action='store_true', help='Disable automatic mixed precision (AMP)')
    parser.add_argument('--precision', type=str, default=None,
                        help='Override mixed precision mode (fp32, mixed_bf16, mixed_fp16, etc.)')
    parser.add_argument('--verbose_logging', action='store_true', help='Enable per-step verbose training logs')
    parser.add_argument('--no_eval_progress_bar', action='store_true', help='Disable evaluation progress bar')
    parser.add_argument('--eval_log_every', type=int, default=None, help='Log every N validation batches during evaluation')
    parser.add_argument('--eval_max_batches', type=int, default=None,
                        help='Maximum number of validation batches per evaluation (0 skips)')
    parser.add_argument('--eval_summary_path', type=str, default=None,
                        help='Optional path to append JSONL evaluation summaries')

    # 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)')
    parser.add_argument('--resume_from', type=str, default=None,
                        help='Path to checkpoint file to resume training from')

    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()
    cli_args = set(sys.argv[1:])

    # 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
    )

    if args.warmup_steps is not None:
        training_config.warmup_steps = args.warmup_steps
    if args.gradient_accumulation_steps is not None:
        training_config.gradient_accumulation_steps = args.gradient_accumulation_steps
    if args.autoencoder_steps is not None:
        training_config.autoencoder_steps = args.autoencoder_steps
    if args.autoencoder_target_accuracy is not None:
        training_config.autoencoder_target_accuracy = args.autoencoder_target_accuracy
    if args.skip_autoencoder_pretrain:
        training_config.pretrain_autoencoder = False
    if args.freeze_autoencoder_after_pretrain:
        training_config.freeze_autoencoder_after_pretrain = True
    if args.autoencoder_freeze_steps is not None:
        training_config.autoencoder_freeze_steps = args.autoencoder_freeze_steps
    if args.save_every is not None:
        training_config.save_every = args.save_every
    if args.eval_every is not None:
        training_config.eval_every = args.eval_every
    if args.keep_n_checkpoints is not None:
        training_config.keep_n_checkpoints = args.keep_n_checkpoints
    if args.disable_curriculum:
        training_config.use_curriculum = False
    if args.initial_seq_length is not None:
        training_config.initial_seq_length = args.initial_seq_length
    if args.no_amp:
        training_config.use_amp = False
    if args.precision is not None:
        training_config.precision = args.precision
    if args.verbose_logging:
        training_config.verbose_logging = True
    if args.no_eval_progress_bar:
        training_config.eval_progress_bar = False
    if args.eval_log_every is not None:
        training_config.eval_log_every = args.eval_log_every
    if args.eval_max_batches is not None:
        training_config.eval_max_batches = args.eval_max_batches
    if args.eval_summary_path is not None:
        training_config.eval_summary_path = args.eval_summary_path
    if args.no_amp:
        training_config.use_amp = False
        if training_config.precision is None:
            training_config.precision = 'fp32'

    if training_config.precision:
        precision_lower = training_config.precision.lower()
        if precision_lower == 'fp32':
            training_config.use_amp = False
            training_config.precision = 'fp32'
        else:
            # For mixed precision modes ensure AMP is enabled unless explicitly disabled
            if getattr(training_config, 'use_amp', True):
                training_config.use_amp = True

    # Special settings for test mode
    if args.test_mode:
        if '--num_epochs' not in cli_args:
            training_config.num_epochs = 1
        if '--autoencoder_steps' not in cli_args:
            training_config.autoencoder_steps = 100
        if '--save_every' not in cli_args:
            training_config.save_every = 50
        if '--eval_every' not in cli_args:
            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
    )

    if args.resume_from:
        resume_path = args.resume_from
        if not os.path.isabs(resume_path):
            resume_path = os.path.join(args.checkpoint_dir, resume_path)
        if not os.path.exists(resume_path):
            raise FileNotFoundError(f"Resume checkpoint not found: {resume_path}")
        print(f"\n=== Resuming from checkpoint: {resume_path} ===\n")
        trainer.load_checkpoint(resume_path)

    # 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()
