File size: 2,639 Bytes
379f378
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f96c703
379f378
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
"""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 "?"