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/vllm.py ADDED
@@ -0,0 +1,204 @@
1
+ """Fast engine for standard VLMs: vLLM (continuous batching, prefix caching,
2
+ CUDA graphs). Optional: ``pip install "simit[vllm]"``.
3
+
4
+ Prompts are built with the model's own chat template (as in the transformers
5
+ backend), so results match it up to numerics.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import itertools
11
+ import queue
12
+ import threading
13
+ import traceback
14
+ from concurrent.futures import Future
15
+ from typing import Optional
16
+
17
+ from PIL import Image
18
+
19
+ from ..utils import strip_thinking, to_pil
20
+ from .base import Backend, GenRequest, GenResult, ScoreRequest
21
+
22
+
23
+ class VLLMBackend(Backend):
24
+ name = "vllm"
25
+ native_image_generation = False
26
+ verify_think = False
27
+
28
+ def __init__(self, model: str, *, gpu_memory_utilization: Optional[float] = None, max_model_len: int = 32768,
29
+ max_images_per_prompt: int = 8, max_image_pixels: Optional[int] = 1024 * 1024,
30
+ tensor_parallel_size: int = 1, think: bool = False, device: Optional[str] = None,
31
+ max_num_seqs: int = 64, **llm_kwargs):
32
+ """``gpu_memory_utilization=None`` takes what is free on the GPU (minus
33
+ headroom), so an image generator loaded first keeps its share."""
34
+ import os
35
+ if device is not None and str(device) not in ("cuda", "cuda:0"):
36
+ raise ValueError("the vLLM engine runs on the first visible GPU; select GPUs with "
37
+ "CUDA_VISIBLE_DEVICES (and tensor_parallel_size for several)")
38
+ if gpu_memory_utilization is None:
39
+ import torch
40
+ free, total = torch.cuda.mem_get_info(0)
41
+ gpu_memory_utilization = max(0.3, min(0.9, (free - 6 * 2**30) / total))
42
+ # Keep vLLM's engine core in this process: a spawned core would re-import
43
+ # the user's __main__ (scripts without an `if __name__ == "__main__"` guard).
44
+ os.environ.setdefault("VLLM_ENABLE_V1_MULTIPROCESSING", "0")
45
+ # Portable defaults: no kernels that JIT-compile with the system nvcc
46
+ # (FlashInfer) or need a matching flash-attn build. Override via kwargs/env.
47
+ os.environ.setdefault("VLLM_USE_FLASHINFER_SAMPLER", "0")
48
+ os.environ.setdefault("VLLM_USE_DEEP_GEMM", "0") # FP8 checkpoints: DeepGEMM JIT needs nvcc >= 12.9
49
+ llm_kwargs.setdefault("attention_backend", "TRITON_ATTN")
50
+ llm_kwargs.setdefault("mm_encoder_attn_backend", "TORCH_SDPA")
51
+ llm_kwargs.setdefault("gdn_prefill_backend", "triton")
52
+ from transformers import AutoProcessor
53
+ from vllm import LLM
54
+
55
+ self.model_id = model
56
+ self.processor = AutoProcessor.from_pretrained(model)
57
+ self.tokenizer = getattr(self.processor, "tokenizer", self.processor)
58
+ self.llm = LLM(model=model, dtype="bfloat16", gpu_memory_utilization=gpu_memory_utilization,
59
+ max_model_len=max_model_len, limit_mm_per_prompt={"image": max_images_per_prompt},
60
+ enable_prefix_caching=True, tensor_parallel_size=tensor_parallel_size,
61
+ max_num_seqs=max_num_seqs, **llm_kwargs)
62
+ self.engine = self.llm.llm_engine
63
+ gc = self.llm.llm_engine.model_config.try_get_generation_config() or {}
64
+ eos = gc.get("eos_token_id")
65
+ eos = eos if isinstance(eos, list) else [eos] if eos is not None else []
66
+ self.eos_ids = {e for e in eos + [self.tokenizer.eos_token_id] if e is not None}
67
+ self.think = think
68
+ self.max_image_pixels = max_image_pixels
69
+ self._ids = itertools.count()
70
+ self._queue: "queue.Queue" = queue.Queue()
71
+ self._live: dict = {}
72
+ self._thread = threading.Thread(target=self._loop, daemon=True, name="simit-vllm")
73
+ self._thread.start()
74
+
75
+ # ------------------------------------------------------------- prompts
76
+ def _image(self, img):
77
+ img = to_pil(img)
78
+ if self.max_image_pixels and img.width * img.height > self.max_image_pixels:
79
+ s = (self.max_image_pixels / (img.width * img.height)) ** 0.5
80
+ img = img.resize((max(28, int(img.width * s)), max(28, int(img.height * s))), Image.BICUBIC)
81
+ return img
82
+
83
+ def _prompt(self, parts, think: bool):
84
+ content, images = [], []
85
+ for p in parts:
86
+ if isinstance(p, Image.Image):
87
+ content.append({"type": "image"})
88
+ images.append(self._image(p))
89
+ elif content and content[-1]["type"] == "text":
90
+ content[-1]["text"] += p
91
+ else:
92
+ content.append({"type": "text", "text": p})
93
+ messages = [{"role": "user", "content": content}]
94
+ try:
95
+ text = self.processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True,
96
+ enable_thinking=think)
97
+ except TypeError:
98
+ text = self.processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
99
+ return text, images
100
+
101
+ @staticmethod
102
+ def _inputs(text, images):
103
+ inp = {"prompt": text}
104
+ if images:
105
+ inp["multi_modal_data"] = {"image": images}
106
+ return inp
107
+
108
+ # ------------------------------------------------------------ requests
109
+ def submit_generate(self, req: GenRequest) -> Future:
110
+ from vllm import SamplingParams
111
+ think = req.think or self.think
112
+ text, images = self._prompt(req.parts, think)
113
+ sp = SamplingParams(max_tokens=req.max_new_tokens, temperature=req.temperature,
114
+ top_p=req.top_p if req.temperature > 0 else 1.0,
115
+ repetition_penalty=req.repetition_penalty or 1.0,
116
+ logprobs=0 if req.logprobs else None)
117
+
118
+ def post(out):
119
+ o = out.outputs[0]
120
+ raw = o.text
121
+ ids = list(o.token_ids)
122
+ if ids and ids[-1] in self.eos_ids: # like the other backends: the stop token is not part of the answer
123
+ ids = ids[:-1]
124
+ lps = None
125
+ if req.logprobs and o.logprobs is not None:
126
+ lps = [step[tid].logprob for tid, step in zip(ids, o.logprobs)]
127
+ text_ = strip_thinking(raw) if (think or "</think>" in raw) else raw.strip()
128
+ return GenResult(text=text_, raw_text=raw, token_logprobs=lps, num_tokens=len(ids))
129
+
130
+ return self._submit(self._inputs(text, images), sp, post)
131
+
132
+ def submit_score(self, req: ScoreRequest) -> Future:
133
+ from vllm import SamplingParams
134
+ text, images = self._prompt(req.parts, think=False)
135
+ n = len(self.tokenizer(req.target, add_special_tokens=False)["input_ids"])
136
+ if n == 0:
137
+ f: Future = Future()
138
+ f.set_result([])
139
+ return f
140
+ sp = SamplingParams(max_tokens=1, temperature=0.0, prompt_logprobs=0)
141
+
142
+ def post(out):
143
+ ids = out.prompt_token_ids[-n:]
144
+ plp = out.prompt_logprobs[-n:]
145
+ return [step[tid].logprob for tid, step in zip(ids, plp)]
146
+
147
+ return self._submit(self._inputs(text + req.target, images), sp, post)
148
+
149
+ def _submit(self, inputs, params, post) -> Future:
150
+ fut: Future = Future()
151
+ self._queue.put((inputs, params, post, fut))
152
+ return fut
153
+
154
+ def close(self):
155
+ self._queue.put(None)
156
+
157
+ # --------------------------------------------------------------- engine
158
+ def _loop(self):
159
+ while True:
160
+ block = not self._live
161
+ try:
162
+ item = self._queue.get(timeout=0.5) if block else self._queue.get_nowait()
163
+ while True:
164
+ if item is None:
165
+ return
166
+ inputs, params, post, fut = item
167
+ if not fut.cancelled():
168
+ rid = str(next(self._ids))
169
+ try:
170
+ self.engine.add_request(rid, inputs, params)
171
+ self._live[rid] = (post, fut)
172
+ except Exception as e:
173
+ fut.set_exception(e)
174
+ item = self._queue.get_nowait()
175
+ except queue.Empty:
176
+ pass
177
+ if not self._live:
178
+ continue
179
+ for rid, (post, fut) in list(self._live.items()):
180
+ if fut.cancelled():
181
+ try:
182
+ self.engine.abort_request([rid])
183
+ except Exception:
184
+ pass
185
+ self._live.pop(rid, None)
186
+ try:
187
+ outputs = self.engine.step()
188
+ except Exception as e:
189
+ tb = traceback.format_exc()
190
+ for post, fut in self._live.values():
191
+ if not fut.done():
192
+ fut.set_exception(RuntimeError(f"{e}\n{tb}"))
193
+ self._live.clear()
194
+ continue
195
+ for out in outputs:
196
+ if not out.finished or out.request_id not in self._live:
197
+ continue
198
+ post, fut = self._live.pop(out.request_id)
199
+ if fut.done():
200
+ continue
201
+ try:
202
+ fut.set_result(post(out))
203
+ except Exception as e:
204
+ fut.set_exception(e)
simit/config.py ADDED
@@ -0,0 +1,112 @@
1
+ """SIMIT hyperparameters. Defaults follow the paper; ``SIMIT.tune`` fits the
2
+ adaptive-budget (ABA) and difficulty-filter (DF) parameters on labeled
3
+ validation data."""
4
+
5
+ from __future__ import annotations
6
+
7
+ import json
8
+ from dataclasses import asdict, dataclass, fields
9
+ from pathlib import Path
10
+ from typing import Optional
11
+
12
+
13
+ @dataclass
14
+ class SIMITConfig:
15
+ # --- budget -----------------------------------------------------------
16
+ #: Maximum number of imagined demonstrations per query (K_max).
17
+ k_max: int = 4
18
+ #: Adaptive budget allocation: fewer demos for confidently answered queries.
19
+ use_aba: bool = True
20
+ #: ABA parameters (target error rate and log-accuracy curve A(K) = A1 + B ln K).
21
+ epsilon: float = 0.14
22
+ A0: float = 0.35
23
+ A1: float = 0.60
24
+ B: float = 0.2
25
+ #: Difficulty filtering: keep demos whose answer confidence lies in [t_low, t_high].
26
+ use_df: bool = True
27
+ t_low: float = 0.2
28
+ t_high: float = 0.9
29
+
30
+ # --- synthesis ----------------------------------------------------------
31
+ # ``None`` means "the backend's default" (e.g. Lance: decomposed synthesis,
32
+ # natural images only, no critic -- see ``Backend.pipeline_defaults``).
33
+ #: Triplet synthesis: "batch" (one call for K triplets, the paper's prompt)
34
+ #: or "decomposed" (question, answer, description in three short calls).
35
+ synthesis: Optional[str] = None
36
+ #: Diversity-encouraging triplet prompt (the paper's "DP").
37
+ diversity_prompt: bool = True
38
+ #: Use the structured skill library (False: realize everything as a natural image).
39
+ use_skills: Optional[bool] = None
40
+ #: Critic verification of every realized image.
41
+ verify: Optional[bool] = None
42
+ #: Minimum critic score (0-100) for an image to be accepted.
43
+ verify_threshold: int = 50
44
+ #: Run the critic in thinking mode (None: backend default; BAGEL: True).
45
+ verify_think: Optional[bool] = None
46
+ verify_max_new_tokens: int = 300
47
+ #: Realization attempts per demonstration slot (breadth-first slot filling).
48
+ attempts_per_slot: int = 4
49
+ #: Repair/regeneration rounds a failed candidate gets in the retry phase.
50
+ repair_retries: int = 2
51
+ verify_rounds: int = 5
52
+ #: Side length of natively generated images.
53
+ image_size: int = 400
54
+ #: Start every remaining candidate attempt at once (up to K_max) instead of only as many as the
55
+ #: budget still needs. Lower latency for a lone query on an otherwise idle GPU, wasted work under
56
+ #: load. None: on for ``imagine``, off for ``imagine_batch`` and ``tune``.
57
+ speculative: Optional[bool] = None
58
+
59
+ # --- answering ------------------------------------------------------------
60
+ answer_max_new_tokens: int = 128
61
+ #: Random seed for image generation (None: random).
62
+ seed: Optional[int] = None
63
+
64
+ def thresholds(self) -> dict[int, float]:
65
+ """ABA confidence thresholds tau_K (Appendix D, Algorithm 1)."""
66
+ return compute_thresholds(self.epsilon, self.A0, self.A1, self.B, self.k_max)
67
+
68
+ def budget(self, p0: Optional[float]) -> int:
69
+ """Number of demonstrations K*(q) for zero-shot confidence ``p0``."""
70
+ if not self.use_aba or p0 is None:
71
+ return self.k_max
72
+ tau = self.thresholds()
73
+ for k in range(self.k_max + 1):
74
+ if p0 >= tau[k]:
75
+ return k
76
+ return self.k_max
77
+
78
+ def in_band(self, c: Optional[float]) -> bool:
79
+ return (not self.use_df) or (c is not None and self.t_low <= c <= self.t_high)
80
+
81
+ def to_dict(self) -> dict:
82
+ return asdict(self)
83
+
84
+ def save(self, path) -> None:
85
+ Path(path).write_text(json.dumps(self.to_dict(), indent=2))
86
+
87
+ @classmethod
88
+ def load(cls, path) -> "SIMITConfig":
89
+ data = json.loads(Path(path).read_text())
90
+ names = {f.name for f in fields(cls)}
91
+ return cls(**{k: v for k, v in data.items() if k in names})
92
+
93
+ def replace(self, **kw) -> "SIMITConfig":
94
+ data = self.to_dict()
95
+ data.update(kw)
96
+ return SIMITConfig(**data)
97
+
98
+
99
+ def compute_thresholds(epsilon: float, A0: float, A1: float, B: float, k_max: int, eta: float = 1e-3) -> dict[int, float]:
100
+ import math
101
+
102
+ def clip(a):
103
+ return min(max(a, eta), 1 - eta)
104
+
105
+ a0 = clip(A0)
106
+ o0 = a0 / (1 - a0)
107
+ tau = {0: 1 - epsilon}
108
+ for k in range(1, k_max + 1):
109
+ ak = clip(A1 + B * math.log(k))
110
+ ok = ak / (1 - ak)
111
+ tau[k] = (1 - epsilon) * o0 / (epsilon * ok + (1 - epsilon) * o0)
112
+ return tau
simit/metrics.py ADDED
@@ -0,0 +1,113 @@
1
+ """Answer-scoring metrics for ``SIMIT.tune`` (higher is better, per sample in [0, 1]).
2
+
3
+ Every metric is ``metric(prediction: str, references: list[str]) -> float``;
4
+ pass any such callable to ``tune(metric=...)``.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import re
10
+ from typing import Callable, Sequence, Union
11
+
12
+ _ARTICLES = {"a", "an", "the"}
13
+ _PUNCT = re.compile(r"[^\w\s]")
14
+ _NUMBER_WORDS = {"none": "0", "zero": "0", "one": "1", "two": "2", "three": "3", "four": "4", "five": "5",
15
+ "six": "6", "seven": "7", "eight": "8", "nine": "9", "ten": "10"}
16
+
17
+
18
+ def normalize(text: str) -> str:
19
+ """Lowercase, drop punctuation and articles, map number words to digits."""
20
+ text = _PUNCT.sub(" ", str(text).lower())
21
+ words = [_NUMBER_WORDS.get(w, w) for w in text.split() if w not in _ARTICLES]
22
+ return " ".join(words)
23
+
24
+
25
+ def _refs(references) -> list[str]:
26
+ if isinstance(references, str):
27
+ return [references]
28
+ return [str(r) for r in references]
29
+
30
+
31
+ def exact_match(prediction: str, references) -> float:
32
+ p = normalize(prediction)
33
+ return float(any(p == normalize(r) for r in _refs(references)))
34
+
35
+
36
+ def vqa_accuracy(prediction: str, references) -> float:
37
+ """VQAv2-style soft accuracy: min(1, #matching annotators / 3)."""
38
+ p = normalize(prediction)
39
+ refs = _refs(references)
40
+ hits = sum(p == normalize(r) for r in refs)
41
+ return min(1.0, hits / 3.0) if len(refs) > 1 else float(hits > 0)
42
+
43
+
44
+ def contains(prediction: str, references) -> float:
45
+ p = normalize(prediction)
46
+ return float(any(normalize(r) and normalize(r) in p for r in _refs(references)))
47
+
48
+
49
+ def _levenshtein(a: str, b: str) -> int:
50
+ if len(a) < len(b):
51
+ a, b = b, a
52
+ prev = list(range(len(b) + 1))
53
+ for i, ca in enumerate(a, 1):
54
+ cur = [i]
55
+ for j, cb in enumerate(b, 1):
56
+ cur.append(min(prev[j] + 1, cur[j - 1] + 1, prev[j - 1] + (ca != cb)))
57
+ prev = cur
58
+ return prev[-1]
59
+
60
+
61
+ def anls(prediction: str, references, threshold: float = 0.5) -> float:
62
+ """Average normalized Levenshtein similarity (DocVQA / InfographicVQA)."""
63
+ p = " ".join(str(prediction).lower().split())
64
+ best = 0.0
65
+ for r in _refs(references):
66
+ r = " ".join(r.lower().split())
67
+ d = _levenshtein(p, r) / max(len(p), len(r), 1)
68
+ best = max(best, 1 - d if d < threshold else 0.0)
69
+ return best
70
+
71
+
72
+ def relaxed_accuracy(prediction: str, references, tolerance: float = 0.05) -> float:
73
+ """ChartQA: numbers within 5% relative error, otherwise exact match."""
74
+ def num(s):
75
+ s = str(s).strip().rstrip("%").replace(",", "")
76
+ try:
77
+ return float(s)
78
+ except ValueError:
79
+ return None
80
+
81
+ pv = num(prediction)
82
+ for r in _refs(references):
83
+ rv = num(r)
84
+ if pv is not None and rv is not None:
85
+ if (rv == 0 and pv == 0) or (rv != 0 and abs(pv - rv) / abs(rv) <= tolerance):
86
+ return 1.0
87
+ elif normalize(prediction) == normalize(r):
88
+ return 1.0
89
+ return 0.0
90
+
91
+
92
+ _LETTER = re.compile(r"\b([A-J])\b")
93
+
94
+
95
+ def multiple_choice(prediction: str, references) -> float:
96
+ """Compare the first option letter in the prediction with the reference letter."""
97
+ m = _LETTER.search(str(prediction).strip().upper()) or re.match(r"\s*([A-J])", str(prediction).upper())
98
+ letter = m.group(1) if m else str(prediction).strip()[:1].upper()
99
+ return float(any(letter == str(r).strip().upper()[:1] for r in _refs(references)))
100
+
101
+
102
+ METRICS: dict[str, Callable] = {
103
+ "exact_match": exact_match, "vqa_accuracy": vqa_accuracy, "contains": contains, "anls": anls,
104
+ "relaxed_accuracy": relaxed_accuracy, "multiple_choice": multiple_choice,
105
+ }
106
+
107
+
108
+ def get_metric(metric: Union[str, Callable]) -> Callable[[str, Sequence[str]], float]:
109
+ if callable(metric):
110
+ return metric
111
+ if metric not in METRICS:
112
+ raise ValueError(f"unknown metric {metric!r}; choose from {sorted(METRICS)} or pass a callable")
113
+ return METRICS[metric]