iolai26-solve / solver /synth.py
rvpant
Single-shot-with-tools v1: scaffolded prompts, shared pipeline, explanation track
fc85791
Raw
History Blame Contribute Delete
4.51 kB
"""Program-synthesis subagent: LLM proposes grammars in the DSL, the
interpreter executes them, the verifier scores them, and failing pairs are
fed back for refinement (CEGIS), up to R rounds.
The LLM never applies rules — it only emits grammar JSON. All execution is
Interpreter; all selection is verifier.evaluate on the attested pairs.
"""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from typing import List, Optional, Sequence, Tuple
from .dsl.grammar import Grammar, from_json
from .dsl.interpreter import Interpreter
from .llm import LLMClient
from .preprocess import Pair, Puzzle
from .scaffold import analysis_blocks
from .verifier import Verdict, evaluate
PROMPT_DIR = Path(__file__).resolve().parents[1] / "prompts"
MAX_FAILURES_SHOWN = 8
@dataclass
class SynthResult:
grammar: Optional[Grammar]
interpreter: Optional[Interpreter]
verdict: Optional[Verdict]
rounds_used: int = 0
def _load_prompt(name: str) -> str:
return (PROMPT_DIR / f"{name}.md").read_text(encoding="utf-8")
def proposer_prompt(puzzle: Puzzle) -> str:
seg_block, align_block = analysis_blocks(puzzle)
pairs_block = "\n".join(f" {p.src} = {p.tgt}" for p in puzzle.pairs) or " (none)"
hints_block = "\n".join(f" {h}" for h in puzzle.hints) or " (none)"
return _load_prompt("proposer").format(
task_lang=puzzle.task_lang or "the unknown language",
work_lang=puzzle.work_lang or "English",
pairs_block=pairs_block,
hints_block=hints_block,
segmentation_block=seg_block,
alignment_block=align_block,
)
def refine_prompt(puzzle: Puzzle, grammar: Grammar, verdict: Verdict) -> str:
fails = verdict.failures[:MAX_FAILURES_SHOWN]
failures_block = "\n".join(
f" input: {src}\n expected: {gold}\n got: {pred or '(nothing)'}"
for src, gold, pred in fails
)
return _load_prompt("refine").format(
task_lang=puzzle.task_lang or "the unknown language",
failures_block=failures_block,
previous_grammar=grammar.to_json(),
)
def _attested_for_direction(pairs: Sequence[Pair], direction: str) -> List[Tuple[str, str]]:
if direction == "to_task":
return [(p.tgt, p.src) for p in pairs] # work -> task (generation)
return [(p.src, p.tgt) for p in pairs] # task -> work (analysis)
def _predictor(interp: Interpreter, direction: str):
return interp.generate if direction == "to_task" else interp.analyze
def score_grammar(g: Grammar, pairs: Sequence[Pair], direction: str) -> Tuple[Interpreter, Verdict]:
interp = Interpreter(g)
attested = _attested_for_direction(pairs, direction)
return interp, evaluate(_predictor(interp, direction), attested, g.mdl())
def synthesize(
puzzle: Puzzle,
client: LLMClient,
direction: str = "to_task",
rounds: int = 2,
) -> SynthResult:
"""CEGIS loop: propose -> execute -> verify -> refine on failures.
Returns the best grammar seen across rounds (never a later-worse one)."""
if not client.available or not puzzle.pairs:
return SynthResult(None, None, None, 0)
best: SynthResult = SynthResult(None, None, None, 0)
prompt = proposer_prompt(puzzle)
for r in range(rounds + 1):
text = client.generate([prompt])[0]
g = from_json(text)
if g is None:
break
interp, verdict = score_grammar(g, puzzle.pairs, direction)
if best.verdict is None or verdict.score > best.verdict.score:
best = SynthResult(g, interp, verdict, r + 1)
if verdict.em >= 1.0 or r == rounds:
break
prompt = refine_prompt(puzzle, g, verdict)
return best
def synthesize_best_of_n(
puzzle: Puzzle,
client: LLMClient,
direction: str,
n: int,
rounds: int = 1,
) -> SynthResult:
"""Phase-3 test-time scaling hook: N independent proposals (greedy base is
deterministic, so diversity must come from prompt variants), each with a
short CEGIS budget; verifier picks. With greedy decoding, n>1 only helps
once prompt variants or sampling adapters exist — the plumbing is here."""
best = SynthResult(None, None, None, 0)
for _ in range(max(1, n)):
r = synthesize(puzzle, client, direction, rounds)
if r.verdict and (best.verdict is None or r.verdict.score > best.verdict.score):
best = r
if best.verdict and best.verdict.em >= 1.0:
break
return best