Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,12 @@ publish = [
"huggingface_hub>=0.23",
"omegaconf>=2.3",
]
sadeed = [
"pyarabic>=0.6",
"prettytable>=3.9",
"pandas>=2.0",
"pyarrow>=14.0",
]
dev = [
"pytest>=8.0",
"pytest-cov>=5.0",
Expand All @@ -45,10 +51,11 @@ dev = [

[project.scripts]
interscript-ml = "src.cli:main"
interscript-sadeed-eval = "sadeedbench.cli:main"

[tool.setuptools.packages.find]
where = ["src"]
include = ["framework*", "tasks*", "imf*"]
include = ["framework*", "tasks*", "imf*", "sadeedbench*"]

[tool.pytest.ini_options]
testpaths = ["tests"]
Expand Down
26 changes: 7 additions & 19 deletions src/gpu/modal_distill.py
Original file line number Diff line number Diff line change
Expand Up @@ -210,25 +210,13 @@ def _maybe_stitch(spec_id: str, spec: dict, student) -> None:
def paired_bootstrap(deltas: list[float], seed: int = 42, n: int = 1000) -> dict:
"""Sentence-level paired bootstrap over per-item DER deltas
(microkimi eval_compare protocol): a point delta without a CI is
not evidence. Deterministic under the fixed default seed."""
import random
import statistics

if not deltas:
raise ValueError("empty deltas")
rng = random.Random(seed)
size = len(deltas)
means = sorted(
statistics.fmean(deltas[rng.randrange(size)] for _ in range(size))
for _ in range(n)
)
ci = (round(means[int(0.025 * n)], 3), round(means[int(0.975 * n) - 1], 3))
delta = statistics.fmean(deltas)
return {
"delta": round(delta, 4),
"ci95": ci,
"p_leq0": round(sum(1 for m in means if m <= 0) / n, 4),
}
not evidence. Deterministic under the fixed default seed.
Delegates to sadeedbench.bootstrap (the pure shared home)."""
from sadeedbench.bootstrap import bootstrap_means

out = bootstrap_means(deltas, seed=seed, n=n)
out["ci95"] = tuple(round(v, 3) for v in out["ci95"])
return out


def _ensure_src_path() -> None:
Expand Down
4 changes: 4 additions & 0 deletions src/sadeedbench/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
"""SadeedDiac-25 scoring as a standalone tool: the campaign's windowed
DER-CE protocol (zero-skip, haraqat-projected, greedy) packaged so
third parties can re-score any model's predictions — including ONNX
artifacts through any IMF runtime — against the published benchmark."""
47 changes: 47 additions & 0 deletions src/sadeedbench/bootstrap.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
"""Sentence-level paired bootstrap over per-item metric deltas: a
point delta without a CI is not evidence. Deterministic under the
fixed default seed (42, 1,000 resamples, percentile CIs)."""

from __future__ import annotations

import random
import statistics
from collections.abc import Iterable


def bootstrap_means(deltas: list[float], seed: int = 42, n: int = 1000) -> dict:
"""Bootstrap the mean of a delta list (percentile CIs, fixed seed)."""
if not deltas:
raise ValueError("empty deltas")
rng = random.Random(seed)
m = len(deltas)
means = [
statistics.fsum(rng.choice(deltas) for _ in range(m)) / m for _ in range(n)
]
means.sort()
lo = means[int(0.025 * n)]
hi = means[min(n - 1, int(0.975 * n))]
delta = statistics.fsum(deltas) / m
p_leq0 = sum(1 for u in means if u <= 0) / n
return {
"delta": round(delta, 4),
"ci95": (round(lo, 4), round(hi, 4)),
"p_leq0": round(p_leq0, 4),
"n": m,
}


def bootstrap_delta(
a: Iterable[float | None],
b: Iterable[float | None],
seed: int = 42,
n: int = 1000,
) -> dict:
"""Bootstrap the mean of (a - b) over pairwise-aligned items; pairs
with a None on either side are dropped (the aggregate DER call
skips paragraphs with no scorable positions). (candidate,
reference): positive delta = candidate worse."""
deltas = [x - y for x, y in zip(a, b, strict=True) if x is not None and y is not None]
if not deltas:
raise ValueError("no scorable pairs")
return bootstrap_means(deltas, seed=seed, n=n)
77 changes: 77 additions & 0 deletions src/sadeedbench/cli.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,77 @@
"""interscript-sadeed-eval — score predictions on SadeedDiac-25 under
the campaign's windowed DER-CE convention.

interscript-sadeed-eval score \\
--preds final_preds.jsonl --data sadeed-diac-25.parquet \\
[--key student] [--vs reference.jsonl] [--vs-key teacher]

Predictions are read as JSONL rows (any key; --key selects, default
"student" — the training harness's final_preds.jsonl writes
idx/src/teacher/student) or as plain text lines. The parquet carries
the benchmark's input/output columns; fetch SadeedDiac-25 from
https://huggingface.co/datasets/Misraj/SadeedDiac-25. Output is JSON
on stdout; --vs adds a paired-bootstrap delta (candidate minus
reference; positive = candidate worse)."""

from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path


def _read_preds(path: Path, key: str) -> list[str]:
preds: list[str] = []
for line in path.read_text(encoding="utf-8").splitlines():
if not line.strip():
continue
try:
row = json.loads(line)
except json.JSONDecodeError:
preds.append(line)
continue
if not isinstance(row, dict) or key not in row:
raise SystemExit(f"row is not a dict with key {key!r}: {line[:60]}")
preds.append(row[key])
return preds


def _load_gold(path: Path) -> list[str]:
import pandas as pd

table = pd.read_parquet(path)
return table["output"].tolist()


def main(argv: list[str] | None = None) -> int:
parser = argparse.ArgumentParser(prog="interscript-sadeed-eval")
sub = parser.add_subparsers(dest="cmd", required=True)
score = sub.add_parser("score", help="score a predictions file")
score.add_argument("--preds", type=Path, required=True)
score.add_argument("--data", type=Path, required=True,
help="SadeedDiac-25 parquet (input/output columns)")
score.add_argument("--key", default="student",
help="JSONL row key carrying the prediction")
score.add_argument("--vs", type=Path, help="reference predictions file")
score.add_argument("--vs-key", default="teacher")
args = parser.parse_args(argv)

from sadeedbench.bootstrap import bootstrap_delta
from sadeedbench.scoring import per_item_der, score_predictions

gts = _load_gold(args.data)
preds = _read_preds(args.preds, args.key)
result = score_predictions(preds, gts)
if args.vs:
ref = _read_preds(args.vs, args.vs_key)
result["vs"] = bootstrap_delta(
per_item_der(preds, gts), per_item_der(ref, gts)
)
json.dump(result, sys.stdout, ensure_ascii=False)
sys.stdout.write("\n")
return 0


if __name__ == "__main__":
raise SystemExit(main())
48 changes: 48 additions & 0 deletions src/sadeedbench/scoring.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
"""Score predictions against SadeedDiac-25 gold under the campaign's
windowed DER-CE convention.

The convention (docs/paper-a.adoc section 3): every paragraph counts
(zero-skip), haraqat are projected before scoring, and the aggregate is
the Misraj evaluator's Total DER with case endings, reported in
percent. Predictions are whatever a model emitted for the full test
set — same order as the benchmark rows."""

from __future__ import annotations

from typing import Iterable


def score_predictions(preds: list[str], gts: list[str]) -> dict:
"""Aggregate DER-CE over aligned (pred, gt) paragraph pairs."""
from sadeedbench.vendored_sadeed_evaluator import (
ArabicDiacritizationEvaluator,
)

if len(preds) != len(gts):
raise ValueError(f"preds/gts length mismatch: {len(preds)} vs {len(gts)}")
_, _, total_der, _, _ = ArabicDiacritizationEvaluator.caculate_errors_on_sentences(
preds, gts, gt_missing_diacritic_is_error=False
)
return {"der_ce": round(float(total_der), 4), "n": len(preds)}


def per_item_der(preds: list[str], gts: list[str]) -> list[float | None]:
"""Per-paragraph DER for paired bootstrap; None where the paragraph
carries no scorable positions (ZeroDivisionError in the evaluator's
single-item mode — the aggregate call skips those paragraphs)."""
from sadeedbench.vendored_sadeed_evaluator import (
ArabicDiacritizationEvaluator,
)

ders: list[float | None] = []
for pred, gt in zip(preds, gts, strict=True):
try:
_, _, d, _, _ = (
ArabicDiacritizationEvaluator.caculate_errors_on_sentences(
[pred], [gt], gt_missing_diacritic_is_error=False
)
)
ders.append(float(d))
except ZeroDivisionError:
ders.append(None)
return ders
Loading
Loading