diff --git a/docs/source/derived.rst b/docs/source/derived.rst index 2133009c..28ef0574 100644 --- a/docs/source/derived.rst +++ b/docs/source/derived.rst @@ -105,6 +105,12 @@ 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 `). Derived waveforms + cannot depend on string-typed waveforms. + Using NumPy Functions --------------------- diff --git a/docs/source/tendencies.rst b/docs/source/tendencies.rst index 88030e0e..0dd3bd45 100644 --- a/docs/source/tendencies.rst +++ b/docs/source/tendencies.rst @@ -51,6 +51,45 @@ 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 (integer or floating-point) or a string. + +An integer value: + +.. code-block:: yaml + + - {type: constant, value: 3, duration: 2} + +A floating-point value: + +.. code-block:: yaml + + - {type: constant, value: 3.5, duration: 2} + +A string value: + +.. code-block:: yaml + + - {type: constant, value: ohmic, duration: 2} + - {type: constant, value: nbi, duration: 2} + +It is also possible to use ``True`` and ``False`` as the ``value``, in which case the +waveform will be considered an integer type (where ``True = 1``, and ``False = 0``): + +.. code-block:: yaml + + - {type: constant, value: True, duration: 2} + - {type: constant, value: False, duration: 2} + +.. warning:: + Integers and floats may be freely combined within a single waveform (the + resulting waveform will be float-typed), but it's not allowed to combine string + values with numbers in a waveform. + Linear Tendency =============== diff --git a/tests/tendencies/test_constant.py b/tests/tendencies/test_constant.py index c21f87cf..158fdeaa 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,6 +70,59 @@ def test_generate(): assert not tendency.annotations +@pytest.mark.parametrize("value", ["ec", 3, 3.5], ids=["str", "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) == [value, value] + + +@pytest.mark.parametrize("value,expected", [(True, 1), (False, 0)]) +def test_bool_value_treated_as_int(value, expected): + tendency = ConstantTendency(user_start=0, user_duration=1, user_value=value) + assert tendency.value == expected + assert type(tendency.value) is int + assert tendency.value_type is int + assert not tendency.annotations + + +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 + + +@pytest.mark.parametrize("value", [5, 5.5, "ec"], ids=["int", "float", "str"]) +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) + + assert t2.value_type is type(value) + assert type(t2.value) is type(value) + + +@pytest.mark.parametrize("value", [5, 5.5, "ec"], ids=["int", "float", "str"]) +def test_inherited_value_stays_plain_type_across_chain(value): + """start_value/end_value/value must be plain Python types""" + t1 = ConstantTendency(user_value=value, user_start=0, user_duration=1) + t2 = ConstantTendency(user_duration=1) + t2.set_previous_tendency(t1) + t3 = ConstantTendency(user_duration=1) + t3.set_previous_tendency(t2) + + for tendency in (t1, t2, t3): + assert tendency.value_type is type(value) + assert type(tendency.value) is type(value) + assert type(tendency.start_value) is type(value) + assert type(tendency.end_value) is type(value) + + def test_declarative_assignments(): t1 = ConstantTendency(user_duration=1) t2 = ConstantTendency(user_duration=1) diff --git a/tests/tendencies/test_repeat.py b/tests/tendencies/test_repeat.py index 27d861cd..a90e0b41 100644 --- a/tests/tendencies/test_repeat.py +++ b/tests/tendencies/test_repeat.py @@ -177,6 +177,19 @@ def test_too_short(repeat_waveform): assert repeat_tendency.annotations[0]["type"] == "warning" +def test_string_value_not_supported(): + """String values inside a repeat tendency are not allowed""" + repeat_tendency = RepeatTendency( + user_duration=4, + user_waveform=[ + {"user_type": "constant", "user_value": "ec", "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_configuration.py b/tests/test_configuration.py index 88ea1905..59f1810f 100644 --- a/tests/test_configuration.py +++ b/tests/test_configuration.py @@ -3,6 +3,7 @@ import pytest from waveform_editor.configuration import WaveformConfiguration +from waveform_editor.derived_waveform import DerivedWaveform from waveform_editor.tendencies.constant import ConstantTendency from waveform_editor.tendencies.linear import LinearTendency from waveform_editor.tendencies.periodic.sine_wave import SineWaveTendency @@ -109,6 +110,54 @@ def test_replace_waveform(config): config.replace_waveform(waveform3) +def test_replace_waveform_revalidates_dependents(config): + """Test if replacing a waveform re-evaluates the value_type of derived waveforms""" + path = ["ec_launchers"] + + int_waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 3, "user_duration": 1}], + name="A", + ) + config.add_waveform(int_waveform, path) + + derived_b = DerivedWaveform("B: \"'A'\"", "B", config) + config.add_waveform(derived_b, path) + derived_c = DerivedWaveform("C: \"'B'\"", "C", config) + config.add_waveform(derived_c, path) + + assert derived_b.value_type is int + assert derived_c.value_type is int + + flt_waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 3.5, "user_duration": 1}], + name="A", + ) + config.replace_waveform(flt_waveform) + + assert derived_b.value_type is float + assert derived_c.value_type is float + + +def test_add_waveform_revalidates_previously_missing_dependency(config): + """Test if a derived waveform's error is cleared once its dependency is added to + the configuration.""" + path = ["ec_launchers"] + + derived = DerivedWaveform("B: \"'A'\"", "B", config) + config.add_waveform(derived, path) + assert derived.annotations + assert derived.value_type is float + + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 3, "user_duration": 1}], + name="A", + ) + config.add_waveform(waveform, path) + + assert not derived.annotations + assert derived.value_type is int + + def test_remove_waveform(config): """Test if waveforms are removed correctly from configuration.""" diff --git a/tests/test_dependency_graph.py b/tests/test_dependency_graph.py index e5abd67a..58f5aa5c 100644 --- a/tests/test_dependency_graph.py +++ b/tests/test_dependency_graph.py @@ -87,3 +87,20 @@ 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", []) + + assert dg.topological_order() == ["C", "B", "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 cc910e09..3cbf4c78 100644 --- a/tests/test_derived_waveform.py +++ b/tests/test_derived_waveform.py @@ -120,9 +120,10 @@ def test_rename_waveform(filled_config): name = "waveform/2" yaml_str = f"{name}: |\n 'waveform/1'" waveform = DerivedWaveform(yaml_str, name, filled_config) + 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'" @@ -151,3 +152,92 @@ def test_function_access_control(filled_config): else: with pytest.raises(NameError): waveform.get_value(time_ret) + + +def test_derived_waveform_type_matches_original(): + yaml_str = """ + root_group: + wf1: + - {type: constant, value: 3, duration: 2} + wf2: | + 'wf1' + """ + config = WaveformConfiguration() + config.load_yaml(yaml_str) + + assert not config["wf1"].annotations + assert config["wf1"].value_type is int + assert config["wf2"].dependencies == {"wf1"} + assert config["wf2"].value_type == config["wf1"].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: 3, duration: 2} + """ + config = WaveformConfiguration() + config.load_yaml(yaml_str) + + for wf in ["wf1", "wf2", "wf3", "wf4"]: + assert config[wf].value_type is int + assert not config[wf].annotations + + +def test_derived_waveform_type_mixing(): + yaml_str = """ + root_group: + wf1: + - {type: constant, value: 3, duration: 2} + wf2: + - {type: constant, value: test, duration: 2} + derived_waveform: | + 'wf1' + 'wf2' + """ + config = WaveformConfiguration() + config.load_yaml(yaml_str) + + assert config["wf1"].value_type is int + assert config["wf2"].value_type is str + assert config["derived_waveform"].annotations # not allowed to mix str and int + + +def test_derived_waveform_disallows_string_dependency(): + yaml_str = """ + root_group: + wf1: + - {type: constant, value: ohmic, duration: 2} + derived_waveform: | + 'wf1' + """ + config = WaveformConfiguration() + config.load_yaml(yaml_str) + + assert config["wf1"].value_type is str + assert config["derived_waveform"].annotations + + +def test_derived_waveform_int_float_mixing(): + yaml_str = """ + root_group: + wf1: + - {type: constant, value: 3, duration: 2} + wf2: + - {type: constant, value: 3.5, duration: 2} + derived_waveform: | + 'wf1' + 'wf2' + """ + config = WaveformConfiguration() + config.load_yaml(yaml_str) + + assert config["wf1"].value_type is int + assert config["wf2"].value_type is float + assert not config["derived_waveform"].annotations + assert config["derived_waveform"].value_type is float diff --git a/tests/test_exporter.py b/tests/test_exporter.py index c82b6af7..0b32e634 100644 --- a/tests/test_exporter.py +++ b/tests/test_exporter.py @@ -647,6 +647,43 @@ 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""" + + 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: 1, duration: 2} + - {type: constant, value: 0, 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..9942af6f 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,135 @@ 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 + + +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 + + +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 + + +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 diff --git a/waveform_editor/base_waveform.py b/waveform_editor/base_waveform.py index 6d88084b..a1821a4f 100644 --- a/waveform_editor/base_waveform.py +++ b/waveform_editor/base_waveform.py @@ -17,6 +17,7 @@ def __init__(self, yaml_str, name, dd_version): self.metadata = self.get_metadata(dd_version) self.annotations = Annotations() self.units = self.metadata.units if self.metadata else "a.u." + self.value_type = float @abstractmethod def get_value( diff --git a/waveform_editor/configuration.py b/waveform_editor/configuration.py index 83c2d4a1..12b555bd 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() @@ -98,6 +96,7 @@ def add_waveform(self, waveform, path): group.waveforms[waveform.name] = waveform self.waveform_map[waveform.name] = group self._calculate_bounds() + self._revalidate_dependents() self.has_changed = True def rename_waveform(self, old_name, new_name): @@ -141,11 +140,11 @@ def check_safe_to_replace(self, waveform): self.dependency_graph.check_safe_to_replace( waveform.name, waveform.dependencies ) - for dependent_wf in waveform.dependencies: - if dependent_wf not in self.waveform_map: - raise ValueError( - f"Cannot depend on waveform '{dependent_wf}', it does not exist!" - ) + + def _revalidate_dependents(self): + """Re-run validation for all derived waveforms, in dependency order.""" + for name in self.dependency_graph.topological_order(): + self[name].prepare_expression() def _validate_name(self, name): """Check that name doesn't exist already. If it does a ValueError is raised. @@ -175,6 +174,7 @@ def replace_waveform(self, waveform): group = self.waveform_map[waveform.name] group.waveforms[waveform.name] = waveform self._calculate_bounds() + self._revalidate_dependents() self.has_changed = True def remove_waveform(self, name): diff --git a/waveform_editor/dependency_graph.py b/waveform_editor/dependency_graph.py index 86d623c8..8f978823 100644 --- a/waveform_editor/dependency_graph.py +++ b/waveform_editor/dependency_graph.py @@ -101,6 +101,30 @@ 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. Dependencies that are not nodes themselves + do not show up. + + 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 79d623da..89961188 100644 --- a/waveform_editor/derived_waveform.py +++ b/waveform_editor/derived_waveform.py @@ -78,6 +78,9 @@ def prepare_expression(self): """Parse the YAML expression, extract dependencies, transform it for evaluation, and compile it. """ + # This method can be called multiple times so clear stale annotations from a + # previous call before re-validating. + self.annotations.clear() if self.yaml is None: return @@ -93,6 +96,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_dependencies() def rename_dependency(self, old_name, new_name): """Rename a dependency waveform in the expression. @@ -110,6 +114,42 @@ def rename_dependency(self, old_name, new_name): self.yaml = renamer.yaml self.prepare_expression() + def _validate_dependencies(self): + """Warn if a dependency doesn't exist, or if the dependencies that do + exist don't have a compatible type. Mixing int and float dependencies + is allowed and results in a float-typed derived waveform. + """ + 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 str in dependency_types: + self.annotations.add( + 0, "Derived waveforms cannot depend on string-typed waveforms.\n" + ) + elif dependency_types <= {int, float}: + self.value_type = float if float in dependency_types else int + elif len(dependency_types) == 1: + self.value_type = dependency_types.pop() + else: + 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, or be a mix of int and float. Found: {type_names}\n", + ) + 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..017418d1 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 value type of the this tendency. May be float, int, or str.", + ) def __init__(self, **kwargs): super().__init__() @@ -223,7 +227,7 @@ def _get_value_and_derivative(self, time): """Get the value and derivative of the tendency at a given time.""" _, value_array = self.get_value(np.array([time])) derivative_array = self.get_derivative(np.array([time])) - return value_array[0], derivative_array[0] + return value_array.item(0), derivative_array.item(0) @abstractmethod def get_value( diff --git a/waveform_editor/tendencies/constant.py b/waveform_editor/tendencies/constant.py index 6b1cd4af..0f001096 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.Number( + user_value = param.ClassSelector( + class_=(int, float, str), default=None, doc="The constant value of the tendency provided by the user.", ) @@ -33,7 +34,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: @@ -65,12 +66,18 @@ def _calc_values(self): else: value = self.user_value + if isinstance(value, bool): + value = int(value) + # Update state and cast to bool, as param does not like numpy booleans - values_changed = bool(self.value != value) + values_changed = bool( + self.value != value or type(self.value) is not type(value) + ) if values_changed: self.value = value # Ensure watchers are called after both values are updated 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/tendencies/repeat.py b/waveform_editor/tendencies/repeat.py index 68a87c8a..e74899e1 100644 --- a/waveform_editor/tendencies/repeat.py +++ b/waveform_editor/tendencies/repeat.py @@ -29,6 +29,11 @@ def __init__(self, **kwargs): self.waveform = Waveform(waveform=waveform, is_repeated=True) self.period = 1 super().__init__(**kwargs) + if self.waveform.value_type is str: + self.waveform.tendencies = [] + error_msg = "String 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 17011175..7a7a69b7 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,19 @@ from waveform_editor.tendencies.repeat import RepeatTendency from waveform_editor.tendencies.smooth import SmoothTendency +IDS_DATATYPES = { + float: {IDSDataType.FLT}, + str: {IDSDataType.STR}, + int: {IDSDataType.INT, IDSDataType.FLT}, # An int is also valid for a float field +} + +NUMPY_DTYPE_MAP = { + float: float, + int: float, + str: object, +} + + tendency_map = { "linear": LinearTendency, "sine-wave": SineWaveTendency, @@ -97,7 +111,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 +127,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 +194,44 @@ 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 + + value_types = set(tendency.value_type for tendency in self.tendencies) + if len(value_types) == 1: + self.value_type = value_types.pop() + elif value_types == {int, float}: + self.value_type = float + else: + type_names = ", ".join(sorted(t.__name__ for t in value_types)) + error_msg = ( + f"Cannot mix string and numerical tendency value types within a single " + f"waveform. Found: {type_names}." + ) + self.annotations.add(0, error_msg) + return + + # If a valid DD path is chosen, check if the value_type matches the DD type + if ( + self.metadata is not None + and self.metadata.data_type not in IDS_DATATYPES[self.value_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."""