Download modeling_diffurefill.py from Asilarkness/DiffuRefill-1B-base: direct link, hf CLI and curl.
- Browser
- Download file 6.64 kB
-
https://huggingface.co/Asilarkness/DiffuRefill-1B-base/resolve/main/modeling_diffurefill.py
- Command line
-
hf download hf://Asilarkness/DiffuRefill-1B-base/modeling_diffurefill.py
-
curl -L -o modeling_diffurefill.py https://huggingface.co/Asilarkness/DiffuRefill-1B-base/resolve/main/modeling_diffurefill.py
6.64 kB
| """DiffuRefill-1B: a 1.08B-parameter masked-diffusion language model. | |
| Standalone inference code: no dependency on the training stack. | |
| from modeling_diffurefill import DiffuRefill | |
| m = DiffuRefill.from_pretrained("Asilarkness/DiffuRefill-1B-base").cuda() | |
| print(m.generate("The capital of France is", max_new_tokens=32)) | |
| The network is a bidirectional transformer (no causal mask) trained to fill in | |
| [MASK] tokens (MDLM objective). Text is generated by starting from a row of | |
| masks after the prompt and committing tokens over several passes. | |
| """ | |
| import json, math | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| TOKENIZER = "openbmb/MiniCPM4-0.5B" | |
| def _rope_cache(T, dh, device, base=10000.0): | |
| inv = 1.0 / (base ** (torch.arange(0, dh, 2, device=device).float() / dh)) | |
| f = torch.outer(torch.arange(T, device=device).float(), inv) | |
| return torch.cos(f)[None, None], torch.sin(f)[None, None] | |
| def _rope(x, cos, sin): | |
| x1, x2 = x[..., ::2], x[..., 1::2] | |
| return torch.stack([x1 * cos - x2 * sin, x1 * sin + x2 * cos], -1).flatten(-2) | |
| class Block(nn.Module): | |
| def __init__(self, dim, heads, ff): | |
| super().__init__() | |
| self.dim, self.heads = dim, heads | |
| self.n1 = nn.RMSNorm(dim) | |
| self.qkv = nn.Linear(dim, 3 * dim, bias=False) | |
| self.o = nn.Linear(dim, dim, bias=False) | |
| self.n2 = nn.RMSNorm(dim) | |
| self.w1 = nn.Linear(dim, ff, bias=False) | |
| self.w3 = nn.Linear(dim, ff, bias=False) | |
| self.w2 = nn.Linear(ff, dim, bias=False) | |
| def forward(self, x, cos, sin): | |
| B, T, D = x.shape | |
| H, dh = self.heads, D // self.heads | |
| q, k, v = self.qkv(self.n1(x)).split(D, -1) | |
| q = _rope(q.view(B, T, H, dh).transpose(1, 2), cos, sin) | |
| k = _rope(k.view(B, T, H, dh).transpose(1, 2), cos, sin) | |
| v = v.view(B, T, H, dh).transpose(1, 2) | |
| a = F.scaled_dot_product_attention(q, k, v) # bidirectional | |
| x = x + self.o(a.transpose(1, 2).reshape(B, T, D)) | |
| h = self.n2(x) | |
| return x + self.w2(F.silu(self.w1(h)) * self.w3(h)) | |
| class DiffuRefill(nn.Module): | |
| def __init__(self, vocab_size=73760, dim=2048, layers=18, heads=16, ff=5632, max_len=2048): | |
| super().__init__() | |
| self.cfg = dict(vocab_size=vocab_size, dim=dim, layers=layers, heads=heads, ff=ff, max_len=max_len) | |
| self.mask_id, self.pad_id = vocab_size - 1, vocab_size - 2 | |
| self.tok = nn.Embedding(vocab_size, dim) | |
| self.blocks = nn.ModuleList([Block(dim, heads, ff) for _ in range(layers)]) | |
| self.nf = nn.RMSNorm(dim) | |
| self._tk = None | |
| # ------------------------------------------------------------------ io | |
| def from_pretrained(cls, repo="Asilarkness/DiffuRefill-1B-base", dtype=torch.bfloat16): | |
| from huggingface_hub import hf_hub_download | |
| from safetensors.torch import load_file | |
| cfg = json.load(open(hf_hub_download(repo, "config.json"))) | |
| m = cls(**{k: cfg[k] for k in ("vocab_size", "dim", "layers", "heads", "ff", "max_len")}) | |
| m.load_state_dict(load_file(hf_hub_download(repo, "model.safetensors"))) | |
| return m.to(dtype).eval() | |
| def tokenizer(self): | |
| if self._tk is None: | |
| from transformers import AutoTokenizer | |
| self._tk = AutoTokenizer.from_pretrained(TOKENIZER, trust_remote_code=True) | |
| return self._tk | |
| # ------------------------------------------------------------- network | |
| def forward(self, ids): | |
| """ids (B, T) -> logits (B, T, vocab). The output head is tied to the embedding.""" | |
| T = ids.shape[1] | |
| cos, sin = _rope_cache(T, self.cfg["dim"] // self.cfg["heads"], ids.device) | |
| x = self.tok(ids) | |
| for b in self.blocks: | |
| x = b(x, cos.to(x.dtype), sin.to(x.dtype)) | |
| return F.linear(self.nf(x), self.tok.weight) | |
| # ------------------------------------------------------------ sampling | |
| def generate(self, prompt, max_new_tokens=64, per_pass=1, block=32, temperature=0.0, return_ids=False): | |
| """Fill `max_new_tokens` masks after the prompt. | |
| Blocks of `block` positions are decoded left to right; inside a block, | |
| each pass commits the `per_pass` most confident masked positions. | |
| per_pass=1 is the slowest and the most coherent setting (one token per | |
| pass; it removes the repetition loops that parallel commits cause); | |
| larger values trade quality for speed. | |
| """ | |
| tk = self.tokenizer | |
| dev = self.tok.weight.device | |
| bos = tk.bos_token_id if tk.bos_token_id is not None else 1 | |
| p = [bos] + [i for i in tk(prompt, add_special_tokens=False).input_ids if i < self.pad_id] | |
| x = torch.tensor([p + [self.mask_id] * max_new_tokens], device=dev) | |
| P = len(p) | |
| for b0 in range(P, P + max_new_tokens, block): | |
| b1 = min(b0 + block, P + max_new_tokens) | |
| while (x[0, b0:b1] == self.mask_id).any(): | |
| lg = self(x)[0, b0:b1, :self.pad_id].float() | |
| if temperature > 0: | |
| probs = F.softmax(lg / temperature, -1) | |
| tok = torch.multinomial(probs, 1)[:, 0] | |
| conf = probs.gather(1, tok[:, None])[:, 0] | |
| else: | |
| conf, tok = F.softmax(lg, -1).max(-1) | |
| masked = x[0, b0:b1] == self.mask_id | |
| conf = torch.where(masked, conf, torch.full_like(conf, -1.0)) | |
| k = min(per_pass, int(masked.sum())) | |
| idx = conf.topk(k).indices | |
| x[0, b0 + idx] = tok[idx] | |
| out = x[0, P:].tolist() | |
| return out if return_ids else tk.decode(out, skip_special_tokens=True) | |
| def loglikelihood(self, context, continuation): | |
| """log p(continuation | context) by the left-to-right chain rule: | |
| token i is predicted with tokens < i visible and tokens >= i masked. | |
| This is the scorer used for the reported benchmark numbers.""" | |
| tk = self.tokenizer | |
| dev = self.tok.weight.device | |
| bos = tk.bos_token_id if tk.bos_token_id is not None else 1 | |
| c = [bos] + tk(context, add_special_tokens=False).input_ids | |
| x = tk(continuation, add_special_tokens=False).input_ids | |
| n = len(x) | |
| rows = torch.tensor([c + x[:i] + [self.mask_id] * (n - i) for i in range(n)], device=dev) | |
| lg = self(rows)[torch.arange(n, device=dev), torch.arange(len(c), len(c) + n, device=dev), :self.pad_id] | |
| return float(F.log_softmax(lg.float(), -1).gather(1, torch.tensor(x, device=dev)[:, None]).sum()) | |