From df601b432bc23eecfb69f9e176ff53ece78f9e0f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dani=C3=ABl=20de=20Kok?= Date: Thu, 24 Sep 2026 14:12:08 +0000 Subject: [PATCH] kernels: add support for layer selectors Until now, one could only register layers by device (+ capability) and the kernelization mode. It can be useful to select kernels in a more fine-grained manner. For instance, the user might want to select the best kernel based on input shapes, available device memory, and other properties. This change makes it possible to register a selector for a kernel, where the selector is a user-provided function that determines the layer repository to use based on the model, device, and module being kernelized. For example: ```python def select_silu_and_mul(module, *, device_type, mode): if device_type.type != "cuda": return None repo = LayerRepository( repo_id="kernels-community/activation", layer_name="SiluAndMul", version=1, ) return repo, Mode.FALLBACK with use_kernel_mapping({"SiluAndMul": select_silu_and_mul}): model = kernelize(model, mode=Mode.TRAINING | Mode.TORCH_COMPILE, device="cuda") ``` --- docs/source/api/layers.md | 4 + kernels/src/kernels/__init__.py | 2 + kernels/src/kernels/layer/globals.py | 6 +- kernels/src/kernels/layer/kernelize.py | 116 +++++++-- kernels/src/kernels/layer/layer.py | 41 ++-- kernels/src/kernels/layer/repos.py | 38 ++- kernels/tests/test_layer.py | 325 ++++++++++++++++++++++++- 7 files changed, 493 insertions(+), 39 deletions(-) diff --git a/docs/source/api/layers.md b/docs/source/api/layers.md index eff53dadc..556ce5fbd 100644 --- a/docs/source/api/layers.md +++ b/docs/source/api/layers.md @@ -75,3 +75,7 @@ ### LockedLayerRepository [[autodoc]] kernels.LockedLayerRepository + +### KernelLayerSelectorProtocol + +[[autodoc]] kernels.KernelLayerSelectorProtocol diff --git a/kernels/src/kernels/__init__.py b/kernels/src/kernels/__init__.py index 7d9bc7d0e..b2eae820e 100644 --- a/kernels/src/kernels/__init__.py +++ b/kernels/src/kernels/__init__.py @@ -28,6 +28,7 @@ use_kernel_mapping, use_kernelized_func, ) +from kernels.layer.repos import KernelLayerSelectorProtocol from kernels.load import ( get_kernel, get_local_kernel, @@ -50,6 +51,7 @@ "Device", "ROCMProperties", "FuncRepository", + "KernelLayerSelectorProtocol", "LayerRepository", "LoadedKernel", "LocalFuncRepository", diff --git a/kernels/src/kernels/layer/globals.py b/kernels/src/kernels/layer/globals.py index fd0a88d78..809f62df2 100644 --- a/kernels/src/kernels/layer/globals.py +++ b/kernels/src/kernels/layer/globals.py @@ -1,8 +1,10 @@ import os from contextvars import ContextVar -from .repos import DeviceRepos +from .repos import DeviceRepos, KernelLayerSelectorProtocol _DISABLE_KERNEL_MAPPING: bool = bool(int(os.environ.get("DISABLE_KERNEL_MAPPING", "0"))) -_KERNEL_MAPPING: ContextVar[dict[str, dict[str, DeviceRepos]]] = ContextVar("_KERNEL_MAPPING", default={}) +_KERNEL_MAPPING: ContextVar[dict[str, dict[str, DeviceRepos] | KernelLayerSelectorProtocol]] = ContextVar( + "_KERNEL_MAPPING", default={} +) diff --git a/kernels/src/kernels/layer/kernelize.py b/kernels/src/kernels/layer/kernelize.py index 11188236c..8f8117667 100644 --- a/kernels/src/kernels/layer/kernelize.py +++ b/kernels/src/kernels/layer/kernelize.py @@ -1,6 +1,8 @@ from __future__ import annotations +import inspect import logging +from collections.abc import Mapping from copy import deepcopy from typing import TYPE_CHECKING @@ -8,7 +10,7 @@ from .globals import _KERNEL_MAPPING from .layer import kernelize_layer from .mode import Mode -from .repos import DeviceRepos, RepositoryProtocol +from .repos import DeviceRepos, KernelLayerSelectorProtocol, RepositoryProtocol if TYPE_CHECKING: import torch @@ -23,7 +25,8 @@ def use_kernel_mapping( dict[ Device | str, RepositoryProtocol | dict[Mode, RepositoryProtocol], - ], + ] + | KernelLayerSelectorProtocol, ], *, inherit_mapping: bool = True, @@ -35,8 +38,9 @@ def use_kernel_mapping( kernel configurations for different parts of your code. Args: - mapping (`dict[str, dict[Union[Device, str], Union[LayerRepositoryProtocol, dict[Mode, LayerRepositoryProtocol]]]]`): - The kernel mapping to apply. Maps layer names to device-specific kernel configurations. + mapping (`dict[str, Union[dict[Union[Device, str], Union[RepositoryProtocol, dict[Mode, RepositoryProtocol]]], KernelLayerSelectorProtocol]]`): + The kernel mapping to apply. Maps layer names to device-specific kernel configurations, or to a + [`KernelLayerSelectorProtocol`] callable that selects the kernel for each module. inherit_mapping (`bool`, *optional*, defaults to `True`): When `True`, the current mapping will be extended by `mapping` inside the context. When `False`, only `mapping` is used inside the context. @@ -79,6 +83,20 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: model = kernelize(model, mode=Mode.TRAINING | Mode.TORCH_COMPILE, device="cuda") # Outside the context, original mappings are restored + + # Use a selector that chooses the kernel per module + def select_silu_and_mul(module, *, device_type, mode): + if device_type.type != "cuda": + return None + repo = LayerRepository( + repo_id="kernels-community/activation", + layer_name="SiluAndMul", + version=1, + ) + return repo, Mode.FALLBACK + + with use_kernel_mapping({"SiluAndMul": select_silu_and_mul}): + model = kernelize(model, mode=Mode.TRAINING | Mode.TORCH_COMPILE, device="cuda") ``` """ @@ -89,7 +107,12 @@ def __enter__(self): self.token = _KERNEL_MAPPING.set(deepcopy(_KERNEL_MAPPING.get())) else: self.token = _KERNEL_MAPPING.set({}) - register_kernel_mapping(mapping) + try: + register_kernel_mapping(mapping) + except BaseException: + # __exit__ is not called when __enter__ raises. + _KERNEL_MAPPING.reset(self.token) + raise def __exit__(self, exc_type, exc_value, traceback): _KERNEL_MAPPING.reset(self.token) @@ -103,7 +126,8 @@ def register_kernel_mapping( dict[ Device | str, RepositoryProtocol | dict[Mode, RepositoryProtocol], - ], + ] + | KernelLayerSelectorProtocol, ], inherit_mapping: bool = True, ): @@ -114,9 +138,11 @@ def register_kernel_mapping( depending on the device and mode. This should be used in conjunction with [`kernelize`]. Args: - mapping (`dict[str, dict[Union[Device, str], Union[RepositoryProtocol, dict[Mode, RepositoryProtocol]]]]`): + mapping (`dict[str, Union[dict[Union[Device, str], Union[RepositoryProtocol, dict[Mode, RepositoryProtocol]]], KernelLayerSelectorProtocol]]`): The kernel mapping to register globally. Maps layer names to device-specific kernels. - The mapping can specify different kernels for different modes (training, inference, etc.). + The mapping can specify different kernels for different modes (training, inference, etc.), + or map a layer name to a [`KernelLayerSelectorProtocol`] callable that selects the kernel for + each module. inherit_mapping (`bool`, *optional*, defaults to `True`): When `True`, the current mapping will be extended by `mapping`. When `False`, the existing mappings are erased before adding `mapping`. @@ -155,24 +181,80 @@ def register_kernel_mapping( } } register_kernel_mapping(advanced_mapping) + + # Mapping with a selector that chooses the kernel per module + def select_rms_norm(module, *, device_type, mode): + if device_type.type != "cuda": + return None + repo = LayerRepository( + repo_id="kernels-community/layer_norm", + layer_name="LlamaRMSNorm", + version=1, + ) + return repo, Mode.FALLBACK + + register_kernel_mapping({"LlamaRMSNorm": select_rms_norm}) ``` """ + # Validate the new mapping first, to avoid that the mapping is in an + # inconsistent state after a failure. + for layer_name, value in mapping.items(): + _validate_mapping_value(layer_name, value) + if not inherit_mapping: _KERNEL_MAPPING.set({}) # Merge with existing mappings. for new_kernel, new_device_repos in mapping.items(): - device_repo = _KERNEL_MAPPING.get().setdefault(new_kernel, {}) - for new_device, new_repo in new_device_repos.items(): - device = Device(type=new_device) if isinstance(new_device, str) else new_device + if not isinstance(new_device_repos, Mapping): + # Validated to be a kernel selector. + _KERNEL_MAPPING.get()[new_kernel] = new_device_repos + else: + device_repo = _KERNEL_MAPPING.get().get(new_kernel, None) + if not isinstance(device_repo, dict): + device_repo = {} + _KERNEL_MAPPING.get()[new_kernel] = device_repo + for new_device, new_repo in new_device_repos.items(): + device = Device(type=new_device) if isinstance(new_device, str) else new_device + + if isinstance(new_repo, dict): + kernel_options = new_repo + else: + kernel_options = {Mode.FALLBACK: new_repo} + + feature_repos = device_repo.setdefault(device.type, DeviceRepos.create_repo(device)) + feature_repos.insert(device, kernel_options) + + +def _validate_mapping_value(layer_name: str, value: object) -> None: + if isinstance(value, Mapping): + return + + # If we don't have a mapping, it must be a kernel selector. A kernel + # selector must comply with the protocol and also not be a class type. + if not isinstance(value, KernelLayerSelectorProtocol) or isinstance(value, type): + raise TypeError( + f"Kernel mapping for `{layer_name}` must be a dict of device-specific kernels " + f"or a kernel selector, got `{type(value).__name__}`" + ) - if isinstance(new_repo, dict): - kernel_options = new_repo - else: - kernel_options = {Mode.FALLBACK: new_repo} + # The protocol check only verifies that the right methods exist, not their + # signatures. So check that the selector has a valid signature. + try: + signature = inspect.signature(value) + except (TypeError, ValueError) as e: + raise TypeError( + f"Cannot inspect the signature of the kernel selector for `{layer_name}`, " + f"wrap it in a Python function that accepts `(module, *, device_type, mode)`: {e}" + ) from None - feature_repos = device_repo.setdefault(device.type, DeviceRepos.create_repo(device)) - feature_repos.insert(device, kernel_options) + try: + signature.bind(None, device_type=None, mode=None) + except TypeError as e: + raise TypeError( + f"Kernel selector for `{layer_name}` must accept `(module, *, device_type, mode)`, " + f"but has signature `{signature}`: {e}" + ) from None def kernelize( diff --git a/kernels/src/kernels/layer/layer.py b/kernels/src/kernels/layer/layer.py index c02d59462..8d2fed8ff 100644 --- a/kernels/src/kernels/layer/layer.py +++ b/kernels/src/kernels/layer/layer.py @@ -26,7 +26,7 @@ from .device import Device from .globals import _DISABLE_KERNEL_MAPPING, _KERNEL_MAPPING from .mode import Mode -from .repos import RepositoryProtocol, _select_repository +from .repos import KernelLayerSelectorProtocol, RepositoryProtocol, _select_repository if TYPE_CHECKING: from torch import nn @@ -484,27 +484,32 @@ def kernelize_layer(module: "nn.Module", *, mode: Mode, device_type: Device, use _replace_forward(module, module_class) return - # Get kernel options for the device - property_repos = kernel.get(device_type.type) + if isinstance(kernel, KernelLayerSelectorProtocol): + repo_with_mode = kernel(module, device_type=device_type, mode=mode) + else: + # Get kernel options for the device + property_repos = kernel.get(device_type.type) - if property_repos is None: - if not use_fallback: - raise ValueError(f"No layer mapping for `{layer_name}` with device type `{device_type}`") - _replace_forward(module, module_class) - return + if property_repos is None: + if not use_fallback: + raise ValueError(f"No layer mapping for `{layer_name}` with device type `{device_type}`") + _replace_forward(module, module_class) + return - repos = property_repos.repos + repos = property_repos.repos - if repos is None: - if not use_fallback: - raise ValueError(f"No layer mapping for `{layer_name}` device `{device_type}` with the right properties") - _replace_forward(module, module_class) - return + if repos is None: + if not use_fallback: + raise ValueError( + f"No layer mapping for `{layer_name}` device `{device_type}` with the right properties" + ) + _replace_forward(module, module_class) + return - repo_with_mode = _select_repository( - repos, - mode=mode, - ) + repo_with_mode = _select_repository( + repos, + mode=mode, + ) if repo_with_mode is None: if not use_fallback: diff --git a/kernels/src/kernels/layer/repos.py b/kernels/src/kernels/layer/repos.py index c53c833bb..75c19fa73 100644 --- a/kernels/src/kernels/layer/repos.py +++ b/kernels/src/kernels/layer/repos.py @@ -1,7 +1,7 @@ import sys from abc import ABC, abstractmethod from functools import lru_cache -from typing import TYPE_CHECKING, Protocol, Type +from typing import TYPE_CHECKING, Protocol, Type, runtime_checkable from ._interval_tree import IntervalTree from .device import CUDAProperties, Device, ROCMProperties @@ -266,6 +266,42 @@ def _select_repository( return None +@runtime_checkable +class KernelLayerSelectorProtocol(Protocol): + """ + Callable that selects the kernel repository for a layer at kernelization time. + + A selector can be used in a kernel mapping instead of the per-device dictionary. [`kernelize`] calls + the selector for every module with the mapped layer name, so the selector can base its decision on + the module instance itself, the device type, and the kernelization mode. + + Selectors should be stateless and must not hold references to modules; the selection should only depend + on the `module`, `device_type`, and `mode` arguments. + + An example can be found in the documentation of [`use_kernel_mapping`]. + """ + + def __call__( + self, module: "nn.Module", *, device_type: Device, mode: Mode + ) -> tuple[RepositoryProtocol, Mode] | None: + """ + Select the kernel repository for a module. + + Args: + module (`nn.Module`): + The module that is being kernelized. + device_type ([`Device`]): + The device that kernels are loaded for. + mode ([`Mode`]): + The mode that the module is kernelized for. + + Returns: + `tuple[RepositoryProtocol, Mode] | None`: The repository and the mode that it supports, or `None` + when no kernel should be used for the module. + """ + ... + + @lru_cache def _find_capability() -> int: import torch diff --git a/kernels/tests/test_layer.py b/kernels/tests/test_layer.py index 41186af62..6eef0a005 100644 --- a/kernels/tests/test_layer.py +++ b/kernels/tests/test_layer.py @@ -1,7 +1,8 @@ +import functools import logging import sys from contextlib import nullcontext -from types import SimpleNamespace +from types import MappingProxyType, SimpleNamespace import pytest import torch @@ -28,6 +29,7 @@ _KERNEL_MAPPING, _validate_layer, ) +from kernels.layer.repos import RepositoryProtocol @pytest.fixture @@ -1385,3 +1387,324 @@ def test_local_overrides_layer(monkeypatch, local_kernel_path): f"kernels-test/silu-and-mul={str(local_kernel_path)}:kernels-test/non-existing2=/non/existing", ) kernelize(model, device="cuda", mode=Mode.INFERENCE) + + +class _SelectorTestRepo: + """Repository that loads a local layer class, so that selector tests do not need the Hub.""" + + def __init__(self, layer): + self.layer = layer + + def load(self): + return self.layer + + def __repr__(self): + return f"_SelectorTestRepo({self.layer.__name__})" + + +class _ReLUSelectorKernel(nn.Module): + def forward(self, input: torch.Tensor) -> torch.Tensor: + return F.relu(input) + + +class _ReLUSelectorKernelNoBackward(nn.Module): + has_backward = False + + def forward(self, input: torch.Tensor) -> torch.Tensor: + return F.relu(input) + + +def test_selector_registered_as_is(): + def selector(module: nn.Module, *, device_type: Device, mode: Mode) -> tuple[RepositoryProtocol, Mode] | None: + return None + + mapping_before = _KERNEL_MAPPING.get() + + with use_kernel_mapping({}, inherit_mapping=False): + register_kernel_mapping({"ReLU": selector}) + assert _KERNEL_MAPPING.get()["ReLU"] is selector + + with use_kernel_mapping({"ReLU": selector}, inherit_mapping=False): + assert _KERNEL_MAPPING.get()["ReLU"] is selector + + assert _KERNEL_MAPPING.get() is mapping_before + + +def test_selector_receives_arguments(device): + repo = _SelectorTestRepo(_ReLUSelectorKernel) + calls = [] + + def selector(module: nn.Module, *, device_type: Device, mode: Mode) -> tuple[RepositoryProtocol, Mode] | None: + calls.append((module, device_type, mode)) + return repo, Mode.FALLBACK + + relu = ReLUWithKernel().to(device) + with use_kernel_mapping({"ReLU": selector}, inherit_mapping=False): + kernelize(relu, device=device, mode=Mode.INFERENCE) + + assert len(calls) == 1 + module, device_type, mode = calls[0] + assert module is relu + assert isinstance(device_type, Device) + assert device_type.type == device + assert mode == Mode.INFERENCE + + +def test_selector_kernel_is_used(device): + repo = _SelectorTestRepo(_ReLUSelectorKernel) + + def selector(module: nn.Module, *, device_type: Device, mode: Mode) -> tuple[RepositoryProtocol, Mode] | None: + return repo, Mode.FALLBACK + + relu = ReLUWithKernel().to(device) + with use_kernel_mapping({"ReLU": selector}, inherit_mapping=False): + kernelize(relu, device=device, mode=Mode.INFERENCE) + + X = torch.randn(10, 32, device=device) + Y = relu(X) + assert relu.n_calls == 0 + torch.testing.assert_close(Y, F.relu(X)) + + +def test_selector_per_instance(device): + repo = _SelectorTestRepo(_ReLUSelectorKernel) + + def selector(module: nn.Module, *, device_type: Device, mode: Mode) -> tuple[RepositoryProtocol, Mode] | None: + return (repo, Mode.FALLBACK) if getattr(module, "use_kernel", False) else None + + with_kernel = ReLUWithKernel().to(device) + with_kernel.use_kernel = True + without_kernel = ReLUWithKernel().to(device) + model = nn.Sequential(with_kernel, without_kernel) + + with use_kernel_mapping({"ReLU": selector}, inherit_mapping=False): + kernelize(model, device=device, mode=Mode.INFERENCE) + + model(torch.randn(10, 32, device=device)) + assert with_kernel.n_calls == 0 + assert without_kernel.n_calls == 1 + + +def test_selector_returns_none(device): + def selector(module: nn.Module, *, device_type: Device, mode: Mode) -> tuple[RepositoryProtocol, Mode] | None: + return None + + relu = ReLUWithKernel().to(device) + with use_kernel_mapping({"ReLU": selector}, inherit_mapping=False): + kernelize(relu, device=device, mode=Mode.INFERENCE) + + with pytest.raises(ValueError, match="No repository for `ReLU`"): + kernelize(relu, device=device, mode=Mode.INFERENCE, use_fallback=False) + + relu(torch.randn(10, 32, device=device)) + assert relu.n_calls == 1 + + +def test_selector_mode_validation(device): + repo = _SelectorTestRepo(_ReLUSelectorKernelNoBackward) + + def selector(module: nn.Module, *, device_type: Device, mode: Mode) -> tuple[RepositoryProtocol, Mode] | None: + return repo, Mode.TRAINING + + relu = ReLUWithKernel().to(device) + with use_kernel_mapping({"ReLU": selector}, inherit_mapping=False): + with pytest.raises(ValueError, match="does not support backward"): + kernelize(relu, device=device, mode=Mode.TRAINING) + + +def test_selector_and_dict_override_each_other(): + cpu_repo = _SelectorTestRepo(_ReLUSelectorKernel) + mps_repo = _SelectorTestRepo(_ReLUSelectorKernel) + + def selector(module: nn.Module, *, device_type: Device, mode: Mode) -> tuple[RepositoryProtocol, Mode] | None: + return None + + with use_kernel_mapping({"ReLU": {"cpu": cpu_repo}}, inherit_mapping=False): + register_kernel_mapping({"ReLU": selector}) + assert _KERNEL_MAPPING.get()["ReLU"] is selector + + register_kernel_mapping({"ReLU": {"mps": mps_repo}}) + device_repos = _KERNEL_MAPPING.get()["ReLU"] + assert isinstance(device_repos, dict) + # Device entries from before the selector was registered must not come back. + assert set(device_repos.keys()) == {"mps"} + assert device_repos["mps"].repos[Mode.FALLBACK] is mps_repo + + # Registering another dict merges with the existing device entries. + register_kernel_mapping({"ReLU": {"cpu": cpu_repo}}) + assert set(_KERNEL_MAPPING.get()["ReLU"].keys()) == {"cpu", "mps"} + + +def test_selector_nested_contexts(): + repo = _SelectorTestRepo(_ReLUSelectorKernel) + + def outer_selector( + module: nn.Module, *, device_type: Device, mode: Mode + ) -> tuple[RepositoryProtocol, Mode] | None: + return None + + def inner_selector( + module: nn.Module, *, device_type: Device, mode: Mode + ) -> tuple[RepositoryProtocol, Mode] | None: + return None + + with use_kernel_mapping({"ReLU": outer_selector}, inherit_mapping=False): + with use_kernel_mapping({"SiluAndMul": {"cpu": repo}}): + assert _KERNEL_MAPPING.get()["ReLU"] is outer_selector + + with use_kernel_mapping({"ReLU": inner_selector}): + assert _KERNEL_MAPPING.get()["ReLU"] is inner_selector + + assert _KERNEL_MAPPING.get()["ReLU"] is outer_selector + + with use_kernel_mapping({"SiluAndMul": {"cpu": repo}}, inherit_mapping=False): + assert "ReLU" not in _KERNEL_MAPPING.get() + + assert _KERNEL_MAPPING.get()["ReLU"] is outer_selector + + +def test_selector_callable_instance(device): + class Selector: + def __init__(self, repo: _SelectorTestRepo): + self.repo = repo + + def __call__( + self, module: nn.Module, *, device_type: Device, mode: Mode + ) -> tuple[RepositoryProtocol, Mode] | None: + return self.repo, Mode.FALLBACK + + relu = ReLUWithKernel().to(device) + with use_kernel_mapping({"ReLU": Selector(_SelectorTestRepo(_ReLUSelectorKernel))}, inherit_mapping=False): + # Entering an inheriting context deep-copies the mapping, including + # callable selector instances. The copied selector must still work. + with use_kernel_mapping({}): + kernelize(relu, device=device, mode=Mode.INFERENCE) + + relu(torch.randn(10, 32, device=device)) + assert relu.n_calls == 0 + + +@pytest.mark.parametrize( + "value", + [_SelectorTestRepo(_ReLUSelectorKernel), _ReLUSelectorKernel, 42], + ids=["bare-repo", "class", "int"], +) +def test_invalid_mapping_value_rejected(value): + def selector(module: nn.Module, *, device_type: Device, mode: Mode) -> tuple[RepositoryProtocol, Mode] | None: + return None + + match = "must be a dict of device-specific kernels or a kernel selector" + mapping_before = _KERNEL_MAPPING.get() + + with pytest.raises(TypeError, match=match): + with use_kernel_mapping({"ReLU": value}): + pass + assert _KERNEL_MAPPING.get() is mapping_before + + with use_kernel_mapping({}, inherit_mapping=False): + # Valid entries must not be registered when another entry is invalid. + with pytest.raises(TypeError, match=match): + register_kernel_mapping({"SiluAndMul": selector, "ReLU": value}) + assert _KERNEL_MAPPING.get() == {} + + +def _selector_missing_mode(module: nn.Module, *, device_type: Device) -> tuple[RepositoryProtocol, Mode] | None: + return None + + +def _selector_wrong_name(module: nn.Module, *, device: Device, mode: Mode) -> tuple[RepositoryProtocol, Mode] | None: + return None + + +def _selector_keyword_only_module( + *, module: nn.Module, device_type: Device, mode: Mode +) -> tuple[RepositoryProtocol, Mode] | None: + return None + + +def _selector_positional_only_device_type( + module: nn.Module, device_type: Device, /, mode: Mode +) -> tuple[RepositoryProtocol, Mode] | None: + return None + + +@pytest.mark.parametrize( + "selector", + [ + _selector_missing_mode, + _selector_wrong_name, + _selector_keyword_only_module, + _selector_positional_only_device_type, + ], +) +def test_selector_invalid_signature_rejected(selector): + with use_kernel_mapping({}, inherit_mapping=False): + with pytest.raises(TypeError, match=r"must accept `\(module, \*, device_type, mode\)`"): + register_kernel_mapping({"ReLU": selector}) + assert "ReLU" not in _KERNEL_MAPPING.get() + + +class _SelectorWithInvalidSignature: + __signature__ = "not a signature" + + def __call__( + self, module: nn.Module, *, device_type: Device, mode: Mode + ) -> tuple[RepositoryProtocol, Mode] | None: + return None + + +@pytest.mark.parametrize( + "selector", + [max, _SelectorWithInvalidSignature()], + ids=["builtin-without-signature", "invalid-signature"], +) +def test_selector_uninspectable_signature_rejected(selector): + with use_kernel_mapping({}, inherit_mapping=False): + with pytest.raises(TypeError, match="Cannot inspect the signature of the kernel selector for `ReLU`"): + register_kernel_mapping({"ReLU": selector}) + assert "ReLU" not in _KERNEL_MAPPING.get() + + +def test_selector_flexible_signatures_accepted(): + def with_kwargs(module: nn.Module, **kwargs: object) -> tuple[RepositoryProtocol, Mode] | None: + return None + + def with_extra_default( + module: nn.Module, *, device_type: Device, mode: Mode, verbose: bool = False + ) -> tuple[RepositoryProtocol, Mode] | None: + return None + + def positional_or_keyword( + module: nn.Module, device_type: Device, mode: Mode + ) -> tuple[RepositoryProtocol, Mode] | None: + return None + + class Selector: + def __call__( + self, module: nn.Module, *, device_type: Device, mode: Mode + ) -> tuple[RepositoryProtocol, Mode] | None: + return None + + def with_prefix( + prefix: str, module: nn.Module, *, device_type: Device, mode: Mode + ) -> tuple[RepositoryProtocol, Mode] | None: + return None + + selectors = [ + with_kwargs, + with_extra_default, + positional_or_keyword, + Selector(), + functools.partial(with_prefix, "relu"), + ] + + for selector in selectors: + with use_kernel_mapping({"ReLU": selector}, inherit_mapping=False): + assert _KERNEL_MAPPING.get()["ReLU"] is selector + + +def test_mapping_accepts_non_dict_mapping(): + repo = _SelectorTestRepo(_ReLUSelectorKernel) + + with use_kernel_mapping({"ReLU": MappingProxyType({"cpu": repo})}, inherit_mapping=False): + assert _KERNEL_MAPPING.get()["ReLU"]["cpu"].repos[Mode.FALLBACK] is repo