simit 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,844 @@
1
+ """Continuous-batching inference engine for MoT unified multimodal models.
2
+
3
+ One background thread owns the GPU. Every step it packs, into a single
4
+ forward pass:
5
+
6
+ * one decode token for every running text generation,
7
+ * prefill chunks (text: causal, image: bidirectional) of newly admitted
8
+ requests, within a token budget,
9
+ * one denoising step (conditional and unconditional CFG branch) for every
10
+ running image generation.
11
+
12
+ KV lives in a contiguous arena; each request owns a region. Long text
13
+ prefixes (skill and router prompts) are kept after a request finishes and
14
+ copied into new requests that share them (automatic prefix caching).
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ import hashlib
20
+ import itertools
21
+ import math
22
+ import queue
23
+ import threading
24
+ import time
25
+ import traceback
26
+ from collections import OrderedDict
27
+ from concurrent.futures import Future
28
+ from dataclasses import dataclass, field
29
+ from typing import Any, Callable, Optional
30
+
31
+ import torch
32
+ import torch.nn.functional as F
33
+
34
+ from .modeling import AttnBatch, KVArena
35
+
36
+
37
+ # ------------------------------------------------------------------ prompts
38
+ @dataclass
39
+ class TextChunk:
40
+ ids: list[int]
41
+
42
+
43
+ @dataclass
44
+ class ImageChunk:
45
+ """An understanding image: ``n`` tokens (start marker, ViT tokens, end
46
+ marker) sharing one rope position. ``encode`` returns their embeddings."""
47
+
48
+ n: int
49
+ key: str # content hash, for prefix-cache keys
50
+ payload: Any # backend-specific preprocessed image
51
+
52
+
53
+ @dataclass
54
+ class TextJob:
55
+ chunks: list # TextChunk / ImageChunk
56
+ max_new_tokens: int = 0 # 0 for scoring
57
+ temperature: float = 0.0
58
+ top_p: float = 1.0
59
+ repetition_penalty: float = 1.0
60
+ logprobs: bool = False
61
+ score_len: int = 0 # >0: teacher-forced scoring of the last `score_len` tokens
62
+ stop_ids: tuple = ()
63
+ on_tokens: Optional[Callable[[list], None]] = None # streaming: called with all output ids so far
64
+ future: Future = field(default_factory=Future)
65
+
66
+
67
+ @dataclass
68
+ class ImageJob:
69
+ chunks: list # conditioning prompt (text chunks)
70
+ height: int
71
+ width: int
72
+ seed: Optional[int] = None
73
+ future: Future = field(default_factory=Future)
74
+
75
+
76
+ # ------------------------------------------------------------- arena regions
77
+ class RegionAllocator:
78
+ def __init__(self, capacity: int):
79
+ self.free = [(0, capacity)] # sorted (start, size)
80
+
81
+ def alloc(self, size: int) -> Optional[int]:
82
+ for i, (start, n) in enumerate(self.free):
83
+ if n >= size:
84
+ if n == size:
85
+ self.free.pop(i)
86
+ else:
87
+ self.free[i] = (start + size, n - size)
88
+ return start
89
+ return None
90
+
91
+ def release(self, start: int, size: int):
92
+ if size <= 0:
93
+ return
94
+ self.free.append((start, size))
95
+ self.free.sort()
96
+ merged = []
97
+ for s, n in self.free:
98
+ if merged and merged[-1][0] + merged[-1][1] == s:
99
+ merged[-1] = (merged[-1][0], merged[-1][1] + n)
100
+ else:
101
+ merged.append((s, n))
102
+ self.free = merged
103
+
104
+
105
+ _G = 16 # prefix-cache granularity (tokens)
106
+
107
+
108
+ def _prefix_hashes(keys: list) -> list[tuple[int, str]]:
109
+ """Chained hashes of every _G-token prefix: [(length, hash), ...]."""
110
+ out, h = [], hashlib.sha1()
111
+ for i, k in enumerate(keys, 1):
112
+ h.update(str(k).encode() + b",")
113
+ if i % _G == 0:
114
+ out.append((i, h.copy().hexdigest()))
115
+ return out
116
+
117
+
118
+ @dataclass
119
+ class _CacheEntry:
120
+ start: int
121
+ length: int
122
+ owner: Any = None # live request still using the region
123
+ ready: bool = False
124
+
125
+
126
+ class PrefixCache:
127
+ def __init__(self, max_tokens: int):
128
+ self.max_tokens = max_tokens
129
+ self.entries: "OrderedDict[int, _CacheEntry]" = OrderedDict() # id -> entry (LRU order)
130
+ self.index: dict[str, tuple[int, int]] = {} # hash -> (entry id, length)
131
+ self._ids = itertools.count()
132
+ self.tokens = 0
133
+ self.hits = 0
134
+ self.hit_tokens = 0
135
+
136
+ def lookup(self, hashes):
137
+ for length, h in reversed(hashes):
138
+ ref = self.index.get(h)
139
+ if ref and ref[0] in self.entries:
140
+ e = self.entries[ref[0]]
141
+ self.entries.move_to_end(ref[0])
142
+ return e, length
143
+ return None, 0
144
+
145
+ def add(self, start: int, length: int, hashes, owner) -> int:
146
+ eid = next(self._ids)
147
+ self.entries[eid] = _CacheEntry(start, length, owner=owner)
148
+ for n, h in hashes:
149
+ if n <= length:
150
+ self.index[h] = (eid, n)
151
+ self.tokens += length
152
+ return eid
153
+
154
+ def evictable(self):
155
+ for eid, e in self.entries.items():
156
+ if e.owner is None:
157
+ yield eid, e
158
+
159
+ def remove(self, eid):
160
+ e = self.entries.pop(eid)
161
+ self.tokens -= e.length
162
+ return e
163
+
164
+
165
+ # --------------------------------------------------------------- running state
166
+ class _Text:
167
+ def __init__(self, job: TextJob, region: int, cap: int, keys: list, n_prompt: int):
168
+ self.job, self.region, self.cap = job, region, cap
169
+ self.keys, self.n_prompt = keys, n_prompt
170
+ self.written = 0 # tokens whose KV is in the region
171
+ self.pos = 0 # next rope position (1-D)
172
+ self.chunk_i, self.chunk_off = 0, 0
173
+ self.out_ids: list[int] = []
174
+ self.out_lp: list[float] = []
175
+ self.score_lp: list[float] = []
176
+ self.last_token: Optional[int] = None
177
+ self.cache_eid: Optional[int] = None
178
+ self.done = False
179
+ if job.score_len: # rows (key indices) whose next-token distribution scores the targets
180
+ self.score_lo = n_prompt - job.score_len - 1
181
+ self.score_hi = n_prompt - 2
182
+
183
+ @property
184
+ def prefilled(self):
185
+ return self.chunk_i >= len(self.job.chunks)
186
+
187
+
188
+ class _Img:
189
+ def __init__(self, job: ImageJob, region, cap, uncond_region, n_lat, h, w, pos_ids, noise, steps):
190
+ self.job, self.region, self.cap = job, region, cap
191
+ self.uncond_region = uncond_region
192
+ self.n_lat, self.h, self.w = n_lat, h, w
193
+ self.lat_pos_ids = pos_ids
194
+ self.x = noise
195
+ self.steps = steps # list of (t, dt)
196
+ self.step_i = 0
197
+ self.written = 0
198
+ self.pos = 0
199
+ self.chunk_i, self.chunk_off = 0, 0
200
+ self.n_prompt = sum(len(c.ids) for c in job.chunks)
201
+
202
+ @property
203
+ def prefilled(self):
204
+ return self.chunk_i >= len(self.job.chunks)
205
+
206
+
207
+ def _positions_tensor(positions, device):
208
+ """[T] for 1-D rope, [3, T] for multimodal rope (positions given as triples)."""
209
+ t = torch.tensor(positions, dtype=torch.long)
210
+ if t.dim() == 2:
211
+ t = t.T.contiguous()
212
+ return t.to(device, non_blocking=True)
213
+
214
+
215
+ class _DecodeGraphs:
216
+ """CUDA graphs for pure-decode steps (one new token per sequence), bucketed
217
+ by batch size. Padding rows write into a reserved dummy region."""
218
+
219
+ BUCKETS = (1, 2, 4, 8, 16, 24, 32, 48, 64)
220
+
221
+ def __init__(self, engine, max_k: int):
222
+ self.engine = engine
223
+ self.max_k = max_k
224
+ self.graphs = {}
225
+ self.pool = None
226
+ self.dummy_slot = engine.arena.capacity - 1 # last arena slot, never allocated
227
+ self.failed = False
228
+
229
+ def bucket(self, b):
230
+ for n in self.BUCKETS:
231
+ if n >= b:
232
+ return n
233
+ return None
234
+
235
+ def _capture(self, n):
236
+ eng = self.engine
237
+ lm, dev = eng.model.lm, eng.device
238
+ head_dim = lm.cfg["hidden_size"] // lm.cfg["num_attention_heads"]
239
+ st = {
240
+ "ids": torch.zeros(n, dtype=torch.long, device=dev),
241
+ "cos": torch.zeros(n, head_dim, dtype=eng.model.dtype, device=dev),
242
+ "sin": torch.zeros(n, head_dim, dtype=eng.model.dtype, device=dev),
243
+ "slots": torch.full((n,), self.dummy_slot, dtype=torch.long, device=dev),
244
+ "cu": torch.arange(n + 1, dtype=torch.int32, device=dev),
245
+ "kstart": torch.full((n + 1,), self.dummy_slot, dtype=torch.int32, device=dev),
246
+ "kused": torch.ones(n, dtype=torch.int32, device=dev),
247
+ }
248
+ batch = AttnBatch(slots=st["slots"], n_causal=n,
249
+ causal=(st["cu"], st["kstart"], st["kused"], 1, self.max_k))
250
+
251
+ def run():
252
+ x = lm.model.embed_tokens(st["ids"])
253
+ h = lm(x, None, batch, eng.arena, cos_sin=(st["cos"], st["sin"]))
254
+ return lm.lm_head(h).float()
255
+
256
+ s = torch.cuda.Stream(device=dev)
257
+ s.wait_stream(torch.cuda.current_stream(dev))
258
+ with torch.cuda.stream(s):
259
+ for _ in range(2):
260
+ run()
261
+ torch.cuda.current_stream(dev).wait_stream(s)
262
+ g = torch.cuda.CUDAGraph()
263
+ with torch.cuda.graph(g, pool=self.pool, stream=torch.cuda.current_stream(dev),
264
+ capture_error_mode="thread_local"):
265
+ out = run()
266
+ self.pool = g.pool()
267
+ st["out"] = out
268
+ self.graphs[n] = (g, st)
269
+
270
+ def run(self, ids, positions, slots, kstart, kused):
271
+ b = len(ids)
272
+ n = self.bucket(b)
273
+ if n is None or self.failed:
274
+ return None
275
+ try:
276
+ if n not in self.graphs:
277
+ self._capture(n)
278
+ except Exception:
279
+ self.failed = True
280
+ return None
281
+ g, st = self.graphs[n]
282
+ dev = self.engine.device
283
+ lm = self.engine.model.lm
284
+ cos, sin = lm.rotary(_positions_tensor(positions, dev), self.engine.model.dtype)
285
+ st["ids"][:b].copy_(torch.tensor(ids, dtype=torch.long), non_blocking=True)
286
+ st["cos"][:b].copy_(cos)
287
+ st["sin"][:b].copy_(sin)
288
+ st["slots"][:b].copy_(torch.tensor(slots, dtype=torch.long), non_blocking=True)
289
+ st["slots"][b:].fill_(self.dummy_slot)
290
+ ks = torch.tensor(kstart + [self.dummy_slot] * (n - b + 1), dtype=torch.int32)
291
+ st["kstart"].copy_(ks, non_blocking=True)
292
+ ku = torch.tensor(kused + [1] * (n - b), dtype=torch.int32)
293
+ st["kused"].copy_(ku, non_blocking=True)
294
+ g.replay()
295
+ return st["out"][:b]
296
+
297
+
298
+ def auto_kv_tokens(model, headroom_gib: float = 10.0) -> int:
299
+ """KV arena size from the GPU memory that is free -- including memory PyTorch
300
+ still caches from an earlier, closed engine -- minus headroom for activations."""
301
+ device = model.device
302
+ free, _ = torch.cuda.mem_get_info(device)
303
+ free += torch.cuda.memory_reserved(device) - torch.cuda.memory_allocated(device)
304
+ cfg = model.lm.cfg
305
+ per_token = 2 * cfg["num_hidden_layers"] * cfg["num_key_value_heads"] * (
306
+ cfg["hidden_size"] // cfg["num_attention_heads"]) * 2
307
+ tokens = int((free - headroom_gib * 2 ** 30) * 0.8 / per_token)
308
+ if tokens < 16384:
309
+ raise RuntimeError(f"not enough free GPU memory for the KV cache ({free / 2 ** 30:.1f} GiB free); "
310
+ "close other engines on this GPU or pass kv_cache_tokens")
311
+ return tokens
312
+
313
+
314
+ class MoTEngine:
315
+ """Owns the model on one device and serves TextJob / ImageJob requests."""
316
+
317
+ def __init__(self, model, *, kv_tokens: int, max_step_tokens: int = 16384,
318
+ max_image_step_tokens: int = 12288, prefix_cache_fraction: float = 0.3,
319
+ min_cache_prefix: int = 256, use_cuda_graphs: bool = True):
320
+ self.model = model # adapter: see BagelModel for the interface
321
+ self.device = model.device
322
+ lm = model.lm
323
+ cfg = lm.cfg
324
+ self.arena = KVArena(lm.num_layers, kv_tokens, cfg["num_key_value_heads"],
325
+ cfg["hidden_size"] // cfg["num_attention_heads"], model.dtype, self.device)
326
+ self.alloc = RegionAllocator(kv_tokens - 1) # last slot: scratch for padded CUDA-graph rows
327
+ self.graphs = _DecodeGraphs(self, max_k=min(kv_tokens, 32768)) if use_cuda_graphs else None
328
+ self.cache = PrefixCache(int(kv_tokens * prefix_cache_fraction))
329
+ self.max_step_tokens = max_step_tokens
330
+ self.max_image_step_tokens = max_image_step_tokens
331
+ self.min_cache_prefix = min_cache_prefix
332
+ self.lock = threading.RLock() # allocator + prefix cache
333
+ self.inbox = {"text": queue.Queue(), "image": queue.Queue()}
334
+ self.pending = {"text": [], "image": []}
335
+ self.texts: list[_Text] = []
336
+ self.images: list[_Img] = []
337
+ self.stats = {"steps": 0, "tokens": 0, "prefix_hit_tokens": 0, "image_steps": 0}
338
+ self._stop = False
339
+ # Text (latency-critical decoding) and image generation run in separate
340
+ # threads on separate CUDA streams so that they overlap on the GPU.
341
+ self._threads = [threading.Thread(target=self._run, args=(kind,), daemon=True, name=f"simit-mot-{kind}")
342
+ for kind in ("text", "image")]
343
+ for t in self._threads:
344
+ t.start()
345
+
346
+ # -------------------------------------------------------------- submit
347
+ def submit(self, job) -> Future:
348
+ self.inbox["text" if isinstance(job, TextJob) else "image"].put(job)
349
+ return job.future
350
+
351
+ def close(self):
352
+ """Stop the engine threads and give the KV arena and graph memory back."""
353
+ self._stop = True
354
+ for q in self.inbox.values():
355
+ q.put(None)
356
+ for t in self._threads:
357
+ if t is not threading.current_thread():
358
+ t.join(timeout=10)
359
+ self.arena = None
360
+ self.graphs = None
361
+ self.cache = None
362
+ if torch.cuda.is_available():
363
+ torch.cuda.empty_cache()
364
+
365
+ # ---------------------------------------------------------------- loop
366
+ def _active(self, kind):
367
+ return self.texts if kind == "text" else self.images
368
+
369
+ def _run(self, kind):
370
+ torch.set_grad_enabled(False)
371
+ # Decode is on every request's critical path: its kernels get scheduling
372
+ # priority over the (compute-bound) diffusion kernels.
373
+ lo, hi = torch.cuda.Stream.priority_range()
374
+ stream = torch.cuda.Stream(device=self.device, priority=hi if kind == "text" else lo)
375
+ inbox, pending = self.inbox[kind], self.pending[kind]
376
+ with torch.cuda.stream(stream):
377
+ while not self._stop:
378
+ idle = not (pending or self._active(kind))
379
+ try:
380
+ job = inbox.get(timeout=0.2 if idle else 0)
381
+ if job is not None:
382
+ pending.append(job)
383
+ while True:
384
+ job = inbox.get_nowait()
385
+ if job is not None:
386
+ pending.append(job)
387
+ except queue.Empty:
388
+ pass
389
+ if not (pending or self._active(kind)):
390
+ continue
391
+ try:
392
+ self._admit(kind)
393
+ if self._active(kind):
394
+ with torch.inference_mode():
395
+ self._step(kind)
396
+ except Exception as e: # never kill the engine thread; fail the affected requests
397
+ tb = traceback.format_exc()
398
+ active = self._active(kind)
399
+ for r in active:
400
+ if not r.job.future.done():
401
+ r.job.future.set_exception(RuntimeError(f"engine step failed: {e}\n{tb}"))
402
+ self._release(r)
403
+ active.clear()
404
+ torch.cuda.empty_cache()
405
+
406
+ # ------------------------------------------------------------ admission
407
+ def _region(self, size: int) -> Optional[int]:
408
+ start = self.alloc.alloc(size)
409
+ while start is None:
410
+ victim = next(self.cache.evictable(), None)
411
+ if victim is None:
412
+ return None
413
+ e = self.cache.remove(victim[0])
414
+ self.alloc.release(e.start, e.length)
415
+ start = self.alloc.alloc(size)
416
+ return start
417
+
418
+ def _keys(self, chunks):
419
+ keys, text_prefix = [], 0
420
+ seen_image = False
421
+ for c in chunks:
422
+ if isinstance(c, TextChunk):
423
+ keys.extend(c.ids)
424
+ if not seen_image:
425
+ text_prefix += len(c.ids)
426
+ else:
427
+ seen_image = True
428
+ keys.extend(f"{c.key}:{i}" for i in range(c.n))
429
+ return keys, text_prefix
430
+
431
+ def _admit(self, kind):
432
+ still = []
433
+ for job in self.pending[kind]:
434
+ if job.future.cancelled():
435
+ continue
436
+ with self.lock:
437
+ ok = self._admit_text(job) if kind == "text" else self._admit_image(job)
438
+ if not ok:
439
+ still.append(job)
440
+ self.pending[kind][:] = still
441
+
442
+ def _admit_text(self, job: TextJob) -> bool:
443
+ keys, text_prefix = self._keys(job.chunks)
444
+ n_prompt = len(keys)
445
+ cap = n_prompt + job.max_new_tokens + 1
446
+ hashes = _prefix_hashes(keys[:text_prefix]) if text_prefix >= _G else []
447
+ entry, hit = self.cache.lookup([(n, h) for n, h in hashes if n < n_prompt])
448
+ if entry is not None and not entry.ready:
449
+ return False # an identical long prefix is being computed right now: wait and reuse it
450
+ start = self._region(cap)
451
+ if start is None:
452
+ return False
453
+ seq = _Text(job, start, cap, keys, n_prompt)
454
+ if entry is not None and hit >= _G:
455
+ self.arena.copy(entry.start, start, hit)
456
+ self._skip(seq, hit)
457
+ self.stats["prefix_hit_tokens"] += hit
458
+ if text_prefix >= self.min_cache_prefix:
459
+ seq.cache_eid = self.cache.add(start, text_prefix, hashes, owner=seq)
460
+ self.texts.append(seq)
461
+ return True
462
+
463
+ def _skip(self, seq, n: int):
464
+ """Mark the first ``n`` prompt tokens (text only) as already written."""
465
+ seq.written = n
466
+ left = n
467
+ while left > 0:
468
+ c = seq.job.chunks[seq.chunk_i]
469
+ take = min(left, len(c.ids) - seq.chunk_off)
470
+ seq.chunk_off += take
471
+ left -= take
472
+ if seq.chunk_off == len(c.ids):
473
+ seq.chunk_i, seq.chunk_off = seq.chunk_i + 1, 0
474
+ seq.pos = self.model.position_after(seq.job.chunks, n)
475
+
476
+ def _admit_image(self, job: ImageJob) -> bool:
477
+ m = self.model
478
+ h, w = job.height // m.latent_downsample, job.width // m.latent_downsample
479
+ n_lat = h * w
480
+ n_prompt = sum(len(c.ids) for c in job.chunks)
481
+ cap = n_prompt + n_lat + 2
482
+ start = self._region(cap)
483
+ if start is None:
484
+ return False
485
+ uncond = self._region(n_lat + 2)
486
+ if uncond is None:
487
+ self.alloc.release(start, cap)
488
+ return False
489
+ gen = torch.Generator(device="cpu")
490
+ gen.manual_seed(job.seed if job.seed is not None else torch.seed() % (2 ** 31))
491
+ noise = torch.randn(n_lat, m.latent_dim, generator=gen).to(self.device, torch.float32)
492
+ self.images.append(_Img(job, start, cap, uncond, n_lat, h, w, m.latent_pos_ids(h, w).to(self.device),
493
+ noise, m.denoise_schedule()))
494
+ return True
495
+
496
+ # ---------------------------------------------------------------- step
497
+ def _drop_cancelled(self, kind):
498
+ active = self._active(kind)
499
+ live = []
500
+ for r in active:
501
+ if r.job.future.cancelled():
502
+ self._release(r)
503
+ else:
504
+ live.append(r)
505
+ active[:] = live
506
+
507
+ def _step(self, kind):
508
+ self._drop_cancelled(kind)
509
+ m, dev = self.model, self.device
510
+ budget = self.max_step_tokens if kind == "text" else self.max_image_step_tokens
511
+ texts = self.texts if kind == "text" else []
512
+ images = self.images if kind == "image" else []
513
+ causal, full = [], [] # pieces: dict(kind, ...)
514
+ # 1) decode tokens
515
+ for s in texts:
516
+ if s.prefilled and s.job.score_len == 0 and s.last_token is not None:
517
+ causal.append(dict(kind="decode", seq=s, ids=[s.last_token], pos=m.text_positions(s.pos, 1),
518
+ logits=True))
519
+ budget -= 1
520
+ # 2) denoising steps (as many images as fit the step budget; the rest wait a step)
521
+ for g in images:
522
+ if g.prefilled:
523
+ t = g.steps[g.step_i][0]
524
+ need = (g.n_lat + 2) * (2 if m.cfg_active(t) else 1)
525
+ if need > budget and full:
526
+ continue
527
+ full.append(dict(kind="latent", job=g, cond=True))
528
+ if m.cfg_active(t):
529
+ full.append(dict(kind="latent", job=g, cond=False))
530
+ budget -= need
531
+ # 3) prefill chunks within the remaining budget
532
+ for r in texts + images:
533
+ if r.prefilled:
534
+ continue
535
+ while not r.prefilled and budget > 0:
536
+ c = r.job.chunks[r.chunk_i]
537
+ if isinstance(c, ImageChunk):
538
+ if c.n > budget and (causal or full):
539
+ break
540
+ full.append(dict(kind="image", req=r, chunk=c, pos=r.pos, start_written=r.written))
541
+ r.written += c.n
542
+ r.pos = m.position_after_image(r.pos, c)
543
+ r.chunk_i += 1
544
+ budget -= c.n
545
+ else:
546
+ take = min(len(c.ids) - r.chunk_off, max(budget, 1))
547
+ ids = c.ids[r.chunk_off:r.chunk_off + take]
548
+ pos = m.text_positions(r.pos, take)
549
+ r.chunk_off += take
550
+ if r.chunk_off == len(c.ids):
551
+ r.chunk_i, r.chunk_off = r.chunk_i + 1, 0
552
+ final = r.prefilled
553
+ causal.append(dict(kind="prefill", req=r, ids=ids, pos=pos, start_written=r.written, final=final))
554
+ r.written += take
555
+ r.pos += take
556
+ budget -= take
557
+ if budget <= 0:
558
+ break
559
+ if not causal and not full:
560
+ return
561
+ self._forward(causal, full, kind)
562
+
563
+ def _forward(self, causal, full, kind):
564
+ m, dev = self.model, self.device
565
+ lm = m.lm
566
+ if (self.graphs is not None and not full and causal and all(p["kind"] == "decode" for p in causal)
567
+ and max(p["seq"].written + 1 for p in causal) <= self.graphs.max_k):
568
+ seqs = [p["seq"] for p in causal]
569
+ logits = self.graphs.run([p["ids"][0] for p in causal], [p["pos"][0] for p in causal],
570
+ [s.region + s.written for s in seqs], [s.region for s in seqs],
571
+ [s.written + 1 for s in seqs])
572
+ if logits is not None:
573
+ self.stats["steps"] += 1
574
+ self.stats["tokens"] += len(seqs)
575
+ self._consume_logits(logits, [(i, 1, p) for i, p in enumerate(causal)])
576
+ self._finish(kind)
577
+ return
578
+ embeds, positions, slots = [], [], []
579
+ logit_rows, latent_rows = [], [] # (row_start, n, piece)
580
+ q_lens_c, kstart_c, kused_c = [], [], []
581
+ q_lens_f, kstart_f, kused_f = [], [], []
582
+ token_ids, token_rows = [], []
583
+ image_pieces = []
584
+ row = 0
585
+
586
+ def add_text(ids, pos, region, written, n_used):
587
+ nonlocal row
588
+ token_ids.extend(ids)
589
+ token_rows.extend(range(row, row + len(ids)))
590
+ embeds.append(None)
591
+ positions.extend(pos)
592
+ slots.extend(range(region + written, region + written + len(ids)))
593
+ row += len(ids)
594
+
595
+ for p in causal:
596
+ if p["kind"] == "decode":
597
+ s = p["seq"]
598
+ start_row = row
599
+ add_text(p["ids"], p["pos"], s.region, s.written, s.written + 1)
600
+ q_lens_c.append(1); kstart_c.append(s.region); kused_c.append(s.written + 1)
601
+ logit_rows.append((start_row, 1, p))
602
+ else:
603
+ r = p["req"]
604
+ start_row = row
605
+ n = len(p["ids"])
606
+ add_text(p["ids"], p["pos"], r.region, p["start_written"], p["start_written"] + n)
607
+ q_lens_c.append(n); kstart_c.append(r.region); kused_c.append(p["start_written"] + n)
608
+ if isinstance(r, _Text):
609
+ a = p["start_written"]
610
+ if r.job.score_len:
611
+ lo, hi = max(a, r.score_lo), min(a + n - 1, r.score_hi)
612
+ if lo <= hi:
613
+ p["score_from"] = lo
614
+ logit_rows.append((start_row + lo - a, hi - lo + 1, p))
615
+ elif p["final"]:
616
+ logit_rows.append((start_row + n - 1, 1, p))
617
+ n_causal = row
618
+ for p in full:
619
+ if p["kind"] == "image":
620
+ r, c = p["req"], p["chunk"]
621
+ start_row = row
622
+ image_pieces.append((start_row, p))
623
+ positions.extend(m.image_positions(p["pos"], c))
624
+ slots.extend(range(r.region + p["start_written"], r.region + p["start_written"] + c.n))
625
+ embeds.append(("image", start_row, c))
626
+ row += c.n
627
+ q_lens_f.append(c.n); kstart_f.append(r.region); kused_f.append(p["start_written"] + c.n)
628
+ # A latent block "[start_of_image] latents [end_of_image]" is attended
629
+ # bidirectionally, so its queries can be split into two pseudo-sequences
630
+ # over the same keys: the two markers (understanding expert) here and the
631
+ # latents (generation expert) last, giving each expert contiguous rows.
632
+ blocks = []
633
+ for p in full:
634
+ if p["kind"] == "latent":
635
+ g = p["job"]
636
+ region = g.region if p["cond"] else g.uncond_region
637
+ prefix = g.n_prompt if p["cond"] else 0
638
+ lpos = m.latent_positions(g.pos if p["cond"] else 0, g.h, g.w)
639
+ blocks.append((p, g, region, prefix, lpos))
640
+ for p, g, region, prefix, lpos in blocks:
641
+ n = g.n_lat + 2
642
+ token_ids.extend([m.start_of_image, m.end_of_image]); token_rows.extend([row, row + 1])
643
+ positions.extend([lpos[0], lpos[-1]])
644
+ slots.extend([region + prefix, region + prefix + n - 1])
645
+ row += 2
646
+ q_lens_f.append(2); kstart_f.append(region); kused_f.append(prefix + n)
647
+ n_und = row
648
+ for p, g, region, prefix, lpos in blocks:
649
+ n = g.n_lat + 2
650
+ embeds.append(("latent", row, g))
651
+ positions.extend(lpos[1:-1])
652
+ slots.extend(range(region + prefix + 1, region + prefix + 1 + g.n_lat))
653
+ latent_rows.append((row, g.n_lat, p))
654
+ row += g.n_lat
655
+ q_lens_f.append(g.n_lat); kstart_f.append(region); kused_f.append(prefix + n)
656
+
657
+ T = row
658
+ hidden = lm.cfg["hidden_size"]
659
+ x = torch.empty(T, hidden, dtype=m.dtype, device=dev)
660
+ if token_ids:
661
+ ids_t = torch.tensor(token_ids, dtype=torch.long, device=dev)
662
+ rows_t = torch.tensor(token_rows, dtype=torch.long, device=dev)
663
+ x[rows_t] = lm.model.embed_tokens(ids_t)
664
+ imgs = [e for e in embeds if e is not None and e[0] == "image"]
665
+ if imgs:
666
+ vit = m.encode_images([c.payload for _, _, c in imgs])
667
+ for (_, start_row, c), emb in zip(imgs, vit):
668
+ # start marker / ViT tokens / end marker
669
+ x[start_row] = lm.model.embed_tokens.weight[m.start_of_image]
670
+ x[start_row + 1:start_row + 1 + emb.shape[0]] = emb
671
+ x[start_row + c.n - 1] = lm.model.embed_tokens.weight[m.end_of_image]
672
+ for e in embeds:
673
+ if e is not None and e[0] == "latent":
674
+ _, start_row, g = e
675
+ t = g.steps[g.step_i][0]
676
+ x[start_row:start_row + g.n_lat] = m.embed_latents(g.x, t, g.lat_pos_ids)
677
+
678
+ pos_t = _positions_tensor(positions, dev)
679
+ slots_t = torch.tensor(slots, dtype=torch.long, device=dev)
680
+
681
+ def meta(q_lens, kstart, kused):
682
+ if not q_lens:
683
+ return None
684
+ cu = torch.tensor([0] + list(itertools.accumulate(q_lens)), dtype=torch.int32, device=dev)
685
+ ks = torch.tensor(kstart + [kstart[-1] + kused[-1]], dtype=torch.int32, device=dev)
686
+ ku = torch.tensor(kused, dtype=torch.int32, device=dev)
687
+ return (cu, ks, ku, max(q_lens), max(kused))
688
+
689
+ batch = AttnBatch(slots=slots_t, n_causal=n_causal, causal=meta(q_lens_c, kstart_c, kused_c),
690
+ full=meta(q_lens_f, kstart_f, kused_f), n_und=n_und if blocks else None)
691
+ h = lm(x, pos_t, batch, self.arena)
692
+ self.stats["steps" if kind == "text" else "image_steps"] += 1
693
+ self.stats["tokens"] += T
694
+
695
+ # ---- text outputs
696
+ if logit_rows:
697
+ rows = torch.cat([torch.arange(r0, r0 + n, device=dev) for r0, n, _ in logit_rows])
698
+ logits = lm.lm_head(h[rows]).float()
699
+ self._consume_logits(logits, logit_rows)
700
+ # ---- latent outputs
701
+ if latent_rows:
702
+ self._consume_latents(h, latent_rows)
703
+ with self.lock:
704
+ for s in self.texts if kind == "text" else []:
705
+ if s.cache_eid is not None:
706
+ e = self.cache.entries.get(s.cache_eid)
707
+ if e is not None and not e.ready and s.written >= e.length:
708
+ e.ready = True
709
+ self._finish(kind)
710
+
711
+ # ------------------------------------------------------------- outputs
712
+ def _consume_logits(self, logits, logit_rows):
713
+ off = 0
714
+ sample_rows, sample_seqs = [], []
715
+ for r0, n, p in logit_rows:
716
+ seq = p["seq"] if p["kind"] == "decode" else p["req"]
717
+ if seq.job.score_len:
718
+ # the row at key index i predicts keys[i + 1]
719
+ i0 = p["score_from"]
720
+ targets = torch.tensor(seq.keys[i0 + 1:i0 + 1 + n], device=logits.device)
721
+ lp = torch.log_softmax(logits[off:off + n], dim=-1)
722
+ seq.score_lp.extend(lp[torch.arange(n, device=logits.device), targets].tolist())
723
+ if i0 + n - 1 >= seq.score_hi:
724
+ seq.done = True
725
+ else:
726
+ sample_rows.append(off)
727
+ sample_seqs.append(seq)
728
+ off += n
729
+ if not sample_rows:
730
+ return
731
+ lg = logits[sample_rows]
732
+ tokens, lps = self._sample(lg, sample_seqs)
733
+ for seq, tok, lp in zip(sample_seqs, tokens, lps):
734
+ if seq.last_token is not None:
735
+ seq.pos += 1 # the decode token just consumed was written at seq.pos
736
+ seq.written += 1
737
+ if tok in seq.job.stop_ids or len(seq.out_ids) >= seq.job.max_new_tokens:
738
+ seq.done = True
739
+ continue
740
+ seq.out_ids.append(tok)
741
+ if seq.job.logprobs:
742
+ seq.out_lp.append(lp)
743
+ seq.last_token = tok
744
+ if len(seq.out_ids) >= seq.job.max_new_tokens:
745
+ seq.done = True
746
+ if seq.job.on_tokens is not None and len(seq.out_ids) % 4 == 0:
747
+ try:
748
+ seq.job.on_tokens(seq.out_ids)
749
+ except Exception:
750
+ pass
751
+
752
+ def _sample(self, logits, seqs):
753
+ raw_lp = torch.log_softmax(logits, dim=-1)
754
+ lg = logits.clone()
755
+ for i, s in enumerate(seqs):
756
+ rp = s.job.repetition_penalty
757
+ if rp != 1.0 and s.out_ids:
758
+ seen = torch.tensor(sorted(set(s.out_ids)), device=lg.device)
759
+ vals = lg[i, seen]
760
+ lg[i, seen] = torch.where(vals < 0, vals * rp, vals / rp)
761
+ temps = torch.tensor([max(s.job.temperature, 0.0) for s in seqs], device=lg.device)
762
+ greedy = temps == 0
763
+ tokens = lg.argmax(-1)
764
+ if not bool(greedy.all()):
765
+ probs = torch.softmax(lg / temps.clamp(min=1e-5)[:, None], dim=-1)
766
+ for i, s in enumerate(seqs):
767
+ if s.job.top_p < 1.0 and not greedy[i]:
768
+ sp, si = probs[i].sort(descending=True)
769
+ keep = (sp.cumsum(0) - sp) <= s.job.top_p
770
+ filt = torch.zeros_like(probs[i]).scatter(0, si[keep], sp[keep])
771
+ probs[i] = filt / filt.sum()
772
+ sampled = torch.multinomial(probs, 1).squeeze(1)
773
+ tokens = torch.where(greedy, tokens, sampled)
774
+ chosen_lp = raw_lp.gather(1, tokens[:, None]).squeeze(1)
775
+ return tokens.tolist(), chosen_lp.tolist()
776
+
777
+ def _consume_latents(self, h, latent_rows):
778
+ m = self.model
779
+ by_job: dict[int, dict] = {}
780
+ r_first = latent_rows[0][0] # latent rows are contiguous and in order
781
+ vel = m.latents_to_velocity(h[r_first:r_first + sum(n for _, n, _ in latent_rows)])
782
+ for r0, n, p in latent_rows:
783
+ v = vel[r0 - r_first:r0 - r_first + n]
784
+ by_job.setdefault(id(p["job"]), {"job": p["job"]})["cond" if p["cond"] else "uncond"] = v
785
+ for d in by_job.values():
786
+ g = d["job"]
787
+ t, dt = g.steps[g.step_i]
788
+ v = d["cond"]
789
+ if "uncond" in d:
790
+ v = m.apply_cfg(v, d["uncond"])
791
+ g.x = g.x - v.float() * dt
792
+ g.step_i += 1
793
+
794
+ def _finish(self, kind):
795
+ m = self.model
796
+ if kind == "text":
797
+ keep = []
798
+ for s in self.texts:
799
+ if s.done:
800
+ self._release(s)
801
+ if not s.job.future.done():
802
+ s.job.future.set_result(s)
803
+ else:
804
+ keep.append(s)
805
+ self.texts[:] = keep
806
+ return
807
+ done_imgs, keep = [], []
808
+ for g in self.images:
809
+ (done_imgs if g.step_i >= len(g.steps) else keep).append(g)
810
+ self.images[:] = keep
811
+ if done_imgs:
812
+ pics = m.decode_latents([(g.x, g.h, g.w) for g in done_imgs])
813
+ for g, pic in zip(done_imgs, pics):
814
+ self._release(g)
815
+ if not g.job.future.done():
816
+ g.job.future.set_result(pic)
817
+
818
+ def _release(self, r):
819
+ with self.lock:
820
+ self._release_locked(r)
821
+
822
+ def _release_locked(self, r):
823
+ if isinstance(r, _Text):
824
+ if r.cache_eid is not None and r.cache_eid in self.cache.entries and r.written >= self.cache.entries[r.cache_eid].length:
825
+ e = self.cache.entries[r.cache_eid]
826
+ e.owner, e.ready = None, True
827
+ # keep only the cached prefix; give back the rest of the region
828
+ self.alloc.release(r.region + e.length, r.cap - e.length)
829
+ self._trim_cache()
830
+ else:
831
+ if r.cache_eid is not None and r.cache_eid in self.cache.entries:
832
+ self.cache.remove(r.cache_eid)
833
+ self.alloc.release(r.region, r.cap)
834
+ else:
835
+ self.alloc.release(r.region, r.cap)
836
+ self.alloc.release(r.uncond_region, r.n_lat + 2)
837
+
838
+ def _trim_cache(self):
839
+ while self.cache.tokens > self.cache.max_tokens:
840
+ victim = next(self.cache.evictable(), None)
841
+ if victim is None:
842
+ return
843
+ e = self.cache.remove(victim[0])
844
+ self.alloc.release(e.start, e.length)