Repository navigation
[Func] Allow proper preservations/reference across copy and pickle
#873
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| 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 | ||
|
|
@@ -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. | ||
|
|
@@ -417,6 +418,47 @@ def decorator(ty): | |
| return decorator | ||
|
|
||
|
|
||
| class CompileableContextVar: | ||
| """ | ||
| 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") | ||
|
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 | ||
|
|
@@ -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 | ||
|
|
@@ -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__}`.") | ||
|
|
@@ -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: | ||
|
|
@@ -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 | ||
|
|
||
|
|
@@ -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
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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) | ||
|
|
@@ -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 | ||
There was a problem hiding this comment.
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