Skip to content
Draft
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
4 changes: 4 additions & 0 deletions docs/source/api/layers.md
Original file line number Diff line number Diff line change
Expand Up @@ -75,3 +75,7 @@
### LockedLayerRepository

[[autodoc]] kernels.LockedLayerRepository

### KernelLayerSelectorProtocol

[[autodoc]] kernels.KernelLayerSelectorProtocol
2 changes: 2 additions & 0 deletions kernels/src/kernels/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -50,6 +51,7 @@
"Device",
"ROCMProperties",
"FuncRepository",
"KernelLayerSelectorProtocol",
"LayerRepository",
"LoadedKernel",
"LocalFuncRepository",
Expand Down
6 changes: 4 additions & 2 deletions kernels/src/kernels/layer/globals.py
Original file line number Diff line number Diff line change
@@ -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={}
)
116 changes: 99 additions & 17 deletions kernels/src/kernels/layer/kernelize.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,16 @@
from __future__ import annotations

import inspect
import logging
from collections.abc import Mapping
from copy import deepcopy
from typing import TYPE_CHECKING

from .device import Device
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
Expand All @@ -23,7 +25,8 @@ def use_kernel_mapping(
dict[
Device | str,
RepositoryProtocol | dict[Mode, RepositoryProtocol],
],
]
| KernelLayerSelectorProtocol,
],
*,
inherit_mapping: bool = True,
Expand All @@ -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.
Expand Down Expand Up @@ -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")
```
"""

Expand All @@ -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)
Expand All @@ -103,7 +126,8 @@ def register_kernel_mapping(
dict[
Device | str,
RepositoryProtocol | dict[Mode, RepositoryProtocol],
],
]
| KernelLayerSelectorProtocol,
],
inherit_mapping: bool = True,
):
Expand All @@ -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`.
Expand Down Expand Up @@ -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(
Expand Down
41 changes: 23 additions & 18 deletions kernels/src/kernels/layer/layer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
38 changes: 37 additions & 1 deletion kernels/src/kernels/layer/repos.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand Down
Loading
Loading