Download model.py from shalev396/tiny-shakespeare-chat: direct link, hf CLI and curl.
- Browser
- Download file 13.6 kB
-
https://huggingface.co/shalev396/tiny-shakespeare-chat/resolve/main/model.py
- Command line
-
hf download hf://shalev396/tiny-shakespeare-chat/model.py
-
curl -L -o model.py https://huggingface.co/shalev396/tiny-shakespeare-chat/resolve/main/model.py
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) + ")") | |
| def from_text(cls, text: str) -> "CharTokenizer": | |
| return cls(sorted(set(text))) | |
| def vocab_size(self) -> int: | |
| return len(self.itos) | |
| 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)) | |
| 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 | |
| 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) | |