Download src/generate.py from Reizxn/makeitwork1: direct link, hf CLI and curl.
- Browser
- Download file 3.54 kB
-
https://huggingface.co/Reizxn/makeitwork1/resolve/main/src/generate.py
- Command line
-
hf download hf://Reizxn/makeitwork1/src/generate.py
-
curl -L -o generate.py https://huggingface.co/Reizxn/makeitwork1/resolve/main/src/generate.py
3.54 kB
| """ | |
| Inference and generation script for Retriever500M. | |
| Loads a checkpoint and generates text to verify the model has learned | |
| code structure during base pretraining. | |
| Usage: | |
| python src/generate.py [--checkpoint PATH] [--prompt "text"] [--tokens N] | |
| """ | |
| import argparse | |
| import os | |
| import sys | |
| import torch | |
| sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) | |
| from model import ModelConfig, Retriever500M | |
| from tokenizers import Tokenizer | |
| PROJECT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) | |
| CHECKPOINT_DIR = os.path.join(PROJECT_DIR, "checkpoints") | |
| TOKENIZER_PATH = os.path.join(PROJECT_DIR, "tokenizer", "tokenizer.json") | |
| def load_model(checkpoint_path: str, device: torch.device) -> tuple[Retriever500M, ModelConfig]: | |
| """Load model from checkpoint.""" | |
| print(f"Loading checkpoint: {checkpoint_path}") | |
| ckpt = torch.load(checkpoint_path, map_location=device, weights_only=False) | |
| config = ModelConfig(**ckpt["config"]) | |
| model = Retriever500M(config).to(device) | |
| model.load_state_dict(ckpt["model_state_dict"]) | |
| model.eval() | |
| print(f" Step: {ckpt.get('step', '?')}") | |
| print(f" Loss: {ckpt.get('loss', '?')}") | |
| print(f" Params: {model.count_parameters() / 1e6:.1f}M") | |
| return model, config | |
| def generate( | |
| model: Retriever500M, | |
| tokenizer: Tokenizer, | |
| prompt: str, | |
| max_new_tokens: int = 128, | |
| temperature: float = 0.8, | |
| top_k: int = 50, | |
| device: torch.device = None, | |
| ) -> str: | |
| """Generate text from a prompt.""" | |
| if device is None: | |
| device = next(model.parameters()).device | |
| # Encode prompt | |
| encoded = tokenizer.encode(prompt) | |
| input_ids = torch.tensor([encoded.ids], dtype=torch.long, device=device) | |
| # Generate | |
| with torch.no_grad(): | |
| output_ids = model.generate( | |
| input_ids, | |
| max_new_tokens=max_new_tokens, | |
| temperature=temperature, | |
| top_k=top_k, | |
| ) | |
| # Decode | |
| output_text = tokenizer.decode(output_ids[0].tolist()) | |
| return output_text | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Generate text with Retriever500M") | |
| parser.add_argument("--checkpoint", type=str, default=os.path.join(CHECKPOINT_DIR, "latest.pt")) | |
| parser.add_argument("--prompt", type=str, default="def fibonacci(n):\n ", help="Generation prompt") | |
| parser.add_argument("--tokens", type=int, default=128, help="Max new tokens") | |
| parser.add_argument("--temperature", type=float, default=0.8) | |
| parser.add_argument("--top_k", type=int, default=50) | |
| args = parser.parse_args() | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| print(f"Device: {device}") | |
| # Load tokenizer | |
| tokenizer = Tokenizer.from_file(TOKENIZER_PATH) | |
| # Load model | |
| model, config = load_model(args.checkpoint, device) | |
| # Generate | |
| prompts = [ | |
| args.prompt, | |
| "def quicksort(arr):\n ", | |
| "function fetchData(url) {\n ", | |
| "import torch\nimport torch.nn as nn\n\nclass Model(nn.Module):\n ", | |
| "fn main() {\n println!", | |
| ] | |
| print("\n" + "=" * 60) | |
| print("GENERATION SAMPLES") | |
| print("=" * 60) | |
| for prompt in prompts: | |
| print(f"\n--- Prompt: {prompt!r} ---") | |
| text = generate(model, tokenizer, prompt, args.tokens, args.temperature, args.top_k, device) | |
| print(text) | |
| print("-" * 60) | |
| if __name__ == "__main__": | |
| main() | |