iolai26-solve / solver /fallback.py
rvpant
Notebook-style MODEL_ID in script.py; move metrics into solver/; untrack dev-only eval/tests/data
f96c703
Raw
History Blame Contribute Delete
2.64 kB
"""chrF-floor fallback: never return an empty or wildly-off answer.
The geometric-mean metric means one empty answer costs far more than a wrong
but plausible one. Fallback ladder (best available wins):
1. analogy from the closest attested source (transfers its target with the
observed source->query edit applied),
2. the attested target of the most chrF-similar attested source,
3. echo of query content words mapped through alignment,
4. the raw query text itself (last resort: shares characters with gold more
often than an empty string does).
"""
from __future__ import annotations
from typing import List, Optional, Tuple
from . import analogy
from .metrics import chrf
from .align import align as build_align, best_translation
from .preprocess import Pair, strip_punct, tokenize
def closest_attested(query: str, sources: List[str]) -> Tuple[int, float]:
"""Index and similarity of the attested source closest to the query."""
best_i, best_s = -1, -1.0
for i, s in enumerate(sources):
sc = chrf(query, s)
if sc > best_s:
best_i, best_s = i, sc
return best_i, best_s
def fallback_answer(query: str, pairs: List[Pair], direction: str = "to_work") -> str:
"""direction: 'to_work' = translate task->work (analysis);
'to_task' = work->task (generation). Pairs are (task, work)."""
if direction == "to_task":
srcs = [p.tgt for p in pairs]
tgts = [p.src for p in pairs]
flipped = [Pair(src=p.tgt, tgt=p.src) for p in pairs]
else:
srcs = [p.src for p in pairs]
tgts = [p.tgt for p in pairs]
flipped = pairs
query = query.strip()
if not query:
return tgts[0] if tgts else "?"
if srcs:
i, sim = closest_attested(query, srcs)
if i >= 0:
# 1. analogy transfer: apply the srcs[i]->query edit to tgts[i]
transfer = analogy.solve(srcs[i], query, tgts[i])
if transfer and sim > 0.3:
return transfer[0]
# 2. echo the closest attested target
if sim > 0.15 and tgts[i]:
return tgts[i]
# 3. word-by-word through alignment
amap = build_align(flipped)
words = [strip_punct(t) for t in tokenize(query)]
mapped = [best_translation(amap, w) or w for w in words if w]
if mapped:
return " ".join(mapped)
# 4. absolute floor
return query
def ensure_nonempty(ans: Optional[str], query: str, pairs: List[Pair], direction: str = "to_work") -> str:
if ans and str(ans).strip():
return str(ans).strip()
return fallback_answer(query, pairs, direction) or "?"