gpt-5horts / chat.py
BrodyMakezAI's picture
Upload 3 files
1a4c9b3 verified
Raw History Blame Contribute Delete
6.17 kB
"""
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()