DiffuRefill-1B-base / modeling_diffurefill.py
Asilarkness's picture
DiffuRefill-1B base: weights, code, evaluation, technical report
9a57d8a verified
Raw History Blame Contribute Delete
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
@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())