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
simit/backends/hf.py
ADDED
|
@@ -0,0 +1,336 @@
|
|
|
1
|
+
"""Any Hugging Face ``transformers`` vision-language model.
|
|
2
|
+
|
|
3
|
+
Requests are batched by a worker thread: whatever is queued when the GPU
|
|
4
|
+
becomes free runs as one padded ``generate`` call (grouped by sampling
|
|
5
|
+
settings), so concurrent pipeline steps share forward passes.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import inspect
|
|
11
|
+
import queue
|
|
12
|
+
import threading
|
|
13
|
+
import traceback
|
|
14
|
+
from concurrent.futures import Future
|
|
15
|
+
from typing import Optional
|
|
16
|
+
|
|
17
|
+
import torch
|
|
18
|
+
from PIL import Image
|
|
19
|
+
|
|
20
|
+
from ..utils import strip_thinking, to_pil
|
|
21
|
+
from .base import Backend, GenRequest, GenResult, ScoreRequest
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class _RowStop:
|
|
25
|
+
"""Per-row stopping (each request keeps its own max_new_tokens)."""
|
|
26
|
+
|
|
27
|
+
def __init__(self, prompt_len: int, limits: list[int], eos: set[int]):
|
|
28
|
+
self.prompt_len, self.limits, self.eos = prompt_len, torch.tensor(limits), eos
|
|
29
|
+
|
|
30
|
+
def __call__(self, input_ids, scores, **kw):
|
|
31
|
+
n = input_ids.shape[1] - self.prompt_len
|
|
32
|
+
done = self.limits.to(input_ids.device) <= n
|
|
33
|
+
if self.eos:
|
|
34
|
+
last = input_ids[:, -1]
|
|
35
|
+
for e in self.eos:
|
|
36
|
+
done |= last == e
|
|
37
|
+
return done
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class HFBackend(Backend):
|
|
41
|
+
name = "hf"
|
|
42
|
+
native_image_generation = False
|
|
43
|
+
verify_think = False
|
|
44
|
+
|
|
45
|
+
def __init__(self, model, processor=None, *, device_map="auto", dtype=torch.bfloat16,
|
|
46
|
+
max_batch_size: int = 16, max_batch_tokens: int = 65536,
|
|
47
|
+
max_image_pixels: Optional[int] = 1024 * 1024, think: bool = False,
|
|
48
|
+
attn_implementation: Optional[str] = None, trust_remote_code: bool = False, **load_kwargs):
|
|
49
|
+
import os
|
|
50
|
+
# FP8 checkpoints: use transformers' Triton kernels, not DeepGEMM (which JIT-compiles with the system nvcc).
|
|
51
|
+
os.environ.setdefault("TRANSFORMERS_DISABLE_DEEPGEMM_LINEAR", "1")
|
|
52
|
+
from transformers import AutoProcessor
|
|
53
|
+
if isinstance(model, str):
|
|
54
|
+
self.model_id = model
|
|
55
|
+
kw = dict(dtype=dtype, device_map=device_map, trust_remote_code=trust_remote_code, **load_kwargs)
|
|
56
|
+
if attn_implementation:
|
|
57
|
+
kw["attn_implementation"] = attn_implementation
|
|
58
|
+
self.model = _load_model(model, kw)
|
|
59
|
+
self.processor = processor or AutoProcessor.from_pretrained(model, trust_remote_code=trust_remote_code)
|
|
60
|
+
else:
|
|
61
|
+
if processor is None:
|
|
62
|
+
raise ValueError("pass the model's processor along with an already-loaded model")
|
|
63
|
+
self.model, self.processor = model, processor
|
|
64
|
+
self.model_id = getattr(model.config, "_name_or_path", "model")
|
|
65
|
+
self.model.eval()
|
|
66
|
+
self.tokenizer = getattr(self.processor, "tokenizer", self.processor)
|
|
67
|
+
self.tokenizer.padding_side = "left"
|
|
68
|
+
if self.tokenizer.pad_token_id is None:
|
|
69
|
+
self.tokenizer.pad_token = self.tokenizer.eos_token
|
|
70
|
+
self.think = think
|
|
71
|
+
self.max_batch_size = max_batch_size
|
|
72
|
+
self.max_batch_tokens = max_batch_tokens
|
|
73
|
+
self.max_image_pixels = max_image_pixels
|
|
74
|
+
gc = getattr(self.model, "generation_config", None)
|
|
75
|
+
eos = getattr(gc, "eos_token_id", None) if gc is not None else None
|
|
76
|
+
eos = eos if isinstance(eos, list) else [eos] if eos is not None else []
|
|
77
|
+
if self.tokenizer.eos_token_id is not None:
|
|
78
|
+
eos.append(self.tokenizer.eos_token_id)
|
|
79
|
+
self.eos_ids = sorted({e for e in eos if e is not None})
|
|
80
|
+
self.device = next(self.model.parameters()).device
|
|
81
|
+
try:
|
|
82
|
+
self._logits_to_keep = "logits_to_keep" in inspect.signature(self.model.forward).parameters
|
|
83
|
+
except (TypeError, ValueError):
|
|
84
|
+
self._logits_to_keep = False
|
|
85
|
+
self._queue: "queue.Queue" = queue.Queue()
|
|
86
|
+
self._thread = threading.Thread(target=self._loop, daemon=True, name="simit-hf")
|
|
87
|
+
self._thread.start()
|
|
88
|
+
|
|
89
|
+
# ------------------------------------------------------------- prompts
|
|
90
|
+
def _image(self, img: Image.Image) -> Image.Image:
|
|
91
|
+
img = to_pil(img)
|
|
92
|
+
if self.max_image_pixels and img.width * img.height > self.max_image_pixels:
|
|
93
|
+
s = (self.max_image_pixels / (img.width * img.height)) ** 0.5
|
|
94
|
+
img = img.resize((max(28, int(img.width * s)), max(28, int(img.height * s))), Image.BICUBIC)
|
|
95
|
+
return img
|
|
96
|
+
|
|
97
|
+
def _render(self, parts, think: bool, add_generation_prompt: bool = True):
|
|
98
|
+
content, images = [], []
|
|
99
|
+
for p in parts:
|
|
100
|
+
if isinstance(p, Image.Image):
|
|
101
|
+
content.append({"type": "image"})
|
|
102
|
+
images.append(self._image(p))
|
|
103
|
+
elif content and content[-1]["type"] == "text":
|
|
104
|
+
content[-1]["text"] += p
|
|
105
|
+
else:
|
|
106
|
+
content.append({"type": "text", "text": p})
|
|
107
|
+
messages = [{"role": "user", "content": content}]
|
|
108
|
+
try:
|
|
109
|
+
text = self.processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=add_generation_prompt,
|
|
110
|
+
enable_thinking=think)
|
|
111
|
+
except TypeError:
|
|
112
|
+
text = self.processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=add_generation_prompt)
|
|
113
|
+
return text, images
|
|
114
|
+
|
|
115
|
+
def _inputs(self, texts, images):
|
|
116
|
+
kw = dict(text=texts, return_tensors="pt", padding=True)
|
|
117
|
+
flat = [im for ims in images for im in ims]
|
|
118
|
+
if flat:
|
|
119
|
+
kw["images"] = flat
|
|
120
|
+
return self.processor(**kw).to(self.device)
|
|
121
|
+
|
|
122
|
+
# ------------------------------------------------------------ requests
|
|
123
|
+
def submit_generate(self, req: GenRequest) -> Future:
|
|
124
|
+
fut: Future = Future()
|
|
125
|
+
self._queue.put(("gen", req, fut))
|
|
126
|
+
return fut
|
|
127
|
+
|
|
128
|
+
def submit_score(self, req: ScoreRequest) -> Future:
|
|
129
|
+
fut: Future = Future()
|
|
130
|
+
self._queue.put(("score", req, fut))
|
|
131
|
+
return fut
|
|
132
|
+
|
|
133
|
+
def close(self):
|
|
134
|
+
self._queue.put(None)
|
|
135
|
+
|
|
136
|
+
# --------------------------------------------------------------- worker
|
|
137
|
+
def _loop(self):
|
|
138
|
+
torch.set_grad_enabled(False)
|
|
139
|
+
while True:
|
|
140
|
+
item = self._queue.get()
|
|
141
|
+
if item is None:
|
|
142
|
+
return
|
|
143
|
+
items = [item]
|
|
144
|
+
try:
|
|
145
|
+
while True:
|
|
146
|
+
nxt = self._queue.get_nowait()
|
|
147
|
+
if nxt is None:
|
|
148
|
+
self._queue.put(None)
|
|
149
|
+
break
|
|
150
|
+
items.append(nxt)
|
|
151
|
+
except queue.Empty:
|
|
152
|
+
pass
|
|
153
|
+
items = [it for it in items if not it[2].cancelled()]
|
|
154
|
+
gens = [it for it in items if it[0] == "gen"]
|
|
155
|
+
scores = [it for it in items if it[0] == "score"]
|
|
156
|
+
# Long generations (code, captions) must not hold short ones hostage.
|
|
157
|
+
groups: dict = {}
|
|
158
|
+
for it in gens:
|
|
159
|
+
r = it[1]
|
|
160
|
+
key = (r.temperature > 0, round(r.temperature, 3), r.top_p, r.repetition_penalty,
|
|
161
|
+
r.think or self.think, r.logprobs, r.max_new_tokens > 256)
|
|
162
|
+
groups.setdefault(key, []).append(it)
|
|
163
|
+
# Cheap work first (scoring, short answers like routing), long generations last.
|
|
164
|
+
for chunk in self._chunks(scores):
|
|
165
|
+
self._safe(self._run_score, chunk)
|
|
166
|
+
for key, group in sorted(groups.items(), key=lambda kv: max(it[1].max_new_tokens for it in kv[1])):
|
|
167
|
+
for chunk in self._chunks(group):
|
|
168
|
+
self._safe(self._run_gen, chunk)
|
|
169
|
+
|
|
170
|
+
def _cost(self, req) -> int:
|
|
171
|
+
"""Rough prompt length in tokens (images at ~one token per 28x28 pixels)."""
|
|
172
|
+
n = 0
|
|
173
|
+
for p in req.parts:
|
|
174
|
+
if isinstance(p, Image.Image):
|
|
175
|
+
px = p.width * p.height
|
|
176
|
+
n += min(px, self.max_image_pixels or px) // 784 + 1
|
|
177
|
+
else:
|
|
178
|
+
n += len(p) // 3 + 1
|
|
179
|
+
return n + len(getattr(req, "target", "")) // 3 + getattr(req, "max_new_tokens", 0)
|
|
180
|
+
|
|
181
|
+
def _chunks(self, items):
|
|
182
|
+
"""Batches of at most ``max_batch_size`` requests and ``max_batch_tokens`` tokens."""
|
|
183
|
+
batch, tokens = [], 0
|
|
184
|
+
for it in items:
|
|
185
|
+
c = self._cost(it[1])
|
|
186
|
+
if batch and (len(batch) >= self.max_batch_size or tokens + c > self.max_batch_tokens):
|
|
187
|
+
yield batch
|
|
188
|
+
batch, tokens = [], 0
|
|
189
|
+
batch.append(it)
|
|
190
|
+
tokens += c
|
|
191
|
+
if batch:
|
|
192
|
+
yield batch
|
|
193
|
+
|
|
194
|
+
def _safe(self, fn, items):
|
|
195
|
+
try:
|
|
196
|
+
fn(items)
|
|
197
|
+
except (torch.cuda.OutOfMemoryError, RuntimeError) as e:
|
|
198
|
+
# cuBLAS reports a failed workspace allocation as its own error
|
|
199
|
+
if not isinstance(e, torch.cuda.OutOfMemoryError) and "CUBLAS_STATUS_" not in str(e):
|
|
200
|
+
tb = traceback.format_exc()
|
|
201
|
+
for _, _, fut in items:
|
|
202
|
+
if not fut.done():
|
|
203
|
+
fut.set_exception(RuntimeError(f"{e}\n{tb}"))
|
|
204
|
+
return
|
|
205
|
+
torch.cuda.empty_cache()
|
|
206
|
+
if len(items) > 1: # retry in halves, and keep later batches below this size
|
|
207
|
+
self.max_batch_tokens = max(2048, min(self.max_batch_tokens,
|
|
208
|
+
sum(self._cost(it[1]) for it in items) // 2))
|
|
209
|
+
mid = len(items) // 2
|
|
210
|
+
self._safe(fn, items[:mid])
|
|
211
|
+
self._safe(fn, items[mid:])
|
|
212
|
+
else:
|
|
213
|
+
items[0][2].set_exception(RuntimeError("CUDA out of memory for a single request"))
|
|
214
|
+
except Exception as e:
|
|
215
|
+
tb = traceback.format_exc()
|
|
216
|
+
for _, _, fut in items:
|
|
217
|
+
if not fut.done():
|
|
218
|
+
fut.set_exception(RuntimeError(f"{e}\n{tb}"))
|
|
219
|
+
|
|
220
|
+
def _run_gen(self, items):
|
|
221
|
+
reqs = [r for _, r, _ in items]
|
|
222
|
+
rendered = [self._render(r.parts, think=r.think or self.think) for r in reqs]
|
|
223
|
+
inputs = self._inputs([t for t, _ in rendered], [ims for _, ims in rendered])
|
|
224
|
+
r0 = reqs[0]
|
|
225
|
+
prompt_len = inputs["input_ids"].shape[1]
|
|
226
|
+
limits = [r.max_new_tokens for r in reqs]
|
|
227
|
+
from transformers import StoppingCriteriaList
|
|
228
|
+
gen_kw = dict(max_new_tokens=max(limits), return_dict_in_generate=True, output_logits=r0.logprobs,
|
|
229
|
+
stopping_criteria=StoppingCriteriaList([_RowStop(prompt_len, limits, set(self.eos_ids))]),
|
|
230
|
+
pad_token_id=self.tokenizer.pad_token_id, eos_token_id=self.eos_ids or None)
|
|
231
|
+
if r0.temperature > 0:
|
|
232
|
+
gen_kw.update(do_sample=True, temperature=r0.temperature, top_p=r0.top_p, top_k=0)
|
|
233
|
+
else:
|
|
234
|
+
gen_kw.update(do_sample=False, temperature=None, top_p=None, top_k=None)
|
|
235
|
+
if r0.repetition_penalty not in (None, 1.0):
|
|
236
|
+
gen_kw["repetition_penalty"] = r0.repetition_penalty
|
|
237
|
+
out = self.model.generate(**inputs, **gen_kw)
|
|
238
|
+
seqs = out.sequences[:, prompt_len:]
|
|
239
|
+
lps = None
|
|
240
|
+
if r0.logprobs:
|
|
241
|
+
logits = torch.stack(out.logits, dim=1).float() # [B, T, V]
|
|
242
|
+
lps = torch.log_softmax(logits, dim=-1).gather(-1, seqs[:, :logits.shape[1], None]).squeeze(-1)
|
|
243
|
+
for i, ((_, r, fut), row) in enumerate(zip(items, seqs)):
|
|
244
|
+
ids = row.tolist()
|
|
245
|
+
n = len(ids)
|
|
246
|
+
for j, t in enumerate(ids):
|
|
247
|
+
if t in self.eos_ids or t == self.tokenizer.pad_token_id and j > 0 and ids[j - 1] in self.eos_ids:
|
|
248
|
+
n = j
|
|
249
|
+
break
|
|
250
|
+
n = min(n, r.max_new_tokens)
|
|
251
|
+
ids = ids[:n]
|
|
252
|
+
raw = self.tokenizer.decode(ids, skip_special_tokens=False)
|
|
253
|
+
text = self.tokenizer.decode(ids, skip_special_tokens=True)
|
|
254
|
+
text = strip_thinking(text) if (r.think or self.think or "<think>" in text or "</think>" in text) else text.strip()
|
|
255
|
+
res = GenResult(text=text, raw_text=raw, num_tokens=n,
|
|
256
|
+
token_logprobs=lps[i, :n].tolist() if lps is not None else None)
|
|
257
|
+
if not fut.done():
|
|
258
|
+
fut.set_result(res)
|
|
259
|
+
|
|
260
|
+
def _run_score(self, items):
|
|
261
|
+
reqs = [r for _, r, _ in items]
|
|
262
|
+
prompts, images, targets = [], [], []
|
|
263
|
+
for r in reqs:
|
|
264
|
+
text, ims = self._render(r.parts, think=False)
|
|
265
|
+
prompts.append(text)
|
|
266
|
+
images.append(ims)
|
|
267
|
+
targets.append(self.tokenizer(r.target, add_special_tokens=False)["input_ids"])
|
|
268
|
+
# Score "prompt + target" with right-aligned targets: append target ids
|
|
269
|
+
# to each (left-padded) prompt, then read logits one position earlier.
|
|
270
|
+
enc = self._inputs(prompts, images)
|
|
271
|
+
ids, mask = enc["input_ids"], enc["attention_mask"]
|
|
272
|
+
T = max(len(t) for t in targets)
|
|
273
|
+
pad = self.tokenizer.pad_token_id
|
|
274
|
+
tgt = torch.full((len(reqs), T), pad, dtype=ids.dtype, device=ids.device)
|
|
275
|
+
tmask = torch.zeros((len(reqs), T), dtype=mask.dtype, device=mask.device)
|
|
276
|
+
for i, t in enumerate(targets):
|
|
277
|
+
if t:
|
|
278
|
+
tgt[i, :len(t)] = torch.tensor(t, device=ids.device)
|
|
279
|
+
tmask[i, :len(t)] = 1
|
|
280
|
+
for key, val in list(enc.items()): # other per-token fields (token types...) extend with zeros
|
|
281
|
+
if key not in ("input_ids", "attention_mask") and torch.is_tensor(val) and val.shape == ids.shape:
|
|
282
|
+
enc[key] = torch.cat([val, torch.zeros_like(tgt, dtype=val.dtype)], dim=1)
|
|
283
|
+
enc["input_ids"] = torch.cat([ids, tgt], dim=1)
|
|
284
|
+
enc["attention_mask"] = torch.cat([mask, tmask], dim=1)
|
|
285
|
+
P = ids.shape[1]
|
|
286
|
+
if self._logits_to_keep: # only the T+1 positions that predict target tokens, not the whole prompt
|
|
287
|
+
logits = self.model(**enc, logits_to_keep=T + 1).logits
|
|
288
|
+
start = 0
|
|
289
|
+
else:
|
|
290
|
+
logits = self.model(**enc).logits
|
|
291
|
+
start = P - 1
|
|
292
|
+
for i, ((_, r, fut), t) in enumerate(zip(items, targets)):
|
|
293
|
+
if not t:
|
|
294
|
+
fut.set_result([])
|
|
295
|
+
continue
|
|
296
|
+
lg = logits[i, start:start + len(t)].float()
|
|
297
|
+
lp = torch.log_softmax(lg, dim=-1).gather(-1, torch.tensor(t, device=lg.device)[:, None]).squeeze(-1)
|
|
298
|
+
fut.set_result(lp.tolist())
|
|
299
|
+
|
|
300
|
+
|
|
301
|
+
def _dense_fp8_fix(model_id: str, kw: dict) -> None:
|
|
302
|
+
"""Pre-quantized FP8 checkpoints of dense models can list MoE router names
|
|
303
|
+
("...mlp.gate") as modules not to convert; transformers matches those as
|
|
304
|
+
prefixes, which also skips the dense "...mlp.gate_proj" and loads its FP8
|
|
305
|
+
weights unscaled (garbage outputs). Drop them when the model has no experts."""
|
|
306
|
+
if "config" in kw:
|
|
307
|
+
return
|
|
308
|
+
try:
|
|
309
|
+
from transformers import AutoConfig
|
|
310
|
+
config = AutoConfig.from_pretrained(model_id, trust_remote_code=kw.get("trust_remote_code", False))
|
|
311
|
+
except Exception:
|
|
312
|
+
return
|
|
313
|
+
quant = getattr(config, "quantization_config", None)
|
|
314
|
+
text = getattr(config, "text_config", config)
|
|
315
|
+
if not isinstance(quant, dict) or not quant.get("modules_to_not_convert") or getattr(text, "num_experts", 0):
|
|
316
|
+
return
|
|
317
|
+
keep = [m for m in quant["modules_to_not_convert"] if not m.endswith(".mlp.gate")]
|
|
318
|
+
if len(keep) != len(quant["modules_to_not_convert"]):
|
|
319
|
+
quant["modules_to_not_convert"] = keep
|
|
320
|
+
kw["config"] = config
|
|
321
|
+
|
|
322
|
+
|
|
323
|
+
def _load_model(model_id: str, kw: dict):
|
|
324
|
+
import transformers
|
|
325
|
+
_dense_fp8_fix(model_id, kw)
|
|
326
|
+
errors = []
|
|
327
|
+
for name in ("AutoModelForMultimodalLM", "AutoModelForImageTextToText", "AutoModelForVision2Seq",
|
|
328
|
+
"AutoModelForCausalLM"):
|
|
329
|
+
cls = getattr(transformers, name, None)
|
|
330
|
+
if cls is None:
|
|
331
|
+
continue
|
|
332
|
+
try:
|
|
333
|
+
return cls.from_pretrained(model_id, **kw)
|
|
334
|
+
except (ValueError, KeyError) as e: # architecture not registered for this auto class
|
|
335
|
+
errors.append(f"{name}: {str(e).splitlines()[0]}")
|
|
336
|
+
raise ValueError(f"could not load {model_id} with transformers: " + "; ".join(errors))
|
|
@@ -0,0 +1,136 @@
|
|
|
1
|
+
"""Text-to-image tools for models that cannot generate images themselves."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import inspect
|
|
6
|
+
import queue
|
|
7
|
+
import threading
|
|
8
|
+
import traceback
|
|
9
|
+
from concurrent.futures import Future
|
|
10
|
+
from typing import Callable, Optional
|
|
11
|
+
|
|
12
|
+
import torch
|
|
13
|
+
from PIL import Image
|
|
14
|
+
|
|
15
|
+
from .base import Backend, ImageRequest
|
|
16
|
+
|
|
17
|
+
# Sampler settings for distilled checkpoints the defaults would waste steps on.
|
|
18
|
+
_PRESETS = {
|
|
19
|
+
"flux.2-klein": dict(num_inference_steps=4, guidance_scale=1.0),
|
|
20
|
+
"flux.1-schnell": dict(num_inference_steps=4, guidance_scale=0.0),
|
|
21
|
+
"sdxl-turbo": dict(num_inference_steps=2, guidance_scale=0.0),
|
|
22
|
+
"sd-turbo": dict(num_inference_steps=2, guidance_scale=0.0),
|
|
23
|
+
}
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class DiffusersImageGenerator(Backend):
|
|
27
|
+
"""Wraps a diffusers text-to-image pipeline; requests of equal size are
|
|
28
|
+
generated in one batch."""
|
|
29
|
+
|
|
30
|
+
name = "diffusers"
|
|
31
|
+
native_image_generation = True
|
|
32
|
+
|
|
33
|
+
def __init__(self, model, device: Optional[str] = None, dtype=torch.bfloat16, max_batch_size: int = 8,
|
|
34
|
+
**call_kwargs):
|
|
35
|
+
if isinstance(model, str):
|
|
36
|
+
from diffusers import DiffusionPipeline
|
|
37
|
+
from diffusers.pipelines import pipeline_utils
|
|
38
|
+
self.model_id = model
|
|
39
|
+
dtype_kw = "dtype" if hasattr(pipeline_utils, "_resolve_dtype") else "torch_dtype" # renamed in diffusers 0.40
|
|
40
|
+
self.pipe = DiffusionPipeline.from_pretrained(model, **{dtype_kw: dtype})
|
|
41
|
+
self.pipe.to(device or ("cuda" if torch.cuda.is_available() else "cpu"))
|
|
42
|
+
else:
|
|
43
|
+
self.pipe = model
|
|
44
|
+
self.model_id = getattr(getattr(model, "config", None), "_name_or_path", type(model).__name__)
|
|
45
|
+
self.pipe.set_progress_bar_config(disable=True)
|
|
46
|
+
preset = next((v for k, v in _PRESETS.items() if k in self.model_id.lower()), {})
|
|
47
|
+
self.call_kwargs = {**preset, **call_kwargs}
|
|
48
|
+
self.max_batch_size = max_batch_size
|
|
49
|
+
self._queue: "queue.Queue" = queue.Queue()
|
|
50
|
+
self._thread = threading.Thread(target=self._loop, daemon=True, name="simit-imagegen")
|
|
51
|
+
self._thread.start()
|
|
52
|
+
|
|
53
|
+
def submit_image(self, req: ImageRequest) -> Future:
|
|
54
|
+
fut: Future = Future()
|
|
55
|
+
self._queue.put((req, fut))
|
|
56
|
+
return fut
|
|
57
|
+
|
|
58
|
+
def close(self):
|
|
59
|
+
self._queue.put(None)
|
|
60
|
+
|
|
61
|
+
def _loop(self):
|
|
62
|
+
torch.set_grad_enabled(False)
|
|
63
|
+
while True:
|
|
64
|
+
item = self._queue.get()
|
|
65
|
+
if item is None:
|
|
66
|
+
return
|
|
67
|
+
items = [item]
|
|
68
|
+
try:
|
|
69
|
+
while len(items) < 64:
|
|
70
|
+
nxt = self._queue.get_nowait()
|
|
71
|
+
if nxt is None:
|
|
72
|
+
self._queue.put(None)
|
|
73
|
+
break
|
|
74
|
+
items.append(nxt)
|
|
75
|
+
except queue.Empty:
|
|
76
|
+
pass
|
|
77
|
+
by_size: dict = {}
|
|
78
|
+
for req, fut in items:
|
|
79
|
+
if not fut.cancelled():
|
|
80
|
+
by_size.setdefault((_round(req.width), _round(req.height)), []).append((req, fut))
|
|
81
|
+
for (w, h), group in by_size.items():
|
|
82
|
+
for i in range(0, len(group), self.max_batch_size):
|
|
83
|
+
self._run(group[i:i + self.max_batch_size], w, h)
|
|
84
|
+
|
|
85
|
+
def _run(self, group, w, h):
|
|
86
|
+
try:
|
|
87
|
+
gens = None
|
|
88
|
+
if any(r.seed is not None for r, _ in group):
|
|
89
|
+
dev = getattr(self.pipe, "device", "cpu")
|
|
90
|
+
gens = [torch.Generator(device=dev).manual_seed(r.seed if r.seed is not None else torch.seed() % 2 ** 31)
|
|
91
|
+
for r, _ in group]
|
|
92
|
+
images = self.pipe(prompt=[r.prompt for r, _ in group], width=w, height=h, generator=gens,
|
|
93
|
+
**self.call_kwargs).images
|
|
94
|
+
for (_, fut), im in zip(group, images):
|
|
95
|
+
if not fut.done():
|
|
96
|
+
fut.set_result(im.convert("RGB"))
|
|
97
|
+
except torch.cuda.OutOfMemoryError:
|
|
98
|
+
torch.cuda.empty_cache()
|
|
99
|
+
if len(group) > 1: # e.g. a VLM engine sharing the GPU: retry in halves
|
|
100
|
+
mid = len(group) // 2
|
|
101
|
+
self._run(group[:mid], w, h)
|
|
102
|
+
self._run(group[mid:], w, h)
|
|
103
|
+
else:
|
|
104
|
+
for _, fut in group:
|
|
105
|
+
if not fut.done():
|
|
106
|
+
fut.set_exception(RuntimeError("CUDA out of memory generating a single image"))
|
|
107
|
+
except Exception as e:
|
|
108
|
+
tb = traceback.format_exc()
|
|
109
|
+
for _, fut in group:
|
|
110
|
+
if not fut.done():
|
|
111
|
+
fut.set_exception(RuntimeError(f"{e}\n{tb}"))
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
class CallableImageGenerator(Backend):
|
|
115
|
+
"""Any ``fn(prompt, width, height) -> PIL.Image`` (e.g. an API client)."""
|
|
116
|
+
|
|
117
|
+
name = "callable"
|
|
118
|
+
native_image_generation = True
|
|
119
|
+
|
|
120
|
+
def __init__(self, fn: Callable):
|
|
121
|
+
self.fn = fn
|
|
122
|
+
self._pool = None
|
|
123
|
+
|
|
124
|
+
def submit_image(self, req: ImageRequest) -> Future:
|
|
125
|
+
from concurrent.futures import ThreadPoolExecutor
|
|
126
|
+
if self._pool is None:
|
|
127
|
+
self._pool = ThreadPoolExecutor(max_workers=4, thread_name_prefix="simit-imagefn")
|
|
128
|
+
return self._pool.submit(self.fn, req.prompt, req.width, req.height)
|
|
129
|
+
|
|
130
|
+
def close(self):
|
|
131
|
+
if self._pool is not None:
|
|
132
|
+
self._pool.shutdown(wait=False)
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
def _round(v: int, multiple: int = 16) -> int:
|
|
136
|
+
return max(multiple, int(round(v / multiple)) * multiple)
|