From a1b9686f1b41e4746b00ab3d6a34ae85ddc3a210 Mon Sep 17 00:00:00 2001 From: Sebbe Blokhuizen Date: Wed, 22 Jul 2026 16:27:05 +0200 Subject: [PATCH 1/2] allow typed waveforms --- tests/tendencies/test_constant.py | 33 +++++ tests/test_derived_waveform.py | 53 +++++++- tests/test_exporter.py | 38 ++++++ tests/test_waveform.py | 163 +++++++++++++++++++++++++ waveform_editor/base_waveform.py | 2 + waveform_editor/configuration.py | 5 +- waveform_editor/derived_waveform.py | 30 ++++- waveform_editor/tendencies/base.py | 8 +- waveform_editor/tendencies/constant.py | 5 +- waveform_editor/waveform.py | 82 +++++++++++-- 10 files changed, 405 insertions(+), 14 deletions(-) diff --git a/tests/tendencies/test_constant.py b/tests/tendencies/test_constant.py index c21f87cf..1a754101 100644 --- a/tests/tendencies/test_constant.py +++ b/tests/tendencies/test_constant.py @@ -69,6 +69,39 @@ def test_generate(): assert not tendency.annotations +def test_categorical_string_value(): + tendency = ConstantTendency(user_start=0, user_duration=1, user_value="ec") + assert tendency.value == "ec" + assert tendency.value_type is str + assert not tendency.annotations + + time, values = tendency.get_value() + assert np.all(time == np.array([0, 1])) + assert list(values) == ["ec", "ec"] + + +def test_categorical_bool_value(): + tendency = ConstantTendency(user_start=0, user_duration=1, user_value=True) + assert tendency.value is True + assert tendency.value_type is bool + assert not tendency.annotations + + time, values = tendency.get_value() + assert np.all(time == np.array([0, 1])) + assert list(values) == [True, True] + + +def test_categorical_int_value(): + tendency = ConstantTendency(user_start=0, user_duration=1, user_value=3) + assert tendency.value == 3 + assert tendency.value_type is int + assert not tendency.annotations + + time, values = tendency.get_value() + assert np.all(time == np.array([0, 1])) + assert list(values) == [3, 3] + + def test_declarative_assignments(): t1 = ConstantTendency(user_duration=1) t2 = ConstantTendency(user_duration=1) diff --git a/tests/test_derived_waveform.py b/tests/test_derived_waveform.py index cc910e09..1d2f9cf3 100644 --- a/tests/test_derived_waveform.py +++ b/tests/test_derived_waveform.py @@ -37,6 +37,7 @@ def const_waveform(config): const_value = 3 yaml_str = f"{name}: {const_value}" waveform = DerivedWaveform(yaml_str, name, config) + waveform.prepare_expression() config.add_waveform(waveform, ["root_group"]) return waveform, const_value, name, config @@ -78,6 +79,7 @@ def test_dependent_waveform(filled_config): name = "waveform/2" yaml_str = f"{name}: |\n 'waveform/1'" waveform = DerivedWaveform(yaml_str, name, filled_config) + waveform.prepare_expression() assert waveform.dependencies == {"waveform/1"} time_ret, value_ret = waveform.get_value() assert time_ret[0] == 5 @@ -92,6 +94,7 @@ def test_dependent_waveform_calc(filled_config): name = "waveform/2" yaml_str = f'{name}: |\n "waveform/1" * 10' waveform = DerivedWaveform(yaml_str, name, filled_config) + waveform.prepare_expression() assert waveform.dependencies == {"waveform/1"} time_ret, value_ret = waveform.get_value() assert time_ret[0] == 5 @@ -106,6 +109,7 @@ def test_dependent_waveform_numpy(filled_config): name = "waveform/2" yaml_str = f"{name}: |\n maximum('waveform/1' * 10, 150)" waveform = DerivedWaveform(yaml_str, name, filled_config) + waveform.prepare_expression() assert waveform.dependencies == {"waveform/1"} time_ret, value_ret = waveform.get_value() assert time_ret[0] == 5 @@ -120,9 +124,11 @@ def test_rename_waveform(filled_config): name = "waveform/2" yaml_str = f"{name}: |\n 'waveform/1'" waveform = DerivedWaveform(yaml_str, name, filled_config) + waveform.prepare_expression() + filled_config.add_waveform(waveform, ["root_group"]) assert waveform.dependencies == {"waveform/1"} assert waveform.get_yaml_string() == "'waveform/1'" - waveform.rename_dependency("waveform/1", "waveform/3") + filled_config.rename_waveform("waveform/1", "waveform/3") assert waveform.dependencies == {"waveform/3"} assert waveform.get_yaml_string() == "'waveform/3'" @@ -144,6 +150,7 @@ def test_function_access_control(filled_config): name = "waveform/2" yaml_str = f"{name}: |\n {expr}" waveform = DerivedWaveform(yaml_str, name, filled_config) + waveform.prepare_expression() time_ret = np.linspace(filled_config.start, filled_config.end, 100) if allowed: _, result = waveform.get_value(time_ret) @@ -151,3 +158,47 @@ def test_function_access_control(filled_config): else: with pytest.raises(NameError): waveform.get_value(time_ret) + + +def test_derived_waveform_type_matches_original(config): + original_name = "wf1" + original = Waveform( + waveform=[{"user_type": "constant", "user_value": 3, "line_number": 1}], + name=original_name, + ) + assert not original.annotations + assert original.value_type is int + config.add_waveform(original, ["root_group"]) + + derived_name = "wf2" + yaml_str = f"{derived_name}: |\n '{original_name}'" + derived = DerivedWaveform(yaml_str, derived_name, config) + derived.prepare_expression() + assert derived.dependencies == {original_name} + assert derived.value_type == original.value_type + + +def test_derived_waveform_type_mixing(config): + wf1_name = "wf1" + wf1 = Waveform( + waveform=[{"user_type": "constant", "user_value": 3, "line_number": 1}], + name=wf1_name, + ) + assert not wf1.annotations + assert wf1.value_type is int + config.add_waveform(wf1, ["root_group"]) + + wf2_name = "wf2" + wf2 = Waveform( + waveform=[{"user_type": "constant", "user_value": "test", "line_number": 2}], + name=wf2_name, + ) + assert not wf2.annotations + assert wf2.value_type is str + config.add_waveform(wf2, ["root_group"]) + + derived_name = "derived_waveform" + yaml_str = f"{derived_name}: |\n '{wf1_name}' + '{wf2_name}'" + derived = DerivedWaveform(yaml_str, derived_name, config) + derived.prepare_expression() + assert derived.annotations # not allowed to mix str and int type waveforms diff --git a/tests/test_exporter.py b/tests/test_exporter.py index c82b6af7..54a92c21 100644 --- a/tests/test_exporter.py +++ b/tests/test_exporter.py @@ -647,6 +647,44 @@ def test_export_constant(tmp_path): assert np.array_equal(ids.beam[2].phase.angle, [3.3e3] * 3) +def test_export_typed_waveforms(tmp_path): + """Check that constant waveforms of each supported value type (float, int, + str, bool) are exported correctly to their respective IDS nodes.""" + + yaml_str = """ + globals: + dd_version: 4.0.0 + core_profiles: + core_profiles/profiles_1d/electrons/temperature_validity: + - {type: constant, value: 0, duration: 2} + - {type: constant, value: 1, duration: 2} + core_profiles/profiles_1d/grid/psi_magnetic_axis: + - {type: constant, value: 1.5, duration: 2} + - {type: constant, value: 3.0, duration: 2} + core_profiles/profiles_1d/ion(1)/name: + - {type: constant, value: D, duration: 2} + - {type: constant, value: He, duration: 2} + core_profiles/profiles_1d/ion(1)/multiple_states_flag: + - {type: constant, value: true, duration: 2} + - {type: constant, value: false, duration: 2} + """ + file_path = f"{tmp_path}/test.nc" + times = np.array([0, 2.0]) + _export_ids(file_path, yaml_str, times) + + with imas.DBEntry(file_path, "r", dd_version="4.0.0") as dbentry: + core_profiles = dbentry.get("core_profiles", autoconvert=False) + assert core_profiles.profiles_1d[0].grid.psi_magnetic_axis == 1.5 + assert core_profiles.profiles_1d[0].electrons.temperature_validity == 0 + assert core_profiles.profiles_1d[0].ion[0].multiple_states_flag == 1 + assert core_profiles.profiles_1d[0].ion[0].name == "D" + + assert core_profiles.profiles_1d[1].grid.psi_magnetic_axis == 3.0 + assert core_profiles.profiles_1d[1].electrons.temperature_validity == 1 + assert core_profiles.profiles_1d[1].ion[0].multiple_states_flag == 0 + assert core_profiles.profiles_1d[1].ion[0].name == "He" + + def test_example_yaml(tmp_path): """Test for an example YAML file if all IDSs are correctly filled.""" diff --git a/tests/test_waveform.py b/tests/test_waveform.py index 6735d8a3..1846e576 100644 --- a/tests/test_waveform.py +++ b/tests/test_waveform.py @@ -7,6 +7,8 @@ from waveform_editor.tendencies.smooth import SmoothTendency from waveform_editor.waveform import Waveform +DD_VERSION = "3.42.0" + def test_empty(): waveform = Waveform() @@ -273,3 +275,164 @@ def test_overlap_derivatives(): expected = [2, 2, -1.5, -1.5, -1.5, -1.5, -1.5] values = waveform.get_derivative(np.linspace(0, 3, 7)) assert np.allclose(values, expected) + + +def test_multiple_tendencies_mixed(): + waveform = Waveform( + waveform=[ + { + "user_type": "constant", + "user_value": "ec", + "user_duration": 2, + "line_number": 1, + }, + { + "user_type": "constant", + "user_value": 3, + "user_duration": 2, + "line_number": 2, + }, + ] + ) + assert waveform.annotations + + +def test_dtype_flt_dd_path(): + """Test float field types.""" + + flt_dd_path = "ec_launchers/beam(1)/phase/angle" + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": "test", "line_number": 1}], + name=flt_dd_path, + dd_version=DD_VERSION, + ) + assert waveform.annotations + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 1, "line_number": 1}], + name=flt_dd_path, + dd_version=DD_VERSION, + ) + assert not waveform.annotations + assert waveform.value_type is int + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 2.5, "line_number": 1}], + name=flt_dd_path, + dd_version=DD_VERSION, + ) + assert not waveform.annotations + assert waveform.value_type is float + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": True, "line_number": 1}], + name=flt_dd_path, + dd_version=DD_VERSION, + ) + assert waveform.annotations + + +def test_dtype_int_dd_path(): + """Test int field types.""" + + int_dd_path = "pulse_schedule/ec/mode" + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": "test", "line_number": 1}], + name=int_dd_path, + dd_version=DD_VERSION, + ) + assert waveform.annotations + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 1, "line_number": 1}], + name=int_dd_path, + dd_version=DD_VERSION, + ) + assert not waveform.annotations + assert waveform.value_type is int + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 2.5, "line_number": 1}], + name=int_dd_path, + dd_version=DD_VERSION, + ) + assert waveform.annotations + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": True, "line_number": 1}], + name=int_dd_path, + dd_version=DD_VERSION, + ) + assert not waveform.annotations + assert waveform.value_type is bool + + +def test_dtype_str_dd_path(): + """Test string field types.""" + str_dd_path = "ec_launchers/ids_properties/comment" + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": "test", "line_number": 1}], + name=str_dd_path, + dd_version=DD_VERSION, + ) + assert not waveform.annotations + assert waveform.value_type is str + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 1, "line_number": 1}], + name=str_dd_path, + dd_version=DD_VERSION, + ) + assert waveform.annotations + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 2.5, "line_number": 1}], + name=str_dd_path, + dd_version=DD_VERSION, + ) + assert waveform.annotations + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": True, "line_number": 1}], + name=str_dd_path, + dd_version=DD_VERSION, + ) + assert waveform.annotations + + +def test_no_metadata_allows_any_type(): + """A waveform whose path does not resolve to any DD node is not restricted + to any particular value type.""" + name = "not_a_real_ids/path" + waveform = Waveform( + waveform=[ + {"user_type": "constant", "user_value": "anything", "line_number": 1} + ], + name=name, + ) + assert waveform.metadata is None + assert not waveform.annotations + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 1, "line_number": 1}], + name=name, + ) + assert waveform.metadata is None + assert not waveform.annotations + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 2.5, "line_number": 1}], + name=name, + ) + assert waveform.metadata is None + assert not waveform.annotations + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": True, "line_number": 1}], + name=name, + ) + assert waveform.metadata is None + assert not waveform.annotations diff --git a/waveform_editor/base_waveform.py b/waveform_editor/base_waveform.py index 6d88084b..f98d6571 100644 --- a/waveform_editor/base_waveform.py +++ b/waveform_editor/base_waveform.py @@ -9,6 +9,8 @@ class BaseWaveform(ABC): + value_type = float + def __init__(self, yaml_str, name, dd_version): yaml_dict = YAML().load(yaml_str) self.yaml = yaml_dict[name] if yaml_dict else None diff --git a/waveform_editor/configuration.py b/waveform_editor/configuration.py index 83c2d4a1..98e4e082 100644 --- a/waveform_editor/configuration.py +++ b/waveform_editor/configuration.py @@ -313,7 +313,10 @@ def parse_waveform(self, yaml_str): The parsed waveform object. """ self.parser.parse_errors = [] - return self.parser.parse_waveform(yaml_str) + waveform = self.parser.parse_waveform(yaml_str) + if isinstance(waveform, DerivedWaveform): + waveform.prepare_expression() + return waveform def _to_commented_map(self): """Return the configuration as a nested CommentedMap.""" diff --git a/waveform_editor/derived_waveform.py b/waveform_editor/derived_waveform.py index 79d623da..b1ad9110 100644 --- a/waveform_editor/derived_waveform.py +++ b/waveform_editor/derived_waveform.py @@ -72,7 +72,6 @@ def __init__(self, yaml_str, name, config, dd_version=None): self.dependencies = set() self.is_constant = False self.expression = None - self.prepare_expression() def prepare_expression(self): """Parse the YAML expression, extract dependencies, transform it for @@ -93,6 +92,7 @@ def prepare_expression(self): self.is_constant = not extractor.string_nodes self.expression = ast.unparse(modified_tree) self.dependencies = set(extractor.string_nodes) + self._validate_type() def rename_dependency(self, old_name, new_name): """Rename a dependency waveform in the expression. @@ -110,6 +110,34 @@ def rename_dependency(self, old_name, new_name): self.yaml = renamer.yaml self.prepare_expression() + def _validate_type(self): + """Warn if a dependency doesn't exist, or if the dependencies that do + exist don't share a common type. + """ + if not self.dependencies: + return + + dependency_types = set() + missing = set() + for dependency in self.dependencies: + try: + dependency_types.add(self.config[dependency].value_type) + except KeyError: + missing.add(dependency) + + if missing: + self.annotations.add(0, f"Unknown dependency: {sorted(missing)!r}\n") + return + + if len(dependency_types) > 1: + self.annotations.add( + 0, + "All dependencies of a derived waveform must have the same " + f"type. Found: {dependency_types}\n", + ) + else: + self.value_type = dependency_types.pop() + def _build_eval_context(self, time: np.ndarray) -> dict: """Build the evaluation context dictionary with dependencies resolved. diff --git a/waveform_editor/tendencies/base.py b/waveform_editor/tendencies/base.py index 50b43198..a016b516 100644 --- a/waveform_editor/tendencies/base.py +++ b/waveform_editor/tendencies/base.py @@ -61,8 +61,8 @@ class BaseTendency(param.Parameterized): values from the start value of this tendency. """, ) - start_value = param.Number(default=0.0, doc="Value at self.start") - end_value = param.Number(default=0.0, doc="Value at self.end") + start_value = param.Parameter(default=0.0, doc="Value at self.start") + end_value = param.Parameter(default=0.0, doc="Value at self.end") start_derivative = param.Number(default=0.0, doc="Derivative at self.start") end_derivative = param.Number(default=0.0, doc="Derivative at self.end") @@ -77,6 +77,10 @@ class BaseTendency(param.Parameterized): ) annotations = param.ClassSelector(class_=Annotations, default=Annotations()) allow_zero_duration = False + value_type = param.Parameter( + default=float, + doc="The type of the value this tendency produces. May be float, int, str, or bool.", + ) def __init__(self, **kwargs): super().__init__() diff --git a/waveform_editor/tendencies/constant.py b/waveform_editor/tendencies/constant.py index 6b1cd4af..a236ca6e 100644 --- a/waveform_editor/tendencies/constant.py +++ b/waveform_editor/tendencies/constant.py @@ -10,7 +10,7 @@ class ConstantTendency(BaseTendency): Constant tendency class for a constant signal. """ - user_value = param.Number( + user_value = param.Parameter( default=None, doc="The constant value of the tendency provided by the user.", ) @@ -33,7 +33,7 @@ def get_value( """ if time is None: time = np.array([self.start, self.end]) - values = self.value * np.ones(len(time)) + values = np.full(len(time), self.value) return time, values def get_derivative(self, time: np.ndarray) -> np.ndarray: @@ -73,4 +73,5 @@ def _calc_values(self): self.param.update( values_changed=values_changed, start_value_set=self.user_value is not None, + value_type=type(value), ) diff --git a/waveform_editor/waveform.py b/waveform_editor/waveform.py index 17011175..b0385b98 100644 --- a/waveform_editor/waveform.py +++ b/waveform_editor/waveform.py @@ -1,6 +1,7 @@ import io import numpy as np +from imas.ids_data_type import IDSDataType from ruamel.yaml import YAML from ruamel.yaml.comments import CommentedSeq @@ -15,6 +16,26 @@ from waveform_editor.tendencies.repeat import RepeatTendency from waveform_editor.tendencies.smooth import SmoothTendency +IDS_DATATYPE_MAP = { + float: IDSDataType.FLT, + str: IDSDataType.STR, + int: IDSDataType.INT, + bool: IDSDataType.INT, # Booleans don't exist in DD +} + +# Numpy dtype to build the evaluated values array with, keyed by value_type. Ints are +# evaluated as floats. Str/bool are categorical: held as a step across gaps rather than +# interpolated, so they use dtype=object -- NOT dtype=str, which numpy would fix at a +# single character's width (silently truncating any longer values written into it +# later) rather than sizing to what's actually assigned. +NUMPY_DTYPE_MAP = { + float: float, + int: float, + str: object, + bool: object, +} + + tendency_map = { "linear": LinearTendency, "sine-wave": SineWaveTendency, @@ -97,7 +118,13 @@ def _evaluate_tendencies(self, time, eval_derivatives=False): Returns: numpy array containing the computed values. """ - values = np.zeros_like(time, dtype=float) + dtype = float if eval_derivatives else NUMPY_DTYPE_MAP[self.value_type] + is_categorical = dtype is object + values = ( + np.empty(len(time), dtype=object) + if is_categorical + else np.zeros_like(time, dtype=dtype) + ) for i, tendency in enumerate(self.tendencies): mask = (time >= tendency.start) & (time <= tendency.end) @@ -107,17 +134,18 @@ def _evaluate_tendencies(self, time, eval_derivatives=False): else: _, values[mask] = tendency.get_value(time[mask]) - # Handle gaps between tendencies, we linearly interpolate between the - # gap values. + # Handle gaps between tendencies: interpolate for numeric values, hold + # the previous value for categorical ones. if i and tendency.prev_tendency.end < tendency.start: prev_tendency = tendency.prev_tendency mask = (time < tendency.start) & (time > prev_tendency.end) - slope = (tendency.start_value - prev_tendency.end_value) / ( - tendency.start - prev_tendency.end - ) if np.any(mask): if eval_derivatives: - values[mask] = slope + values[mask] = ( + tendency.start_value - prev_tendency.end_value + ) / (tendency.start - prev_tendency.end) + elif is_categorical: + values[mask] = prev_tendency.end_value else: values[mask] = np.interp( time[mask], @@ -173,11 +201,51 @@ def _process_waveform(self, waveform): self.tendencies[i - 1].set_next_tendency(self.tendencies[i]) self.tendencies[i].set_previous_tendency(self.tendencies[i - 1]) + self._validate_value_type() self.update_annotations() for tendency in self.tendencies: tendency.param.watch(self.update_annotations, "annotations") + def _validate_value_type(self): + """Determine this waveform's value type from its tendencies and set + ``self.value_type`` to reflect it. + """ + if not self.tendencies: + return + + self.value_type = self.tendencies[0].value_type + for tendency in self.tendencies[1:]: + if {tendency.value_type, self.value_type} <= {int, float}: + if tendency.value_type is float: + self.value_type = float + continue + if tendency.value_type != self.value_type: + error_msg = ( + f"Cannot mix {self.value_type.__name__} and " + f"{tendency.value_type.__name__} values within a single " + "waveform.\n" + ) + self.annotations.add(tendency.line_number, error_msg) + + # If a valid DD path is chosen, check if the value_type matches the DD type + if self.metadata is None: + return + + # An int value is also valid for a float field + int_for_flt = ( + self.value_type is int and self.metadata.data_type is IDSDataType.FLT + ) + if ( + not int_for_flt + and IDS_DATATYPE_MAP[self.value_type] != self.metadata.data_type + ): + error_msg = ( + "Type is not valid here: this waveform expects a " + f"{self.metadata.data_type}.\n" + ) + self.annotations.add(self.tendencies[0].line_number, error_msg) + def update_annotations(self, event=None): """Merges the annotations of the individual tendencies into the annotations of this waveform.""" From 64e95ca6ebdd42610e24818f763a81a4d8d88131 Mon Sep 17 00:00:00 2001 From: Sebbe Blokhuizen Date: Thu, 23 Jul 2026 15:48:00 +0200 Subject: [PATCH 2/2] add docs and minor fixes --- docs/source/derived.rst | 5 +++ docs/source/tendencies.rst | 16 ++++++++++ tests/tendencies/test_constant.py | 44 +++++++++++++------------- tests/tendencies/test_repeat.py | 14 ++++++++ tests/test_dependency_graph.py | 19 +++++++++++ tests/test_derived_waveform.py | 37 +++++++++++++++++----- tests/test_exporter.py | 3 +- waveform_editor/configuration.py | 11 ++----- waveform_editor/dependency_graph.py | 23 ++++++++++++++ waveform_editor/derived_waveform.py | 5 ++- waveform_editor/tendencies/base.py | 2 +- waveform_editor/tendencies/constant.py | 8 ++++- waveform_editor/tendencies/repeat.py | 13 ++++++++ waveform_editor/waveform.py | 5 --- 14 files changed, 157 insertions(+), 48 deletions(-) diff --git a/docs/source/derived.rst b/docs/source/derived.rst index 2133009c..3d196f3c 100644 --- a/docs/source/derived.rst +++ b/docs/source/derived.rst @@ -105,6 +105,11 @@ In the example below, waveform ``test/3`` is the sum of the waveforms ``test/1`` :width: 600px :align: center +.. note:: + A derived waveform's value type mirrors that of the waveform(s) it depends on. + When an expression references multiple waveforms, they must all share the same + value type (see :ref:`Value Types `). + Using NumPy Functions --------------------- diff --git a/docs/source/tendencies.rst b/docs/source/tendencies.rst index 88030e0e..843a41da 100644 --- a/docs/source/tendencies.rst +++ b/docs/source/tendencies.rst @@ -51,6 +51,22 @@ If the ``value`` is not specified, it will be set to the last value of the previ - {type: linear, to: 3, duration: 10} - {type: constant, duration: 10} +.. _constant-value-types: + +Value Types +----------- + +The ``value`` of a constant tendency may be a number, a string, or a boolean: + +.. code-block:: yaml + + - {type: constant, value: ohmic, duration: 2} + - {type: constant, value: nbi, duration: 2} + +.. warning:: + Integers and floats may be freely combined within a single waveform, but other + value types may not be combined with each other. + Linear Tendency =============== diff --git a/tests/tendencies/test_constant.py b/tests/tendencies/test_constant.py index 1a754101..7cc3e508 100644 --- a/tests/tendencies/test_constant.py +++ b/tests/tendencies/test_constant.py @@ -1,4 +1,5 @@ import numpy as np +import pytest from waveform_editor.tendencies.constant import ConstantTendency @@ -69,37 +70,36 @@ def test_generate(): assert not tendency.annotations -def test_categorical_string_value(): - tendency = ConstantTendency(user_start=0, user_duration=1, user_value="ec") - assert tendency.value == "ec" - assert tendency.value_type is str +@pytest.mark.parametrize( + "value", ["ec", True, 3, 3.5], ids=["str", "bool", "int", "float"] +) +def test_categorical_value(value): + tendency = ConstantTendency(user_start=0, user_duration=1, user_value=value) + assert tendency.value == value + assert tendency.value_type is type(value) assert not tendency.annotations time, values = tendency.get_value() assert np.all(time == np.array([0, 1])) - assert list(values) == ["ec", "ec"] + assert list(values) == [value, value] -def test_categorical_bool_value(): - tendency = ConstantTendency(user_start=0, user_duration=1, user_value=True) - assert tendency.value is True - assert tendency.value_type is bool - assert not tendency.annotations - - time, values = tendency.get_value() - assert np.all(time == np.array([0, 1])) - assert list(values) == [True, True] +def test_unsupported_value_type(): + tendency = ConstantTendency(user_start=0, user_duration=1, user_value=[1, 2, 3]) + assert tendency.annotations + assert tendency.value == 0.0 -def test_categorical_int_value(): - tendency = ConstantTendency(user_start=0, user_duration=1, user_value=3) - assert tendency.value == 3 - assert tendency.value_type is int - assert not tendency.annotations +@pytest.mark.parametrize( + "value", [5, 5.5, "ec", True], ids=["int", "float", "str", "bool"] +) +def test_inherited_value(value): + t1 = ConstantTendency(user_value=value, user_start=0, user_duration=1) + t2 = ConstantTendency(user_duration=1) + t2.set_previous_tendency(t1) - time, values = tendency.get_value() - assert np.all(time == np.array([0, 1])) - assert list(values) == [3, 3] + assert t2.value_type is type(value) + assert type(t2.value) is type(value) def test_declarative_assignments(): diff --git a/tests/tendencies/test_repeat.py b/tests/tendencies/test_repeat.py index 27d861cd..547a6262 100644 --- a/tests/tendencies/test_repeat.py +++ b/tests/tendencies/test_repeat.py @@ -177,6 +177,20 @@ def test_too_short(repeat_waveform): assert repeat_tendency.annotations[0]["type"] == "warning" +@pytest.mark.parametrize("value", ["ec", True], ids=["str", "bool"]) +def test_categorical_value_not_supported(value): + """Categorical values inside a repeat tendency are not allowed""" + repeat_tendency = RepeatTendency( + user_duration=4, + user_waveform=[ + {"user_type": "constant", "user_value": value, "user_duration": 1}, + ], + ) + assert repeat_tendency.annotations + times = np.linspace(0, 4, 9) + repeat_tendency.get_value(times) + + def test_period(repeat_waveform): """Check values when period is provided.""" repeat_waveform["user_period"] = 1 diff --git a/tests/test_dependency_graph.py b/tests/test_dependency_graph.py index e5abd67a..7c4a129d 100644 --- a/tests/test_dependency_graph.py +++ b/tests/test_dependency_graph.py @@ -87,3 +87,22 @@ def test_detect_cycles_with_start_node(): dg.graph["C"] = {"A"} with pytest.raises(RuntimeError): dg.detect_cycles("A") + + +def test_topological_order(): + dg = DependencyGraph() + dg.add_node("A", ["B"]) + dg.add_node("B", ["C"]) + dg.add_node("C", []) + + order = dg.topological_order() + assert set(order) == {"A", "B", "C"} + assert order.index("C") < order.index("B") < order.index("A") + + +def test_topological_order_ignores_leaf_dependencies(): + """Dependencies that are not nodes themselves should not show up.""" + dg = DependencyGraph() + dg.add_node("A", ["leaf"]) + + assert dg.topological_order() == ["A"] diff --git a/tests/test_derived_waveform.py b/tests/test_derived_waveform.py index 1d2f9cf3..15c46eb9 100644 --- a/tests/test_derived_waveform.py +++ b/tests/test_derived_waveform.py @@ -37,7 +37,6 @@ def const_waveform(config): const_value = 3 yaml_str = f"{name}: {const_value}" waveform = DerivedWaveform(yaml_str, name, config) - waveform.prepare_expression() config.add_waveform(waveform, ["root_group"]) return waveform, const_value, name, config @@ -79,7 +78,6 @@ def test_dependent_waveform(filled_config): name = "waveform/2" yaml_str = f"{name}: |\n 'waveform/1'" waveform = DerivedWaveform(yaml_str, name, filled_config) - waveform.prepare_expression() assert waveform.dependencies == {"waveform/1"} time_ret, value_ret = waveform.get_value() assert time_ret[0] == 5 @@ -94,7 +92,6 @@ def test_dependent_waveform_calc(filled_config): name = "waveform/2" yaml_str = f'{name}: |\n "waveform/1" * 10' waveform = DerivedWaveform(yaml_str, name, filled_config) - waveform.prepare_expression() assert waveform.dependencies == {"waveform/1"} time_ret, value_ret = waveform.get_value() assert time_ret[0] == 5 @@ -109,7 +106,6 @@ def test_dependent_waveform_numpy(filled_config): name = "waveform/2" yaml_str = f"{name}: |\n maximum('waveform/1' * 10, 150)" waveform = DerivedWaveform(yaml_str, name, filled_config) - waveform.prepare_expression() assert waveform.dependencies == {"waveform/1"} time_ret, value_ret = waveform.get_value() assert time_ret[0] == 5 @@ -124,7 +120,6 @@ def test_rename_waveform(filled_config): name = "waveform/2" yaml_str = f"{name}: |\n 'waveform/1'" waveform = DerivedWaveform(yaml_str, name, filled_config) - waveform.prepare_expression() filled_config.add_waveform(waveform, ["root_group"]) assert waveform.dependencies == {"waveform/1"} assert waveform.get_yaml_string() == "'waveform/1'" @@ -150,7 +145,6 @@ def test_function_access_control(filled_config): name = "waveform/2" yaml_str = f"{name}: |\n {expr}" waveform = DerivedWaveform(yaml_str, name, filled_config) - waveform.prepare_expression() time_ret = np.linspace(filled_config.start, filled_config.end, 100) if allowed: _, result = waveform.get_value(time_ret) @@ -173,11 +167,39 @@ def test_derived_waveform_type_matches_original(config): derived_name = "wf2" yaml_str = f"{derived_name}: |\n '{original_name}'" derived = DerivedWaveform(yaml_str, derived_name, config) - derived.prepare_expression() assert derived.dependencies == {original_name} assert derived.value_type == original.value_type +def test_derived_waveform_chain_type_order_independent(): + yaml_str = """ + root_group: + wf1: | + 'wf2' + 'wf3' + wf2: | + 'wf3' + wf3: | + 'wf4' + wf4: + - {type: constant, value: hello, duration: 2} + """ + config = WaveformConfiguration() + config.load_yaml(yaml_str) + + wf4 = config["wf4"] + wf3 = config["wf3"] + wf2 = config["wf2"] + wf1 = config["wf1"] + assert wf4.value_type is str + assert wf3.value_type is str + assert wf2.value_type is str + assert wf1.value_type is str + assert not wf4.annotations + assert not wf3.annotations + assert not wf2.annotations + assert not wf1.annotations + + def test_derived_waveform_type_mixing(config): wf1_name = "wf1" wf1 = Waveform( @@ -200,5 +222,4 @@ def test_derived_waveform_type_mixing(config): derived_name = "derived_waveform" yaml_str = f"{derived_name}: |\n '{wf1_name}' + '{wf2_name}'" derived = DerivedWaveform(yaml_str, derived_name, config) - derived.prepare_expression() assert derived.annotations # not allowed to mix str and int type waveforms diff --git a/tests/test_exporter.py b/tests/test_exporter.py index 54a92c21..2b1d9d60 100644 --- a/tests/test_exporter.py +++ b/tests/test_exporter.py @@ -648,8 +648,7 @@ def test_export_constant(tmp_path): def test_export_typed_waveforms(tmp_path): - """Check that constant waveforms of each supported value type (float, int, - str, bool) are exported correctly to their respective IDS nodes.""" + """Check that constant waveforms of each supported value type""" yaml_str = """ globals: diff --git a/waveform_editor/configuration.py b/waveform_editor/configuration.py index 98e4e082..8827beba 100644 --- a/waveform_editor/configuration.py +++ b/waveform_editor/configuration.py @@ -66,10 +66,8 @@ def load_yaml(self, yaml_str): try: self.parser.load_yaml(yaml_str) self._calculate_bounds() - for name, group in self.waveform_map.items(): - waveform = group[name] - if isinstance(waveform, DerivedWaveform): - waveform.prepare_expression() + for name in self.dependency_graph.topological_order(): + self[name].prepare_expression() self.has_changed = False except Exception as e: self.clear() @@ -313,10 +311,7 @@ def parse_waveform(self, yaml_str): The parsed waveform object. """ self.parser.parse_errors = [] - waveform = self.parser.parse_waveform(yaml_str) - if isinstance(waveform, DerivedWaveform): - waveform.prepare_expression() - return waveform + return self.parser.parse_waveform(yaml_str) def _to_commented_map(self): """Return the configuration as a nested CommentedMap.""" diff --git a/waveform_editor/dependency_graph.py b/waveform_editor/dependency_graph.py index 86d623c8..f6e318db 100644 --- a/waveform_editor/dependency_graph.py +++ b/waveform_editor/dependency_graph.py @@ -101,6 +101,29 @@ def rename_node(self, old_name, new_name): dependencies.add(new_name) return dependents + def topological_order(self): + """Return the nodes in dependency-first order: a node's dependencies always + appear before the node itself. + + Returns: + List of node names. + """ + visited = set() + result = [] + + def visit(node): + if node in visited: + return + visited.add(node) + for neighbor in self.graph.get(node, []): + visit(neighbor) + if node in self.graph: + result.append(node) + + for node in self.graph: + visit(node) + return result + def detect_cycles(self, start_node=None): """Detect cycles in the graph, optionally starting from a specific node. Raises RuntimeError if a circular dependency is found. diff --git a/waveform_editor/derived_waveform.py b/waveform_editor/derived_waveform.py index b1ad9110..6be7ae10 100644 --- a/waveform_editor/derived_waveform.py +++ b/waveform_editor/derived_waveform.py @@ -72,11 +72,13 @@ def __init__(self, yaml_str, name, config, dd_version=None): self.dependencies = set() self.is_constant = False self.expression = None + self.prepare_expression() def prepare_expression(self): """Parse the YAML expression, extract dependencies, transform it for evaluation, and compile it. """ + self.annotations.clear() if self.yaml is None: return @@ -130,10 +132,11 @@ def _validate_type(self): return if len(dependency_types) > 1: + type_names = sorted(t.__name__ for t in dependency_types) self.annotations.add( 0, "All dependencies of a derived waveform must have the same " - f"type. Found: {dependency_types}\n", + f"type. Found: {type_names}\n", ) else: self.value_type = dependency_types.pop() diff --git a/waveform_editor/tendencies/base.py b/waveform_editor/tendencies/base.py index a016b516..19eb2ca8 100644 --- a/waveform_editor/tendencies/base.py +++ b/waveform_editor/tendencies/base.py @@ -79,7 +79,7 @@ class BaseTendency(param.Parameterized): allow_zero_duration = False value_type = param.Parameter( default=float, - doc="The type of the value this tendency produces. May be float, int, str, or bool.", + doc="The value type of the this tendency. May be float, int, str, or bool.", ) def __init__(self, **kwargs): diff --git a/waveform_editor/tendencies/constant.py b/waveform_editor/tendencies/constant.py index a236ca6e..b8ae330d 100644 --- a/waveform_editor/tendencies/constant.py +++ b/waveform_editor/tendencies/constant.py @@ -10,7 +10,8 @@ class ConstantTendency(BaseTendency): Constant tendency class for a constant signal. """ - user_value = param.Parameter( + user_value = param.ClassSelector( + class_=(bool, int, float, str), default=None, doc="The constant value of the tendency provided by the user.", ) @@ -65,6 +66,11 @@ def _calc_values(self): else: value = self.user_value + # If value is inherited from previous tendency, value normalize it back to + # a plain Python type + if isinstance(value, np.generic): + value = value.item() + # Update state and cast to bool, as param does not like numpy booleans values_changed = bool(self.value != value) if values_changed: diff --git a/waveform_editor/tendencies/repeat.py b/waveform_editor/tendencies/repeat.py index 68a87c8a..c3edce35 100644 --- a/waveform_editor/tendencies/repeat.py +++ b/waveform_editor/tendencies/repeat.py @@ -28,7 +28,20 @@ def __init__(self, **kwargs): self.waveform = Waveform(waveform=waveform, is_repeated=True) self.period = 1 + # Categorical values are not supported inside a repeat tendency. + has_categorical_value = any( + t.value_type in (str, bool) for t in self.waveform.tendencies + ) + if has_categorical_value: + self.waveform.tendencies = [] super().__init__(**kwargs) + if has_categorical_value: + error_msg = ( + "Categorical (str/bool) values are not supported inside a repeat " + "tendency.\n" + ) + self.annotations.add(self.line_number, error_msg) + return if not self.waveform.tendencies: error_msg = "There are no tendencies in the repeated waveform.\n" self.annotations.add(self.line_number, error_msg) diff --git a/waveform_editor/waveform.py b/waveform_editor/waveform.py index b0385b98..e8fb5c64 100644 --- a/waveform_editor/waveform.py +++ b/waveform_editor/waveform.py @@ -23,11 +23,6 @@ bool: IDSDataType.INT, # Booleans don't exist in DD } -# Numpy dtype to build the evaluated values array with, keyed by value_type. Ints are -# evaluated as floats. Str/bool are categorical: held as a step across gaps rather than -# interpolated, so they use dtype=object -- NOT dtype=str, which numpy would fix at a -# single character's width (silently truncating any longer values written into it -# later) rather than sizing to what's actually assigned. NUMPY_DTYPE_MAP = { float: float, int: float,