merge/model-development-into-main #2
No files matched your search
@@ -0,0 +1,77 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
|
||||
from PythonModels.systems.test_mql_baseline import (
|
||||
TestMqlBaselineRun,
|
||||
run_test_mql_baseline_passthrough,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TestMqlBaselinePathConfig:
|
||||
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 TestMqlBaselineExecutionConfig:
|
||||
write_summary: bool = True
|
||||
data_paths: tuple[str, ...] | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TestMqlBaselineScriptConfig:
|
||||
paths: TestMqlBaselinePathConfig = field(default_factory=TestMqlBaselinePathConfig)
|
||||
execution: TestMqlBaselineExecutionConfig = field(default_factory=TestMqlBaselineExecutionConfig)
|
||||
|
||||
|
||||
def _default_output_dir() -> Path:
|
||||
pythonmodels_root = Path(__file__).resolve().parents[1]
|
||||
timestamp = datetime.now(UTC).strftime("test_mql_baseline_%Y%m%d_%H%M%S_%f")
|
||||
return pythonmodels_root / "runs" / timestamp
|
||||
|
||||
|
||||
def format_test_mql_baseline_summary(run: TestMqlBaselineRun) -> str:
|
||||
return "\n".join(
|
||||
[
|
||||
"Model: test_mql",
|
||||
"Mode: AMESim baseline passthrough",
|
||||
f"Samples: {run.sample_count}",
|
||||
f"Output schema signals: {run.output_schema.signal_count}",
|
||||
f"Compared signals: {run.signal_count}",
|
||||
f"Observation bindings: {run.observation_catalog.binding_count}",
|
||||
f"Max absolute error: {run.comparison.max_abs_error}",
|
||||
f"Max relative error: {run.comparison.max_rel_error}",
|
||||
]
|
||||
) + "\n"
|
||||
|
||||
|
||||
def run_test_mql_baseline(config: TestMqlBaselineScriptConfig | None = None):
|
||||
config = config or TestMqlBaselineScriptConfig()
|
||||
run = run_test_mql_baseline_passthrough(
|
||||
config.paths.archive_path,
|
||||
data_paths=config.execution.data_paths,
|
||||
)
|
||||
output_dir = config.paths.output_dir or _default_output_dir()
|
||||
if config.execution.write_summary:
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
(output_dir / "test_mql_baseline_summary.txt").write_text(
|
||||
format_test_mql_baseline_summary(run),
|
||||
encoding="utf-8",
|
||||
)
|
||||
return run, output_dir
|
||||
|
||||
|
||||
def main() -> None:
|
||||
run, output_dir = run_test_mql_baseline()
|
||||
print(format_test_mql_baseline_summary(run), end="")
|
||||
print(f"Output directory: {output_dir}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,88 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from PythonModels.scripts.run_test_mql_baseline import (
|
||||
TestMqlBaselineExecutionConfig,
|
||||
TestMqlBaselinePathConfig,
|
||||
TestMqlBaselineScriptConfig,
|
||||
format_test_mql_baseline_summary,
|
||||
run_test_mql_baseline,
|
||||
)
|
||||
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[1]
|
||||
TEST_MQL_AME = REPO_ROOT / "AmesimModels" / "test_mql.ame"
|
||||
|
||||
|
||||
class RunTestMqlBaselineScriptTests(unittest.TestCase):
|
||||
def test_baseline_script_writes_summary_for_selected_paths(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
output_dir = Path(tmpdir)
|
||||
run, returned_output_dir = run_test_mql_baseline(
|
||||
TestMqlBaselineScriptConfig(
|
||||
paths=TestMqlBaselinePathConfig(
|
||||
archive_path=TEST_MQL_AME,
|
||||
output_dir=output_dir,
|
||||
),
|
||||
execution=TestMqlBaselineExecutionConfig(
|
||||
data_paths=("press@pn_c1_8", "vol1@pn_brp2_8"),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
summary_path = output_dir / "test_mql_baseline_summary.txt"
|
||||
summary = summary_path.read_text(encoding="utf-8")
|
||||
|
||||
self.assertEqual(returned_output_dir, output_dir)
|
||||
self.assertEqual(run.sample_count, 1002)
|
||||
self.assertEqual(run.signal_count, 2)
|
||||
self.assertEqual(run.comparison.max_abs_error, 0.0)
|
||||
self.assertIn("Mode: AMESim baseline passthrough", summary)
|
||||
self.assertIn("Samples: 1002", summary)
|
||||
self.assertIn("Compared signals: 2", summary)
|
||||
self.assertIn("Max absolute error: 0.0", summary)
|
||||
|
||||
def test_baseline_script_can_skip_summary_file(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
output_dir = Path(tmpdir)
|
||||
run, returned_output_dir = run_test_mql_baseline(
|
||||
TestMqlBaselineScriptConfig(
|
||||
paths=TestMqlBaselinePathConfig(
|
||||
archive_path=TEST_MQL_AME,
|
||||
output_dir=output_dir,
|
||||
),
|
||||
execution=TestMqlBaselineExecutionConfig(
|
||||
write_summary=False,
|
||||
data_paths=("press@pn_c1_8",),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual(returned_output_dir, output_dir)
|
||||
self.assertEqual(run.signal_count, 1)
|
||||
self.assertFalse((output_dir / "test_mql_baseline_summary.txt").exists())
|
||||
|
||||
def test_format_test_mql_baseline_summary(self) -> None:
|
||||
run, _ = run_test_mql_baseline(
|
||||
TestMqlBaselineScriptConfig(
|
||||
paths=TestMqlBaselinePathConfig(archive_path=TEST_MQL_AME),
|
||||
execution=TestMqlBaselineExecutionConfig(
|
||||
write_summary=False,
|
||||
data_paths=("press@pn_c1_8",),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
summary = format_test_mql_baseline_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()
|
||||
Reference in new issue
Block a user