Skip to content

Commit 6bf4a87

Browse files
author
Ronald Tse
committed
fix(lint): ruff clean on main; drop MTP save lines from logit-mode distill
The F821 was a real latent bug: MTP checkpoint lines from the #114 wiring were pasted into the logit-mode distiller save block where mtp_head never exists — any logit-mode run would NameError at its first step checkpoint. The lines belong only in distill_sequence. The rest: import sorting, unused zipfile, quoted annotation, E501 signature wraps, zip strict=, context manager for models.yaml.
1 parent ac1e04e commit 6bf4a87

10 files changed

Lines changed: 22 additions & 18 deletions

File tree

‎benchmarks/imf-runtime/modal-bench.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -69,6 +69,7 @@ def main(
6969
if not filename:
7070
import yaml
7171

72-
index = yaml.safe_load(open("models.yaml", encoding="utf-8"))
72+
with open("models.yaml", encoding="utf-8") as fh:
73+
index = yaml.safe_load(fh)
7374
filename = index["models"][model_id]["filename"]
7475
print(bench.remote(model_id, filename))

‎runtime/src/interscript_ml/model.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,8 +6,8 @@
66

77
import numpy as np
88

9-
from interscript_ml.tokens import EOS_ID, PAD_ID, decode, encode
109
from interscript_ml.loader import load_manifest, verify_and_read
10+
from interscript_ml.tokens import EOS_ID, PAD_ID, decode, encode
1111

1212

1313
class Model:
@@ -45,7 +45,7 @@ def __init__(self, zip_path: Path | str):
4545
self._output_names = [o.name for o in self._decoder.get_outputs()]
4646

4747
@classmethod
48-
def load(cls, path_or_id: Path | str, index_url: str | None = None) -> "Model":
48+
def load(cls, path_or_id: Path | str, index_url: str | None = None) -> Model:
4949
"""Accepts a zip path OR a model id from models.yaml (dynamic
5050
fetch: download -> verify -> cache)."""
5151
candidate = str(path_or_id)

‎runtime/tests/test_model.py‎

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -19,11 +19,10 @@
1919
ort = pytest.importorskip("onnxruntime")
2020
onnx = pytest.importorskip("onnx")
2121

22+
import numpy as np # noqa: E402
2223
from interscript_ml import Model, ModelFormatError, decode, encode # noqa: E402
2324
from onnx import TensorProto, helper, numpy_helper # noqa: E402
2425

25-
import numpy as np # noqa: E402
26-
2726

2827
def _graph(opset: int = 14) -> bytes:
2928
graph = helper.make_graph(

‎runtime/tests/test_registry.py‎

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -3,17 +3,14 @@
33
from __future__ import annotations
44

55
import hashlib
6-
import zipfile
6+
import os # noqa: E402
77
from pathlib import Path
88

99
import pytest
1010
import yaml
11-
1211
from interscript_ml.registry import RegistryError, resolve
1312
from tests_helpers import build_tiny_zip
1413

15-
import os # noqa: E402
16-
1714

1815
def _index_file(tmp_path: Path, zip_path: Path, sha256: str | None = None) -> Path:
1916
index = {

‎scripts/publish_model.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -30,9 +30,10 @@
3030
sys.path.insert(0, str(REPO_ROOT / "src"))
3131
sys.path.insert(0, str(REPO_ROOT / "scripts"))
3232

33-
from imf.validator import validate_zip # noqa: E402
3433
from split_release import split as split_zip # noqa: E402
3534

35+
from imf.validator import validate_zip # noqa: E402
36+
3637
# GitHub hard-caps release assets at 2,147,483,648 bytes; split well below.
3738
SPLIT_THRESHOLD = 2_000_000_000
3839
DEFAULT_REPO = "interscript/interscript-ml"

‎src/gpu/modal_distill.py‎

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -429,8 +429,6 @@ def val_loss() -> float:
429429
ck.mkdir(exist_ok=True)
430430
torch.save(student.state_dict(), ck / "student.pt")
431431
torch.save(optimizer.state_dict(), ck / "optim.pt")
432-
if mtp_head is not None:
433-
torch.save(mtp_head.state_dict(), ck / "mtp_head.pt")
434432
CHECKPOINTS.commit()
435433

436434
vl = val_loss()

‎src/gpu/modal_export.py‎

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -356,7 +356,9 @@ def stage(event: str) -> None:
356356
timeout=5 * 3600,
357357
volumes={**CHECKPOINT_VOLUMES, **DATASET_VOLUMES, "/outputs": MODELS_VOLUME},
358358
)
359-
def margin_model(model_id: str, precisions: list[str], limit: int = 0, dump_positions: bool = False) -> dict[str, str]:
359+
def margin_model(
360+
model_id: str, precisions: list[str], limit: int = 0, dump_positions: bool = False
361+
) -> dict[str, str]:
360362
"""Margin analysis alone over already-exported zips — read-only for the
361363
zips (diagnostic JSON only); validates published artifacts without
362364
touching their metadata."""
@@ -414,7 +416,9 @@ def parity(model: str, precisions: str = "fp32,fp16,int8", limit: int = 0) -> No
414416

415417

416418
@app.local_entrypoint()
417-
def margins(model: str, precisions: str = "fp32,fp16,int8", limit: int = 0, dump_positions: bool = False) -> None:
419+
def margins(
420+
model: str, precisions: str = "fp32,fp16,int8", limit: int = 0, dump_positions: bool = False
421+
) -> None:
418422
reports = margin_model.remote(model, precisions.split(","), limit, dump_positions)
419423
for precision, status in reports.items():
420424
print(f"{model} [{precision}] {status}")

‎src/gpu/mtp.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,8 +11,8 @@
1111
from __future__ import annotations
1212

1313
import torch
14-
from torch import nn
1514
import torch.nn.functional as F
15+
from torch import nn
1616

1717

1818
class MTPHead(nn.Module):

‎src/imf/parity.py‎

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -253,7 +253,9 @@ def _onnx_forced_logits(enc_sess, dec_sess, source: str, target_ids: list[int]):
253253
return dict(zip(out_names, out, strict=True))["logits"][0]
254254

255255

256-
def run_margin_analysis(model, zip_path, pairs, max_len: int = 256, dump_positions: Path | str | None = None) -> MarginReport:
256+
def run_margin_analysis(
257+
model, zip_path, pairs, max_len: int = 256, dump_positions: Path | str | None = None
258+
) -> MarginReport:
257259
"""pairs: iterable of (source_text, gold_target) — the same probe set
258260
the CER parity gate uses. Teacher-forces both sides and measures the
259261
argmax flip rate, reference top1−top2 margin quantiles, and KL
@@ -294,7 +296,7 @@ def run_margin_analysis(model, zip_path, pairs, max_len: int = 256, dump_positio
294296
dump = Path(dump_positions)
295297
dump.parent.mkdir(parents=True, exist_ok=True)
296298
with dump.open("w", encoding="utf-8") as out:
297-
for pair_idx, (m, f) in enumerate(zip(margin_chunks, flip_chunks)):
299+
for pair_idx, (m, f) in enumerate(zip(margin_chunks, flip_chunks, strict=True)):
298300
out.write(_json.dumps(
299301
{"pair": pair_idx, "tokens": int(f.size),
300302
"flip_positions": [int(x) for x in np.nonzero(f)[0]],

‎tests/test_imf_parity.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -168,7 +168,9 @@ def test_margin_report_json_roundtrip(gated_zip: Path, fixture_model, tmp_path:
168168
assert data["flip_rate"] == report.flip_rate
169169

170170

171-
def test_margin_analysis_dumps_per_pair_positions(gated_zip: Path, fixture_model, tmp_path: Path) -> None:
171+
def test_margin_analysis_dumps_per_pair_positions(
172+
gated_zip: Path, fixture_model, tmp_path: Path
173+
) -> None:
172174
"""TODO.training-work/05: the flip bootstrap needs per-pair token and
173175
flip counts, not just aggregates."""
174176
dump = tmp_path / "positions.jsonl"

0 commit comments

Comments
 (0)