diff --git a/CHANGELOG.md b/CHANGELOG.md index dedc1c17f9..174f45f118 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -22,6 +22,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Optional/secondary fields of discrete parameter classes are now keyword-only ### Added +- `SequenceParameter` class for modeling parameters whose values are configurable-length + token sequences from a predefined alphabet +- `EncoderProtocol` as a public interface for specifying parameter encoders - `coefficients` attribute for `DiscreteSumConstraint`, enabling weighted sums. Follows the same pattern as `ContinuousLinearConstraint.coefficients` - `simplex_coefficients` keyword argument to `SubspaceDiscrete.from_simplex` for @@ -30,7 +33,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - `CandidatesProtocol` as an interface for candidates generation - `EmptyCandidates`, `TableCandidates` and `ProductCandidates` classes implementing `CandidatesProtocol` -- `DiscreteParameter.is_finite` property +- `DiscreteParameter.__len__` and `DiscreteParameter.is_finite` property - `SubspaceDiscrete.batch_constraints` field for storing batch-level constraints - `SubspaceDiscrete.from_dataframe` now accepts `batch_constraints` - `validate_parameter_input` now accepts an `allow_empty` flag to permit zero-row input diff --git a/baybe/exceptions.py b/baybe/exceptions.py index c049d414b8..ca7df414e6 100644 --- a/baybe/exceptions.py +++ b/baybe/exceptions.py @@ -180,7 +180,7 @@ class UnsupportedEarlyFilteringError(Exception): class InfiniteSpaceError(Exception): - """An operation requires a finite search space but the space is infinite.""" + """An operation requires a finite space but the space is infinite.""" # Collect leftover original slotted classes processed by `attrs.define` diff --git a/baybe/parameters/__init__.py b/baybe/parameters/__init__.py index 93e62b6ee9..2c1eda16c0 100644 --- a/baybe/parameters/__init__.py +++ b/baybe/parameters/__init__.py @@ -2,6 +2,7 @@ from baybe.parameters.categorical import CategoricalParameter, TaskParameter from baybe.parameters.custom import CustomDiscreteParameter +from baybe.parameters.encoding import EncoderProtocol from baybe.parameters.enum import ( CategoricalEncoding, CustomEncoding, @@ -11,6 +12,7 @@ NumericalContinuousParameter, NumericalDiscreteParameter, ) +from baybe.parameters.sequence import SequenceParameter from baybe.parameters.substance import SubstanceParameter from baybe.utils.metadata import MeasurableMetadata @@ -19,9 +21,11 @@ "CategoricalParameter", "CustomDiscreteParameter", "CustomEncoding", + "EncoderProtocol", "MeasurableMetadata", "NumericalContinuousParameter", "NumericalDiscreteParameter", + "SequenceParameter", "SubstanceEncoding", "SubstanceParameter", "TaskParameter", diff --git a/baybe/parameters/base.py b/baybe/parameters/base.py index 8625b6b41e..8c45faaf19 100644 --- a/baybe/parameters/base.py +++ b/baybe/parameters/base.py @@ -16,6 +16,7 @@ from narwhals.stable.v2.dependencies import is_into_series from typing_extensions import override +from baybe.exceptions import InfiniteSpaceError from baybe.serialization import ( SerialMixin, ) @@ -131,11 +132,18 @@ class DiscreteParameter(Parameter, ABC): def values(self) -> tuple: """The values the parameter can take.""" + def __len__(self) -> int: + """Return the number of values the parameter can take.""" + return len(self.values) + @property def is_finite(self) -> bool: """Indicates whether the parameter has a finite number of values.""" - len(self.values) # <-- raises an error if the parameter is infinite - return True + try: + len(self) + return True + except InfiniteSpaceError: + return False @property def active_values(self) -> tuple: @@ -242,7 +250,7 @@ def summary(self) -> dict: return dict( Name=self.name, Type=self.__class__.__name__, - nValues=len(self.values), + nValues=len(self), ) @@ -261,7 +269,7 @@ class _EncodedDiscreteParameter(DiscreteParameter, ABC): _active_values: tuple[str | bool, ...] | None = field( default=None, converter=optional_c( - Converter( # type: ignore[misc, call-overload] + Converter( # type: ignore[misc] nonstring_to_tuple, takes_self=True, takes_field=True ) ), @@ -301,10 +309,10 @@ def _validate_active_values( # noqa: DOC101, DOC103 ) if len(set(content)) != len(content): raise ValueError("The active parameter values must be unique.") - if not all(v in self.values for v in content): + if invalid := [v for v in content if not self.is_in_range(v)]: raise ValueError( - f"All active values must be valid parameter choices from: " - f"{self.values}, provided: {content}" + f"All active values must be valid parameter choices. " + f"Provided invalid values: {invalid}" ) @override diff --git a/baybe/parameters/categorical.py b/baybe/parameters/categorical.py index 76117d8bb5..f40d3ceba0 100644 --- a/baybe/parameters/categorical.py +++ b/baybe/parameters/categorical.py @@ -49,7 +49,7 @@ class CategoricalParameter(_EncodedDiscreteParameter): encoding: CategoricalEncoding = field( default=CategoricalEncoding.OHE, converter=CategoricalEncoding, kw_only=True ) - # See base class. + """The encoding used the generate the parameters computational representation.""" @override @property @@ -108,7 +108,7 @@ class TaskParameter(CategoricalParameter): """Parameter class for task parameters.""" encoding: CategoricalEncoding = field(default=CategoricalEncoding.INT, init=False) - # See base class. + """The encoding used the generate the parameters computational representation.""" # Collect leftover original slotted classes processed by `attrs.define` diff --git a/baybe/parameters/encoding.py b/baybe/parameters/encoding.py new file mode 100644 index 0000000000..a858def26a --- /dev/null +++ b/baybe/parameters/encoding.py @@ -0,0 +1,133 @@ +"""Sequence encoders.""" + +from __future__ import annotations + +import gc +from typing import TYPE_CHECKING, Protocol, runtime_checkable + +import narwhals.stable.v2 as nw +from attrs import define, field +from attrs.validators import instance_of +from exceptiongroup import ExceptionGroup + +from baybe.settings import active_settings +from baybe.utils.dataframe import _df_with_backend + +if TYPE_CHECKING: + from narwhals.stable.v2.typing import IntoDataFrame, IntoSeries + + +@runtime_checkable +class EncoderProtocol(Protocol): + """Type protocol specifying the interface encoders need to implement.""" + + # Use slots so that derived classes also remain slotted + # See also: https://www.attrs.org/en/stable/glossary.html#term-slotted-classes + __slots__ = () + + def __call__(self, series: IntoSeries, /) -> IntoDataFrame: + """Encode a given series of values. + + Args: + series: A series in an arbitrary backend, containing the values to encode. + + Returns: + A dataframe containing the encoded representations of the input values, in + the same backend as the input series and with the same row order. + """ + + +@define +class _Encoder: + """A narwhals wrapper for user-specified encoders to hide their native backend. + + Wraps a user-provided instance of an :class:`EncoderProtocol` and automatically + infers which native dataframe backend it expects via trial and error. The inferred + backend is cached after the first successful call so that trial-and-error detection + runs only once. + """ + + _encoder: EncoderProtocol = field( + alias="encoder", validator=instance_of(EncoderProtocol) + ) + """The user-provided encoder.""" + + _implementation: nw.Implementation | None = field( + default=None, init=False, eq=False + ) + """The inferred native backend, cached after the first successful call.""" + + def __call__(self, series: nw.Series, /) -> nw.DataFrame: + """Encode a narwhals series, inferring the required backend if not yet known. + + On the first call, tries available backends in order (starting with + :attr:`~baybe.settings.Settings.default_dataframe_backend`) until the wrapped + callable succeeds. The successful backend is cached for all subsequent calls. + + Args: + series: A narwhals series containing the values to encode. + + Returns: + A narwhals dataframe containing the encoded representations, collected into + the same backend used for the input series. + """ + if self._implementation is None: + self._implementation, result = self._infer_backend(series) + return result + + return self._encode(series, self._implementation) + + def _infer_backend( + self, series: nw.Series, / + ) -> tuple[nw.Implementation, nw.DataFrame]: + """Infer the encoder's backend by trial-and-error across available backends. + + Args: + series: The series to use for probing. + + Returns: + A tuple containing: + + * the first backend for which the encoder succeeds + * the resulting encoded dataframe + + Raises: + ExceptionGroup: If the encoder raises an exception for every tried backend. + """ + preferred = active_settings.default_dataframe_backend + ordered = [preferred] + [b for b in nw.Implementation if b is not preferred] + backends = [b for b in ordered if _is_backend_imported(b)] + + exceptions: list[Exception] = [] + for backend in backends: + try: + result = self._encode(series, backend) + return backend, result + except Exception as ex: # noqa: BLE001 + exceptions.append(ex) + + raise ExceptionGroup("The encoder failed for all tried backends", exceptions) + + def _encode(self, series: nw.Series, backend: nw.Implementation, /) -> nw.DataFrame: + """Call the encoder with the given series using the specified backend. + + Args: + series: The series to encode. + backend: The native backend to convert the series to before encoding. + + Returns: + A narwhals dataframe in the original series backend. + """ + native_series = _df_with_backend(series, backend).to_native() + result = nw.from_native(self._encoder(native_series), eager_only=True) + return _df_with_backend(result, series.implementation) + + +def _is_backend_imported(backend: nw.Implementation) -> bool: + """Check if the given native backend has already been imported.""" + getter = getattr(nw.dependencies, f"get_{backend.value}", None) + return getter is not None and getter() is not None + + +# Collect leftover original slotted classes processed by `attrs.define` +gc.collect() diff --git a/baybe/parameters/sequence.py b/baybe/parameters/sequence.py new file mode 100644 index 0000000000..24204879ba --- /dev/null +++ b/baybe/parameters/sequence.py @@ -0,0 +1,144 @@ +"""Sequence parameters.""" + +from __future__ import annotations + +import gc +from collections.abc import Sequence +from functools import cached_property +from itertools import chain, product + +import narwhals.stable.v2 as nw +from attrs import Attribute, Converter, define, field +from attrs.validators import deep_iterable, ge, instance_of, min_len, optional +from typing_extensions import override + +from baybe.exceptions import InfiniteSpaceError +from baybe.parameters.base import _JOIN_KEY, _EncodedDiscreteParameter +from baybe.parameters.encoding import _Encoder +from baybe.utils.conversion import nonstring_to_tuple + + +@define(frozen=True, slots=False) +class SequenceParameter(_EncodedDiscreteParameter): + """Parameter class for sequence parameters.""" + + alphabet: tuple[str, ...] = field( + converter=Converter( # type: ignore[misc] + lambda value, self, field: tuple( + sorted(nonstring_to_tuple(value, self, field)) + ), + takes_self=True, + takes_field=True, + ), + validator=deep_iterable( + member_validator=(instance_of(str), min_len(1)), + iterable_validator=min_len(1), + ), + ) + """The alphabet defining the tokens used to construct the sequences.""" + + encoder: _Encoder = field( + converter=lambda v: v if isinstance(v, _Encoder) else _Encoder(encoder=v), + ) + """The encoder used to map sequence values to their computational representation. + + Can be implemented in any narwhals-supported dataframe backend since an + automatically applied wrapper layer handles backend conversion when necessary. + """ + + min_length: int = field( + default=1, validator=(instance_of(int), ge(1)), kw_only=True + ) + """The minimum token length of the constructed sequences.""" + + max_length: int | None = field( + default=None, validator=optional(instance_of(int)), kw_only=True + ) + """Optional maximum token length of the constructed sequences.""" + + @max_length.validator + def _validate_max_length( # noqa: DOC101, DOC103 + self, _: Attribute, value: int | None + ) -> None: + if value is not None and value < self.min_length: + raise ValueError( + f"The maximum sequence length ({value}) must be greater than or " + f"equal to the minimum sequence length ({self.min_length})." + ) + + @override + @property + def is_finite(self) -> bool: + return self.max_length is not None + + @override + def __len__(self) -> int: + if not self.is_finite: + raise InfiniteSpaceError( + f"Cannot determine the number of sequences of a " + f"'{self.__class__.__name__}' that has no explicit maximum length." + ) + assert self.max_length is not None + return sum( + len(self.alphabet) ** length + for length in range(self.min_length, self.max_length + 1) + ) + + @override + @cached_property + def values(self) -> tuple[tuple[str, ...], ...]: + if not self.is_finite: + raise InfiniteSpaceError( + f"Cannot enumerate the sequences of a '{self.__class__.__name__}' " + "that has no explicit maximum length." + ) + assert self.max_length is not None + all_values = map( + tuple, + chain.from_iterable( + product(self.alphabet, repeat=length) + for length in range(self.min_length, self.max_length + 1) + ), + ) + return tuple(all_values) + + @property + @override + def comp_rep_columns(self) -> tuple[str, ...]: + # TODO: Override can be dropped once method is removed from base class + raise NotImplementedError() + + @override + def _encoding_table(self, values: nw.Series, /) -> nw.DataFrame: + return self.encoder(values).with_columns(values.rename(_JOIN_KEY)) + + @override + def is_in_range(self, item: Sequence[str]) -> bool: + try: + item = nonstring_to_tuple(item) + except (TypeError, ValueError): + return False + length = len(item) + if length < self.min_length or ( + self.max_length is not None and length > self.max_length + ): + return False + + return all(ch in self.alphabet for ch in item) + + @override + def summary(self) -> dict: + information: dict[str, object] = dict( + Name=self.name, + Type=self.__class__.__name__, + Alphabet=self.alphabet, + MinLength=self.min_length, + ) + if self.max_length is not None: + information["MaxLength"] = self.max_length + information["nValues"] = len(self) + return information + + +# Collect leftover original slotted classes processed by `attrs.define` +gc.collect() diff --git a/baybe/parameters/substance.py b/baybe/parameters/substance.py index 8160f1aa4d..c35be650dd 100644 --- a/baybe/parameters/substance.py +++ b/baybe/parameters/substance.py @@ -55,7 +55,7 @@ class SubstanceParameter(_EncodedDiscreteParameter): encoding: SubstanceEncoding = field( default=SubstanceEncoding.MORDRED, converter=SubstanceEncoding, kw_only=True ) - # See base class. + """The encoding used the generate the parameters computational representation.""" decorrelate: bool | float = field( default=True, validator=validate_decorrelation, kw_only=True diff --git a/baybe/searchspace/core.py b/baybe/searchspace/core.py index b58c19494b..4185390aee 100644 --- a/baybe/searchspace/core.py +++ b/baybe/searchspace/core.py @@ -276,7 +276,7 @@ def n_tasks(self) -> int: if (task_param := self._task_parameter) is None: # When there are no task parameters, we effectively have a single task return 1 - return len(task_param.values) + return len(task_param) @property def n_subsets(self) -> int: diff --git a/baybe/serialization/core.py b/baybe/serialization/core.py index 21123511ec..ab1b33804a 100644 --- a/baybe/serialization/core.py +++ b/baybe/serialization/core.py @@ -102,6 +102,15 @@ def unstructure_base(obj: Any) -> dict[str, Any]: pass dct = hook(obj) + + # The hook must return a dictionary since it is intended to be applied to + # serializable container classes (e.g. attrs classes) + if not isinstance(dct, dict): + raise TypeError( + f"Expected a type of '{base.__name__}' that supports serialization. " + f"Passed object: {obj!r} (type: {type(obj).__name__})", + ) + return _add_type_to_dict(dct, obj.__class__.__name__) return unstructure_base diff --git a/baybe/utils/conversion.py b/baybe/utils/conversion.py index ed5f622532..f11eb9a8ea 100644 --- a/baybe/utils/conversion.py +++ b/baybe/utils/conversion.py @@ -32,13 +32,36 @@ def fraction_to_float(value: str | float | Fraction, /) -> float: return float(value) -def nonstring_to_tuple(x: Sequence[_T], self: type, field: Attribute) -> tuple[_T, ...]: - """Convert a sequence to tuple but raise an exception for string input.""" +def nonstring_to_tuple( + x: Sequence[_T], + self: object = None, + field: Attribute | None = None, + /, +) -> tuple[_T, ...]: + """Convert a sequence to tuple but raise an exception for string input. + + Can be used for plain conversion or as a converter for an attrs field. + + Args: + x: The sequence to be converted. + self: The object owning the field, used for error reporting. When provided, + its class name is included in the error message. + field: The field descriptor, used for error reporting. When provided, its + alias is included in the error message. + + Returns: + The tuple representation of the given sequence. + + Raises: + ValueError: If the provided value is a string. + """ if isinstance(x, str): - raise ValueError( - f"Argument passed to '{field.alias}' of class '{self.__class__.__name__}' " - f"must be a sequence but cannot be a string." + context = ( + "Argument" + + (f" '{field.alias}'" if field is not None else "") + + (f" of class '{self.__class__.__name__}'" if self is not None else "") ) + raise ValueError(f"{context} must be a sequence but cannot be a string.") return tuple(x) diff --git a/baybe/utils/dataframe.py b/baybe/utils/dataframe.py index 5bc70e9f18..1f3c4f6d43 100644 --- a/baybe/utils/dataframe.py +++ b/baybe/utils/dataframe.py @@ -11,6 +11,7 @@ import numpy as np import pandas as pd from narwhals.testing import assert_frame_equal +from narwhals.typing import IntoDataFrame, IntoSeries from typing_extensions import assert_never from baybe.exceptions import InputDataTypeWarning, SearchSpaceMatchWarning @@ -311,6 +312,42 @@ def df_uncorrelated_features( return data +_SeriesOrFrameT = TypeVar( + "_SeriesOrFrameT", nw.Series, nw.DataFrame, IntoSeries, IntoDataFrame +) + + +def _df_with_backend( + obj: _SeriesOrFrameT, backend: nw.Implementation, / +) -> _SeriesOrFrameT: + """Convert a native/narwhals Series/DataFrame to a different native backend. + + Args: + obj: The native/narwhals Series/DataFrame to convert. + backend: The target backend to convert to. + + Returns: + The input object converted to specified backend. + """ + # TODO: Replace once built-in solution is available + # https://github.com/narwhals-dev/narwhals/issues/3812 + + if is_native := not isinstance(obj, nw.Series | nw.DataFrame): + obj = nw.from_native(obj, allow_series=True) + + if isinstance(obj, nw.Series): + name = obj.name + frame = nw.from_dict(obj.to_frame().to_dict(), backend=backend) + result = frame[name] + else: + result = nw.from_dict(obj.to_dict(), backend=backend) + + if is_native: + return result.to_native() + + return result + + def add_noise_to_perturb_degenerate_rows( df: pd.DataFrame, noise_ratio: float = 0.001 ) -> pd.DataFrame: diff --git a/baybe/utils/validation.py b/baybe/utils/validation.py index 10f44301b1..ffc70c10ae 100644 --- a/baybe/utils/validation.py +++ b/baybe/utils/validation.py @@ -211,16 +211,9 @@ def validate_parameter_input( ) # Check if all rows have valid inputs matching allowed parameter values - if p.is_numerical: - valid = ( - not numerical_measurements_must_be_within_tolerance - or data[p.name].map(p.is_in_range).all() - ) - else: - from baybe.parameters.base import _EncodedDiscreteParameter - - assert isinstance(p, _EncodedDiscreteParameter) - valid = data[p.name].isin(p.values).all() + valid = ( + p.is_numerical and not numerical_measurements_must_be_within_tolerance + ) or data[p.name].map(p.is_in_range).all() if not valid: raise ValueError( f"The provided dataframe has invalid values for parameter '{p.name}'. " diff --git a/tests/test_encoding.py b/tests/test_encoding.py new file mode 100644 index 0000000000..4a1e2a028e --- /dev/null +++ b/tests/test_encoding.py @@ -0,0 +1,93 @@ +"""Tests for parameter encoders.""" + +from __future__ import annotations + +from typing import Any + +import narwhals.stable.v2 as nw +import pandas as pd +import polars as pl +import pytest +from exceptiongroup import ExceptionGroup +from typing_extensions import Never + +from baybe.parameters.encoding import _Encoder + + +def _pandas_encoder(series: pd.Series) -> pd.DataFrame: + """Pandas passthrough encoder.""" + assert isinstance(series, pd.Series), f"Expected pd.Series, got {type(series)}" + return series.to_frame() + + +def _polars_encoder(series: pl.Series) -> pl.DataFrame: + """Polars passthrough encoder.""" + assert isinstance(series, pl.Series), f"Expected pl.Series, got {type(series)}" + return series.to_frame() + + +def _nw_series(backend: str = "polars") -> nw.Series: + """Create a narwhals Series.""" + return nw.new_series(name="x", values=["A", "B"], backend=backend) + + +@pytest.mark.parametrize("series_backend", ["pandas", "polars"]) +@pytest.mark.parametrize("encoder_backend", ["pandas", "polars"]) +def test_wrapped_encoder_accepts_any_input_backend(series_backend, encoder_backend): + """A wrapped encoder handles any input/encoder backend combination.""" + user_encoder = globals()[f"_{encoder_backend}_encoder"] + encoder = _Encoder(encoder=user_encoder) + assert encoder._implementation is None + + series = _nw_series(backend=series_backend) + encoded = encoder(series) + + # The wrapper correctly infers the user-native backend. + assert encoder._implementation is nw.Implementation.from_string(encoder_backend) + + # The returned output backend matches the input series backend, regardless of the + # encoder's native backend. + assert nw.get_native_namespace(encoded) is nw.get_native_namespace(series) + + +def test_callable_not_retried_after_backend_cached(): + """Once the backend is cached, the encoder is queried once per call.""" + call_count = 0 + + def counting_encoder(series: pd.Series) -> pd.DataFrame: + nonlocal call_count + call_count += 1 + assert isinstance(series, pd.Series), f"Expected pd.Series, got {type(series)}" + return pd.DataFrame() + + encoder = _Encoder(encoder=counting_encoder) + + # The first call may trigger several inner calls to determine the backend + encoder(_nw_series()) + initial_call_count = call_count + assert initial_call_count > 0 + + # Once the backend is cached, only one call is required to encode + encoder(_nw_series()) + assert call_count == initial_call_count + 1 + + +def test_broken_encoder_raises_exception_group(): + """An encoder that fails for all backends raises an ExceptionGroup.""" + + def broken_encoder(_: Any) -> Never: + raise ValueError("intentional failure") + + enc = _Encoder(encoder=broken_encoder) + with pytest.raises(ExceptionGroup, match="failed for all tried backends"): + enc(_nw_series()) + + +def test_equality_after_encoding(): + """Encoders still compare equal once the implementation backend has been cached.""" + encoder = lambda s: s.to_frame() # noqa: E731 + e1 = _Encoder(encoder) + e2 = _Encoder(encoder) + assert e1 == e2 + e1(_nw_series()) + assert e1 == e2 diff --git a/tests/test_integration.py b/tests/test_integration.py index c1ff3132f4..b3e8806c7d 100644 --- a/tests/test_integration.py +++ b/tests/test_integration.py @@ -19,9 +19,7 @@ p1 = NumericalDiscreteParameter("p1", [1, 2]) t1 = NumericalTarget("t1") objective = t1.to_objective() -measurements = pd.DataFrame( - {p1.name: p1.values, t1.name: np.random.random(len(p1.values))} -) +measurements = pd.DataFrame({p1.name: p1.values, t1.name: np.random.random(len(p1))}) @pytest.fixture(name="searchspace") diff --git a/tests/test_parameters.py b/tests/test_parameters.py index 2c0d68f12e..1587c3a2e4 100644 --- a/tests/test_parameters.py +++ b/tests/test_parameters.py @@ -14,6 +14,7 @@ from baybe.parameters.categorical import CategoricalParameter from baybe.parameters.custom import CustomDiscreteParameter from baybe.parameters.numerical import NumericalDiscreteParameter +from baybe.parameters.sequence import SequenceParameter from baybe.settings import active_settings if CHEM_INSTALLED: @@ -125,6 +126,19 @@ def _list(name, values): ), id="custom", ), + pytest.param( + SequenceParameter( + name="seq", + alphabet=("A", "C"), + encoder=lambda values: pd.DataFrame( + {"seq": ["".join(v) for v in values]} + ), + min_length=1, + max_length=1, + ), + pd.DataFrame({"seq": ["A", "C"]}, index=pd.Index(["A", "C"])), + id="sequence", + ), pytest.param( SubstanceParameter( name="solvent", diff --git a/tests/test_sequence_parameter.py b/tests/test_sequence_parameter.py new file mode 100644 index 0000000000..cd21d11172 --- /dev/null +++ b/tests/test_sequence_parameter.py @@ -0,0 +1,99 @@ +"""Tests for SequenceParameter behavior not covered by shared parameter tests.""" + +from __future__ import annotations + +import pytest + +from baybe.exceptions import InfiniteSpaceError +from baybe.parameters.sequence import SequenceParameter + +_DNA = ("A", "C", "G", "T") +_encoder = lambda series: series.to_frame() # noqa: E731 + + +@pytest.mark.parametrize( + ("range", "expected"), + [ + ((1, 1), 4), + ((1, 2), 20), + ((2, 2), 16), + ((1, 5), 1364), + ((6, 6), 4096), + ], +) +def test_length(range, expected): + """The parameter correctly computes the number of values in its range.""" + p = SequenceParameter( + name="seq", + alphabet=_DNA, + min_length=range[0], + max_length=range[1], + encoder=_encoder, + ) + assert len(p) == expected + + +@pytest.mark.parametrize("max_length", [None, 1]) +def test_is_finite(max_length): + """The parameter correctly indicates if it is finite or infinite.""" + p = SequenceParameter("seq", _DNA, max_length=max_length, encoder=_encoder) + assert p.is_finite is (max_length is not None) + + +def test_values_raises_without_max_length(): + """Accessing finite-based properties of an infinite parameter raises an error.""" + p = SequenceParameter(name="seq", alphabet=_DNA, encoder=_encoder) + with pytest.raises(InfiniteSpaceError): + p.values + with pytest.raises(InfiniteSpaceError): + len(p) + + +def test_values_construction(): + """The generated sequences match with the specified range.""" + p = SequenceParameter( + name="seq", + alphabet=("A", "BC", "D"), + encoder=_encoder, + min_length=1, + max_length=2, + ) + assert p.values == ( + ("A",), + ("BC",), + ("D",), + ("A", "A"), + ("A", "BC"), + ("A", "D"), + ("BC", "A"), + ("BC", "BC"), + ("BC", "D"), + ("D", "A"), + ("D", "BC"), + ("D", "D"), + ) + + +@pytest.mark.parametrize( + ("item", "expected"), + [ + pytest.param(("A", "C", "G", "T"), True, id="valid"), + pytest.param(("A", "C", "G", "T"), True, id="valid_non_tuple"), + pytest.param("ACGT", False, id="plain_string"), + pytest.param(42, False, id="wrong_type"), + pytest.param(("A", "X"), False, id="out_of_alphabet"), + pytest.param(("A", "C", "G", "T", "A"), False, id="too_long"), + pytest.param(("A",), False, id="too_short"), + ], +) +def test_is_in_range(item, expected): + """In-range check validates element-level alphabet membership and length.""" + p = SequenceParameter("seq", _DNA, encoder=_encoder, min_length=2, max_length=4) + assert p.is_in_range(item) is expected + + +def test_equality_modulo_alphabet_ordering(): + """Alphabet ordering does not affect equality of parameters.""" + p1 = SequenceParameter("seq", ("A", "C", "G", "T"), encoder=_encoder) + p2 = SequenceParameter("seq", ("T", "G", "C", "A"), encoder=_encoder) + assert p1 == p2 diff --git a/tests/validation/test_parameter_validation.py b/tests/validation/test_parameter_validation.py index f86c515f3a..4987fc7432 100644 --- a/tests/validation/test_parameter_validation.py +++ b/tests/validation/test_parameter_validation.py @@ -18,6 +18,7 @@ NumericalContinuousParameter, NumericalDiscreteParameter, ) +from baybe.parameters.sequence import SequenceParameter from baybe.parameters.substance import SubstanceParameter from baybe.parameters.validation import validate_decorrelation from baybe.utils.interval import InfiniteIntervalError @@ -115,6 +116,35 @@ def test_invalid_encoding_categorical_parameter(): CategoricalParameter(name="invalid_encoding", values=["A", "B"], encoding="enc") +@pytest.mark.parametrize( + ("override", "error"), + [ + param({"alphabet": set()}, ValueError, id="empty_alphabet"), + param({"alphabet": {""}}, ValueError, id="element_too_short"), + param({"alphabet": {1}}, TypeError, id="nonstring_element"), + param({"encoder": object()}, TypeError, id="non_encoder_object"), + param({"min_length": ""}, TypeError, id="non_int_min_length"), + param({"max_length": ""}, TypeError, id="non_int_max_length"), + param({"min_length": 0}, ValueError, id="min_length_zero"), + param({"max_length": 0}, ValueError, id="max_length_zero"), + param({"min_length": 3, "max_length": 2}, ValueError, id="max_less_than_min"), + ], +) +def test_invalid_sequence_parameter(override, error): + """Providing invalid arguments to SequenceParameter raises an exception.""" + defaults = { + "name": "seq", + "alphabet": {"A", "C", "G", "T"}, + "encoder": lambda values: pd.DataFrame( + {"seq": [tuple(v) for v in values]}, index=pd.Index(values) + ), + "min_length": 1, + "max_length": 2, + } + with pytest.raises(error): + SequenceParameter(**{**defaults, **override}) + + @pytest.mark.parametrize( ("values", "active_values", "error"), [