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