import argparse
import torch

from mk6.config import ModelConfig, TrainingConfig
from mk6.data import (
    NextTokenDataset,
    build_dataloader,
    create_synthetic_corpus,
    create_train_val_split,
    load_texts,
)
from mk6.model import MK6Model
from mk6.repro import get_device, print_system_info
from mk6.tokenizer import ByteTokenizer
from mk6.trainer import Trainer


def parse_args():
    parser = argparse.ArgumentParser(description="Train MK6 salience-native model")
    parser.add_argument("--preset", default="small_5060", choices=["probe", "smoke", "small_5060", "h100_base"])
    parser.add_argument("--train_file", type=str, default=None)
    parser.add_argument("--val_file", type=str, default=None)
    parser.add_argument("--test_mode", action="store_true")
    parser.add_argument("--checkpoint_dir", type=str, default=None)
    parser.add_argument("--max_seq_len", type=int, default=None)
    parser.add_argument("--batch_size", type=int, default=None)
    parser.add_argument("--num_epochs", type=int, default=None)
    parser.add_argument("--learning_rate", type=float, default=None)
    parser.add_argument("--device", type=str, default=None)
    parser.add_argument("--resume_from", type=str, default=None)
    parser.add_argument("--max_train_steps", type=int, default=None)
    parser.add_argument("--disable_salience", action="store_true")
    parser.add_argument("--disable_moe", action="store_true")
    parser.add_argument("--selective_update_groups", type=int, default=None)
    parser.add_argument("--selective_update_active_groups", type=int, default=None)
    parser.add_argument("--step_timeout_seconds", type=float, default=None)
    parser.add_argument("--data_timeout_seconds", type=float, default=None)
    parser.add_argument("--heartbeat_path", type=str, default=None)
    parser.add_argument("--metrics_jsonl_path", type=str, default=None)
    parser.add_argument("--verbose_timing", action="store_true")
    parser.add_argument("--eval_max_batches", type=int, default=None)
    parser.add_argument("--final_eval_on_stop", action="store_true")
    return parser.parse_args()


def main():
    args = parse_args()
    print_system_info()

    model_config = ModelConfig.preset(args.preset)
    training_config = TrainingConfig.preset(args.preset)

    if args.checkpoint_dir is not None:
        training_config.checkpoint_dir = args.checkpoint_dir
    if args.max_seq_len is not None:
        model_config.max_seq_len = args.max_seq_len
    if args.batch_size is not None:
        training_config.batch_size = args.batch_size
    if args.num_epochs is not None:
        training_config.num_epochs = args.num_epochs
    if args.learning_rate is not None:
        training_config.learning_rate = args.learning_rate
    if args.max_train_steps is not None:
        training_config.max_train_steps = args.max_train_steps
    if args.selective_update_groups is not None:
        training_config.selective_update_groups = args.selective_update_groups
    if args.selective_update_active_groups is not None:
        training_config.selective_update_active_groups = args.selective_update_active_groups
    if args.step_timeout_seconds is not None:
        training_config.step_timeout_seconds = args.step_timeout_seconds
    if args.data_timeout_seconds is not None:
        training_config.data_timeout_seconds = args.data_timeout_seconds
    if args.heartbeat_path is not None:
        training_config.heartbeat_path = args.heartbeat_path
    if args.metrics_jsonl_path is not None:
        training_config.metrics_jsonl_path = args.metrics_jsonl_path
    if args.verbose_timing:
        training_config.verbose_timing = True
    if args.eval_max_batches is not None:
        training_config.eval_max_batches = args.eval_max_batches
    if args.final_eval_on_stop:
        training_config.final_eval_on_stop = True
    if args.disable_salience:
        model_config.use_salience = False
        training_config.salience_aux_weight = 0.0
    if args.disable_moe:
        model_config.use_moe = False
        training_config.moe_aux_weight = 0.0

    if args.test_mode:
        train_texts, val_texts = create_synthetic_corpus()
    else:
        if args.train_file is None:
            raise ValueError("Provide --train_file or use --test_mode")
        train_texts = load_texts(args.train_file)
        if args.val_file is not None:
            val_texts = load_texts(args.val_file)
        else:
            train_texts, val_texts = create_train_val_split(train_texts, val_ratio=0.1, seed=training_config.seed)

    tokenizer = ByteTokenizer()
    train_dataset = NextTokenDataset(train_texts, tokenizer, seq_len=model_config.max_seq_len)
    val_dataset = NextTokenDataset(val_texts, tokenizer, seq_len=model_config.max_seq_len)
    train_loader = build_dataloader(train_dataset, training_config.batch_size, True, training_config.num_workers)
    val_loader = build_dataloader(val_dataset, training_config.batch_size, False, training_config.num_workers)

    model = MK6Model(model_config)
    if model_config.compile_model and hasattr(torch, "compile"):
        model = torch.compile(model)

    trainer = Trainer(
        model=model,
        model_config=model_config,
        training_config=training_config,
        train_loader=train_loader,
        val_loader=val_loader,
        tokenizer=tokenizer,
        device=get_device(args.device),
    )
    if args.resume_from:
        trainer.load_checkpoint(args.resume_from)
    trainer.train()
    print(f"Training complete. Best validation loss: {trainer.best_val_loss:.4f}")
    print(f"Artifacts saved to: {training_config.checkpoint_dir}")


if __name__ == "__main__":
    main()
