ONNX
onnxruntime
onnx-mlir
quantization
fp32
ONNX_Models / tests /test_ad01_compiled_mlir_evaluation.py
purejomo's picture
Finalize public ONNX/ONNX-MLIR validation release
ed3aeeb
Raw
History Blame Contribute Delete
5.09 kB
from __future__ import annotations
from dataclasses import dataclass
import numpy as np
from scripts.stages.evaluate_ad01_compiled_mlir import (
compiled_metric_rows,
evaluate_feature_vectors,
merge_file_score_rows,
merge_metric_rows,
)
@dataclass
class _Input:
name: str = "input_1"
class _Reference:
def __init__(self) -> None:
self.input_shapes: list[tuple[int, ...]] = []
def get_inputs(self) -> list[_Input]:
return [_Input()]
def run(self, _outputs: object, feeds: dict[str, np.ndarray]) -> list[np.ndarray]:
self.input_shapes.append(feeds["input_1"].shape)
return [feeds["input_1"] * np.float32(0.5)]
class _Compiled:
def __init__(self) -> None:
self.input_shapes: list[tuple[int, ...]] = []
def run(self, inputs: list[np.ndarray]) -> list[np.ndarray]:
self.input_shapes.append(inputs[0].shape)
return [inputs[0] * np.float32(0.5)]
class _ChangedCompiled(_Compiled):
def run(self, inputs: list[np.ndarray]) -> list[np.ndarray]:
value = super().run(inputs)[0].copy()
value[0, 0] += np.float32(1.0)
return [value]
def test_compiled_rows_reuse_official_metric_and_preserve_fidelity() -> None:
# The official path invokes every one of a file's 196 feature vectors with
# batch one. A [196, 640] compiler batch is deliberately not substituted:
# it has a distinct FP32 numerical fidelity result at the pinned tolerance.
vectors = np.arange(196 * 640, dtype=np.float64).reshape(196, 640) / 1000.0
reference = _Reference()
compiled = _Compiled()
evaluated = evaluate_feature_vectors(
vectors,
filename="normal_id_01_00000000.wav",
variant="fp32",
reference=reference,
compiled=compiled,
quantization=None,
fp32_atol=1e-5,
fp32_rtol=1e-5,
)
assert evaluated["fidelity_rows"] == 196
assert evaluated["fidelity_matching_rows"] == 196
assert evaluated["fidelity_mismatching_rows"] == 0
assert evaluated["fidelity_bitwise_equal_rows"] == 196
assert evaluated["fidelity_total_elements"] == 196 * 640
assert len(evaluated["fidelity_digest"]) == 64
assert evaluated["fidelity_mismatches"] == []
assert reference.input_shapes == [(1, 640)] * 196
assert compiled.input_shapes == [(1, 640)] * 196
assert evaluated["compiled_score"] == evaluated["onnxruntime_score"]
mismatch = evaluate_feature_vectors(
vectors[:1],
filename="anomaly_id_01_00000000.wav",
variant="fp32",
reference=_Reference(),
compiled=_ChangedCompiled(),
quantization=None,
fp32_atol=1e-5,
fp32_rtol=1e-5,
)
assert mismatch["fidelity_rows"] == 1
assert mismatch["fidelity_matching_rows"] == 0
assert mismatch["fidelity_mismatching_rows"] == 1
assert len(mismatch["fidelity_mismatches"]) == 1
assert mismatch["fidelity_mismatches"][0]["status"] == "FAIL"
file_rows = [
{"machine_id": "id_01", "label": 0, "compiled_score": 0.1},
{"machine_id": "id_01", "label": 0, "compiled_score": 0.2},
{"machine_id": "id_01", "label": 1, "compiled_score": 0.8},
{"machine_id": "id_01", "label": 1, "compiled_score": 0.9},
]
metrics = compiled_metric_rows(
file_rows, "compiled_score", "fp32", max_fpr=0.1
)
average = next(row for row in metrics if row["machine_id"] == "Average")
assert average["variant"] == "fp32"
assert average["auc"] == 1.0
assert average["pauc"] == 1.0
fp32_scores = [{
"filename": "a.wav", "machine_id": "id_01", "label": "0",
"feature_vectors": "196", "onnxruntime_score": "0.1",
"compiled_score": "0.2",
}]
quantized_scores = [{
"filename": "a.wav", "machine_id": "id_01", "label": "0",
"feature_vectors": "196", "onnxruntime_score": "0.3",
"compiled_score": "0.4",
}]
merged_scores = merge_file_score_rows(fp32_scores, quantized_scores)
assert merged_scores == [{
"filename": "a.wav", "machine_id": "id_01", "label": "0",
"feature_vectors": "196", "fp32_onnxruntime_score": "0.1",
"fp32_compiled_score": "0.2",
"public_quantized_onnxruntime_score": "0.3",
"public_quantized_compiled_score": "0.4",
}]
metric_inputs = []
for runtime in ("onnxruntime_reference", "compiled_mlir"):
for machine_id in ("id_01", "id_02", "id_03", "id_04", "Average"):
metric_inputs.append({
"runtime": runtime, "variant": "fp32", "machine_id": machine_id,
"auc": "0.9", "pauc": "0.8", "max_fpr": "0.1",
})
quantized_metric_inputs = [
{**row, "variant": "public_quantized"} for row in metric_inputs
]
merged_metrics = merge_metric_rows(metric_inputs, quantized_metric_inputs)
assert len(merged_metrics) == 20
assert {row["variant"] for row in merged_metrics} == {
"fp32_onnxruntime", "fp32_compiled",
"public_quantized_onnxruntime", "public_quantized_compiled",
}