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
30 changes: 30 additions & 0 deletions docs/adr/0001-run-scripts-are-provenance-artifacts.adoc
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
= ADR-0001: Run scripts are provenance artifacts
Status: Accepted (2026-08-29)

== Context

The windowed zero-skip Sadeed protocol (split at word boundaries, 2x
window generation cap, haraqat projection) was pasted verbatim into
`train_arabic_r5.py` through `train_arabic_r8.py`. `docs/RESULTS.md`
cites those scripts as the record of what produced each published
number. ml-models now owns one deep, tested copy of the protocol
(`src/harness/sadeed.py`), vendored here as `sadeed_harness.py`.

== Decision

- Past run scripts (r5-r8, eval_*) are **never rewritten** to consume
the shared harness: they are provenance artifacts. Their inline
copies stand as historical record of what ran.
- New run scripts (r9+) import `sadeed_harness` — no new inline copies
of the protocol.
- Protocol changes land in `sadeed_harness.py` (kept in sync with
ml-models `src/harness/sadeed.py`) and are noted in the RESULTS
section of any future run.

== Consequences

- The four inline copies remain in the tree deliberately; architecture
reviews should not re-suggest rewriting them.
- The harness carries its own unit tests in ml-models; a vendored copy
that drifts from the source is a bug — sync both or vendor via
submodule later.
106 changes: 106 additions & 0 deletions sadeed_harness.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,106 @@
"""The Sadeed windowed harness — the campaign's central measurement
protocol, in one place.

Windowed zero-skip (the published protocol): inputs over the byte
budget split at word boundaries; greedy decode with a 2x-window
generation cap (diacritized output runs 1.4-1.6x input — a shorter cap
silently truncates the hardest paragraphs, which the evaluator then
skips: survivorship bias, not quality); window predictions stitched and
haraqat projected onto the input letters so output structure always
matches ground truth.

Previously pasted verbatim into rababa train_arabic_r5/r6/r7/r8 and
modal_distill.evaluate_der — five copies of a protocol every published
Arabic number depends on. rababa scripts should vendor this single
module instead of carrying inline copies.
"""

from __future__ import annotations

import re

DIACRITICS_RE = re.compile("[ؐ-ًؚ-ٰٟۖ-ۜ۟-۪ۨ-ۭ]")


def split_windows(text: str, budget: int = 1400) -> list[str]:
"""Split at word boundaries so no window exceeds the byte budget."""
if len(text.encode("utf-8")) <= budget:
return [text]
windows: list[str] = []
current: list[str] = []
n = 0
for word in text.split():
cost = len(word.encode("utf-8")) + 1
if current and n + cost > budget:
windows.append(" ".join(current))
current, n = [], 0
current.append(word)
n += cost
if current:
windows.append(" ".join(current))
return windows


def project_haraqat(pred: str, text: str) -> str:
"""Project predicted haraqat onto the input's letters, so the output
structure matches ground truth even when the prediction inserts or
drops letters (zero-skip contract)."""
from difflib import SequenceMatcher

pred_haraqat = [""]
for ch in pred:
if DIACRITICS_RE.match(ch):
pred_haraqat[-1] += ch
else:
pred_haraqat.append("")
pred_haraqat = pred_haraqat[1:]
pred_letters = [c for c in pred if not DIACRITICS_RE.match(c)]
text_letters = [c for c in text if not DIACRITICS_RE.match(c)]
sm = SequenceMatcher(None, text_letters, pred_letters, autojunk=False)
out: list[str] = []
for op, i1, i2, j1, _j2 in sm.get_opcodes():
if op == "equal":
for k in range(i2 - i1):
out.append(text_letters[i1 + k] + pred_haraqat[j1 + k])
else:
for k in range(i1, i2):
out.append(text_letters[k])
return "".join(out)


def strip_diacritics(text: str) -> str:
return DIACRITICS_RE.sub("", text)


def windowed_paragraphs(model, tokenizer, inputs, window: int = 1400,
batch_size: int = 8, device: str = "cuda") -> list[str]:
"""The full protocol over a paragraph list: split, greedy-decode each
window (2x-window cap) under bf16 autocast, stitch, project. Returns
haraqat-projected paragraphs ready for the Misraj evaluator."""
import torch

windows: list[str] = []
counts: list[int] = []
for text in inputs:
ws = split_windows(text, window)
counts.append(len(ws))
windows.extend(ws)

preds: list[str] = []
with torch.no_grad():
for i in range(0, len(windows), batch_size):
batch = windows[i : i + batch_size]
enc = tokenizer(
batch, return_tensors="pt", padding=True, truncation=True,
max_length=window,
).to(device)
with torch.autocast("cuda", torch.bfloat16):
gen = model.generate(**enc, max_new_tokens=window * 2, num_beams=1)
preds.extend(tokenizer.batch_decode(gen, skip_special_tokens=True))

paragraphs: list[str] = []
k = 0
for text, c in zip(inputs, counts, strict=True):
paragraphs.append(project_haraqat(" ".join(preds[k : k + c]), text))
k += c
return paragraphs
Loading