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
7 changes: 6 additions & 1 deletion src/gpu/modal_export.py
Original file line number Diff line number Diff line change
Expand Up @@ -294,7 +294,12 @@ def stage(event: str) -> None:
MODELS_VOLUME.commit()

stage(f"start model={model_id} pairs={len(pairs)}")
reference = reference_decode(model, [src for src, _ in pairs], max_len=128)
reference = reference_decode(
model,
[src for src, _ in pairs],
max_len=128,
resume_path=Path("/outputs/imf") / model_id / "reference_progress.jsonl",
)
stage("reference-decode done")

out_dir = Path("/outputs/imf") / model_id
Expand Down
34 changes: 32 additions & 2 deletions src/imf/parity.py
Original file line number Diff line number Diff line change
Expand Up @@ -118,9 +118,39 @@ def _sessions_from_zip(zip_path: Path):
return enc, dec


def reference_decode(model, sources, max_len: int = 256) -> list[list[int]]:
def reference_decode(model, sources, max_len: int = 256, resume_path=None):
"""Torch-reference greedy decode of many inputs, computed once and
shared across precision variants by run_parity."""
shared across precision variants by run_parity.

resume_path: append-only JSONL of per-input results. Container-level
kills (variable-lifetime, traceback-less — observed five times on
ara-diac2) then cost only the tail: a relaunch resumes from the
saved prefix instead of redoing hours of decode."""
if resume_path is not None:
import json as _json
from pathlib import Path as _Path

resume_path = _Path(resume_path)
done: dict[int, list[int]] = {}
if resume_path.exists():
for line in resume_path.read_text(encoding="utf-8").splitlines():
if line.strip():
row = _json.loads(line)
done[row["i"]] = row["tokens"]
results: list[list[int]] = [[]] * len(sources)
with resume_path.open("a", encoding="utf-8") as fh:
for i, source in enumerate(sources):
if i in done:
results[i] = done[i]
continue
tokens = _torch_greedy_tokens(model, source, max_len)
results[i] = tokens
fh.write(
_json.dumps({"i": i, "tokens": tokens}, ensure_ascii=False) + "\n"
)
if i % 100 == 0:
fh.flush()
return results
return [_torch_greedy_tokens(model, source, max_len) for source in sources]


Expand Down
Loading