import argparse
import os
import sys

import torch

from mk6.config import ModelConfig
from mk6.model import MK6Model
from mk6.repro import get_device
from mk6.tokenizer import ByteTokenizer


def parse_args():
    parser = argparse.ArgumentParser(description="Generate text with MK6")
    parser.add_argument("--checkpoint", required=True)
    parser.add_argument("--tokenizer", default=None)
    parser.add_argument("--prompt", default="")
    parser.add_argument("--max_new_tokens", type=int, default=128)
    parser.add_argument("--temperature", type=float, default=0.9)
    parser.add_argument("--top_p", type=float, default=0.95)
    parser.add_argument("--device", default=None)
    return parser.parse_args()


def main():
    if hasattr(sys.stdout, "reconfigure"):
        sys.stdout.reconfigure(encoding="utf-8", errors="replace")
    args = parse_args()
    device = get_device(args.device)
    checkpoint = torch.load(args.checkpoint, map_location=device)
    model_config = ModelConfig(**checkpoint["model_config"])
    model = MK6Model(model_config)
    model.load_state_dict(checkpoint["model_state_dict"])
    model.to(device)
    model.eval()

    tokenizer_path = args.tokenizer or os.path.join(os.path.dirname(args.checkpoint), "tokenizer.json")
    tokenizer = ByteTokenizer.load(tokenizer_path)
    prompt_ids = tokenizer.encode(args.prompt, add_bos=True, add_eos=False)
    input_ids = torch.tensor([prompt_ids], dtype=torch.long, device=device)

    generated = model.generate(
        input_ids=input_ids,
        max_new_tokens=args.max_new_tokens,
        temperature=args.temperature,
        top_p=args.top_p,
        eos_token_id=tokenizer.eos_token_id,
    )
    print(tokenizer.decode(generated[0].tolist()))


if __name__ == "__main__":
    main()
