Skip to content

Commit 7871d0a

Browse files
committed
FEAT: add serialization workspace loader
1 parent e68a987 commit 7871d0a

3 files changed

Lines changed: 210 additions & 0 deletions

File tree

src/ampform_dpd/io/serialization/__init__.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,3 +7,7 @@
77
This module is in preview, see https://github.com/ComPWA/ampform-dpd/issues/133 for
88
updates.
99
"""
10+
11+
from ampform_dpd.io.serialization.workspace import Workspace, load_workspace
12+
13+
__all__ = ["Workspace", "load_workspace"]
Lines changed: 146 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,146 @@
1+
from __future__ import annotations
2+
3+
import json
4+
from collections import Counter
5+
from collections.abc import Mapping
6+
from dataclasses import dataclass
7+
from pathlib import Path
8+
from types import MappingProxyType
9+
from typing import TYPE_CHECKING, Any, cast
10+
11+
from ampform_dpd.io.serialization.amplitude import formulate
12+
from ampform_dpd.io.serialization.dynamics import (
13+
PropagatorDynamicsBuilder,
14+
formulate_dynamics,
15+
formulate_form_factor,
16+
)
17+
18+
if TYPE_CHECKING:
19+
from os import PathLike
20+
21+
from ampform_dpd import AmplitudeModel, DefinedExpression
22+
from ampform_dpd.io.serialization.format import ModelDefinition
23+
24+
25+
@dataclass(frozen=True)
26+
class Workspace:
27+
"""Backend-independent formulation of a serialized amplitude model."""
28+
29+
definition: Mapping[str, Any]
30+
distributions: Mapping[str, AmplitudeModel]
31+
functions: Mapping[str, DefinedExpression]
32+
kinematics: Mapping[str, Any]
33+
reference_points: tuple[Mapping[str, Any], ...]
34+
checksums: tuple[Mapping[str, Any], ...]
35+
36+
37+
def load_workspace(
38+
source: str | PathLike[str] | Mapping[str, Any],
39+
*,
40+
builders: Mapping[str, PropagatorDynamicsBuilder] | None = None,
41+
) -> Workspace:
42+
"""Load and formulate every distribution in a serialized model."""
43+
definition = _load_definition(source)
44+
_raise_on_duplicate_names(definition)
45+
distributions = {
46+
distribution["name"]: formulate(
47+
_select_distribution(definition, distribution),
48+
additional_builders=dict(builders) if builders is not None else None,
49+
)
50+
for distribution in definition["distributions"]
51+
}
52+
functions = _formulate_functions(definition, builders)
53+
kinematics = {
54+
distribution["name"]: distribution["decay_description"]["kinematics"]
55+
for distribution in definition["distributions"]
56+
}
57+
checksums = definition.get("misc", {}).get("amplitude_model_checksums", [])
58+
return Workspace(
59+
definition=_freeze(definition),
60+
distributions=MappingProxyType(distributions),
61+
functions=MappingProxyType(functions),
62+
kinematics=_freeze(kinematics),
63+
reference_points=tuple(
64+
_freeze(point) for point in definition.get("parameter_points", [])
65+
),
66+
checksums=tuple(_freeze(checksum) for checksum in checksums),
67+
)
68+
69+
70+
def _load_definition(
71+
source: str | PathLike[str] | Mapping[str, Any],
72+
) -> ModelDefinition:
73+
if isinstance(source, Mapping):
74+
return cast("ModelDefinition", dict(source))
75+
with Path(source).open() as stream:
76+
return cast("ModelDefinition", json.load(stream))
77+
78+
79+
def _raise_on_duplicate_names(definition: ModelDefinition) -> None:
80+
for collection_name in ("distributions", "functions"):
81+
names = [item["name"] for item in definition[collection_name]]
82+
duplicates = sorted(name for name, count in Counter(names).items() if count > 1)
83+
if duplicates:
84+
msg = f"Duplicate {collection_name} names: {', '.join(duplicates)}"
85+
raise ValueError(msg)
86+
87+
88+
def _select_distribution(
89+
definition: ModelDefinition, distribution: Mapping[str, Any]
90+
) -> ModelDefinition:
91+
selected = dict(definition)
92+
selected["distributions"] = [dict(distribution)]
93+
return cast("ModelDefinition", selected)
94+
95+
96+
def _formulate_functions(
97+
definition: ModelDefinition,
98+
builders: Mapping[str, PropagatorDynamicsBuilder] | None,
99+
) -> dict[str, DefinedExpression]:
100+
formulated = {}
101+
for function in definition["functions"]:
102+
name = function["name"]
103+
formulated[name] = _formulate_function(name, definition, builders)
104+
return formulated
105+
106+
107+
def _formulate_function(
108+
name: str,
109+
definition: ModelDefinition,
110+
builders: Mapping[str, PropagatorDynamicsBuilder] | None,
111+
) -> DefinedExpression:
112+
for distribution in definition["distributions"]:
113+
model = _select_distribution(definition, distribution)
114+
for chain in distribution["decay_description"]["chains"]:
115+
for propagator in chain["propagators"]:
116+
if propagator.get("parametrization") == name:
117+
single_propagator_chain = dict(chain)
118+
single_propagator_chain["propagators"] = [propagator]
119+
return formulate_dynamics(
120+
cast("Any", single_propagator_chain),
121+
model,
122+
additional_definitions=(
123+
dict(builders) if builders is not None else None
124+
),
125+
)
126+
for vertex in chain["vertices"]:
127+
if vertex.get("formfactor") == name:
128+
return formulate_form_factor(vertex, model)
129+
function_type = next(
130+
function["type"]
131+
for function in definition["functions"]
132+
if function["name"] == name
133+
)
134+
msg = (
135+
f"Cannot formulate function {name!r} of type {function_type!r}: "
136+
"it has no propagator or form-factor context"
137+
)
138+
raise NotImplementedError(msg)
139+
140+
141+
def _freeze(value: Any) -> Any:
142+
if isinstance(value, Mapping):
143+
return MappingProxyType({key: _freeze(item) for key, item in value.items()})
144+
if isinstance(value, list):
145+
return tuple(_freeze(item) for item in value)
146+
return value
Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,60 @@
1+
from __future__ import annotations
2+
3+
import json
4+
from copy import deepcopy
5+
from types import MappingProxyType
6+
from typing import TYPE_CHECKING
7+
8+
import pytest
9+
10+
from ampform_dpd.io.serialization import load_workspace
11+
12+
if TYPE_CHECKING:
13+
from pathlib import Path
14+
15+
from ampform_dpd.io.serialization.format import ModelDefinition
16+
17+
18+
@pytest.mark.parametrize("from_path", [False, True], ids=["mapping", "path"])
19+
def test_load_workspace(
20+
model_definition: ModelDefinition, tmp_path: Path, from_path: bool
21+
):
22+
source = model_definition
23+
if from_path:
24+
source = tmp_path / "model.json"
25+
source.write_text(json.dumps(model_definition))
26+
workspace = load_workspace(source)
27+
assert tuple(workspace.distributions) == ("default_model",)
28+
assert set(workspace.functions) == {
29+
item["name"] for item in model_definition["functions"]
30+
}
31+
assert isinstance(workspace.definition, MappingProxyType)
32+
with pytest.raises(TypeError):
33+
workspace.distributions["new"] = workspace.distributions["default_model"] # ty: ignore[invalid-assignment]
34+
35+
36+
def test_loads_multiple_distributions(model_definition: ModelDefinition):
37+
definition = deepcopy(model_definition)
38+
second = deepcopy(definition["distributions"][0])
39+
second["name"] = "second"
40+
definition["distributions"].append(second)
41+
workspace = load_workspace(definition)
42+
assert tuple(workspace.distributions) == ("default_model", "second")
43+
44+
45+
@pytest.mark.parametrize("collection", ["distributions", "functions"])
46+
def test_rejects_duplicate_names(model_definition: ModelDefinition, collection: str):
47+
definition = deepcopy(model_definition)
48+
if collection == "distributions":
49+
definition["distributions"].append(deepcopy(definition["distributions"][0]))
50+
else:
51+
definition["functions"].append(deepcopy(definition["functions"][0]))
52+
with pytest.raises(ValueError, match=rf"Duplicate {collection} names"):
53+
load_workspace(definition)
54+
55+
56+
def test_reports_unsupported_unreferenced_function(model_definition: ModelDefinition):
57+
definition = deepcopy(model_definition)
58+
definition["functions"].append({"name": "orphan", "type": "Unknown"})
59+
with pytest.raises(NotImplementedError, match=r"orphan.*Unknown"):
60+
load_workspace(definition)

0 commit comments

Comments
 (0)