初始化项目,由ModelHub XC社区提供模型
Model: rpant/iolai26-solve Source: Original Platform
This commit is contained in:
382
solver/pipeline.py
Normal file
382
solver/pipeline.py
Normal file
@@ -0,0 +1,382 @@
|
||||
"""Pipeline (shared by script.py and eval/dev_harness.py): symbolic pass, then
|
||||
LLM answering (lean or scaffolded), merge, never-empty guarantee."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, List, Optional, Sequence
|
||||
|
||||
from .budget import Budget
|
||||
from .direct import align_answers, solve_direct
|
||||
from .llm import LLMClient, NullClient
|
||||
from .matching import solve_match_letters_llm
|
||||
from .preprocess import Puzzle, parse_puzzle
|
||||
from .router import solve_puzzle_ex
|
||||
from .scaffold import build_scaffold, light_hint
|
||||
|
||||
CONF_KEEP = 0.5 # symbolic answers at/above this verifier fit override the LLM
|
||||
LLM_BATCH = 4 # puzzles per generation batch (T4 KV-cache friendly)
|
||||
DEEP_MODE_MAX_ROWS = 30 # at/below this many puzzles, spend more per puzzle
|
||||
|
||||
_DIGITS_RX = re.compile(r"\d[\d,. ]*")
|
||||
_LETTER_RX = re.compile(r"\b([A-Z])\b")
|
||||
_PREAMBLE_RX = re.compile(
|
||||
r"^(?:the\s+)?(?:answer|translation|result)\s*(?:is|:)\s*", re.IGNORECASE)
|
||||
_TERMINAL_PUNCT = (".", "!", "?")
|
||||
|
||||
|
||||
@dataclass
|
||||
class PuzzleResult:
|
||||
row_id: str
|
||||
answers: List[str]
|
||||
explanation: str = ""
|
||||
confs: List[float] = field(default_factory=list)
|
||||
methods: List[str] = field(default_factory=list)
|
||||
llm_used: bool = False
|
||||
raw: str = ""
|
||||
|
||||
|
||||
def clean_llm_answer(ans: str, task_type: str, option_labels: Sequence[str] = ()) -> str:
|
||||
"""Per-task-type format guard on a parsed LLM answer line. Conservative:
|
||||
only rewrites when the expected shape is unambiguous."""
|
||||
a = _PREAMBLE_RX.sub("", ans.strip()).strip().strip("'\"“”")
|
||||
if task_type == "text_to_num":
|
||||
m = _DIGITS_RX.search(a)
|
||||
if m:
|
||||
digits = re.sub(r"[,. ]", "", m.group(0))
|
||||
if digits.isdigit():
|
||||
return digits
|
||||
elif task_type == "match_letters":
|
||||
if len(a) > 2: # "B. the bird sleeps" or "option B" -> "B"
|
||||
candidates = _LETTER_RX.findall(a)
|
||||
wanted = [c for c in candidates if not option_labels or c in option_labels]
|
||||
if len(set(wanted)) == 1:
|
||||
return wanted[0]
|
||||
return a
|
||||
|
||||
|
||||
def _vote(cands: Sequence[Optional[str]], anchor: Optional[str]) -> Optional[str]:
|
||||
"""Greedy-anchored vote: keep `anchor` (the greedy answer) unless at least
|
||||
two sampled candidates agree on the same normalised form AND that form
|
||||
outnumbers the anchor's support. Monotone — it can only fire on genuine
|
||||
agreement, so it never replaces greedy with a lone sample."""
|
||||
from collections import Counter
|
||||
cands = [c for c in cands if c and str(c).strip()]
|
||||
if anchor is None:
|
||||
anchor = cands[0] if cands else None
|
||||
if len(cands) < 3:
|
||||
return anchor
|
||||
norm = lambda s: re.sub(r"\s+", " ", str(s).strip().lower())
|
||||
groups: dict = {}
|
||||
for c in cands:
|
||||
groups.setdefault(norm(c), []).append(c)
|
||||
anchor_support = len(groups.get(norm(anchor), [])) if anchor else 0
|
||||
best_key = max(groups, key=lambda k: len(groups[k]))
|
||||
if len(groups[best_key]) >= 2 and len(groups[best_key]) > anchor_support:
|
||||
return Counter(groups[best_key]).most_common(1)[0][0]
|
||||
return anchor
|
||||
|
||||
|
||||
def induce_format(answer: str, puzzle: Puzzle, item_idx: int) -> str:
|
||||
"""Nudge an answer toward the dataset's surface convention (fable §3):
|
||||
if the attested answers on this item's side overwhelmingly end in a
|
||||
terminal punctuation mark or start with a capital, mirror that. Converts
|
||||
chrF-close answers into EM hits, which the geometric-mean scoring rewards
|
||||
twice. Conservative: only fires on a near-unanimous (>=85%) convention,
|
||||
only ADDS a missing terminal mark or leading capital, never strips."""
|
||||
if puzzle.task_type not in ("translation", "fill_blanks"):
|
||||
return answer
|
||||
a = answer.strip()
|
||||
if not a:
|
||||
return answer
|
||||
it = puzzle.items[item_idx] if item_idx < len(puzzle.items) else None
|
||||
direction = getattr(it, "direction", None)
|
||||
# answer side: to_work -> work-language (tgt); else task-language (src)
|
||||
side = [p.tgt for p in puzzle.pairs] if direction == "to_work" \
|
||||
else [p.src for p in puzzle.pairs]
|
||||
side = [s.strip() for s in side if s and s.strip()]
|
||||
if len(side) < 4:
|
||||
return answer
|
||||
n = len(side)
|
||||
# terminal punctuation: only if a single mark dominates
|
||||
for mark in _TERMINAL_PUNCT:
|
||||
if sum(1 for s in side if s.endswith(mark)) / n >= 0.85:
|
||||
if not a.endswith(_TERMINAL_PUNCT):
|
||||
a = a + mark
|
||||
break
|
||||
# leading capitalization
|
||||
if sum(1 for s in side if s[:1].isupper()) / n >= 0.85:
|
||||
if a[:1].islower():
|
||||
a = a[:1].upper() + a[1:]
|
||||
return a
|
||||
|
||||
|
||||
def _symbolic_explanation(puzzle: Puzzle, methods: Sequence[str]) -> str:
|
||||
used = [m for m in dict.fromkeys(methods) if m not in ("none", "fallback")]
|
||||
if not used:
|
||||
return ("- Answered by nearest-attested analogy over the given examples "
|
||||
"(no reliable rule could be verified).")
|
||||
tmpl = {
|
||||
"numeral system (verified)": (
|
||||
"- Induced each morpheme's numeric value from the attested numerals "
|
||||
"and verified the system reproduces every given example; applied it "
|
||||
"to each query item (smaller-before-larger multiplies, otherwise "
|
||||
"values add)."),
|
||||
"table completion": (
|
||||
"- Learned the mapping between the paradigm-table columns from the "
|
||||
"attested rows (leave-one-out verified) and applied it to each "
|
||||
"incomplete row."),
|
||||
"template substitution": (
|
||||
"- For each query, took the closest attested sentence and swapped "
|
||||
"the differing words through morpheme alignments induced from "
|
||||
"minimal pairs in the data."),
|
||||
"induced grammar": (
|
||||
"- Induced a lexicon and affix rules that reproduce the attested "
|
||||
"pairs exactly, then applied them mechanically to the query items."),
|
||||
"optimal matching": (
|
||||
"- Scored every form-meaning pair by shared-morpheme/shared-word "
|
||||
"consistency and picked the globally optimal assignment."),
|
||||
}
|
||||
return "\n".join(tmpl.get(m, f"- Solved by {m}.") for m in used)
|
||||
|
||||
|
||||
def run_pipeline(rows: Sequence[dict], client: Optional[LLMClient] = None,
|
||||
budget: Optional[Budget] = None, verbose: bool = True,
|
||||
conf_keep: float = CONF_KEEP, llm_batch: int = LLM_BATCH,
|
||||
max_new_tokens: Optional[int] = None,
|
||||
checkpoint: Optional[Callable[[List["PuzzleResult"]], None]] = None,
|
||||
lean: bool = False,
|
||||
use_match_assignment: bool = True,
|
||||
vote_samples: int = 0,
|
||||
vote_temp: float = 0.5,
|
||||
hint: bool = False
|
||||
) -> List[PuzzleResult]:
|
||||
"""`checkpoint`, when given, is called with the (complete, valid) results
|
||||
after the symbolic pass and after every LLM batch — so a crash at ANY
|
||||
later point still leaves a full submission on disk."""
|
||||
client = client or NullClient()
|
||||
budget = budget or Budget()
|
||||
|
||||
# adaptive mode: the hidden test is small (an IOL contest reformatted into
|
||||
# a handful of multi-item rows), so default to spending more per puzzle.
|
||||
# A large row count flips us to coverage mode (shorter generations, serve
|
||||
# the most puzzles).
|
||||
deep_mode = len(rows) <= DEEP_MODE_MAX_ROWS
|
||||
if max_new_tokens is None:
|
||||
max_new_tokens = 2048 if deep_mode else 1024
|
||||
|
||||
def log(msg: str) -> None:
|
||||
if verbose:
|
||||
print(msg, flush=True)
|
||||
|
||||
def save() -> None:
|
||||
if checkpoint is not None:
|
||||
try:
|
||||
checkpoint(results)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# ---- 1. symbolic pass ----
|
||||
results: List[PuzzleResult] = []
|
||||
puzzles: List[Optional[Puzzle]] = []
|
||||
for i, row in enumerate(rows):
|
||||
rid = str(row.get("id", i))
|
||||
try:
|
||||
p = parse_puzzle(row)
|
||||
answers, confs, methods = solve_puzzle_ex(p, NullClient(), budget,
|
||||
puzzles_left=len(rows) - i)
|
||||
except Exception:
|
||||
p = None
|
||||
answers, confs, methods = [str(row.get("query", "?")).strip() or "?"], [0.0], ["fallback"]
|
||||
puzzles.append(p)
|
||||
results.append(PuzzleResult(rid, answers, "", confs, methods))
|
||||
# make every result submission-valid NOW (explanations + non-empty), so
|
||||
# each checkpoint from here on is a complete fallback submission
|
||||
for i, r in enumerate(results):
|
||||
r.answers = [str(a).strip() or "?" for a in r.answers]
|
||||
r.explanation = (_symbolic_explanation(puzzles[i], r.methods)
|
||||
if puzzles[i] is not None else
|
||||
"- No parseable structure found; answered by "
|
||||
"closest-example analogy.")
|
||||
log(f"symbolic pass done in {budget.elapsed():.1f}s")
|
||||
save()
|
||||
|
||||
# arm the per-token wall-clock abort so generation can't overrun the budget
|
||||
try:
|
||||
client.deadline = budget.start + budget.total - budget.safety
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# ---- 1b. match_letters assignment pass (model logprobs -> Hungarian) ----
|
||||
# Free-form generation answers match_letters with the identity permutation
|
||||
# (scores ~0). Solve it as an assignment from the model's own distribution
|
||||
# instead. Solved puzzles get high confidence so the free-form LLM pass
|
||||
# skips them; a declined puzzle falls through to that pass unchanged.
|
||||
# Gated by use_match_assignment: when off, match_letters puzzles go through
|
||||
# the normal free-form LLM pass.
|
||||
if not use_match_assignment:
|
||||
log("match_letters assignment pass disabled; using free-form LLM path")
|
||||
if use_match_assignment and getattr(client, "can_score", False):
|
||||
n_assigned = 0
|
||||
for i, p in enumerate(puzzles):
|
||||
if p is None or p.task_type != "match_letters" or budget.exhausted():
|
||||
continue
|
||||
try:
|
||||
letters = solve_match_letters_llm(p, client)
|
||||
except Exception as e:
|
||||
log(f" match_letters solver failed on {results[i].row_id}: "
|
||||
f"{type(e).__name__}: {e}")
|
||||
letters = None
|
||||
r = results[i]
|
||||
if letters and len(letters) == len(r.answers):
|
||||
r.answers = [str(x).strip() or "?" for x in letters]
|
||||
r.confs = [0.9] * len(letters)
|
||||
r.methods = ["llm-assignment"] * len(letters)
|
||||
r.llm_used = True
|
||||
n_assigned += 1
|
||||
if n_assigned:
|
||||
log(f"match_letters assignment solver used on {n_assigned} puzzle(s)")
|
||||
save()
|
||||
|
||||
# ---- 2. scaffolded LLM pass ----
|
||||
# In lean mode symbolic is a pure fallback: the model answers EVERY puzzle
|
||||
# and its answer wins wherever it produced one (symbolic stands only for the
|
||||
# items it left blank). Otherwise the model runs only on low-confidence
|
||||
# puzzles and verified symbolic answers override it.
|
||||
if lean:
|
||||
need = [i for i in range(len(results)) if puzzles[i] is not None]
|
||||
else:
|
||||
need = [i for i, r in enumerate(results)
|
||||
if puzzles[i] is not None and any(c < conf_keep for c in r.confs)]
|
||||
if deep_mode:
|
||||
# weakest-first: a budget cutoff then drops the puzzles we could help least
|
||||
need.sort(key=lambda i: sum(min(c, conf_keep) for c in results[i].confs)
|
||||
/ max(len(results[i].confs), 1))
|
||||
else:
|
||||
# coverage: shortest prompt first maximizes puzzles served per second
|
||||
need.sort(key=lambda i: len(puzzles[i].context) + len(puzzles[i].query))
|
||||
from .direct import LEAN_SYSTEM, SYSTEM
|
||||
sys_prompt = LEAN_SYSTEM if lean else SYSTEM
|
||||
log(f"LLM pass ({'lean' if lean else 'scaffold'}, "
|
||||
f"{'deep' if deep_mode else 'coverage'}, "
|
||||
f"max_new_tokens={max_new_tokens}): {len(need)}/{len(rows)} puzzles "
|
||||
f"need the model; client={'yes' if client.available else 'no'}")
|
||||
parse_ok = parse_fail = 0
|
||||
if client.available and need:
|
||||
done = 0
|
||||
for start in range(0, len(need), llm_batch):
|
||||
if budget.exhausted():
|
||||
log(f"budget cutoff after {done} LLM puzzles")
|
||||
break
|
||||
batch = need[start : start + llm_batch]
|
||||
if lean:
|
||||
# lean: no scaffold, unless the optional light hint is enabled
|
||||
scaffolds = [light_hint(puzzles[i]) if hint else "" for i in batch]
|
||||
else:
|
||||
scaffolds = []
|
||||
for i in batch:
|
||||
p, r = puzzles[i], results[i]
|
||||
try:
|
||||
scaffolds.append(build_scaffold(p, r.answers, r.confs, r.methods))
|
||||
except Exception:
|
||||
scaffolds.append("")
|
||||
try:
|
||||
outs = solve_direct([puzzles[i] for i in batch], client, scaffolds,
|
||||
max_new_tokens=max_new_tokens, system=sys_prompt,
|
||||
lean=lean)
|
||||
except Exception as e:
|
||||
# systemic generation failure (bad load, driver, etc.) — the
|
||||
# symbolic answers already on every item are the submission
|
||||
log(f"LLM batch failed ({type(e).__name__}: {e}); "
|
||||
f"keeping symbolic answers for the rest")
|
||||
break
|
||||
for i, (direct, expl, raw, found) in zip(batch, outs):
|
||||
r = results[i]
|
||||
p = puzzles[i]
|
||||
r.llm_used = True
|
||||
r.raw = raw
|
||||
if found:
|
||||
parse_ok += 1
|
||||
else:
|
||||
# no answer block parsed even after salvage: do NOT let
|
||||
# reasoning prose overwrite the symbolic answers
|
||||
parse_fail += 1
|
||||
continue
|
||||
labels = sorted(p.lettered) if p.lettered else ()
|
||||
aligned = align_answers(direct, len(r.answers))
|
||||
for j, d in enumerate(aligned):
|
||||
# lean: the model's answer wins wherever it gave one;
|
||||
# scaffold: only where symbolic isn't confident
|
||||
if d and (lean or r.confs[j] < conf_keep):
|
||||
cleaned = clean_llm_answer(d, p.task_type, labels)
|
||||
if cleaned:
|
||||
# lean mode ships the model's answer as-is (no
|
||||
# punctuation/casing induction)
|
||||
if not lean:
|
||||
cleaned = induce_format(cleaned, p, j)
|
||||
r.answers[j] = cleaned.strip() or r.answers[j]
|
||||
r.methods[j] = "llm"
|
||||
if expl:
|
||||
r.explanation = expl
|
||||
done += len(batch)
|
||||
log(f" llm {done}/{len(need)} ok={parse_ok} fail={parse_fail} "
|
||||
f"t={budget.elapsed():.0f}s")
|
||||
save()
|
||||
log(f"LLM pass done: {parse_ok} parsed, {parse_fail} unparsable "
|
||||
f"(kept symbolic); {budget.elapsed():.0f}s elapsed")
|
||||
|
||||
# ---- 2b. light greedy-anchored self-consistency voting (lean only) ----
|
||||
# The greedy answers are already checkpointed; sampled passes can only
|
||||
# displace an item on genuine agreement (see _vote), so this is monotone and
|
||||
# budget-gated — if the clock runs out we simply keep the greedy answers.
|
||||
if lean and vote_samples > 0 and client.available and need and not budget.exhausted():
|
||||
ballots = {i: [list(results[i].answers)] for i in need} # greedy = ballot 0
|
||||
done_votes = 0
|
||||
for _s in range(vote_samples):
|
||||
if budget.exhausted():
|
||||
break
|
||||
failed = False
|
||||
for start in range(0, len(need), llm_batch):
|
||||
if budget.exhausted():
|
||||
break
|
||||
batch = need[start : start + llm_batch]
|
||||
try:
|
||||
outs = solve_direct([puzzles[i] for i in batch], client,
|
||||
["" for _ in batch],
|
||||
max_new_tokens=max_new_tokens,
|
||||
system=sys_prompt, lean=True,
|
||||
sample=True, temperature=vote_temp)
|
||||
except Exception as e:
|
||||
log(f"vote pass failed ({type(e).__name__}: {e}); "
|
||||
f"keeping greedy answers")
|
||||
failed = True
|
||||
break
|
||||
for i, (direct, _e, _raw, _f) in zip(batch, outs):
|
||||
p = puzzles[i]
|
||||
labels = sorted(p.lettered) if p.lettered else ()
|
||||
ballot = []
|
||||
for j in range(len(results[i].answers)):
|
||||
d = direct[j] if j < len(direct) else None
|
||||
ballot.append(clean_llm_answer(d, p.task_type, labels)
|
||||
if d else None)
|
||||
ballots[i].append(ballot)
|
||||
if failed:
|
||||
break
|
||||
done_votes += 1
|
||||
for i in need: # re-vote (greedy-anchored) after each sample pass
|
||||
r = results[i]
|
||||
anchor = ballots[i][0]
|
||||
for j in range(len(r.answers)):
|
||||
v = _vote([b[j] for b in ballots[i] if j < len(b)],
|
||||
anchor[j] if j < len(anchor) else None)
|
||||
if v and str(v).strip():
|
||||
r.answers[j] = str(v).strip()
|
||||
save()
|
||||
log(f"voting: {done_votes}/{vote_samples} sample pass(es); "
|
||||
f"t={budget.elapsed():.0f}s")
|
||||
|
||||
# ---- 3. final never-empty guarantee ----
|
||||
for r in results:
|
||||
r.answers = [str(a).strip() or "?" for a in r.answers]
|
||||
return results
|
||||
Reference in New Issue
Block a user