Nova / scripts /generate_text.py
kings1's picture
Upload folder using huggingface_hub
23ea6bd verified
Raw History Blame Contribute Delete
2.74 kB
"""
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()