diff --git a/PythonModels/scripts/run_test_mql_full_state_comparison.py b/PythonModels/scripts/run_test_mql_full_state_comparison.py new file mode 100644 index 0000000..1ac96b4 --- /dev/null +++ b/PythonModels/scripts/run_test_mql_full_state_comparison.py @@ -0,0 +1,193 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from datetime import UTC, datetime +from pathlib import Path + +from PythonModels.core.solver import SolveIVPConfig +from PythonModels.reporting.amesim_results import AmesimResults, load_test_mql_amesim_results +from PythonModels.reporting.test_mql_comparison import ( + TestMqlComparisonResult, + write_test_mql_comparison_csv, +) +from PythonModels.reporting.test_mql_output_schema import ( + TestMqlOutputSchema, + build_test_mql_output_schema, +) +from PythonModels.reporting.test_mql_output_validation import ( + TestMqlValidatedOutput, + compare_validated_test_mql_output, + validate_test_mql_output, +) +from PythonModels.systems.test_mql import TestMqlSimulationResult, TestMqlSystem + + +DEFAULT_FULL_STATE_COMPARISON_DATA_PATHS = ( + "press@pn_c1_8", + "vol@pn_c1_8", + "vol1@pn_brp2_8", + "vvol1@pn_brp2_8", + "x1@mass_friction_endstops_10", + "v1@mass_friction_endstops_10", + "acc1@mass_friction_endstops_10", + "x1@mass_friction_endstops_18", + "v1@mass_friction_endstops_18", + "acc1@mass_friction_endstops_18", +) + + +@dataclass(frozen=True) +class TestMqlFullStateComparisonRun: + system: TestMqlSystem + amesim_results: AmesimResults + output_schema: TestMqlOutputSchema + result: TestMqlSimulationResult + output: TestMqlValidatedOutput + comparison: TestMqlComparisonResult + + @property + def sample_count(self) -> int: + return len(self.output.times) + + @property + def signal_count(self) -> int: + return len(self.output.data_paths) + + +@dataclass(frozen=True) +class TestMqlFullStateComparisonPathConfig: + archive_path: Path = field( + default_factory=lambda: Path(__file__).resolve().parents[2] + / "AmesimModels" + / "test_mql.ame" + ) + output_dir: Path | None = None + + +@dataclass(frozen=True) +class TestMqlFullStateComparisonExecutionConfig: + write_summary: bool = True + write_comparison_csv: bool = True + data_paths: tuple[str, ...] | None = DEFAULT_FULL_STATE_COMPARISON_DATA_PATHS + solver: SolveIVPConfig = field( + default_factory=lambda: SolveIVPConfig(t_stop=1.0e-5, max_step=1.0e-6) + ) + t_eval: tuple[float, ...] | None = (0.0, 1.0e-5) + inlet_node_pressure_pa: float = 15.31e6 + resistance_boundary_pressure_pa: float = 15.29e6 + inlet_node_temperature_k: float = 293.15 + resistance_boundary_temperature_k: float = 293.15 + + +@dataclass(frozen=True) +class TestMqlFullStateComparisonScriptConfig: + paths: TestMqlFullStateComparisonPathConfig = field( + default_factory=TestMqlFullStateComparisonPathConfig + ) + execution: TestMqlFullStateComparisonExecutionConfig = field( + default_factory=TestMqlFullStateComparisonExecutionConfig + ) + + +def _default_output_dir() -> Path: + pythonmodels_root = Path(__file__).resolve().parents[1] + timestamp = datetime.now(UTC).strftime("test_mql_full_state_%Y%m%d_%H%M%S_%f") + return pythonmodels_root / "runs" / timestamp + + +def run_test_mql_full_state_comparison( + config: TestMqlFullStateComparisonScriptConfig | None = None, +) -> tuple[TestMqlFullStateComparisonRun, Path]: + config = config or TestMqlFullStateComparisonScriptConfig() + system = TestMqlSystem(archive_path=config.paths.archive_path) + amesim_results = load_test_mql_amesim_results(config.paths.archive_path) + output_schema = build_test_mql_output_schema(amesim_results) + selected_paths = config.execution.data_paths + spec = system.discover_pneumatic_branch_topology().chamber_segment_specs[0] + result = system.simulate_full_state_series_from_spec( + spec, + inlet_node_pressure_pa=config.execution.inlet_node_pressure_pa, + resistance_boundary_pressure_pa=( + config.execution.resistance_boundary_pressure_pa + ), + inlet_node_temperature_k=config.execution.inlet_node_temperature_k, + resistance_boundary_temperature_k=( + config.execution.resistance_boundary_temperature_k + ), + config=config.execution.solver, + t_eval=list(config.execution.t_eval) if config.execution.t_eval is not None else None, + data_paths=selected_paths, + ) + series_by_data_path = { + data_path: result.series[data_path] + for data_path in result.series + if data_path != "time" + } + output = validate_test_mql_output( + times=result.t, + series_by_data_path=series_by_data_path, + schema=output_schema, + data_paths=selected_paths, + ) + comparison = compare_validated_test_mql_output( + times=output.times, + series_by_data_path=output.series_by_data_path, + schema=output_schema, + amesim_results=amesim_results, + data_paths=output.data_paths, + ) + run = TestMqlFullStateComparisonRun( + system=system, + amesim_results=amesim_results, + output_schema=output_schema, + result=result, + output=output, + comparison=comparison, + ) + output_dir = config.paths.output_dir or _default_output_dir() + if config.execution.write_summary or config.execution.write_comparison_csv: + output_dir.mkdir(parents=True, exist_ok=True) + if config.execution.write_summary: + (output_dir / "test_mql_full_state_comparison_summary.txt").write_text( + format_test_mql_full_state_comparison_summary(run), + encoding="utf-8", + ) + if config.execution.write_comparison_csv: + write_test_mql_comparison_csv( + output_dir=output_dir, + python_times=output.times, + python_series_by_data_path=output.series_by_data_path, + amesim_results=amesim_results, + data_paths=output.data_paths, + ) + return run, output_dir + + +def format_test_mql_full_state_comparison_summary( + run: TestMqlFullStateComparisonRun, +) -> str: + lines = [ + "Model: test_mql", + "Mode: Python 132 full-state closure comparison", + f"Samples: {run.sample_count}", + f"Compared signals: {run.signal_count}", + f"Output schema signals: {run.output_schema.signal_count}", + f"Max absolute error: {run.comparison.max_abs_error}", + f"Max relative error: {run.comparison.max_rel_error}", + ] + for metric in run.comparison.metrics: + lines.append( + f" - {metric.data_path}: max_abs_error={metric.max_abs_error}, " + f"final_abs_error={metric.final_abs_error}" + ) + return "\n".join(lines) + "\n" + + +def main() -> None: + run, output_dir = run_test_mql_full_state_comparison() + print(format_test_mql_full_state_comparison_summary(run), end="") + print(f"Output directory: {output_dir}") + + +if __name__ == "__main__": + main() diff --git a/tests/test_run_test_mql_full_state_comparison.py b/tests/test_run_test_mql_full_state_comparison.py new file mode 100644 index 0000000..a770b7a --- /dev/null +++ b/tests/test_run_test_mql_full_state_comparison.py @@ -0,0 +1,102 @@ +from __future__ import annotations + +import tempfile +import unittest +from pathlib import Path + +from PythonModels.core.solver import SolveIVPConfig +from PythonModels.scripts.run_test_mql_full_state_comparison import ( + TestMqlFullStateComparisonExecutionConfig, + TestMqlFullStateComparisonPathConfig, + TestMqlFullStateComparisonScriptConfig, + format_test_mql_full_state_comparison_summary, + run_test_mql_full_state_comparison, +) + + +REPO_ROOT = Path(__file__).resolve().parents[1] +TEST_MQL_AME = REPO_ROOT / "AmesimModels" / "test_mql.ame" + + +class RunTestMqlFullStateComparisonScriptTests(unittest.TestCase): + def test_full_state_script_writes_summary_and_comparison_csv(self) -> None: + with tempfile.TemporaryDirectory() as tmpdir: + output_dir = Path(tmpdir) + run, returned_output_dir = run_test_mql_full_state_comparison( + TestMqlFullStateComparisonScriptConfig( + paths=TestMqlFullStateComparisonPathConfig( + archive_path=TEST_MQL_AME, + output_dir=output_dir, + ), + execution=TestMqlFullStateComparisonExecutionConfig( + data_paths=("press@pn_c1_8", "vol1@pn_brp2_8"), + solver=SolveIVPConfig(t_start=0.0, t_stop=0.0), + t_eval=(0.0,), + ), + ) + ) + summary_path = output_dir / "test_mql_full_state_comparison_summary.txt" + csv_path = output_dir / "test_mql_amesim_comparison.csv" + csv_summary_path = output_dir / "test_mql_amesim_comparison_summary.txt" + summary = summary_path.read_text(encoding="utf-8") + csv_header = csv_path.read_text(encoding="utf-8").splitlines()[0] + csv_summary_exists = csv_summary_path.exists() + + self.assertEqual(returned_output_dir, output_dir) + self.assertEqual(run.sample_count, 1) + self.assertEqual(run.signal_count, 2) + self.assertAlmostEqual(run.comparison.max_abs_error, 0.0, delta=1.0e-8) + self.assertTrue(csv_summary_exists) + self.assertIn("Mode: Python 132 full-state closure comparison", summary) + self.assertIn("Compared signals: 2", summary) + self.assertIn("python.press@pn_c1_8", csv_header) + self.assertIn("amesim.vol1@pn_brp2_8", csv_header) + + def test_full_state_script_can_skip_artifact_files(self) -> None: + with tempfile.TemporaryDirectory() as tmpdir: + output_dir = Path(tmpdir) + run, returned_output_dir = run_test_mql_full_state_comparison( + TestMqlFullStateComparisonScriptConfig( + paths=TestMqlFullStateComparisonPathConfig( + archive_path=TEST_MQL_AME, + output_dir=output_dir, + ), + execution=TestMqlFullStateComparisonExecutionConfig( + write_summary=False, + write_comparison_csv=False, + data_paths=("press@pn_c1_8",), + solver=SolveIVPConfig(t_start=0.0, t_stop=0.0), + t_eval=(0.0,), + ), + ) + ) + + self.assertEqual(returned_output_dir, output_dir) + self.assertEqual(run.signal_count, 1) + self.assertFalse((output_dir / "test_mql_full_state_comparison_summary.txt").exists()) + self.assertFalse((output_dir / "test_mql_amesim_comparison.csv").exists()) + + def test_format_test_mql_full_state_comparison_summary(self) -> None: + run, _ = run_test_mql_full_state_comparison( + TestMqlFullStateComparisonScriptConfig( + paths=TestMqlFullStateComparisonPathConfig(archive_path=TEST_MQL_AME), + execution=TestMqlFullStateComparisonExecutionConfig( + write_summary=False, + write_comparison_csv=False, + data_paths=("press@pn_c1_8",), + solver=SolveIVPConfig(t_start=0.0, t_stop=0.0), + t_eval=(0.0,), + ), + ) + ) + + summary = format_test_mql_full_state_comparison_summary(run) + + self.assertIn("Model: test_mql", summary) + self.assertIn("Output schema signals: 858", summary) + self.assertIn("Compared signals: 1", summary) + self.assertTrue(summary.endswith("\n")) + + +if __name__ == "__main__": + unittest.main()