""" chat.py Interactive REPL for generating YouTube-Shorts-style comments from your trained checkpoint (comment_gpt.pt). Type a prompt, get a comment back. Empty prompt = generate from scratch. Usage: python chat.py --ckpt comment_gpt.pt """ import argparse import torch import torch.nn as nn from torch.nn import functional as F BLOCK_SIZE = 128 N_LAYER = 6 N_HEAD = 6 N_EMBD = 384 DROPOUT = 0.1 def get_device(): if torch.backends.mps.is_available(): return "mps" if torch.cuda.is_available(): return "cuda" return "cpu" class Head(nn.Module): def __init__(self, head_size): super().__init__() self.key = nn.Linear(N_EMBD, head_size, bias=False) self.query = nn.Linear(N_EMBD, head_size, bias=False) self.value = nn.Linear(N_EMBD, head_size, bias=False) self.register_buffer("tril", torch.tril(torch.ones(BLOCK_SIZE, BLOCK_SIZE))) self.dropout = nn.Dropout(DROPOUT) def forward(self, x): B, T, C = x.shape k = self.key(x) q = self.query(x) wei = q @ k.transpose(-2, -1) * (C ** -0.5) wei = wei.masked_fill(self.tril[:T, :T] == 0, float("-inf")) wei = F.softmax(wei, dim=-1) wei = self.dropout(wei) v = self.value(x) return wei @ v class MultiHeadAttention(nn.Module): def __init__(self, num_heads, head_size): super().__init__() self.heads = nn.ModuleList([Head(head_size) for _ in range(num_heads)]) self.proj = nn.Linear(N_EMBD, N_EMBD) self.dropout = nn.Dropout(DROPOUT) def forward(self, x): out = torch.cat([h(x) for h in self.heads], dim=-1) return self.dropout(self.proj(out)) class FeedForward(nn.Module): def __init__(self, n_embd): super().__init__() self.net = nn.Sequential( nn.Linear(n_embd, 4 * n_embd), nn.GELU(), nn.Linear(4 * n_embd, n_embd), nn.Dropout(DROPOUT), ) def forward(self, x): return self.net(x) class Block(nn.Module): def __init__(self, n_embd, n_head): super().__init__() head_size = n_embd // n_head self.sa = MultiHeadAttention(n_head, head_size) self.ffwd = FeedForward(n_embd) self.ln1 = nn.LayerNorm(n_embd) self.ln2 = nn.LayerNorm(n_embd) def forward(self, x): x = x + self.sa(self.ln1(x)) x = x + self.ffwd(self.ln2(x)) return x class CommentGPT(nn.Module): def __init__(self, vocab_size): super().__init__() self.token_embedding = nn.Embedding(vocab_size, N_EMBD) self.position_embedding = nn.Embedding(BLOCK_SIZE, N_EMBD) self.blocks = nn.Sequential(*[Block(N_EMBD, N_HEAD) for _ in range(N_LAYER)]) self.ln_f = nn.LayerNorm(N_EMBD) self.lm_head = nn.Linear(N_EMBD, vocab_size) self.vocab_size = vocab_size def forward(self, idx, targets=None): B, T = idx.shape tok_emb = self.token_embedding(idx) pos_emb = self.position_embedding(torch.arange(T, device=idx.device)) x = tok_emb + pos_emb x = self.blocks(x) x = self.ln_f(x) logits = self.lm_head(x) return logits, None @torch.no_grad() def generate(self, idx, max_new_tokens, temperature=0.8, top_k=40): for _ in range(max_new_tokens): idx_cond = idx[:, -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) idx_next = torch.multinomial(probs, num_samples=1) idx = torch.cat((idx, idx_next), dim=1) return idx class CharTokenizer: def __init__(self, stoi, itos): self.stoi = stoi self.itos = {int(k): v for k, v in itos.items()} def encode(self, s): # skip characters not seen during training instead of crashing return [self.stoi[c] for c in s if c in self.stoi] def decode(self, ids): return "".join(self.itos[i] for i in ids) def load_model(ckpt_path, device): ckpt = torch.load(ckpt_path, map_location=device) model = CommentGPT(ckpt["vocab_size"]).to(device) model.load_state_dict(ckpt["model_state"]) model.eval() tokenizer = CharTokenizer(ckpt["stoi"], ckpt["itos"]) return model, tokenizer def generate_comment(model, tokenizer, device, prompt="", max_new_tokens=200, temperature=0.8, top_k=40): if prompt: ids = tokenizer.encode(prompt) if not ids: ids = [0] else: ids = [0] context = torch.tensor([ids], dtype=torch.long, device=device) out = model.generate(context, max_new_tokens=max_new_tokens, temperature=temperature, top_k=top_k)[0].tolist() text = tokenizer.decode(out) # cut at the first <|end|> after the prompt so you get one clean comment text = text.split("<|end|>")[0].strip() return text def main(): p = argparse.ArgumentParser() p.add_argument("--ckpt", default="comment_gpt.pt") p.add_argument("--temperature", type=float, default=0.8) p.add_argument("--top_k", type=int, default=40) p.add_argument("--max_new_tokens", type=int, default=200) args = p.parse_args() device = get_device() print(f"using device: {device}") print(f"loading {args.ckpt}...") model, tokenizer = load_model(args.ckpt, device) print("loaded. type a prompt (or leave blank) and hit enter. ctrl+c to quit.\n") while True: try: prompt = input("> ") except (KeyboardInterrupt, EOFError): print("\nbye") break comment = generate_comment( model, tokenizer, device, prompt=prompt, max_new_tokens=args.max_new_tokens, temperature=args.temperature, top_k=args.top_k, ) print(comment) print() if __name__ == "__main__": main()