| 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" |
|
|
|
|