From d9449ab4e8ef9c3618e2284312e62b9d09556fa8 Mon Sep 17 00:00:00 2001 From: Ronald Tse Date: Mon, 17 Aug 2026 22:07:47 +0800 Subject: [PATCH] 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. --- src/gpu/modal_distill.py | 49 +++++++++++++++++++++++++++++++++++----- 1 file changed, 43 insertions(+), 6 deletions(-) diff --git a/src/gpu/modal_distill.py b/src/gpu/modal_distill.py index 618033f..8a551a3 100644 --- a/src/gpu/modal_distill.py +++ b/src/gpu/modal_distill.py @@ -277,17 +277,54 @@ def greedy(model, text: str, max_len: int = 256) -> str: out = model.generate(**ids, max_new_tokens=max_len, num_beams=1) return tokenizer.batch_decode(out, skip_special_tokens=True)[0].strip() - def metrics(model) -> dict: - der_sum = cer_sum = n = 0.0 + import random + + def per_pair_der(model) -> list[float]: + ders = [] for src, tgt in pairs: pred = greedy(model, src) gold_n, pred_n = _nikud_only(tgt), _nikud_only(pred) - der_sum += _edit_distance(pred_n, gold_n) / max(1, len(gold_n)) + ders.append( + 100 * _edit_distance(pred_n, gold_n) / max(1, len(gold_n)) + ) + return ders + + def metrics(model) -> dict: + ders = per_pair_der(model) + cer_sum = n = 0.0 + for src, tgt in pairs: + pred = greedy(model, src) cer_sum += _edit_distance(list(pred), list(tgt)) / max(1, len(tgt)) n += 1 - return {"der": round(100 * der_sum / n, 2), "cer": round(100 * cer_sum / n, 2), "n": int(n)} - - return {"teacher": metrics(teacher), "student": metrics(student)} + return { + "der": round(sum(ders) / len(ders), 2), + "cer": round(100 * cer_sum / n, 2), + "n": int(n), + } + + teacher_ders = per_pair_der(teacher) + student_ders = per_pair_der(student) + deltas = [s - t for s, t in zip(student_ders, teacher_ders, strict=True)] + + # Paired bootstrap (microkimi eval_compare protocol): sentences are + # the resampling unit; a point-estimate delta without a CI is not + # evidence the student regressed. + rng = random.Random(42) + means = [] + for _ in range(1000): + sample = [deltas[rng.randrange(len(deltas))] for _ in deltas] + means.append(sum(sample) / len(sample)) + means.sort() + ci = (round(means[24], 3), round(means[974], 3)) + p_value = sum(1 for m in means if m <= 0) / len(means) + + return { + "teacher": metrics(teacher), + "student": metrics(student), + "paired_delta_pp": round(sum(deltas) / len(deltas), 3), + "bootstrap_ci95": ci, + "p_value_student_worse": round(p_value, 4), + } @app.local_entrypoint()