From 8a371ffc1caa11b7726b2ce8ab7ec569738d5a3d Mon Sep 17 00:00:00 2001 From: Johannes Kasimir Date: Wed, 5 Aug 2026 11:03:13 +0200 Subject: [PATCH] refactor: WIP --- .../essreduce/src/ess/reduce/parameter.py | 181 +++-------- .../src/ess/reduce/parameter_models.py | 283 ++++++++++++++++++ packages/essreduce/src/ess/reduce/workflow.py | 162 ++++++++-- 3 files changed, 453 insertions(+), 173 deletions(-) create mode 100644 packages/essreduce/src/ess/reduce/parameter_models.py diff --git a/packages/essreduce/src/ess/reduce/parameter.py b/packages/essreduce/src/ess/reduce/parameter.py index 8137bce57..560b4a18a 100644 --- a/packages/essreduce/src/ess/reduce/parameter.py +++ b/packages/essreduce/src/ess/reduce/parameter.py @@ -2,12 +2,11 @@ # Copyright (c) 2023 Scipp contributors (https://github.com/scipp) from __future__ import annotations -from collections.abc import MutableMapping -from dataclasses import dataclass -from enum import Enum -from typing import Generic, Self, TypeVar +from collections.abc import Callable, MutableMapping +from dataclasses import dataclass, replace +from typing import Any, Generic, TypeVar -import scipp as sc +from sciline import Pipeline from sciline._utils import key_name from sciline.typing import Key @@ -21,158 +20,51 @@ class KeepDefaultType: keep_default = KeepDefaultType() -@dataclass -class Parameter(Generic[T]): - """Interface between workflow parameters and widgets. +@dataclass(frozen=True) +class ParameterSpec(Generic[T]): + """Specification for configuring a Sciline workflow key.""" - .. versionadded:: RELEASE_PLACEHOLDER - """ + model: Any + category: str + title: str | None = None + description: str | None = None + default: T | KeepDefaultType = keep_default + transform: Callable[[T], Any] | None = None + apply: Callable[[Pipeline, T], Pipeline] | None = None + use_workflow_default: bool = True + key: Key | None = None - name: str - description: str - default: T - optional: bool = False - """If True, widget has radio buttons switch between "None" and param widget.""" - switchable: bool = False - """If True, widget has checkbox to enable/disable parameter.""" + def bind(self, key: Key) -> ParameterSpec[T]: + return replace(self, key=key) - def with_default(self: Self, default: T | KeepDefaultType = keep_default) -> Self: - if default == keep_default: + def with_default( + self, default: T | KeepDefaultType = keep_default + ) -> ParameterSpec[T]: + if default is keep_default or not self.use_workflow_default: return self - # TODO I think some subclasses currently cannot be instantiated with this method - return type(self)( - self.name, self.description, default, self.optional, self.switchable - ) - - @classmethod - def from_type( - cls: type[Self], - t: type[T], - default: T | None = None, - optional: bool = False, - switchable: bool = False, - ) -> Self: - # TODO __doc__ not correct when using NewType - # TODO __doc__ not correct when using Generic - # use sciline type->string helper - return cls( - name=key_name(t), - description=t.__doc__, - default=default, - optional=optional, - switchable=switchable, - ) + return replace(self, default=default) + def with_apply( + self, apply: Callable[[Pipeline, T], Pipeline] | None + ) -> ParameterSpec[T]: + return replace(self, apply=apply) -@dataclass(kw_only=True) -class ParamWithOptions(Parameter[T]): - options: Enum - - @classmethod - def from_enum(cls: type[Self], t: type[T], default: T) -> Self: - return cls( - name=t.__name__, - description=t.__doc__, - options=t.__members__, - default=default, + @property + def name(self) -> str: + return self.title or ( + key_name(self.key) if self.key is not None else 'Parameter' ) -@dataclass -class FilenameParameter(Parameter[str]): - """Widget for entering a filename or selecting one in a file dialog.""" - - # TODO need specifics for different file types, nexus, ... - - -@dataclass -class MultiFilenameParameter(Parameter[tuple[str, ...]]): - """Widget for entering multiple filenames or selecting multiple in a file dialog.""" - - -@dataclass(kw_only=True) -class BinEdgesParameter(Parameter[sc.Variable]): - """Widget for entering bin edges.""" - - dim: str - start: float | None = None - stop: float | None = None - nbins: int = 1 - unit: str | None = "undefined" # If "undefined", the unit is deduced from the dim - log: bool = False - - def __init__( - self, - t: type[T], - dim: str, - start: float | None = None, - stop: float | None = None, - nbins: int = 1, - unit: str | None = "undefined", - log: bool = False, - ): - self.dim = dim - self.start = start - self.stop = stop - self.nbins = nbins - self.unit = unit - self.log = log - super().__init__(name=key_name(t), description=t.__doc__, default=None) - - -@dataclass -class BooleanParameter(Parameter[bool]): - pass - - -@dataclass -class StringParameter(Parameter[str]): - pass - - -@dataclass -class MultiStringParameter(Parameter[tuple[str, ...]]): - """Widget for entering multiple strings.""" - - -@dataclass(kw_only=True) -class ParamWithBounds(Parameter[T]): - bounds: tuple[T, T] - - -@dataclass(kw_only=True) -class ScalarParameter(Parameter[T]): - """Fixed unit displayed in widget""" - - unit: str - - -@dataclass(kw_only=True) -class ScalarParamWithUnitOptions(Parameter[T]): - """User can select between compatible units""" - - unit_options: list[str] - - -@dataclass(kw_only=True) -class Vector2dParameter(Parameter[sc.Variable]): - """Widget for entering a 2d vector.""" - - -@dataclass(kw_only=True) -class Vector3dParameter(Parameter[sc.Variable]): - """Widget for entering a 3d vector.""" - - class ParameterRegistry(MutableMapping): def __init__(self): - self._parameters = {} + self._parameters: dict[Key, ParameterSpec] = {} - def __getitem__(self, key: Key) -> Parameter: + def __getitem__(self, key: Key) -> ParameterSpec: return self._parameters[key] - def __setitem__(self, key: Key, value: Parameter): - self._parameters[key] = value + def __setitem__(self, key: Key, value: ParameterSpec): + self._parameters[key] = value.bind(key) def __delitem__(self, key: Key): del self._parameters[key] @@ -185,6 +77,3 @@ def __len__(self) -> int: parameter_registry = ParameterRegistry() - - -parameter_mappers = {} diff --git a/packages/essreduce/src/ess/reduce/parameter_models.py b/packages/essreduce/src/ess/reduce/parameter_models.py new file mode 100644 index 000000000..2b398c7ca --- /dev/null +++ b/packages/essreduce/src/ess/reduce/parameter_models.py @@ -0,0 +1,283 @@ +# SPDX-License-Identifier: BSD-3-Clause +# Copyright (c) 2025 Scipp contributors (https://github.com/scipp) +""" +Models for data reduction workflow inputs. + +These models define inputs to Sciline workflows. They can be used for creating +user interfaces for configuring workflows as well as validation of the inputs +before values are applied to a workflow. +""" + +from __future__ import annotations + +import json +from abc import ABC +from enum import StrEnum +from pathlib import Path + +import scipp as sc +from pydantic import BaseModel, Field, field_validator, model_validator + + +def parse_number_list(value: str) -> list[float]: + """Parse a comma-separated string of numbers into floats.""" + value = value.strip() + if not value: + return [] + try: + parsed = json.loads(f"[{value}]") + except json.JSONDecodeError as e: + raise ValueError(f"Invalid number list: {e}") from e + if any(isinstance(x, bool) or not isinstance(x, int | float) for x in parsed): + raise ValueError("All entries must be numbers") + return [float(x) for x in parsed] + + +class RangeModel(BaseModel, ABC): + """Base model for range with common fields and validation.""" + + start: float = Field(default=0.0, description="Start of the range.") + stop: float = Field(default=10.0, description="Stop of the range.") + + @field_validator('stop') + @classmethod + def stop_must_be_greater_than_start(cls, v, info): + start = info.data.get('start') + if start is not None and v <= start: + raise ValueError('stop must be greater than start') + return v + + def get_start(self) -> sc.Variable: + """Get the start of the range as a scipp variable.""" + return sc.scalar(self.start, unit=self.unit.value) + + def get_stop(self) -> sc.Variable: + """Get the stop of the range as a scipp variable.""" + return sc.scalar(self.stop, unit=self.unit.value) + + +class Scale(StrEnum): + """Allowed scales for data reduction.""" + + LINEAR = 'linear' + LOG = 'log' + + +class EdgesModel(BaseModel, ABC): + """Base model for edges with common fields and validation.""" + + start: float = Field(default=1.0, description="Start of the edges.") + stop: float = Field(default=10.0, description="Stop of the edges.") + num_bins: int = Field(default=100, ge=1, le=10000, description="Number of bins.") + scale: Scale = Field( + default=Scale.LINEAR, + description="Scale of the edges, either 'linear' or 'log'.", + ) + + @field_validator('stop') + @classmethod + def stop_must_be_greater_than_start(cls, v, info): + start = info.data.get('start') + if start is not None and v <= start: + raise ValueError('stop must be greater than start') + return v + + @model_validator(mode='after') + def start_must_be_positive_if_log(self): + if self.scale == Scale.LOG and self.start <= 0: + raise ValueError("start must be positive when scale is 'log'") + return self + + +class TimeUnit(StrEnum): + """Allowed units for time.""" + + NS = 'ns' + US = 'us' + MICROSECOND = 'μs' + MS = 'ms' + S = 's' + + +class WavelengthUnit(StrEnum): + """Allowed units for wavelength.""" + + ANGSTROM = 'Å' + NANOMETER = 'nm' + + +class DspacingUnit(StrEnum): + """Allowed units for d-spacing.""" + + ANGSTROM = 'Å' + NANOMETER = 'nm' + + +class LengthUnit(StrEnum): + """Allowed units for length.""" + + METER = 'm' + CENTIMETER = 'cm' + MILLIMETER = 'mm' + + +class AngleUnit(StrEnum): + """Allowed units for angles.""" + + DEGREE = 'deg' + RADIAN = 'rad' + + +class QUnit(StrEnum): + """Allowed units for Q.""" + + INVERSE_ANGSTROM = '1/Å' + INVERSE_NANOMETER = '1/nm' + + +class Angle(BaseModel): + """Model for an angle value.""" + + value: float = Field(default=0.0, description="Angle value.") + unit: AngleUnit = Field(default=AngleUnit.DEGREE, description="Unit of the angle.") + + def get_value(self) -> sc.Variable: + """Get the angle as a scipp scalar.""" + return sc.scalar(self.value, unit=self.unit.value) + + +class Filename(BaseModel): + """Model for a filename.""" + + value: Path = Field(..., description="Path to the file.") + + +class WavelengthRange(RangeModel): + """Model for wavelength range.""" + + unit: WavelengthUnit = Field( + default=WavelengthUnit.ANGSTROM, description="Unit of the wavelength range." + ) + + +class TOARange(RangeModel): + """Time of arrival range filter settings.""" + + enabled: bool = Field(default=False, description="Enable the range filter.") + unit: TimeUnit = Field( + default=TimeUnit.MICROSECOND, description="Unit of the interval bounds." + ) + + @property + def range(self) -> tuple[sc.Variable, sc.Variable]: + """Time of arrival range as a tuple of scipp scalars.""" + return (self.get_start(), self.get_stop()) + + +class TOAEdges(EdgesModel): + """Model for time of arrival edges.""" + + unit: TimeUnit = Field(default=TimeUnit.MS, description="Unit of the edges.") + + def get_edges(self) -> sc.Variable: + """Get the edges as a scipp variable.""" + return make_edges(model=self, dim='time_of_arrival', unit=self.unit.value) + + +class WavelengthEdges(EdgesModel): + """Model for wavelength edges.""" + + unit: WavelengthUnit = Field( + default=WavelengthUnit.ANGSTROM, description="Unit of the edges." + ) + + def get_edges(self) -> sc.Variable: + """Get the edges as a scipp variable.""" + return make_edges(model=self, dim='wavelength', unit=self.unit.value) + + +class WavelengthRangeFilter(RangeModel): + """Wavelength range filter settings for detector view.""" + + enabled: bool = Field(default=False, description="Enable the range filter.") + unit: WavelengthUnit = Field( + default=WavelengthUnit.ANGSTROM, description="Unit of the range bounds." + ) + + @property + def range(self) -> tuple[sc.Variable, sc.Variable]: + """Wavelength range as a tuple of scipp scalars.""" + return (self.get_start(), self.get_stop()) + + +class DspacingEdges(EdgesModel): + """Model for d-spacing edges.""" + + unit: DspacingUnit = Field( + default=DspacingUnit.ANGSTROM, description="Unit of the edges." + ) + + def get_edges(self) -> sc.Variable: + """Get the edges as a scipp variable.""" + return make_edges(model=self, dim='dspacing', unit=self.unit.value) + + +class TwoTheta(EdgesModel): + """Model for two-theta edges.""" + + unit: AngleUnit = Field(default=AngleUnit.DEGREE, description="Unit of the edges.") + + def get_edges(self) -> sc.Variable: + """Get the edges as a scipp variable.""" + return make_edges(model=self, dim='two_theta', unit=self.unit.value) + + +class ThetaEdges(EdgesModel): + """Model for theta bin edges.""" + + unit: AngleUnit = Field( + default=AngleUnit.DEGREE, description="Unit of the theta bin edges." + ) + + def get_edges(self) -> sc.Variable: + """Get the edges as a scipp variable.""" + return make_edges(model=self, dim='theta', unit=self.unit.value) + + +class QEdges(EdgesModel): + """Model for Q edges.""" + + unit: QUnit = Field( + default=QUnit.INVERSE_ANGSTROM, description="Unit of the Q edges." + ) + + def get_edges(self) -> sc.Variable: + """Get the edges as a scipp variable.""" + return make_edges(model=self, dim='Q', unit=self.unit.value) + + +class EnergyUnit(StrEnum): + """Allowed units for energy transfer.""" + + MILLI_EV = 'meV' + MICRO_EV = 'μeV' + + +class EnergyEdges(EdgesModel): + """Model for energy transfer edges.""" + + unit: EnergyUnit = Field( + default=EnergyUnit.MILLI_EV, description="Unit of the energy transfer edges." + ) + + def get_edges(self) -> sc.Variable: + """Get the edges as a scipp variable.""" + return make_edges(model=self, dim='ΔE', unit=self.unit.value) + + +def make_edges(*, model: EdgesModel, dim: str, unit: str) -> sc.Variable: + """Convert the edges to a scipp variable.""" + op = {Scale.LINEAR: sc.linspace, Scale.LOG: sc.geomspace}[model.scale] + return op( + dim=dim, start=model.start, stop=model.stop, num=model.num_bins + 1, unit=unit + ) diff --git a/packages/essreduce/src/ess/reduce/workflow.py b/packages/essreduce/src/ess/reduce/workflow.py index 641246161..b35632e66 100644 --- a/packages/essreduce/src/ess/reduce/workflow.py +++ b/packages/essreduce/src/ess/reduce/workflow.py @@ -2,25 +2,83 @@ # Copyright (c) 2023 Scipp contributors (https://github.com/scipp) from __future__ import annotations -from collections.abc import Callable, MutableSet, Sequence +from collections.abc import Callable, Mapping, MutableSet, Sequence +from dataclasses import dataclass, field from typing import Any, TypeVar -import networkx as nx from sciline import Pipeline from sciline._utils import key_name +from sciline.handler import HandleAsComputeTimeException from sciline.typing import Key -from .parameter import Parameter, keep_default, parameter_mappers, parameter_registry +from .parameter import ( + ParameterRegistry, + ParameterSpec, + keep_default, + parameter_registry, +) T = TypeVar("T") +WorkflowFactory = Callable[..., Pipeline] + + +@dataclass(frozen=True) +class WorkflowSpec: + """Workflow factory and metadata needed to generate user interfaces.""" + + factory: WorkflowFactory + parameters: Mapping[Key, ParameterSpec] = field(default_factory=ParameterRegistry) + typical_outputs: tuple[Key, ...] | None = None + name: str | None = None + title: str | None = None + description: str | None = None + version: str | None = None + + def __post_init__(self) -> None: + if self.name is None: + object.__setattr__(self, 'name', self.factory.__name__) + + def __call__(self, *args: Any, **kwargs: Any) -> Pipeline: + return self.factory(*args, **kwargs) + + @property + def __name__(self) -> str: + return self.factory.__name__ + + def create_workflow(self) -> Pipeline: + return self.factory() + + @classmethod + def from_factory( + cls, + factory: WorkflowFactory, + *, + parameters: Mapping[Key, ParameterSpec] | None = None, + typical_outputs: Sequence[Key] | None = None, + name: str | None = None, + title: str | None = None, + description: str | None = None, + version: str | None = None, + ) -> WorkflowSpec: + return cls( + factory=factory, + parameters=parameters if parameters is not None else ParameterRegistry(), + typical_outputs=None if typical_outputs is None else tuple(typical_outputs), + name=name, + title=title, + description=description, + version=version, + ) class WorkflowRegistry(MutableSet): def __init__(self): - self._workflows: dict[str, type] = {} + self._workflows: dict[str, WorkflowSpec] = {} def __contains__(self, item: object) -> bool: - return item in self._workflows.values() + return item in self._workflows.values() or any( + item is spec.factory for spec in self._workflows.values() + ) def __iter__(self): return iter(self._workflows.values()) @@ -28,20 +86,60 @@ def __iter__(self): def __len__(self) -> int: return len(self._workflows) - def add(self, value: type) -> None: + def add(self, value: WorkflowFactory | WorkflowSpec) -> None: + if isinstance(value, WorkflowSpec): + spec = value + else: + spec = WorkflowSpec.from_factory(value) + key = spec.factory.__qualname__ + self._workflows[key] = spec + + def discard(self, value: WorkflowFactory | WorkflowSpec) -> None: + self._workflows = { + k: v + for k, v in self._workflows.items() + if v != value and v.factory is not value + } + + def get(self, value: WorkflowFactory | WorkflowSpec) -> WorkflowSpec: + if isinstance(value, WorkflowSpec): + return value key = value.__qualname__ - self._workflows[key] = value - - def discard(self, value: type) -> None: - self._workflows = {k: v for k, v in self._workflows.items() if v != value} + try: + return self._workflows[key] + except KeyError as e: + raise KeyError(f"Workflow {value.__qualname__!r} is not registered.") from e workflow_registry = WorkflowRegistry() -def register_workflow(cls: Callable[[], Pipeline]) -> Callable[[], Pipeline]: - workflow_registry.add(cls) - return cls +def register_workflow( + *, + parameters: Mapping[Key, ParameterSpec] | None = None, + typical_outputs: Sequence[Key] | None = None, + name: str | None = None, + title: str | None = None, + description: str | None = None, + version: str | None = None, +) -> Callable[[WorkflowFactory], WorkflowFactory]: + """Register a workflow factory for generated user interfaces.""" + + def decorator(factory: WorkflowFactory) -> WorkflowFactory: + workflow_registry.add( + WorkflowSpec.from_factory( + factory, + parameters=parameters, + typical_outputs=typical_outputs, + name=name, + title=title, + description=description, + version=version, + ) + ) + return factory + + return decorator def _get_defaults_from_workflow(workflow: Pipeline) -> dict[Key, Any]: @@ -49,8 +147,10 @@ def _get_defaults_from_workflow(workflow: Pipeline) -> dict[Key, Any]: return {key: values["value"] for key, values in nodes.items() if "value" in values} -def get_typical_outputs(pipeline: Pipeline) -> tuple[Key, ...]: - if (typical_outputs := getattr(pipeline, "typical_outputs", None)) is None: +def get_typical_outputs( + pipeline: Pipeline, typical_outputs: Sequence[Key] | None = None +) -> tuple[Key, ...]: + if typical_outputs is None: graph = pipeline.underlying_graph sink_nodes = [node for node, degree in graph.out_degree if degree == 0] return sorted(_with_pretty_names(sink_nodes), key=lambda x: x[0]) @@ -69,27 +169,35 @@ def _with_pretty_names(outputs: Sequence[Key]) -> tuple[tuple[str, Key], ...]: def get_parameters( - pipeline: Pipeline, outputs: tuple[Key, ...] -) -> dict[Key, Parameter]: + pipeline: Pipeline, + outputs: tuple[Key, ...], + parameters: Mapping[Key, ParameterSpec] = parameter_registry, +) -> dict[Key, ParameterSpec]: """Return a dictionary of parameters for the workflow.""" - subgraph = set(outputs) - graph = pipeline.underlying_graph - for key in outputs: - subgraph.update(nx.ancestors(graph, key)) + required_keys = set( + pipeline.get(outputs, handler=HandleAsComputeTimeException()).keys() + ) defaults = _get_defaults_from_workflow(pipeline) return { - key: param.with_default(defaults.get(key, keep_default)) - for key, param in parameter_registry.items() - if key in subgraph + key: spec.with_default(defaults.get(key, keep_default)) + for key, spec in parameters.items() + if key in required_keys } -def assign_parameter_values(pipeline: Pipeline, values: dict[Key, Any]) -> Pipeline: +def assign_parameter_values( + pipeline: Pipeline, + values: dict[Key, Any], + parameters: Mapping[Key, ParameterSpec] = parameter_registry, +) -> Pipeline: """Set a value for a parameter in the pipeline.""" pipeline = pipeline.copy() for key, value in values.items(): - if (mapper := parameter_mappers.get(key)) is not None: - pipeline = mapper(pipeline, value) + spec = parameters[key] + if spec.transform is not None: + value = spec.transform(value) + if spec.apply is not None: + pipeline = spec.apply(pipeline, value) else: pipeline[key] = value return pipeline