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 ADDED
@@ -0,0 +1,39 @@
1
+ """SIMIT-ICL: self-improving vision-language models via imagination at test time.
2
+
3
+ from simit import SIMIT
4
+ model = SIMIT.from_pretrained("ByteDance-Seed/BAGEL-7B-MoT")
5
+ demos = model.imagine(image, question)
6
+ answer = model.answer(image, question, demos)
7
+ """
8
+
9
+ __version__ = "0.1.0"
10
+
11
+ # Heavy modules (torch, transformers) load lazily so that light imports such
12
+ # as ``from simit import Skill`` -- and the render worker processes -- stay fast.
13
+ _LAZY = {
14
+ "SIMIT": ("simit.api", "SIMIT"),
15
+ "SIMITConfig": ("simit.config", "SIMITConfig"),
16
+ "Demo": ("simit.types", "Demo"),
17
+ "Skill": ("simit.skills.base", "Skill"),
18
+ "SpecError": ("simit.skills.base", "SpecError"),
19
+ "Realization": ("simit.skills.base", "Realization"),
20
+ "extract_fenced_block": ("simit.skills.base", "extract_fenced_block"),
21
+ "save_demos": ("simit.types", "save_demos"),
22
+ "load_demos": ("simit.types", "load_demos"),
23
+ "default_skills": ("simit.skills.builtin", "default_skills"),
24
+ "metrics": ("simit.metrics", None),
25
+ }
26
+
27
+ __all__ = list(_LAZY) + ["__version__"]
28
+
29
+
30
+ def __getattr__(name):
31
+ if name not in _LAZY:
32
+ raise AttributeError(f"module 'simit' has no attribute {name!r}")
33
+ import importlib
34
+
35
+ module_name, attr = _LAZY[name]
36
+ module = importlib.import_module(module_name)
37
+ value = module if attr is None else getattr(module, attr)
38
+ globals()[name] = value
39
+ return value
simit/api.py ADDED
@@ -0,0 +1,258 @@
1
+ """The user-facing ``SIMIT`` class."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import logging
6
+ from collections import OrderedDict
7
+ from typing import Callable, Iterable, Optional, Sequence, Union
8
+
9
+ from PIL import Image
10
+
11
+ from .backends import load_backend, load_image_generator, registered_backend
12
+ from .backends.base import Backend
13
+ from .config import SIMITConfig
14
+ from .pipeline.core import Pipeline, Runtime
15
+ from .skills.base import Skill
16
+ from .skills.builtin import default_skills
17
+ from .skills.pool import RenderPool
18
+ from .types import Demo, Imagination, ZeroShot
19
+ from .utils import ImageLike, image_hash, to_pil
20
+
21
+ log = logging.getLogger("simit")
22
+
23
+
24
+ class SIMIT:
25
+ """A vision-language model that imagines its own in-context demonstrations.
26
+
27
+ >>> model = SIMIT.from_pretrained("ByteDance-Seed/BAGEL-7B-MoT")
28
+ >>> demos = model.imagine(image, question)
29
+ >>> answer = model.answer(image, question, demos)
30
+ """
31
+
32
+ def __init__(self, backend: Backend, *, image_generator=None, skills: Optional[Sequence[Skill]] = None,
33
+ config: Optional[SIMITConfig] = None, render_workers: Optional[int] = None,
34
+ model_id: str = "", verbose: bool = False):
35
+ if verbose:
36
+ logging.basicConfig(level=logging.INFO, format="%(asctime)s [simit] %(message)s")
37
+ log.setLevel(logging.INFO)
38
+ self.backend = backend
39
+ self.image_generator = image_generator
40
+ self.model_id = model_id or getattr(backend, "model_id", backend.name)
41
+ self._runtime = Runtime()
42
+ self._pool = RenderPool(render_workers)
43
+ self._pipe = Pipeline(backend, skills if skills is not None else default_skills(),
44
+ config or SIMITConfig(), image_generator=image_generator, pool=self._pool,
45
+ runtime=self._runtime)
46
+ self._zs_cache: "OrderedDict[tuple, ZeroShot]" = OrderedDict()
47
+
48
+ # ------------------------------------------------------------ loading
49
+ @classmethod
50
+ def from_pretrained(cls, model: Union[str, object], *, processor=None,
51
+ image_generator: Union[str, object, None] = None, skills: Optional[Sequence[Skill]] = None,
52
+ config: Optional[SIMITConfig] = None, device: Optional[str] = None,
53
+ image_generator_device: Optional[str] = None, render_workers: Optional[int] = None,
54
+ verbose: bool = False, **backend_kwargs) -> "SIMIT":
55
+ """Load a model by Hugging Face id / local path, or wrap an already
56
+ loaded ``transformers`` model (pass its ``processor`` too).
57
+
58
+ ``image_generator``: a diffusers model id (e.g.
59
+ ``"black-forest-labs/FLUX.2-klein-4B"``), a diffusers pipeline, or a
60
+ callable ``(prompt, width, height) -> PIL.Image``. It is used for
61
+ natural images when the model cannot generate images itself; omit it
62
+ to use the structured skills only."""
63
+ gen = None
64
+ gen_device = image_generator_device or _second_device(device)
65
+ if image_generator is not None and not isinstance(model, Backend) and registered_backend(model) is None:
66
+ # Load the tool first so the VLM (vLLM / device_map="auto") sizes itself around it.
67
+ gen = load_image_generator(image_generator, device=gen_device)
68
+ backend = load_backend(model, processor=processor, device=device, **backend_kwargs)
69
+ if image_generator is not None and backend.native_image_generation:
70
+ log.info("%s generates images natively; ignoring image_generator", backend.name)
71
+ elif image_generator is not None and gen is None:
72
+ gen = load_image_generator(image_generator, device=gen_device)
73
+ model_id = model if isinstance(model, str) else getattr(getattr(model, "config", None), "_name_or_path", "")
74
+ return cls(backend, image_generator=gen, skills=skills, config=config, render_workers=render_workers,
75
+ model_id=model_id, verbose=verbose)
76
+
77
+ # --------------------------------------------------------- properties
78
+ @property
79
+ def config(self) -> SIMITConfig:
80
+ return self._pipe.config
81
+
82
+ @config.setter
83
+ def config(self, cfg: SIMITConfig):
84
+ self._pipe.config = cfg
85
+
86
+ @property
87
+ def skills(self) -> list[Skill]:
88
+ return list(self._pipe.skills)
89
+
90
+ def add_skill(self, skill: Skill, replace: bool = False) -> None:
91
+ """Register a custom skill (it joins the router's categories)."""
92
+ if not isinstance(skill, Skill):
93
+ raise TypeError("expected a simit.Skill instance")
94
+ current = [s for s in self._pipe.skills if not (replace and s.name == skill.name)]
95
+ self._pipe.set_skills(current + [skill])
96
+
97
+ def remove_skill(self, name: str) -> None:
98
+ self._pipe.set_skills([s for s in self._pipe.skills if s.name != name])
99
+
100
+ # ------------------------------------------------------------ helpers
101
+ def _run(self, coro):
102
+ return self._runtime.run(coro)
103
+
104
+ def _mnt(self, max_new_tokens):
105
+ return max_new_tokens or self.config.answer_max_new_tokens
106
+
107
+ async def _zero_shot(self, image: Image.Image, question: str, mnt: int) -> ZeroShot:
108
+ key = (image_hash(image), question, mnt)
109
+ if key in self._zs_cache:
110
+ self._zs_cache.move_to_end(key)
111
+ return self._zs_cache[key]
112
+ zs = await self._pipe.zero_shot(image, question, mnt)
113
+ self._zs_cache[key] = zs
114
+ while len(self._zs_cache) > 4096:
115
+ self._zs_cache.popitem(last=False)
116
+ return zs
117
+
118
+ async def _imagine(self, image, question, k, fill, mnt, alone: bool, time_limit=None, on_demo=None,
119
+ on_progress=None):
120
+ image = to_pil(image)
121
+ zs = await self._zero_shot(image, question, mnt)
122
+ spec = self.config.speculative
123
+ if on_progress is not None:
124
+ self._pipe.on_call = on_progress
125
+ try:
126
+ return await self._pipe.imagine(image, question, k=k, fill=fill, max_new_tokens=mnt, zero_shot=zs,
127
+ speculative=alone if spec is None else spec, time_limit=time_limit,
128
+ on_demo=on_demo)
129
+ finally:
130
+ if on_progress is not None:
131
+ self._pipe.on_call = None
132
+
133
+ async def _answer(self, image, question, demos, mnt):
134
+ image = to_pil(image)
135
+ if not demos:
136
+ return (await self._zero_shot(image, question, mnt)).answer
137
+ return await self._pipe.answer(image, question, demos, mnt)
138
+
139
+ # -------------------------------------------------------------- public
140
+ def imagine(self, image: ImageLike, question: str, k: Optional[int] = None, *,
141
+ return_details: bool = False, max_new_tokens: Optional[int] = None,
142
+ time_limit: Optional[float] = None, on_demo: Optional[Callable[[Demo], None]] = None,
143
+ on_progress: Optional[Callable[[str], None]] = None):
144
+ """Imagine demonstrations similar to the query ``(image, question)``.
145
+
146
+ By default the number of demos is chosen adaptively from the model's
147
+ zero-shot confidence (0 for queries it is already sure about) and
148
+ candidates outside the confidence band are filtered out. ``k=n``
149
+ forces exactly ``n`` verified demos. ``return_details=True`` returns an
150
+ :class:`Imagination` (zero-shot answer, budget, all candidates, stats).
151
+ A single call runs its candidate attempts in parallel (see
152
+ ``SIMITConfig.speculative``); use :meth:`imagine_batch` for many queries.
153
+
154
+ ``time_limit`` (seconds, not counting the zero-shot pass) stops imagining
155
+ and returns the demos accepted so far; ``on_demo(demo)`` is called as
156
+ each demo is accepted and ``on_progress(stage)`` as each model call
157
+ starts ("synthesis", "route", "image", "render:<skill>", "critic", ...),
158
+ e.g. to show progress in a UI."""
159
+ res: Imagination = self._run(self._imagine(image, question, k, False, self._mnt(max_new_tokens), True,
160
+ time_limit, on_demo, on_progress))
161
+ return res if return_details else res.demos
162
+
163
+ def answer(self, image: ImageLike, question: str, demos: Optional[Sequence[Demo]] = None, *,
164
+ max_new_tokens: Optional[int] = None) -> str:
165
+ """Answer the query using the demonstrations in context (no demos:
166
+ the zero-shot answer)."""
167
+ return self._run(self._answer(image, question, list(demos or []), self._mnt(max_new_tokens)))
168
+
169
+ def greedy(self, image: ImageLike, question: str, *, max_new_tokens: Optional[int] = None) -> str:
170
+ """The standard zero-shot greedy answer (cached, shared with ``imagine``)."""
171
+ return self.zero_shot(image, question, max_new_tokens=max_new_tokens).answer
172
+
173
+ def zero_shot(self, image: ImageLike, question: str, *, max_new_tokens: Optional[int] = None) -> ZeroShot:
174
+ """The greedy answer with its confidence ``p0`` (what the adaptive budget
175
+ uses: ``model.config.budget(zs.confidence)`` demos). Cached and reused by
176
+ a following ``imagine`` of the same query."""
177
+ return self._run(self._zero_shot(to_pil(image), question, self._mnt(max_new_tokens)))
178
+
179
+ def __call__(self, image: ImageLike, question: str, k: Optional[int] = None, **kw) -> str:
180
+ """Imagine, then answer with the imagined demonstrations."""
181
+ return self.answer(image, question, self.imagine(image, question, k, **kw), **kw)
182
+
183
+ def imagine_batch(self, queries: Iterable, k: Optional[int] = None, *, return_details: bool = False,
184
+ max_new_tokens: Optional[int] = None, concurrency: int = 16) -> list:
185
+ """``imagine`` for many ``(image, question)`` pairs at once; their model
186
+ calls are batched together (much faster than a loop)."""
187
+ import asyncio
188
+ queries = [(q["image"], q["question"]) if isinstance(q, dict) else tuple(q) for q in queries]
189
+ mnt = self._mnt(max_new_tokens)
190
+
191
+ async def run():
192
+ sem = asyncio.Semaphore(concurrency)
193
+
194
+ async def one(img, q):
195
+ async with sem:
196
+ return await self._imagine(img, q, k, False, mnt, False)
197
+ return await asyncio.gather(*[one(img, q) for img, q in queries])
198
+
199
+ res = self._run(run())
200
+ return res if return_details else [r.demos for r in res]
201
+
202
+ def answer_batch(self, queries: Iterable, demos: Optional[Sequence[Sequence[Demo]]] = None, *,
203
+ max_new_tokens: Optional[int] = None) -> list[str]:
204
+ import asyncio
205
+ queries = [(q["image"], q["question"]) if isinstance(q, dict) else tuple(q) for q in queries]
206
+ demos = demos if demos is not None else [[] for _ in queries]
207
+ mnt = self._mnt(max_new_tokens)
208
+
209
+ async def run():
210
+ return await asyncio.gather(*[self._answer(img, q, list(d), mnt) for (img, q), d in zip(queries, demos)])
211
+ return self._run(run())
212
+
213
+ def tune(self, val_set, metric: Union[str, Callable] = "exact_match", n_trials: int = 100, *,
214
+ k_max_choices: Optional[Sequence[int]] = None, cache_dir=None, seed: int = 0,
215
+ concurrency: int = 16, max_new_tokens: Optional[int] = None, apply: bool = True,
216
+ max_candidates: Optional[float] = None, max_runtime_ratio: Optional[float] = None):
217
+ """Fit ABA + DF hyperparameters on labeled validation data.
218
+
219
+ ``val_set``: iterable of ``(image, question, answer_or_answers)`` or
220
+ dicts with ``image``/``question``/``answer(s)`` (and optionally their own
221
+ ``metric``, for validation sets that mix benchmarks). ``metric``: a name from
222
+ :mod:`simit.metrics` or ``fn(prediction, references) -> float``.
223
+ Synthesis runs once per query (cached, optionally on disk under
224
+ ``cache_dir``); each Optuna trial only re-selects and re-answers.
225
+
226
+ ``max_candidates`` caps the mean number of candidates synthesized per
227
+ query (the dominant runtime cost), e.g. ``1.0`` for roughly a quarter
228
+ of K_max=4 as in the paper. ``max_runtime_ratio`` (e.g. ``25``) caps the
229
+ estimated batched runtime relative to zero-shot instead, from timings
230
+ measured on this machine while tuning. ``None`` for both optimizes
231
+ accuracy alone."""
232
+ from .tune import run_tune
233
+ return run_tune(self, val_set, metric=metric, n_trials=n_trials, k_max_choices=k_max_choices,
234
+ cache_dir=cache_dir, seed=seed, concurrency=concurrency, max_new_tokens=max_new_tokens,
235
+ apply=apply, max_candidates=max_candidates, max_runtime_ratio=max_runtime_ratio)
236
+
237
+ def close(self):
238
+ self._pool.close()
239
+ self.backend.close()
240
+ if self.image_generator is not None:
241
+ self.image_generator.close()
242
+ self._runtime.close()
243
+
244
+ def __repr__(self):
245
+ return (f"SIMIT(model={self.model_id!r}, backend={self.backend.name}, skills={len(self.skills)}, "
246
+ f"image_generator={'native' if self.backend.native_image_generation else type(self.image_generator).__name__ if self.image_generator else None})")
247
+
248
+
249
+ def _second_device(device: Optional[str]) -> Optional[str]:
250
+ """Put a separate image generator on another GPU when there is one."""
251
+ try:
252
+ import torch
253
+ n = torch.cuda.device_count()
254
+ except Exception:
255
+ return None
256
+ if n >= 2:
257
+ return f"cuda:{n - 1}"
258
+ return device
@@ -0,0 +1,112 @@
1
+ """Model backends. ``load_backend`` picks one from a model id/path/object:
2
+
3
+ * BAGEL checkpoints -> :class:`~simit.backends.bagel.BagelBackend` (native image generation)
4
+ * Lance checkpoints -> :class:`~simit.backends.lance.LanceBackend` (native image generation)
5
+ * anything else -> :class:`~simit.backends.hf.HFBackend` (any transformers VLM)
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import json
11
+ from pathlib import Path
12
+ from typing import Optional
13
+
14
+ from .base import Backend, GenRequest, GenResult, ImageRequest, ScoreRequest
15
+
16
+ __all__ = ["Backend", "GenRequest", "GenResult", "ScoreRequest", "ImageRequest", "load_backend",
17
+ "load_image_generator", "register_backend", "registered_backend"]
18
+
19
+ # name -> (detector(model_id, config) -> bool, factory(model, **kw) -> Backend)
20
+ _REGISTRY: dict = {}
21
+
22
+
23
+ def register_backend(name: str, detect, factory):
24
+ """Add a backend for a new model family (checked before the generic HF backend)."""
25
+ _REGISTRY[name] = (detect, factory)
26
+
27
+
28
+ def _read_config(model_id: str) -> dict:
29
+ path = Path(model_id)
30
+ try:
31
+ if path.exists():
32
+ f = path / "config.json"
33
+ return json.loads(f.read_text()) if f.exists() else {}
34
+ from huggingface_hub import hf_hub_download
35
+ return json.loads(Path(hf_hub_download(model_id, "config.json")).read_text())
36
+ except Exception:
37
+ return {}
38
+
39
+
40
+ def _is_bagel(model_id: str, cfg: dict) -> bool:
41
+ return cfg.get("model_type") == "bagel" or "BagelForConditionalGeneration" in cfg.get("architectures", [])
42
+
43
+
44
+ def _is_lance(model_id: str, cfg: dict) -> bool:
45
+ return cfg.get("model_name") == "Lance" or Path(model_id).name.lower() == "lance"
46
+
47
+
48
+ def _bagel(model, **kw):
49
+ from .bagel import BagelBackend
50
+ return BagelBackend(model, **kw)
51
+
52
+
53
+ def _lance(model, **kw):
54
+ from .lance import LanceBackend
55
+ return LanceBackend(model, **kw)
56
+
57
+
58
+ register_backend("bagel", _is_bagel, _bagel)
59
+ register_backend("lance", _is_lance, _lance)
60
+
61
+
62
+ def registered_backend(model) -> Optional[str]:
63
+ """Name of the registered model-family backend (e.g. ``"bagel"``) for a
64
+ model id/path, or None for generic transformers/vLLM models."""
65
+ if not isinstance(model, str):
66
+ return None
67
+ cfg = _read_config(model)
68
+ return next((name for name, (detect, _) in _REGISTRY.items() if detect(model, cfg)), None)
69
+
70
+
71
+ def _vllm_available() -> bool:
72
+ import importlib.util
73
+ return importlib.util.find_spec("vllm") is not None
74
+
75
+
76
+ def load_backend(model, processor=None, device: Optional[str] = None, engine: str = "auto", **kw) -> Backend:
77
+ """``engine`` (standard VLMs only): ``"transformers"``, ``"vllm"``, or
78
+ ``"auto"`` (vLLM when installed, otherwise transformers)."""
79
+ if isinstance(model, Backend):
80
+ return model
81
+ if not isinstance(model, str): # an already-loaded transformers model
82
+ from .hf import HFBackend
83
+ return HFBackend(model, processor=processor, **kw)
84
+ cfg = _read_config(model)
85
+ for name, (detect, factory) in _REGISTRY.items():
86
+ if detect(model, cfg):
87
+ if device is not None:
88
+ kw["device"] = device
89
+ return factory(model, **kw)
90
+ if engine not in ("auto", "transformers", "vllm"):
91
+ raise ValueError(f"engine must be 'auto', 'transformers' or 'vllm', got {engine!r}")
92
+ if engine == "vllm" or (engine == "auto" and _vllm_available() and processor is None):
93
+ from .vllm import VLLMBackend
94
+ return VLLMBackend(model, device=device, **kw)
95
+ from .hf import HFBackend
96
+ if device is not None:
97
+ kw.setdefault("device_map", device)
98
+ return HFBackend(model, processor=processor, **kw)
99
+
100
+
101
+ def load_image_generator(gen, device: Optional[str] = None):
102
+ """A text-to-image tool for models without native generation."""
103
+ if gen is None or isinstance(gen, Backend):
104
+ return gen
105
+ from .imagegen import CallableImageGenerator, DiffusersImageGenerator
106
+ if isinstance(gen, str):
107
+ return DiffusersImageGenerator(gen, device=device)
108
+ if hasattr(gen, "__call__") and hasattr(gen, "components"): # a diffusers pipeline object
109
+ return DiffusersImageGenerator(gen, device=device)
110
+ if callable(gen):
111
+ return CallableImageGenerator(gen)
112
+ raise TypeError(f"unsupported image_generator: {type(gen).__name__}")