From 53e2550dd434c7b3621660ea77c63e7a45af3d12 Mon Sep 17 00:00:00 2001 From: Daan van Vugt Date: Thu, 25 Jun 2026 11:09:35 +0200 Subject: [PATCH 1/6] dtype-aware tendency values --- tests/tendencies/test_constant.py | 31 ++++++++++++++++++++++++++ tests/test_waveform.py | 29 ++++++++++++++++++++++++ waveform_editor/tendencies/base.py | 5 +++-- waveform_editor/tendencies/constant.py | 9 ++++++-- waveform_editor/waveform.py | 23 +++++++++++++------ 5 files changed, 86 insertions(+), 11 deletions(-) diff --git a/tests/tendencies/test_constant.py b/tests/tendencies/test_constant.py index c21f87cf..63eda92a 100644 --- a/tests/tendencies/test_constant.py +++ b/tests/tendencies/test_constant.py @@ -86,3 +86,34 @@ def test_declarative_assignments(): assert t2.value == 6 assert not t1.annotations assert not t2.annotations + + +def test_integer_value_not_coerced(): + """Integer inputs are kept as integers, not coerced to float.""" + tendency = ConstantTendency(user_duration=1, user_value=5) + assert tendency.value == 5 + assert isinstance(tendency.value, int) + assert not tendency.is_string + assert not tendency.annotations + + +def test_string_value(): + """A constant tendency can hold a string (non-numeric) value.""" + tendency = ConstantTendency(user_duration=2, user_value="nbi") + assert tendency.value == "nbi" + assert tendency.is_string + assert tendency.start_value == "nbi" + assert tendency.end_value == "nbi" + + _, values = tendency.get_value(np.array([0.0, 1.0, 2.0])) + assert list(values) == ["nbi", "nbi", "nbi"] + assert not tendency.annotations + + +def test_string_value_chained(): + """A value-less string constant inherits the previous string value.""" + prev = ConstantTendency(user_value="ec", user_start=0, user_duration=1) + tendency = ConstantTendency(user_duration=1) + tendency.set_previous_tendency(prev) + assert tendency.value == "ec" + assert tendency.is_string diff --git a/tests/test_waveform.py b/tests/test_waveform.py index 1d38ebad..52a92bd5 100644 --- a/tests/test_waveform.py +++ b/tests/test_waveform.py @@ -320,3 +320,32 @@ 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_string_waveform(): + """A waveform of string constants evaluates as a zero-order-hold step function.""" + waveform = Waveform( + waveform=[ + {"user_type": "constant", "user_value": "ohmic", "user_duration": 2}, + {"user_type": "constant", "user_value": "nbi", "user_duration": 2}, + ] + ) + assert waveform.is_string + + _, values = waveform.get_value(np.array([0.0, 1.0, 2.0, 3.0])) + # The switch happens at t=2; later tendencies take precedence at the boundary. + assert list(values) == ["ohmic", "ohmic", "nbi", "nbi"] + + # Values are held (not interpolated) when extrapolating beyond the domain. + _, extrap = waveform.get_value(np.array([-1.0, 5.0])) + assert list(extrap) == ["ohmic", "nbi"] + + +def test_numeric_waveform_evaluates_to_float(): + """A numeric waveform evaluates to a float array (the int value is not stepwise).""" + waveform = Waveform( + waveform=[{"user_type": "constant", "user_value": 8, "user_duration": 3}] + ) + assert not waveform.is_string + _, values = waveform.get_value(np.array([0.0, 1.0, 2.0, 3.0])) + assert values.dtype == float diff --git a/waveform_editor/tendencies/base.py b/waveform_editor/tendencies/base.py index 50b43198..7bb40535 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,7 @@ class BaseTendency(param.Parameterized): ) annotations = param.ClassSelector(class_=Annotations, default=Annotations()) allow_zero_duration = False + is_string = False def __init__(self, **kwargs): super().__init__() diff --git a/waveform_editor/tendencies/constant.py b/waveform_editor/tendencies/constant.py index 6b1cd4af..62b1e4a7 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.", ) @@ -19,6 +19,11 @@ def __init__(self, **kwargs): self.value = 0.0 super().__init__(**kwargs) + @property + def is_string(self): + """Whether this constant holds a string (non-numeric) value.""" + return isinstance(self.value, str) + def get_value( self, time: np.ndarray | None = None ) -> tuple[np.ndarray, np.ndarray]: @@ -33,7 +38,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: diff --git a/waveform_editor/waveform.py b/waveform_editor/waveform.py index fd1c7ed6..9ecb98f3 100644 --- a/waveform_editor/waveform.py +++ b/waveform_editor/waveform.py @@ -69,6 +69,11 @@ def _bind_imports(self): tendency.resolver = resolver tendency.default_path = self.name + @property + def is_string(self): + """Whether this waveform produces string (non-numeric) values.""" + return any(t.is_string for t in self.tendencies) + def get_value( self, time: np.ndarray | None = None ) -> tuple[np.ndarray, np.ndarray]: @@ -125,7 +130,11 @@ def _evaluate_tendencies(self, time, eval_derivatives=False): Returns: numpy array containing the computed values. """ - values = np.zeros_like(time, dtype=float) + is_string = self.is_string and not eval_derivatives + if is_string: + values = np.empty(len(time), dtype=object) + else: + values = np.zeros_like(time, dtype=float) for i, tendency in enumerate(self.tendencies): mask = (time >= tendency.start) & (time <= tendency.end) @@ -135,17 +144,17 @@ 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 (hold for strings, interpolate otherwise). 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_string: + values[mask] = prev_tendency.end_value else: values[mask] = np.interp( time[mask], From 0cf91b80dc61512bf12eab075e6375fa73cd906d Mon Sep 17 00:00:00 2001 From: Daan van Vugt Date: Thu, 25 Jun 2026 18:53:59 +0200 Subject: [PATCH 2/6] Generalize categorical tendency values, validate mixing, document - Rename is_string -> is_categorical so booleans are treated like strings (held as a zero-order-hold step rather than coerced to float and interpolated) - Reject mixing categorical (string/boolean) and numeric values within a single waveform, since the gaps between them cannot be interpolated - Document the constant tendency value types (float/int/categorical) in the docs --- docs/source/tendencies.rst | 21 +++++++++++++ tests/tendencies/test_constant.py | 21 +++++++++++-- tests/test_waveform.py | 41 ++++++++++++++++++++++++-- waveform_editor/tendencies/base.py | 2 +- waveform_editor/tendencies/constant.py | 8 +++-- waveform_editor/waveform.py | 31 +++++++++++++++---- 6 files changed, 109 insertions(+), 15 deletions(-) diff --git a/docs/source/tendencies.rst b/docs/source/tendencies.rst index 72d610aa..bb0a693a 100644 --- a/docs/source/tendencies.rst +++ b/docs/source/tendencies.rst @@ -53,6 +53,27 @@ 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} +Value types +----------- + +The ``value`` is type-aware and is not restricted to floating-point numbers: + +* **Floats** are interpolated across gaps between tendencies, as usual. +* **Integers** are preserved as integers (they are not coerced to float). +* **Categorical values** (strings or booleans) describe non-numeric signals, such + as a heating scheme (``"nbi"``, ``"ec"``) or an on/off flag. These are held as a + step (zero-order hold) rather than interpolated: across a gap, the previous value + is carried forward, and the same applies when extrapolating beyond the waveform. + +.. code-block:: yaml + + - {type: constant, value: ohmic, duration: 2} + - {type: constant, value: nbi, duration: 2} + +.. note:: + Categorical and numeric values cannot be mixed within a single waveform, because + the gaps between them cannot be interpolated. Doing so raises a validation error. + Linear Tendency =============== diff --git a/tests/tendencies/test_constant.py b/tests/tendencies/test_constant.py index 63eda92a..b557dfec 100644 --- a/tests/tendencies/test_constant.py +++ b/tests/tendencies/test_constant.py @@ -93,7 +93,7 @@ def test_integer_value_not_coerced(): tendency = ConstantTendency(user_duration=1, user_value=5) assert tendency.value == 5 assert isinstance(tendency.value, int) - assert not tendency.is_string + assert not tendency.is_categorical assert not tendency.annotations @@ -101,7 +101,7 @@ def test_string_value(): """A constant tendency can hold a string (non-numeric) value.""" tendency = ConstantTendency(user_duration=2, user_value="nbi") assert tendency.value == "nbi" - assert tendency.is_string + assert tendency.is_categorical assert tendency.start_value == "nbi" assert tendency.end_value == "nbi" @@ -110,10 +110,25 @@ def test_string_value(): assert not tendency.annotations +def test_boolean_value(): + """A constant tendency can hold a boolean value, kept as a (categorical) bool.""" + tendency = ConstantTendency(user_duration=2, user_value=True) + assert tendency.value is True + assert tendency.is_categorical + # start/end values round-trip through numpy, so compare by value (np.bool_) + assert bool(tendency.start_value) is True + assert bool(tendency.end_value) is True + + _, values = tendency.get_value(np.array([0.0, 1.0, 2.0])) + assert [bool(v) for v in values] == [True, True, True] + assert values.dtype == bool + assert not tendency.annotations + + def test_string_value_chained(): """A value-less string constant inherits the previous string value.""" prev = ConstantTendency(user_value="ec", user_start=0, user_duration=1) tendency = ConstantTendency(user_duration=1) tendency.set_previous_tendency(prev) assert tendency.value == "ec" - assert tendency.is_string + assert tendency.is_categorical diff --git a/tests/test_waveform.py b/tests/test_waveform.py index 52a92bd5..0d6b3e19 100644 --- a/tests/test_waveform.py +++ b/tests/test_waveform.py @@ -330,7 +330,7 @@ def test_string_waveform(): {"user_type": "constant", "user_value": "nbi", "user_duration": 2}, ] ) - assert waveform.is_string + assert waveform.is_categorical _, values = waveform.get_value(np.array([0.0, 1.0, 2.0, 3.0])) # The switch happens at t=2; later tendencies take precedence at the boundary. @@ -341,11 +341,48 @@ def test_string_waveform(): assert list(extrap) == ["ohmic", "nbi"] +def test_boolean_waveform(): + """A waveform of boolean constants evaluates as a zero-order-hold step function.""" + waveform = Waveform( + waveform=[ + {"user_type": "constant", "user_value": False, "user_duration": 2}, + {"user_type": "constant", "user_value": True, "user_duration": 2}, + ] + ) + assert waveform.is_categorical + + _, values = waveform.get_value(np.array([0.0, 1.0, 2.0, 3.0])) + assert [bool(v) for v in values] == [False, False, True, True] + + def test_numeric_waveform_evaluates_to_float(): """A numeric waveform evaluates to a float array (the int value is not stepwise).""" waveform = Waveform( waveform=[{"user_type": "constant", "user_value": 8, "user_duration": 3}] ) - assert not waveform.is_string + assert not waveform.is_categorical _, values = waveform.get_value(np.array([0.0, 1.0, 2.0, 3.0])) assert values.dtype == float + + +def test_mixing_categorical_and_numeric_is_flagged(): + """Mixing categorical (string/bool) and numeric tendencies adds an annotation.""" + waveform = Waveform( + waveform=[ + {"user_type": "constant", "user_value": "nbi", "user_duration": 2}, + {"user_type": "constant", "user_value": 3, "user_duration": 2}, + ] + ) + assert waveform.annotations + assert any("mix" in a["text"].lower() for a in waveform.annotations) + + +def test_homogeneous_categorical_not_flagged(): + """A waveform of only categorical values is not flagged as mixed.""" + waveform = Waveform( + waveform=[ + {"user_type": "constant", "user_value": "nbi", "user_duration": 2}, + {"user_type": "constant", "user_value": "ec", "user_duration": 2}, + ] + ) + assert not waveform.annotations diff --git a/waveform_editor/tendencies/base.py b/waveform_editor/tendencies/base.py index 7bb40535..f92246d3 100644 --- a/waveform_editor/tendencies/base.py +++ b/waveform_editor/tendencies/base.py @@ -77,7 +77,7 @@ class BaseTendency(param.Parameterized): ) annotations = param.ClassSelector(class_=Annotations, default=Annotations()) allow_zero_duration = False - is_string = False + is_categorical = False def __init__(self, **kwargs): super().__init__() diff --git a/waveform_editor/tendencies/constant.py b/waveform_editor/tendencies/constant.py index 62b1e4a7..01c8d7fd 100644 --- a/waveform_editor/tendencies/constant.py +++ b/waveform_editor/tendencies/constant.py @@ -20,9 +20,11 @@ def __init__(self, **kwargs): super().__init__(**kwargs) @property - def is_string(self): - """Whether this constant holds a string (non-numeric) value.""" - return isinstance(self.value, str) + def is_categorical(self): + """Whether this constant holds a non-numeric (categorical) value, e.g. a + string or boolean, that is held as a step rather than interpolated.""" + value = self.value + return isinstance(value, bool) or not isinstance(value, (int, float)) def get_value( self, time: np.ndarray | None = None diff --git a/waveform_editor/waveform.py b/waveform_editor/waveform.py index 9ecb98f3..c78c52ab 100644 --- a/waveform_editor/waveform.py +++ b/waveform_editor/waveform.py @@ -70,9 +70,10 @@ def _bind_imports(self): tendency.default_path = self.name @property - def is_string(self): - """Whether this waveform produces string (non-numeric) values.""" - return any(t.is_string for t in self.tendencies) + def is_categorical(self): + """Whether this waveform produces non-numeric (categorical) values, e.g. + strings or booleans, which are held as steps rather than interpolated.""" + return any(t.is_categorical for t in self.tendencies) def get_value( self, time: np.ndarray | None = None @@ -130,8 +131,8 @@ def _evaluate_tendencies(self, time, eval_derivatives=False): Returns: numpy array containing the computed values. """ - is_string = self.is_string and not eval_derivatives - if is_string: + is_categorical = self.is_categorical and not eval_derivatives + if is_categorical: values = np.empty(len(time), dtype=object) else: values = np.zeros_like(time, dtype=float) @@ -153,7 +154,7 @@ def _evaluate_tendencies(self, time, eval_derivatives=False): values[mask] = ( tendency.start_value - prev_tendency.end_value ) / (tendency.start - prev_tendency.end) - elif is_string: + elif is_categorical: values[mask] = prev_tendency.end_value else: values[mask] = np.interp( @@ -210,11 +211,29 @@ 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_types() + self.update_annotations() for tendency in self.tendencies: tendency.param.watch(self.update_annotations, "annotations") + def _validate_value_types(self): + """Categorical (e.g. string or boolean) and numeric tendencies cannot be + mixed within a single waveform, as the gaps between them cannot be + interpolated. Flag the first tendency whose type breaks the pattern.""" + if not self.tendencies: + return + first_is_categorical = self.tendencies[0].is_categorical + for tendency in self.tendencies[1:]: + if tendency.is_categorical != first_is_categorical: + error_msg = ( + "Cannot mix categorical (e.g. string or boolean) and numeric " + "values within a single waveform.\n" + ) + self.annotations.add(tendency.line_number, error_msg) + break + def update_annotations(self, event=None): """Merges the annotations of the individual tendencies into the annotations of this waveform.""" From efc0a3c52e9462189a02614b6371c6e243903b2e Mon Sep 17 00:00:00 2001 From: Daan van Vugt Date: Thu, 25 Jun 2026 19:02:35 +0200 Subject: [PATCH 3/6] docs: clarify integers are interpolated as floats, not categorical --- docs/source/tendencies.rst | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/docs/source/tendencies.rst b/docs/source/tendencies.rst index bb0a693a..ee55ef7f 100644 --- a/docs/source/tendencies.rst +++ b/docs/source/tendencies.rst @@ -58,8 +58,9 @@ Value types The ``value`` is type-aware and is not restricted to floating-point numbers: -* **Floats** are interpolated across gaps between tendencies, as usual. -* **Integers** are preserved as integers (they are not coerced to float). +* **Numeric values** (floats and integers) produce a floating-point signal that is + interpolated across gaps between tendencies, as usual. Integers are treated as + floats during evaluation. * **Categorical values** (strings or booleans) describe non-numeric signals, such as a heating scheme (``"nbi"``, ``"ec"``) or an on/off flag. These are held as a step (zero-order hold) rather than interpolated: across a gap, the previous value From fc4860013410ec435cc2e9dbcad4f5c864677253 Mon Sep 17 00:00:00 2001 From: Daan van Vugt Date: Fri, 10 Jul 2026 11:57:12 +0200 Subject: [PATCH 4/6] Evaluate numeric constant values as float on all eval paths np.full preserved int dtype for integer constants, so get_value() (the time=None path) returned an int64 array while the explicit-time path upcast to float. Force float for numeric values; categorical values (str/bool) keep their native dtype. --- waveform_editor/tendencies/constant.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/waveform_editor/tendencies/constant.py b/waveform_editor/tendencies/constant.py index 01c8d7fd..720d8283 100644 --- a/waveform_editor/tendencies/constant.py +++ b/waveform_editor/tendencies/constant.py @@ -40,7 +40,10 @@ def get_value( """ if time is None: time = np.array([self.start, self.end]) - values = np.full(len(time), self.value) + # Numeric values evaluate as floats (ints included); categorical values keep + # their native dtype (str stays str, bool stays bool). + dtype = None if self.is_categorical else float + values = np.full(len(time), self.value, dtype=dtype) return time, values def get_derivative(self, time: np.ndarray) -> np.ndarray: From 2ff479418623c1002e624335b54c6437cc97d7e9 Mon Sep 17 00:00:00 2001 From: Daan van Vugt Date: Thu, 25 Jun 2026 11:40:34 +0200 Subject: [PATCH 5/6] Add steps (zero-order-hold) tendency A piecewise-constant tendency that holds each value from its breakpoint until the next, supporting non-numeric (string/boolean) values. Registered as 'steps'. - start is inferred from the first breakpoint; an explicit end/duration may be given so the final value is held for a real duration (defaults to the last time) - shared time/value validation and start/end resolution are factored out of PiecewiseLinearTendency via overridable _coerce_value / _resolve_time_bounds hooks, so the steps tendency only overrides what genuinely differs --- docs/source/tendencies.rst | 26 ++++++ tests/tendencies/test_steps.py | 109 ++++++++++++++++++++++++ waveform_editor/tendencies/piecewise.py | 37 +++++++- waveform_editor/tendencies/steps.py | 58 +++++++++++++ waveform_editor/waveform.py | 2 + 5 files changed, 228 insertions(+), 4 deletions(-) create mode 100644 tests/tendencies/test_steps.py create mode 100644 waveform_editor/tendencies/steps.py diff --git a/docs/source/tendencies.rst b/docs/source/tendencies.rst index ee55ef7f..c95bb395 100644 --- a/docs/source/tendencies.rst +++ b/docs/source/tendencies.rst @@ -212,6 +212,32 @@ Parameters .. warning:: This tendency does **not** accept the common ``start``, ``duration``, or ``end`` parameters. These are derived directly from the required ``time`` list. +.. _steps-tendency: + +Steps Tendency +============== + +Defines a piecewise-constant (zero-order-hold) signal. Each ``time`` is the start of a step that holds the corresponding ``value`` until the next breakpoint, and the final value is held until the tendency's ``end``. Unlike the :ref:`Piecewise Linear Tendency `, the values are never interpolated, so they may be :ref:`categorical ` (strings or booleans) as well as numeric. + +Parameters +---------- +* ``time``: A list of step start times. Must be strictly monotonically increasing and must have at least 1 point. +* ``value``: A list of corresponding values, one per breakpoint in ``time``. Must have the same length as ``time``. Values may be numbers, strings, or booleans. +* ``duration``, ``end``: Optionally extend the tendency past the last breakpoint, so the final value is held for a real duration. If omitted, the tendency ends at the last ``time`` (the final value then has zero duration within the tendency). The ``start`` is always inferred from the first ``time`` and may not be set explicitly. + +.. code-block:: yaml + + - {type: steps, time: [0, 2, 4], value: [1, 3, 5], end: 6} + +A common use is a non-numeric signal, such as the active heating scheme: + +.. code-block:: yaml + + - {type: steps, time: [0, 10, 20], value: [ohmic, nbi, ec], end: 30} + +.. warning:: + Categorical and numeric values cannot be mixed within a single waveform. + .. _import-tendency: Import diff --git a/tests/tendencies/test_steps.py b/tests/tendencies/test_steps.py new file mode 100644 index 00000000..256041fc --- /dev/null +++ b/tests/tendencies/test_steps.py @@ -0,0 +1,109 @@ +import numpy as np + +from waveform_editor.tendencies.steps import StepsTendency +from waveform_editor.waveform import Waveform + + +def test_filled(): + """Time/value arrays keep their native dtype and are not interpolated.""" + tendency = StepsTendency(user_time=[0, 2, 4], user_value=[1, 3, 5]) + assert np.all(tendency.time == np.array([0, 2, 4])) + assert np.all(tendency.value == np.array([1, 3, 5])) + assert not tendency.is_categorical + assert not tendency.annotations + + +def test_zero_order_hold(): + """Each value is held from its breakpoint until the next; ends extrapolate.""" + tendency = StepsTendency(user_time=[0, 2, 4], user_value=[1, 3, 5]) + t = np.array([-1.0, 0.0, 1.0, 2.0, 3.0, 4.0, 5.0]) + _, values = tendency.get_value(t) + assert list(values) == [1, 1, 1, 3, 3, 5, 5] + + +def test_derivative_is_zero(): + tendency = StepsTendency(user_time=[0, 2, 4], user_value=[1, 3, 5]) + assert np.all(tendency.get_derivative(np.array([0.0, 1.0, 2.0, 3.0])) == 0) + + +def test_string_values(): + """A steps tendency can hold string values.""" + tendency = StepsTendency(user_time=[0, 2, 4], user_value=["ohmic", "nbi", "ec"]) + assert tendency.is_categorical + _, values = tendency.get_value(np.array([0.0, 1.0, 2.0, 3.0, 4.0])) + assert list(values) == ["ohmic", "ohmic", "nbi", "nbi", "ec"] + assert not tendency.annotations + + +def test_boolean_values(): + """A steps tendency can hold boolean values, which are categorical (held).""" + tendency = StepsTendency(user_time=[0, 2, 4], user_value=[False, True, False]) + assert tendency.is_categorical + _, values = tendency.get_value(np.array([0.0, 1.0, 2.0, 3.0, 4.0])) + assert [bool(v) for v in values] == [False, False, True, True, False] + assert not tendency.annotations + + +def test_non_monotonic_time(): + tendency = StepsTendency(user_time=[0, 2, 1], user_value=[1, 2, 3]) + assert tendency.annotations + + +def test_mismatched_lengths(): + tendency = StepsTendency(user_time=[0, 2], user_value=[1, 2, 3]) + assert tendency.annotations + + +def test_start_and_end_inference(): + """Start comes from the first breakpoint; end defaults to the last breakpoint.""" + tendency = StepsTendency(user_time=[0, 2, 4], user_value=[1, 3, 5]) + assert tendency.start == 0 + assert tendency.end == 4 + + +def test_explicit_end_holds_final_value(): + """An explicit end holds the final value for its full duration.""" + tendency = StepsTendency(user_time=[0, 2, 4], user_value=[1, 3, 5], user_end=6) + assert tendency.start == 0 + assert tendency.end == 6 + assert not tendency.annotations + _, values = tendency.get_value(np.array([4.0, 5.0, 6.0])) + assert list(values) == [5, 5, 5] + + +def test_end_before_last_time_is_flagged(): + """An end that precedes the last breakpoint is an error.""" + tendency = StepsTendency(user_time=[0, 2, 4], user_value=[1, 3, 5], user_end=3) + assert tendency.annotations + + +def test_start_not_allowed(): + """Providing an explicit start is rejected; it is inferred from the time list.""" + tendency = StepsTendency(user_time=[0, 2, 4], user_value=[1, 3, 5], user_start=1) + assert tendency.annotations + assert tendency.start == 0 + + +def test_in_waveform(): + """A steps tendency drives a zero-order-hold waveform (numeric and string).""" + numeric = Waveform( + waveform=[ + {"user_type": "steps", "user_time": [0, 2, 4], "user_value": [1, 3, 5]} + ] + ) + assert not numeric.is_categorical + _, values = numeric.get_value(np.array([0.0, 1.0, 2.0, 3.0, 4.0])) + assert list(values) == [1, 1, 3, 3, 5] + + string = Waveform( + waveform=[ + { + "user_type": "steps", + "user_time": [0, 2, 4], + "user_value": ["a", "b", "c"], + } + ] + ) + assert string.is_categorical + _, values = string.get_value(np.array([1.0, 3.0, 5.0])) + assert list(values) == ["a", "b", "c"] diff --git a/waveform_editor/tendencies/piecewise.py b/waveform_editor/tendencies/piecewise.py index e557bc18..3222fa9a 100644 --- a/waveform_editor/tendencies/piecewise.py +++ b/waveform_editor/tendencies/piecewise.py @@ -21,12 +21,11 @@ class PiecewiseLinearTendency(BaseTendency): def __init__(self, user_time=None, user_value=None, **kwargs): self.pre_check_annotations = Annotations() time, value = self._validate_time_value(user_time, user_value) - self._remove_user_time_params(kwargs) + time_bounds = self._resolve_time_bounds(time, kwargs) super().__init__( - user_start=time[0], - user_end=time[-1], time=time, value=value, + **time_bounds, **kwargs, ) self.annotations.add_annotations(self.pre_check_annotations) @@ -34,6 +33,22 @@ def __init__(self, user_time=None, user_value=None, **kwargs): self.start_value_set = True self.param.update(values_changed=True) + def _resolve_time_bounds(self, time, kwargs): + """Determine the tendency's start/end from the time array. The piecewise-linear + tendency derives its full interval from the time list, so the common time + parameters (``start``, ``duration``, ``end``) are not accepted. + + Args: + time: The validated time array. + kwargs: The remaining keyword arguments. Disallowed time parameters are + removed in place. + + Returns: + A dict of the ``user_start``/``user_end`` keyword arguments to pass on. + """ + self._remove_user_time_params(kwargs) + return {"user_start": time[0], "user_end": time[-1]} + def get_value( self, time: np.ndarray | None = None ) -> tuple[np.ndarray, np.ndarray]: @@ -105,7 +120,7 @@ def _validate_time_value(self, time, value): try: time = np.asarray_chkfinite(time, dtype=float) - value = np.asarray_chkfinite(value, dtype=float) + value = self._coerce_value(value) is_monotonic = np.all(np.diff(time) > 0) if not is_monotonic: error_msg = "The provided time array is not monotonically increasing.\n" @@ -119,6 +134,20 @@ def _validate_time_value(self, time, value): else: return self.time, self.value + def _coerce_value(self, value): + """Coerce the value array to the dtype expected by this tendency. The + piecewise-linear tendency interpolates its values, so they must be finite + floats. Subclasses may override this (e.g. a step function keeps the native + dtype, allowing non-numeric values). + + Args: + value: The values defined on each time step. + + Returns: + The coerced value array. + """ + return np.asarray_chkfinite(value, dtype=float) + def _remove_user_time_params(self, kwargs): """Remove user_start, user_duration, and user_end if they are passed as kwargs, and add error messages as annotations. These variables will be set from the diff --git a/waveform_editor/tendencies/steps.py b/waveform_editor/tendencies/steps.py new file mode 100644 index 00000000..81f5e999 --- /dev/null +++ b/waveform_editor/tendencies/steps.py @@ -0,0 +1,58 @@ +import numpy as np + +from waveform_editor.tendencies.piecewise import PiecewiseLinearTendency + + +class StepsTendency(PiecewiseLinearTendency): + """A piecewise-constant (zero-order-hold) tendency. + + Each ``time`` is the start of a step holding the corresponding ``value`` until the + next breakpoint; the final value is held until the tendency's ``end``. Unlike the + piecewise-linear tendency, the values are never interpolated, so they may be + non-numeric (e.g. strings or booleans). + + The start is inferred from the first breakpoint. An explicit ``end`` (or + ``duration``) may be given so the final value is held for a real duration; if + neither is given, the tendency ends at the last breakpoint. + """ + + @property + def is_categorical(self): + # Strings ("U"/"S"), objects ("O") and booleans ("b") are non-numeric and are + # held as a step rather than interpolated. + return self.value.dtype.kind in ("U", "S", "O", "b") + + def _coerce_value(self, value): + # A step function is never interpolated, so keep the native dtype (which may be + # non-numeric) instead of coercing to float like the piecewise-linear tendency. + return np.asarray(value) + + def _resolve_time_bounds(self, time, kwargs): + # The start is fixed at the first breakpoint, but (unlike piecewise-linear) an + # explicit end/duration is allowed so the final value can be held for a real + # duration. Without one, the tendency ends at the last breakpoint. + line_number = kwargs.get("line_number", 0) + if "user_start" in kwargs: + kwargs.pop("user_start") + self.pre_check_annotations.add( + line_number, "'start' is not allowed in a steps tendency\n" + ) + end = kwargs.get("user_end") + if end is None and kwargs.get("user_duration") is None: + kwargs["user_end"] = time[-1] + elif isinstance(end, (int, float)) and end < time[-1]: + self.pre_check_annotations.add( + line_number, + "The `end` of a steps tendency must not precede the last time.\n", + ) + return {"user_start": time[0]} + + def get_value(self, time: np.ndarray | None = None): + if time is None: + return self.time, self.value + indices = np.searchsorted(self.time, time, side="right") - 1 + indices = np.clip(indices, 0, len(self.value) - 1) + return time, self.value[indices] + + def get_derivative(self, time: np.ndarray) -> np.ndarray: + return np.zeros_like(time, dtype=float) diff --git a/waveform_editor/waveform.py b/waveform_editor/waveform.py index c78c52ab..7e7f92c9 100644 --- a/waveform_editor/waveform.py +++ b/waveform_editor/waveform.py @@ -15,6 +15,7 @@ from waveform_editor.tendencies.piecewise import PiecewiseLinearTendency from waveform_editor.tendencies.repeat import RepeatTendency from waveform_editor.tendencies.smooth import SmoothTendency +from waveform_editor.tendencies.steps import StepsTendency tendency_map = { "linear": LinearTendency, @@ -29,6 +30,7 @@ "constant": ConstantTendency, "smooth": SmoothTendency, "piecewise": PiecewiseLinearTendency, + "steps": StepsTendency, "repeat": RepeatTendency, "import": ImportTendency, # `reference` kept as an alias for the import tendency type. From 166b751c9eb40fa342ed9b2f7664e79f42504353 Mon Sep 17 00:00:00 2001 From: Daan van Vugt Date: Thu, 25 Jun 2026 13:56:17 +0200 Subject: [PATCH 6/6] turn expressions into tendencies too Replace DerivedWaveform with an ExpressionTendency on the single Waveform class, integrated with develop's import/static waveforms. Bare scalars are now classified by content: a number is a constant, a string that references another waveform or uses an operator/function is an expression, and a plain word is a literal string constant. Explicit {value:}/{expression:} override the heuristic. Example config uses the compact bare form. --- tests/test_configuration.py | 2 + tests/test_derived_waveform.py | 153 --------------- tests/test_expression.py | 109 +++++++++++ tests/test_yaml/example.yaml | 1 + tests/test_yaml_parser.py | 42 ++++- waveform_editor/base_waveform.py | 17 ++ waveform_editor/configuration.py | 14 +- waveform_editor/derived_waveform.py | 177 ------------------ waveform_editor/gui/editor.py | 5 +- waveform_editor/gui/plotter_edit.py | 5 +- .../gui/shape_editor/coil_currents.py | 3 +- waveform_editor/tendencies/expression.py | 135 +++++++++++++ waveform_editor/waveform.py | 41 +++- waveform_editor/yaml/yaml_globals.py | 9 + waveform_editor/yaml/yaml_parser.py | 70 ++++++- 15 files changed, 417 insertions(+), 366 deletions(-) delete mode 100644 tests/test_derived_waveform.py create mode 100644 tests/test_expression.py delete mode 100644 waveform_editor/derived_waveform.py create mode 100644 waveform_editor/tendencies/expression.py diff --git a/tests/test_configuration.py b/tests/test_configuration.py index 2bf51be9..4105d5e7 100644 --- a/tests/test_configuration.py +++ b/tests/test_configuration.py @@ -255,6 +255,7 @@ def test_dump_comments(): yaml_str = dedent(""" globals: + version: 2 dd_version: 3.42.0 imports: ec_launchers: imas:hdf5?path=test_md @@ -284,6 +285,7 @@ def test_dump_globals(): dumped_yaml = config.dump() expected_dump = dedent(""" globals: + version: 2 dd_version: 3.41.0 imports: ec_launchers: imas:mdsplus?path=test diff --git a/tests/test_derived_waveform.py b/tests/test_derived_waveform.py deleted file mode 100644 index cc910e09..00000000 --- a/tests/test_derived_waveform.py +++ /dev/null @@ -1,153 +0,0 @@ -import numpy as np -import pytest - -from waveform_editor.configuration import WaveformConfiguration -from waveform_editor.derived_waveform import DerivedWaveform -from waveform_editor.waveform import Waveform - - -@pytest.fixture -def config(): - config = WaveformConfiguration() - config.add_group("root_group", []) - return config - - -@pytest.fixture -def filled_config(config): - waveform_list = [ - { - "user_type": "linear", - "user_from": 10, - "user_to": 20, - "user_start": 5, - "user_end": 15, - "line_number": 1, - } - ] - - waveform = Waveform(waveform=waveform_list, name="waveform/1") - config.add_waveform(waveform, ["root_group"]) - return config - - -@pytest.fixture -def const_waveform(config): - name = "waveform/1" - const_value = 3 - yaml_str = f"{name}: {const_value}" - waveform = DerivedWaveform(yaml_str, name, config) - config.add_waveform(waveform, ["root_group"]) - return waveform, const_value, name, config - - -def test_const_waveform(const_waveform): - waveform, const_value, name, _ = const_waveform - - assert waveform.name == name - assert waveform.yaml == const_value - assert waveform.dependencies == set() - assert waveform.get_yaml_string() == str(const_value) - time = np.linspace(0, 1, 1000) - time_ret, value_ret = waveform.get_value() - assert np.all(time_ret == time) - assert np.all(value_ret == const_value) - time = np.linspace(0, 100, num=101) - time_ret, value_ret = waveform.get_value(time) - assert np.all(time == time_ret) - assert np.all(value_ret == const_value) - - -def test_bounds(const_waveform): - derived_waveform, const_value, _, config = const_waveform - start = 5 - end = 15 - config.start = start - config.end = end - - time_ret, value_ret = derived_waveform.get_value() - assert time_ret[0] == start - assert time_ret[-1] == end - assert np.all(value_ret == const_value) - - _, value_ret = derived_waveform.get_value(np.array([0, 5, 10, 15])) - assert np.all(value_ret == const_value) - - -def test_dependent_waveform(filled_config): - name = "waveform/2" - yaml_str = f"{name}: |\n 'waveform/1'" - waveform = DerivedWaveform(yaml_str, name, filled_config) - assert waveform.dependencies == {"waveform/1"} - time_ret, value_ret = waveform.get_value() - assert time_ret[0] == 5 - assert time_ret[-1] == 15 - assert value_ret[0] == 10 - assert value_ret[-1] == 20 - _, value_ret = waveform.get_value(np.array([0, 5, 10, 15, 20])) - assert np.all(value_ret == [10, 10, 15, 20, 20]) - - -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) - assert waveform.dependencies == {"waveform/1"} - time_ret, value_ret = waveform.get_value() - assert time_ret[0] == 5 - assert time_ret[-1] == 15 - assert value_ret[0] == 100 - assert value_ret[-1] == 200 - _, value_ret = waveform.get_value(np.array([0, 5, 10, 15, 20])) - assert np.all(value_ret == [100, 100, 150, 200, 200]) - - -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) - assert waveform.dependencies == {"waveform/1"} - time_ret, value_ret = waveform.get_value() - assert time_ret[0] == 5 - assert time_ret[-1] == 15 - assert value_ret[0] == 150 - assert value_ret[-1] == 200 - _, value_ret = waveform.get_value(np.array([0, 5, 10, 15, 20])) - assert np.all(value_ret == [150, 150, 150, 200, 200]) - - -def test_rename_waveform(filled_config): - name = "waveform/2" - yaml_str = f"{name}: |\n 'waveform/1'" - waveform = DerivedWaveform(yaml_str, name, filled_config) - assert waveform.dependencies == {"waveform/1"} - assert waveform.get_yaml_string() == "'waveform/1'" - waveform.rename_dependency("waveform/1", "waveform/3") - assert waveform.dependencies == {"waveform/3"} - assert waveform.get_yaml_string() == "'waveform/3'" - - -def test_function_access_control(filled_config): - test_exprs = [ - ('max("waveform/1")', False), - ('sum("waveform/1")', False), - ('eval("waveform/1")', False), - ('dot("waveform/1", "waveform/1")', False), - ('linalg.norm("waveform/1")', False), - ('linalg.inv("waveform/1")', False), - ('sin("waveform/1")', True), - ('log("waveform/1" + 1)', True), - ('maximum("waveform/1", 10)', True), - ] - - for expr, allowed in test_exprs: - name = "waveform/2" - yaml_str = f"{name}: |\n {expr}" - waveform = DerivedWaveform(yaml_str, name, filled_config) - time_ret = np.linspace(filled_config.start, filled_config.end, 100) - if allowed: - _, result = waveform.get_value(time_ret) - assert result is not None - else: - with pytest.raises(NameError): - waveform.get_value(time_ret) diff --git a/tests/test_expression.py b/tests/test_expression.py new file mode 100644 index 00000000..183f3553 --- /dev/null +++ b/tests/test_expression.py @@ -0,0 +1,109 @@ +import numpy as np +import pytest + +from waveform_editor.configuration import WaveformConfiguration +from waveform_editor.waveform import Waveform + + +@pytest.fixture +def config(): + config = WaveformConfiguration() + config.add_group("root_group", []) + return config + + +@pytest.fixture +def filled_config(config): + waveform_list = [ + { + "user_type": "linear", + "user_from": 10, + "user_to": 20, + "user_start": 5, + "user_end": 15, + "line_number": 1, + } + ] + waveform = Waveform(waveform=waveform_list, name="waveform/1") + config.add_waveform(waveform, ["root_group"]) + return config + + +def make_expression(config, name, expr, add=True): + """Create an expression waveform bound to ``config``.""" + waveform = Waveform( + waveform=[{"user_expression": expr, "line_number": 0}], + name=name, + config=config, + ) + if add: + config.add_waveform(waveform, ["root_group"]) + return waveform + + +def test_constant_expression(config): + waveform = make_expression(config, "waveform/1", "3") + assert waveform.is_expression + assert not waveform.is_categorical + assert waveform.dependencies == set() + _, value = waveform.get_value(np.linspace(0, 100, 101)) + assert np.all(value == 3) + + +def test_dependent_waveform(filled_config): + waveform = make_expression(filled_config, "waveform/2", '"waveform/1"') + assert waveform.dependencies == {"waveform/1"} + time, value = waveform.get_value() + assert time[0] == 5 + assert time[-1] == 15 + assert value[0] == 10 + assert value[-1] == 20 + _, value = waveform.get_value(np.array([0, 5, 10, 15, 20])) + assert np.all(value == [10, 10, 15, 20, 20]) + + +def test_dependent_waveform_calc(filled_config): + waveform = make_expression(filled_config, "waveform/2", '"waveform/1" * 10') + assert waveform.dependencies == {"waveform/1"} + _, value = waveform.get_value(np.array([0, 5, 10, 15, 20])) + assert np.all(value == [100, 100, 150, 200, 200]) + + +def test_dependent_waveform_numpy(filled_config): + waveform = make_expression( + filled_config, "waveform/2", 'maximum("waveform/1" * 10, 150)' + ) + assert waveform.dependencies == {"waveform/1"} + _, value = waveform.get_value(np.array([0, 5, 10, 15, 20])) + assert np.all(value == [150, 150, 150, 200, 200]) + + +def test_rename_dependency(filled_config): + waveform = make_expression(filled_config, "waveform/2", '"waveform/1"', add=False) + assert waveform.dependencies == {"waveform/1"} + waveform.rename_dependency("waveform/1", "waveform/3") + assert waveform.dependencies == {"waveform/3"} + + +def test_function_access_control(filled_config): + test_exprs = [ + ('max("waveform/1")', False), + ('sum("waveform/1")', False), + ('eval("waveform/1")', False), + ('dot("waveform/1", "waveform/1")', False), + ('linalg.norm("waveform/1")', False), + ('linalg.inv("waveform/1")', False), + ('sin("waveform/1")', True), + ('log("waveform/1" + 1)', True), + ('maximum("waveform/1", 10)', True), + ] + + time = np.linspace(filled_config.start, filled_config.end, 100) + for expr, allowed in test_exprs: + waveform = make_expression(filled_config, "waveform/2", expr, add=False) + if allowed: + _, result = waveform.get_value(time) + assert result is not None + else: + with pytest.raises(NameError): + waveform.get_value(time) diff --git a/tests/test_yaml/example.yaml b/tests/test_yaml/example.yaml index 4980a1ab..92d5a3f0 100644 --- a/tests/test_yaml/example.yaml +++ b/tests/test_yaml/example.yaml @@ -1,4 +1,5 @@ globals: + version: 2 dd_version: 4.0.0 imports: {} dummy_waveform: diff --git a/tests/test_yaml_parser.py b/tests/test_yaml_parser.py index c682e8af..969c92c2 100644 --- a/tests/test_yaml_parser.py +++ b/tests/test_yaml_parser.py @@ -2,7 +2,7 @@ from pytest import approx from waveform_editor.configuration import WaveformConfiguration -from waveform_editor.derived_waveform import DerivedWaveform +from waveform_editor.static_waveform import StaticWaveform from waveform_editor.tendencies.constant import ConstantTendency from waveform_editor.tendencies.linear import LinearTendency from waveform_editor.tendencies.periodic.sawtooth_wave import SawtoothWaveTendency @@ -132,19 +132,47 @@ def test_scientific_notation(yaml_parser): def test_constant_shorthand_notation(yaml_parser): - """Test if shorthand notation is parsed correctly.""" + """A bare scalar is shorthand for a single constant tendency.""" waveforms = {"waveform: 5": 5, "waveform: 1.23": 1.23} - for waveform, expected_value in waveforms.items(): - waveform = yaml_parser.parse_waveform(waveform) - assert isinstance(waveform, DerivedWaveform) - assert waveform.yaml == expected_value - assert not waveform.annotations + for waveform_str, expected_value in waveforms.items(): + waveform = yaml_parser.parse_waveform(waveform_str) + assert isinstance(waveform, Waveform) + assert not waveform.is_expression assert waveform.dependencies == set() + assert waveform.tendencies[0].value == expected_value + assert not waveform.annotations assert not yaml_parser.parse_errors +def test_bare_word_is_literal_string(yaml_parser): + """A bare plain word is a literal string constant, not an expression.""" + waveform = yaml_parser.parse_waveform("waveform: total") + assert isinstance(waveform, StaticWaveform) + assert not waveform.is_expression + assert waveform.value == "total" + assert not yaml_parser.parse_errors + + +def test_bare_number_is_constant(yaml_parser): + """A bare number is a constant tendency, not an expression.""" + waveform = yaml_parser.parse_waveform("waveform: 3.5") + assert isinstance(waveform, Waveform) + assert not waveform.is_expression + assert waveform.tendencies[0].value == 3.5 + assert not yaml_parser.parse_errors + + +def test_bare_reference_is_expression(yaml_parser): + """A bare string that references another waveform (quoted) is an expression.""" + waveform = yaml_parser.parse_waveform('waveform: \'"other/1" * 2\'') + assert isinstance(waveform, Waveform) + assert waveform.is_expression + assert waveform.dependencies == {"other/1"} + assert not yaml_parser.parse_errors + + def test_load_yaml(config): """Test if yaml is loaded correctly.""" yaml_str = """ diff --git a/waveform_editor/base_waveform.py b/waveform_editor/base_waveform.py index 6d88084b..0cdc196c 100644 --- a/waveform_editor/base_waveform.py +++ b/waveform_editor/base_waveform.py @@ -28,6 +28,23 @@ def get_value( def get_yaml_string(self) -> str: raise NotImplementedError + @property + def dependencies(self): + """Names of other waveforms this waveform depends on. Empty unless the waveform + contains expressions (overridden by :class:`Waveform`).""" + return set() + + @property + def is_expression(self): + """Whether this waveform is computed from an expression. False by default.""" + return False + + def prepare_expression(self): # noqa: B027 + """Re-parse expression tendencies. No-op for waveforms without expressions.""" + + def rename_dependency(self, old_name, new_name): # noqa: B027 + """Rename a referenced waveform. No-op for waveforms without expressions.""" + def get_metadata(self, dd_version): """Parses the name of the waveform and returns the IDS metadata for this waveform. The name must be formatted as follows: ``/`` diff --git a/waveform_editor/configuration.py b/waveform_editor/configuration.py index 276f8c5c..d4094c27 100644 --- a/waveform_editor/configuration.py +++ b/waveform_editor/configuration.py @@ -6,7 +6,6 @@ from ruamel.yaml.comments import CommentedMap from waveform_editor.dependency_graph import DependencyGraph -from waveform_editor.derived_waveform import DerivedWaveform from waveform_editor.group import WaveformGroup from waveform_editor.import_resolver import ImportResolver from waveform_editor.yaml.yaml_globals import YamlGlobals @@ -94,8 +93,7 @@ def load_yaml(self, yaml_str): self._calculate_bounds() for name, group in self.waveform_map.items(): waveform = group[name] - if isinstance(waveform, DerivedWaveform): - waveform.prepare_expression() + waveform.prepare_expression() self.has_changed = False except Exception as e: self.clear() @@ -119,7 +117,7 @@ def add_waveform(self, waveform, path): f"The group {group.name!r} already contains {waveform.name!r}." ) - if isinstance(waveform, DerivedWaveform): + if waveform.dependencies: self.dependency_graph.add_node(waveform.name, waveform.dependencies) group.waveforms[waveform.name] = waveform self.waveform_map[waveform.name] = group @@ -193,7 +191,7 @@ def replace_waveform(self, waveform): f"Waveform '{waveform.name}' does not exist in the configuration." ) - if isinstance(waveform, DerivedWaveform): + if waveform.dependencies: self.dependency_graph.replace_node(waveform.name, waveform.dependencies) elif waveform.name in self.dependency_graph: self.dependency_graph.remove_node(waveform.name) @@ -235,9 +233,7 @@ def remove_group(self, path): for wf_name, grp in self.waveform_map.items(): if wf_name not in to_remove: wf = grp[wf_name] - if isinstance(wf, DerivedWaveform) and to_remove.intersection( - wf.dependencies - ): + if to_remove.intersection(wf.dependencies): raise RuntimeError( f"Cannot remove group {group.name}. " f"{wf.name!r} depends on a waveform in it." @@ -354,7 +350,7 @@ def _calculate_bounds(self): for name in self.waveform_map: waveform = self[name] - if not isinstance(waveform, DerivedWaveform) and waveform.tendencies: + if not waveform.dependencies and waveform.tendencies: min_start = min(min_start, waveform.tendencies[0].start) max_end = max(max_end, waveform.tendencies[-1].end) diff --git a/waveform_editor/derived_waveform.py b/waveform_editor/derived_waveform.py deleted file mode 100644 index 79d623da..00000000 --- a/waveform_editor/derived_waveform.py +++ /dev/null @@ -1,177 +0,0 @@ -import ast - -import numpy as np -from asteval import Interpreter - -from waveform_editor.base_waveform import BaseWaveform - -NUMPY_UFUNCS = {} -for name in np.__all__: - obj = getattr(np, name) - if isinstance(obj, np.ufunc): - NUMPY_UFUNCS[name] = obj - - -class DependencyRenamer(ast.NodeTransformer): - """AST transformer to rename string constants.""" - - def __init__(self, rename_from, rename_to, yaml): - self.rename_from = rename_from - self.rename_to = rename_to - self.yaml = yaml - - def visit_Constant(self, node): - """ - Replace string constants equal to `rename_from` with `rename_to` - and update the YAML source lines accordingly. - """ - if isinstance(node.value, str) and node.value == self.rename_from: - split_yaml = self.yaml.splitlines() - line_number = node.lineno - 1 - line = split_yaml[line_number] - split_yaml[line_number] = ( - line[: node.col_offset] - + line[node.col_offset : node.end_col_offset].replace( - self.rename_from, self.rename_to - ) - + line[node.end_col_offset :] - ) - self.yaml = "\n".join(split_yaml) - return ast.copy_location(ast.Constant(value=self.rename_to), node) - return node - - -class ExpressionExtractor(ast.NodeTransformer): - """ - AST transformer extracting all string constants from expressions - and replacing them with Name nodes for later evaluation. - """ - - def __init__(self): - self.string_nodes = [] - - def visit_Constant(self, node): - if isinstance(node.value, str): - self.string_nodes.append(node.value) - return ast.copy_location( - ast.Subscript( - value=ast.Name(id="__w", ctx=ast.Load()), - slice=ast.Constant(value=node.value), - ctx=ast.Load(), - ), - node, - ) - else: - return node - - -class DerivedWaveform(BaseWaveform): - def __init__(self, yaml_str, name, config, dd_version=None): - super().__init__(yaml_str, name, dd_version) - self.config = config - 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. - """ - if self.yaml is None: - return - - try: - tree = ast.parse(str(self.yaml), mode="eval") - except Exception as e: - self.annotations.add(0, f"Could not parse or evaluate the waveform: {e}") - self.expression = None - return - - extractor = ExpressionExtractor() - modified_tree = ast.fix_missing_locations(extractor.visit(tree)) - self.is_constant = not extractor.string_nodes - self.expression = ast.unparse(modified_tree) - self.dependencies = set(extractor.string_nodes) - - def rename_dependency(self, old_name, new_name): - """Rename a dependency waveform in the expression. - - Args: - old_name: Original dependency name. - new_name: New dependency name. - """ - if old_name not in self.dependencies: - return - - tree = ast.parse(self.yaml, mode="eval") - renamer = DependencyRenamer(old_name, new_name, self.yaml) - ast.fix_missing_locations(renamer.visit(tree)) - self.yaml = renamer.yaml - self.prepare_expression() - - def _build_eval_context(self, time: np.ndarray) -> dict: - """Build the evaluation context dictionary with dependencies resolved. - - Args: - time: The time array on which to generate points. - - Returns: - dict: Mapping dependency names to waveform values at given times. - """ - eval_context = {} - - for name in self.dependencies: - eval_context[name] = self.config[name].get_value(time)[1] - return eval_context - - def get_value( - self, time: np.ndarray | None = None - ) -> tuple[np.ndarray, np.ndarray]: - """Evaluate the derived waveform expression at specified times. - - Args: - time: Array of time points. Defaults to 1000 points between config.start - and config.end. - - Returns: - Tuple containing the time and the derived waveform values. - """ - if time is None: - # TODO: properly handle time for plotting - time = np.linspace(self.config.start, self.config.end, 1000) - if self.expression is None: - return time, np.zeros_like(time) - - eval_context = self._build_eval_context(time) - sym_table = NUMPY_UFUNCS.copy() - sym_table["__w"] = eval_context - aeval = Interpreter( - symtable=sym_table, - minimal=True, - use_numpy=False, - ) - - # Don't print the entire NumPy array in the error alert message - with np.printoptions(threshold=10): - result = aeval.eval(self.expression, raise_errors=True) - - # If derived waveform is a constant, ensure an array is returned - if self.is_constant: - return time, np.full_like(time, result, dtype=float) - - # Ensure the result is a 1D array - if not isinstance(result, np.ndarray): - raise ValueError("The derived waveform is not a 1D array.") - result = np.asarray(result) - if result.shape != time.shape: - raise ValueError( - f"The shape of the derived waveform {result.shape} does not match the " - f"shape of the time array {time.shape}" - ) - - return time, result - - def get_yaml_string(self): - """Returns the current YAML expression string.""" - return str(self.yaml) diff --git a/waveform_editor/gui/editor.py b/waveform_editor/gui/editor.py index 97ab4624..627ba434 100644 --- a/waveform_editor/gui/editor.py +++ b/waveform_editor/gui/editor.py @@ -4,7 +4,6 @@ import param from panel.viewable import Viewer -from waveform_editor.derived_waveform import DerivedWaveform from waveform_editor.import_waveform import ImportWaveform from waveform_editor.static_waveform import StaticWaveform from waveform_editor.waveform import Waveform @@ -14,7 +13,7 @@ class WaveformEditor(Viewer): """A Panel interface for waveform editing.""" waveform = param.ClassSelector( - class_=(Waveform, DerivedWaveform, ImportWaveform, StaticWaveform), + class_=(Waveform, ImportWaveform, StaticWaveform), doc="Waveform currently being edited. Use `set_waveform` to change.", ) stored_string = param.String( @@ -113,7 +112,7 @@ def handle_exceptions(self, waveform): ) self.alert_type = "warning" else: - if isinstance(waveform, DerivedWaveform): + if waveform.dependencies: try: self.config.check_safe_to_replace(waveform) except Exception as e: diff --git a/waveform_editor/gui/plotter_edit.py b/waveform_editor/gui/plotter_edit.py index 77f9ada9..83b14afe 100644 --- a/waveform_editor/gui/plotter_edit.py +++ b/waveform_editor/gui/plotter_edit.py @@ -8,7 +8,6 @@ from panel.viewable import Viewer from ruamel.yaml import YAML -from waveform_editor.derived_waveform import DerivedWaveform from waveform_editor.import_waveform import ImportWaveform from waveform_editor.static_waveform import StaticWaveform from waveform_editor.tendencies.piecewise import PiecewiseLinearTendency @@ -20,7 +19,7 @@ class PlotterEdit(Viewer): """Class to plot a single waveform in edit mode.""" plotted_waveform: Waveform = param.ClassSelector( - class_=(Waveform, DerivedWaveform, ImportWaveform, StaticWaveform), + class_=(Waveform, ImportWaveform, StaticWaveform), allow_refs=True, ) @@ -45,7 +44,7 @@ def update_plot(self): if self._update_plot_from_drag: return # Skip update triggered from a drag-and-drop - if isinstance(self.plotted_waveform, DerivedWaveform): + if self.plotted_waveform is not None and self.plotted_waveform.is_expression: try: self.pane.object = self.main_curve() except Exception as e: diff --git a/waveform_editor/gui/shape_editor/coil_currents.py b/waveform_editor/gui/shape_editor/coil_currents.py index bbe6c47f..01efc5bd 100644 --- a/waveform_editor/gui/shape_editor/coil_currents.py +++ b/waveform_editor/gui/shape_editor/coil_currents.py @@ -7,7 +7,6 @@ from bokeh.models.widgets.tables import NumberFormatter from panel.viewable import Viewer -from waveform_editor.derived_waveform import DerivedWaveform from waveform_editor.settings import settings from waveform_editor.tendencies.piecewise import PiecewiseLinearTendency @@ -175,7 +174,7 @@ def _store_coil_currents(self, group_name="Coil Currents"): new_waveforms_created = True else: waveform = config[name] - if isinstance(waveform, DerivedWaveform): + if waveform.is_expression: pn.state.notifications.error( f"Could not store coil current in waveform {name!r}, " "because it is a derived waveform" diff --git a/waveform_editor/tendencies/expression.py b/waveform_editor/tendencies/expression.py new file mode 100644 index 00000000..a308ffcb --- /dev/null +++ b/waveform_editor/tendencies/expression.py @@ -0,0 +1,135 @@ +import ast + +import numpy as np +from asteval import Interpreter + +from waveform_editor.tendencies.base import BaseTendency + +NUMPY_UFUNCS = {} +for _name in np.__all__: + _obj = getattr(np, _name) + if isinstance(_obj, np.ufunc): + NUMPY_UFUNCS[_name] = _obj + + +class DependencyRenamer(ast.NodeTransformer): + """AST transformer to rename a string constant (a waveform reference).""" + + def __init__(self, rename_from, rename_to, source): + self.rename_from = rename_from + self.rename_to = rename_to + self.source = source + + def visit_Constant(self, node): + if isinstance(node.value, str) and node.value == self.rename_from: + lines = self.source.splitlines() + i = node.lineno - 1 + line = lines[i] + lines[i] = ( + line[: node.col_offset] + + line[node.col_offset : node.end_col_offset].replace( + self.rename_from, self.rename_to + ) + + line[node.end_col_offset :] + ) + self.source = "\n".join(lines) + return ast.copy_location(ast.Constant(value=self.rename_to), node) + return node + + +class ExpressionExtractor(ast.NodeTransformer): + """Replace string constants (waveform references) with ``__w[...]`` lookups.""" + + def __init__(self): + self.string_nodes = [] + + def visit_Constant(self, node): + if isinstance(node.value, str): + self.string_nodes.append(node.value) + return ast.copy_location( + ast.Subscript( + value=ast.Name(id="__w", ctx=ast.Load()), + slice=ast.Constant(value=node.value), + ctx=ast.Load(), + ), + node, + ) + return node + + +class ExpressionTendency(BaseTendency): + """A tendency whose values are computed from an expression over other waveforms. + + References to other waveforms are written as quoted strings (e.g. ``"a" * 10``). + Without any references the expression is a constant. The owning waveform resolves + dependencies through the configuration passed in ``config``. + """ + + def __init__(self, user_expression=None, config=None, **kwargs): + self.config = config + self.source = user_expression + self.dependencies = set() + self.is_constant = False + self.expression = None + super().__init__(**kwargs) + self.prepare_expression() + + def prepare_expression(self): + """Parse the expression, extract dependencies and compile it for evaluation.""" + if self.source is None: + return + try: + tree = ast.parse(str(self.source), mode="eval") + except Exception as e: + self.annotations.add(self.line_number, f"Could not parse expression: {e}") + self.expression = None + return + extractor = ExpressionExtractor() + modified = ast.fix_missing_locations(extractor.visit(tree)) + self.is_constant = not extractor.string_nodes + self.expression = ast.unparse(modified) + self.dependencies = set(extractor.string_nodes) + + def rename_dependency(self, old_name, new_name): + if old_name not in self.dependencies: + return + tree = ast.parse(str(self.source), mode="eval") + renamer = DependencyRenamer(old_name, new_name, str(self.source)) + ast.fix_missing_locations(renamer.visit(tree)) + self.source = renamer.source + self.prepare_expression() + + def _calc_start_end_values(self): + # Boundary values depend on the configuration and dependencies, which are not + # resolvable at construction time, so they are not computed eagerly. + pass + + def get_value(self, time: np.ndarray | None = None): + if time is None: + time = np.linspace(self.config.start, self.config.end, 1000) + if self.expression is None: + return time, np.zeros_like(time) + + eval_context = { + name: self.config[name].get_value(time)[1] for name in self.dependencies + } + sym_table = NUMPY_UFUNCS.copy() + sym_table["__w"] = eval_context + aeval = Interpreter(symtable=sym_table, minimal=True, use_numpy=False) + + with np.printoptions(threshold=10): + result = aeval.eval(self.expression, raise_errors=True) + + if self.is_constant: + return time, np.full_like(time, result, dtype=float) + + result = np.asarray(result) + if result.shape != time.shape: + raise ValueError( + f"The shape of the derived waveform {result.shape} does not match the " + f"shape of the time array {time.shape}" + ) + return time, result + + def get_derivative(self, time: np.ndarray) -> np.ndarray: + return np.zeros_like(time, dtype=float) diff --git a/waveform_editor/waveform.py b/waveform_editor/waveform.py index 7e7f92c9..e2d629ad 100644 --- a/waveform_editor/waveform.py +++ b/waveform_editor/waveform.py @@ -6,6 +6,7 @@ from waveform_editor.base_waveform import BaseWaveform from waveform_editor.tendencies.constant import ConstantTendency +from waveform_editor.tendencies.expression import ExpressionTendency from waveform_editor.tendencies.import_tendency import ImportTendency from waveform_editor.tendencies.linear import LinearTendency from waveform_editor.tendencies.periodic.sawtooth_wave import SawtoothWaveTendency @@ -53,8 +54,9 @@ def __init__( super().__init__(yaml_str, name, dd_version) self.line_number = line_number self.is_repeated = is_repeated - # Used to reach the import resolver for {ref: ...} tendencies (None when the - # waveform is built outside a configuration, e.g. a repeated sub-waveform). + # Used to reach the import resolver for {ref: ...} tendencies and to resolve + # expression dependencies (None when the waveform is built outside a + # configuration, e.g. a repeated sub-waveform). self.config = config if waveform is not None: self._process_waveform(waveform) @@ -77,6 +79,31 @@ def is_categorical(self): strings or booleans, which are held as steps rather than interpolated.""" return any(t.is_categorical for t in self.tendencies) + @property + def is_expression(self): + """Whether this waveform is computed from an expression.""" + return any(isinstance(t, ExpressionTendency) for t in self.tendencies) + + @property + def dependencies(self): + """Names of other waveforms this waveform references through expressions.""" + deps = set() + for tendency in self.tendencies: + deps |= getattr(tendency, "dependencies", set()) + return deps + + def prepare_expression(self): + """Re-parse any expression tendencies (e.g. after dependencies change).""" + for tendency in self.tendencies: + if isinstance(tendency, ExpressionTendency): + tendency.prepare_expression() + + def rename_dependency(self, old_name, new_name): + """Rename a referenced waveform in any expression tendencies.""" + for tendency in self.tendencies: + if isinstance(tendency, ExpressionTendency): + tendency.rename_dependency(old_name, new_name) + def get_value( self, time: np.ndarray | None = None ) -> tuple[np.ndarray, np.ndarray]: @@ -93,6 +120,12 @@ def get_value( if not self.tendencies: return np.array([]), np.array([]) + # A sole expression tendency spans the whole domain and is evaluated directly. + if len(self.tendencies) == 1 and isinstance( + self.tendencies[0], ExpressionTendency + ): + return self.tendencies[0].get_value(time) + self._bind_imports() if time is None: @@ -326,6 +359,10 @@ def _handle_tendency(self, entry): Returns: The created tendency or None, if the tendency cannot be created """ + if "user_expression" in entry: + entry.pop("user_type", None) + return ExpressionTendency(config=self.config, **entry) + if self._has_type_error(entry): return None else: diff --git a/waveform_editor/yaml/yaml_globals.py b/waveform_editor/yaml/yaml_globals.py index e5c4e390..e97049b8 100644 --- a/waveform_editor/yaml/yaml_globals.py +++ b/waveform_editor/yaml/yaml_globals.py @@ -6,8 +6,17 @@ logger = logging.getLogger(__name__) +# Configuration schema version. Bumped to 2 when bare-string values became literal +# string constants instead of derived-waveform expressions. +CURRENT_SCHEMA_VERSION = 2 + class YamlGlobals(param.Parameterized): + version = param.Integer( + default=CURRENT_SCHEMA_VERSION, + label="Schema Version", + doc="Waveform Editor configuration schema version.", + ) dd_version = param.Selector( label="DD Version", default=LATEST_DD_VERSION, diff --git a/waveform_editor/yaml/yaml_parser.py b/waveform_editor/yaml/yaml_parser.py index e4f5f752..1b083c30 100644 --- a/waveform_editor/yaml/yaml_parser.py +++ b/waveform_editor/yaml/yaml_parser.py @@ -1,3 +1,4 @@ +import ast import logging import re from io import StringIO @@ -7,10 +8,10 @@ from imas.ids_path import IDSPath from ruamel.yaml import YAML -from waveform_editor.derived_waveform import DerivedWaveform from waveform_editor.import_waveform import ImportWaveform from waveform_editor.static_waveform import StaticWaveform from waveform_editor.waveform import Waveform +from waveform_editor.yaml.yaml_globals import CURRENT_SCHEMA_VERSION logger = logging.getLogger(__name__) @@ -22,6 +23,29 @@ def _is_import_entry(entry): ) +def _looks_like_expression(value): + """Whether a bare string value should be read as an expression rather than a + literal string constant. + + Waveform references are quoted strings, so an expression is anything *dynamic* (it + contains a quoted reference) or *functional* (it uses a call or an operator). A + plain word (``nbi``) or an unparseable string is a literal constant. The explicit + ``{value: ...}`` / ``{expression: ...}`` forms override this heuristic. + """ + try: + tree = ast.parse(value, mode="eval") + except SyntaxError: + return False + for node in ast.walk(tree): + if isinstance(node, ast.Constant) and isinstance(node.value, str): + return True # a quoted waveform reference -> dynamic + if isinstance( + node, (ast.BinOp, ast.UnaryOp, ast.BoolOp, ast.Compare, ast.Call) + ): + return True # uses an operator or function -> functional + return False + + def _import_is_non_scalar(name, entry, dd_version): """Whether an import must be an ImportWaveform rather than a 0D segment. @@ -90,6 +114,17 @@ def load_yaml(self, yaml_str): yaml_data = self.yaml.load(yaml_str) if yaml_str else {} globals = yaml_data.get("globals", {}) + file_version = globals.get("version") + if file_version is None or file_version < CURRENT_SCHEMA_VERSION: + logger.warning( + "Configuration schema version (%s) is older than the current version " + "(%s). A bare string is now an expression only if it references " + "another waveform or uses an operator/function; a plain word is a " + "literal constant. Use `{value: ...}` or `{expression: ...}` to be " + "explicit.", + file_version, + CURRENT_SCHEMA_VERSION, + ) self.config.globals.set_globals(globals) if not isinstance(yaml_data, dict): @@ -183,14 +218,14 @@ def parse_waveform(self, yaml_str): raise yaml.YAMLError("Cannot have an empty waveform.") if not isinstance(waveform, (list, int, float, str)): raise yaml.YAMLError( - "Waveform must either be a list of tendencies, " - "a single constant value (int/float), or a derived waveform (str)." + "Waveform must either be a list of tendencies or a bare constant " + "value (number or string)." ) line_number = waveform_yaml.get("line_number", 0) dd_version = self.config.globals.dd_version if isinstance(waveform, list): - # A single {value: } entry is a static constant (e.g. an - # identifier name) -- strings don't fit the numeric tendency value flow. + # A single {value: } entry is a static string constant (e.g. + # an identifier name). if ( len(waveform) == 1 and isinstance(waveform[0], dict) @@ -216,7 +251,7 @@ def parse_waveform(self, yaml_str): name=name, dd_version=dd_version, ) - waveform = Waveform( + return Waveform( waveform=waveform, yaml_str=yaml_str, line_number=line_number, @@ -224,11 +259,26 @@ def parse_waveform(self, yaml_str): dd_version=dd_version, config=self.config, ) + # A bare scalar is shorthand. A number is a constant; a string that + # references other waveforms or uses operators/functions is an expression, + # otherwise it is a static (literal) string constant. + if isinstance(waveform, str): + if _looks_like_expression(waveform): + entry = {"user_expression": waveform, "line_number": line_number} + else: + return StaticWaveform( + waveform, yaml_str=yaml_str, name=name, dd_version=dd_version + ) else: - waveform = DerivedWaveform( - yaml_str, name, self.config, dd_version=dd_version - ) - return waveform + entry = {"user_value": waveform, "line_number": line_number} + return Waveform( + waveform=[entry], + yaml_str=yaml_str, + line_number=line_number, + name=name, + dd_version=dd_version, + config=self.config, + ) except yaml.YAMLError as e: self.parse_errors.append(str(e)) empty_waveform = Waveform()