Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions kernels/src/kernels/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
CUDAProperties,
Device,
FuncRepository,
KernelizeFallback,
LayerRepository,
LocalFuncRepository,
LocalLayerRepository,
Expand Down Expand Up @@ -48,6 +49,7 @@
"Benchmark",
"CUDAProperties",
"Device",
"KernelizeFallback",
"ROCMProperties",
"FuncRepository",
"LayerRepository",
Expand Down
2 changes: 2 additions & 0 deletions kernels/src/kernels/layer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
use_kernel_mapping,
)
from .layer import (
KernelizeFallback,
LayerRepository,
LocalLayerRepository,
LockedLayerRepository,
Expand All @@ -23,6 +24,7 @@
__all__ = [
"CUDAProperties",
"Device",
"KernelizeFallback",
"ROCMProperties",
"FuncRepository",
"LayerRepository",
Expand Down
14 changes: 9 additions & 5 deletions kernels/src/kernels/layer/kernelize.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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.
Expand All @@ -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.
Expand Down Expand Up @@ -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)

Expand Down
52 changes: 44 additions & 8 deletions kernels/src/kernels/layer/layer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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: ...
Expand Down Expand Up @@ -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]

Expand All @@ -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
Expand All @@ -488,15 +513,15 @@ 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

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
Expand All @@ -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
Expand All @@ -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)

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

In this case, since there is a layer/device/repo registered, it might be good to emit a logging.warning that there is no correct build variant.

@SunMarc SunMarc Sep 24, 2026 •

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

okay will do ! maybe logging.info instead ? I feel like it will be too verbose otherwise, especially in transformers

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,
Expand Down Expand Up @@ -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)

Expand All @@ -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:
Expand Down
29 changes: 29 additions & 0 deletions kernels/tests/test_layer.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
CUDAProperties,
Device,
FuncRepository,
KernelizeFallback,
LayerRepository,
LocalLayerRepository,
Mode,
Expand Down Expand Up @@ -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):
Expand Down
Loading