from __future__ import annotations import numpy as np from scripts.stages.compare_onnx_mlir_compiled import comparison from scripts.stages.onnx_mlir_compiled_runtime import NUMPY_TO_OM_DTYPE, OM_DTYPE_TO_NUMPY def test_onnx_dtype_mapping_round_trip() -> None: for code, dtype in OM_DTYPE_TO_NUMPY.items(): assert NUMPY_TO_OM_DTYPE[dtype] == code def test_comparison_pass_and_fail() -> None: reference = np.asarray([[1.0, 2.0]], dtype=np.float32) close = np.asarray([[1.0, 2.0 + 1e-6]], dtype=np.float32) far = np.asarray([[1.0, 2.1]], dtype=np.float32) assert comparison(reference, close, atol=1e-5, rtol=1e-5, domain="real")["status"] == "PASS" assert comparison(reference, far, atol=1e-5, rtol=1e-5, domain="real")["status"] == "FAIL" def test_raw_integer_requires_exact_match() -> None: reference = np.asarray([[1, 2]], dtype=np.int8) same = reference.copy() changed = np.asarray([[1, 3]], dtype=np.int8) assert comparison(reference, same, atol=0.0, rtol=0.0, domain="raw_integer")["status"] == "PASS" assert comparison(reference, changed, atol=0.0, rtol=0.0, domain="raw_integer")["status"] == "FAIL"