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