""" Script 3: Interactive Text Generation & Inference for Nova 1.0 """ import os import argparse import torch from config.model_config import Nova1Config from src.tokenizer.bpe_tokenizer import Nova1Tokenizer from src.model.nova1_hrm import Nova1HRM from src.inference.generate import Nova1Generator def main(): parser = argparse.ArgumentParser(description="Generate text with Nova 1.0") parser.add_argument("--checkpoint", type=str, default="checkpoints/nova1_final.pt", help="Path to Nova 1.0 checkpoint") default_tok = "checkpoints/nova1_gemini_tokenizer.json" if os.path.exists("checkpoints/nova1_gemini_tokenizer.json") else "checkpoints/nova1_tokenizer.json" parser.add_argument("--tokenizer_path", type=str, default=default_tok, help="Path to tokenizer JSON") parser.add_argument("--prompt", type=str, default="Once upon a time in a small village,", help="Generation prompt") parser.add_argument("--max_tokens", type=int, default=60, help="Max new tokens to generate") parser.add_argument("--temperature", type=float, default=0.8, help="Sampling temperature") parser.add_argument("--top_k", type=int, default=40, help="Top-K sampling") parser.add_argument("--top_p", type=float, default=0.9, help="Top-P nucleus sampling") parser.add_argument("--N_cycles", type=int, default=None, help="Test-time reasoning depth N cycles") parser.add_argument("--T_steps", type=int, default=None, help="Test-time low-level steps T") args = parser.parse_args() if not os.path.exists(args.tokenizer_path): print(f"Error: Tokenizer not found at '{args.tokenizer_path}'.") return tokenizer = Nova1Tokenizer.load(args.tokenizer_path) if not os.path.exists(args.checkpoint): print(f"Error: Model checkpoint not found at '{args.checkpoint}'.") return checkpoint = torch.load(args.checkpoint, map_location="cpu", weights_only=False) config: Nova1Config = checkpoint.get("config", Nova1Config(vocab_size=tokenizer.vocab_size)) model = Nova1HRM(config) model.load_state_dict(checkpoint["model_state"]) print(f"Loaded Nova 1.0 model checkpoint from '{args.checkpoint}'.") generator = Nova1Generator(model=model, tokenizer=tokenizer, config=config) print(f"\n--- Nova 1.0 Text Generation ---") print(f"Prompt: {args.prompt}") output_text = generator.generate( prompt=args.prompt, max_new_tokens=args.max_tokens, temperature=args.temperature, top_k=args.top_k, top_p=args.top_p, N_cycles=args.N_cycles, T_steps=args.T_steps ) print("\n--- Output ---") print(output_text) print("--------------\n") if __name__ == "__main__": main()