From 937cdd651e4f656e65053d027c71b51f1e1411ec Mon Sep 17 00:00:00 2001 From: Vincent Wang Date: Thu, 29 Feb 2024 23:03:57 +0800 Subject: [PATCH] [ORTMODULE] Support Register Custom Triton Kernel (#19690) Add support for registering custom Triton kernel function. --- .../python/training/ort_triton/__init__.py | 1 + .../python/training/ort_triton/triton_op_executor.py | 12 +++++++++++- 2 files changed, 12 insertions(+), 1 deletion(-) diff --git a/orttraining/orttraining/python/training/ort_triton/__init__.py b/orttraining/orttraining/python/training/ort_triton/__init__.py index fbb59d1354..5f2d0c62ff 100644 --- a/orttraining/orttraining/python/training/ort_triton/__init__.py +++ b/orttraining/orttraining/python/training/ort_triton/__init__.py @@ -9,6 +9,7 @@ from functools import wraps from onnxruntime.capi import _pybind_state as _C from .kernel import * # noqa: F403 +from .triton_op_executor import register_triton_kernel # noqa: F401 from .triton_op_executor import call_triton_by_name, call_triton_by_onnx, get_config diff --git a/orttraining/orttraining/python/training/ort_triton/triton_op_executor.py b/orttraining/orttraining/python/training/ort_triton/triton_op_executor.py index f16abc7125..e104ea13c5 100644 --- a/orttraining/orttraining/python/training/ort_triton/triton_op_executor.py +++ b/orttraining/orttraining/python/training/ort_triton/triton_op_executor.py @@ -23,6 +23,8 @@ from ._utils import gen_unique_name, next_power_of_2 _DEBUG_MODE = "ORTMODULE_TRITON_DEBUG" in os.environ and int(os.getenv("ORTMODULE_TRITON_DEBUG")) == 1 +_CUSTOM_KERNELS = dict() + @functools.lru_cache(None) def _gen_module_internal(sorted_graph: SortedGraph) -> Tuple[str, str, ModuleType]: @@ -113,7 +115,10 @@ def call_triton_by_name(func_name: str, *tensors, **kwargs): """ torch_tensors = [_from_dlpack(tensor) if tensor is not None else None for tensor in tensors] - func = getattr(sys.modules[".".join(__name__.split(".")[:-1])], func_name) + func = getattr(sys.modules[".".join(__name__.split(".")[:-1])], func_name, None) + if func is None: + func = _CUSTOM_KERNELS.get(func_name) + assert func is not None, f"Function {func_name} is not found in the registered kernels." output = func(*torch_tensors, **kwargs) if output is not None: if isinstance(output, tuple): @@ -138,3 +143,8 @@ def call_triton_by_onnx(onnx_key: int, onnx_str: bytes, *tensors): if isinstance(output, tuple): return tuple([to_dlpack(tensor) for tensor in output]) return to_dlpack(output) + + +def register_triton_kernel(fn): + _CUSTOM_KERNELS[fn.__name__] = fn + return fn