File size: 3,540 Bytes
3738348 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 | """
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()
|