"""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())