-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathevaluation.py
More file actions
183 lines (155 loc) · 6.46 KB
/
Copy pathevaluation.py
File metadata and controls
183 lines (155 loc) · 6.46 KB
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
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
"""Run a learning system against a dataset and score it.
The harness teaches lessons one at a time. After each lesson, it scores
*every question whose source lessons have all been taught* and records
the per-kind accuracy — producing a forgetting curve not just for recall
but for composition and negatives too.
Every lesson, every question, every answer is streamed into the
`RunRecorder` (if one is supplied). `events.jsonl` is the ground-truth
log; `metrics.json` / `curves.json` / `predictions.jsonl` are derived
from it at close time.
"""
from __future__ import annotations
import time
from dataclasses import dataclass, field, asdict
from architectures.base import LearningSystem
from recorder import RunRecorder
from tasks import Answer, Dataset, Question
@dataclass
class EvalResult:
arch_name: str
seed: int
n_lessons: int
recall_accuracy: float
composition_accuracy: float
composition_3hop_accuracy: float
negative_accuracy: float
overall_accuracy: float
curves: dict[str, list[float]] = field(default_factory=dict)
@property
def forgetting_curve(self) -> list[float]:
"""Back-compat alias for the recall curve."""
return self.curves.get("recall", [])
def to_dict(self) -> dict:
return asdict(self)
def _score(question: Question, answer: Answer) -> bool:
if question.kind == "negative":
return answer.predicted_is_true == question.gold_is_true
return answer.predicted_object == question.gold_object
def _answer_with_chaining(system: "LearningSystem", q: Question) -> Answer:
"""Produce an answer for any question kind, chaining recall queries for 3-hop.
composition_3hop is scored by calling system.answer() three times with
synthetic recall questions, threading the predicted object through the
three relation hops. This keeps architectures from needing to know about
composition_3hop — they only need to handle 'recall' correctly.
All other question kinds (recall, composition, negative) are dispatched
directly to the architecture's own answer() method, preserving its
internal 2-hop implementation for composition if it has one.
"""
if q.kind != "composition_3hop":
return system.answer(q)
assert q.relation2 is not None and q.relation3 is not None
q1 = Question(qid=-1, kind="recall", subject=q.subject, relation=q.relation,
source_lessons=())
a1 = system.answer(q1)
if a1.predicted_object is None:
return Answer(predicted_object=None)
q2 = Question(qid=-1, kind="recall", subject=a1.predicted_object, relation=q.relation2,
source_lessons=())
a2 = system.answer(q2)
if a2.predicted_object is None:
return Answer(predicted_object=None)
q3 = Question(qid=-1, kind="recall", subject=a2.predicted_object, relation=q.relation3,
source_lessons=())
a3 = system.answer(q3)
return a3
def _applicable(question: Question, current_lesson: int) -> bool:
"""A question is applicable once every lesson it depends on has been taught."""
return all(sl <= current_lesson for sl in question.source_lessons)
def _prediction_payload(question: Question, answer: Answer) -> object:
if question.kind == "negative":
return answer.predicted_is_true
return answer.predicted_object
def evaluate(
system: LearningSystem,
dataset: Dataset,
recorder: RunRecorder | None = None,
verbose: bool = False,
) -> EvalResult:
curves: dict[str, list[float]] = {
"recall": [],
"composition": [],
"composition_3hop": [],
"negative": [],
"overall_applicable": [],
}
for li, lesson in enumerate(dataset.lessons):
if recorder is not None:
recorder.log_event(
"lesson_started",
lesson_idx=li,
facts=[[f.subject, f.relation, f.object] for f in lesson.facts],
)
t0 = time.time()
system.learn(lesson)
train_time = time.time() - t0
if recorder is not None:
recorder.log_event("lesson_finished", lesson_idx=li, wall_time_s=train_time)
# Score every question whose source lessons have all been taught.
kind_counts: dict[str, list[int]] = {
"recall": [0, 0],
"composition": [0, 0],
"composition_3hop": [0, 0],
"negative": [0, 0],
}
for q in dataset.questions:
if not _applicable(q, li):
continue
a = _answer_with_chaining(system, q)
ok = _score(q, a)
if recorder is not None:
recorder.log_event(
"question_answered",
checkpoint_lesson=li,
qid=q.qid,
kind=q.kind,
prediction=_prediction_payload(q, a),
correct=ok,
)
kind_counts[q.kind][0] += int(ok)
kind_counts[q.kind][1] += 1
def _pct(pair: list[int]) -> float:
return pair[0] / pair[1] if pair[1] else 0.0
checkpoint = {
"recall": _pct(kind_counts["recall"]),
"composition": _pct(kind_counts["composition"]),
"composition_3hop": _pct(kind_counts["composition_3hop"]),
"negative": _pct(kind_counts["negative"]),
}
total_ok = sum(c[0] for c in kind_counts.values())
total_n = sum(c[1] for c in kind_counts.values())
checkpoint["overall_applicable"] = total_ok / total_n if total_n else 0.0
for k, v in checkpoint.items():
curves[k].append(v)
if recorder is not None:
recorder.log_event("checkpoint_evaluated", lesson_idx=li, **checkpoint)
if verbose:
print(
f" after lesson {li}: "
f"recall={checkpoint['recall']:.2%} "
f"compose={checkpoint['composition']:.2%} "
f"negative={checkpoint['negative']:.2%} "
f"(train {train_time:.2f}s)"
)
def _last(key: str) -> float:
return curves[key][-1] if curves[key] else 0.0
return EvalResult(
arch_name=system.name,
seed=system.seed,
n_lessons=len(dataset.lessons),
recall_accuracy=_last("recall"),
composition_accuracy=_last("composition"),
composition_3hop_accuracy=_last("composition_3hop"),
negative_accuracy=_last("negative"),
overall_accuracy=_last("overall_applicable"),
curves=curves,
)