Download sampler.py from basically-experimental/Notio-3.7M-RNN-v1: direct link, hf CLI and curl.
- Browser
- Download file 3.1 kB
-
https://huggingface.co/basically-experimental/Notio-3.7M-RNN-v1/resolve/main/sampler.py
- Command line
-
hf download hf://basically-experimental/Notio-3.7M-RNN-v1/sampler.py
-
curl -L -o sampler.py https://huggingface.co/basically-experimental/Notio-3.7M-RNN-v1/resolve/main/sampler.py
3.1 kB
| """Sampler: load a Notio checkpoint and generate completions statefully. | |
| Prompt is fed once (positions 0..P-1), then generation continues token by | |
| token carrying the RNN state with advancing position offsets. Decodes via | |
| the runtime decoder: tags -> effects, display text for humans. | |
| """ | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| PROJ = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(PROJ)) | |
| sys.path.insert(0, str(PROJ / "scripts")) | |
| import tokenizer as tk | |
| from model import Notio | |
| tk.load_vocab() | |
| def detach_states(states): | |
| out = [] | |
| for st in states: | |
| out.append(tuple(s.detach() for s in st) if isinstance(st, tuple) else st.detach()) | |
| return out | |
| def load_model(checkpoint, device=None): | |
| """Load checkpoint -> (model, device, ck_dict). Weights-only disabled: | |
| our checkpoints embed NotioConfig.""" | |
| device = device or ("cuda" if torch.cuda.is_available() else "cpu") | |
| ck = torch.load(checkpoint, map_location=device, weights_only=False) | |
| m = Notio(ck["cfg"]) | |
| m.load_state_dict(ck["state_dict"]) | |
| m.to(device) | |
| m.eval() | |
| return m, device, ck | |
| def encode_prompt(text): | |
| """Text -> id list, prefixed with <bos>. Accepts literal whitespace.""" | |
| if isinstance(text, str): | |
| text = text.encode() | |
| return [1] + list(tk.encode_bytes(text)) # <bos> | |
| def ids_to_display(ids): | |
| """id list -> human text via the runtime decoder (tags mapped/dropped).""" | |
| out = bytes(np.asarray(ids, dtype=np.uint8)).translate(tk.TABLE_DEC) | |
| for sent, tok in tk.SENT_DEC: | |
| out = out.replace(sent, tok) | |
| return tk.display(out).decode() | |
| def generate(m, device, prompt_ids, max_tokens=900, temperature=0.9, top_k=0): | |
| # NOTE: learned story length is ~700-900 ids; too-small max_tokens chops | |
| # stories off before the model's natural <eos> (see src/diag_eos.py) | |
| """Stateful completion. Stops at <eos> or max_tokens. Returns id list.""" | |
| block = m.cfg.layer1.block_size | |
| if len(prompt_ids) >= block: | |
| return list(prompt_ids) # nothing left to generate | |
| if len(prompt_ids) + max_tokens > block: | |
| max_tokens = block - len(prompt_ids) # wpe capacity is the ceiling | |
| ids = torch.tensor([prompt_ids], dtype=torch.long, device=device) | |
| logits, states = m(ids) # one pass over the prompt, pos 0..P-1 | |
| states = detach_states(states) | |
| pos = len(prompt_ids) | |
| for _ in range(max_tokens): | |
| logits, states = m(ids[:, -1:], states, pos_offset=pos) | |
| states = detach_states(states) | |
| pos += 1 | |
| logits = logits[:, -1, :] / max(temperature, 1e-6) | |
| if top_k > 0: | |
| v, _ = torch.topk(logits, top_k) | |
| logits = torch.where(logits < v[:, -1:], | |
| torch.full_like(logits, -float("inf")), logits) | |
| probs = F.softmax(logits, dim=-1) | |
| nxt = torch.multinomial(probs, 1) | |
| ids = torch.cat([ids, nxt], dim=1) | |
| if nxt.item() == 2: # <eos> | |
| break | |
| return ids[0].tolist() | |