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/__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__}")
|