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.
- simit/__init__.py +39 -0
- simit/api.py +258 -0
- simit/backends/__init__.py +112 -0
- simit/backends/bagel.py +359 -0
- simit/backends/base.py +114 -0
- simit/backends/hf.py +336 -0
- simit/backends/imagegen.py +136 -0
- simit/backends/lance.py +338 -0
- simit/backends/mot/engine.py +844 -0
- simit/backends/mot/kernels.py +130 -0
- simit/backends/mot/modeling.py +514 -0
- simit/backends/mot/wan_vae.py +872 -0
- simit/backends/vllm.py +204 -0
- simit/config.py +112 -0
- simit/metrics.py +113 -0
- simit/pipeline/core.py +586 -0
- simit/pipeline/prompts.py +167 -0
- simit/skills/__init__.py +6 -0
- simit/skills/assets/MERMAID_LICENSE +21 -0
- simit/skills/assets/mermaid.min.js +3587 -0
- simit/skills/base.py +172 -0
- simit/skills/builtin.py +614 -0
- simit/skills/helpers.py +197 -0
- simit/skills/pool.py +262 -0
- simit/skills/prompts.py +2374 -0
- simit/skills/renderers.py +1292 -0
- simit/skills/routing.py +89 -0
- simit/skills/sandbox.py +188 -0
- simit/skills/web.py +253 -0
- simit/tune.py +323 -0
- simit/types.py +96 -0
- simit/utils.py +67 -0
- simit-0.1.0.dist-info/METADATA +388 -0
- simit-0.1.0.dist-info/RECORD +38 -0
- simit-0.1.0.dist-info/WHEEL +5 -0
- simit-0.1.0.dist-info/licenses/LICENSE +202 -0
- simit-0.1.0.dist-info/licenses/src/simit/skills/assets/MERMAID_LICENSE +21 -0
- simit-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -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)
|