Skip to content
Merged
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
137 changes: 131 additions & 6 deletions kernels/src/kernels/layer/layer.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,11 @@
from __future__ import annotations

import copy
import functools
import importlib
import inspect
import logging
from contextvars import ContextVar
from enum import Flag, auto
from inspect import Parameter, Signature
from pathlib import Path
Expand Down Expand Up @@ -305,9 +309,6 @@ def __str__(self) -> str:
return f"`{self._repo_id}` (revision: {commit}), layer `{self.layer_name}`)"


_CACHED_LAYER: dict[RepositoryProtocol, Type["nn.Module"]] = {}


def replace_kernel_forward_from_hub(cls, layer_name: str, condition: Callable[["nn.Module"], bool] | None = None):
"""
Function that prepares a layer class to use kernels from the Hugging Face Hub.
Expand Down Expand Up @@ -417,6 +418,47 @@ def decorator(ty):
return decorator


class CompileableContextVar:

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

it's needed for torch compile but i think there are works to make it native within dynamo, still need it for BC either way

"""
Inspired by transformers `output_capturing.py` ensuring compatibility with torch compile for older torch versions.

The main difference is the call order, i.e. nested scopes and `get` before `set`.
"""

def __init__(self, name):
self.context_var = ContextVar(name, default=None)
self.global_var = None
self.uses_global_var = False

def get(self):
import torch

if self.uses_global_var or torch.compiler.is_compiling():
return self.global_var
return self.context_var.get()

def set(self, value):
import torch

if torch.compiler.is_compiling():
previous = self.global_var
self.global_var = value
self.uses_global_var = True
return previous

return self.context_var.set(value)

def reset(self, token):
if self.uses_global_var:
self.global_var = token
self.uses_global_var = token is not None
else:
self.context_var.reset(token)


_ACTIVE_KERNEL_FUNCS = CompileableContextVar("_ACTIVE_KERNEL_FUNCS")

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

also let me know whether to move things to other files as im not super familiar with the repo structure I thought its simpler to just add it directly for now



def use_kernelized_func(*args: Callable):
"""
This decorator attaches the target function within the module as a plain
Expand Down Expand Up @@ -469,19 +511,38 @@ def decorator(cls):
)

orig_init = cls.__init__
orig_forward = cls.forward

def new_init(self, *args, **kwargs):
orig_init(self, *args, **kwargs)

# Register new function as non-submodule within the modules dict
# Give each model instance its own kernel wrappers
hidden_kernels = self.__dict__.setdefault("_kernel_funcs", {})
for fn in decorator_args:
name = getattr(fn, "__name__", None) or getattr(fn, "kernel_layer_name", None)
assert name is not None

hidden_kernels[name] = fn
hidden_kernels[name] = type(fn)()

@functools.wraps(orig_forward)
def new_forward(self, *args, **kwargs):
# Route global function calls to this model's private wrapper
active_kernel_funcs = dict(_ACTIVE_KERNEL_FUNCS.get() or {})

for fn in decorator_args:
name = getattr(fn, "__name__", None) or getattr(fn, "kernel_layer_name", None)
assert name is not None

active_kernel_funcs[id(fn)] = self._kernel_funcs[name]

token = _ACTIVE_KERNEL_FUNCS.set(active_kernel_funcs)
try:
return orig_forward(self, *args, **kwargs)
finally:
_ACTIVE_KERNEL_FUNCS.reset(token)

cls.__init__ = new_init
cls.forward = new_forward
return cls

return decorator
Expand Down Expand Up @@ -599,8 +660,9 @@ def _validate_layer(*, check_cls, cls, repo: RepositoryProtocol):

# ... or predefined member variables.
torch_module_members = {name for name, _ in inspect.getmembers(nn.Module)}
func_module_members = {"__copy__", "__deepcopy__"} # Separate overrides needed for `Func`
cls_members = {name for name, _ in inspect.getmembers(cls)}
difference = cls_members - torch_module_members
difference = cls_members - torch_module_members - func_module_members
# verify if : difference ⊄ {"can_torch_compile", "has_backward"}
if not difference <= {"can_torch_compile", "has_backward"}:
raise TypeError(f"{repo} must not contain additional members compared to `{check_cls.__name__}`.")
Expand Down Expand Up @@ -681,6 +743,9 @@ def _validate_layer_has_mode(
return True


_CACHED_LAYER: dict[RepositoryProtocol, Type["nn.Module"]] = {}


def _get_layer_memoize(repo: RepositoryProtocol, module_class: Type["nn.Module"]) -> Type["nn.Module"]:
layer = _CACHED_LAYER.get(repo, None)
if layer is not None:
Expand All @@ -693,6 +758,12 @@ def _get_layer_memoize(repo: RepositoryProtocol, module_class: Type["nn.Module"]
return layer


def _rebuild_kernel_func(module_name: str, func_name: str):
"""For pickle kept outside as base function to rebuild its own local function"""
module = importlib.import_module(module_name)
return type(getattr(module, func_name))()


def _create_func_module(func: Callable) -> Type["nn.Module"]:
from torch import nn

Expand All @@ -702,8 +773,60 @@ class Func(nn.Module):
has_backward = getattr(func, "has_backward", True)

def forward(self, *args, **kwargs):
# Dispatch global calls to the active model's private wrapper
if (active_kernel_funcs := _ACTIVE_KERNEL_FUNCS.get()) is not None:
if (kernel_func := active_kernel_funcs.get(id(self))) is not None:
return kernel_func(*args, **kwargs)

return func(*args, **kwargs)

def __copy__(self):
result = type(self)()
result.__dict__.update(self.__dict__)

if isinstance(forward := self.__dict__.get("forward"), MethodType):
# Rebind fwd set by `kernelize` to the copy
if forward.__self__ is self:
result.forward = MethodType(forward.__func__, result)

return result

def __deepcopy__(self, memo):
result = type(self)()
memo[id(self)] = result
result.__dict__.update(copy.deepcopy(self.__dict__, memo))
return result

def __setstate__(self, state):
state, forward_func = state
self.__dict__.update(state)

if forward_func is not None:
# Rebind fwd set by `kernelize` to the restored instance
self.forward = MethodType(forward_func, self)

def __reduce__(self):
module = importlib.import_module(func.__module__)

# Global case => just resolve by name
if getattr(module, func.__name__, None) is self:
return func.__name__

# The own private wrapper case => rebuild a separate instance and restore its state
forward_func = None
state = self.__dict__.copy()
if isinstance(forward := state.get("forward"), MethodType):
if forward.__self__ is self:
# Store the function separately and rebind it on restore
forward_func = forward.__func__
del state["forward"]

return (
_rebuild_kernel_func,
(func.__module__, func.__name__),
(state, forward_func),
)
Comment on lines +800 to +828

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

pickle is the complex case as this is just due to pickle 😅 I think we can keep it since some do use pickle but they should use safetensors


# Use function signature with args prepended by self to support
# module validation.
func_sig = inspect.signature(func)
Expand All @@ -713,5 +836,7 @@ def forward(self, *args, **kwargs):
parameters=new_args,
return_annotation=func_sig.return_annotation,
)
# pickle needs to resolve by its original function's module, also see `__reduce__`
Func.__module__ = func.__module__

return Func
123 changes: 122 additions & 1 deletion kernels/tests/test_func.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,8 @@
import copy
import logging
import pickle
from pathlib import Path
from types import MethodType

import pytest
import torch
Expand All @@ -21,7 +24,18 @@
from kernels.layer.func import LockedFuncRepository


# A function + layer that we can map arbitrary functions to for testing.
# Base modules used as a replacement forward in tests
class AddOne(nn.Module):
def forward(self, x):
return x + 1


class TimesThree(nn.Module):
def forward(self, x):
return x * 3


# Functions + layers used to test function kernelization
@use_kernel_forward_from_hub("surprise_me")
def surprise_me(x: torch.Tensor):
return x
Expand All @@ -33,6 +47,29 @@ def forward(self, x: torch.Tensor):
return surprise_me(x)


@use_kernel_forward_from_hub("double_me")
def double_me(x: torch.Tensor):
return x * 2


@use_kernelized_func(double_me)
class Inner(nn.Module):
def forward(self, x):
return double_me(x)


# To check nested modules
@use_kernelized_func(surprise_me)
class Outer(nn.Module):
def __init__(self):
super().__init__()
self.inner = Inner()

def forward(self, x):
# The second call also verifies that Inner restored the outer context.
return surprise_me(self.inner(surprise_me(x)))


def test_decorator():
@use_kernel_forward_from_hub("identity_func")
def identity(x):
Expand Down Expand Up @@ -181,3 +218,87 @@ def forward(self, x: torch.Tensor):
def _silu_and_mul(x: torch.Tensor) -> torch.Tensor:
d = x.shape[-1] // 2
return F.silu(x[..., :d]) * x[..., d:]


# Imitates the kernels exchange
def _bind_forward(wrapper, forward):
wrapper.forward = MethodType(forward, wrapper)


def test_kernel_func_is_per_instance():
a, b = SurpriseMe(), SurpriseMe()

a_func = a._kernel_funcs["surprise_me"]
b_func = b._kernel_funcs["surprise_me"]

assert a_func is not b_func
assert a_func is not surprise_me
assert b_func is not surprise_me

_bind_forward(a_func, AddOne.forward)

x = torch.arange(4).float()
torch.testing.assert_close(a(x), x + 1)
torch.testing.assert_close(b(x), x)


@pytest.mark.parametrize("copy_fn", [copy.copy, copy.deepcopy])
def test_kernel_func_copy_is_independent(copy_fn):
model = SurpriseMe()
kernel_func = model._kernel_funcs["surprise_me"]
_bind_forward(kernel_func, AddOne.forward)

copied = copy_fn(kernel_func)

assert copied is not kernel_func
assert copied is not surprise_me
assert copied.__dict__["forward"].__self__ is copied

x = torch.arange(4).float()
torch.testing.assert_close(copied(x), x + 1)
torch.testing.assert_close(kernel_func(x), x + 1)


@pytest.mark.parametrize(
"restore_fn",
[
copy.deepcopy,
lambda model: pickle.loads(pickle.dumps(model)),
],
)
def test_kernel_func_serialization_is_independent(restore_fn):
model = SurpriseMe()
_bind_forward(model._kernel_funcs["surprise_me"], AddOne.forward)

restored = restore_fn(model)

assert restored._kernel_funcs["surprise_me"] is not model._kernel_funcs["surprise_me"]
assert restored._kernel_funcs["surprise_me"] is not surprise_me
assert restored._kernel_funcs["surprise_me"].__dict__["forward"].__self__ is restored._kernel_funcs["surprise_me"]

x = torch.arange(4).float()
torch.testing.assert_close(restored(x), x + 1)

# Resetting the restored model must not affect the original
with use_kernel_mapping({"surprise_me": {}}, inherit_mapping=False):
kernelize(restored, device="cpu", mode=Mode.INFERENCE)

torch.testing.assert_close(restored(x), x)
torch.testing.assert_close(model(x), x + 1)


@pytest.mark.parametrize("compile", [False, True])
def test_kernel_func_nested_dispatch(compile):
model = Outer()

_bind_forward(model._kernel_funcs["surprise_me"], AddOne.forward)
_bind_forward(model.inner._kernel_funcs["double_me"], TimesThree.forward)

if compile:
# We only need to know whether it's safe around get/set so the backend is not relevant
model = torch.compile(model, backend="eager", fullgraph=True)

x = torch.tensor(1.0)

# Outer (+1) -> Inner (*3) -> Outer (+1): 1 -> 2 -> 6 -> 7
torch.testing.assert_close(model(x), torch.tensor(7.0))
Loading