69 lines
2.5 KiB
Python
69 lines
2.5 KiB
Python
#!/usr/bin/env python3
|
|
"""Generate text from a trained checkpoint."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
|
|
import torch
|
|
from transformers import AutoTokenizer
|
|
|
|
from model import GPT, GPTConfig
|
|
|
|
|
|
def load_checkpoint(path: str, map_location: str | torch.device = "cpu"):
|
|
try:
|
|
return torch.load(path, map_location=map_location, weights_only=False)
|
|
except TypeError:
|
|
return torch.load(path, map_location=map_location)
|
|
|
|
|
|
def main() -> None:
|
|
p = argparse.ArgumentParser()
|
|
p.add_argument("--checkpoint", type=str, required=True)
|
|
p.add_argument("--tokenizer_path", type=str, required=True)
|
|
p.add_argument("--prompt", type=str, default="The purpose of education is")
|
|
p.add_argument("--max_new_tokens", type=int, default=128)
|
|
p.add_argument("--temperature", type=float, default=0.8)
|
|
p.add_argument("--top_k", type=int, default=50)
|
|
p.add_argument("--top_p", type=float, default=0.95)
|
|
p.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu")
|
|
p.add_argument("--precision", choices=["bf16", "fp16", "fp32"], default="bf16")
|
|
p.add_argument("--compile", action="store_true")
|
|
args = p.parse_args()
|
|
|
|
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_path, use_fast=True)
|
|
ckpt = load_checkpoint(args.checkpoint, map_location="cpu")
|
|
config = GPTConfig(**ckpt["model_config"])
|
|
if config.vocab_size != len(tokenizer):
|
|
raise ValueError(f"checkpoint vocab_size={config.vocab_size}, tokenizer vocab_size={len(tokenizer)}")
|
|
|
|
device = torch.device(args.device)
|
|
model = GPT(config)
|
|
model.load_state_dict(ckpt["model"], strict=True)
|
|
model.to(device)
|
|
model.eval()
|
|
if args.compile:
|
|
model = torch.compile(model)
|
|
|
|
ids = tokenizer.encode(args.prompt, add_special_tokens=False)
|
|
x = torch.tensor([ids], dtype=torch.long, device=device)
|
|
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[args.precision]
|
|
use_amp = args.precision != "fp32" and device.type == "cuda"
|
|
|
|
with torch.no_grad():
|
|
with torch.autocast(device_type=device.type, dtype=dtype, enabled=use_amp):
|
|
y = model.generate(
|
|
x,
|
|
max_new_tokens=args.max_new_tokens,
|
|
temperature=args.temperature,
|
|
top_k=args.top_k,
|
|
top_p=args.top_p,
|
|
eos_token_id=tokenizer.eos_token_id,
|
|
)
|
|
print(tokenizer.decode(y[0].tolist(), skip_special_tokens=True))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|