Skip to content

Commit f93ac41

Browse files
Ronald Tseronaldtse
authored andcommitted
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 37022b4 commit f93ac41

1 file changed

Lines changed: 42 additions & 5 deletions

File tree

src/gpu/modal_distill.py

Lines changed: 42 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -485,17 +485,54 @@ def greedy(model, text: str, max_len: int = 256) -> str:
485485
out = model.generate(**ids, max_new_tokens=max_len, num_beams=1)
486486
return tokenizer.batch_decode(out, skip_special_tokens=True)[0].strip()
487487

488-
def metrics(model) -> dict:
489-
der_sum = cer_sum = n = 0.0
488+
import random
489+
490+
def per_pair_der(model) -> list[float]:
491+
ders = []
490492
for src, tgt in pairs:
491493
pred = greedy(model, src)
492494
gold_n, pred_n = _nikud_only(tgt), _nikud_only(pred)
493-
der_sum += _edit_distance(pred_n, gold_n) / max(1, len(gold_n))
495+
ders.append(
496+
100 * _edit_distance(pred_n, gold_n) / max(1, len(gold_n))
497+
)
498+
return ders
499+
500+
def metrics(model) -> dict:
501+
ders = per_pair_der(model)
502+
cer_sum = n = 0.0
503+
for src, tgt in pairs:
504+
pred = greedy(model, src)
494505
cer_sum += _edit_distance(list(pred), list(tgt)) / max(1, len(tgt))
495506
n += 1
496-
return {"der": round(100 * der_sum / n, 2), "cer": round(100 * cer_sum / n, 2), "n": int(n)}
507+
return {
508+
"der": round(sum(ders) / len(ders), 2),
509+
"cer": round(100 * cer_sum / n, 2),
510+
"n": int(n),
511+
}
497512

498-
return {"teacher": metrics(teacher), "student": metrics(student)}
513+
teacher_ders = per_pair_der(teacher)
514+
student_ders = per_pair_der(student)
515+
deltas = [s - t for s, t in zip(student_ders, teacher_ders, strict=True)]
516+
517+
# Paired bootstrap (microkimi eval_compare protocol): sentences are
518+
# the resampling unit; a point-estimate delta without a CI is not
519+
# evidence the student regressed.
520+
rng = random.Random(42)
521+
means = []
522+
for _ in range(1000):
523+
sample = [deltas[rng.randrange(len(deltas))] for _ in deltas]
524+
means.append(sum(sample) / len(sample))
525+
means.sort()
526+
ci = (round(means[24], 3), round(means[974], 3))
527+
p_value = sum(1 for m in means if m <= 0) / len(means)
528+
529+
return {
530+
"teacher": metrics(teacher),
531+
"student": metrics(student),
532+
"paired_delta_pp": round(sum(deltas) / len(deltas), 3),
533+
"bootstrap_ci95": ci,
534+
"p_value_student_worse": round(p_value, 4),
535+
}
499536

500537

501538

0 commit comments

Comments
 (0)