初始化项目,由ModelHub XC社区提供模型
Model: rpant/iolai26-solve Source: Original Platform
This commit is contained in:
126
solver/synth.py
Normal file
126
solver/synth.py
Normal file
@@ -0,0 +1,126 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user