Skip to content

Commit d9449ab

Browse files
author
Ronald Tse
committed
feat(distill): paired bootstrap CI on the before/after DER
microkimi's eval_compare protocol: sentences are the resampling unit, 1000 bootstrap resamples, 95% CI + one-sided p-value for 'student is worse'. A point-estimate DER delta alone is not evidence of a regression; the interval is.
1 parent 2e692f4 commit d9449ab

1 file changed

Lines changed: 43 additions & 6 deletions

File tree

‎src/gpu/modal_distill.py‎

Lines changed: 43 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -277,17 +277,54 @@ def greedy(model, text: str, max_len: int = 256) -> str:
277277
out = model.generate(**ids, max_new_tokens=max_len, num_beams=1)
278278
return tokenizer.batch_decode(out, skip_special_tokens=True)[0].strip()
279279

280-
def metrics(model) -> dict:
281-
der_sum = cer_sum = n = 0.0
280+
import random
281+
282+
def per_pair_der(model) -> list[float]:
283+
ders = []
282284
for src, tgt in pairs:
283285
pred = greedy(model, src)
284286
gold_n, pred_n = _nikud_only(tgt), _nikud_only(pred)
285-
der_sum += _edit_distance(pred_n, gold_n) / max(1, len(gold_n))
287+
ders.append(
288+
100 * _edit_distance(pred_n, gold_n) / max(1, len(gold_n))
289+
)
290+
return ders
291+
292+
def metrics(model) -> dict:
293+
ders = per_pair_der(model)
294+
cer_sum = n = 0.0
295+
for src, tgt in pairs:
296+
pred = greedy(model, src)
286297
cer_sum += _edit_distance(list(pred), list(tgt)) / max(1, len(tgt))
287298
n += 1
288-
return {"der": round(100 * der_sum / n, 2), "cer": round(100 * cer_sum / n, 2), "n": int(n)}
289-
290-
return {"teacher": metrics(teacher), "student": metrics(student)}
299+
return {
300+
"der": round(sum(ders) / len(ders), 2),
301+
"cer": round(100 * cer_sum / n, 2),
302+
"n": int(n),
303+
}
304+
305+
teacher_ders = per_pair_der(teacher)
306+
student_ders = per_pair_der(student)
307+
deltas = [s - t for s, t in zip(student_ders, teacher_ders, strict=True)]
308+
309+
# Paired bootstrap (microkimi eval_compare protocol): sentences are
310+
# the resampling unit; a point-estimate delta without a CI is not
311+
# evidence the student regressed.
312+
rng = random.Random(42)
313+
means = []
314+
for _ in range(1000):
315+
sample = [deltas[rng.randrange(len(deltas))] for _ in deltas]
316+
means.append(sum(sample) / len(sample))
317+
means.sort()
318+
ci = (round(means[24], 3), round(means[974], 3))
319+
p_value = sum(1 for m in means if m <= 0) / len(means)
320+
321+
return {
322+
"teacher": metrics(teacher),
323+
"student": metrics(student),
324+
"paired_delta_pp": round(sum(deltas) / len(deltas), 3),
325+
"bootstrap_ci95": ci,
326+
"p_value_student_worse": round(p_value, 4),
327+
}
291328

292329

293330
@app.local_entrypoint()

0 commit comments

Comments
 (0)