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/lance.py
ADDED
|
@@ -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()
|