Download scripts/generate_text.py from kings1/Nova: direct link, hf CLI and curl.
- Browser
- Download file 2.74 kB
-
https://huggingface.co/kings1/Nova/resolve/main/scripts/generate_text.py
- Command line
-
hf download hf://kings1/Nova/scripts/generate_text.py
-
curl -L -o generate_text.py https://huggingface.co/kings1/Nova/resolve/main/scripts/generate_text.py
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() | |