merge/model-development-into-main #2
No files matched your search
@@ -0,0 +1,161 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from math import isfinite
|
||||||
|
|
||||||
|
from PythonModels.reporting.amesim_results import AmesimResults
|
||||||
|
from PythonModels.reporting.test_mql_comparison import (
|
||||||
|
TestMqlComparisonResult,
|
||||||
|
compare_test_mql_series,
|
||||||
|
)
|
||||||
|
from PythonModels.reporting.test_mql_output_schema import TestMqlOutputSchema
|
||||||
|
|
||||||
|
|
||||||
|
class TestMqlOutputValidationError(ValueError):
|
||||||
|
"""Raised when a Python test_mql output does not satisfy the AMESim output contract."""
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class TestMqlValidatedOutput:
|
||||||
|
times: tuple[float, ...]
|
||||||
|
series_by_data_path: dict[str, tuple[float, ...]]
|
||||||
|
data_paths: tuple[str, ...]
|
||||||
|
|
||||||
|
def series(self, data_path: str) -> tuple[float, ...]:
|
||||||
|
if data_path not in self.series_by_data_path:
|
||||||
|
raise KeyError(data_path)
|
||||||
|
return self.series_by_data_path[data_path]
|
||||||
|
|
||||||
|
|
||||||
|
def validate_test_mql_output(
|
||||||
|
*,
|
||||||
|
times: tuple[float, ...] | list[float],
|
||||||
|
series_by_data_path: dict[str, tuple[float, ...] | list[float]],
|
||||||
|
schema: TestMqlOutputSchema,
|
||||||
|
data_paths: tuple[str, ...] | list[str] | None = None,
|
||||||
|
allow_extra_paths: bool = False,
|
||||||
|
require_all_schema_paths: bool = False,
|
||||||
|
) -> TestMqlValidatedOutput:
|
||||||
|
validated_times = _validate_time_axis(times)
|
||||||
|
selected_paths = _select_paths(
|
||||||
|
series_by_data_path=series_by_data_path,
|
||||||
|
schema=schema,
|
||||||
|
data_paths=data_paths,
|
||||||
|
allow_extra_paths=allow_extra_paths,
|
||||||
|
require_all_schema_paths=require_all_schema_paths,
|
||||||
|
)
|
||||||
|
validated_series = {
|
||||||
|
data_path: _validate_series(
|
||||||
|
data_path=data_path,
|
||||||
|
values=series_by_data_path[data_path],
|
||||||
|
expected_count=len(validated_times),
|
||||||
|
)
|
||||||
|
for data_path in selected_paths
|
||||||
|
}
|
||||||
|
return TestMqlValidatedOutput(
|
||||||
|
times=validated_times,
|
||||||
|
series_by_data_path=validated_series,
|
||||||
|
data_paths=selected_paths,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def compare_validated_test_mql_output(
|
||||||
|
*,
|
||||||
|
times: tuple[float, ...] | list[float],
|
||||||
|
series_by_data_path: dict[str, tuple[float, ...] | list[float]],
|
||||||
|
schema: TestMqlOutputSchema,
|
||||||
|
amesim_results: AmesimResults,
|
||||||
|
data_paths: tuple[str, ...] | list[str] | None = None,
|
||||||
|
allow_extra_paths: bool = False,
|
||||||
|
require_all_schema_paths: bool = False,
|
||||||
|
relative_floor: float = 1.0e-12,
|
||||||
|
) -> TestMqlComparisonResult:
|
||||||
|
validated = validate_test_mql_output(
|
||||||
|
times=times,
|
||||||
|
series_by_data_path=series_by_data_path,
|
||||||
|
schema=schema,
|
||||||
|
data_paths=data_paths,
|
||||||
|
allow_extra_paths=allow_extra_paths,
|
||||||
|
require_all_schema_paths=require_all_schema_paths,
|
||||||
|
)
|
||||||
|
return compare_test_mql_series(
|
||||||
|
python_times=validated.times,
|
||||||
|
python_series_by_data_path=validated.series_by_data_path,
|
||||||
|
amesim_results=amesim_results,
|
||||||
|
data_paths=validated.data_paths,
|
||||||
|
relative_floor=relative_floor,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_time_axis(times: tuple[float, ...] | list[float]) -> tuple[float, ...]:
|
||||||
|
if not times:
|
||||||
|
raise TestMqlOutputValidationError("Python time axis is empty.")
|
||||||
|
validated = tuple(_finite_float("time", value) for value in times)
|
||||||
|
previous = validated[0]
|
||||||
|
for value in validated[1:]:
|
||||||
|
if value < previous:
|
||||||
|
raise TestMqlOutputValidationError("Python time axis must be monotonically increasing.")
|
||||||
|
previous = value
|
||||||
|
return validated
|
||||||
|
|
||||||
|
|
||||||
|
def _select_paths(
|
||||||
|
*,
|
||||||
|
series_by_data_path: dict[str, tuple[float, ...] | list[float]],
|
||||||
|
schema: TestMqlOutputSchema,
|
||||||
|
data_paths: tuple[str, ...] | list[str] | None,
|
||||||
|
allow_extra_paths: bool,
|
||||||
|
require_all_schema_paths: bool,
|
||||||
|
) -> tuple[str, ...]:
|
||||||
|
schema_paths = set(schema.data_paths())
|
||||||
|
provided_paths = set(series_by_data_path)
|
||||||
|
if not allow_extra_paths:
|
||||||
|
extra_paths = sorted(provided_paths - schema_paths)
|
||||||
|
if extra_paths:
|
||||||
|
raise TestMqlOutputValidationError(
|
||||||
|
f"Python output contains Data_Path values outside test_mql schema: {extra_paths}"
|
||||||
|
)
|
||||||
|
if require_all_schema_paths:
|
||||||
|
missing_schema_paths = sorted(schema_paths - provided_paths)
|
||||||
|
if missing_schema_paths:
|
||||||
|
raise TestMqlOutputValidationError(
|
||||||
|
f"Python output is missing required test_mql schema Data_Path values: {missing_schema_paths}"
|
||||||
|
)
|
||||||
|
selected_paths = tuple(data_paths) if data_paths is not None else tuple(sorted(provided_paths & schema_paths))
|
||||||
|
if not selected_paths:
|
||||||
|
raise TestMqlOutputValidationError("no test_mql schema Data_Path values are available.")
|
||||||
|
unknown_selected = [data_path for data_path in selected_paths if data_path not in schema_paths]
|
||||||
|
if unknown_selected:
|
||||||
|
raise TestMqlOutputValidationError(
|
||||||
|
f"Requested Data_Path values are outside test_mql schema: {unknown_selected}"
|
||||||
|
)
|
||||||
|
missing_selected = [data_path for data_path in selected_paths if data_path not in series_by_data_path]
|
||||||
|
if missing_selected:
|
||||||
|
raise TestMqlOutputValidationError(
|
||||||
|
f"Python output is missing selected Data_Path values: {missing_selected}"
|
||||||
|
)
|
||||||
|
return selected_paths
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_series(
|
||||||
|
*,
|
||||||
|
data_path: str,
|
||||||
|
values: tuple[float, ...] | list[float],
|
||||||
|
expected_count: int,
|
||||||
|
) -> tuple[float, ...]:
|
||||||
|
if len(values) != expected_count:
|
||||||
|
raise TestMqlOutputValidationError(
|
||||||
|
f"Python series length mismatch for {data_path!r}: "
|
||||||
|
f"{len(values)} values for {expected_count} time samples."
|
||||||
|
)
|
||||||
|
return tuple(_finite_float(data_path, value) for value in values)
|
||||||
|
|
||||||
|
|
||||||
|
def _finite_float(label: str, value: float) -> float:
|
||||||
|
try:
|
||||||
|
numeric_value = float(value)
|
||||||
|
except (TypeError, ValueError) as exc:
|
||||||
|
raise TestMqlOutputValidationError(f"{label!r} contains a non-numeric value: {value!r}") from exc
|
||||||
|
if not isfinite(numeric_value):
|
||||||
|
raise TestMqlOutputValidationError(f"{label!r} contains a non-finite value: {value!r}")
|
||||||
|
return numeric_value
|
||||||
@@ -0,0 +1,181 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from PythonModels.reporting.amesim_results import load_test_mql_amesim_results
|
||||||
|
from PythonModels.reporting.test_mql_observations import build_test_mql_observation_catalog
|
||||||
|
from PythonModels.reporting.test_mql_output_schema import build_test_mql_output_schema
|
||||||
|
from PythonModels.reporting.test_mql_output_validation import (
|
||||||
|
TestMqlOutputValidationError,
|
||||||
|
compare_validated_test_mql_output,
|
||||||
|
validate_test_mql_output,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||||
|
TEST_MQL_AME = REPO_ROOT / "AmesimModels" / "test_mql.ame"
|
||||||
|
|
||||||
|
|
||||||
|
class TestMqlOutputValidationTests(unittest.TestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls) -> None:
|
||||||
|
cls.amesim_results = load_test_mql_amesim_results(TEST_MQL_AME)
|
||||||
|
cls.observation_catalog = build_test_mql_observation_catalog(cls.amesim_results)
|
||||||
|
cls.schema = build_test_mql_output_schema(
|
||||||
|
cls.amesim_results,
|
||||||
|
observation_catalog=cls.observation_catalog,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_validates_selected_python_output_series(self) -> None:
|
||||||
|
times = [0.0, 0.1, 0.2]
|
||||||
|
series = {
|
||||||
|
"press@pn_c1_8": [1.0, 2.0, 3.0],
|
||||||
|
"vol1@pn_brp2_8": [4, 5, 6],
|
||||||
|
}
|
||||||
|
|
||||||
|
validated = validate_test_mql_output(
|
||||||
|
times=times,
|
||||||
|
series_by_data_path=series,
|
||||||
|
schema=self.schema,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(validated.times, (0.0, 0.1, 0.2))
|
||||||
|
self.assertEqual(validated.data_paths, ("press@pn_c1_8", "vol1@pn_brp2_8"))
|
||||||
|
self.assertEqual(validated.series("press@pn_c1_8"), (1.0, 2.0, 3.0))
|
||||||
|
self.assertEqual(validated.series("vol1@pn_brp2_8"), (4.0, 5.0, 6.0))
|
||||||
|
|
||||||
|
def test_can_validate_explicit_data_path_order(self) -> None:
|
||||||
|
validated = validate_test_mql_output(
|
||||||
|
times=(0.0, 0.1),
|
||||||
|
series_by_data_path={
|
||||||
|
"press@pn_c1_8": (10.0, 11.0),
|
||||||
|
"vol1@pn_brp2_8": (20.0, 21.0),
|
||||||
|
},
|
||||||
|
schema=self.schema,
|
||||||
|
data_paths=("vol1@pn_brp2_8", "press@pn_c1_8"),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(validated.data_paths, ("vol1@pn_brp2_8", "press@pn_c1_8"))
|
||||||
|
|
||||||
|
def test_rejects_output_paths_outside_schema_by_default(self) -> None:
|
||||||
|
with self.assertRaisesRegex(TestMqlOutputValidationError, "outside test_mql schema"):
|
||||||
|
validate_test_mql_output(
|
||||||
|
times=(0.0,),
|
||||||
|
series_by_data_path={"not_a_signal@not_an_owner": (1.0,)},
|
||||||
|
schema=self.schema,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_can_ignore_extra_output_paths_when_requested(self) -> None:
|
||||||
|
validated = validate_test_mql_output(
|
||||||
|
times=(0.0,),
|
||||||
|
series_by_data_path={
|
||||||
|
"press@pn_c1_8": (1.0,),
|
||||||
|
"not_a_signal@not_an_owner": (2.0,),
|
||||||
|
},
|
||||||
|
schema=self.schema,
|
||||||
|
allow_extra_paths=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(validated.data_paths, ("press@pn_c1_8",))
|
||||||
|
|
||||||
|
def test_rejects_missing_selected_data_path(self) -> None:
|
||||||
|
with self.assertRaisesRegex(TestMqlOutputValidationError, "missing selected"):
|
||||||
|
validate_test_mql_output(
|
||||||
|
times=(0.0,),
|
||||||
|
series_by_data_path={"press@pn_c1_8": (1.0,)},
|
||||||
|
schema=self.schema,
|
||||||
|
data_paths=("press@pn_c1_8", "vol1@pn_brp2_8"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_rejects_incomplete_full_schema_output(self) -> None:
|
||||||
|
with self.assertRaisesRegex(TestMqlOutputValidationError, "missing required"):
|
||||||
|
validate_test_mql_output(
|
||||||
|
times=(0.0,),
|
||||||
|
series_by_data_path={"press@pn_c1_8": (1.0,)},
|
||||||
|
schema=self.schema,
|
||||||
|
require_all_schema_paths=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_rejects_length_mismatch(self) -> None:
|
||||||
|
with self.assertRaisesRegex(TestMqlOutputValidationError, "length mismatch"):
|
||||||
|
validate_test_mql_output(
|
||||||
|
times=(0.0, 0.1),
|
||||||
|
series_by_data_path={"press@pn_c1_8": (1.0,)},
|
||||||
|
schema=self.schema,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_rejects_non_finite_values(self) -> None:
|
||||||
|
with self.assertRaisesRegex(TestMqlOutputValidationError, "non-finite"):
|
||||||
|
validate_test_mql_output(
|
||||||
|
times=(0.0,),
|
||||||
|
series_by_data_path={"press@pn_c1_8": (float("nan"),)},
|
||||||
|
schema=self.schema,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_rejects_non_monotonic_time_axis(self) -> None:
|
||||||
|
with self.assertRaisesRegex(TestMqlOutputValidationError, "monotonically increasing"):
|
||||||
|
validate_test_mql_output(
|
||||||
|
times=(0.0, 0.2, 0.1),
|
||||||
|
series_by_data_path={"press@pn_c1_8": (1.0, 2.0, 3.0)},
|
||||||
|
schema=self.schema,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_compares_validated_amesim_baseline_with_zero_error(self) -> None:
|
||||||
|
selected_paths = (
|
||||||
|
"press@pn_c1_8",
|
||||||
|
"dm2@pn_morifice_11",
|
||||||
|
"x1@mass_friction_endstops_10",
|
||||||
|
)
|
||||||
|
baseline = self.observation_catalog.baseline_series_by_data_path(
|
||||||
|
self.amesim_results,
|
||||||
|
selected_paths,
|
||||||
|
)
|
||||||
|
|
||||||
|
comparison = compare_validated_test_mql_output(
|
||||||
|
times=self.amesim_results.times,
|
||||||
|
series_by_data_path=baseline,
|
||||||
|
schema=self.schema,
|
||||||
|
amesim_results=self.amesim_results,
|
||||||
|
data_paths=selected_paths,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(len(comparison.metrics), len(selected_paths))
|
||||||
|
self.assertEqual(comparison.max_abs_error, 0.0)
|
||||||
|
self.assertEqual(comparison.max_rel_error, 0.0)
|
||||||
|
|
||||||
|
def test_compares_validated_output_and_reports_perturbation_error(self) -> None:
|
||||||
|
selected_paths = ("press@pn_c1_8",)
|
||||||
|
baseline = self.observation_catalog.baseline_series_by_data_path(
|
||||||
|
self.amesim_results,
|
||||||
|
selected_paths,
|
||||||
|
)
|
||||||
|
perturbed = dict(baseline)
|
||||||
|
values = list(perturbed["press@pn_c1_8"])
|
||||||
|
values[-1] += 10.0
|
||||||
|
perturbed["press@pn_c1_8"] = tuple(values)
|
||||||
|
|
||||||
|
comparison = compare_validated_test_mql_output(
|
||||||
|
times=self.amesim_results.times,
|
||||||
|
series_by_data_path=perturbed,
|
||||||
|
schema=self.schema,
|
||||||
|
amesim_results=self.amesim_results,
|
||||||
|
data_paths=selected_paths,
|
||||||
|
)
|
||||||
|
|
||||||
|
metric = comparison.metric("press@pn_c1_8")
|
||||||
|
self.assertAlmostEqual(metric.final_abs_error, 10.0)
|
||||||
|
self.assertAlmostEqual(metric.max_abs_error, 10.0)
|
||||||
|
|
||||||
|
def test_compare_validated_output_reuses_schema_validation(self) -> None:
|
||||||
|
with self.assertRaisesRegex(TestMqlOutputValidationError, "outside test_mql schema"):
|
||||||
|
compare_validated_test_mql_output(
|
||||||
|
times=(0.0,),
|
||||||
|
series_by_data_path={"not_a_signal@not_an_owner": (1.0,)},
|
||||||
|
schema=self.schema,
|
||||||
|
amesim_results=self.amesim_results,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in new issue
Block a user