"""
Text generation script using the trained model.
"""

import torch
import argparse
import sys
import os

sys.path.append(os.path.dirname(os.path.abspath(__file__)))

from config.config import Config
from model.architecture import NovelAIModel
from utils.tokenizer import SimpleTokenizer
import json


def main():
    parser = argparse.ArgumentParser(description='Generate text using the trained model')
    parser.add_argument('--checkpoint', type=str, required=True, help='Path to model checkpoint')
    parser.add_argument('--prompt', type=str, required=True, help='Input prompt')
    parser.add_argument('--max_new_tokens', type=int, default=50, help='Maximum tokens to generate')
    parser.add_argument('--temperature', type=float, default=1.0, help='Sampling temperature')
    parser.add_argument('--top_k', type=int, default=None, help='Top-k sampling')
    parser.add_argument('--top_p', type=float, default=1.0, help='Nucleus sampling')
    parser.add_argument('--seed', type=int, default=42, help='Random seed')
    
    args = parser.parse_args()
    
    # Set seed
    torch.manual_seed(args.seed)
    if torch.cuda.is_available():
        torch.cuda.manual_seed_all(args.seed)
    
    # Load checkpoint
    print(f"Loading checkpoint from {args.checkpoint}...")
    checkpoint = torch.load(args.checkpoint, map_location='cpu')
    
    # Load config if available
    config_dir = os.path.dirname(args.checkpoint)
    config_path = os.path.join(config_dir, 'config.json')
    if os.path.exists(config_path):
        config = Config.load(config_path)
    else:
        # Try to get config from checkpoint or use defaults
        if 'config' in checkpoint:
            config = Config.from_dict(checkpoint['config'])
        else:
            config = Config.default()
            # Update from checkpoint metadata if available
            if 'train_metrics' in checkpoint:
                # Infer some settings from checkpoint
                pass
    
    # Set device
    device = 'cuda' if torch.cuda.is_available() else 'cpu'
    print(f"Using device: {device}")
    
    # Create tokenizer (would need to be saved with model in production)
    # For now, create a basic one
    tokenizer = SimpleTokenizer(is_char_level=False)
    
    # Rebuild tokenizer vocab if saved (in production, save/load tokenizer separately)
    # For now, we'll need to handle this - in a real implementation,
    # you'd save the tokenizer with the model
    
    # Create model
    print("Creating model...")
    formula_config = {
        'w1': config.model.formula_w1,
        'w2': config.model.formula_w2,
        'w3': config.model.formula_w3,
        'lambda_decay': config.model.formula_lambda,
        'k_fatigue': config.model.formula_k,
        'learnable_weights': config.model.formula_learnable_weights,
        'learnable_decay': config.model.formula_learnable_decay,
    }
    
    model = NovelAIModel(
        vocab_size=config.model.vocab_size,
        embedding_dim=config.model.embedding_dim,
        num_layers=config.model.num_layers,
        num_heads=config.model.num_heads,
        feedforward_dim=config.model.feedforward_dim,
        max_seq_length=config.model.max_seq_length,
        dropout=config.model.dropout,
        formula_config=formula_config,
        use_memory_buffer=config.model.use_memory_buffer,
        memory_buffer_size=config.model.memory_buffer_size,
        use_formula_attention=config.model.use_formula_attention,
    )
    
    # Load weights
    model.load_state_dict(checkpoint['model_state_dict'])
    model = model.to(device)
    model.eval()
    
    # Tokenize prompt
    # Note: In production, you'd load the same tokenizer used for training
    prompt_ids = tokenizer.encode(args.prompt, add_special_tokens=False)
    input_ids = torch.tensor([prompt_ids], dtype=torch.long).to(device)
    
    print(f"\nPrompt: {args.prompt}")
    print(f"Tokenized length: {len(prompt_ids)}")
    print("\nGenerating...")
    
    # Generate
    with torch.no_grad():
        generated = model.generate(
            input_ids=input_ids,
            max_new_tokens=args.max_new_tokens,
            temperature=args.temperature,
            top_k=args.top_k,
            top_p=args.top_p,
            do_sample=args.temperature > 0,
            pad_token_id=tokenizer.pad_token_id
        )
    
    # Decode
    generated_ids = generated[0].cpu().tolist()
    generated_text = tokenizer.decode(generated_ids, skip_special_tokens=True)
    
    print(f"\nGenerated text:\n{generated_text}")


if __name__ == '__main__':
    main()



