shalev396's picture
Load weights on CPU before moving to the device (ZeroGPU-compatible)
13bff61 unverified
Raw History Blame Contribute Delete
13.6 kB
"""Tiny Shakespeare Chat: a character-level GPT built from scratch, then chat-tuned.
Single source of truth for the tokenizer, the architecture, the chat prompt format and
sampling. It is used by the Hugging Face Space (`space/app.py`), the Inference Endpoint
(`handler.py`), the training code (`training/src/`) and anyone who downloads this repo:
import model
predictor = model.load("path/to/this/repo", device="cpu")
predictor.predict("How fares the king?") # -> "GLOUCESTER: ..."
for piece in predictor.stream("Speak to me of love"): print(piece, end="")
"""
from __future__ import annotations
import math
import re
from collections.abc import Iterable, Iterator
from pathlib import Path
import torch
from huggingface_hub import PyTorchModelHubMixin
from torch import nn
from torch.nn import functional as F
REPO_ID = "shalev396/tiny-shakespeare-chat"
FRAMEWORK = "pytorch"
WEIGHTS_FILE = "model.safetensors"
# Chat markers. Each one is a single reserved token id (0, 1, 2); corpus characters follow.
USER, BOT, END = "<|user|>", "<|bot|>", "<|end|>"
SPECIAL_TOKENS = (USER, BOT, END)
# Generation defaults (the Space, the endpoint and the notebook all use these).
MAX_NEW_TOKENS = 200
TEMPERATURE = 0.8
TOP_K = 40
HISTORY_TURNS = 3 # earlier (user, bot) turns kept in the prompt
REPLY_ROOM = 16 # prompt is cropped to block_size - REPLY_ROOM tokens
def cuda_available() -> bool:
return torch.cuda.is_available()
# --------------------------------------------------------------------------- tokenizer
class CharTokenizer:
"""Character-level vocabulary plus the three chat markers as single reserved ids.
`encode` splits the text on the marker strings first, so each marker becomes ONE id, and
char-encodes everything in between. Characters that are not in the vocabulary are dropped.
"""
def __init__(self, chars: Iterable[str], specials: Iterable[str] = SPECIAL_TOKENS):
self.chars = list(chars)
self.specials = list(specials)
self.itos = {i: tok for i, tok in enumerate(self.specials + self.chars)}
self.stoi = {tok: i for i, tok in self.itos.items()}
self._special_set = set(self.specials)
self._pattern = re.compile("(" + "|".join(re.escape(t) for t in self.specials) + ")")
@classmethod
def from_text(cls, text: str) -> "CharTokenizer":
return cls(sorted(set(text)))
@property
def vocab_size(self) -> int:
return len(self.itos)
@property
def end_id(self) -> int:
return self.stoi[END]
def encode(self, text: str) -> list[int]:
ids: list[int] = []
for part in self._pattern.split(text):
if not part:
continue
if part in self._special_set:
ids.append(self.stoi[part])
else:
ids.extend(self.stoi[ch] for ch in part if ch in self.stoi)
return ids
def decode(self, ids: Iterable[int]) -> str:
return "".join(self.itos[int(i)] for i in ids if int(i) in self.itos)
def format_pair(user: str, bot: str) -> str:
"""One chat-formatted exchange, exactly as in the chat-tuning data."""
return f"{USER} {user} {END} {BOT} {bot} {END}"
def normalize_history(history) -> list[tuple[str, str]]:
"""Accept [[user, bot], ...] pairs or Gradio/OpenAI-style [{"role", "content"}, ...] messages."""
pairs: list[tuple[str, str]] = []
pending: str | None = None
for item in history or []:
if isinstance(item, dict):
role, content = item.get("role"), item.get("content")
if isinstance(content, list): # Gradio 6 content blocks: [{"type": "text", "text": ...}]
content = "".join(c.get("text", "") for c in content if isinstance(c, dict))
if not isinstance(content, str):
continue
if role == "user":
pending = content
elif role == "assistant" and pending is not None:
pairs.append((pending, content))
pending = None
elif isinstance(item, (list, tuple)) and len(item) == 2:
user, bot = item
if isinstance(user, str) and isinstance(bot, str):
pairs.append((user, bot))
else:
raise ValueError("history must be a list of [user, bot] pairs or {role, content} messages")
return pairs
def build_prompt(message: str, history=None, history_turns: int = HISTORY_TURNS) -> str:
"""`<|user|> msg <|end|> <|bot|>`, preceded by up to `history_turns` earlier exchanges."""
earlier = normalize_history(history)[-history_turns:] if history_turns > 0 else []
parts = [format_pair(u, b) + "\n" for u, b in earlier]
parts.append(f"{USER} {message} {END} {BOT}")
return "".join(parts)
# --------------------------------------------------------------------------- architecture
class CausalSelfAttention(nn.Module):
"""Multi-head causal self-attention through the fused SDPA kernel."""
def __init__(self, n_embd: int, n_head: int, dropout: float):
super().__init__()
if n_embd % n_head:
raise ValueError("n_embd must be divisible by n_head")
self.n_head = n_head
self.dropout = dropout
self.c_attn = nn.Linear(n_embd, 3 * n_embd, bias=False) # q, k, v for every head at once
self.c_proj = nn.Linear(n_embd, n_embd, bias=False)
self.resid_dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
B, T, C = x.shape
q, k, v = self.c_attn(x).split(C, dim=2)
q, k, v = (t.view(B, T, self.n_head, C // self.n_head).transpose(1, 2) for t in (q, k, v))
y = F.scaled_dot_product_attention(q, k, v, dropout_p=self.dropout if self.training else 0.0,
is_causal=True)
return self.resid_dropout(self.c_proj(y.transpose(1, 2).contiguous().view(B, T, C)))
class MLP(nn.Module):
"""Position-wise feed-forward: expand 4x, GELU, project back."""
def __init__(self, n_embd: int, dropout: float):
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: torch.Tensor) -> torch.Tensor:
return self.dropout(self.c_proj(self.gelu(self.c_fc(x))))
class Block(nn.Module):
"""Pre-LayerNorm transformer block: x + attn(ln(x)), then x + mlp(ln(x))."""
def __init__(self, n_embd: int, n_head: int, dropout: float):
super().__init__()
self.ln_1 = nn.LayerNorm(n_embd)
self.attn = CausalSelfAttention(n_embd, n_head, dropout)
self.ln_2 = nn.LayerNorm(n_embd)
self.mlp = MLP(n_embd, dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x + self.attn(self.ln_1(x))
return x + self.mlp(self.ln_2(x))
class GPT(
nn.Module,
PyTorchModelHubMixin,
tags=["ml-lab", "text-generation", "gpt", "from-scratch", "character-level"],
repo_url="https://github.com/shalev396/ml-lab",
pipeline_tag="text-generation",
license="mit",
):
"""Decoder-only transformer (nanoGPT-style) over the char + chat-marker vocabulary.
token + learned position embeddings -> `n_layer` pre-LN blocks -> LayerNorm -> LM head whose
weight is tied to the token embedding. The vocabulary (`chars`, `specials`) and `block_size`
are constructor args, so PyTorchModelHubMixin stores them in config.json.
"""
def __init__(self, chars: list[str], specials: list[str] = list(SPECIAL_TOKENS), n_layer: int = 6,
n_head: int = 6, n_embd: int = 384, block_size: int = 256, dropout: float = 0.2):
super().__init__()
self.chars, self.specials = list(chars), list(specials)
self.vocab_size = len(self.specials) + len(self.chars)
self.block_size = block_size
self.wte = nn.Embedding(self.vocab_size, n_embd)
self.wpe = nn.Embedding(block_size, n_embd)
self.drop = nn.Dropout(dropout)
self.blocks = nn.ModuleList(Block(n_embd, n_head, dropout) for _ in range(n_layer))
self.ln_f = nn.LayerNorm(n_embd)
self.lm_head = nn.Linear(n_embd, self.vocab_size, bias=False)
self.lm_head.weight = self.wte.weight # weight tying
self.apply(self._init_weights)
for name, p in self.named_parameters(): # GPT-2: scale residual projections
if name.endswith("c_proj.weight"):
nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * n_layer))
@staticmethod
def _init_weights(module: nn.Module) -> None:
if isinstance(module, (nn.Linear, nn.Embedding)):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if isinstance(module, nn.Linear) and module.bias is not None:
nn.init.zeros_(module.bias)
def tokenizer(self) -> CharTokenizer:
return CharTokenizer(self.chars, self.specials)
def forward(self, idx: torch.Tensor, targets: torch.Tensor | None = None,
last_only: bool = False) -> tuple[torch.Tensor, torch.Tensor | None]:
"""(B, T) ids -> (logits, loss). `last_only` returns logits for the last position only."""
T = idx.shape[1]
if T > self.block_size:
raise ValueError(f"sequence length {T} > block_size {self.block_size}")
x = self.drop(self.wte(idx) + self.wpe(torch.arange(T, device=idx.device)))
for block in self.blocks:
x = block(x)
x = self.ln_f(x)
if last_only and targets is None:
return self.lm_head(x[:, [-1], :]), None
logits = self.lm_head(x)
loss = None if targets is None else F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))
return logits, loss
def num_params(self) -> int:
return sum(p.numel() for p in self.parameters()) # tied weight counted once
@torch.inference_mode()
def sample(model: GPT, prompt_ids: list[int], *, max_new_tokens: int = MAX_NEW_TOKENS,
temperature: float = TEMPERATURE, top_k: int | None = TOP_K, stop_id: int | None = None,
generator: torch.Generator | None = None) -> Iterator[int]:
"""Autoregressive sampling (temperature + top-k). Yields new token ids; stops at `stop_id`."""
device = next(model.parameters()).device
idx = torch.tensor([prompt_ids], dtype=torch.long, device=device)
for _ in range(int(max_new_tokens)):
logits, _ = model(idx[:, -model.block_size:], last_only=True)
logits = logits[:, -1, :].float() / max(float(temperature), 1e-5)
if top_k:
v, _ = torch.topk(logits, min(int(top_k), logits.size(-1)))
logits[logits < v[:, [-1]]] = -float("inf")
next_id = torch.multinomial(F.softmax(logits, dim=-1), 1, generator=generator)
if stop_id is not None and int(next_id) == stop_id:
return
idx = torch.cat([idx, next_id], dim=1)
yield int(next_id)
# --------------------------------------------------------------------------- inference
class Predictor:
"""Loads the chat-tuned GPT once and answers chat messages in play-style verse."""
def __init__(self, model_dir: str | Path, device: str = "cpu"):
model_dir = Path(model_dir)
if not (model_dir / WEIGHTS_FILE).is_file():
raise FileNotFoundError(
f"{WEIGHTS_FILE} not found in {model_dir}: train/export the model first "
f"or download it with huggingface_hub.snapshot_download('{REPO_ID}')."
)
self.device = torch.device(device)
self.model = GPT.from_pretrained(model_dir, map_location="cpu", strict=True)
self.model.to(self.device).eval()
self.tokenizer = self.model.tokenizer()
self.block_size = self.model.block_size
def prompt_ids(self, message: str, history=None) -> list[int]:
ids = self.tokenizer.encode(build_prompt(message, history))
return ids[-(self.block_size - REPLY_ROOM):] # keep room for the reply to start
def stream(self, message: str, history=None, max_new_tokens: int = MAX_NEW_TOKENS,
temperature: float = TEMPERATURE, top_k: int = TOP_K,
seed: int | None = None) -> Iterator[str]:
"""Yield the reply one character at a time (stops at `<|end|>` or `max_new_tokens`)."""
if not isinstance(message, str) or not message.strip():
raise ValueError("message must be a non-empty string")
generator = None
if seed is not None:
generator = torch.Generator(device=self.device.type).manual_seed(int(seed))
for token in sample(self.model, self.prompt_ids(message.strip(), history),
max_new_tokens=max_new_tokens, temperature=temperature, top_k=top_k,
stop_id=self.tokenizer.end_id, generator=generator):
yield self.tokenizer.decode([token])
def predict(self, message: str, history=None, max_new_tokens: int = MAX_NEW_TOKENS,
temperature: float = TEMPERATURE, top_k: int = TOP_K, seed: int | None = None) -> str:
"""Chat message (+ optional history) -> the whole reply as a plain string."""
return "".join(self.stream(message, history, max_new_tokens, temperature, top_k, seed)).strip()
def load(model_dir: str | Path, device: str = "cpu") -> Predictor:
return Predictor(model_dir, device)