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.
@@ -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