diff --git a/kernels/src/kernels/__init__.py b/kernels/src/kernels/__init__.py index 7d9bc7d0..b0fb477d 100644 --- a/kernels/src/kernels/__init__.py +++ b/kernels/src/kernels/__init__.py @@ -13,6 +13,7 @@ CUDAProperties, Device, FuncRepository, + KernelizeFallback, LayerRepository, LocalFuncRepository, LocalLayerRepository, @@ -48,6 +49,7 @@ "Benchmark", "CUDAProperties", "Device", + "KernelizeFallback", "ROCMProperties", "FuncRepository", "LayerRepository", diff --git a/kernels/src/kernels/layer/__init__.py b/kernels/src/kernels/layer/__init__.py index 217e3d0f..a46f23be 100644 --- a/kernels/src/kernels/layer/__init__.py +++ b/kernels/src/kernels/layer/__init__.py @@ -11,6 +11,7 @@ use_kernel_mapping, ) from .layer import ( + KernelizeFallback, LayerRepository, LocalLayerRepository, LockedLayerRepository, @@ -23,6 +24,7 @@ __all__ = [ "CUDAProperties", "Device", + "KernelizeFallback", "ROCMProperties", "FuncRepository", "LayerRepository", diff --git a/kernels/src/kernels/layer/kernelize.py b/kernels/src/kernels/layer/kernelize.py index 11188236..a06d4df2 100644 --- a/kernels/src/kernels/layer/kernelize.py +++ b/kernels/src/kernels/layer/kernelize.py @@ -6,7 +6,7 @@ from .device import Device from .globals import _KERNEL_MAPPING -from .layer import kernelize_layer +from .layer import KernelizeFallback, kernelize_layer from .mode import Mode from .repos import DeviceRepos, RepositoryProtocol @@ -180,7 +180,7 @@ def kernelize( *, mode: Mode, device: str | "torch.device" | None = None, - use_fallback: bool = True, + use_fallback: bool | KernelizeFallback = True, ): """ Replace layer forward methods with optimized kernel implementations. @@ -197,9 +197,10 @@ def kernelize( device (`Union[str, torch.device]`, *optional*): The device type to load kernels for. Supported device types are: "cuda", "mps", "npu", "rocm", "tpu", "xpu". The device type will be inferred from the model parameters when not provided. - use_fallback (`bool`, *optional*, defaults to `True`): - Whether to use the original forward method of modules when no compatible kernel could be found. - If set to `False`, an exception will be raised in such cases. + use_fallback (`bool | KernelizeFallback`, *optional*, defaults to `True`): + Cases in which to use the original forward method. `True` is equivalent to `KernelizeFallback.DEFAULT`: it allows + fallback when no compatible mapping or mode exists. `False` is equivalent to `KernelizeFallback.NONE` and raises + instead. `KernelizeFallback.ALL` also falls back when kernel loading fails. Returns: `nn.Module`: The kernelized model with optimized kernel implementations. @@ -259,6 +260,9 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: assert isinstance(device_type, Device) + if isinstance(use_fallback, bool): + use_fallback = KernelizeFallback.DEFAULT if use_fallback else KernelizeFallback.NONE + for _, module in model.named_modules(): module_class = type(module) diff --git a/kernels/src/kernels/layer/layer.py b/kernels/src/kernels/layer/layer.py index c02d5946..23b5a138 100644 --- a/kernels/src/kernels/layer/layer.py +++ b/kernels/src/kernels/layer/layer.py @@ -2,6 +2,7 @@ import inspect import logging +from enum import Flag, auto from inspect import Parameter, Signature from pathlib import Path from types import MethodType, ModuleType @@ -34,6 +35,30 @@ logger = logging.getLogger(__name__) +class KernelizeFallback(Flag): + """Cases in which kernelization may keep the original layer forward. + + - `NONE`: Raise instead of falling back. + - `NO_LAYER`: No kernel mapping exists for the layer. + - `NO_DEVICE`: The layer has no mapping for the requested device type. + - `NO_COMPATIBLE_PROPERTIES`: No mapping matches the device properties. + - `NO_COMPATIBLE_MODE`: No repository or loaded kernel supports the requested mode. + - `CANNOT_LOAD`: Kernel loading fails. + - `DEFAULT`: Fall back for missing mappings, properties, or mode; raise on a loading error. + - `ALL`: Fall back in every case above, including missing files during loading. + """ + + NONE = 0 + NO_LAYER = auto() + NO_DEVICE = auto() + NO_COMPATIBLE_PROPERTIES = auto() + NO_COMPATIBLE_MODE = auto() + CANNOT_LOAD = auto() + + DEFAULT = NO_LAYER | NO_DEVICE | NO_COMPATIBLE_PROPERTIES | NO_COMPATIBLE_MODE + ALL = DEFAULT | CANNOT_LOAD + + class LayerRepositoryProtocol(RepositoryProtocol, Protocol): @property def layer_name(self) -> str: ... @@ -462,7 +487,7 @@ def new_init(self, *args, **kwargs): return decorator -def kernelize_layer(module: "nn.Module", *, mode: Mode, device_type: Device, use_fallback): +def kernelize_layer(module: "nn.Module", *, mode: Mode, device_type: Device, use_fallback: KernelizeFallback): module_class = type(module) layer_name = module_class.kernel_layer_name # type: ignore[attr-defined] @@ -479,7 +504,7 @@ def kernelize_layer(module: "nn.Module", *, mode: Mode, device_type: Device, use f"Check if the layer name matches one of the kernels in the mapping or add the kernel " f"you want to use to the mapping. Defaulting to original forward implementation." ) - if not use_fallback: + if KernelizeFallback.NO_LAYER not in use_fallback: raise ValueError(f"No layer mapping for `{layer_name}`") _replace_forward(module, module_class) return @@ -488,7 +513,7 @@ def kernelize_layer(module: "nn.Module", *, mode: Mode, device_type: Device, use property_repos = kernel.get(device_type.type) if property_repos is None: - if not use_fallback: + if KernelizeFallback.NO_DEVICE not in use_fallback: raise ValueError(f"No layer mapping for `{layer_name}` with device type `{device_type}`") _replace_forward(module, module_class) return @@ -496,7 +521,7 @@ def kernelize_layer(module: "nn.Module", *, mode: Mode, device_type: Device, use repos = property_repos.repos if repos is None: - if not use_fallback: + if KernelizeFallback.NO_COMPATIBLE_PROPERTIES not in use_fallback: raise ValueError(f"No layer mapping for `{layer_name}` device `{device_type}` with the right properties") _replace_forward(module, module_class) return @@ -507,7 +532,7 @@ def kernelize_layer(module: "nn.Module", *, mode: Mode, device_type: Device, use ) if repo_with_mode is None: - if not use_fallback: + if KernelizeFallback.NO_COMPATIBLE_MODE not in use_fallback: raise ValueError(f"No repository for `{layer_name}` for configuration mode={mode}") _replace_forward(module, module_class) return @@ -517,7 +542,18 @@ def kernelize_layer(module: "nn.Module", *, mode: Mode, device_type: Device, use logging.info(f"Using function/layer from repo {repo}") logging.debug(f"kernelize mode: {mode}, repo mode: {repo_mode}") - layer = _get_layer_memoize(repo, module_class) + try: + layer = _get_layer_memoize(repo, module_class) + except FileNotFoundError: + if KernelizeFallback.CANNOT_LOAD not in use_fallback: + raise + logger.info( + "Kernel for layer `%s` on %s could not be loaded; using the original forward.", + layer_name, + device_type.type, + ) + _replace_forward(module, module_class) + return # Ideally we would do validation on the mapping where we check that # e.g. if a repo class is registered for TRAINING | TORCH_COMPILE, @@ -590,7 +626,7 @@ def _conditionally_replace_forward( module: "nn.Module", layer: Type["nn.Module"], mode: Mode, - use_fallback: bool, + use_fallback: KernelizeFallback, ): module_class = type(module) @@ -603,7 +639,7 @@ def _conditionally_replace_forward( needs_fallback_for_backward = Mode.TRAINING in mode and not getattr(layer, "has_backward", True) if needs_fallback_for_compile or needs_fallback_for_backward: - if use_fallback: + if KernelizeFallback.NO_COMPATIBLE_MODE in use_fallback: if needs_fallback_for_compile: logging.info("Layer does not support torch.compile, using fallback") if needs_fallback_for_backward: diff --git a/kernels/tests/test_layer.py b/kernels/tests/test_layer.py index 41186af6..3ad253b4 100644 --- a/kernels/tests/test_layer.py +++ b/kernels/tests/test_layer.py @@ -13,6 +13,7 @@ CUDAProperties, Device, FuncRepository, + KernelizeFallback, LayerRepository, LocalLayerRepository, Mode, @@ -425,6 +426,34 @@ class SiluAndMulWithKernelFallback(SiluAndMul): kernelize(silu_and_mul, device="cuda", mode=Mode.INFERENCE) +def test_missing_kernel_build_falls_back(monkeypatch, caplog): + repo = LayerRepository( + repo_id="kernels-test/no-compatible-build", + layer_name="SiluAndMul", + version=1, + ) + + def unavailable_build(): + raise FileNotFoundError("Cannot find a build variant for this system") + + monkeypatch.setattr(repo, "load", unavailable_build) + layer = SiluAndMulWithKernel() + mapping = {"SiluAndMul": {"cuda": repo}} + + with ( + use_kernel_mapping(mapping, inherit_mapping=False), + caplog.at_level(logging.INFO, logger="kernels.layer.layer"), + ): + kernelize(layer, device="cuda", mode=Mode.INFERENCE, use_fallback=KernelizeFallback.ALL) + assert "Kernel for layer `SiluAndMul` on cuda could not be loaded" in caplog.text + + with use_kernel_mapping(mapping, inherit_mapping=False): + with pytest.raises(FileNotFoundError, match="Cannot find a build variant"): + kernelize(layer, device="cuda", mode=Mode.INFERENCE) + with pytest.raises(FileNotFoundError, match="Cannot find a build variant"): + kernelize(layer, device="cuda", mode=Mode.INFERENCE, use_fallback=False) + + def test_kernel_condition_skips_kernelization(caplog): @use_kernel_forward_from_hub("SiluAndMulNonExisting", condition=lambda module: False) class SiluAndMulConditionSkipped(SiluAndMul):