267 lines
9.9 KiB
Python
267 lines
9.9 KiB
Python
"""match_letters: optimal one-to-one assignment (pure-python Hungarian) over
|
|
either surface-similarity scores or the model's next-token log-probs."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import Dict, List, Optional, Sequence, Tuple
|
|
|
|
from .align import align as build_align
|
|
from .preprocess import Pair, Puzzle, strip_punct, tokenize
|
|
|
|
|
|
def hungarian(cost: List[List[float]]) -> List[int]:
|
|
"""Minimum-cost perfect matching on a square cost matrix.
|
|
Returns assignment: row i -> column result[i]. Jonker-style O(n^3)
|
|
shortest augmenting path implementation."""
|
|
n = len(cost)
|
|
if n == 0:
|
|
return []
|
|
INF = float("inf")
|
|
u = [0.0] * (n + 1)
|
|
v = [0.0] * (n + 1)
|
|
p = [0] * (n + 1) # p[j] = row matched to column j (1-indexed)
|
|
way = [0] * (n + 1)
|
|
for i in range(1, n + 1):
|
|
p[0] = i
|
|
j0 = 0
|
|
minv = [INF] * (n + 1)
|
|
used = [False] * (n + 1)
|
|
while True:
|
|
used[j0] = True
|
|
i0, delta, j1 = p[j0], INF, 0
|
|
for j in range(1, n + 1):
|
|
if not used[j]:
|
|
cur = cost[i0 - 1][j - 1] - u[i0] - v[j]
|
|
if cur < minv[j]:
|
|
minv[j] = cur
|
|
way[j] = j0
|
|
if minv[j] < delta:
|
|
delta = minv[j]
|
|
j1 = j
|
|
for j in range(n + 1):
|
|
if used[j]:
|
|
u[p[j]] += delta
|
|
v[j] -= delta
|
|
else:
|
|
minv[j] -= delta
|
|
j0 = j1
|
|
if p[j0] == 0:
|
|
break
|
|
while j0:
|
|
j1 = way[j0]
|
|
p[j0] = p[j1]
|
|
j0 = j1
|
|
ans = [0] * n
|
|
for j in range(1, n + 1):
|
|
if p[j]:
|
|
ans[p[j] - 1] = j - 1
|
|
return ans
|
|
|
|
|
|
def _char_ngrams(s: str, nmin: int = 2, nmax: int = 4) -> set:
|
|
s = s.casefold().replace(" ", "")
|
|
return {s[i : i + n] for n in range(nmin, nmax + 1) for i in range(len(s) - n + 1)}
|
|
|
|
|
|
def _sim(a: str, b: str) -> float:
|
|
ga, gb = _char_ngrams(a), _char_ngrams(b)
|
|
if not ga or not gb:
|
|
return 0.0
|
|
return len(ga & gb) / max(len(ga | gb), 1)
|
|
|
|
|
|
def score_matrix(
|
|
forms: Sequence[str], meanings: Sequence[str], pairs: List[Pair]
|
|
) -> List[List[float]]:
|
|
"""Higher = better match. Combines token-level alignment evidence from the
|
|
attested pairs with surface similarity to attested forms sharing meaning
|
|
words."""
|
|
amap = build_align(pairs) if pairs else {}
|
|
# index attested: meaning word -> attested source strings
|
|
attested_by_word: Dict[str, List[str]] = {}
|
|
for p in pairs:
|
|
for w in tokenize(p.tgt):
|
|
w = strip_punct(w).casefold()
|
|
if w:
|
|
attested_by_word.setdefault(w, []).append(p.src)
|
|
|
|
S = []
|
|
for f in forms:
|
|
f_toks = [strip_punct(t).casefold() for t in tokenize(f)]
|
|
row = []
|
|
for m in meanings:
|
|
m_words = [strip_punct(w).casefold() for w in tokenize(m)]
|
|
score = 0.0
|
|
# alignment evidence: form tokens aligned to meaning words
|
|
for ft in f_toks:
|
|
for tgt, s in amap.get(ft, []):
|
|
if tgt in m_words:
|
|
score += s
|
|
# surface similarity to attested sources of these meaning words
|
|
for w in m_words:
|
|
for src in attested_by_word.get(w, []):
|
|
score += 0.5 * _sim(f, src)
|
|
row.append(score)
|
|
S.append(row)
|
|
return S
|
|
|
|
|
|
def _self_sim(texts: Sequence[str], char_level: bool) -> List[List[float]]:
|
|
"""Pairwise similarity within one side: shared char n-grams for unknown
|
|
forms, shared content words for meanings."""
|
|
n = len(texts)
|
|
feats = []
|
|
for t in texts:
|
|
if char_level:
|
|
feats.append(_char_ngrams(t, 3, 5))
|
|
else:
|
|
stop = {"the", "a", "an", "of", "is", "it", "he", "she", "they",
|
|
"are", "in", "to", "for", "with", "and", "or"}
|
|
feats.append({w for w in (strip_punct(x).casefold() for x in tokenize(t))
|
|
if w and w not in stop})
|
|
S = [[0.0] * n for _ in range(n)]
|
|
for i in range(n):
|
|
for j in range(i + 1, n):
|
|
inter = len(feats[i] & feats[j])
|
|
union = len(feats[i] | feats[j]) or 1
|
|
S[i][j] = S[j][i] = inter / union
|
|
return S
|
|
|
|
|
|
def structural_scores(forms: Sequence[str], meanings: Sequence[str],
|
|
iters: int = 4) -> List[List[float]]:
|
|
"""Structure-matching signal for zero-lexical-evidence matching: forms
|
|
sharing morphemes should map to meanings sharing words. Soft assignment
|
|
power iteration X <- Sf @ X @ Sm (a light quadratic-assignment relaxation)
|
|
starting from uniform. Returns an (n_forms x n_meanings) score matrix."""
|
|
nf, nm = len(forms), len(meanings)
|
|
if nf == 0 or nm == 0:
|
|
return [[0.0] * nm for _ in range(nf)]
|
|
Sf = _self_sim(forms, char_level=True)
|
|
Sm = _self_sim(meanings, char_level=False)
|
|
# seed with degree-profile agreement: a form clustered with k others
|
|
# should map to a meaning clustered with ~k others. (A uniform seed is a
|
|
# degenerate fixed point — every row converges to the same profile.)
|
|
def profile(S, i):
|
|
return sorted((v for v in S[i] if v > 0.05), reverse=True)[:6]
|
|
|
|
X = []
|
|
for i in range(nf):
|
|
pf = profile(Sf, i)
|
|
row = []
|
|
for j in range(nm):
|
|
pm = profile(Sm, j)
|
|
d = sum(abs(a - b) for a, b in zip(pf, pm)) + abs(len(pf) - len(pm))
|
|
row.append(1.0 / (1.0 + d))
|
|
X.append(row)
|
|
for _ in range(iters):
|
|
# Y = Sf @ X @ Sm (tiny n: pure-python is fine)
|
|
T = [[sum(Sf[i][k] * X[k][j] for k in range(nf)) for j in range(nm)]
|
|
for i in range(nf)]
|
|
Y = [[sum(T[i][k] * Sm[k][j] for k in range(nm)) for j in range(nm)]
|
|
for i in range(nf)]
|
|
# row-normalize to keep the iteration bounded
|
|
X = []
|
|
for row in Y:
|
|
z = sum(row) or 1.0
|
|
X.append([v / z for v in row])
|
|
return X
|
|
|
|
|
|
def solve_matching(
|
|
forms: Sequence[str], meanings: Sequence[str], pairs: List[Pair]
|
|
) -> List[Tuple[str, str]]:
|
|
"""Optimal assignment of forms to meanings. Combines lexical/alignment
|
|
evidence (when attested pairs exist) with the structural signal (always).
|
|
Pads to square with zero scores when lengths differ."""
|
|
n = max(len(forms), len(meanings))
|
|
S = score_matrix(forms, meanings, pairs)
|
|
S2 = structural_scores(forms, meanings)
|
|
cost = [[0.0] * n for _ in range(n)]
|
|
for i in range(n):
|
|
for j in range(n):
|
|
s = 0.0
|
|
if i < len(forms) and j < len(meanings):
|
|
s = S[i][j] + 3.0 * len(meanings) * S2[i][j]
|
|
cost[i][j] = -s
|
|
assign = hungarian(cost)
|
|
out = []
|
|
for i, f in enumerate(forms):
|
|
j = assign[i]
|
|
out.append((f, meanings[j] if j < len(meanings) else ""))
|
|
return out
|
|
|
|
|
|
# --------------------------------------------------------------------------
|
|
# LLM-scored assignment. Free-form generation tends to answer match_letters with
|
|
# the identity permutation (A, B, C, ...) — a valid permutation that scores ~0.
|
|
# Instead we score each (item, option-letter) pair from the model's next-token
|
|
# log-probs and take the optimal one-to-one assignment, so the bijection is
|
|
# enforced exactly rather than hoped for. The assignment engine is our existing
|
|
# Hungarian; only the score source changes (surface features -> the model's own
|
|
# distribution).
|
|
# --------------------------------------------------------------------------
|
|
|
|
_MATCH_SYSTEM = (
|
|
"You match items to their correct counterparts in a linguistics problem. "
|
|
"Reply with one option letter only."
|
|
)
|
|
|
|
|
|
def letters_from_scores(scores: List[List[float]], letters: Sequence[str]
|
|
) -> List[str]:
|
|
"""Given a per-item score over option letters (item x option), return one
|
|
letter per item. When items and options are equinumerous (the IOL matching
|
|
shape) the assignment is a strict bijection via Hungarian; otherwise it
|
|
degrades to an independent per-item argmax."""
|
|
n_items, n_opt = len(scores), len(letters)
|
|
if n_items == 0 or n_opt == 0:
|
|
return []
|
|
if n_items == n_opt:
|
|
cost = [[-scores[i][j] for j in range(n_opt)] for i in range(n_items)]
|
|
assign = hungarian(cost)
|
|
return [letters[assign[i]] for i in range(n_items)]
|
|
return [letters[max(range(n_opt), key=lambda j: scores[i][j])]
|
|
for i in range(n_items)]
|
|
|
|
|
|
def solve_match_letters_llm(puzzle: Puzzle, client) -> Optional[List[str]]:
|
|
"""Assign each numbered item to an option letter using the model's
|
|
next-token log-probs, then Hungarian. Returns one letter per puzzle item
|
|
(in item order), or None if it declines: no scoring backend, not a
|
|
well-formed lettered-matching shape, or fewer than 3 items/options."""
|
|
if not getattr(client, "can_score", False):
|
|
return None
|
|
items = puzzle.items
|
|
if not puzzle.lettered or len(items) < 3:
|
|
return None
|
|
letters = sorted(puzzle.lettered)
|
|
if len(letters) < 3:
|
|
return None
|
|
|
|
tok = client.tok
|
|
cand_ids: List[List[int]] = []
|
|
for L in letters:
|
|
ids = set()
|
|
for form in (L, " " + L):
|
|
t = tok.encode(form, add_special_tokens=False)
|
|
if t:
|
|
ids.add(t[0])
|
|
cand_ids.append(sorted(ids))
|
|
if any(not g for g in cand_ids):
|
|
return None
|
|
|
|
ctx = (puzzle.context or "").strip()
|
|
prompts = []
|
|
for it in items:
|
|
num = it.number or "?"
|
|
form = (it.text or "").strip()
|
|
prompts.append(
|
|
f"{ctx}\n\nWhich lettered option corresponds to item {num} "
|
|
f"({form})? Reply with the option letter only.")
|
|
|
|
scores = client.score_next_logprobs(prompts, cand_ids, system=_MATCH_SYSTEM)
|
|
if not scores or len(scores) != len(items):
|
|
return None
|
|
return letters_from_scores(scores, letters)
|