File size: 6,641 Bytes
9a57d8a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 | """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
@classmethod
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()
@property
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
@torch.no_grad()
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)
@torch.no_grad()
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())
|