From 968c52b53b7f716b5f7f71b4f08a8129d43f2a98 Mon Sep 17 00:00:00 2001 From: kalama-ai Date: Tue, 7 Jul 2026 16:31:18 +0200 Subject: [PATCH] Add transfer-learning override for the GP kernel factory - Add `TaskParameter.override_transfer_learning_mode` and the `TransferLearningMode` enum to force `IndexKernel`/`PositiveIndexKernel` - Resolve the task kernel in `GaussianProcessSurrogate._resolve_kernel`, combining it with a task-free base kernel - Raise `IncompatibleOverrideError` for unsupported kernels/factories --- CHANGELOG.md | 4 + baybe/exceptions.py | 8 + baybe/kernels/base.py | 36 ++- baybe/kernels/composite.py | 10 +- baybe/parameters/__init__.py | 2 + baybe/parameters/categorical.py | 15 +- baybe/parameters/enum.py | 11 + baybe/searchspace/core.py | 7 +- baybe/surrogates/gaussian_process/core.py | 168 +++++++++++++- tests/conftest.py | 13 ++ tests/hypothesis_strategies/kernels.py | 11 +- tests/hypothesis_strategies/parameters.py | 10 +- tests/test_iterations.py | 15 ++ tests/test_kernel_factories.py | 263 +++++++++++++++++++++- tests/test_kernels.py | 70 +++++- tests/test_reduced_searchspace.py | 8 + 16 files changed, 623 insertions(+), 28 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5e316267c7..78626ec490 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -36,6 +36,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 inverting the specification to keep the complement - `DiscreteSelectionConstraint` as the condition-based filtering constraint (inclusion-by-default; replaces `DiscreteExcludeConstraint`) +- `TaskParameter.override_transfer_learning_mode` (and the corresponding + `TransferLearningMode` enum) for selecting the kernel that models the task + correlations in transfer learning, taking precedence over the task kernel of the + configured kernel factory ### Changed - `BOTORCH` GP preset now includes `BetaPrior(2.5, 1.5)` for the task covariance diff --git a/baybe/exceptions.py b/baybe/exceptions.py index 1018d0cbb7..407c88a26f 100644 --- a/baybe/exceptions.py +++ b/baybe/exceptions.py @@ -74,6 +74,14 @@ class IncompatibleArgumentError(IncompatibilityError): """An incompatible argument was passed to a callable.""" +class IncompatibleOverrideError(IncompatibilityError): + """An override conflicts with another specification.""" + + +class _UnsupportedSearchSpaceAttributeError(AttributeError): + """Access to a blocked attribute on a reduced search space was attempted.""" + + class NonGaussianityError(Exception): """An operation assuming Gaussianity is attempted on a non-Gaussian distribution.""" diff --git a/baybe/kernels/base.py b/baybe/kernels/base.py index 980fedaa0c..7a07f97096 100644 --- a/baybe/kernels/base.py +++ b/baybe/kernels/base.py @@ -7,7 +7,7 @@ from itertools import chain from typing import TYPE_CHECKING, Any -from attrs import define, field +from attrs import define, evolve, field from attrs.converters import optional as optional_c from attrs.validators import deep_iterable, instance_of from attrs.validators import optional as optional_v @@ -99,6 +99,28 @@ def to_factory(self) -> PlainKernelFactory: return PlainKernelFactory(self) + def _without_parameter( + self, name: str, searchspace: SearchSpace, / + ) -> Kernel | None: + """Return a copy of the kernel that no longer acts on the specified parameter. + + Args: + name: The name of the parameter to ignore. + searchspace: The search space, used to enumerate the remaining parameter + names when the kernel does not explicitly specify the ones it acts on. + + Raises: + TypeError: If the kernel structure does not support removing a single + parameter unambiguously. + + Returns: + The reduced kernel, or ``None`` if removing the parameter leaves the + kernel with no parameters to act on. + """ + raise TypeError( + f"Cannot remove a parameter from kernel '{self.__class__.__name__}'. " + ) + @abstractmethod def _get_dimensions( self, searchspace: SearchSpace @@ -239,6 +261,18 @@ def _get_dimensions( ) return active_dims, ard_num_dims + @override + def _without_parameter( + self, name: str, searchspace: SearchSpace, / + ) -> Kernel | None: + if self.parameter_names is None: + remaining = tuple(n for n in searchspace.parameter_names if n != name) + elif name in self.parameter_names: + remaining = tuple(n for n in self.parameter_names if n != name) + else: + return self + return evolve(self, parameter_names=remaining) if remaining else None + @define(frozen=True) class CompositeKernel(Kernel, ABC): diff --git a/baybe/kernels/composite.py b/baybe/kernels/composite.py index 2274416c25..431f0fd953 100644 --- a/baybe/kernels/composite.py +++ b/baybe/kernels/composite.py @@ -4,7 +4,7 @@ from functools import reduce from operator import add, mul -from attrs import define, field +from attrs import define, evolve, field from attrs.converters import optional as optional_c from attrs.validators import deep_iterable, gt, instance_of, min_len from attrs.validators import optional as optional_v @@ -12,6 +12,7 @@ from baybe.kernels.base import CompositeKernel, Kernel from baybe.priors.base import Prior +from baybe.searchspace.core import SearchSpace from baybe.settings import active_settings from baybe.utils.basic import to_tuple from baybe.utils.validation import finite_float @@ -42,6 +43,13 @@ class ScaleKernel(CompositeKernel): If ``False``, the output scale is frozen at its initial value and excluded from optimization.""" + @override + def _without_parameter( + self, name: str, searchspace: SearchSpace, / + ) -> Kernel | None: + stripped = self.base_kernel._without_parameter(name, searchspace) + return None if stripped is None else evolve(self, base_kernel=stripped) + @override def to_gpytorch(self, *args, **kwargs): import torch diff --git a/baybe/parameters/__init__.py b/baybe/parameters/__init__.py index 93e62b6ee9..259f3c049a 100644 --- a/baybe/parameters/__init__.py +++ b/baybe/parameters/__init__.py @@ -6,6 +6,7 @@ CategoricalEncoding, CustomEncoding, SubstanceEncoding, + TransferLearningMode, ) from baybe.parameters.numerical import ( NumericalContinuousParameter, @@ -25,4 +26,5 @@ "SubstanceEncoding", "SubstanceParameter", "TaskParameter", + "TransferLearningMode", ] diff --git a/baybe/parameters/categorical.py b/baybe/parameters/categorical.py index 119d220cad..a98f4dc1a9 100644 --- a/baybe/parameters/categorical.py +++ b/baybe/parameters/categorical.py @@ -5,12 +5,13 @@ import numpy as np import pandas as pd +from attr.converters import optional as optional_c from attrs import Converter, define, field from attrs.validators import deep_iterable, instance_of, min_len from typing_extensions import override from baybe.parameters.base import _DiscreteLabelLikeParameter -from baybe.parameters.enum import CategoricalEncoding +from baybe.parameters.enum import CategoricalEncoding, TransferLearningMode from baybe.settings import active_settings from baybe.utils.conversion import nonstring_to_tuple, sort_tuple from baybe.utils.validation import validate_unique_values @@ -87,6 +88,18 @@ class TaskParameter(CategoricalParameter): encoding: CategoricalEncoding = field(default=CategoricalEncoding.INT, init=False) # See base class. + override_transfer_learning_mode: TransferLearningMode | None = field( + default=None, + converter=optional_c(TransferLearningMode), + ) + """Optional override for how the task dimension is modeled. + + Only applies to :class:`.GaussianProcessSurrogate`. When ``None``, the surrogate's + kernel factory decides how the task dimension is treated. When set, the surrogate + attaches the requested task kernel to a task-free base kernel derived from the + configured factory. + """ + # Collect leftover original slotted classes processed by `attrs.define` gc.collect() diff --git a/baybe/parameters/enum.py b/baybe/parameters/enum.py index 622fa4af58..45a32bb2c1 100644 --- a/baybe/parameters/enum.py +++ b/baybe/parameters/enum.py @@ -174,3 +174,14 @@ class SubstanceEncoding(ParameterEncoding): WHIM = "WHIM" """:class:`skfp.fingerprints.WHIMFingerprint`""" + + +class TransferLearningMode(Enum): + """Transfer learning modes for :class:`.TaskParameter`.""" + + INDEX_KERNEL = "INDEX_KERNEL" + """:class:`gpytorch.kernels.IndexKernel` for arbitrary correlations.""" + + POSITIVE_INDEX_KERNEL = "POSITIVE_INDEX_KERNEL" + """:class:`botorch.models.kernels.positive_index.PositiveIndexKernel` for positive + correlations.""" diff --git a/baybe/searchspace/core.py b/baybe/searchspace/core.py index b43533817c..90093bbee9 100644 --- a/baybe/searchspace/core.py +++ b/baybe/searchspace/core.py @@ -16,7 +16,10 @@ from baybe.constraints import validate_constraints from baybe.constraints.base import Constraint -from baybe.exceptions import InfeasibilityError +from baybe.exceptions import ( + InfeasibilityError, + _UnsupportedSearchSpaceAttributeError, +) from baybe.parameters import TaskParameter from baybe.parameters.base import Parameter from baybe.searchspace.continuous import SubspaceContinuous @@ -616,7 +619,7 @@ def __getattribute__(self, name: str): allowed = object.__getattribute__(self, "_ALLOWED_ATTRIBUTES") if name in allowed: return object.__getattribute__(self, name) - raise AttributeError( + raise _UnsupportedSearchSpaceAttributeError( f"'{object.__getattribute__(self, '__class__').__name__}' does not " f"support attribute '{name}'. Only parameter information is available." ) diff --git a/baybe/surrogates/gaussian_process/core.py b/baybe/surrogates/gaussian_process/core.py index 8afb92b7f1..625227cfb7 100644 --- a/baybe/surrogates/gaussian_process/core.py +++ b/baybe/surrogates/gaussian_process/core.py @@ -10,17 +10,25 @@ from typing import TYPE_CHECKING, ClassVar import pandas as pd -from attrs import Converter, define, field +from attrs import Converter, define, field, fields from attrs.converters import optional as optional_c from attrs.converters import pipe from attrs.validators import instance_of, is_callable, optional -from typing_extensions import Self, override - -from baybe.exceptions import DeprecationError, ModelNotTrainedError +from typing_extensions import Self, assert_never, override + +from baybe.exceptions import ( + DeprecationError, + IncompatibleOverrideError, + IncompatibleSearchSpaceError, + ModelNotTrainedError, + _UnsupportedSearchSpaceAttributeError, +) from baybe.kernels.base import Kernel +from baybe.kernels.basic import IndexKernel, PositiveIndexKernel from baybe.objectives.base import Objective from baybe.parameters.base import Parameter from baybe.parameters.categorical import TaskParameter +from baybe.parameters.enum import TransferLearningMode from baybe.searchspace.core import SearchSpace from baybe.surrogates.base import Surrogate from baybe.surrogates.gaussian_process.components.fit_criterion import ( @@ -29,6 +37,7 @@ ) from baybe.surrogates.gaussian_process.components.generic import ( GPComponentType, + PlainGPComponentFactory, to_component_factory, ) from baybe.surrogates.gaussian_process.components.kernel import ( @@ -82,6 +91,14 @@ def task_idx(self) -> int | None: """The computational column index of the task parameter, if available.""" return self.searchspace.task_idx + @property + def tl_override(self) -> TransferLearningMode | None: + """The task parameter's transfer learning override, if any.""" + task_param = self.searchspace._task_parameter + return ( + None if task_param is None else task_param.override_transfer_learning_mode + ) + @property def is_multitask(self) -> bool: """Indicates if model is to be operated in a multi-task context.""" @@ -167,6 +184,10 @@ class GaussianProcessSurrogate(Surrogate): * :class:`baybe.kernels.base.Kernel` * :obj:`.components.kernel.KernelFactoryProtocol` * :class:`gpytorch.kernels.Kernel` + + If a :class:`.TaskParameter` sets ``override_transfer_learning_mode``, this must + reduce to a task-free BayBE kernel or an :class:`.IncompatibleOverrideError` is + raised. """ mean_factory: MeanFactoryProtocol | None = field( @@ -375,6 +396,131 @@ def _posterior(self, candidates_comp_scaled: Tensor, /) -> Posterior: assert self._model is not None return self._model.posterior(candidates_comp_scaled) + def _resolve_kernel(self, context: _ModelContext) -> GPyTorchKernel: + """Resolve the GP kernel, dispatching on task parameter overrides. + + Args: + context: The model context providing searchspace information. + + Raises: + IncompatibleOverrideError: If a transfer learning override is combined + with a kernel or kernel factory that cannot be reduced to a task-free + base kernel operating on parameter names. + + Returns: + The constructed gpytorch kernel. + """ + searchspace = context.searchspace + task_param = searchspace._task_parameter + tl_override = context.tl_override + + if tl_override is None: + # No override: let the factory handle everything (default path) + kernel_factory = self.kernel_factory or BayBEKernelFactory() + kernel = kernel_factory( + searchspace, context.objective, context.measurements + ) + if isinstance(kernel, Kernel): + kernel = kernel.to_gpytorch(searchspace=searchspace) + return kernel + + assert task_param is not None # a set override implies a task parameter + + # Override is set: assemble the prescribed task kernel + n_tasks = searchspace.n_tasks + task_kernel_cls: type[IndexKernel] + match tl_override: + case TransferLearningMode.POSITIVE_INDEX_KERNEL: + task_kernel_cls = PositiveIndexKernel + case TransferLearningMode.INDEX_KERNEL: + task_kernel_cls = IndexKernel + case _: + assert_never(tl_override) + task_kernel_spec = task_kernel_cls( + num_tasks=n_tasks, rank=n_tasks, parameter_names=(task_param.name,) + ) + + # Default factory (None or a `BayBEKernelFactory` without a custom parameter + # selector): reuse the ICM machinery on the full searchspace, which builds the + # task-excluded base kernel and combines it with the prescribed task kernel. + # This avoids the reduced searchspace, on which the default factory's numerical + # kernel cannot resolve its active dimensions. A `BayBEKernelFactory` carrying a + # custom selector falls through to the general factory path below, so the + # selector is honored (or loudly rejected when it cannot be). + if self.kernel_factory is None or ( + isinstance(self.kernel_factory, BayBEKernelFactory) + and self.kernel_factory.parameter_selector is None + ): + icm = ICMKernelFactory(task_kernel_or_factory=task_kernel_spec) + kernel = icm(searchspace, context.objective, context.measurements) + if isinstance(kernel, Kernel): + kernel = kernel.to_gpytorch(searchspace=searchspace) + return kernel + + # Otherwise, build a task-free base kernel and attach the prescribed task + # kernel manually. + effective_factory = self.kernel_factory + incompatible_message = ( + f"The '{TaskParameter.__name__}' '{task_param.name}' specifies " + f"'{fields(TaskParameter).override_transfer_learning_mode.name}=" + f"{tl_override.name}', which requires a " + f"kernel (factory) that yields a task-free BayBE kernel operating on " + f"parameter names. The provided kernel factory " + f"'{type(effective_factory).__name__}' does not satisfy this (e.g., it " + f"returns a raw gpytorch kernel or already operates on the task " + f"parameter)." + ) + + if isinstance(self.kernel_factory, PlainGPComponentFactory): + # A fixed kernel was provided: strip the task parameter directly. + component = self.kernel_factory.component + if not isinstance(component, Kernel): + raise IncompatibleOverrideError(incompatible_message) + try: + base_spec = component._without_parameter(task_param.name, searchspace) + except TypeError as ex: + raise IncompatibleOverrideError(incompatible_message) from ex + else: + # Call the factory on a reduced (task-free) searchspace so that it + # produces only the base kernel. Factories that need computational + # information unavailable on the reduced space, or that return a raw + # gpytorch kernel, are not supported. + reduced_searchspace = searchspace._drop_parameters({task_param.name}) + try: + factory_kernel = effective_factory( + reduced_searchspace, context.objective, context.measurements + ) + except ( + IncompatibleSearchSpaceError, + _UnsupportedSearchSpaceAttributeError, + ) as ex: + raise IncompatibleOverrideError(incompatible_message) from ex + if not isinstance(factory_kernel, Kernel): + raise IncompatibleOverrideError(incompatible_message) + # Normalize to an explicitly task-free spec. + try: + base_spec = factory_kernel._without_parameter( + task_param.name, searchspace + ) + except TypeError as ex: + raise IncompatibleOverrideError(incompatible_message) from ex + + # Stripping left no base kernel: return only the prescribed task kernel. + if base_spec is None: + return task_kernel_spec.to_gpytorch(searchspace=searchspace) + + # Combine base and task kernel via the ICM machinery (as in the default-factory + # branch above), which converts both on the full searchspace and validates the + # dimension partitioning between base and task kernel. + icm = ICMKernelFactory( + base_kernel_or_factory=base_spec, + task_kernel_or_factory=task_kernel_spec, + ) + kernel = icm(searchspace, context.objective, context.measurements) + if isinstance(kernel, Kernel): + kernel = kernel.to_gpytorch(searchspace=searchspace) + return kernel + def _resolve_components( self, context: _ModelContext ) -> tuple[GPyTorchKernel, GPyTorchMean, GPyTorchLikelihood, FitCriterion]: @@ -390,20 +536,15 @@ def _resolve_components( Returns: A tuple of (kernel, mean, likelihood, criterion). """ - kernel_factory = self.kernel_factory or BayBEKernelFactory() mean_factory = self.mean_factory or BayBEMeanFactory() likelihood_factory = self.likelihood_factory or BayBELikelihoodFactory() criterion_factory = self.fit_criterion_factory or BayBEFitCriterionFactory() - mean = mean_factory( - context.searchspace, context.objective, context.measurements - ) + kernel = self._resolve_kernel(context) - kernel = kernel_factory( + mean = mean_factory( context.searchspace, context.objective, context.measurements ) - if isinstance(kernel, Kernel): - kernel = kernel.to_gpytorch(searchspace=context.searchspace) likelihood = likelihood_factory( context.searchspace, context.objective, context.measurements @@ -431,9 +572,14 @@ def _fit(self, train_x: Tensor, train_y: Tensor) -> None: context = _ModelContext(self._searchspace, self._objective, self._measurements) + # Check for custom kernel + multi-task clash (only relevant when no + # override_transfer_learning_mode is set, since the override mechanism + # handles task kernel attachment explicitly). + has_tl_override = context.tl_override is not None if ( context.is_multitask and self._custom_kernel + and not has_tl_override and not strtobool(os.getenv("BAYBE_DISABLE_CUSTOM_KERNEL_WARNING", "False")) ): raise DeprecationError( diff --git a/tests/conftest.py b/tests/conftest.py index dcfcef621f..e6ab22883d 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -55,6 +55,7 @@ SubstanceEncoding, TaskParameter, ) +from baybe.parameters.enum import TransferLearningMode from baybe.parameters.substance import SubstanceParameter from baybe.priors import GammaPrior from baybe.recommenders.meta.base import MetaRecommender @@ -312,6 +313,18 @@ def fixture_parameters( values=("A", "B", "C"), active_values=("A", "B"), ), + TaskParameter( + name="Task_index_override", + values=("A", "B", "C"), + active_values=("A", "B"), + override_transfer_learning_mode=TransferLearningMode.INDEX_KERNEL, + ), + TaskParameter( + name="Task_positive_index_override", + values=("A", "B", "C"), + active_values=("A", "B"), + override_transfer_learning_mode=TransferLearningMode.POSITIVE_INDEX_KERNEL, + ), ] if CHEM_INSTALLED: diff --git a/tests/hypothesis_strategies/kernels.py b/tests/hypothesis_strategies/kernels.py index cedf2668f8..341b14f7fe 100644 --- a/tests/hypothesis_strategies/kernels.py +++ b/tests/hypothesis_strategies/kernels.py @@ -131,17 +131,12 @@ def index_kernels( draw: st.DrawFn, parameter_names: Sequence[str] | None = None, ): - """A strategy that generates index kernels.""" + """A strategy that generates index kernels (regular and positive).""" + cls = draw(st.sampled_from([IndexKernel, PositiveIndexKernel])) num_tasks = draw(st.integers(min_value=2, max_value=5)) rank = draw(st.integers(min_value=1, max_value=num_tasks)) names = draw(active_parameter_names(parameter_names)) - if draw(st.booleans()): - return PositiveIndexKernel( - parameter_names=names, - num_tasks=num_tasks, - rank=rank, - ) - return IndexKernel(parameter_names=names, num_tasks=num_tasks, rank=rank) + return cls(parameter_names=names, num_tasks=num_tasks, rank=rank) def base_kernels(parameter_names: Sequence[str] | None = None): diff --git a/tests/hypothesis_strategies/parameters.py b/tests/hypothesis_strategies/parameters.py index 004ef185ed..0b34632f52 100644 --- a/tests/hypothesis_strategies/parameters.py +++ b/tests/hypothesis_strategies/parameters.py @@ -12,6 +12,7 @@ TaskParameter, ) from baybe.parameters.custom import CustomDiscreteParameter +from baybe.parameters.enum import TransferLearningMode from baybe.parameters.numerical import ( NumericalContinuousParameter, NumericalDiscreteParameter, @@ -160,8 +161,15 @@ def task_parameters(draw: st.DrawFn): values = draw(categories) active_values = draw(_active_values(values)) param_metadata = draw(measurable_metadata()) + override_transfer_learning_mode = draw( + st.one_of(st.none(), st.sampled_from(TransferLearningMode)) + ) return TaskParameter( - name=name, values=values, active_values=active_values, metadata=param_metadata + name=name, + values=values, + active_values=active_values, + metadata=param_metadata, + override_transfer_learning_mode=override_transfer_learning_mode, ) diff --git a/tests/test_iterations.py b/tests/test_iterations.py index f61d43fc13..9d9ee6396f 100644 --- a/tests/test_iterations.py +++ b/tests/test_iterations.py @@ -377,6 +377,21 @@ def test_kernel_factories(ongoing_campaign, n_iterations, batch_size): run_iterations(ongoing_campaign, n_iterations, batch_size) +@pytest.mark.parametrize( + "parameter_names", + [ + param(["Categorical_1", "Num_disc_1", "Task_index_override"], id="index"), + param( + ["Categorical_1", "Num_disc_1", "Task_positive_index_override"], + id="positive_index", + ), + ], +) +def test_transfer_learning_override(ongoing_campaign, n_iterations, batch_size): + """A task parameter override survives a full fit/recommend loop.""" + run_iterations(ongoing_campaign, n_iterations, batch_size) + + @pytest.mark.slow @pytest.mark.parametrize( "surrogate_model", diff --git a/tests/test_kernel_factories.py b/tests/test_kernel_factories.py index f9927dfa93..a4904b4358 100644 --- a/tests/test_kernel_factories.py +++ b/tests/test_kernel_factories.py @@ -2,17 +2,31 @@ from contextlib import nullcontext +import gpytorch import pandas as pd import pytest +from botorch.models.kernels.positive_index import ( + PositiveIndexKernel as GPyTorchPositiveIndexKernel, +) +from gpytorch.kernels import IndexKernel as GPyTorchIndexKernel from pytest import param -from baybe.exceptions import IncompatibleSearchSpaceError -from baybe.parameters.categorical import CategoricalParameter, TaskParameter +from baybe.exceptions import IncompatibleOverrideError, IncompatibleSearchSpaceError +from baybe.kernels.basic import IndexKernel, MaternKernel +from baybe.kernels.composite import ScaleKernel +from baybe.parameters.categorical import ( + CategoricalParameter, + TaskParameter, +) +from baybe.parameters.enum import TransferLearningMode from baybe.parameters.numerical import ( NumericalContinuousParameter, NumericalDiscreteParameter, ) +from baybe.parameters.selectors import TypeSelector from baybe.searchspace.core import SearchSpace +from baybe.surrogates import GaussianProcessSurrogate +from baybe.surrogates.gaussian_process.components.kernel import ICMKernelFactory from baybe.surrogates.gaussian_process.presets.baybe import ( BayBEKernelFactory, _BayBENumericalKernelFactory, @@ -24,6 +38,42 @@ _SELECT_ALL = lambda parameter: True # noqa: E731 +def _matern_kernel_factory(searchspace, objective, measurements): + """A callable kernel factory returning a plain BayBE Matern kernel. + + The kernel sets no parameter names. When combined with a transfer learning + override, it is called on the reduced (task-free) search space to build the base + kernel. + """ + return MaternKernel(parameter_names=tuple(searchspace.parameter_names)) + + +def _gpytorch_returning_factory(searchspace, objective, measurements): + """A callable kernel factory returning a raw gpytorch kernel. + + Such factories are unsupported in combination with an override, since a raw + gpytorch kernel does not operate on parameter names. + """ + return gpytorch.kernels.MaternKernel(nu=2.5) + + +def _unnamed_matern_factory(searchspace, objective, measurements): + """A callable kernel factory returning an unnamed BayBE kernel. + + A ``parameter_names=None`` kernel spans all columns on full-space conversion, + so the override path must normalize it to stay task-free. + """ + return MaternKernel() + + +def _task_named_matern_factory(searchspace, objective, measurements): + """A callable kernel factory returning a kernel that names the task parameter. + + The override path must strip the task name so the base kernel stays task-free. + """ + return MaternKernel(parameter_names=("x", "Task")) + + @pytest.mark.parametrize( ("factory", "parameters", "error"), [ @@ -74,3 +124,212 @@ def test_factory_parameter_kind_validation(factory, parameters, error): else pytest.raises(error, match="does not support") ): factory(searchspace, objective, measurements) + + +def _make_dispatch_context(override_mode): + """Build a `_ModelContext` with a numerical + task parameter for dispatch tests.""" + from baybe.surrogates.gaussian_process.core import _ModelContext + + task_param = TaskParameter( + "Task", + ["A", "B", "C"], + active_values=["A"], + override_transfer_learning_mode=override_mode, + ) + num_param = NumericalDiscreteParameter("x", [1, 2, 3, 4, 5]) + searchspace = SearchSpace.from_product([num_param, task_param]) + objective = NumericalTarget("y").to_objective() + measurements = pd.DataFrame() + return _ModelContext(searchspace, objective, measurements) + + +@pytest.mark.parametrize( + ("override_mode", "kernel_or_factory", "expected_task_kernel_cls", "has_base"), + [ + param( + None, + None, + GPyTorchPositiveIndexKernel, + True, + id="no_override+default_factory", + ), + param( + None, + ICMKernelFactory( + task_kernel_or_factory=IndexKernel( + num_tasks=3, rank=3, parameter_names=("Task",) + ) + ), + GPyTorchIndexKernel, + True, + id="no_override+custom_icm_index_kernel_escape_hatch", + ), + param( + TransferLearningMode.POSITIVE_INDEX_KERNEL, + MaternKernel(), + GPyTorchPositiveIndexKernel, + True, + id="positive_index_override+bare_baybe_matern", + ), + param( + TransferLearningMode.INDEX_KERNEL, + MaternKernel(), + GPyTorchIndexKernel, + True, + id="index_override+bare_baybe_matern", + ), + param( + TransferLearningMode.POSITIVE_INDEX_KERNEL, + MaternKernel(parameter_names=("x", "Task")), + GPyTorchPositiveIndexKernel, + True, + id="positive_index_override+baybe_matern_with_task_name", + ), + param( + TransferLearningMode.POSITIVE_INDEX_KERNEL, + ScaleKernel(MaternKernel()), + GPyTorchPositiveIndexKernel, + True, + id="positive_index_override+scaled_baybe_matern", + ), + param( + TransferLearningMode.INDEX_KERNEL, + IndexKernel(num_tasks=3, rank=3, parameter_names=("Task",)), + GPyTorchIndexKernel, + False, + id="index_override+task_only_index_kernel", + ), + param( + TransferLearningMode.INDEX_KERNEL, + ScaleKernel(IndexKernel(num_tasks=3, rank=3, parameter_names=("Task",))), + GPyTorchIndexKernel, + False, + id="index_override+scaled_task_only_index_kernel", + ), + param( + TransferLearningMode.POSITIVE_INDEX_KERNEL, + _matern_kernel_factory, + GPyTorchPositiveIndexKernel, + True, + id="positive_index_override+callable_factory", + ), + param( + TransferLearningMode.INDEX_KERNEL, + _matern_kernel_factory, + GPyTorchIndexKernel, + True, + id="index_override+callable_factory", + ), + param( + TransferLearningMode.INDEX_KERNEL, + _unnamed_matern_factory, + GPyTorchIndexKernel, + True, + id="index_override+unnamed_callable_factory", + ), + param( + TransferLearningMode.INDEX_KERNEL, + _task_named_matern_factory, + GPyTorchIndexKernel, + True, + id="index_override+task_named_callable_factory", + ), + param( + TransferLearningMode.POSITIVE_INDEX_KERNEL, + None, + GPyTorchPositiveIndexKernel, + True, + id="positive_index_override+default_factory", + ), + param( + TransferLearningMode.INDEX_KERNEL, + None, + GPyTorchIndexKernel, + True, + id="index_override+default_factory", + ), + ], +) +def test_resolve_kernel_dispatch_success( + monkeypatch, override_mode, kernel_or_factory, expected_task_kernel_cls, has_base +): + """`_resolve_kernel` produces the expected task kernel for supported inputs.""" + monkeypatch.setenv("BAYBE_DISABLE_CUSTOM_KERNEL_WARNING", "True") + + context = _make_dispatch_context(override_mode) + surrogate = GaussianProcessSurrogate(kernel_or_factory=kernel_or_factory) + + kernel = surrogate._resolve_kernel(context) + + if has_base: + # The resolved kernel is a product of base * task kernel + base_kernel, task_kernel = kernel.kernels + else: + # Stripping left no non-task parameters -> only the task kernel remains + base_kernel, task_kernel = None, kernel + assert isinstance(task_kernel, expected_task_kernel_cls) + + # Override branches must partition active dims so the task kernel acts exactly + # on the task column, while the base kernel stays task-free. + if override_mode is not None: + assert task_kernel.active_dims is not None + assert set(task_kernel.active_dims.tolist()) == {context.task_idx} + if base_kernel is not None: + assert base_kernel.active_dims is not None + assert context.task_idx not in base_kernel.active_dims.tolist() + + +@pytest.mark.parametrize( + ("override_mode", "kernel_or_factory"), + [ + param( + TransferLearningMode.INDEX_KERNEL, + gpytorch.kernels.MaternKernel(nu=2.5), + id="index_override+bare_gpytorch_matern", + ), + param( + TransferLearningMode.POSITIVE_INDEX_KERNEL, + MaternKernel(parameter_names=("x",)) + * IndexKernel(num_tasks=3, rank=3, parameter_names=("Task",)), + id="override+product_kernel", + ), + param( + TransferLearningMode.POSITIVE_INDEX_KERNEL, + ICMKernelFactory( + task_kernel_or_factory=IndexKernel( + num_tasks=3, rank=3, parameter_names=("Task",) + ) + ), + id="override+task_aware_factory", + ), + param( + TransferLearningMode.INDEX_KERNEL, + _gpytorch_returning_factory, + id="override+factory_returning_gpytorch", + ), + param( + TransferLearningMode.INDEX_KERNEL, + BayBEKernelFactory( + parameter_selector=TypeSelector((TaskParameter,), exclude=True) + ), + # On a searchspace without a substance parameter, the default factory + # yields a gpytorch kernel and cannot be combined with an override. + id="override+default_factory_with_selector_unsupported", + ), + ], +) +def test_resolve_kernel_dispatch_raises(monkeypatch, override_mode, kernel_or_factory): + """`_resolve_kernel` raises for inputs incompatible with an override. + + This covers raw gpytorch kernels, composite kernels, and task-aware factories. + """ + monkeypatch.setenv("BAYBE_DISABLE_CUSTOM_KERNEL_WARNING", "True") + + context = _make_dispatch_context(override_mode) + kwargs = ( + {} if kernel_or_factory is None else {"kernel_or_factory": kernel_or_factory} + ) + surrogate = GaussianProcessSurrogate(**kwargs) + + with pytest.raises(IncompatibleOverrideError): + surrogate._resolve_kernel(context) diff --git a/tests/test_kernels.py b/tests/test_kernels.py index 425fc25243..b8d54b7b17 100644 --- a/tests/test_kernels.py +++ b/tests/test_kernels.py @@ -13,7 +13,11 @@ from baybe.kernels.base import BasicKernel, Kernel from baybe.kernels.basic import IndexKernel, MaternKernel, RBFKernel from baybe.kernels.composite import AdditiveKernel, ProductKernel, ScaleKernel -from baybe.parameters import NumericalContinuousParameter +from baybe.parameters import ( + NumericalContinuousParameter, + NumericalDiscreteParameter, + TaskParameter, +) from baybe.searchspace.core import SearchSpace from tests.hypothesis_strategies.kernels import kernels @@ -223,3 +227,67 @@ def test_mul_constant_produces_constant_scale_kernel(left, right, searchspace): optimizer = torch.optim.SGD(gpytorch_kernel.parameters(), lr=0.1) optimizer.step() assert gpytorch_kernel.outputscale.item() == initial_outputscale + + +@pytest.fixture(name="task_searchspace") +def fixture_task_searchspace() -> SearchSpace: + """A search space with a numerical (``x``) and a task (``Task``) parameter.""" + return SearchSpace.from_product( + [ + NumericalDiscreteParameter("x", [1, 2, 3]), + TaskParameter("Task", ["A", "B", "C"]), + ] + ) + + +@pytest.mark.parametrize( + ("kernel", "expected"), + [ + param( + MaternKernel(parameter_names=("x", "Task")), + MaternKernel(parameter_names=("x",)), + id="named_kernel_drops_parameter", + ), + param( + MaternKernel(), + MaternKernel(parameter_names=("x",)), + id="unnamed_kernel_pins_to_remaining", + ), + param( + MaternKernel(parameter_names=("x",)), + MaternKernel(parameter_names=("x",)), + id="absent_parameter_leaves_kernel_unchanged", + ), + param( + IndexKernel(num_tasks=3, rank=3, parameter_names=("Task",)), + None, + id="sole_parameter_collapses_to_none", + ), + param( + ScaleKernel(MaternKernel()), + ScaleKernel(MaternKernel(parameter_names=("x",))), + id="scale_kernel_recurses_into_base", + ), + param( + ScaleKernel(IndexKernel(num_tasks=3, rank=3, parameter_names=("Task",))), + None, + id="scale_kernel_over_sole_parameter_collapses_to_none", + ), + ], +) +def test_without_parameter(kernel, expected, task_searchspace): + """Removing a parameter reduces basic and scaled kernels as expected.""" + assert kernel._without_parameter("Task", task_searchspace) == expected + + +@pytest.mark.parametrize( + "kernel", + [ + param(MaternKernel() * RBFKernel(), id="product_kernel"), + param(MaternKernel() + RBFKernel(), id="additive_kernel"), + ], +) +def test_without_parameter_unsupported(kernel, task_searchspace): + """Kernels that cannot be reduced unambiguously raise ``TypeError``.""" + with pytest.raises(TypeError, match="Cannot remove a parameter"): + kernel._without_parameter("Task", task_searchspace) diff --git a/tests/test_reduced_searchspace.py b/tests/test_reduced_searchspace.py index dd5b9d7778..98bbd4284c 100644 --- a/tests/test_reduced_searchspace.py +++ b/tests/test_reduced_searchspace.py @@ -82,6 +82,14 @@ def test_reduced_blocked_attributes(reduced_searchspace): getattr(reduced_searchspace, attr) +def test_reduced_blocked_attribute_exception_type(reduced_searchspace): + """Blocked attribute access raises the dedicated ``AttributeError`` subclass.""" + from baybe.exceptions import _UnsupportedSearchSpaceAttributeError + + with pytest.raises(_UnsupportedSearchSpaceAttributeError): + reduced_searchspace.get_comp_rep_parameter_indices("x") + + def test_reduced_repr(reduced_searchspace): """Verify that repr does not crash on a reduced search space.""" result = repr(reduced_searchspace)