| from __future__ import annotations |
|
|
| import json |
|
|
| import numpy as np |
| import pytest |
|
|
| from scripts.stages.evaluate_sp01_mlcommons_kws import ( |
| canonical_digest, |
| load_checkpoint, |
| official_int8_input, |
| quality_decision, |
| ) |
|
|
|
|
| def test_official_int8_input_matches_scale_offset_cast_without_local_ptq() -> None: |
| feature = np.asarray([-2.0, 0.0, 1.0, 2.0], dtype=np.float32) |
| actual = official_int8_input(feature, 0.5, 3) |
| assert actual.dtype == np.int8 |
| assert actual.tolist() == [-1, 3, 5, 7] |
|
|
|
|
| def test_quality_decision_uses_inclusive_mlcommons_threshold() -> None: |
| assert quality_decision(9, 10, 0.9)["threshold_met"] is True |
| assert quality_decision(899, 1000, 0.9)["threshold_met"] is False |
|
|
|
|
| def test_checkpoint_requires_exact_fingerprint(tmp_path) -> None: |
| path = tmp_path / "checkpoint.jsonl" |
| path.write_text(json.dumps({"fingerprint": "one", "index": 0, "label_id": 2}) + "\n") |
| assert load_checkpoint(path, "one")[0]["label_id"] == 2 |
| with pytest.raises(ValueError, match="fingerprint mismatch"): |
| load_checkpoint(path, "two") |
|
|
|
|
| def test_canonical_digest_is_key_order_independent() -> None: |
| assert canonical_digest({"a": 1, "b": 2}) == canonical_digest({"b": 2, "a": 1}) |
|
|