"""
Text generation script for MK3.

Usage:
    python generate.py --checkpoint checkpoints/best_model.pt --prompt "Hello world"
"""

import argparse
import torch
import os
import sys

from training.config import ModelConfig
from calm.continuous_model import ContinuousAutoregressiveModel
from utils.tokenizer import SimpleTokenizer


def parse_args():
    parser = argparse.ArgumentParser(description='Generate text with MK3 model')

    parser.add_argument('--checkpoint', type=str, required=True, help='Path to model checkpoint')
    parser.add_argument('--tokenizer', type=str, default=None, help='Path to tokenizer (default: checkpoint_dir/tokenizer.json)')
    parser.add_argument('--prompt', type=str, default='', help='Input prompt')
    parser.add_argument('--max_new_vectors', type=int, default=10, help='Number of vectors to generate')
    parser.add_argument('--temperature', type=float, default=1.0, help='Sampling temperature')
    parser.add_argument('--top_p', type=float, default=0.9, help='Nucleus sampling threshold')
    parser.add_argument('--device', type=str, default=None, help='Device (cuda/cpu)')

    return parser.parse_args()


def main():
    args = parse_args()

    # Determine device
    if args.device:
        device = torch.device(args.device)
    else:
        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

    print(f"Using device: {device}")

    # Load checkpoint
    print(f"\nLoading checkpoint from {args.checkpoint}")
    checkpoint = torch.load(args.checkpoint, map_location=device)

    # Load model config
    model_config_dict = checkpoint['model_config']
    model_config = ModelConfig(**model_config_dict)

    print("\nModel configuration:")
    for key, value in model_config.to_dict().items():
        print(f"  {key}: {value}")

    # Load tokenizer
    if args.tokenizer:
        tokenizer_path = args.tokenizer
    else:
        checkpoint_dir = os.path.dirname(args.checkpoint)
        tokenizer_path = os.path.join(checkpoint_dir, 'tokenizer.json')

    print(f"\nLoading tokenizer from {tokenizer_path}")
    tokenizer = SimpleTokenizer()
    tokenizer.load(tokenizer_path)
    print(f"Vocabulary size: {tokenizer.vocab_size}")

    # Build model
    print("\nBuilding model...")
    salience_config = {
        'w1': model_config.salience_w1,
        'w2': model_config.salience_w2,
        'w3': model_config.salience_w3,
        'lambda_decay': model_config.salience_lambda,
        'k_fatigue': model_config.salience_k_fatigue,
        'temperature': model_config.salience_temperature,
        'learnable_weights': model_config.learnable_salience_params,
        'learnable_decay': model_config.learnable_salience_params,
        'learnable_temperature': model_config.learnable_salience_params,
        'normalization': model_config.salience_normalization,
        'use_dimensional_scaling': model_config.use_dimensional_scaling,
    }

    model = ContinuousAutoregressiveModel(
        vocab_size=model_config.vocab_size,
        embedding_dim=model_config.embedding_dim,
        vector_dim=model_config.vector_dim,
        chunk_size=model_config.chunk_size,
        num_layers=model_config.num_layers,
        num_heads=model_config.num_heads,
        feedforward_dim=model_config.feedforward_dim,
        max_seq_length=model_config.max_seq_length,
        dropout=model_config.dropout,
        salience_config=salience_config,
        memory_buffer_size=model_config.memory_buffer_size
    )

    # Load weights
    model.load_state_dict(checkpoint['model_state_dict'])
    model.to(device)
    model.eval()

    print("Model loaded successfully!")

    # Encode prompt
    if args.prompt:
        prompt_ids = tokenizer.encode(args.prompt, add_special_tokens=True)
    else:
        # Use BOS token if no prompt
        prompt_ids = [tokenizer.bos_token_id]

    print(f"\n{'=' * 60}")
    print(f"Prompt: {args.prompt if args.prompt else '<empty>'}")
    print(f"{'=' * 60}\n")

    # Convert to tensor
    prompt_tensor = torch.tensor([prompt_ids], dtype=torch.long, device=device)

    # Generate
    print("Generating...")
    with torch.no_grad():
        generated_ids = model.generate(
            initial_tokens=prompt_tensor,
            max_new_vectors=args.max_new_vectors,
            temperature=args.temperature,
            top_p=args.top_p
        )

    # Decode
    generated_text = tokenizer.decode(generated_ids[0], skip_special_tokens=True)

    print(f"\n{'=' * 60}")
    print("Generated Text:")
    print(f"{'=' * 60}")
    print(generated_text)
    print(f"{'=' * 60}\n")

    # Print info
    num_tokens_generated = generated_ids.shape[1] - len(prompt_ids)
    num_vectors = args.max_new_vectors
    speedup_factor = model_config.chunk_size

    print(f"Tokens generated: {num_tokens_generated}")
    print(f"Vectors generated: {num_vectors}")
    print(f"Speedup factor: {speedup_factor}x (compared to token-by-token)")


if __name__ == '__main__':
    main()
