File size: 8,077 Bytes
aad1ed3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
#!/usr/bin/env python3
"""Tests for the extractive reader. Pure python, no model, no GPU."""
import sys
import time
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from src.reader import AnswerType, LexicalSpanReader, _stem, detect_answer_type  # noqa: E402
from src.router import Passage, Route, Router, RouterConfig  # noqa: E402

PASS = FAIL = 0
def chk(c, l, d=""):
    global PASS, FAIL
    if c: PASS += 1; print(f"  ok   {l}")
    else:  FAIL += 1; print(f"  FAIL {l}  {d}")

R = LexicalSpanReader()
P = [Passage("c1", "Quick Answer. Eagles fly 30 to 55 mph and dive at over 100 mph. "
                   "Eagles can soar for hours on warm air currents, which conserves energy.", 0.91),
     Passage("c2", "The bald eagle is the national bird of the United States. "
                   "It was adopted as the emblem in 1782.", 0.72)]


print("[stemmer: plural must collapse onto singular]")
for a, b in [("eagles", "eagle"), ("boxes", "box"), ("dishes", "dish"),
             ("cities", "city"), ("countries", "country"), ("states", "state")]:
    chk(_stem(a) == _stem(b), f"{a} == {b}", (_stem(a), _stem(b)))
print("  -- and must NOT over-strip --")
for w in ("class", "glass", "was", "is", "bird", "bus"):
    chk(_stem(w) == w, f"{w} unchanged", _stem(w))


print("\n[answer-type detection across scripts]")
for q, e in [("how fast does an eagle travel", AnswerType.NUMERIC),
             ("how many people live there", AnswerType.NUMERIC),
             ("what year was it adopted", AnswerType.NUMERIC),
             ("who is the president", AnswerType.PERSON),
             ("where is mount everest", AnswerType.LOCATION),
             ("effects of climate change", AnswerType.DESCRIPTION),
             ("कितने लोग रहते हैं", AnswerType.NUMERIC),
             ("யார் இதை இயக்கினார்", AnswerType.PERSON),
             ("کتنے لوگ رہتے ہیں", AnswerType.NUMERIC)]:
    chk(detect_answer_type(q) == e, f"{q[:34]:36s} -> {e}", detect_answer_type(q))


print("\n[extraction finds the right span]")
for q, want in [("how fast does an eagle travel", "30 to 55 mph"),
                ("what year was the bald eagle adopted", "1782"),
                ("what is the national bird of the united states", "bald eagle"),
                ("how do eagles conserve energy", "soar")]:
    s = R.read(q, P)
    chk(want.lower() in s.text.lower(), f"{q[:42]:44s} -> {want}", s.text[:60])
    chk(s.chunk_id in ("c1", "c2"), "cites a real chunk")


print("\n[NUMERIC questions prefer spans containing digits]")
s = R.read("how fast does an eagle travel", P)
chk(any(ch.isdigit() for ch in s.text), "digit-bearing span chosen", s.text)
# the query shares only 'eagle' with the answer sentence, so pure overlap would
# have picked the bald-eagle sentence -- the numeric prior is what rescues it
chk("bald" not in s.text.lower(), "numeric prior overrode raw lexical overlap")


print("\n[confidence separates answerable from unanswerable]")
good = R.read("what is the national bird of the united states", P).score
bad = R.read("what is the capital of mongolia", P).score
chk(good > 0.8, "confident on an answerable query", good)
chk(bad < 0.2, "unconfident on an unanswerable one", bad)
chk(good - bad > 0.5, "wide separation -> usable for conformal", good - bad)


print("\n[degenerate inputs]")
chk(R.read("anything", [])[1] == 0.0 if isinstance(R.read("anything", []), tuple)
    else R.read("anything", []).score == 0.0, "no passages -> zero confidence")
chk(R.read("", P).score >= 0.0, "empty query does not crash")
chk(R.read("x", [Passage("c", "", 0.5)]).score == 0.0, "empty passage text")
long_q = " ".join(["word"] * 200)
chk(R.read(long_q, P).score >= 0.0, "very long query does not crash")


print("\n[latency budget: the reader must be single-digit ms]")
ts = []
for _ in range(500):
    t = time.perf_counter_ns()
    R.read("how fast does an eagle travel", P)
    ts.append((time.perf_counter_ns() - t) / 1e6)
ts.sort()
p50, p100 = ts[250], ts[-1]
chk(p50 < 5.0, f"P50 {p50:.3f} ms < 5 ms")
chk(p100 < 20.0, f"P100 {p100:.3f} ms < 20 ms")
print("")
print("[agglutinative Indic queries actually select a span]")
_q   = "ಕಾರ್ಪೋರೇಷನ್ ಎಂದರೇನು"
_hit = "ಕಾರ್ಪೋರೇಷನ್‌ಗಳು ಎಂದು ಕರೆಯಲಾಗುತ್ತದೆ."
_other = "ಸ್ಟಾಕ್ ಮಾರುಕಟ್ಟೆ ತೆರೆಯಿತು."
_p = [Passage("k:0", _other, 0.60), Passage("k:1", _hit, 0.61)]
_s = LexicalSpanReader().read(_q, _p)
# Before prefix folding every candidate scored 0, so the reader returned
# whichever sorted first and reported conf 0.00 -- it was not choosing.
chk(_s.score > 0.0, "confidence is non-zero on an inflected Indic match", _s.score)
chk(_s.chunk_id == "k:1", "the passage containing the term is chosen", _s.chunk_id)

print(f"       P50 {p50:.3f}  P70 {ts[350]:.3f}  P100 {p100:.3f} ms")


print("\n[wires into the router end to end]")
def gen(q, ps, cap, constrained):
    return (f"generated<={cap}", 0.75)

router = Router(RouterConfig(tau_retrieval=0.30, tau_extract=0.60), R, gen)
d = router.route("what is the national bird of the united states", P)
chk(d.route is Route.EXTRACT, "confident query takes the EXTRACT path", d.route)
chk("bald eagle" in d.answer.lower(), "answer carried through", d.answer[:50])
chk(d.total_ms < 10, f"end-to-end {d.total_ms:.2f} ms", d.timings_ms)

d = router.route("what is the capital of mongolia", P)
chk(d.route is not Route.EXTRACT, "unconfident query does not extract", d.route)

d = router.route("q", [Passage("c", "text", 0.1)])
chk(d.route is Route.ABSTAIN, "low retrieval score still abstains first")


print("\n[route mix on a mixed workload resembles the measured bands]")
qs = [("what is the national bird of the united states", P),
      ("how fast does an eagle travel", P),
      ("what year was the bald eagle adopted", P),
      ("how do eagles conserve energy", P),
      ("what is the capital of mongolia", P)]
routes = [router.route(q, ps).route for q, ps in qs]
n_ex = sum(r is Route.EXTRACT for r in routes)
chk(n_ex >= 2, f"{n_ex}/5 extracted", [r.value for r in routes])


# ---------------------------------------------------------------- prior weight
print("\n[retrieval prior]")
from src.reader import LexicalSpanReader as _LSR  # noqa: E402

_G = "এটা ঈগলে ৩০ থেকে ৫৫ মাইল প্রতি ঘণ্টা গতিত উৰিব পাৰে।"
_D = "আপুনি কিমান আগতীয়াকৈ সংৰক্ষণ কৰে আৰু ৰে'লখন কিমান ভৰ্তি হৈ আছে তাৰ ওপৰত নিৰ্ভৰ কৰে।"
_Q = "এটা ঈগলে কিমান দ্ৰুতগতিত ভ্ৰমণ কৰে"

_hi = [Passage("gold", _G, 0.9, "as"), Passage("dist", _D, 0.1, "as")]
chk(_LSR(prior_weight=0.0).read(_Q, _hi).chunk_id in ("gold", "dist"),
    "prior_weight 0 leaves selection to lexical overlap")
chk(_LSR(prior_weight=1.0).read(_Q, _hi).chunk_id == "gold",
    "prior_weight 1 follows the retriever")

# A passage the retriever loves must not win with zero lexical support.
_lo = [Passage("gold", _G, 0.1, "as"), Passage("junk", "কলা ভাল ফল।", 0.9, "as")]
chk(_LSR(prior_weight=1.0).read(_Q, _lo).chunk_id == "gold",
    "multiplicative blend: prior cannot rescue a no-overlap passage")

# Equal scores must behave exactly like the prior being off.
_eq = [Passage("a", _G, 0.5, "as"), Passage("b", _D, 0.5, "as")]
chk(_LSR(prior_weight=1.0).read(_Q, _eq).chunk_id == _LSR(prior_weight=0.0).read(_Q, _eq).chunk_id,
    "flat retrieval scores are a no-op")
chk(_LSR().prior_weight == 0.0, "default is OFF -- opt in after measuring")

print(f"\n{'='*54}\n  {PASS} passed, {FAIL} failed\n{'='*54}")

sys.exit(1 if FAIL else 0)