From f16e68c4cb44f1fcd2c86ada124c4a2c82eab03c Mon Sep 17 00:00:00 2001 From: IIIllllIlIlllII <306726845+IIIllllIlIlllII@users.noreply.github.com> Date: Thu, 27 Aug 2026 13:45:41 +0000 Subject: [PATCH 1/2] Recalibrate v8 raw parity thresholds on 50 seeds The v8 raw-logit thresholds were calibrated on five deterministic seeds. Each test compares about five million values, so the observed maximum is an extreme-value statistic that keeps growing with the seed count: it rises by 60% between 5 and 50 seeds on both branches. Both branches exceed their original threshold before 50 seeds, so check_onnx_parity.py --num-tests 20 failed on an otherwise healthy export: o2m max@5 1.741886e-03 max@50 2.788305e-03 threshold 2.0e-03 nms-free max@5 2.990723e-03 max@50 4.793882e-03 threshold 3.5e-03 Set both thresholds to twice the maximum observed over 50 seeds and record the measurement in the validation report. Co-Authored-By: Claude Opus 5 --- docs/onnx/tr_hash_v8_validation_report.md | 38 ++++++++++++++++++----- scripts/check_onnx_parity.py | 10 ++++-- tests/test_detector_export.py | 4 +-- 3 files changed, 40 insertions(+), 12 deletions(-) diff --git a/docs/onnx/tr_hash_v8_validation_report.md b/docs/onnx/tr_hash_v8_validation_report.md index 5750c66..0f95250 100644 --- a/docs/onnx/tr_hash_v8_validation_report.md +++ b/docs/onnx/tr_hash_v8_validation_report.md @@ -119,14 +119,8 @@ Decoded-output drift is substantially lower: | O2M | `6.00814819e-05` | `3.02493572e-05` | | NMS-free | `5.73396683e-05` | `1.35712326e-05` | -The calibrated v8 raw-logit thresholds are: - -| Branch | Raw-logit tolerance | Justification | -|---|---:|---| -| O2M | `0.002` | Covers observed max drift up to `1.74e-03` with decoded box drift near `6e-05` | -| NMS-free | `0.0035` | Covers observed max drift up to `3.00e-03` with decoded box drift near `5.7e-05` | - -Calibrated parity output: +The thresholds first derived from these five seeds were `0.002` (O2M) and +`0.0035` (NMS-free): ```text O2M, tolerance 0.002: @@ -146,6 +140,34 @@ NMS-free, tolerance 0.0035: Parity PASSED: branch=nms-free, tolerance=0.0035 ``` +### Sample-size sensitivity + +Five seeds are not enough to bound the maximum. Each test compares roughly five +million values, so the observed maximum is an extreme-value statistic that keeps +growing with the seed count. Re-measured on the same checkpoint: + +| Branch | max over 5 seeds | max over 50 seeds | Growth | Old threshold | +|---|---:|---:|---:|---:| +| O2M | `1.741886e-03` | `2.788305e-03` | `+60.1%` | `2.0e-03` | +| NMS-free | `2.990723e-03` | `4.793882e-03` | `+60.3%` | `3.5e-03` | + +Both branches exceed their original threshold well before 50 seeds, so +`check_onnx_parity.py --num-tests 20` failed on an otherwise healthy export. +The thresholds are therefore twice the maximum observed over 50 seeds: + +| Branch | Observed max (50 seeds) | Threshold | Headroom | +|---|---:|---:|---:| +| O2M | `2.788305e-03` | `6.0e-03` | `2.2x` | +| NMS-free | `4.793882e-03` | `1.0e-02` | `2.1x` | + +Reproduce either column by running the parity check with `--num-tests 5` and +`--num-tests 50`. + +Measured with PyTorch `2.13.0+cu130`, ONNX `1.21.0`, ONNX Runtime `1.24.4` on +`CPUExecutionProvider`. The five-seed maxima reproduce the values above to three +significant digits despite the different runtime versions, so this drift +originates in graph operation ordering rather than in the runtime build. + ## Benchmarks Benchmarks used batch size 1, `10` warmup runs, and `50` measured runs. diff --git a/scripts/check_onnx_parity.py b/scripts/check_onnx_parity.py index 908f02f..a4566f4 100644 --- a/scripts/check_onnx_parity.py +++ b/scripts/check_onnx_parity.py @@ -21,9 +21,15 @@ from complexity.generative.detection.hub import load_detector_checkpoint DEFAULT_PARITY_TOLERANCE = 1e-4 + +# Each test compares about five million values, so the observed maximum is an +# extreme-value statistic that keeps growing with the seed count: it rises by +# 60% between 5 and 50 seeds on both branches. The thresholds below are twice +# the maximum observed over 50 deterministic seeds. +# See docs/onnx/tr_hash_v8_validation_report.md. V8_PARITY_TOLERANCES = { - "o2m": 2e-3, - "nms-free": 3.5e-3, + "o2m": 6e-3, + "nms-free": 1e-2, } diff --git a/tests/test_detector_export.py b/tests/test_detector_export.py index 060588d..5189611 100644 --- a/tests/test_detector_export.py +++ b/tests/test_detector_export.py @@ -85,8 +85,8 @@ def test_auto_export_selects_the_production_branch() -> None: @pytest.mark.parametrize( ("branch", "expected"), ( - ("o2m", 2e-3), - ("nms-free", 3.5e-3), + ("o2m", 6e-3), + ("nms-free", 1e-2), ), ) def test_v8_exports_use_branch_calibrated_parity_tolerances( From 3dcdab5ba55eaeecb5902055d9ff0e524bd1aa2c Mon Sep 17 00:00:00 2001 From: IIIllllIlIlllII <306726845+IIIllllIlIlllII@users.noreply.github.com> Date: Thu, 27 Aug 2026 13:47:36 +0000 Subject: [PATCH 2/2] Add decoded-output parity gates for Vision v8 ONNX exports Raw-logit tolerance only approximates deployment behaviour: softmax absorbs a translation shared by all DFL bins, while the stride amplifies drift on coarse grids, so raw and decoded drift are not monotonically related. check_onnx_parity.py now evaluates three independent gates on the same deterministic inputs: raw logits, normalized decoded boxes, and sigmoid quality-class scores. Each reports its own max and mean difference, fails on its own, and is named in the summary when it fails. Decoding goes through complexity.deploy.onnx_detector so the gates exercise the deployment code path instead of a separate implementation. Decoded thresholds are twice the maximum observed over 50 deterministic seeds. decoded-box is shared across branches (their box drift differs by 19%); decoded-score is branch-specific (o2m drift is 2.2x nms-free). Exports whose sidecar carries no v8 decode metadata skip the decoded gates rather than failing them, and keep the strict 1e-4 raw threshold. Closes #12 Co-Authored-By: Claude Opus 5 --- docs/onnx/tr_hash_v8_validation_report.md | 84 ++++++- scripts/check_onnx_parity.py | 268 +++++++++++++++++++--- tests/test_detector_export.py | 129 ++++++++++- 3 files changed, 439 insertions(+), 42 deletions(-) diff --git a/docs/onnx/tr_hash_v8_validation_report.md b/docs/onnx/tr_hash_v8_validation_report.md index 0f95250..9514167 100644 --- a/docs/onnx/tr_hash_v8_validation_report.md +++ b/docs/onnx/tr_hash_v8_validation_report.md @@ -168,6 +168,84 @@ Measured with PyTorch `2.13.0+cu130`, ONNX `1.21.0`, ONNX Runtime `1.24.4` on significant digits despite the different runtime versions, so this drift originates in graph operation ordering rather than in the runtime build. +## Decoded Parity Gates + +Raw-logit tolerance is only a proxy for deployment behaviour, and the section +above shows it is also the least stable quantity to gate on. `check_onnx_parity.py` +therefore evaluates three independent gates on the same deterministic inputs: + +| Gate | Compares | Role | +|---|---|---| +| `raw` | exported logits | coarse export-integrity check | +| `decoded-box` | normalized xyxy boxes after LTRB/DFL decode | deployment guarantee | +| `decoded-score` | sigmoid quality-class scores | deployment guarantee | + +Each gate reports its own max and mean difference and fails independently, and +the CLI names the failing gates in its summary. Decoding uses +`complexity.deploy.onnx_detector`, so the gates exercise the deployment code +path rather than a separate reimplementation. + +### Units + +Decoded box drift is reported in **normalized** box coordinates, matching the +`box_norm` output of the deployment pipeline. The equivalent input-pixel drift +is `640x` larger: a normalized `5.50e-05` is `3.52e-02` px on a 640px input. +The five-seed table under Parity Results is also in normalized coordinates. + +### Stability and thresholds + +Decoded outputs are markedly more stable than raw logits under resampling, +which is the quantitative case for gating on them: + +| Branch | Gate | max over 5 seeds | max over 50 seeds | Growth | +|---|---|---:|---:|---:| +| O2M | decoded-box | `4.904270e-05` | `5.497933e-05` | `+12.1%` | +| O2M | decoded-score | `2.907962e-05` | `3.811717e-05` | `+31.1%` | +| NMS-free | decoded-box | `5.540848e-05` | `6.532669e-05` | `+17.9%` | +| NMS-free | decoded-score | `1.281500e-05` | `1.719594e-05` | `+34.2%` | + +Compare with `+60%` for raw logits on both branches. Thresholds follow the same +rule as the raw gate: twice the maximum observed over 50 deterministic seeds. + +| Branch | Gate | Observed max (50 seeds) | Threshold | Headroom | +|---|---|---:|---:|---:| +| O2M | decoded-box | `5.497933e-05` | `1.3e-04` | `2.4x` | +| O2M | decoded-score | `3.811717e-05` | `8.0e-05` | `2.1x` | +| NMS-free | decoded-box | `6.532669e-05` | `1.3e-04` | `2.0x` | +| NMS-free | decoded-score | `1.719594e-05` | `4.0e-05` | `2.3x` | + +`decoded-box` shares one threshold across branches because both use the same +regression head and their box drift differs by only 19%. `decoded-score` is +branch-specific because O2M score drift is 2.2x that of NMS-free. + +Legacy exports are unaffected: when the ONNX sidecar carries no v8 decode +metadata, the decoded gates are skipped rather than failed, and the raw gate +keeps its strict `1e-4` threshold. `--skip-decoded` forces that behaviour on any +export. + +### Full gate output + +```bash +PYTHONPATH=. python scripts/check_onnx_parity.py models/TR-HASH-Vision-v8-2M-COCO-SFT tr_hash_v8_o2m.onnx --num-tests 50 +PYTHONPATH=. python scripts/check_onnx_parity.py models/TR-HASH-Vision-v8-2M-COCO-SFT/best_nms_free tr_hash_v8_nms_free.onnx --num-tests 50 +``` + +```text +Parity PASSED: branch=o2m + raw tol=6.00e-03 worst_max=2.79e-03 [PASS] + decoded-box tol=1.30e-04 worst_max=5.50e-05 [PASS] + decoded-score tol=8.00e-05 worst_max=3.81e-05 [PASS] + +Parity PASSED: branch=nms-free + raw tol=1.00e-02 worst_max=4.79e-03 [PASS] + decoded-box tol=1.30e-04 worst_max=6.53e-05 [PASS] + decoded-score tol=4.00e-05 worst_max=1.72e-05 [PASS] +``` + +The decoded values differ slightly from the five-seed table under Parity +Results because they are now computed with the shared `onnx_detector` decoder +rather than an ad-hoc one. + ## Benchmarks Benchmarks used batch size 1, `10` warmup runs, and `50` measured runs. @@ -195,5 +273,7 @@ NMS-free was `0.858 ms` slower than O2M on mean latency, a `2.60%` increase. TR-HASH Vision v8 exports successfully to ONNX for both raw prediction branches. Branch behavior matches the expected architecture: NMS-free computes the extra one-to-one head path and is empirically slower than O2M. Deployment -validation should use decoded-output drift and the calibrated v8 raw-logit -thresholds above rather than the legacy `1e-4` raw-logit threshold. +validation should rely on the `decoded-box` and `decoded-score` gates, which +measure the outputs deployments actually consume and are three to five times +more stable across seeds; the `raw` gate remains only as an export-integrity +check. diff --git a/scripts/check_onnx_parity.py b/scripts/check_onnx_parity.py index a4566f4..b4da661 100644 --- a/scripts/check_onnx_parity.py +++ b/scripts/check_onnx_parity.py @@ -1,5 +1,17 @@ """Verify numerical parity between a detector checkpoint and its ONNX export. +Three independent gates are evaluated on the same deterministic inputs: + +* ``raw`` compares the exported logits directly. It is a coarse export-integrity + check: raw drift concentrates in fine-grid regression logits and its observed + maximum is unstable across seeds. +* ``decoded-box`` compares normalized xyxy boxes after LTRB/DFL decode. +* ``decoded-score`` compares sigmoid quality-class scores. + +The decoded gates describe deployment behaviour and carry the real guarantee. +They are only available when the ONNX sidecar exposes v8 decode metadata; +legacy exports keep their strict raw-only behaviour. + Usage: python scripts/check_onnx_parity.py CHECKPOINT model.onnx python scripts/check_onnx_parity.py CHECKPOINT model.onnx --num-tests 10 @@ -10,6 +22,7 @@ import argparse import json import os +from dataclasses import dataclass, replace from pathlib import Path import numpy as np @@ -17,22 +30,59 @@ os.environ["COMPLEXITY_DISABLE_KERNELS"] = "1" +from complexity.deploy.onnx_detector.dfl import decode_dfl_boxes +from complexity.deploy.onnx_detector.grid import GridGeometry, generate_grid_geometry +from complexity.deploy.onnx_detector.metadata import ( + OnnxDetectorMetadata, + metadata_from_mapping, + validate_output_shape, +) from complexity.generative.detection.exporting import ExportBranch, RawDetectorExport from complexity.generative.detection.hub import load_detector_checkpoint +RAW_GATE = "raw" +DECODED_BOX_GATE = "decoded-box" +DECODED_SCORE_GATE = "decoded-score" + DEFAULT_PARITY_TOLERANCE = 1e-4 -# Each test compares about five million values, so the observed maximum is an -# extreme-value statistic that keeps growing with the seed count: it rises by -# 60% between 5 and 50 seeds on both branches. The thresholds below are twice -# the maximum observed over 50 deterministic seeds. -# See docs/onnx/tr_hash_v8_validation_report.md. + +@dataclass(frozen=True) +class ParityTolerances: + """Per-gate absolute tolerances; ``None`` disables a decoded gate.""" + + raw: float + decoded_box: float | None = None + decoded_score: float | None = None + + +# Calibrated on AETHORIA-AI/TR-HASH-Vision-v8-2M-COCO-SFT over 50 deterministic +# seeds, at twice the observed maximum. See docs/onnx/tr_hash_v8_validation_report.md. +V8_TOLERANCES = { + "o2m": ParityTolerances(raw=6e-3, decoded_box=1.3e-4, decoded_score=8e-5), + "nms-free": ParityTolerances(raw=1e-2, decoded_box=1.3e-4, decoded_score=4e-5), +} + +# Raw-only view kept for callers that just need the legacy scalar threshold. V8_PARITY_TOLERANCES = { - "o2m": 6e-3, - "nms-free": 1e-2, + branch: tolerances.raw for branch, tolerances in V8_TOLERANCES.items() } +@dataclass(frozen=True) +class GateResult: + """Outcome of one parity gate on one test input.""" + + name: str + tolerance: float + max_difference: float + mean_difference: float + + @property + def passed(self) -> bool: + return self.max_difference <= self.tolerance + + def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("checkpoint", type=Path, help="PyTorch checkpoint directory") @@ -48,10 +98,27 @@ def parse_args() -> argparse.Namespace: type=float, default=None, help=( - "Max allowed absolute difference. Defaults to calibrated v8 branch " - "thresholds when ONNX metadata is available, otherwise 1e-4." + "Max allowed absolute raw-logit difference. Defaults to calibrated " + "v8 branch thresholds when ONNX metadata is available, otherwise 1e-4." ), ) + parser.add_argument( + "--decoded-box-tolerance", + type=float, + default=None, + help="Max allowed absolute difference on normalized decoded boxes", + ) + parser.add_argument( + "--decoded-score-tolerance", + type=float, + default=None, + help="Max allowed absolute difference on sigmoid quality-class scores", + ) + parser.add_argument( + "--skip-decoded", + action="store_true", + help="Only run the raw-logit gate, even when decode metadata is available", + ) parser.add_argument( "--num-tests", type=int, @@ -88,9 +155,74 @@ def branch_from_sidecar(metadata: dict, requested: ExportBranch) -> ExportBranch def calibrated_parity_tolerance(metadata: dict, branch: ExportBranch) -> float: """Return the default raw-logit parity tolerance for an exported model.""" - if metadata.get("architecture_version") == 8 and branch in V8_PARITY_TOLERANCES: - return V8_PARITY_TOLERANCES[branch] - return DEFAULT_PARITY_TOLERANCE + return calibrated_tolerances(metadata, branch).raw + + +def calibrated_tolerances(metadata: dict, branch: ExportBranch) -> ParityTolerances: + """Return per-gate defaults, falling back to strict raw-only for legacy exports.""" + + if metadata.get("architecture_version") == 8 and branch in V8_TOLERANCES: + return V8_TOLERANCES[branch] + return ParityTolerances(raw=DEFAULT_PARITY_TOLERANCE) + + +def decode_context(metadata: dict) -> tuple[OnnxDetectorMetadata, GridGeometry] | None: + """Build decode metadata and grid geometry, or None when unavailable. + + Legacy sidecars predate the v8 decode contract, so a rejected mapping is a + supported outcome rather than an error: the decoded gates are skipped. + """ + + if not metadata: + return None + try: + resolved = metadata_from_mapping(metadata) + except (KeyError, TypeError, ValueError): + return None + geometry = generate_grid_geometry(resolved.image_size, resolved.grid_sizes) + return resolved, geometry + + +def decoded_outputs( + predictions: np.ndarray, + metadata: OnnxDetectorMetadata, + geometry: GridGeometry, +) -> tuple[np.ndarray, np.ndarray]: + """Return normalized xyxy boxes and sigmoid quality-class scores.""" + + validate_output_shape(metadata, predictions.shape) + regression = predictions[..., : metadata.regression_width] + class_logits = predictions[..., metadata.regression_width :] + boxes = decode_dfl_boxes(regression, metadata, geometry) + boxes_norm = boxes / np.float32(metadata.image_size) + return boxes_norm, _sigmoid(class_logits) + + +def evaluate_gates( + torch_predictions: np.ndarray, + onnx_predictions: np.ndarray, + tolerances: ParityTolerances, + context: tuple[OnnxDetectorMetadata, GridGeometry] | None = None, +) -> list[GateResult]: + """Evaluate every enabled gate; each one reports independently.""" + + results = [_gate(RAW_GATE, tolerances.raw, torch_predictions, onnx_predictions)] + if context is None: + return results + + metadata, geometry = context + torch_boxes, torch_scores = decoded_outputs(torch_predictions, metadata, geometry) + onnx_boxes, onnx_scores = decoded_outputs(onnx_predictions, metadata, geometry) + + if tolerances.decoded_box is not None: + results.append( + _gate(DECODED_BOX_GATE, tolerances.decoded_box, torch_boxes, onnx_boxes) + ) + if tolerances.decoded_score is not None: + results.append( + _gate(DECODED_SCORE_GATE, tolerances.decoded_score, torch_scores, onnx_scores) + ) + return results def check_parity( @@ -99,10 +231,13 @@ def check_parity( *, branch: ExportBranch = "auto", tolerance: float | None = None, + decoded_box_tolerance: float | None = None, + decoded_score_tolerance: float | None = None, + skip_decoded: bool = False, num_tests: int = 5, batch_size: int = 1, ) -> bool: - """Compare PyTorch and ONNX raw outputs on random inputs.""" + """Compare PyTorch and ONNX raw and decoded outputs on random inputs.""" if num_tests <= 0 or batch_size <= 0: raise ValueError("num_tests and batch_size must be positive") @@ -116,21 +251,30 @@ def check_parity( detector = load_detector_checkpoint(checkpoint_path, device="cpu") metadata = sidecar_metadata(onnx_path) resolved_branch = branch_from_sidecar(metadata, branch) - effective_tolerance = ( - calibrated_parity_tolerance(metadata, resolved_branch) - if tolerance is None - else tolerance + tolerances = _resolve_tolerances( + metadata, + resolved_branch, + tolerance, + decoded_box_tolerance, + decoded_score_tolerance, ) + context = None if skip_decoded else decode_context(metadata) + export_model = RawDetectorExport(detector, resolved_branch).eval() image_size = detector.config.image_size print(f"Prediction branch: {export_model.branch}") + if context is None: + print( + "Decoded gates: skipped " + f"({'--skip-decoded' if skip_decoded else 'no v8 decode metadata in sidecar'})" + ) print(f"Loading ONNX model: {onnx_path}") session = ort.InferenceSession(str(onnx_path), providers=["CPUExecutionProvider"]) input_name = session.get_inputs()[0].name output_name = session.get_outputs()[0].name - all_passed = True + worst: dict[str, GateResult] = {} for test_index in range(num_tests): torch.manual_seed(test_index) values = torch.randn(batch_size, 3, image_size, image_size) @@ -138,23 +282,14 @@ def check_parity( pytorch_output = export_model(values).numpy() onnx_output = session.run([output_name], {input_name: values.numpy()})[0] - absolute_difference = np.abs(pytorch_output - onnx_output) - max_difference = float(absolute_difference.max()) - mean_difference = float(absolute_difference.mean()) - passed = max_difference <= effective_tolerance - all_passed &= passed - status = "PASS" if passed else "FAIL" - print( - f" Test {test_index + 1}/{num_tests}: " - f"max_diff={max_difference:.2e}, mean_diff={mean_difference:.2e} [{status}]" - ) + results = evaluate_gates(pytorch_output, onnx_output, tolerances, context) + _print_test(test_index, num_tests, results) + for result in results: + current = worst.get(result.name) + if current is None or result.max_difference > current.max_difference: + worst[result.name] = result - summary = "PASSED" if all_passed else "FAILED" - print( - f"\nParity {summary}: " - f"branch={export_model.branch}, tolerance={effective_tolerance}" - ) - return all_passed + return _print_summary(export_model.branch, worst) def main() -> None: @@ -164,11 +299,76 @@ def main() -> None: args.onnx_model, branch=args.branch, tolerance=args.tolerance, + decoded_box_tolerance=args.decoded_box_tolerance, + decoded_score_tolerance=args.decoded_score_tolerance, + skip_decoded=args.skip_decoded, num_tests=args.num_tests, batch_size=args.batch_size, ) raise SystemExit(0 if success else 1) +def _sigmoid(values: np.ndarray) -> np.ndarray: + return 1.0 / (1.0 + np.exp(-values)) + + +def _gate( + name: str, + tolerance: float, + expected: np.ndarray, + actual: np.ndarray, +) -> GateResult: + difference = np.abs(expected - actual) + return GateResult( + name=name, + tolerance=tolerance, + max_difference=float(difference.max()), + mean_difference=float(difference.mean()), + ) + + +def _resolve_tolerances( + metadata: dict, + branch: ExportBranch, + raw_override: float | None, + box_override: float | None, + score_override: float | None, +) -> ParityTolerances: + tolerances = calibrated_tolerances(metadata, branch) + if raw_override is not None: + tolerances = replace(tolerances, raw=raw_override) + if box_override is not None: + tolerances = replace(tolerances, decoded_box=box_override) + if score_override is not None: + tolerances = replace(tolerances, decoded_score=score_override) + return tolerances + + +def _print_test(test_index: int, num_tests: int, results: list[GateResult]) -> None: + label = f" Test {test_index + 1}/{num_tests}" + padding = " " * len(label) + for position, result in enumerate(results): + status = "PASS" if result.passed else "FAIL" + print( + f"{label if position == 0 else padding} {result.name:<13} " + f"max={result.max_difference:.2e} mean={result.mean_difference:.2e} [{status}]" + ) + + +def _print_summary(branch: str, worst: dict[str, GateResult]) -> bool: + failed = [result.name for result in worst.values() if not result.passed] + summary = "FAILED" if failed else "PASSED" + print(f"\nParity {summary}: branch={branch}") + for result in worst.values(): + status = "PASS" if result.passed else "FAIL" + print( + f" {result.name:<13} tol={result.tolerance:.2e} " + f"worst_max={result.max_difference:.2e} [{status}]" + ) + if failed: + print(f" failing gates: {', '.join(failed)}") + return not failed + + if __name__ == "__main__": main() diff --git a/tests/test_detector_export.py b/tests/test_detector_export.py index 5189611..2e27b87 100644 --- a/tests/test_detector_export.py +++ b/tests/test_detector_export.py @@ -3,6 +3,7 @@ import json from pathlib import Path +import numpy as np import pytest import torch from safetensors.torch import save_file @@ -14,11 +15,19 @@ ) from complexity.generative.detection.exporting import RawDetectorExport from scripts.check_onnx_parity import ( + DECODED_BOX_GATE, + DECODED_SCORE_GATE, DEFAULT_PARITY_TOLERANCE, + RAW_GATE, V8_PARITY_TOLERANCES, + V8_TOLERANCES, + ParityTolerances, branch_from_sidecar, calibrated_parity_tolerance, + calibrated_tolerances, check_parity, + decode_context, + evaluate_gates, ) from scripts.export_onnx import export_onnx @@ -82,21 +91,48 @@ def test_auto_export_selects_the_production_branch() -> None: RawDetectorExport(classic, "nms-free") +def sidecar_mapping(config: TRHashDetectorConfig, branch: str) -> dict: + """Mirror the metadata sidecar written by scripts/export_onnx.py.""" + + return { + "architecture_version": config.architecture_version, + "image_size": config.image_size, + "num_classes": config.num_classes, + "num_cells": config.num_cells, + "regression_width": config.regression_width, + "reg_max": config.reg_max, + "scale_factors": list(config.scale_factors), + "grid_sizes": list(config.grid_sizes), + "p2_head": config.p2_head, + "branch": branch, + "requires_nms": branch == "o2m", + "output_semantics": "raw_ltrb_dfl_and_quality_class_logits", + } + + @pytest.mark.parametrize( - ("branch", "expected"), + ("branch", "raw", "decoded_box", "decoded_score"), ( - ("o2m", 6e-3), - ("nms-free", 1e-2), + ("o2m", 6e-3, 1.3e-4, 8e-5), + ("nms-free", 1e-2, 1.3e-4, 4e-5), ), ) def test_v8_exports_use_branch_calibrated_parity_tolerances( branch: str, - expected: float, + raw: float, + decoded_box: float, + decoded_score: float, ) -> None: metadata = {"architecture_version": 8, "branch": branch} - assert V8_PARITY_TOLERANCES[branch] == expected - assert calibrated_parity_tolerance(metadata, branch) == expected + assert V8_PARITY_TOLERANCES[branch] == raw + assert calibrated_parity_tolerance(metadata, branch) == raw + assert V8_TOLERANCES[branch] == ParityTolerances( + raw=raw, + decoded_box=decoded_box, + decoded_score=decoded_score, + ) + assert calibrated_tolerances(metadata, branch) == V8_TOLERANCES[branch] def test_legacy_or_unlabelled_exports_keep_strict_parity_tolerance() -> None: @@ -107,6 +143,76 @@ def test_legacy_or_unlabelled_exports_keep_strict_parity_tolerance() -> None: ) +def test_legacy_exports_disable_the_decoded_gates() -> None: + legacy = calibrated_tolerances({"architecture_version": 7}, "o2m") + + assert legacy.raw == DEFAULT_PARITY_TOLERANCE + assert legacy.decoded_box is None + assert legacy.decoded_score is None + + +@pytest.mark.parametrize( + "metadata", + ( + {}, + {"architecture_version": 7, "branch": "o2m"}, + {"architecture_version": 8, "branch": "o2m"}, # decode fields missing + ), +) +def test_decode_context_is_unavailable_without_full_v8_metadata(metadata: dict) -> None: + assert decode_context(metadata) is None + + +def test_decoded_gates_are_skipped_rather_than_failed_for_legacy_exports() -> None: + predictions = np.zeros((1, 4, 8), dtype=np.float32) + drifted = predictions + 5e-5 + + results = evaluate_gates( + predictions, + drifted, + calibrated_tolerances({}, "auto"), + decode_context({}), + ) + + assert [result.name for result in results] == [RAW_GATE] + assert results[0].passed + + +def test_decoded_box_gate_catches_drift_the_raw_gate_tolerates() -> None: + """Amplification case: softmax turns a coherent logit tilt into box motion. + + Tilting the DFL logits by ``epsilon * bin_index`` shifts the decoded + expectation by about ``epsilon * Var(bin)``, which is then scaled by the + stride. The raw drift stays under its (deliberately coarse) threshold while + the decoded box drift blows past its own. + """ + + config = tiny_config(end_to_end=False) + metadata = sidecar_mapping(config, "o2m") + context = decode_context(metadata) + assert context is not None + + epsilon = 1e-3 + bins = config.dfl_bins + tilt = epsilon * np.arange(bins, dtype=np.float32) + + baseline = np.zeros((1, config.num_cells, config.prediction_width), dtype=np.float32) + drifted = baseline.copy() + # Same tilt on each of the four LTRB distributions; class logits untouched. + drifted[..., : config.regression_width] += np.tile(tilt, 4) + + tolerances = ParityTolerances(raw=5e-3, decoded_box=1.3e-4, decoded_score=4e-5) + results = { + result.name: result + for result in evaluate_gates(baseline, drifted, tolerances, context) + } + + assert results[RAW_GATE].passed, "raw drift must stay inside its tolerance" + assert not results[DECODED_BOX_GATE].passed, "decoded box drift must be caught" + assert results[DECODED_SCORE_GATE].passed, "gates must fail independently" + assert results[RAW_GATE].max_difference == pytest.approx(epsilon * (bins - 1)) + + def test_sidecar_less_auto_branch_still_defers_to_model_resolution() -> None: assert branch_from_sidecar({}, "auto") == "auto" @@ -131,10 +237,21 @@ def test_dynamic_onnx_export_matches_the_selected_branch(tmp_path: Path, branch: assert metadata["architecture_version"] == 8 assert metadata["branch"] == branch assert metadata["requires_nms"] is (branch == "o2m") + # The sidecar must be rich enough to drive the decoded gates. + assert decode_context(metadata) is not None + # Strict raw threshold, calibrated decoded thresholds: all three gates run. + assert check_parity( + checkpoint, + output, + num_tests=1, + batch_size=2, + tolerance=1e-4, + ) assert check_parity( checkpoint, output, num_tests=1, batch_size=2, tolerance=1e-4, + skip_decoded=True, )