WikiAI / generate.py
Reedy4687's picture
Upload 3 files
6838da9 verified
Raw
History Blame Contribute Delete
6.55 kB
#!/usr/bin/env python3
"""
Generate text from a checkpoint produced by train_wiki_gpt.py.
Only needs:
- ckpt.pt (contains weights + architecture config)
- this script
- pip install torch tiktoken
tiktoken downloads and caches the GPT-2 BPE encoding on first use, so the
machine running this needs internet access once if that cache isn't
already warm.
Usage:
python generate.py --ckpt out/ckpt.pt --prompt "The history of"
python generate.py --ckpt out/ckpt.pt --prompt "In 1969," --max-new-tokens 300 --temperature 0.9
"""
import argparse
import torch
import torch.nn as nn
import torch.nn.functional as F
# --------------------------------------------------------------------------
# Model definition — must match train_wiki_gpt.py exactly so state_dict
# keys line up. Duplicated here (rather than imported) so this script works
# standalone with just ckpt.pt, even on a machine without the training code.
# --------------------------------------------------------------------------
class CausalSelfAttention(nn.Module):
def __init__(self, n_embd, n_head, block_size, dropout):
super().__init__()
assert n_embd % n_head == 0
self.n_head = n_head
self.n_embd = n_embd
self.c_attn = nn.Linear(n_embd, 3 * n_embd, bias=False)
self.c_proj = nn.Linear(n_embd, n_embd, bias=False)
self.attn_dropout = dropout
self.resid_dropout = nn.Dropout(dropout)
def forward(self, x):
B, T, C = x.shape
q, k, v = self.c_attn(x).split(self.n_embd, dim=2)
q = q.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
k = k.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
v = v.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
y = F.scaled_dot_product_attention(q, k, v, is_causal=True, dropout_p=0.0)
y = y.transpose(1, 2).contiguous().view(B, T, C)
return self.resid_dropout(self.c_proj(y))
class MLP(nn.Module):
def __init__(self, n_embd, dropout):
super().__init__()
self.c_fc = nn.Linear(n_embd, 4 * n_embd, bias=False)
self.gelu = nn.GELU()
self.c_proj = nn.Linear(4 * n_embd, n_embd, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x):
return self.dropout(self.c_proj(self.gelu(self.c_fc(x))))
class Block(nn.Module):
def __init__(self, n_embd, n_head, block_size, dropout):
super().__init__()
self.ln_1 = nn.LayerNorm(n_embd)
self.attn = CausalSelfAttention(n_embd, n_head, block_size, dropout)
self.ln_2 = nn.LayerNorm(n_embd)
self.mlp = MLP(n_embd, dropout)
def forward(self, x):
x = x + self.attn(self.ln_1(x))
x = x + self.mlp(self.ln_2(x))
return x
class GPT(nn.Module):
def __init__(self, vocab_size, block_size, n_layer, n_head, n_embd, dropout=0.0):
super().__init__()
self.block_size = block_size
self.tok_emb = nn.Embedding(vocab_size, n_embd)
self.pos_emb = nn.Embedding(block_size, n_embd)
self.drop = nn.Dropout(dropout)
self.blocks = nn.ModuleList(
[Block(n_embd, n_head, block_size, dropout) for _ in range(n_layer)]
)
self.ln_f = nn.LayerNorm(n_embd)
self.head = nn.Linear(n_embd, vocab_size, bias=False)
self.tok_emb.weight = self.head.weight
def forward(self, idx):
B, T = idx.shape
pos = torch.arange(T, device=idx.device)
x = self.drop(self.tok_emb(idx) + self.pos_emb(pos))
for block in self.blocks:
x = block(x)
x = self.ln_f(x)
return self.head(x)
@torch.no_grad()
def generate(self, idx, max_new_tokens, temperature=0.8, top_k=50):
for _ in range(max_new_tokens):
idx_cond = idx[:, -self.block_size:]
logits = self(idx_cond)
logits = logits[:, -1, :] / temperature
if top_k is not None:
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < v[:, [-1]]] = -float("inf")
probs = F.softmax(logits, dim=-1)
next_id = torch.multinomial(probs, num_samples=1)
idx = torch.cat((idx, next_id), dim=1)
return idx
def strip_compile_prefix(state_dict):
"""torch.compile sometimes prefixes keys with '_orig_mod.' — strip it
so the state_dict loads into a plain (uncompiled) model."""
if any(k.startswith("_orig_mod.") for k in state_dict):
return {k.replace("_orig_mod.", "", 1): v for k, v in state_dict.items()}
return state_dict
def main():
parser = argparse.ArgumentParser(description="Generate text from a trained checkpoint.")
parser.add_argument("--ckpt", type=str, default="out/ckpt.pt", help="Path to ckpt.pt")
parser.add_argument("--prompt", type=str, default="The history of", help="Text prompt")
parser.add_argument("--max-new-tokens", type=int, default=200)
parser.add_argument("--temperature", type=float, default=0.8)
parser.add_argument("--top-k", type=int, default=50)
parser.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu")
args = parser.parse_args()
ckpt = torch.load(args.ckpt, map_location=args.device)
model_args = ckpt["args"]
model = GPT(
vocab_size=50304,
block_size=model_args["block_size"],
n_layer=model_args["n_layer"],
n_head=model_args["n_head"],
n_embd=model_args["n_embd"],
).to(args.device)
state_dict = strip_compile_prefix(ckpt["model"])
model.load_state_dict(state_dict)
model.eval()
print(f"Loaded checkpoint from iter {ckpt.get('iter', '?')} "
f"({model_args['n_layer']}L/{model_args['n_head']}H/{model_args['n_embd']}D)")
import tiktoken
enc = tiktoken.get_encoding("gpt2")
idx = torch.tensor([enc.encode_ordinary(args.prompt)], dtype=torch.long, device=args.device)
if args.device == "cuda":
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
out = model.generate(idx, args.max_new_tokens, args.temperature, args.top_k)
else:
out = model.generate(idx, args.max_new_tokens, args.temperature, args.top_k)
print("\n--- Generated text ---")
print(enc.decode(out[0].tolist()))
if __name__ == "__main__":
main()