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/pipeline/core.py
ADDED
|
@@ -0,0 +1,586 @@
|
|
|
1
|
+
"""The SIMIT synthesis pipeline (Fig. 2 of the paper), written as coroutines.
|
|
2
|
+
|
|
3
|
+
Every query runs as its own coroutine on a background event loop; their
|
|
4
|
+
model calls are batched by the backend's engine, and rendering runs in the
|
|
5
|
+
worker pool, so many queries (and many candidates per query) progress
|
|
6
|
+
concurrently.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import asyncio
|
|
12
|
+
import logging
|
|
13
|
+
import math
|
|
14
|
+
import re
|
|
15
|
+
import threading
|
|
16
|
+
import time
|
|
17
|
+
from collections import deque
|
|
18
|
+
from typing import Callable, Optional, Sequence
|
|
19
|
+
|
|
20
|
+
from PIL import Image
|
|
21
|
+
|
|
22
|
+
from ..backends.base import Backend, GenRequest, ImageRequest, ScoreRequest
|
|
23
|
+
from ..config import SIMITConfig
|
|
24
|
+
from ..skills.base import Realization, Skill
|
|
25
|
+
from ..skills.pool import RenderPool
|
|
26
|
+
from ..skills.routing import build_router_prompt, parse_route, pre_route
|
|
27
|
+
from ..types import Demo, Imagination, ZeroShot
|
|
28
|
+
from ..utils import clean_response, to_pil
|
|
29
|
+
from . import prompts as P
|
|
30
|
+
|
|
31
|
+
log = logging.getLogger("simit")
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
class Runtime:
|
|
35
|
+
"""A background asyncio loop shared by all pipeline coroutines."""
|
|
36
|
+
|
|
37
|
+
def __init__(self):
|
|
38
|
+
self.loop = asyncio.new_event_loop()
|
|
39
|
+
self.thread = threading.Thread(target=self.loop.run_forever, daemon=True, name="simit-loop")
|
|
40
|
+
self.thread.start()
|
|
41
|
+
|
|
42
|
+
def run(self, coro):
|
|
43
|
+
if threading.current_thread() is self.thread:
|
|
44
|
+
raise RuntimeError("SIMIT's blocking API cannot be called from its own event loop")
|
|
45
|
+
return asyncio.run_coroutine_threadsafe(coro, self.loop).result()
|
|
46
|
+
|
|
47
|
+
def close(self):
|
|
48
|
+
self.loop.call_soon_threadsafe(self.loop.stop)
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
class SkillContext:
|
|
52
|
+
"""What a ``Skill.realize`` implementation may call."""
|
|
53
|
+
|
|
54
|
+
def __init__(self, pipe: "Pipeline", size: int, allow_composite: bool = True):
|
|
55
|
+
self.pipe, self.size, self.allow_composite = pipe, size, allow_composite
|
|
56
|
+
|
|
57
|
+
async def generate(self, prompt: str, max_new_tokens: int = 512, temperature: float = 0.0,
|
|
58
|
+
images: Optional[Sequence[Image.Image]] = None) -> str:
|
|
59
|
+
parts = list(images or []) + [prompt]
|
|
60
|
+
res = await self.pipe.generate(GenRequest(parts=parts, style="chat", max_new_tokens=max_new_tokens,
|
|
61
|
+
temperature=temperature), tag="skill_spec")
|
|
62
|
+
return res.text
|
|
63
|
+
|
|
64
|
+
async def render(self, skill: Skill, spec) -> Image.Image:
|
|
65
|
+
self.pipe._started(f"render:{skill.name}")
|
|
66
|
+
t0 = time.perf_counter()
|
|
67
|
+
try:
|
|
68
|
+
return await asyncio.wrap_future(self.pipe.pool.submit(skill, spec))
|
|
69
|
+
finally:
|
|
70
|
+
if self.pipe.trace is not None:
|
|
71
|
+
self.pipe.trace.append((f"render:{skill.name}", time.perf_counter() - t0, 0))
|
|
72
|
+
|
|
73
|
+
async def generate_image(self, prompt: str) -> Optional[Image.Image]:
|
|
74
|
+
return await self.pipe.generate_image(prompt, self.size)
|
|
75
|
+
|
|
76
|
+
async def realize(self, request: str, retries: int = 1, allow_composite: bool = False,
|
|
77
|
+
size: Optional[int] = None) -> Optional[Realization]:
|
|
78
|
+
"""Route and realize a sub-request without verification (composite panels)."""
|
|
79
|
+
skill = await self.pipe.route(request, allow_composite=allow_composite)
|
|
80
|
+
ctx = SkillContext(self.pipe, size or self.size, allow_composite)
|
|
81
|
+
try:
|
|
82
|
+
return await skill.realize(ctx, request, retries=retries)
|
|
83
|
+
except Exception as e: # a broken panel must not sink the composite
|
|
84
|
+
log.debug("sub-realization failed: %s", e)
|
|
85
|
+
return None
|
|
86
|
+
|
|
87
|
+
def log(self, msg: str):
|
|
88
|
+
log.debug(msg)
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
_SCORE_RE = re.compile(r"SCORE:\s*(\d{1,3})", re.IGNORECASE)
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def _parse_score(text: str) -> Optional[int]:
|
|
95
|
+
m = _SCORE_RE.findall(text)
|
|
96
|
+
return max(0, min(100, int(m[-1]))) if m else None
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def _norm_q(question: str) -> str:
|
|
100
|
+
return " ".join(re.sub(r"[^\w\s]", " ", question.lower()).split())
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
def _mean_prob(logprobs) -> Optional[float]:
|
|
104
|
+
return math.exp(sum(logprobs) / len(logprobs)) if logprobs else None
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
class Pipeline:
|
|
108
|
+
def __init__(self, backend: Backend, skills: Sequence[Skill], config: SIMITConfig,
|
|
109
|
+
image_generator: Optional[Backend] = None, pool: Optional[RenderPool] = None,
|
|
110
|
+
runtime: Optional[Runtime] = None):
|
|
111
|
+
self.backend = backend
|
|
112
|
+
self.image_generator = image_generator
|
|
113
|
+
self.config = config
|
|
114
|
+
self.pool = pool or RenderPool()
|
|
115
|
+
self.runtime = runtime or Runtime()
|
|
116
|
+
self.trace: Optional[list] = None # set to [] to record (stage, seconds, tokens) per model call
|
|
117
|
+
self.on_call: Optional[Callable[[str], None]] = None # called with a stage tag as each model call starts
|
|
118
|
+
self.set_skills(skills)
|
|
119
|
+
|
|
120
|
+
# ------------------------------------------------------------ setup
|
|
121
|
+
@property
|
|
122
|
+
def can_generate_images(self) -> bool:
|
|
123
|
+
return self.backend.native_image_generation or self.image_generator is not None
|
|
124
|
+
|
|
125
|
+
def set_skills(self, skills: Sequence[Skill]):
|
|
126
|
+
skills = list(skills)
|
|
127
|
+
names = [s.name for s in skills]
|
|
128
|
+
if len(set(names)) != len(names):
|
|
129
|
+
raise ValueError(f"duplicate skill names: {names}")
|
|
130
|
+
if not self.can_generate_images: # structured skills only
|
|
131
|
+
skills = [s for s in skills if s.name != "natural"]
|
|
132
|
+
self.skills = skills
|
|
133
|
+
self.by_name = {s.name: s for s in skills}
|
|
134
|
+
self.router_prompt = build_router_prompt(skills)
|
|
135
|
+
|
|
136
|
+
def option(self, name: str, default):
|
|
137
|
+
value = getattr(self.config, name)
|
|
138
|
+
if value is None:
|
|
139
|
+
value = getattr(self.backend, "pipeline_defaults", {}).get(name, default)
|
|
140
|
+
return value
|
|
141
|
+
|
|
142
|
+
@property
|
|
143
|
+
def use_skills(self) -> bool:
|
|
144
|
+
return bool(self.option("use_skills", True)) or "natural" not in self.by_name
|
|
145
|
+
|
|
146
|
+
@property
|
|
147
|
+
def verify_enabled(self) -> bool:
|
|
148
|
+
return bool(self.option("verify", True))
|
|
149
|
+
|
|
150
|
+
@property
|
|
151
|
+
def verify_think(self) -> bool:
|
|
152
|
+
v = self.config.verify_think
|
|
153
|
+
return self.backend.verify_think if v is None else v
|
|
154
|
+
|
|
155
|
+
# ------------------------------------------------------- model calls
|
|
156
|
+
def _started(self, tag: str):
|
|
157
|
+
if self.on_call is not None:
|
|
158
|
+
try:
|
|
159
|
+
self.on_call(tag)
|
|
160
|
+
except Exception as e: # a progress callback must not break the pipeline
|
|
161
|
+
log.debug("on_call failed: %s", e)
|
|
162
|
+
|
|
163
|
+
async def generate(self, req: GenRequest, tag: str = "generate"):
|
|
164
|
+
self._started(tag)
|
|
165
|
+
t0 = time.perf_counter()
|
|
166
|
+
res = await asyncio.wrap_future(self.backend.submit_generate(req))
|
|
167
|
+
if self.trace is not None:
|
|
168
|
+
self.trace.append((tag, time.perf_counter() - t0, res.num_tokens))
|
|
169
|
+
return res
|
|
170
|
+
|
|
171
|
+
async def score(self, req: ScoreRequest):
|
|
172
|
+
return await asyncio.wrap_future(self.backend.submit_score(req))
|
|
173
|
+
|
|
174
|
+
async def generate_image(self, prompt: str, size: int) -> Optional[Image.Image]:
|
|
175
|
+
target = self.backend if self.backend.native_image_generation else self.image_generator
|
|
176
|
+
if target is None:
|
|
177
|
+
return None
|
|
178
|
+
self._started("image")
|
|
179
|
+
try:
|
|
180
|
+
t0 = time.perf_counter()
|
|
181
|
+
img = await asyncio.wrap_future(target.submit_image(
|
|
182
|
+
ImageRequest(prompt=prompt, width=size, height=size, seed=self._next_seed())))
|
|
183
|
+
if self.trace is not None:
|
|
184
|
+
self.trace.append(("image", time.perf_counter() - t0, 0))
|
|
185
|
+
return img
|
|
186
|
+
except Exception as e:
|
|
187
|
+
log.debug("image generation failed: %s", e)
|
|
188
|
+
return None
|
|
189
|
+
|
|
190
|
+
def _next_seed(self):
|
|
191
|
+
if self.config.seed is None:
|
|
192
|
+
return None
|
|
193
|
+
self._seed = getattr(self, "_seed", self.config.seed) + 1
|
|
194
|
+
return self._seed
|
|
195
|
+
|
|
196
|
+
# --------------------------------------------------- answering (ICL)
|
|
197
|
+
async def zero_shot(self, image, question: str, max_new_tokens: int) -> ZeroShot:
|
|
198
|
+
res = await self.generate(GenRequest(parts=[image, question], style="answer",
|
|
199
|
+
max_new_tokens=max_new_tokens, logprobs=True), tag="zero_shot")
|
|
200
|
+
return ZeroShot(res.text, _mean_prob(res.token_logprobs), res.token_logprobs or ())
|
|
201
|
+
|
|
202
|
+
async def answer(self, image, question: str, demos: Sequence[Demo], max_new_tokens: int) -> str:
|
|
203
|
+
res = await self.generate(GenRequest(parts=icl_parts(demos, image, question), style="answer",
|
|
204
|
+
max_new_tokens=max_new_tokens), tag="answer")
|
|
205
|
+
return res.text
|
|
206
|
+
|
|
207
|
+
async def confidence(self, image, question: str, answer: str) -> Optional[float]:
|
|
208
|
+
lp = await self.score(ScoreRequest(parts=[image, question], target=answer, style="answer"))
|
|
209
|
+
return _mean_prob(lp)
|
|
210
|
+
|
|
211
|
+
# ---------------------------------------------------- triplet synthesis
|
|
212
|
+
@staticmethod
|
|
213
|
+
def _parse_triplets(text: str, final: bool = True) -> list[tuple[str, str, str]]:
|
|
214
|
+
text = clean_response(text).replace("**", "")
|
|
215
|
+
found = []
|
|
216
|
+
for m in P.TRIPLET_RE.finditer(text):
|
|
217
|
+
if not final and m.end() >= len(text.rstrip()):
|
|
218
|
+
break # may still be growing
|
|
219
|
+
q, a, d = m.groups()
|
|
220
|
+
lines = [ln.strip() for ln in q.strip().splitlines() if ln.strip()]
|
|
221
|
+
if len(lines) > 1 and lines[0].endswith("?"):
|
|
222
|
+
q = lines[0] # stray text (often the answer) leaked onto the question's next line
|
|
223
|
+
sep = d.find("\n---")
|
|
224
|
+
found.append(tuple(_unbracket(x) for x in (q, a, d[:sep] if sep != -1 else d)))
|
|
225
|
+
return found
|
|
226
|
+
|
|
227
|
+
async def stream_triplets(self, image, k: int, push, max_attempts: int = 3,
|
|
228
|
+
temperatures: tuple = (0.2, 0.7), tag: str = "synthesis") -> int:
|
|
229
|
+
"""Like ``triplets`` but hands each triplet to ``push`` (which returns
|
|
230
|
+
True if it was kept) as soon as it is complete in the decoded stream,
|
|
231
|
+
so realization starts early; decoding stops once ``k`` were kept."""
|
|
232
|
+
loop = asyncio.get_running_loop()
|
|
233
|
+
base = P.TRIPLET_PROMPT if self.config.diversity_prompt else P.TRIPLET_PROMPT_NO_DIVERSITY
|
|
234
|
+
prompt = base.format(K=k)
|
|
235
|
+
kept = 0
|
|
236
|
+
for attempt in range(max_attempts):
|
|
237
|
+
seen = [0]
|
|
238
|
+
call: list = []
|
|
239
|
+
|
|
240
|
+
def take(found):
|
|
241
|
+
nonlocal kept
|
|
242
|
+
for t in found[seen[0]:]:
|
|
243
|
+
kept += bool(push(t))
|
|
244
|
+
seen[0] = max(seen[0], len(found))
|
|
245
|
+
if kept >= k and call and not call[0].done():
|
|
246
|
+
call[0].cancel() # enough triplets: stop decoding the rest
|
|
247
|
+
|
|
248
|
+
req = GenRequest(parts=[image, prompt], style="chat", max_new_tokens=1800,
|
|
249
|
+
temperature=temperatures[min(attempt, len(temperatures) - 1)],
|
|
250
|
+
on_text=lambda text: loop.call_soon_threadsafe(
|
|
251
|
+
lambda: take(self._parse_triplets(text, final=False))))
|
|
252
|
+
call.append(asyncio.ensure_future(self.generate(req, tag=tag)))
|
|
253
|
+
try:
|
|
254
|
+
res = await call[0]
|
|
255
|
+
except asyncio.CancelledError:
|
|
256
|
+
if kept >= k:
|
|
257
|
+
return kept
|
|
258
|
+
raise
|
|
259
|
+
take(self._parse_triplets(res.text))
|
|
260
|
+
if kept >= k:
|
|
261
|
+
return kept
|
|
262
|
+
prompt = base.format(K=k) + P.FORMAT_REMINDER.format(K=k)
|
|
263
|
+
return kept
|
|
264
|
+
|
|
265
|
+
async def decomposed_triplets(self, image, k: int, push, max_attempts: Optional[int] = None) -> int:
|
|
266
|
+
"""Three short calls per triplet (question -> answer -> description),
|
|
267
|
+
for models that cannot follow the one-shot K-triplet prompt."""
|
|
268
|
+
budget = [max_attempts or k * 6]
|
|
269
|
+
kept = [0]
|
|
270
|
+
counter = [0]
|
|
271
|
+
|
|
272
|
+
def clean(text: str) -> str:
|
|
273
|
+
text = clean_response(text).strip()
|
|
274
|
+
if not text:
|
|
275
|
+
return ""
|
|
276
|
+
line = text.splitlines()[0].strip().strip('"')
|
|
277
|
+
line = re.sub(r"^[-*•]\s*", "", line) # bullets copied from the few-shot list
|
|
278
|
+
line = re.sub(r"^question:\s*", "", line, flags=re.IGNORECASE)
|
|
279
|
+
# chat preambles: "Sure, here's the first one: How many ...?"
|
|
280
|
+
line = re.sub(r"^(sure|okay|ok|here(?:'s| is| are))\b[^:?]{0,60}:\s*", "", line, flags=re.IGNORECASE)
|
|
281
|
+
return re.sub(r"^\d+[.)]\s*", "", line)
|
|
282
|
+
|
|
283
|
+
async def ask(prompt, n):
|
|
284
|
+
res = await self.generate(GenRequest(parts=[image, prompt], style="chat", max_new_tokens=n,
|
|
285
|
+
temperature=0.8), tag="synthesis")
|
|
286
|
+
return res.text
|
|
287
|
+
|
|
288
|
+
async def worker():
|
|
289
|
+
while kept[0] < k and budget[0] > 0:
|
|
290
|
+
budget[0] -= 1
|
|
291
|
+
topic = P.DECOMPOSED_TOPICS[counter[0] % len(P.DECOMPOSED_TOPICS)]
|
|
292
|
+
counter[0] += 1
|
|
293
|
+
q = clean(await ask(P.DECOMPOSED_QUESTION_PROMPT.format(topic=topic), 40))
|
|
294
|
+
if (len(q.split()) < 3 or not q.endswith("?") or P.COORD_LIST_RE.match(q)
|
|
295
|
+
or P.META_LEAK_RE.search(q)):
|
|
296
|
+
continue
|
|
297
|
+
hint = next((h for pat, h in P.ANSWER_TYPE_HINTS if pat.match(q)), "")
|
|
298
|
+
a = clean(await ask(P.DECOMPOSED_ANSWER_PROMPT.format(question=q, type_hint=hint), 16))
|
|
299
|
+
if not a:
|
|
300
|
+
continue
|
|
301
|
+
d = (await ask(P.DECOMPOSED_DESCRIPTION_PROMPT.format(question=q, answer=a), 150)).strip()
|
|
302
|
+
if len(d) < 20 or P.COORD_LIST_RE.match(d):
|
|
303
|
+
continue
|
|
304
|
+
if kept[0] < k:
|
|
305
|
+
kept[0] += bool(push((q, a, d)))
|
|
306
|
+
|
|
307
|
+
await asyncio.gather(*[worker() for _ in range(k)])
|
|
308
|
+
return kept[0]
|
|
309
|
+
|
|
310
|
+
async def triplets(self, image, k: int, max_attempts: int = 3,
|
|
311
|
+
temperatures: tuple = (0.2, 0.7)) -> list[tuple[str, str, str]]:
|
|
312
|
+
base = P.TRIPLET_PROMPT if self.config.diversity_prompt else P.TRIPLET_PROMPT_NO_DIVERSITY
|
|
313
|
+
prompt = base.format(K=k)
|
|
314
|
+
best: list = []
|
|
315
|
+
for attempt in range(max_attempts):
|
|
316
|
+
res = await self.generate(GenRequest(parts=[image, prompt], style="chat", max_new_tokens=1800,
|
|
317
|
+
temperature=temperatures[min(attempt, len(temperatures) - 1)]),
|
|
318
|
+
tag="synthesis_fresh")
|
|
319
|
+
found = self._parse_triplets(res.text)
|
|
320
|
+
if len(found) > len(best):
|
|
321
|
+
best = found
|
|
322
|
+
if len(best) >= k:
|
|
323
|
+
break
|
|
324
|
+
prompt = base.format(K=k) + P.FORMAT_REMINDER.format(K=k)
|
|
325
|
+
return best
|
|
326
|
+
|
|
327
|
+
# ------------------------------------------------------------- routing
|
|
328
|
+
async def route(self, description: str, allow_composite: bool = True) -> Skill:
|
|
329
|
+
if not self.use_skills:
|
|
330
|
+
return self.by_name["natural"]
|
|
331
|
+
skills = [s for s in self.skills if allow_composite or s.name != "html_composite"]
|
|
332
|
+
name = pre_route(description, skills)
|
|
333
|
+
if name is None:
|
|
334
|
+
res = await self.generate(GenRequest(parts=[self.router_prompt.replace("{request}", description)],
|
|
335
|
+
style="chat", max_new_tokens=20), tag="route")
|
|
336
|
+
default = "figure" if "figure" in self.by_name else skills[0].name
|
|
337
|
+
name = parse_route(clean_response(res.text), skills, default)
|
|
338
|
+
return self.by_name[name]
|
|
339
|
+
|
|
340
|
+
# -------------------------------------------------------- verification
|
|
341
|
+
async def critic(self, description: str, image: Image.Image, rubric: str):
|
|
342
|
+
template = P.CRITIC_STRICT_PROMPT if rubric == "strict" else P.CRITIC_STRUCTURED_PROMPT
|
|
343
|
+
res = await self.generate(GenRequest(parts=[to_pil(image), template.format(request=description)],
|
|
344
|
+
style="chat", max_new_tokens=self.config.verify_max_new_tokens,
|
|
345
|
+
think=self.verify_think), tag="critic")
|
|
346
|
+
score = _parse_score(res.text)
|
|
347
|
+
if score is None and res.raw_text: # thinking cut short before the answer: use what it wrote
|
|
348
|
+
score = _parse_score(res.raw_text)
|
|
349
|
+
if score is None and res.num_tokens >= self.config.verify_max_new_tokens:
|
|
350
|
+
# A verbose critic ran out of tokens before its verdict: ask for the score line only.
|
|
351
|
+
critique = (res.text or res.raw_text).strip()
|
|
352
|
+
follow = (template.format(request=description) + "\n\nYour review so far:\n" + critique[-3000:]
|
|
353
|
+
+ "\n\nNow output only the final line: 'SCORE: <integer 0-100>'.")
|
|
354
|
+
res2 = await self.generate(GenRequest(parts=[to_pil(image), follow], style="chat", max_new_tokens=16),
|
|
355
|
+
tag="critic_score")
|
|
356
|
+
score = _parse_score(res2.text)
|
|
357
|
+
return score, res.text
|
|
358
|
+
|
|
359
|
+
async def revise(self, request: str, critique: str) -> str:
|
|
360
|
+
res = await self.generate(GenRequest(parts=[P.REVISE_PROMPT.format(request=request, critique=critique)],
|
|
361
|
+
style="chat", max_new_tokens=200, think=self.verify_think), tag="revise")
|
|
362
|
+
return res.text.strip() or request
|
|
363
|
+
|
|
364
|
+
async def realize_verified(self, triplet, single_shot: bool) -> Optional[Demo]:
|
|
365
|
+
"""Route, realize and verify one triplet. The image is always judged
|
|
366
|
+
against the original description; failed rounds revise the
|
|
367
|
+
description used for the next generation."""
|
|
368
|
+
question, answer, description = triplet
|
|
369
|
+
cfg = self.config
|
|
370
|
+
skill = await self.route(description)
|
|
371
|
+
ctx = SkillContext(self, cfg.image_size)
|
|
372
|
+
retries = 0 if single_shot else cfg.repair_retries
|
|
373
|
+
rounds = 1 if single_shot else (cfg.verify_rounds if self.verify_enabled else 1)
|
|
374
|
+
request, failures = description, 0
|
|
375
|
+
for rnd in range(rounds):
|
|
376
|
+
try:
|
|
377
|
+
real = await skill.realize(ctx, request, retries=retries)
|
|
378
|
+
except Exception as e:
|
|
379
|
+
log.debug("realization error (%s): %s", skill.name, e)
|
|
380
|
+
real = None
|
|
381
|
+
if real is None:
|
|
382
|
+
failures += 1
|
|
383
|
+
if failures >= 2:
|
|
384
|
+
break
|
|
385
|
+
continue
|
|
386
|
+
failures = 0
|
|
387
|
+
score = None
|
|
388
|
+
if self.verify_enabled:
|
|
389
|
+
score, critique = await self.critic(description, real.image, skill.verify_rubric)
|
|
390
|
+
if score is None or score < cfg.verify_threshold:
|
|
391
|
+
if rnd < rounds - 1:
|
|
392
|
+
request = await self.revise(request, critique)
|
|
393
|
+
continue
|
|
394
|
+
return Demo(image=to_pil(real.image), question=question, answer=answer, description=description,
|
|
395
|
+
skill=real.skill, verify_score=score, spec=real.spec_text)
|
|
396
|
+
return None
|
|
397
|
+
|
|
398
|
+
# -------------------------------------------------------------- imagine
|
|
399
|
+
async def imagine(self, image, question: str, *, k: Optional[int] = None, fill: bool = False,
|
|
400
|
+
max_new_tokens: Optional[int] = None, zero_shot: Optional[ZeroShot] = None,
|
|
401
|
+
speculative: bool = False, time_limit: Optional[float] = None,
|
|
402
|
+
on_demo: Optional[Callable[[Demo], None]] = None) -> Imagination:
|
|
403
|
+
"""Imagine demonstrations for one query.
|
|
404
|
+
|
|
405
|
+
* default: adaptive budget K*(p0) and difficulty filtering (online rule:
|
|
406
|
+
stop once K* candidates pass DF or K_max candidates exist);
|
|
407
|
+
* ``k=n``: exactly ``n`` verified demos, no ABA/DF;
|
|
408
|
+
* ``fill=True``: K_max verified candidates, all kept (tuning cache).
|
|
409
|
+
|
|
410
|
+
``time_limit`` (seconds) stops synthesis and returns the demos accepted
|
|
411
|
+
so far; ``on_demo`` is called with each accepted demo as it arrives."""
|
|
412
|
+
cfg = self.config
|
|
413
|
+
image = to_pil(image)
|
|
414
|
+
mnt = max_new_tokens or cfg.answer_max_new_tokens
|
|
415
|
+
zs = zero_shot or await self.zero_shot(image, question, mnt)
|
|
416
|
+
if fill:
|
|
417
|
+
target, cap, accept = cfg.k_max, cfg.k_max, (lambda c: True)
|
|
418
|
+
budget = cfg.k_max
|
|
419
|
+
elif k is not None:
|
|
420
|
+
target, cap, accept = k, k, (lambda c: True)
|
|
421
|
+
budget = k
|
|
422
|
+
else:
|
|
423
|
+
budget = cfg.budget(zs.confidence)
|
|
424
|
+
target, cap, accept = budget, cfg.k_max, cfg.in_band
|
|
425
|
+
if target <= 0:
|
|
426
|
+
return Imagination(demos=[], zero_shot=zs, budget=budget)
|
|
427
|
+
need_conf = fill or (cfg.use_df and k is None)
|
|
428
|
+
demos, cands, stats = await self._fill_slots(image, target, cap, accept, need_conf, speculative,
|
|
429
|
+
time_limit, on_demo)
|
|
430
|
+
return Imagination(demos=demos, zero_shot=zs, budget=budget, candidates=cands, stats=stats)
|
|
431
|
+
|
|
432
|
+
async def _fill_slots(self, image, target, cap, accept, need_conf, speculative=False, time_limit=None,
|
|
433
|
+
on_demo=None):
|
|
434
|
+
"""Breadth-first slot filling (Appendix G.1): every fresh triplet gets
|
|
435
|
+
one cheap attempt before any failure gets the full repair budget."""
|
|
436
|
+
cfg = self.config
|
|
437
|
+
seen: set = set()
|
|
438
|
+
|
|
439
|
+
def novel(triplet) -> bool:
|
|
440
|
+
key = _norm_q(triplet[0])
|
|
441
|
+
if key in seen:
|
|
442
|
+
return False
|
|
443
|
+
seen.add(key)
|
|
444
|
+
return True
|
|
445
|
+
|
|
446
|
+
queue: deque = deque()
|
|
447
|
+
arrived = asyncio.Event()
|
|
448
|
+
|
|
449
|
+
def push(triplet) -> bool:
|
|
450
|
+
if novel(triplet):
|
|
451
|
+
queue.append(triplet)
|
|
452
|
+
arrived.set()
|
|
453
|
+
return True
|
|
454
|
+
return False
|
|
455
|
+
|
|
456
|
+
decomposed = self.option("synthesis", "batch") == "decomposed"
|
|
457
|
+
synth = asyncio.ensure_future(self.decomposed_triplets(image, target, push) if decomposed
|
|
458
|
+
else self.stream_triplets(image, target, push))
|
|
459
|
+
attempts_left = cfg.attempts_per_slot * target
|
|
460
|
+
phase1_left = max(target, attempts_left - target)
|
|
461
|
+
retry_pool: deque = deque()
|
|
462
|
+
accepted: list[Demo] = []
|
|
463
|
+
produced: list[Demo] = []
|
|
464
|
+
running: dict = {}
|
|
465
|
+
stats = {"attempts": 0, "fresh_triplet_calls": 0, "verified": 0, "df_rejected": 0}
|
|
466
|
+
loop = asyncio.get_running_loop()
|
|
467
|
+
deadline = loop.time() + time_limit if time_limit is not None else None
|
|
468
|
+
|
|
469
|
+
async def attempt(triplet, single_shot):
|
|
470
|
+
if triplet is None:
|
|
471
|
+
# One new triplet at a time (as in the paper), sampled so that it
|
|
472
|
+
# does not repeat a question this query already has.
|
|
473
|
+
got: list = []
|
|
474
|
+
for _ in range(3):
|
|
475
|
+
stats["fresh_triplet_calls"] += 1
|
|
476
|
+
keep = lambda t: novel(t) and (got.append(t) or True) # noqa: E731
|
|
477
|
+
if decomposed:
|
|
478
|
+
await self.decomposed_triplets(image, 1, keep, max_attempts=2)
|
|
479
|
+
else:
|
|
480
|
+
await self.stream_triplets(image, 1, keep, max_attempts=1, temperatures=(0.7,),
|
|
481
|
+
tag="synthesis_fresh")
|
|
482
|
+
if got:
|
|
483
|
+
triplet = got[0]
|
|
484
|
+
break
|
|
485
|
+
else:
|
|
486
|
+
return None, None
|
|
487
|
+
demo = await self.realize_verified(triplet, single_shot)
|
|
488
|
+
if demo is not None and need_conf:
|
|
489
|
+
demo.confidence = await self.confidence(demo.image, demo.question, demo.answer)
|
|
490
|
+
return triplet, demo
|
|
491
|
+
|
|
492
|
+
def open_slots():
|
|
493
|
+
if speculative: # every candidate that may still be needed, in parallel
|
|
494
|
+
return cap - len(produced) - len(running)
|
|
495
|
+
return min(target - len(accepted), cap - len(produced)) - len(running)
|
|
496
|
+
|
|
497
|
+
while True:
|
|
498
|
+
if deadline is not None and loop.time() >= deadline:
|
|
499
|
+
stats["timed_out"] = True
|
|
500
|
+
break
|
|
501
|
+
while open_slots() > 0 and attempts_left > 0:
|
|
502
|
+
if phase1_left > 0:
|
|
503
|
+
if queue:
|
|
504
|
+
triplet = queue.popleft()
|
|
505
|
+
elif not synth.done():
|
|
506
|
+
break # more triplets are being decoded right now
|
|
507
|
+
else:
|
|
508
|
+
triplet = None # request a fresh one
|
|
509
|
+
single = True
|
|
510
|
+
phase1_left -= 1
|
|
511
|
+
elif retry_pool:
|
|
512
|
+
triplet, single = retry_pool.popleft(), False
|
|
513
|
+
else:
|
|
514
|
+
break
|
|
515
|
+
attempts_left -= 1
|
|
516
|
+
stats["attempts"] += 1
|
|
517
|
+
task = asyncio.ensure_future(attempt(triplet, single))
|
|
518
|
+
running[task] = triplet
|
|
519
|
+
waiters = list(running)
|
|
520
|
+
arrival = None
|
|
521
|
+
if not synth.done():
|
|
522
|
+
arrived.clear()
|
|
523
|
+
arrival = asyncio.ensure_future(arrived.wait())
|
|
524
|
+
waiters += [synth, arrival]
|
|
525
|
+
if not waiters:
|
|
526
|
+
break
|
|
527
|
+
timeout = max(0.0, deadline - loop.time()) if deadline is not None else None
|
|
528
|
+
done, _ = await asyncio.wait(waiters, timeout=timeout, return_when=asyncio.FIRST_COMPLETED)
|
|
529
|
+
if arrival is not None and not arrival.done():
|
|
530
|
+
arrival.cancel()
|
|
531
|
+
for task in done:
|
|
532
|
+
if task not in running:
|
|
533
|
+
continue
|
|
534
|
+
running.pop(task)
|
|
535
|
+
try:
|
|
536
|
+
triplet, demo = task.result()
|
|
537
|
+
except Exception as e:
|
|
538
|
+
log.debug("attempt crashed: %s", e)
|
|
539
|
+
triplet, demo = None, None
|
|
540
|
+
if demo is None:
|
|
541
|
+
if triplet is not None:
|
|
542
|
+
retry_pool.append(triplet)
|
|
543
|
+
continue
|
|
544
|
+
stats["verified"] += 1
|
|
545
|
+
produced.append(demo)
|
|
546
|
+
if accept(demo.confidence):
|
|
547
|
+
accepted.append(demo)
|
|
548
|
+
if on_demo is not None and len(accepted) <= target:
|
|
549
|
+
try:
|
|
550
|
+
on_demo(demo)
|
|
551
|
+
except Exception as e: # a UI callback must not break synthesis
|
|
552
|
+
log.debug("on_demo failed: %s", e)
|
|
553
|
+
else:
|
|
554
|
+
stats["df_rejected"] += 1
|
|
555
|
+
# a filtered candidate reopens its slot with a fresh budget
|
|
556
|
+
attempts_left += cfg.attempts_per_slot
|
|
557
|
+
phase1_left += cfg.attempts_per_slot - 1
|
|
558
|
+
if len(accepted) >= target or len(produced) >= cap:
|
|
559
|
+
break
|
|
560
|
+
for task in running:
|
|
561
|
+
task.cancel()
|
|
562
|
+
if not synth.done():
|
|
563
|
+
synth.cancel()
|
|
564
|
+
return accepted[:target], produced, stats
|
|
565
|
+
|
|
566
|
+
|
|
567
|
+
def _unbracket(field: str) -> str:
|
|
568
|
+
"""Drop the "[...]" the prompt's template uses around a whole field
|
|
569
|
+
(models sometimes copy it: "Answer: [3]")."""
|
|
570
|
+
field = field.strip()
|
|
571
|
+
if len(field) > 2 and field[0] == "[" and field[-1] == "]" and "[" not in field[1:-1] and "]" not in field[1:-1]:
|
|
572
|
+
return field[1:-1].strip()
|
|
573
|
+
return field
|
|
574
|
+
|
|
575
|
+
|
|
576
|
+
def icl_parts(demos: Sequence[Demo], image, question: str) -> list:
|
|
577
|
+
"""The in-context prompt of Appendix G (SIMIT-ICL)."""
|
|
578
|
+
image = to_pil(image)
|
|
579
|
+
if not demos:
|
|
580
|
+
return [image, question]
|
|
581
|
+
n = len(demos)
|
|
582
|
+
parts: list = [f"Here {'is' if n == 1 else 'are'} {n} example{'s' if n > 1 else ''} to help you answer the question:\n"]
|
|
583
|
+
for i, d in enumerate(demos, 1):
|
|
584
|
+
parts += [f"\nExample {i}:", to_pil(d.image), f"Question: {d.question}\nAnswer: {d.answer}"]
|
|
585
|
+
parts += ["\nNow, answer this question:", image, question]
|
|
586
|
+
return parts
|