From 347c963d5cdbc180980d9a7758ca6aab46da5b52 Mon Sep 17 00:00:00 2001 From: mindest <30493312+mindest@users.noreply.github.com> Date: Tue, 11 Jul 2023 13:55:30 +0800 Subject: [PATCH] [ROCm] Add ROCm Triton TunableOp for GroupNorm (#16196) ### Description - Refactor existing Triton TunableOp-related code (based on work in #15862) - Add GroupNorm Triton implementation --- cmake/onnxruntime_compile_triton_kernel.cmake | 5 +- docs/ORT_Use_Trtion_Kernel.md | 104 ++++++++++++++++ docs/ORT_use_trtion_kernel.md | 98 ---------------- .../rocm/diffusion/group_norm_triton.cuh | 91 ++++++++++++++ .../rocm/diffusion/group_norm_triton.py | 104 ++++++++++++++++ .../rocm/diffusion/group_norm_tunable_op.h | 12 ++ .../core/providers/cuda/triton_kernel.cu | 111 ++++++++++-------- .../providers/rocm/math/softmax_triton.cuh | 14 +-- .../math/{softmax.py => softmax_triton.py} | 2 +- .../rocm/math/softmax_tunable_op.cuh | 16 +-- .../kernels/rocm/group_norm.cu | 55 +++++++++ .../kernel_explorer/kernels/rocm/softmax.cu | 2 +- tools/ci_build/compile_triton.py | 2 +- 13 files changed, 450 insertions(+), 166 deletions(-) create mode 100644 docs/ORT_Use_Trtion_Kernel.md delete mode 100644 docs/ORT_use_trtion_kernel.md create mode 100644 onnxruntime/contrib_ops/rocm/diffusion/group_norm_triton.cuh create mode 100644 onnxruntime/contrib_ops/rocm/diffusion/group_norm_triton.py rename onnxruntime/core/providers/rocm/math/{softmax.py => softmax_triton.py} (98%) diff --git a/cmake/onnxruntime_compile_triton_kernel.cmake b/cmake/onnxruntime_compile_triton_kernel.cmake index b9482ecaac..f59cc6de10 100644 --- a/cmake/onnxruntime_compile_triton_kernel.cmake +++ b/cmake/onnxruntime_compile_triton_kernel.cmake @@ -4,8 +4,9 @@ find_package(Python3 COMPONENTS Interpreter REQUIRED) # set all triton kernel ops that need to be compiled -set(triton_kernel_scripts - "onnxruntime/core/providers/rocm/math/softmax.py" +set(triton_kernel_scripts + "onnxruntime/core/providers/rocm/math/softmax_triton.py" + "onnxruntime/contrib_ops/rocm/diffusion/group_norm_triton.py" ) function(compile_triton_kernel out_triton_kernel_obj_file out_triton_kernel_header_dir) diff --git a/docs/ORT_Use_Trtion_Kernel.md b/docs/ORT_Use_Trtion_Kernel.md new file mode 100644 index 0000000000..6a38606ebf --- /dev/null +++ b/docs/ORT_Use_Trtion_Kernel.md @@ -0,0 +1,104 @@ +## Description + +In some scenarios, the Triton written kernels are more performant than CK or other handwritten kernels, so we implement a framework that enables onnxruntime to use these Triton written kernels. + +Here we use `softmax` op as an example to show how to integrate a Triton written kernel into onnxruntime CUDA/ROCm EP. + +### Write and compile Triton kernel + +We have implemented a softmax kernel using Triton at `onnxruntime/core/providers/rocm/math/softmax_triton.py` + +```python +@triton.jit +def softmax_kernel( + output_ptr, input_ptr, input_row_stride, output_row_stride, n_cols, + BLOCK_SIZE: tl.constexpr +): + # softmax implementations + ... + ... +``` + +This is a very simple implementation. The `n_cols` parameter should be smaller than `BLOCK_SIZE`. And `BLOCK_SIZE` MUST be determined at compile time. + +In order to support different input shapes, we compile multiple kernels with different `BLOCK_SIZE`s. + +Each kernel with different `BLOCK_SIZE` generates different `num_warps` and shared memory usage, we call them `metadata`, and these metadata are needed when launching kernels in onnxruntime. + +We develop a script `tools/ci_build/compile_triton.py` to compile kernel and generate metadata for kernel launching. + +To generate metadata for softmax, we need to add description info and implement a `get_function_table` function in `softmax_triton.py`: + +```python +# kernel dtype and BLOCK_SIZE to generate. +dtypes = ['fp32', 'fp16'] +blocks = [1024, 2048, 4096, 8192, 16384] +name_pattern = 'softmax_{}_{}' +sig_pattern = '*{},*{},i32,i32,i32' +group_pattern = 'softmax_{}' + +""" +SHOULD implement a function that returns a metadata list with format: + +function_table = [ + { + 'name': xx, + 'group': yy, + 'func': func, + 'sig': sig, + 'kwargs': kwargs + } +] + +The kwargs is a dict of {string: int} which is used for kernel constants. For example, BLOCK_SIZE of softmax. +""" + +def get_function_table(): + ... +``` + +When compiling onnxruntime with `--use_triton_kernel` flag, this softmax kernel will be compiled and combined into `libonnxruntime_providers_rocm.so` for ROCm or `libonnxruntime_providers_cuda.so` for CUDA. + +### onnxruntime C++ code modification + +To use the Triton kernels in onnxruntime, we need to implement a C++ op that calls these Triton kernels. + +Similar with CK, we implement a function that returns all possible Triton kernels, and the `TunableOp` will select the best one. + +```cpp +template +auto GetSoftmaxTritonOps() { + std::vector>>> ret; + auto group_name = GetSoftmaxTritonGroupName(); + // here use group_name to get all kernel with same group_name + // for example, 'softmax_fp16' represents a group of kernels with different BLOCK_SIZE for float16 softmax + auto *kernel_list = GetOrtTritonKernelByGroup(group_name); + if (kernel_list == nullptr) { + return ret; + } + + for (auto i : *kernel_list) { + // check params match + ... + } + return ret; +} +``` + +### Test + +Using kernel_explorer, we can test this softmax kernel like: + +```bash +export KERNEL_EXPLORER_BUILD_DIR= + +python onnxruntime/python/tools/kernel_explorer/kernels/softmax_test.py +``` + +and the result shows that `TunableOp` selects `softmax_fp16_2048` which is a Triton written kernel and better than others. + +```text +SoftmaxTunable float16 batch_count=1 softmax_elements=2048 is_log_softmax=0 4.27 us, 1.92 GB/s +softmax_fp16_2048 float16 batch_count=1 softmax_elements=2048 is_log_softmax=0 4.48 us, 1.83 GB/s +... +``` diff --git a/docs/ORT_use_trtion_kernel.md b/docs/ORT_use_trtion_kernel.md deleted file mode 100644 index 7c9aff4bb5..0000000000 --- a/docs/ORT_use_trtion_kernel.md +++ /dev/null @@ -1,98 +0,0 @@ -## Description -In some scenarios, the triton written kernels are more performant than CK or other handwritten kernels, so we implement a framework that onnxruntime can use these triton written kernels. - -Here we use `softmax` op as an example to show how to integrate a triton written kernel into onnxruntime cuda/rocm EP. - -#### Write and compile triton kernel -We have implemented a softmax kernel using triton at `onnxruntime/core/providers/rocm/math/softmax.py` - -``` -@triton.jit -def softmax_kernel( - output_ptr, input_ptr, input_row_stride, output_row_stride, n_cols, - BLOCK_SIZE: tl.constexpr -): - # softmax implementings - .... - .... -``` -This is a very simple implementation. The `n_cols` parameter should be smaller than BLOCK_SIZE. And BLOCK_SIZE MUST be determined at compile time. - -In order to support different input shape, we compile multiple kernels with different BLOCK_SIZE. - -Each kernel with different BLOCK_SIZE generates different num_warps and shared memory usage, we call them as `metadata`, and these metadata are needed when launching kernels in onnxruntime. - -We develop a script `tools/ci_build/compile_triton.py` to compile kernel and generate metadata for kernel launching. - -To generate metadta for softmax, it needs to add description info and implement a `get_function_table` function in `softmax.py`: -``` -# kernel dtype and BLOCK_SIZE to generate. -dtypes = ['fp32', 'fp16'] -blocks = [1024, 2048, 4096, 8192, 16384] -name_pattern = 'softmax_{}_{}' -sig_pattern = '*{},*{},i32,i32,i32' -group_pattern = 'softmax_{}' - -""" -SHOULD implement a function that returns a metadata list with format: - -function_table = [ - {'name': xx, - 'group': yy, - 'func': func, - 'sig': sig, - 'kwargs': kwargs - } -] - -The kwargs is a dict of {string: int} which is used for kernel constants. For example, BLOCK_SIZE of softmax. -""" - -def get_funcion_table(): - ...... - ...... -``` - -When compiling onnxruntime with `--use_triton_kernel` flag, this softmax kernel will be compiled and combined into libonnxruntime_providers_rocm.so for rocm or libonnxruntime_providers_cuda.so for cuda. - -### onnxruntime c++ code moidfication -To use the triton kernels in onnxruntime, we need to implement a c++ op that calls these triton kernels. - -And there are many implements for softmax, we implement a tunable op for softmax to choose the best performant one. - -Same as CK, we implement a function that returns all possible triton kernels, and the tunable op will select best one. - -``` -template -auto GetSoftmaxTritonOps() { - std::vector>>> ret; - auto group_name = GetSoftmaxTritonGroupName(); - // here use group_name to get all kernel with same group_name - // for example, 'softmax_fp16' represents a group of kernels with different BLOCK_SIZE for float16 softmax - auto *kernel_list = GetOrtTritonKernelByGroup(group_name); - if (kernel_list == nullptr) { - return ret; - } - - for (auto i : *kernel_list) { - // check params match - ..... - } - return ret; -} -``` - -### Test -Using kernel_explorer, we can test this softmax kernel. -``` -export KERNEL_EXPLORER_BUILD_DIR=/ws/code/onnxruntime/build_rocm/Debug - -python onnxruntime/python/tools/kernel_explorer/kernels/softmax_test.py -``` -and can see the results, here TunableOp select softmax_fp16_2048 which is a triton written kernel, it's better than others. -``` -SoftmaxTunable float16 batch_count=1 softmax_elements=2048 is_log_softmax=0 4.27 us, 1.92 GB/s -softmax_fp16_2048 float16 batch_count=1 softmax_elements=2048 is_log_softmax=0 4.48 us, 1.83 GB/s -.... -.... -``` diff --git a/onnxruntime/contrib_ops/rocm/diffusion/group_norm_triton.cuh b/onnxruntime/contrib_ops/rocm/diffusion/group_norm_triton.cuh new file mode 100644 index 0000000000..7ba0eefb9c --- /dev/null +++ b/onnxruntime/contrib_ops/rocm/diffusion/group_norm_triton.cuh @@ -0,0 +1,91 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include +#include +#include + +#include "contrib_ops/rocm/diffusion/group_norm_common.h" +#include "core/providers/rocm/triton_kernel.h" + +using namespace onnxruntime::rocm; + +namespace onnxruntime { +namespace contrib { +namespace rocm { + +#ifdef USE_TRITON_KERNEL + +namespace { + +template +std::string GetGroupNormTritonGroupName() { + std::string ret = "GroupNormTriton_"; + std::string swish_suffix = WithSwish ? "Swish_" : "Pass_"; + ret += swish_suffix; + ret += GetDataTypeName(); + return ret; +} + +} // namespace + +template +auto GetTritonGroupNormNHWCTypeStringAndOps() { + std::vector>>> ret; + auto group_name = GetGroupNormTritonGroupName(); + auto* kernel_list = GetOrtTritonKernelByGroup(group_name); + if (kernel_list == nullptr) { + return ret; + } + + for (auto i : *kernel_list) { + // Check params match + auto* metadata = GetOrtTritonKernelMetadata(i); + auto block_size = metadata->constants.at("BLOCK_SIZE"); + auto hw_size = metadata->constants.at("HW_SIZE"); + auto impl = [i, block_size, hw_size](const GroupNormNHWCParams* params) -> Status { + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF( + params->cPerGroup > block_size || params->cPerGroup * 2 <= block_size, + "Arg block_size (", block_size, ") is not the next power of 2 of cPerGroup (", params->cPerGroup, ")."); + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF( + params->hw % hw_size != 0, "Arg hw_size (", hw_size ") is not a divisor of hw (", params->hw, ")."); + if constexpr (WithSwish) { + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(!params->withSwish, "Swish version does not support GN w/o swish."); + } else { + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(params->withSwish, "Pass version does not support GN w/ swish."); + } + // Construct args for launch kernel + struct { + void* X; + void* Y; + const void* gamma; + const void* beta; + int hw; + int c; + int c_per_group; + float eps; + } args = { + (void*)params->src, + (void*)params->dst, + (const void*)params->gamma, + (const void*)params->beta, + params->hw, + params->c, + params->cPerGroup, + params->epsilon}; + + // Grid dim is (batch_count, groups, 1) + return LaunchTritonKernel(params->stream, i, params->n, params->groups, 1, &args, sizeof(args)); + }; + ret.emplace_back(std::make_pair(metadata->name, std::move(impl))); + } + return ret; +} + +#endif // USE_TRITON_KERNEL + +} // namespace rocm +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/rocm/diffusion/group_norm_triton.py b/onnxruntime/contrib_ops/rocm/diffusion/group_norm_triton.py new file mode 100644 index 0000000000..56b3a030b2 --- /dev/null +++ b/onnxruntime/contrib_ops/rocm/diffusion/group_norm_triton.py @@ -0,0 +1,104 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + +from itertools import product + +import triton +import triton.language as tl + + +@triton.jit +def group_norm_kernel( + input_ptr, + output_ptr, + gamma_ptr, + beta_ptr, + img_size, + c, + c_per_group, + eps, + BLOCK_SIZE: tl.constexpr, + HW_SIZE: tl.constexpr, + ACTIVATION_SWISH: tl.constexpr, +): + row_x = tl.program_id(0) + row_y = tl.program_id(1) + stride = img_size * c + input_ptr += row_x * stride + row_y * c_per_group + output_ptr += row_x * stride + row_y * c_per_group + gamma_ptr += row_y * c_per_group + beta_ptr += row_y * c_per_group + + cols = tl.arange(0, BLOCK_SIZE) + hw = tl.arange(0, HW_SIZE) + offsets = hw[:, None] * c + cols[None, :] + mask = (cols < c_per_group)[None, :] + + # Calculate mean and variance + _sum = tl.zeros([HW_SIZE, BLOCK_SIZE], dtype=tl.float32) + _square_sum = tl.zeros([HW_SIZE, BLOCK_SIZE], dtype=tl.float32) + for i in range(tl.cdiv(img_size, HW_SIZE)): + x_ptr = input_ptr + i * HW_SIZE * c + a = tl.load(x_ptr + offsets, mask=mask, other=0.0).to(tl.float32) + _sum += a + _square_sum += a * a + + # Set axis=None (or leave it unspecified) to reduce all axes. + # TODO: In older Triton we have to reduce an axis at a time, but in our case + # for some configs it may have some issue when reducing sequentially along the axes. + group_mean = tl.sum(_sum, axis=None) / (img_size * c_per_group) + group_var = tl.sum(_square_sum, axis=None) / (img_size * c_per_group) - group_mean * group_mean + + rstd = 1 / tl.sqrt(group_var + eps) + + # Normalize and apply linear transformation + gamma = tl.load(gamma_ptr + cols, mask=cols < c_per_group).to(tl.float32) + beta = tl.load(beta_ptr + cols, mask=cols < c_per_group).to(tl.float32) + for i in range(tl.cdiv(img_size, HW_SIZE)): + x_ptr = input_ptr + i * HW_SIZE * c + y_ptr = output_ptr + i * HW_SIZE * c + x = tl.load(x_ptr + offsets, mask=mask, other=0.0).to(tl.float32) + x_hat = (x - group_mean) * rstd + y = x_hat * gamma + beta + if ACTIVATION_SWISH: + y *= tl.sigmoid(y) + tl.store(y_ptr + offsets, y, mask=mask) + + +# We can have more combinations of blocks and hw_sizes, e.g., +# blocks = [16, 32, 64, 128, 256, 512] +# hw_sizes = [8, 16, 32, 64, 128, 256, 512] +# but this will result in too many functions and slow down the compilation. +with_swish = [True, False] +dtypes = ["fp32", "fp16"] +blocks = [16, 32, 64, 128] +hw_sizes = [8, 16, 32, 64, 128, 256] +warps = [1, 2, 4, 8, 16] +name_pattern = "GroupNormTriton_{}_{}_b{}_hw{}_w{}" +sig_pattern = "*{},*{},*fp32,*fp32,i32,i32,i32,fp32" +group_pattern = "GroupNormTriton_{}_{}" + + +def get_function_table(): + func_table = [] + + for swish, dtype, hw_size, warp, b in product(with_swish, dtypes, hw_sizes, warps, blocks): + swish_suffix = "Swish" if swish else "Pass" + name = name_pattern.format(swish_suffix, dtype, b, hw_size, warp) + group = group_pattern.format(swish_suffix, dtype) + sig = sig_pattern.format(dtype, dtype) + kwargs = { + "num_warps": warp, + "constants": {"BLOCK_SIZE": b, "HW_SIZE": hw_size, "ACTIVATION_SWISH": int(swish)}, + } + func_desc = {"name": name, "group": group, "func": group_norm_kernel, "sig": sig, "kwargs": kwargs} + func_table.append(func_desc) + return func_table + + +if __name__ == "__main__": + func_table = get_function_table() + for func_desc in func_table: + print(func_desc) diff --git a/onnxruntime/contrib_ops/rocm/diffusion/group_norm_tunable_op.h b/onnxruntime/contrib_ops/rocm/diffusion/group_norm_tunable_op.h index 7ca38608b8..5e9482cee4 100644 --- a/onnxruntime/contrib_ops/rocm/diffusion/group_norm_tunable_op.h +++ b/onnxruntime/contrib_ops/rocm/diffusion/group_norm_tunable_op.h @@ -11,6 +11,7 @@ #include "contrib_ops/rocm/diffusion/group_norm_common.h" #include "contrib_ops/rocm/diffusion/group_norm_impl.h" #include "contrib_ops/rocm/diffusion/group_norm_impl_kernel.cuh" +#include "contrib_ops/rocm/diffusion/group_norm_triton.cuh" namespace onnxruntime { namespace contrib { @@ -192,6 +193,17 @@ class GroupNormNHWCTunableOp : public TunableOp> { this->RegisterOp(std::move(op)); } #endif // USE_COMPOSABLE_KERNEL + +#ifdef USE_TRITON_KERNEL + for (auto&& [_, op] : GetTritonGroupNormNHWCTypeStringAndOps()) { + ORT_UNUSED_PARAMETER(_); + this->RegisterOp(std::move(op)); + } + for (auto&& [_, op] : GetTritonGroupNormNHWCTypeStringAndOps()) { + ORT_UNUSED_PARAMETER(_); + this->RegisterOp(std::move(op)); + } +#endif } }; diff --git a/onnxruntime/core/providers/cuda/triton_kernel.cu b/onnxruntime/core/providers/cuda/triton_kernel.cu index c3210d5dad..6ffbf0420a 100644 --- a/onnxruntime/core/providers/cuda/triton_kernel.cu +++ b/onnxruntime/core/providers/cuda/triton_kernel.cu @@ -11,55 +11,69 @@ #include "triton_kernel_infos.h" #endif -#define ORT_TRITON_CHECK(status, msg) \ - if ((status) != CUDA_SUCCESS) { \ - ORT_RETURN_IF(true, msg); \ - } +#define ORT_TRITON_CHECK(expr, msg) \ + do { \ + auto status = expr; \ + const char* error_str; \ + if (status != CUDA_SUCCESS) { \ + auto get_status_err_str = cuGetErrorString(status, &error_str); \ + ORT_UNUSED_PARAMETER(get_status_err_str); \ + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, msg, " ", error_str); \ + } \ + } while (0) -#define ORT_TRITON_THROW(status, msg) \ - if ((status) != CUDA_SUCCESS) { \ - ORT_THROW(msg); \ - } +#define ORT_TRITON_THROW(expr, msg) \ + do { \ + auto status = expr; \ + const char* error_str; \ + if (status != CUDA_SUCCESS) { \ + auto get_status_err_str = cuGetErrorString(status, &error_str); \ + ORT_UNUSED_PARAMETER(get_status_err_str); \ + ORT_THROW(msg, error_str); \ + } \ + } while (0) namespace onnxruntime { namespace cuda { namespace { -// a vector of kernel metadata +// A vector of kernel metadata static std::vector ort_triton_kernel_metadata; -// store group_name -> [kernel metadata id vector] +// Store group_name -> [kernel metadata id vector] static std::unordered_map> ort_triton_kernel_group_map; #ifdef USE_TRITON_KERNEL -// store func_name -> kernel metadata id +// Store func_name -> kernel metadata id static std::unordered_map ort_triton_kernel_map; const int GPU_WARP_SIZE = 32; +constexpr int kMaxThreadsPerBlock = 1024; +// Currently the max shared memory per block is hardcoded to 64KB. +constexpr int kMaxSharedMemoryPerBlock = 64 * 1024; Status GetSymbolFromLibrary(const std::string& symbol_name, void** symbol) { - dlerror(); // clear any old error str + dlerror(); // Clear any old error str - // USe RTLD_DEFAULT for search current lib.so - // value of RTLD_DEFAULT differs across posix platforms (-2 on macos, 0 on linux). + // Use RTLD_DEFAULT for search current lib.so. + // Value of RTLD_DEFAULT differs across posix platforms (-2 on macos, 0 on linux). void* handle = RTLD_DEFAULT; *symbol = dlsym(handle, symbol_name.c_str()); char* error_str = dlerror(); if (error_str) { Status status = ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, - "Failed to get symbol " + symbol_name + " with error: " + error_str); + "Failed to get symbol " + symbol_name + " with error: " + error_str); return status; } - // it's possible to get a NULL symbol in our case when Schemas are not custom. + // It's possible to get a NULL symbol in our case when Schemas are not custom. return Status::OK(); } #endif - /* - * Try to load HIP kernels that compiled by triton. + * Try to load CUDA kernels that are compiled by Triton. * They are in hsaco/cubin format, and should use cuModuleLoad to load these kernels. */ void TryToLoadKernel() { @@ -72,15 +86,15 @@ void TryToLoadKernel() { for (int i = 0; i < size; ++i) { auto k_i = kernel_infos[i]; - void *buff; + void* buff; ORT_THROW_IF_ERROR(GetSymbolFromLibrary(k_i.name_start, &buff)); // try to load module and get function CUmodule module; - ORT_TRITON_THROW(cuModuleLoadData(&module, buff), "load module data failed."); + ORT_TRITON_THROW(cuModuleLoadData(&module, buff), "Loading module data failed."); CUfunction function; - ORT_TRITON_THROW(cuModuleGetFunction(&function, module, k_i.func_name), "get funcion from module failed."); + ORT_TRITON_THROW(cuModuleGetFunction(&function, module, k_i.func_name), "Getting function from module failed."); // setup kernel metadata TritonKernelMetaData metadata; @@ -92,7 +106,7 @@ void TryToLoadKernel() { std::string group_name = k_i.group_name; // pass constants - for (auto &kv : k_i.constants) { + for (auto& kv : k_i.constants) { metadata.constants[kv.first] = kv.second; } @@ -100,7 +114,7 @@ void TryToLoadKernel() { ort_triton_kernel_metadata.push_back(metadata); ort_triton_kernel_map[fname] = idx; ort_triton_kernel_group_map[group_name].push_back(idx); - LOGS_DEFAULT(VERBOSE) << "loaded ort triton kernel: " << fname << " idx: " << idx; + LOGS_DEFAULT(VERBOSE) << "Loaded ort triton kernel: " << fname << " idx: " << idx; } #endif @@ -109,63 +123,64 @@ void TryToLoadKernel() { static std::once_flag load_ort_triton_kernel_flag; -} // end of namespace +} // namespace void LoadOrtTritonKernel() { // load kernel should be called only once std::call_once(load_ort_triton_kernel_flag, TryToLoadKernel); } -Status LaunchTritonKernel(cudaStream_t stream, std::string fname, int grid0, int grid1, int grid2, void* args, size_t args_size) { +Status LaunchTritonKernel(cudaStream_t stream, std::string fname, + int grid0, int grid1, int grid2, void* args, size_t args_size) { #ifdef USE_TRITON_KERNEL if (ort_triton_kernel_map.count(fname) == 0) { - // return unsupported status when not found function name in registry - // this error status will be used by tunableOp + // Return unsupported status if function name not found in registry. + // This error status will be used by TunableOp std::ostringstream message_stream; - message_stream << "can't find ort triton kernel name: " << fname; + message_stream << "Can't find ort triton kernel name: " << fname; std::string message = message_stream.str(); TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(true, message); } auto idx = ort_triton_kernel_map[fname]; - auto metadata = ort_triton_kernel_metadata[idx]; - - void* config[] = {CU_LAUNCH_PARAM_BUFFER_POINTER, args, CU_LAUNCH_PARAM_BUFFER_SIZE, &args_size, - CU_LAUNCH_PARAM_END}; - - ORT_TRITON_CHECK(cuLaunchKernel(metadata.func, - grid0, grid1, grid2, - GPU_WARP_SIZE * metadata.num_warps, 1, 1, - metadata.shared_mem_size, - stream, - nullptr, - (void**)&config), "launch kernel failed."); -#endif - + return LaunchTritonKernel(stream, idx, grid0, grid1, grid2, args, args_size); +#else return Status::OK(); +#endif } -Status LaunchTritonKernel(cudaStream_t stream, size_t idx, int grid0, int grid1, int grid2, void* args, size_t args_size) { +Status LaunchTritonKernel(cudaStream_t stream, size_t idx, + int grid0, int grid1, int grid2, void* args, size_t args_size) { #ifdef USE_TRITON_KERNEL if (idx >= ort_triton_kernel_metadata.size()) { - // return unsupported status when not found function name in registry - // this error status will be used by tunableOp + // Return unsupported status when idx exceeds the size of ort_triton_kernel_metadata. + // This error status will be used by TunableOp std::ostringstream message_stream; - message_stream << "can't find ort triton kernel idx: " << idx; + message_stream << "Can't find ort triton kernel idx: " << idx; std::string message = message_stream.str(); TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(true, message); } auto metadata = ort_triton_kernel_metadata[idx]; + int threads_per_block = GPU_WARP_SIZE * metadata.num_warps; + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF( + threads_per_block > kMaxThreadsPerBlock, + "The threads_per_block (", threads_per_block, ") exceeds the max allowed value (", kMaxThreadsPerBlock, ")."); + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF( + metadata.shared_mem_size > kMaxSharedMemoryPerBlock, + "The shared_mem_size (", metadata.shared_mem_size, ") exceeds the max allowed value (", + kMaxSharedMemoryPerBlock, " bytes)."); + void* config[] = {CU_LAUNCH_PARAM_BUFFER_POINTER, args, CU_LAUNCH_PARAM_BUFFER_SIZE, &args_size, CU_LAUNCH_PARAM_END}; ORT_TRITON_CHECK(cuLaunchKernel(metadata.func, grid0, grid1, grid2, - GPU_WARP_SIZE * metadata.num_warps, 1, 1, + threads_per_block, 1, 1, metadata.shared_mem_size, stream, nullptr, - (void**)&config), "launch kernel failed."); + (void**)&config), + "Launching kernel failed."); #endif return Status::OK(); diff --git a/onnxruntime/core/providers/rocm/math/softmax_triton.cuh b/onnxruntime/core/providers/rocm/math/softmax_triton.cuh index 11cf2daf20..737e396855 100644 --- a/onnxruntime/core/providers/rocm/math/softmax_triton.cuh +++ b/onnxruntime/core/providers/rocm/math/softmax_triton.cuh @@ -30,30 +30,30 @@ template auto GetSoftmaxTritonOps() { std::vector>>> ret; auto group_name = GetSoftmaxTritonGroupName(); - auto *kernel_list = GetOrtTritonKernelByGroup(group_name); + auto* kernel_list = GetOrtTritonKernelByGroup(group_name); if (kernel_list == nullptr) { return ret; } for (auto i : *kernel_list) { // check params match - auto *metadata = GetOrtTritonKernelMetadata(i); + auto* metadata = GetOrtTritonKernelMetadata(i); auto block_size = -1; const std::string block_name = "BLOCK_SIZE"; if (metadata->constants.count(block_name) != 0) { block_size = metadata->constants.at(block_name); } - auto impl = [i, block_size](const SoftmaxParams *params) -> Status { + auto impl = [i, block_size](const SoftmaxParams* params) -> Status { if (params->is_log_softmax) { - TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(true, "log_softmax not support"); + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(true, "log_softmax is not supported."); } if (block_size < params->softmax_elements) { - TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(true, "BLOCK_SIZE not support"); + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(true, "BLOCK_SIZE (", block_size, ") is not supported."); } // construct args for launch kernel struct { - void *out; - const void *in; + void* out; + const void* in; int in_stride; int out_stride; int n_cols; diff --git a/onnxruntime/core/providers/rocm/math/softmax.py b/onnxruntime/core/providers/rocm/math/softmax_triton.py similarity index 98% rename from onnxruntime/core/providers/rocm/math/softmax.py rename to onnxruntime/core/providers/rocm/math/softmax_triton.py index 6d9c1634e7..f4c9b64590 100644 --- a/onnxruntime/core/providers/rocm/math/softmax.py +++ b/onnxruntime/core/providers/rocm/math/softmax_triton.py @@ -41,7 +41,7 @@ sig_pattern = "*{},*{},i32,i32,i32" group_pattern = "softmax_{}" -def get_funcion_table(): +def get_function_table(): func_table = [] def get_num_warps(block_size): diff --git a/onnxruntime/core/providers/rocm/math/softmax_tunable_op.cuh b/onnxruntime/core/providers/rocm/math/softmax_tunable_op.cuh index d3c943c315..3daae19c9a 100644 --- a/onnxruntime/core/providers/rocm/math/softmax_tunable_op.cuh +++ b/onnxruntime/core/providers/rocm/math/softmax_tunable_op.cuh @@ -58,13 +58,13 @@ Status SoftmaxWarpwiseStaticSelection(const SoftmaxParams* para dim3 threads(warp_size, warps_per_block, 1); // Launch code would be more elegant if C++ supported FOR CONSTEXPR switch (log2_elements) { - #define LAUNCH_SOFTMAX_WARP_FORWARD(L2E) \ - case L2E: \ - softmax_warp_forward \ - <<stream>>>(params->output, params->input, params->batch_count, \ - params->input_stride, params->softmax_elements, \ - params->is_log_softmax); \ - break; +#define LAUNCH_SOFTMAX_WARP_FORWARD(L2E) \ + case L2E: \ + softmax_warp_forward \ + <<stream>>>(params->output, params->input, params->batch_count, \ + params->input_stride, params->softmax_elements, \ + params->is_log_softmax); \ + break; LAUNCH_SOFTMAX_WARP_FORWARD(0); // 1 LAUNCH_SOFTMAX_WARP_FORWARD(1); // 2 LAUNCH_SOFTMAX_WARP_FORWARD(2); // 4 @@ -132,7 +132,7 @@ class SoftmaxTunableOp : public tunable::TunableOp()) { + for (auto&& [_, op] : GetSoftmaxTritonOps()) { ORT_UNUSED_PARAMETER(_); this->RegisterOp(std::move(op)); } diff --git a/onnxruntime/python/tools/kernel_explorer/kernels/rocm/group_norm.cu b/onnxruntime/python/tools/kernel_explorer/kernels/rocm/group_norm.cu index fd1aaa97fb..0bd47b2c03 100644 --- a/onnxruntime/python/tools/kernel_explorer/kernels/rocm/group_norm.cu +++ b/onnxruntime/python/tools/kernel_explorer/kernels/rocm/group_norm.cu @@ -150,6 +150,51 @@ class CKGroupNormNHWC : public IKernelExplorer { }; #endif // USE_COMPOSABLE_KERNEL +#ifdef USE_TRITON_KERNEL +template +class GroupNormNHWCTriton : public IKernelExplorer { + public: + GroupNormNHWCTriton(DeviceArray& output, DeviceArray& workspace, DeviceArray& input, DeviceArray& gamma, DeviceArray& beta, + int batch_size, int height, int width, int num_channels, int num_groups, float epsilon, bool use_swish) + : params_(TuningContext(), Stream(), static_cast(output.ptr()), static_cast(workspace.ptr()), + static_cast(input.ptr()), static_cast(gamma.ptr()), static_cast(beta.ptr()), + batch_size, height, width, num_channels, num_groups, epsilon, use_swish) { + for (auto&& [name, op] : contrib::rocm::GetTritonGroupNormNHWCTypeStringAndOps()) { + name_strings_.emplace_back(name); + ops_.emplace_back(std::move(op)); + } + } + + void Run() override { + ORT_THROW_IF_ERROR(ops_[selected_op_](¶ms_)); + } + + std::vector ListOps() const { + return name_strings_; + } + + bool SelectOp(const std::string& name) { + for (size_t i = 0; i < ops_.size(); i++) { + if (name_strings_[i] == name) { + selected_op_ = i; + Status status = ops_[i](¶ms_); + return status.IsOK(); + } + } + + ORT_THROW("Cannot find implementation ", name); + } + + private: + using ParamsT = contrib::rocm::GroupNormNHWCParams; + using OpT = rocm::tunable::Op; + ParamsT params_{}; + std::vector ops_; + std::vector name_strings_; + size_t selected_op_{}; +}; +#endif // USE_TRITON_KERNEL + #define REGISTER_OP(name, type, threads_per_block, vec_size) \ py::class_>(m, #name "_" #type "_" #threads_per_block "_" #vec_size) \ .def(py::init class SoftmaxTriton : public IKernelExplorer { public: SoftmaxTriton(DeviceArray& output, DeviceArray& input, int softmax_elements, - int input_stride, int output_stride, int batch_count, bool is_log_softmax) + int input_stride, int output_stride, int batch_count, bool is_log_softmax) : params_(TuningContext(), Stream(), static_cast(output.ptr()), static_cast(input.ptr()), softmax_elements, input_stride, output_stride, batch_count, is_log_softmax) { for (auto&& [name, op] : rocm::GetSoftmaxTritonOps()) { diff --git a/tools/ci_build/compile_triton.py b/tools/ci_build/compile_triton.py index 212b23a0e6..eb476da058 100644 --- a/tools/ci_build/compile_triton.py +++ b/tools/ci_build/compile_triton.py @@ -145,7 +145,7 @@ def main(args): spec = importlib.util.spec_from_file_location(f"module_{i}", f) module = importlib.util.module_from_spec(spec) spec.loader.exec_module(module) - func_tb = module.get_funcion_table() + func_tb = module.get_function_table() m = compile(func_tb, out_dir) metadata.extend(m)