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