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/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