File size: 17,317 Bytes
bcda938
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
"""Fast generation: cached incremental decoding, batched sampling, optional grammar constraints.



Each sequence's state is small and fixed per token: attention layers keep their keys/values, DeltaNet

layers keep their recurrent state (flash-linear-attention cache), and Block AttnRes only mixes within

a token, so it needs no cache. A prompt is processed once (prefill) and each new token then costs one

cheap step instead of re-running the whole conversation.



  outs = generate(model, [prompt_ids, ...], max_new=200, stop_ids=[enc.im_end])

  matcher = GrammarFactory(tokenizer_path, im_end_id).json(schema)   # output guaranteed to parse

"""
import torch
import torch.nn.functional as F

from model import DeltaNet, GatedAttention, attnres_mix, inv_rms, rms


class _GDNStates:
    """The minimal cache interface flash-linear-attention layers expect (len / [] / update)."""

    def __init__(self, n):
        self.layers = [None] * n

    def __len__(self):
        return len(self.layers)

    def __getitem__(self, i):
        return self.layers[i]

    def update(self, layer_idx, recurrent_state=None, conv_state=None, **_):
        cur = self.layers[layer_idx]
        if cur is None:  # first call (prefill): keep the tensors
            self.layers[layer_idx] = dict(recurrent_state=recurrent_state, conv_state=conv_state)
            return
        # later calls: write into the existing buffers, so a replayed CUDA graph carries the state forward
        pairs = [(cur["recurrent_state"], recurrent_state)] + list(zip(cur["conv_state"] or (), conv_state or ()))
        for dst, src in pairs:
            if src is not None and dst.data_ptr() != src.data_ptr():
                dst.copy_(src)


class Cache:
    """Decoding state for a batch: attention keys/values (preallocated to `capacity`), DeltaNet states,

    and each row's current length (rows may have different lengths)."""

    def __init__(self, model, batch, capacity):
        dev, dt = model.embed.weight.device, torch.bfloat16
        self.pos = torch.zeros(batch, dtype=torch.long, device=dev)
        self.max_pos, self.capacity = 0, capacity
        self.kv = {j: (torch.zeros(batch, s.nkv, capacity, s.hd, device=dev, dtype=dt),
                       torch.zeros(batch, s.nkv, capacity, s.hd, device=dev, dtype=dt))
                   for j, s in enumerate(model.sublayers) if isinstance(s, GatedAttention)}
        self.gdn = _GDNStates(len(model.sublayers))

    @staticmethod
    def stack(caches, extra):
        """Combine single-sequence caches of different lengths into one batch with room for `extra` tokens."""
        out = Cache.__new__(Cache)
        out.pos = torch.cat([c.pos for c in caches])
        out.max_pos = max(c.max_pos for c in caches)
        out.capacity = out.max_pos + extra
        out.kv = {}
        for j, (K0, _) in caches[0].kv.items():
            K = K0.new_zeros(len(caches), K0.size(1), out.capacity, K0.size(3))
            V = torch.zeros_like(K)
            for i, c in enumerate(caches):
                n = c.max_pos
                K[i, :, :n], V[i, :, :n] = c.kv[j][0][0, :, :n], c.kv[j][1][0, :, :n]
            out.kv[j] = (K, V)
        out.gdn = _GDNStates(len(caches[0].gdn))
        for i, layer in enumerate(caches[0].gdn.layers):
            if layer is not None:
                conv = layer["conv_state"]
                out.gdn.layers[i] = dict(
                    recurrent_state=torch.cat([c.gdn.layers[i]["recurrent_state"] for c in caches]),
                    conv_state=None if conv is None else tuple(
                        torch.cat([c.gdn.layers[i]["conv_state"][k] for c in caches]) for k in range(len(conv))))
        return out


def _attention(sub, x, K, V, pos):
    """GatedAttention over cached keys/values; x holds t new tokens per row starting at position pos[row]."""
    B, t, _ = x.shape
    q, gate = sub.q_proj(x).split(sub.nh * sub.hd, -1)
    q = q.view(B, t, sub.nh, sub.hd)
    k, v = sub.kv_proj(x).view(B, t, 2, sub.nkv, sub.hd).unbind(2)
    positions = pos[:, None] + torch.arange(t, device=x.device)            # [B, t]
    cos, sin = (b[0, :, 0][positions].unsqueeze(2).to(x.dtype) for b in (sub.cos, sub.sin))

    def rope(z):
        z1, z2 = z.chunk(2, -1)
        return torch.cat([z1 * cos - z2 * sin, z1 * sin + z2 * cos], -1)

    q, k = rope(rms(q)), rope(rms(k))
    rows = torch.arange(B, device=x.device)[:, None].expand(B, t)
    K[rows, :, positions] = k.to(K.dtype)
    V[rows, :, positions] = v.to(V.dtype)
    keys = K.repeat_interleave(sub.nh // sub.nkv, dim=1)
    vals = V.repeat_interleave(sub.nh // sub.nkv, dim=1)
    mask = torch.arange(K.size(2), device=x.device)[None, None, :] <= positions[:, :, None]  # causal per row
    y = F.scaled_dot_product_attention(q.transpose(1, 2), keys.to(q.dtype), vals.to(q.dtype), attn_mask=mask[:, None])
    y = y.transpose(1, 2).reshape(B, t, sub.nh * sub.hd)
    return sub.o_proj(y * torch.sigmoid(gate))


@torch.no_grad()
def step(model, idx, cache):
    """Feed t new tokens per row (idx [B, t]); returns next-token logits [B, vocab] for the last position."""
    t = idx.size(1)
    if cache.max_pos + t > cache.capacity:
        raise ValueError(f"cache full ({cache.capacity} tokens)")
    x0 = rms(model.embed(idx))
    blocks, blocks_inv, partial = [x0], [inv_rms(x0)], None
    for j, sub in enumerate(model.sublayers):
        srcs = blocks if partial is None else blocks + [partial]
        inv = blocks_inv if partial is None else blocks_inv + [inv_rms(partial)]
        h = rms(attnres_mix(srcs, inv, model.attnres_queries[j]))
        if isinstance(sub, GatedAttention):
            out = _attention(sub, h, *cache.kv[j], cache.pos)
        elif isinstance(sub, DeltaNet):
            sub.gdn.layer_idx = j
            out = sub.gdn(h, past_key_values=cache.gdn, use_cache=True)[0]
        else:
            out = sub(h)
        out = out.float()
        partial = out if partial is None else partial + out
        if (j + 1) % model.block_size == 0:
            blocks.append(partial)
            blocks_inv.append(inv_rms(partial))
            partial = None
    srcs = blocks if partial is None else blocks + [partial]
    inv = blocks_inv if partial is None else blocks_inv + [inv_rms(partial)]
    h = rms(attnres_mix(srcs, inv, model.attnres_queries[-1]))
    cache.pos += t
    cache.max_pos += t
    return model._logits(h[:, -1])


def prefill(model, prompt, extra):
    """Process one prompt (list/array of ids); returns (cache with room for `extra` tokens, last logits [1, V])."""
    ids = torch.as_tensor(prompt, dtype=torch.long, device=model.embed.weight.device)[None]
    cache = Cache(model, 1, ids.size(1) + extra)
    return cache, step(model, ids, cache)


def sample(logits, temperature=0.7, top_p=0.9, top_k=0, recent=None, rep_penalty=1.0, bitmask=None, vocab=None):
    """Pick one token per row. recent: [B, n] recently generated ids (-1 = none) for the repetition penalty."""
    logits = logits.float()
    if vocab is not None and vocab < logits.size(1):
        logits[:, vocab:] = -float("inf")  # padding rows of the embedding are not real tokens
    if recent is not None and rep_penalty != 1.0:
        seen = torch.zeros(logits.size(0), logits.size(1) + 1, device=logits.device, dtype=torch.bool)
        seen.scatter_(1, torch.where(recent < 0, logits.size(1), recent), True)
        seen = seen[:, :-1]
        logits = torch.where(seen, torch.where(logits > 0, logits / rep_penalty, logits * rep_penalty), logits)
    if bitmask is not None:  # llguidance bitmask: bit i of word w allows token 32*w + i
        bits = (bitmask.to(logits.device)[:, :, None] >> torch.arange(32, device=logits.device)) & 1
        logits = logits.masked_fill(bits.reshape(bits.size(0), -1)[:, :logits.size(1)] == 0, -float("inf"))
    if temperature <= 0:
        return logits.argmax(-1)
    logits = logits / temperature
    if top_k:
        kth = torch.topk(logits, top_k, dim=-1).values[:, -1:]
        logits = logits.masked_fill(logits < kth, -float("inf"))
    probs = torch.softmax(logits, -1)
    if top_p < 1.0:
        sp, si = probs.sort(-1, descending=True)
        sp = sp.masked_fill(sp.cumsum(-1) - sp > top_p, 0.0)
        probs = torch.zeros_like(probs).scatter_(-1, si, sp)
    return torch.multinomial(probs, 1)[:, 0]


@torch.no_grad()
def generate(model, prompts, max_new=256, temperature=0.7, top_p=0.9, top_k=0, rep_penalty=1.0, stop_ids=(),

             matchers=None, batch_size=32, vocab=None, on_token=None, cuda_graph=True):
    """Sample continuations for many prompts, `batch_size` at a time. Returns a list of token-id lists.

    matchers: optional per-prompt llguidance matchers (None = unconstrained) that force a grammar.

    on_token(row_index, token_id): optional callback, e.g. for streaming a single prompt."""
    from llguidance.torch import allocate_token_bitmask, fill_next_token_bitmask
    results = [None] * len(prompts)
    stop = torch.tensor(list(stop_ids) or [-1], device=model.embed.weight.device)
    for start in range(0, len(prompts), batch_size):
        idx = list(range(start, min(start + batch_size, len(prompts))))
        pre = [prefill(model, prompts[i], max_new) for i in idx]
        cache = Cache.stack([c for c, _ in pre], max_new) if len(pre) > 1 else pre[0][0]
        logits = torch.cat([l for _, l in pre])
        del pre
        B = len(idx)
        rows_m = [matchers[i] if matchers else None for i in idx]
        bitmask = allocate_token_bitmask(B, logits.size(1)) if any(rows_m) else None
        out = [[] for _ in range(B)]
        done = torch.zeros(B, dtype=torch.bool, device=logits.device)
        recent = torch.full((B, 64), -1, dtype=torch.long, device=logits.device)
        graph = _DecodeGraph(model, cache, B) if cuda_graph else None
        for n in range(max_new):
            if bitmask is not None:
                bitmask.fill_(-1)  # all tokens allowed ...
                for r, m in enumerate(rows_m):
                    if m is not None and not m.is_stopped():
                        fill_next_token_bitmask(m, bitmask, r)  # ... except where a grammar forbids them
            nxt = sample(logits, temperature, top_p, top_k, recent, rep_penalty, bitmask, vocab)
            nxt = torch.where(done, stop[0].clamp(min=0), nxt)
            toks = nxt.tolist()
            for r in range(B):
                if done[r]:
                    continue
                if rows_m[r] is not None:
                    rows_m[r].consume_token(toks[r])
                if toks[r] in stop_ids:
                    done[r] = True
                    continue
                out[r].append(toks[r])  # keep it even if it completes the grammar (e.g. the final "}")
                if on_token:
                    on_token(idx[r], toks[r])
                if rows_m[r] is not None and rows_m[r].is_stopped():
                    done[r] = True
            recent = torch.cat([recent[:, 1:], nxt[:, None]], 1)
            if bool(done.all()) or n == max_new - 1:
                break
            logits = graph.step(nxt) if graph else step(model, nxt[:, None], cache)
        for r, i in enumerate(idx):
            results[i] = out[r]
    return results


class _DecodeGraph:
    """Records the one-token decode step as a CUDA graph and replays it: one launch instead of ~3,500

    small kernel launches per token (launch overhead, not the GPU, dominates at this model size).

    The first steps run normally (they also warm up the Triton kernels); if recording fails, it

    quietly keeps running the normal way."""

    WARMUP = 2

    def __init__(self, model, cache, batch):
        self.model, self.cache, self.n, self.graph = model, cache, 0, None
        self.tok = torch.zeros(batch, 1, dtype=torch.long, device=model.embed.weight.device)
        self.failed = False

    def step(self, nxt):
        self.n += 1
        if self.failed or self.n <= self.WARMUP:
            return step(self.model, nxt[:, None], self.cache)
        if self.cache.max_pos + 1 > self.cache.capacity:
            raise ValueError(f"cache full ({self.cache.capacity} tokens)")
        self.tok.copy_(nxt[:, None])
        if self.graph is None:
            try:
                self.graph = torch.cuda.CUDAGraph()
                with torch.cuda.graph(self.graph), torch.autocast("cuda", dtype=torch.bfloat16, cache_enabled=False):
                    self.out = step(self.model, self.tok, self.cache)  # recorded, not run
                self.cache.max_pos -= 1  # the recording pass bumped it; the replay below is the real step
            except Exception as e:  # noqa: BLE001 - fall back to plain steps
                print(f"(CUDA graph unavailable, using normal decoding: {type(e).__name__}: {str(e)[:120]})")
                self.failed, self.graph = True, None
                return step(self.model, nxt[:, None], self.cache)
        self.graph.replay()
        self.cache.max_pos += 1
        return self.out


class GrammarFactory:
    """Builds llguidance matchers for our tokenizer: guaranteed-valid JSON or tool calls."""

    def __init__(self, tokenizer_path, stop_id, n_vocab):
        """n_vocab: the model's (padded) output size, so masks line up with its logits."""
        import llguidance.hf
        from transformers import PreTrainedTokenizerFast
        from tokenizers import Tokenizer
        hf = PreTrainedTokenizerFast(tokenizer_object=Tokenizer.from_file(str(tokenizer_path)))
        self.lltok = llguidance.hf.from_tokenizer(hf, n_vocab=n_vocab, eos_token=stop_id)

    def _matcher(self, grammar):
        from llguidance import LLMatcher
        m = LLMatcher(self.lltok, grammar)
        if m.is_error():
            raise ValueError(m.get_error())
        return m

    COMPACT = {"whitespace_flexible": False, "item_separator": ", ", "key_separator": ": "}

    def json(self, schema=None):
        """Any valid JSON value matching `schema` (a dict; None = any JSON object). Compact whitespace, so a

        reply can't wander off into blank space and run out of tokens before the JSON is closed."""
        from llguidance import LLMatcher
        return self._matcher(LLMatcher.grammar_from_json_schema({**(schema or {"type": "object"}), "x-guidance": self.COMPACT}))

    def tool_calls(self, tools):
        """One or more <tool_call> blocks whose JSON names a listed function with schema-valid arguments.

        tools: OpenAI-style [{"type": "function", "function": {"name", "parameters"}}, ...]. Parameter lists in

        the shorthand some datasets use ({"arg": {"type": "str, optional"}}) are converted to JSON Schema. If a

        schema still can't be compiled, only the function name is enforced."""
        from llguidance import LLMatcher
        fns = [t.get("function", t) for t in tools]
        try:
            return self._tool_matcher(fns, strict_args=True)
        except ValueError:
            return self._tool_matcher(fns, strict_args=False)

    def _tool_matcher(self, fns, strict_args):
        import json as _json
        from llguidance import LLMatcher
        options = []
        for fn in fns:
            params = to_json_schema(fn.get("parameters")) if strict_args else {"type": "object"}
            options.append({"type": "object", "properties": {"name": {"const": fn["name"]}, "arguments": params},
                            "required": ["name", "arguments"], "additionalProperties": False})
        schema = _json.dumps({"anyOf": options, "x-guidance": self.COMPACT})
        lark = (f'start: call ("\\n" call)*\n'
                f'call: "<tool_call>\\n" body "\\n</tool_call>"\n'
                f'body: %json {schema}\n')
        return self._matcher(LLMatcher.grammar_from_lark(lark))


_SHORT_TYPES = {"str": "string", "string": "string", "int": "integer", "integer": "integer", "float": "number",
                "number": "number", "bool": "boolean", "boolean": "boolean", "list": "array", "array": "array",
                "dict": "object", "object": "object"}


def to_json_schema(params):
    """Tool parameters as JSON Schema, accepting either real JSON Schema or the {"arg": {"type": "str, optional"}}

    shorthand. Unknown types become 'any value'. Only declared arguments are allowed."""
    if not params:
        return {"type": "object", "properties": {}, "additionalProperties": False}
    if params.get("type") == "object" or "properties" in params:
        out = dict(params)
        out.setdefault("additionalProperties", False)
        return out
    props, required = {}, []
    for name, spec in params.items():
        spec = spec if isinstance(spec, dict) else {}
        raw = str(spec.get("type", "")).lower()
        base = raw.split(",")[0].strip().split("[")[0]
        props[name] = {"type": _SHORT_TYPES[base]} if base in _SHORT_TYPES else {}
        if "optional" not in raw:
            required.append(name)
    return {"type": "object", "properties": props, "required": required, "additionalProperties": False}