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,338 @@
1
+ """Lance-3B (bytedance-research/Lance) on the batched MoT engine.
2
+
3
+ Lance pairs BAGEL's Mixture-of-Transformers trunk with the Qwen2.5-VL vision
4
+ tower, multimodal (t, h, w) RoPE and the Wan2.2 VAE. Prompt formats and
5
+ decoding follow the setup the SIMIT paper evaluated Lance with: chat-template
6
+ turns with task system prompts, a newline start token inside the assistant
7
+ turn, sampling (T >= 0.8, repetition penalty 1.05) for free generation and
8
+ greedy decoding for confidence scores.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import json
14
+ from pathlib import Path
15
+ from typing import Optional, Union
16
+
17
+ import torch
18
+ from PIL import Image
19
+
20
+ from ..utils import image_hash, strip_thinking, to_pil
21
+ from .base import Backend, GenRequest, GenResult, ImageRequest, ScoreRequest, chained
22
+ from .mot.engine import ImageChunk, ImageJob, MoTEngine, TextChunk, TextJob, auto_kv_tokens
23
+ from .mot.modeling import DeviceMixin, MoTLanguageModel, MRotaryEmbedding, PositionTable, TimestepEmbedder
24
+
25
+ X2T_SYSTEM_PROMPT = "Look at the image carefully and answer the question."
26
+ TEXT_SYSTEM_PROMPT = "You are a helpful assistant."
27
+ T2I_SYSTEM_PROMPT = ("Describe the image by detailing the color, quantity, text, shape, size, texture, "
28
+ "spatial relationships of the objects and background:")
29
+
30
+ _VIT_BUCKET_RESOLUTION = 672
31
+ _VIT_BUCKET_STRIDE = 16
32
+ _VIT_DIVISIBLE_CROP = 28
33
+ _VIT_ASPECT_RATIOS = ("21:9", "16:9", "4:3", "1:1", "3:4", "9:16")
34
+ _VIT_T_SHIFT = 1000 # ViT tokens live far away on the t axis of mRoPE
35
+
36
+
37
+ def _vit_buckets():
38
+ import math
39
+ max_area = _VIT_BUCKET_RESOLUTION ** 2
40
+ out = []
41
+ for name in _VIT_ASPECT_RATIOS:
42
+ ws, hs = (int(v) for v in name.split(":"))
43
+ aspect = ws / hs
44
+ bw1 = round(math.sqrt(max_area * aspect) / _VIT_BUCKET_STRIDE) * _VIT_BUCKET_STRIDE
45
+ bh1 = round(bw1 / aspect / _VIT_BUCKET_STRIDE) * _VIT_BUCKET_STRIDE
46
+ bh2 = round(math.sqrt(max_area / aspect) / _VIT_BUCKET_STRIDE) * _VIT_BUCKET_STRIDE
47
+ bw2 = round(bh2 * aspect / _VIT_BUCKET_STRIDE) * _VIT_BUCKET_STRIDE
48
+ d1, d2 = abs(bw1 / bh1 - aspect), abs(bw2 / bh2 - aspect)
49
+ if d1 < d2 or (d1 == d2 and abs(bh1 * bw1 - max_area) <= abs(bh2 * bw2 - max_area)):
50
+ out.append(((bh1, bw1), bw1 / bh1))
51
+ else:
52
+ out.append(((bh2, bw2), bw2 / bh2))
53
+ return out
54
+
55
+
56
+ _BUCKETS = _vit_buckets()
57
+
58
+
59
+ def _bucket_resize(img: Image.Image) -> Image.Image:
60
+ """Lance's ViT preprocessing: nearest aspect bucket (center crop + bicubic
61
+ resize to ~672^2 px), then center-crop to multiples of 28."""
62
+ from torchvision.transforms import InterpolationMode, RandomResizedCrop
63
+ ratio = img.width / img.height
64
+ (bh, bw), br = min(_BUCKETS, key=lambda b: abs(ratio - b[1]))
65
+ out = RandomResizedCrop(size=(bh, bw), scale=(1.0, 1.0), ratio=(br, br),
66
+ interpolation=InterpolationMode.BICUBIC)(img)
67
+ w, h = out.size
68
+ cw, ch = w - w % _VIT_DIVISIBLE_CROP, h - h % _VIT_DIVISIBLE_CROP
69
+ if (cw, ch) != (w, h):
70
+ left, top = (w - cw) // 2, (h - ch) // 2
71
+ out = out.crop((left, top, left + cw, top + ch))
72
+ return out
73
+
74
+
75
+ class LanceModel(DeviceMixin):
76
+ latent_channels = 48
77
+ latent_downsample = 16
78
+ max_latent_side = 64
79
+ num_timesteps = 30
80
+ timestep_shift = 3.5
81
+ cfg_text_scale = 4.0
82
+ cfg_interval = (0.4, 1.0)
83
+ cfg_renorm_min = 0.0
84
+
85
+ @classmethod
86
+ def load(cls, model_id: str = "bytedance-research/Lance", device: str = "cuda",
87
+ dtype=torch.bfloat16) -> "LanceModel":
88
+ """From a Hub id (downloaded on first use) or a local path."""
89
+ path = Path(model_id)
90
+ if not (path / "Lance_3B").exists():
91
+ from huggingface_hub import snapshot_download
92
+ path = Path(snapshot_download(model_id, allow_patterns=[
93
+ "config.json", "Lance_3B/*", "Qwen2.5-VL-ViT/*", "Wan2.2_VAE.pth"]))
94
+ model = cls(str(path), device=device, dtype=dtype)
95
+ model.model_id = model_id
96
+ return model
97
+
98
+ def __init__(self, root: str, device: str = "cuda", dtype=torch.bfloat16):
99
+ self.model_id = str(root)
100
+ from safetensors import safe_open
101
+ from transformers import AutoTokenizer, Qwen2VLImageProcessor
102
+ from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import Qwen2_5_VLVisionConfig
103
+ from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import Qwen2_5_VisionTransformerPretrainedModel
104
+
105
+ from .mot.wan_vae import LanceWanVAE
106
+
107
+ self.device, self.dtype = torch.device(device), dtype
108
+ root = Path(root)
109
+ ckpt = root / "Lance_3B"
110
+ cfg = json.loads((ckpt / "llm_config.json").read_text())
111
+ hidden = cfg["hidden_size"]
112
+ self.latent_dim = self.latent_channels
113
+ with torch.device("meta"):
114
+ self.lm = MoTLanguageModel(cfg)
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
+ with safe_open(str(ckpt / "model.safetensors"), framework="pt", device=str(self.device)) as f:
120
+ keys = set(f.keys())
121
+ for prefix, mod in (("language_model", self.lm), ("latent_pos_embed", self.latent_pos_embed),
122
+ ("time_embedder", self.time_embedder), ("vae2llm", self.vae2llm),
123
+ ("llm2vae", self.llm2vae)):
124
+ sd = {n: f.get_tensor(f"{prefix}.{n}").to(dtype) for n, _ in mod.named_parameters()
125
+ if f"{prefix}.{n}" in keys}
126
+ if "lm_head.weight" not in sd and prefix == "language_model": # tied embeddings
127
+ sd["lm_head.weight"] = sd["model.embed_tokens.weight"]
128
+ missing = [n for n, _ in mod.named_parameters() if n not in sd]
129
+ if missing:
130
+ raise RuntimeError(f"Lance checkpoint is missing {prefix} weights: {missing[:5]}")
131
+ mod.load_state_dict(sd, strict=False, assign=True)
132
+ mod.eval()
133
+ head_dim = hidden // cfg["num_attention_heads"]
134
+ sections = cfg.get("rope_scaling", {}).get("mrope_section", [16, 24, 24])
135
+ self.lm.rotary = MRotaryEmbedding(head_dim, cfg.get("rope_theta", 1e6), sections).to(self.device)
136
+
137
+ vit_cfg = json.loads((root / "Qwen2.5-VL-ViT" / "config.json").read_text())
138
+ vit_cfg["_attn_implementation"] = "sdpa"
139
+ self.vit = Qwen2_5_VisionTransformerPretrainedModel(Qwen2_5_VLVisionConfig(**vit_cfg))
140
+ with safe_open(str(root / "Qwen2.5-VL-ViT" / "vit.safetensors"), framework="pt", device="cpu") as f:
141
+ sd = {k: f.get_tensor(k) for k in f.keys()}
142
+ missing, unexpected = self.vit.load_state_dict(sd, strict=False)
143
+ if [m for m in missing if "rotary" not in m]:
144
+ raise RuntimeError(f"Lance ViT checkpoint mismatch: missing {missing[:5]}")
145
+ self.vit.to(self.device, dtype).eval()
146
+ self.merge = vit_cfg.get("spatial_merge_size", 2)
147
+ self.image_processor = Qwen2VLImageProcessor()
148
+
149
+ self.vae = LanceWanVAE(str(root / "Wan2.2_VAE.pth"), dtype=torch.bfloat16, device=self.device)
150
+
151
+ tok = AutoTokenizer.from_pretrained(str(ckpt))
152
+ self.tokenizer = tok
153
+ self.bos = tok.convert_tokens_to_ids("<|im_start|>")
154
+ self.eos = tok.convert_tokens_to_ids("<|im_end|>")
155
+ self.start_of_image = tok.convert_tokens_to_ids("<|vision_start|>")
156
+ self.end_of_image = tok.convert_tokens_to_ids("<|vision_end|>")
157
+ self.stop_ids = (self.eos, tok.convert_tokens_to_ids("<|endoftext|>"))
158
+ nl = tok.encode("\n", add_special_tokens=False)
159
+ assert len(nl) == 1
160
+ self.newline = nl[0]
161
+
162
+ # ---------------------------------------------------------- understanding
163
+ def preprocess_image(self, image) -> ImageChunk:
164
+ img = _bucket_resize(to_pil(image))
165
+ out = self.image_processor(images=img, return_tensors="pt")
166
+ pix, grid = out["pixel_values"], out["image_grid_thw"]
167
+ t, h, w = (int(v) for v in grid[0].tolist())
168
+ n = t * h * w // (self.merge ** 2)
169
+ return ImageChunk(n=n + 2, key=image_hash(img), payload=(pix, grid, h // self.merge, w // self.merge))
170
+
171
+ def encode_images(self, payloads):
172
+ pix = torch.cat([p[0] for p in payloads]).to(self.device, self.dtype)
173
+ grid = torch.cat([p[1] for p in payloads]).to(self.device)
174
+ out = self.vit(hidden_states=pix, grid_thw=grid)
175
+ if torch.is_tensor(out):
176
+ merged = out
177
+ else: # BaseModelOutputWithPooling: the merger output is the pooled one
178
+ merged = getattr(out, "pooler_output", None)
179
+ if merged is None:
180
+ merged = out[0]
181
+ sizes = [p[2] * p[3] for p in payloads]
182
+ return list(merged.to(self.dtype).split(sizes))
183
+
184
+ # (t, h, w) rope positions
185
+ def text_positions(self, start: int, n: int):
186
+ return [(p, p, p) for p in range(start, start + n)]
187
+
188
+ def image_positions(self, start: int, chunk):
189
+ _, _, hm, wm = chunk.payload
190
+ m = max(hm, wm)
191
+ pos = [(_VIT_T_SHIFT, start, start)]
192
+ pos += [(_VIT_T_SHIFT + 1, start + 1 + i, start + 1 + j) for i in range(hm) for j in range(wm)]
193
+ pos.append((_VIT_T_SHIFT + m + 1, start + m + 1, start + m + 1))
194
+ return pos
195
+
196
+ def position_after_image(self, pos: int, chunk) -> int:
197
+ _, _, hm, wm = chunk.payload
198
+ return pos + max(hm, wm) + 2
199
+
200
+ def position_after(self, chunks, n: int) -> int:
201
+ return n # the prefix cache only covers leading text
202
+
203
+ # -------------------------------------------------------------- generation
204
+ def latent_positions(self, start: int, h: int, w: int):
205
+ m = max(h, w)
206
+ pos = [(start, start, start)]
207
+ pos += [(start + 1, start + 1 + i, start + 1 + j) for i in range(h) for j in range(w)]
208
+ pos.append((start + m + 1,) * 3)
209
+ return pos
210
+
211
+ def latent_pos_ids(self, h: int, w: int):
212
+ return (torch.arange(h)[:, None] * self.max_latent_side + torch.arange(w)).flatten()
213
+
214
+ def denoise_schedule(self):
215
+ n, s = self.num_timesteps, self.timestep_shift
216
+ t = torch.linspace(1, 0, n + 1) # Lance runs exactly `num_timesteps` Euler steps
217
+ t = s * t / (1 + (s - 1) * t)
218
+ return [(float(t[i]), float(t[i] - t[i + 1])) for i in range(n)]
219
+
220
+ def cfg_active(self, t: float) -> bool:
221
+ lo, hi = self.cfg_interval
222
+ return lo < t <= hi and self.cfg_text_scale > 1.0
223
+
224
+ def embed_latents(self, x, t: float, pos_ids):
225
+ tt = torch.full((x.shape[0],), t, device=x.device)
226
+ out = self.vae2llm(x.to(self.dtype)) + self.time_embedder(tt) + self.latent_pos_embed(pos_ids).to(self.dtype)
227
+ return out.to(self.dtype)
228
+
229
+ def latents_to_velocity(self, h):
230
+ return self.llm2vae(h)
231
+
232
+ def apply_cfg(self, v, v_uncond):
233
+ v, v_uncond = v.float(), v_uncond.float()
234
+ guided = v_uncond + self.cfg_text_scale * (v - v_uncond)
235
+ scale = (v.norm() / (guided.norm() + 1e-8)).clamp(min=self.cfg_renorm_min, max=1.0)
236
+ return guided * scale
237
+
238
+ def decode_latents(self, items):
239
+ out = []
240
+ for x, h, w in items:
241
+ lat = x.reshape(1, h, w, self.latent_channels).permute(0, 3, 1, 2)
242
+ img = self.vae.decode(lat.to(torch.bfloat16)).float()
243
+ img = ((img * 0.5 + 0.5).clamp(0, 1)[0].permute(1, 2, 0) * 255).to(torch.uint8).cpu().numpy()
244
+ out.append(Image.fromarray(img))
245
+ return out
246
+
247
+
248
+ class LanceBackend(Backend):
249
+ name = "lance"
250
+ native_image_generation = True
251
+ verify_think = False
252
+ #: Lance follows short single-turn instructions only (see the paper's
253
+ #: cross-model study): triplets are synthesized in three simple calls,
254
+ #: every description is realized by native generation, no critic.
255
+ pipeline_defaults = {"synthesis": "decomposed", "use_skills": False, "verify": False}
256
+ #: Greedy decoding degenerates on this checkpoint; free generation samples.
257
+ min_temperature = 0.8
258
+ repetition_penalty = 1.05
259
+
260
+ def __init__(self, model_id: Union[str, "LanceModel"] = "bytedance-research/Lance", device: Optional[str] = None,
261
+ kv_cache_tokens: Optional[int] = None, max_step_tokens: int = 16384,
262
+ image_steps: Optional[int] = None, use_cuda_graphs: bool = True):
263
+ """``model_id``: a Hub id / local path, or a :class:`LanceModel` loaded earlier
264
+ (e.g. on CPU at startup, then moved per request). ``device``: where to put
265
+ the weights (default ``cuda:0``; for a loaded model, where it already is)."""
266
+ if isinstance(model_id, LanceModel):
267
+ self.model = model_id
268
+ if device is not None and self.model.device != torch.device(device):
269
+ self.model.to(device)
270
+ else:
271
+ self.model = LanceModel.load(model_id, device=device or "cuda:0")
272
+ self.model_id = self.model.model_id
273
+ if image_steps is not None: # fewer denoising steps: faster, slightly lower fidelity
274
+ self.model.num_timesteps = int(image_steps)
275
+ if kv_cache_tokens is None:
276
+ kv_cache_tokens = auto_kv_tokens(self.model)
277
+ self.engine = MoTEngine(self.model, kv_tokens=kv_cache_tokens, max_step_tokens=max_step_tokens,
278
+ use_cuda_graphs=use_cuda_graphs)
279
+
280
+ def _encode(self, text: str):
281
+ return self.model.tokenizer.encode(text, add_special_tokens=False)
282
+
283
+ def _chunks(self, parts):
284
+ m = self.model
285
+ system = X2T_SYSTEM_PROMPT if any(isinstance(p, Image.Image) for p in parts) else TEXT_SYSTEM_PROMPT
286
+ chunks, buf = [], f"<|im_start|>system\n{system}<|im_end|>\n<|im_start|>user\n"
287
+ for p in parts:
288
+ if isinstance(p, Image.Image):
289
+ chunks.append(TextChunk(self._encode(buf)))
290
+ chunks.append(m.preprocess_image(p))
291
+ buf = ""
292
+ else:
293
+ buf += p
294
+ buf += "<|im_end|>\n<|im_start|>assistant"
295
+ chunks.append(TextChunk(self._encode(buf) + [m.newline])) # start inside the assistant turn
296
+ return [c for c in chunks if not isinstance(c, TextChunk) or c.ids]
297
+
298
+ def submit_generate(self, req: GenRequest):
299
+ if req.logprobs: # confidence: greedy, as the paper's Lance scores were computed
300
+ temperature, rp = 0.0, 1.0
301
+ else:
302
+ temperature = max(req.temperature, self.min_temperature)
303
+ rp = req.repetition_penalty or self.repetition_penalty
304
+ job = TextJob(chunks=self._chunks(req.parts), max_new_tokens=req.max_new_tokens, temperature=temperature,
305
+ top_p=req.top_p, repetition_penalty=rp, logprobs=req.logprobs, stop_ids=self.model.stop_ids)
306
+ tok = self.model.tokenizer
307
+ if req.on_text is not None:
308
+ job.on_tokens = lambda ids: req.on_text(tok.decode(ids, skip_special_tokens=True))
309
+
310
+ def post(seq):
311
+ raw = tok.decode(seq.out_ids, skip_special_tokens=True)
312
+ text = strip_thinking(raw) if req.think else raw.strip()
313
+ return GenResult(text=text, raw_text=raw, token_logprobs=seq.out_lp if req.logprobs else None,
314
+ num_tokens=len(seq.out_ids))
315
+
316
+ return chained(self.engine.submit(job), post)
317
+
318
+ def submit_score(self, req: ScoreRequest):
319
+ target = self._encode(req.target)
320
+ if not target:
321
+ from concurrent.futures import Future
322
+ f = Future()
323
+ f.set_result([])
324
+ return f
325
+ chunks = self._chunks(req.parts) + [TextChunk(target)]
326
+ job = TextJob(chunks=chunks, score_len=len(target), stop_ids=self.model.stop_ids)
327
+ return chained(self.engine.submit(job), lambda seq: seq.score_lp)
328
+
329
+ def submit_image(self, req: ImageRequest):
330
+ m = self.model
331
+ w, h = (max(16, int(round(v / 16)) * 16) for v in (req.width, req.height))
332
+ text = (f"system\n{T2I_SYSTEM_PROMPT}<|im_end|>\n<|im_start|>user\n{req.prompt}<|im_end|>\n"
333
+ f"<|im_start|>assistant\n")
334
+ job = ImageJob(chunks=[TextChunk([m.bos] + self._encode(text) + [m.eos])], height=h, width=w, seed=req.seed)
335
+ return self.engine.submit(job)
336
+
337
+ def close(self):
338
+ self.engine.close()