|
|
| """
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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() |