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/bagel.py
ADDED
|
@@ -0,0 +1,359 @@
|
|
|
1
|
+
"""BAGEL-7B-MoT (ByteDance-Seed/BAGEL-7B-MoT) on the batched MoT engine.
|
|
2
|
+
|
|
3
|
+
Prompt formats reproduce the SIMIT research setup: ``chat`` requests use the
|
|
4
|
+
serving template the synthesis pipeline was developed with (system prompt,
|
|
5
|
+
user/assistant turns, repetition penalty 1.05); ``answer`` requests use
|
|
6
|
+
BAGEL's interleaved inference format (each text part its own turn), which
|
|
7
|
+
the evaluation and confidence scores were computed with.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
import json
|
|
13
|
+
import math
|
|
14
|
+
import os
|
|
15
|
+
from pathlib import Path
|
|
16
|
+
from typing import Optional, Union
|
|
17
|
+
|
|
18
|
+
import torch
|
|
19
|
+
from PIL import Image
|
|
20
|
+
|
|
21
|
+
from ..utils import image_hash, strip_thinking, to_pil
|
|
22
|
+
from .base import Backend, GenRequest, GenResult, ImageRequest, ScoreRequest, chained
|
|
23
|
+
from .mot.engine import ImageChunk, ImageJob, MoTEngine, TextChunk, TextJob, auto_kv_tokens
|
|
24
|
+
from .mot.modeling import (DeviceMixin, FluxAutoEncoder, MLPConnector, MoTLanguageModel, PositionTable, SiglipNaViT,
|
|
25
|
+
TimestepEmbedder)
|
|
26
|
+
|
|
27
|
+
DEFAULT_SYSTEM_PROMPT = "You are BAGEL, a helpful assistant created by ByteDance."
|
|
28
|
+
CHAT_THINK_PROMPT = ("You should first think about the reasoning process in the mind "
|
|
29
|
+
"and then provide the user with the answer.")
|
|
30
|
+
ANSWER_THINK_PROMPT = ("You should first think about the reasoning process in the mind and then provide the "
|
|
31
|
+
"user with the answer. \nThe reasoning process is enclosed within <think> </think> tags, "
|
|
32
|
+
"i.e. <think> reasoning process here </think> answer here")
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
# -------------------------------------------------------------- image prep
|
|
36
|
+
def _resize(img: Image.Image, max_size: int, min_size: int, stride: int, max_pixels: int) -> Image.Image:
|
|
37
|
+
"""BAGEL's MaxLongEdgeMinShortEdgeResize (bicubic, antialiased)."""
|
|
38
|
+
def divisible(v):
|
|
39
|
+
return max(stride, int(round(v / stride) * stride))
|
|
40
|
+
|
|
41
|
+
def scaled(w, h, s):
|
|
42
|
+
return divisible(round(w * s)), divisible(round(h * s))
|
|
43
|
+
|
|
44
|
+
w, h = img.size
|
|
45
|
+
s = min(max_size / max(w, h), 1.0)
|
|
46
|
+
s = max(s, min_size / min(w, h))
|
|
47
|
+
nw, nh = scaled(w, h, s)
|
|
48
|
+
if nw * nh > max_pixels:
|
|
49
|
+
nw, nh = scaled(nw, nh, max_pixels / (nw * nh))
|
|
50
|
+
if max(nw, nh) > max_size:
|
|
51
|
+
nw, nh = scaled(nw, nh, max_size / max(nw, nh))
|
|
52
|
+
if (nw, nh) == (w, h):
|
|
53
|
+
return img
|
|
54
|
+
from torchvision.transforms import InterpolationMode
|
|
55
|
+
from torchvision.transforms import functional as TF
|
|
56
|
+
return TF.resize(img, [nh, nw], InterpolationMode.BICUBIC, antialias=True)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def _patchify(img: Image.Image, patch: int, max_side: int):
|
|
60
|
+
from torchvision.transforms import functional as TF
|
|
61
|
+
t = TF.to_tensor(img)
|
|
62
|
+
t = (t - 0.5) / 0.5
|
|
63
|
+
c, h, w = t.shape
|
|
64
|
+
hp, wp = h // patch, w // patch
|
|
65
|
+
patches = t.reshape(c, hp, patch, wp, patch).permute(1, 3, 2, 4, 0).reshape(hp * wp, patch * patch * c)
|
|
66
|
+
pos = (torch.arange(hp)[:, None] * max_side + torch.arange(wp)).flatten()
|
|
67
|
+
return patches, pos
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
class BagelModel(DeviceMixin):
|
|
71
|
+
"""Weights + modality-specific pieces the MoT engine calls into."""
|
|
72
|
+
|
|
73
|
+
latent_patch = 2
|
|
74
|
+
latent_channels = 16
|
|
75
|
+
vae_downsample = 8
|
|
76
|
+
num_timesteps = 50
|
|
77
|
+
timestep_shift = 3.0
|
|
78
|
+
cfg_text_scale = 4.0
|
|
79
|
+
cfg_interval = (0.4, 1.0)
|
|
80
|
+
cfg_renorm_min = 0.0
|
|
81
|
+
|
|
82
|
+
@classmethod
|
|
83
|
+
def load(cls, model_id: str = "ByteDance-Seed/BAGEL-7B-MoT", device: str = "cuda",
|
|
84
|
+
dtype=torch.bfloat16) -> "BagelModel":
|
|
85
|
+
"""From a Hub id (downloaded on first use) or a local path."""
|
|
86
|
+
path = Path(model_id)
|
|
87
|
+
if not path.exists():
|
|
88
|
+
from huggingface_hub import snapshot_download
|
|
89
|
+
path = Path(snapshot_download(model_id, allow_patterns=[
|
|
90
|
+
"*.json", "ema.safetensors", "ae.safetensors", "*.txt", "tokenizer*", "vocab*", "merges*"]))
|
|
91
|
+
model = cls(str(path), device=device, dtype=dtype)
|
|
92
|
+
model.model_id = model_id
|
|
93
|
+
return model
|
|
94
|
+
|
|
95
|
+
def __init__(self, path: str, device: str = "cuda", dtype=torch.bfloat16):
|
|
96
|
+
self.model_id = str(path)
|
|
97
|
+
self.device = torch.device(device)
|
|
98
|
+
self.dtype = dtype
|
|
99
|
+
path = Path(path)
|
|
100
|
+
cfg = json.loads((path / "config.json").read_text())
|
|
101
|
+
llm_cfg = json.loads((path / "llm_config.json").read_text())
|
|
102
|
+
vit_cfg = json.loads((path / "vit_config.json").read_text())
|
|
103
|
+
self.vit_patch = vit_cfg["patch_size"]
|
|
104
|
+
self.vit_max_side = cfg.get("vit_max_num_patch_per_side", 70)
|
|
105
|
+
self.latent_downsample = self.vae_downsample * self.latent_patch
|
|
106
|
+
self.latent_dim = self.latent_patch ** 2 * self.latent_channels
|
|
107
|
+
self.max_latent_side = 64
|
|
108
|
+
|
|
109
|
+
hidden = llm_cfg["hidden_size"]
|
|
110
|
+
with torch.device("meta"):
|
|
111
|
+
self.lm = MoTLanguageModel(llm_cfg)
|
|
112
|
+
self.vit = SiglipNaViT(vit_cfg, vit_cfg["num_hidden_layers"] - 1)
|
|
113
|
+
self.connector = MLPConnector(vit_cfg["hidden_size"], hidden)
|
|
114
|
+
self.vit_pos_embed = PositionTable(self.vit_max_side, hidden)
|
|
115
|
+
self.latent_pos_embed = PositionTable(self.max_latent_side, hidden)
|
|
116
|
+
self.time_embedder = TimestepEmbedder(hidden)
|
|
117
|
+
self.vae2llm = torch.nn.Linear(self.latent_dim, hidden)
|
|
118
|
+
self.llm2vae = torch.nn.Linear(hidden, self.latent_dim)
|
|
119
|
+
modules = {
|
|
120
|
+
"language_model": self.lm, "vit_model": self.vit, "connector": self.connector,
|
|
121
|
+
"vit_pos_embed": self.vit_pos_embed, "latent_pos_embed": self.latent_pos_embed,
|
|
122
|
+
"time_embedder": self.time_embedder, "vae2llm": self.vae2llm, "llm2vae": self.llm2vae,
|
|
123
|
+
}
|
|
124
|
+
self._load(path / "ema.safetensors", modules)
|
|
125
|
+
self.lm.rotary = type(self.lm.rotary)(hidden // llm_cfg["num_attention_heads"],
|
|
126
|
+
llm_cfg.get("rope_theta", 1e6)).to(self.device)
|
|
127
|
+
with torch.device("meta"):
|
|
128
|
+
self.vae = FluxAutoEncoder()
|
|
129
|
+
self._load(path / "ae.safetensors", {"": self.vae}, dtype=torch.float32)
|
|
130
|
+
|
|
131
|
+
from transformers import Qwen2TokenizerFast # AutoTokenizer would warn about model_type "bagel"
|
|
132
|
+
tok = Qwen2TokenizerFast.from_pretrained(str(path))
|
|
133
|
+
missing = [t for t in ("<|im_start|>", "<|im_end|>", "<|vision_start|>", "<|vision_end|>")
|
|
134
|
+
if t not in tok.get_vocab()]
|
|
135
|
+
if missing:
|
|
136
|
+
tok.add_tokens(missing)
|
|
137
|
+
self.tokenizer = tok
|
|
138
|
+
self.bos = tok.convert_tokens_to_ids("<|im_start|>")
|
|
139
|
+
self.eos = tok.convert_tokens_to_ids("<|im_end|>")
|
|
140
|
+
self.start_of_image = tok.convert_tokens_to_ids("<|vision_start|>")
|
|
141
|
+
self.end_of_image = tok.convert_tokens_to_ids("<|vision_end|>")
|
|
142
|
+
self.stop_ids = (self.eos, tok.convert_tokens_to_ids("<|endoftext|>"))
|
|
143
|
+
|
|
144
|
+
def _load(self, file: Path, modules: dict, dtype=None):
|
|
145
|
+
from safetensors import safe_open
|
|
146
|
+
dtype = dtype or self.dtype
|
|
147
|
+
with safe_open(str(file), framework="pt", device=str(self.device)) as f:
|
|
148
|
+
keys = set(f.keys())
|
|
149
|
+
for prefix, mod in modules.items():
|
|
150
|
+
pre = f"{prefix}." if prefix else ""
|
|
151
|
+
sd = {}
|
|
152
|
+
for name, _ in list(mod.named_parameters()) + list(mod.named_buffers()):
|
|
153
|
+
key = pre + name
|
|
154
|
+
if key in keys:
|
|
155
|
+
sd[name] = f.get_tensor(key).to(dtype)
|
|
156
|
+
missing = [n for n, _ in mod.named_parameters() if n not in sd]
|
|
157
|
+
if missing:
|
|
158
|
+
raise RuntimeError(f"{file.name}: missing weights for {prefix}: {missing[:5]}")
|
|
159
|
+
mod.load_state_dict(sd, strict=False, assign=True)
|
|
160
|
+
mod.eval()
|
|
161
|
+
|
|
162
|
+
# ---------------------------------------------------------- understanding
|
|
163
|
+
def preprocess_image(self, image: Image.Image) -> ImageChunk:
|
|
164
|
+
img = _resize(to_pil(image), 1024, 512, 16, 14 * 14 * 9 * 1024) # BAGEL's VAE-side resize first
|
|
165
|
+
img = _resize(img, 980, 224, self.vit_patch, 14 * 14 * 9 * 1024)
|
|
166
|
+
patches, pos = _patchify(img, self.vit_patch, self.vit_max_side)
|
|
167
|
+
return ImageChunk(n=patches.shape[0] + 2, key=image_hash(img), payload=(patches, pos))
|
|
168
|
+
|
|
169
|
+
def encode_images(self, payloads):
|
|
170
|
+
patches = torch.cat([p for p, _ in payloads]).to(self.device, self.dtype)
|
|
171
|
+
pos = torch.cat([q for _, q in payloads]).to(self.device)
|
|
172
|
+
lens = [p.shape[0] for p, _ in payloads]
|
|
173
|
+
cu = torch.tensor([0] + list(_cumsum(lens)), dtype=torch.int32, device=self.device)
|
|
174
|
+
x = self.vit(patches, pos, cu, max(lens))
|
|
175
|
+
x = self.connector(x) + self.vit_pos_embed(pos).to(x.dtype)
|
|
176
|
+
return list(x.split(lens))
|
|
177
|
+
|
|
178
|
+
# 1-D rope: text positions increase per token, an image block shares one.
|
|
179
|
+
def text_positions(self, start: int, n: int):
|
|
180
|
+
return list(range(start, start + n))
|
|
181
|
+
|
|
182
|
+
def image_positions(self, start: int, chunk):
|
|
183
|
+
return [start] * chunk.n
|
|
184
|
+
|
|
185
|
+
def latent_positions(self, start: int, h: int, w: int):
|
|
186
|
+
return [start] * (h * w + 2)
|
|
187
|
+
|
|
188
|
+
def position_after_image(self, pos: int, chunk) -> int:
|
|
189
|
+
return pos + 1
|
|
190
|
+
|
|
191
|
+
def position_after(self, chunks, n: int) -> int:
|
|
192
|
+
pos = 0
|
|
193
|
+
for c in chunks:
|
|
194
|
+
if n <= 0:
|
|
195
|
+
break
|
|
196
|
+
if isinstance(c, TextChunk):
|
|
197
|
+
take = min(n, len(c.ids))
|
|
198
|
+
pos += take
|
|
199
|
+
n -= take
|
|
200
|
+
else:
|
|
201
|
+
pos += 1
|
|
202
|
+
n -= c.n
|
|
203
|
+
return pos
|
|
204
|
+
|
|
205
|
+
# ------------------------------------------------------------- generation
|
|
206
|
+
def latent_pos_ids(self, h: int, w: int):
|
|
207
|
+
return (torch.arange(h)[:, None] * self.max_latent_side + torch.arange(w)).flatten()
|
|
208
|
+
|
|
209
|
+
def denoise_schedule(self):
|
|
210
|
+
n, s = self.num_timesteps, self.timestep_shift
|
|
211
|
+
t = torch.linspace(1, 0, n)
|
|
212
|
+
t = s * t / (1 + (s - 1) * t)
|
|
213
|
+
return [(float(t[i]), float(t[i] - t[i + 1])) for i in range(n - 1)]
|
|
214
|
+
|
|
215
|
+
def cfg_active(self, t: float) -> bool:
|
|
216
|
+
lo, hi = self.cfg_interval
|
|
217
|
+
return lo < t <= hi and self.cfg_text_scale > 1.0
|
|
218
|
+
|
|
219
|
+
def embed_latents(self, x, t: float, pos_ids):
|
|
220
|
+
tt = torch.full((x.shape[0],), t, device=x.device)
|
|
221
|
+
out = self.vae2llm(x.to(self.dtype)) + self.time_embedder(tt) + self.latent_pos_embed(pos_ids).to(self.dtype)
|
|
222
|
+
return out.to(self.dtype)
|
|
223
|
+
|
|
224
|
+
def latents_to_velocity(self, h):
|
|
225
|
+
return self.llm2vae(h)
|
|
226
|
+
|
|
227
|
+
def apply_cfg(self, v, v_uncond):
|
|
228
|
+
v, v_uncond = v.float(), v_uncond.float()
|
|
229
|
+
guided = v_uncond + self.cfg_text_scale * (v - v_uncond)
|
|
230
|
+
scale = (v.norm() / (guided.norm() + 1e-8)).clamp(min=self.cfg_renorm_min, max=1.0) # "global" renorm
|
|
231
|
+
return guided * scale
|
|
232
|
+
|
|
233
|
+
def decode_latents(self, items):
|
|
234
|
+
out = []
|
|
235
|
+
p, c = self.latent_patch, self.latent_channels
|
|
236
|
+
for x, h, w in items:
|
|
237
|
+
lat = x.reshape(1, h, w, p, p, c)
|
|
238
|
+
lat = torch.einsum("nhwpqc->nchpwq", lat).reshape(1, c, h * p, w * p)
|
|
239
|
+
img = self.vae.decode(lat.float())
|
|
240
|
+
img = ((img * 0.5 + 0.5).clamp(0, 1)[0].permute(1, 2, 0) * 255).to(torch.uint8).cpu().numpy()
|
|
241
|
+
out.append(Image.fromarray(img))
|
|
242
|
+
return out
|
|
243
|
+
|
|
244
|
+
|
|
245
|
+
def _cumsum(xs):
|
|
246
|
+
total = 0
|
|
247
|
+
for x in xs:
|
|
248
|
+
total += x
|
|
249
|
+
yield total
|
|
250
|
+
|
|
251
|
+
|
|
252
|
+
class BagelBackend(Backend):
|
|
253
|
+
name = "bagel"
|
|
254
|
+
native_image_generation = True
|
|
255
|
+
verify_think = True # the SIMIT critic runs BAGEL in thinking mode
|
|
256
|
+
|
|
257
|
+
def __init__(self, model_id: Union[str, "BagelModel"] = "ByteDance-Seed/BAGEL-7B-MoT", device: Optional[str] = None,
|
|
258
|
+
kv_cache_tokens: Optional[int] = None, max_step_tokens: int = 16384,
|
|
259
|
+
image_steps: Optional[int] = None, use_cuda_graphs: bool = True):
|
|
260
|
+
"""``model_id``: a Hub id / local path, or a :class:`BagelModel` loaded earlier
|
|
261
|
+
(e.g. on CPU at startup, then moved per request). ``device``: where to put
|
|
262
|
+
the weights (default ``cuda:0``; for a loaded model, where it already is)."""
|
|
263
|
+
if isinstance(model_id, BagelModel):
|
|
264
|
+
self.model = model_id
|
|
265
|
+
if device is not None and self.model.device != torch.device(device):
|
|
266
|
+
self.model.to(device)
|
|
267
|
+
else:
|
|
268
|
+
self.model = BagelModel.load(model_id, device=device or "cuda:0")
|
|
269
|
+
self.model_id = self.model.model_id
|
|
270
|
+
if image_steps is not None: # fewer denoising steps: faster, slightly lower fidelity
|
|
271
|
+
self.model.num_timesteps = int(image_steps)
|
|
272
|
+
if kv_cache_tokens is None:
|
|
273
|
+
kv_cache_tokens = auto_kv_tokens(self.model)
|
|
274
|
+
self.engine = MoTEngine(self.model, kv_tokens=kv_cache_tokens, max_step_tokens=max_step_tokens,
|
|
275
|
+
use_cuda_graphs=use_cuda_graphs)
|
|
276
|
+
|
|
277
|
+
# ------------------------------------------------------------- prompts
|
|
278
|
+
def _encode(self, text: str) -> list[int]:
|
|
279
|
+
return self.model.tokenizer.encode(text, add_special_tokens=False)
|
|
280
|
+
|
|
281
|
+
def _chunks(self, parts, style: str, think: bool, open_answer: bool = True):
|
|
282
|
+
m = self.model
|
|
283
|
+
chunks = []
|
|
284
|
+
if style == "chat":
|
|
285
|
+
system = DEFAULT_SYSTEM_PROMPT + (f" {CHAT_THINK_PROMPT}" if think else "")
|
|
286
|
+
has_image = any(isinstance(p, Image.Image) for p in parts)
|
|
287
|
+
if not has_image:
|
|
288
|
+
prompt = "".join(parts)
|
|
289
|
+
text = f"<|im_start|>{system}<|im_end|><|im_start|>{prompt}<|im_end|><|im_start|>"
|
|
290
|
+
return [TextChunk(self._encode(text))]
|
|
291
|
+
buf = f"<|im_start|>system\n{system}<|im_end|>\n<|im_start|>user\n"
|
|
292
|
+
after_image = False
|
|
293
|
+
for p in parts:
|
|
294
|
+
if isinstance(p, Image.Image):
|
|
295
|
+
chunks.append(TextChunk(self._encode(buf)))
|
|
296
|
+
chunks.append(m.preprocess_image(p))
|
|
297
|
+
buf, after_image = "", True
|
|
298
|
+
else:
|
|
299
|
+
buf += ("\n" if after_image else "") + p
|
|
300
|
+
after_image = False
|
|
301
|
+
buf += "<|im_end|>\n<|im_start|>assistant\n"
|
|
302
|
+
chunks.append(TextChunk(self._encode(buf)))
|
|
303
|
+
return [c for c in chunks if not isinstance(c, TextChunk) or c.ids]
|
|
304
|
+
# "answer": every text part is its own <|im_start|>...<|im_end|> turn
|
|
305
|
+
if think:
|
|
306
|
+
chunks.append(TextChunk([m.bos] + self._encode(ANSWER_THINK_PROMPT) + [m.eos]))
|
|
307
|
+
for p in parts:
|
|
308
|
+
if isinstance(p, Image.Image):
|
|
309
|
+
chunks.append(m.preprocess_image(p))
|
|
310
|
+
else:
|
|
311
|
+
chunks.append(TextChunk([m.bos] + self._encode(p) + [m.eos]))
|
|
312
|
+
if open_answer:
|
|
313
|
+
chunks.append(TextChunk([m.bos]))
|
|
314
|
+
return chunks
|
|
315
|
+
|
|
316
|
+
# -------------------------------------------------------------- requests
|
|
317
|
+
def submit_generate(self, req: GenRequest):
|
|
318
|
+
rp = req.repetition_penalty
|
|
319
|
+
if rp is None:
|
|
320
|
+
rp = 1.05 if req.style == "chat" else 1.0
|
|
321
|
+
job = TextJob(chunks=self._chunks(req.parts, req.style, req.think), max_new_tokens=req.max_new_tokens,
|
|
322
|
+
temperature=req.temperature, top_p=req.top_p, repetition_penalty=rp,
|
|
323
|
+
logprobs=req.logprobs, stop_ids=self.model.stop_ids)
|
|
324
|
+
tok = self.model.tokenizer
|
|
325
|
+
if req.on_text is not None:
|
|
326
|
+
job.on_tokens = lambda ids: req.on_text(tok.decode(ids, skip_special_tokens=True))
|
|
327
|
+
|
|
328
|
+
def post(seq):
|
|
329
|
+
raw = tok.decode(seq.out_ids, skip_special_tokens=False)
|
|
330
|
+
raw = raw.split("<|im_end|>")[0].replace("<|im_start|>", "")
|
|
331
|
+
text = strip_thinking(raw) if req.think else raw.strip()
|
|
332
|
+
return GenResult(text=text, raw_text=raw, token_logprobs=seq.out_lp if req.logprobs else None,
|
|
333
|
+
num_tokens=len(seq.out_ids))
|
|
334
|
+
|
|
335
|
+
return chained(self.engine.submit(job), post)
|
|
336
|
+
|
|
337
|
+
def submit_score(self, req: ScoreRequest):
|
|
338
|
+
target = self._encode(req.target)
|
|
339
|
+
if not target:
|
|
340
|
+
from concurrent.futures import Future
|
|
341
|
+
f = Future()
|
|
342
|
+
f.set_result([])
|
|
343
|
+
return f
|
|
344
|
+
chunks = self._chunks(req.parts, req.style, think=False, open_answer=False)
|
|
345
|
+
if req.style == "chat":
|
|
346
|
+
chunks[-1] = TextChunk(chunks[-1].ids + target)
|
|
347
|
+
else:
|
|
348
|
+
chunks.append(TextChunk([self.model.bos] + target))
|
|
349
|
+
job = TextJob(chunks=chunks, score_len=len(target), stop_ids=self.model.stop_ids)
|
|
350
|
+
return chained(self.engine.submit(job), lambda seq: seq.score_lp)
|
|
351
|
+
|
|
352
|
+
def submit_image(self, req: ImageRequest):
|
|
353
|
+
w, h = (max(16, int(round(v / 16)) * 16) for v in (req.width, req.height))
|
|
354
|
+
text = f"<|im_start|>{req.prompt}<|im_end|><|im_start|>"
|
|
355
|
+
job = ImageJob(chunks=[TextChunk(self._encode(text))], height=h, width=w, seed=req.seed)
|
|
356
|
+
return self.engine.submit(job)
|
|
357
|
+
|
|
358
|
+
def close(self):
|
|
359
|
+
self.engine.close()
|
simit/backends/base.py
ADDED
|
@@ -0,0 +1,114 @@
|
|
|
1
|
+
"""Backend interface: every model (UMM or plain VLM) serves three request
|
|
2
|
+
kinds, each submitted asynchronously and batched by the backend itself."""
|
|
3
|
+
|
|
4
|
+
from __future__ import annotations
|
|
5
|
+
|
|
6
|
+
from concurrent.futures import Future
|
|
7
|
+
from dataclasses import dataclass, field
|
|
8
|
+
from typing import Callable, Optional, Sequence, Union
|
|
9
|
+
|
|
10
|
+
from PIL import Image
|
|
11
|
+
|
|
12
|
+
Part = Union[str, Image.Image]
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
@dataclass
|
|
16
|
+
class GenRequest:
|
|
17
|
+
"""Generate text for one user turn of interleaved text/images.
|
|
18
|
+
|
|
19
|
+
``style`` selects the prompt format: ``"chat"`` for the synthesis
|
|
20
|
+
pipeline's instruction prompts, ``"answer"`` for answering benchmark-style
|
|
21
|
+
questions (zero-shot and in-context). Backends may format them alike."""
|
|
22
|
+
|
|
23
|
+
parts: Sequence[Part]
|
|
24
|
+
style: str = "chat"
|
|
25
|
+
max_new_tokens: int = 256
|
|
26
|
+
temperature: float = 0.0
|
|
27
|
+
top_p: float = 1.0
|
|
28
|
+
repetition_penalty: Optional[float] = None # None -> backend default for the style
|
|
29
|
+
think: bool = False
|
|
30
|
+
logprobs: bool = False
|
|
31
|
+
#: Optional ``fn(partial_text)`` called (from a backend thread) as tokens
|
|
32
|
+
#: are decoded; backends without streaming simply never call it.
|
|
33
|
+
on_text: Optional[Callable[[str], None]] = None
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
@dataclass
|
|
37
|
+
class GenResult:
|
|
38
|
+
text: str # answer, thinking removed
|
|
39
|
+
raw_text: str = ""
|
|
40
|
+
token_logprobs: Optional[list[float]] = None
|
|
41
|
+
num_tokens: int = 0
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
@dataclass
|
|
45
|
+
class ScoreRequest:
|
|
46
|
+
"""Log-probabilities of ``target`` as the answer to ``parts`` (teacher forcing)."""
|
|
47
|
+
|
|
48
|
+
parts: Sequence[Part]
|
|
49
|
+
target: str
|
|
50
|
+
style: str = "answer"
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
@dataclass
|
|
54
|
+
class ImageRequest:
|
|
55
|
+
prompt: str
|
|
56
|
+
width: int = 400
|
|
57
|
+
height: int = 400
|
|
58
|
+
seed: Optional[int] = None
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
class Backend:
|
|
62
|
+
"""Base class. Subclasses implement the three ``submit_*`` methods; each
|
|
63
|
+
returns a ``concurrent.futures.Future`` and may batch freely."""
|
|
64
|
+
|
|
65
|
+
name: str = "backend"
|
|
66
|
+
#: True for unified models that generate images natively.
|
|
67
|
+
native_image_generation: bool = False
|
|
68
|
+
#: Default for running the verification critic in thinking mode.
|
|
69
|
+
verify_think: bool = False
|
|
70
|
+
#: Defaults for ``SIMITConfig`` fields left as ``None`` (synthesis, use_skills, verify).
|
|
71
|
+
pipeline_defaults: dict = {}
|
|
72
|
+
|
|
73
|
+
def submit_generate(self, req: GenRequest) -> Future:
|
|
74
|
+
raise NotImplementedError
|
|
75
|
+
|
|
76
|
+
def submit_score(self, req: ScoreRequest) -> Future:
|
|
77
|
+
raise NotImplementedError
|
|
78
|
+
|
|
79
|
+
def submit_image(self, req: ImageRequest) -> Future:
|
|
80
|
+
raise NotImplementedError(f"{self.name} cannot generate images")
|
|
81
|
+
|
|
82
|
+
def close(self):
|
|
83
|
+
pass
|
|
84
|
+
|
|
85
|
+
# Synchronous conveniences -------------------------------------------
|
|
86
|
+
def generate(self, reqs: list[GenRequest]) -> list[GenResult]:
|
|
87
|
+
futs = [self.submit_generate(r) for r in reqs]
|
|
88
|
+
return [f.result() for f in futs]
|
|
89
|
+
|
|
90
|
+
def score(self, reqs: list[ScoreRequest]) -> list[list[float]]:
|
|
91
|
+
futs = [self.submit_score(r) for r in reqs]
|
|
92
|
+
return [f.result() for f in futs]
|
|
93
|
+
|
|
94
|
+
def generate_images(self, reqs: list[ImageRequest]) -> list[Image.Image]:
|
|
95
|
+
futs = [self.submit_image(r) for r in reqs]
|
|
96
|
+
return [f.result() for f in futs]
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def chained(fut: Future, fn) -> Future:
|
|
100
|
+
"""A future resolving to ``fn(fut.result())``."""
|
|
101
|
+
out: Future = Future()
|
|
102
|
+
|
|
103
|
+
def done(f):
|
|
104
|
+
if out.done():
|
|
105
|
+
return
|
|
106
|
+
try:
|
|
107
|
+
out.set_result(fn(f.result()))
|
|
108
|
+
except BaseException as e: # propagate engine/postprocess errors to the caller
|
|
109
|
+
if not out.done():
|
|
110
|
+
out.set_exception(e)
|
|
111
|
+
|
|
112
|
+
fut.add_done_callback(done)
|
|
113
|
+
out.add_done_callback(lambda o: fut.cancel() if o.cancelled() else None) # cancellation reaches the engine
|
|
114
|
+
return out
|