| """
|
| 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
|
|
|
|
|
| encoded = tokenizer.encode(prompt)
|
| input_ids = torch.tensor([encoded.ids], dtype=torch.long, device=device)
|
|
|
|
|
| with torch.no_grad():
|
| output_ids = model.generate(
|
| input_ids,
|
| max_new_tokens=max_new_tokens,
|
| temperature=temperature,
|
| top_k=top_k,
|
| )
|
|
|
|
|
| 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}")
|
|
|
|
|
| tokenizer = Tokenizer.from_file(TOKENIZER_PATH)
|
|
|
|
|
| model, config = load_model(args.checkpoint, device)
|
|
|
|
|
| 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()
|
|
|