ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

ONNX Runtime 集成 Triton Kernel 实战指南:以 Softmax 为例打通 CUDA/ROCm 的编译、注册与调优全链路

ONNX Runtime 集成 Triton Kernel 实战指南:以 Softmax 为例打通 CUDA/ROCm 的编译、注册与调优全链路 ONNX Runtime 集成 Triton Kernel 实战指南以 Softmax 为例打通 CUDA/ROCm 的编译、注册与调优全链路【免费下载链接】onnxruntimeONNX Runtime: cross-platform, high performance ML inferencing and training accelerator项目地址: https://gitcode.com/GitHub_Trending/on/onnxruntime导读本文基于 ONNX Runtime 官方文档 ORT_Use_Triton_Kernel.md系统讲解如何将 Triton 语言编写的 Kernel 集成进 ONNX Runtime 的 CUDA/ROCm Execution ProviderEP并以softmax算子为完整示例覆盖 Triton Kernel 编写、多BLOCK_SIZE编译与元数据生成、C 算子接入TunableOp自动调优、以及基于 kernel_explorer 的验证测试全流程。读完本文你将掌握 ONNX Runtime 自定义高性能算子开发中的Triton 路线用 Python 写 kernel用构建脚本产出 GPU 二进制并嵌入 provider 动态库再用 C 统一注册、启动与调优从而在 CK 或手写 CUDA kernel 表现欠佳的场景下获得更优性能。一、动机为什么 ONNX Runtime 需要 Triton Kernel在某些算子上Triton 编写的 Kernel 比 Composable KernelCK或其他手写 Kernel 性能更好。Triton 是一种面向 GPU 的类 Python 领域专用语言DSL其编译器和自动并行化能力可以让开发者以接近标量代码的写法获得接近手写 CUDA 的性能。为此ONNX Runtime 实现了一套框架USE_TRITON_KERNEL使 CUDA/ROCm EP 能够直接编译、加载并调度 Triton 编写的 Kernel。整套框架的核心设计目标有三个编译期生成在构建 ONNX Runtime 时用 Python 侧工具将.py中的 Triton kernel 编译成 GPU 二进制CUDA 的cubin或 ROCm 的hsaco并嵌入 provider 动态库运行期加载C 侧通过 driver API 加载二进制、解析函数句柄并保留每个 kernel 的启动元数据num_warps、共享内存大小、编译期常量等自动调优所有同类型 kernel 通过TunableOp机制统一注册推理时自动 benchmark 并选择最优实现。从仓库源码看该框架的运行时核心位于 onnxruntime/core/providers/cuda/triton_kernel.cu 与 onnxruntime/core/providers/cuda/triton_kernel.h编译工具链位于 tools/ci_build/compile_triton.py。二、整体工作流从 Python Kernel 到 GPU 上运行集成一个 Triton kernel 需要经过以下四个阶段本文后续小节将逐一展开编写用triton.jit在 Python 文件中编写 kernel并实现get_function_table()返回待编译 kernel 的描述清单编译构建时运行tools/ci_build/compile_triton.py逐个调用triton.compile产出二进制文件cubin/hsaco通过objcopy转成.o、ar归档成静态库同时生成一个 C 头文件triton_kernel_infos.h记录每个 kernel 的符号名、函数名、分组、num_warps、共享内存大小和编译期常量链接与加载归档文件与头文件随 provider 库libonnxruntime_providers_cuda.so或libonnxruntime_providers_rocm.so一起编译链接EP 初始化时调用LoadOrtTritonKernel()用cuModuleLoadData/cuModuleGetFunction解析出可调用的CUfunction调度与调优C 算子通过GetOrtTritonKernelByGroup()按分组取出候选 kernel包装成tunable::Op交给TunableOp运行时自动选择性能最优者。三、第一步编写 Triton Kernel以 Softmax 为例文档中给出的 softmax kernel 示例位于onnxruntime/core/providers/rocm/math/softmax_triton.py该文件为文档所示示例路径当前仓库中实际包含的同类 Triton kernel 示例可参见 onnxruntime/contrib_ops/cuda/sparse/sparse_attention_v1/sparse_attention_triton.py 与 onnxruntime/contrib_ops/cuda/sparse/sparse_attention_v2/sparse_attention_v2_triton.py。核心 kernel 结构如下triton.jit def softmax_kernel( output_ptr, input_ptr, input_row_stride, output_row_stride, n_cols, BLOCK_SIZE: tl.constexpr ): # softmax implementations ... ...这是一个非常简单的实现其中有两条关键约束n_cols必须小于BLOCK_SIZE即每个 programblock处理的列数不能超过常量块大小BLOCK_SIZE必须是tl.constexpr编译期常量Triton 要求该值在编译时确定因为它直接决定循环展开、线程/共享内存布局等代码生成策略。由于推理时输入形状尤其是列数是变化的单一BLOCK_SIZE无法覆盖所有场景。因此框架的做法是为同一 kernel 编译多个不同BLOCK_SIZE的版本运行时按实际形状选择匹配的版本。每个BLOCK_SIZE版本编译后会生成不同的num_warps和共享内存占用这些启动参数被称为 kernel 的metadata是 onnxruntime 在 launch kernel 时必需的输入。四、第二步声明描述信息并实现get_function_table要让编译脚本知道该编译哪些 kernel、以什么签名编译、归到哪个组需要在 kernel 所在的 Python 模块中追加两样东西模块级的编译描述常量以及一个返回函数表的get_function_table()函数。文档中 softmax 的声明如下# 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_{} def get_function_table(): ...各字段含义字段含义示例值说明dtypes需要编译的 kernel 数据类型fp32、fp16分别对应 float 与 halfblocks需要编译的BLOCK_SIZE集合102416384覆盖不同列数规模name_patternkernel 实例命名模板softmax_{}_{}按 (dtype, block) 展开如softmax_fp16_2048sig_patternTriton 编译签名模板*{},*{},i32,i32,i32表示两个指针加三个 32 位整数参数group_pattern分组模板softmax_{}按 dtype 分组如softmax_fp16代表一组不同BLOCK_SIZE的 fp16 softmax kernelget_function_table()必须返回一个元数据列表元素格式固定为function_table [ { name: xx, # kernel 实例名如 softmax_fp16_2048 group: yy, # 分组名如 softmax_fp16 func: func, # triton.jit 装饰的 kernel 函数对象 sig: sig, # 编译签名 kwargs: kwargs # {string: int} 形式的编译期常量例如 {BLOCK_SIZE: 2048} } ]其中kwargs是{字符串: 整数}字典用于传入 Triton 的编译期常量对 softmax 而言即BLOCK_SIZE它会被原样保留进最终元数据供运行时在 launch 前还原常量参数。五、第三步编译脚本compile_triton.py深度解析编译脚本位于 tools/ci_build/compile_triton.py命令行参数如下--header 生成的头文件名默认 triton_kernel_infos.h --ort_root onnxruntime 根目录默认 onnxruntime --script_files 包含 get_function_table 的 Python 脚本可传多个 --obj_file 输出的归档文件名默认 triton_kernel_infos.a其执行流程可拆解为四个环节1. 动态加载脚本并汇总函数表对应main()脚本用importlib.util.spec_from_file_location逐个加载--script_files指定的模块并调用其中的get_function_table()得到全部待编译 kernel 描述spec importlib.util.spec_from_file_location(fmodule_{i}, f) module importlib.util.module_from_spec(spec) spec.loader.exec_module(module) func_tb module.get_function_table()2. 调用triton.compile编译并提取二进制对应compile()对函数表中的每一项以func、signaturesig及**kwargs调用triton.compile并从编译结果ret.asm中提取 GPU 二进制与元数据若存在hsaco_pathROCm 平台将编译产物复制为{name}.hsaco若存在cubinCUDA 平台将字节流写为{name}.cubin否则抛出异常提示未找到 ROCm 或 CUDA 编译产物。同时记录每个 kernel 的func_nameret.metadata[name]、num_warpsret.metadata[num_warps]、sharedret.metadata[shared]即共享内存字节数以及kwargs中的constants。3. 用 objcopy/ar 把二进制嵌入可链接对象对应convert_lib_to_obj()与archive_obj_files()由于动态库无法直接包含裸.cubin/.hsaco文件脚本使用objcopy将二进制文件转换为带符号的 ELF 目标文件objcopy -I binary -O elf64-x86-64 -B i386:x86-64 {name}.cubin {name}.o转换后的.o会被ar rcs归档进triton_kernel_infos.a最终随 provider 库一并链接。二进制文件以_binary_{name}_start形式导出符号供 C 侧通过dlsym定位。4. 生成 C 元数据头文件对应convert_and_save()脚本为每个 kernel 生成一行_TritonKernelInfo初始化器包含struct _TritonKernelInfo { const char* name_start; // _binary_xxx_start 符号名 const char* func_name; // Triton 编译出的内部函数名 const char* group_name; // 分组名 const char* name; // kernel 实例名 int num_warps; // 线程束数量 int shared; // 共享内存字节数 std::unordered_mapstd::string, int constants; // 编译期常量 };最终写出kernel_infos[]数组并保存为triton_kernel_infos.h。对test_compile_triton.py等自动化测试见 tools/ci_build/test_compile_triton.py该脚本同样作为核心被测对象。六、第四步构建集成与运行时加载6.1 构建期--use_triton_kernel标志构建 ONNX Runtime 时传入--use_triton_kernel标志对应 CMake 宏USE_TRITON_KERNEL上述 softmax kernel 即会被编译并合并进 provider 动态库ROCm 平台libonnxruntime_providers_rocm.soCUDA 平台libonnxruntime_providers_cuda.so构建侧集成点可从 cmake/onnxruntime_providers_cuda.cmake 中看到对 onnxruntime/core/providers/cuda/triton_kernel.h 的引用。此外tritonPython 包是构建环境的前置依赖CI 镜像中固定为triton3.5.0见 tools/ci_build/github/linux/docker/scripts/manylinux/requirements.txt版本需与编译脚本兼容。6.2 运行时EP 初始化时加载 kernel在 CUDA EP 的构造函数中USE_TRITON_KERNEL宏开启时会在初始化末尾调用LoadOrtTritonKernel()见 onnxruntime/core/providers/cuda/cuda_execution_provider.cc该调用由std::call_once保证只执行一次triton_kernel.cu。TryToLoadKernel()是加载核心triton_kernel.cu其过程为遍历编译期生成的kernel_infos[]数组通过dlsym(RTLD_DEFAULT, name_start)获取嵌入的二进制起始地址调用cuModuleLoadData将cubin/hsaco数据加载为 CUDA module调用cuModuleGetFunction取得 Triton 编译出的CUfunction将num_warps、共享内存大小、常量映射等填入TritonKernelMetaData登记进三个全局注册表ort_triton_kernel_metadata按序号索引的元数据向量ort_triton_kernel_map按实例名如softmax_fp16_2048索引ort_triton_kernel_group_map按分组名如softmax_fp16映射到一组 kernel 序号。6.3 元数据结构TritonKernelMetaData定义在 triton_kernel.hstruct TritonKernelMetaData { int num_warps; int shared_mem_size; CUfunction func; std::unordered_mapstd::string, int constants; std::string name; };同文件还提供了GetDataTypeNameT()的映射fp32/fp16/fp64/bf16用于在 C 侧按模板类型拼出与 Python 侧一致的 kernel/分组名。七、C 算子接入注册 Triton Kernel 并交给 TunableOp要真正在 onnxruntime 中使用这些 Triton kernel还需要实现一个调用它们的 C 算子。与 CK 的做法类似框架约定实现一个返回该算子所有候选 Triton kernel的函数然后由TunableOp自动 benchmark 并选择最优实现。文档给出的示意代码如下template typename T, typename OutputT auto GetSoftmaxTritonOps() { std::vectorstd::pairstd::string, tunable::OpSoftmaxParamsT, OutputT ret; auto group_name GetSoftmaxTritonGroupNameT(); // 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; }这段代码的关键点GetSoftmaxTritonGroupNameT()依据模板类型拼出分组名例如softmax_fp16声明于 triton_kernel.hGetOrtTritonKernelByGroup(group_name)返回该分组下所有 kernel 的序号向量实现见 triton_kernel.cu若分组不存在返回nullptr遍历候选 kernel 时逐一校验参数是否匹配例如输入列数是否小于对应BLOCK_SIZE、数据类型是否一致将匹配者包装为tunable::Op加入返回列表。最终调度时算子调用LaunchTritonKerneltriton_kernel.cu完成启动它支持按序号size_t idx或按实例名std::string fname两种方式定位 kernel并将CU_LAUNCH_PARAM_BUFFER_POINTER/CU_LAUNCH_PARAM_BUFFER_SIZE打包后通过cuLaunchKernel提交。启动前有两项运行时护栏线程数校验threads_per_block 32 * num_warps不得超过 1024kMaxThreadsPerBlock共享内存校验shared_mem_size不得超过硬编码上限 64KBkMaxSharedMemoryPerBlock见 triton_kernel.cu。任一校验不通过都会返回TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF状态交由TunableOp决策如回退到其他实现。若编译时未开启USE_TRITON_KERNEL两个LaunchTritonKernel重载为空实现直接返回Status::OK()保证未启用该功能的构建不受影响。八、测试与验证kernel_explorer 实测调优结果借助 onnxruntime 的 kernel_explorer 工具可以单独跑 Triton softmax kernel 的基准测试。文档给出的测试方式为export KERNEL_EXPLORER_BUILD_DIRONNXRUNTIME_BUILD_DIR python onnxruntime/python/tools/kernel_explorer/kernels/softmax_test.py测试输出示例节选SoftmaxTunable float16 batch_count1 softmax_elements2048 is_log_softmax0 4.27 us, 1.92 GB/s softmax_fp16_2048 float16 batch_count1 softmax_elements2048 is_log_softmax0 4.48 us, 1.83 GB/s ...结果解读SoftmaxTunable是TunableOp封装层含调优开销后的整体表现softmax_fp16_2048是本次被选中的 Triton kernel 实例该场景下fp16、softmax_elements2048TunableOp最终选择了 Triton 编写的softmax_fp16_2048且其单次执行耗时4.48 us优于同组其他候选实现TunableOp的调优结论并非固定值——它会针对不同输入形状如不同batch_count、softmax_elements分别评测并缓存最优选择这正是框架编译多个BLOCK_SIZE版本并分组注册的意义所在。九、关键约束与使用注意事项综合文档与源码实现使用本框架时需特别注意以下约束BLOCK_SIZE编译期确定Triton kernel 中的常量参数如BLOCK_SIZE必须声明为tl.constexpr并保证调用侧n_cols小于BLOCK_SIZE否则运行结果不正确形状覆盖靠多版本为支持不同输入形状应像 softmax 示例一样枚举多组BLOCK_SIZE编译多版本运行时按参数匹配选择kwargs中常量会被固化get_function_table()返回的kwargs如{BLOCK_SIZE: 2048}在编译期即确定并随元数据进入运行时用于还原 kernel 常量num_warps与共享内存有上限运行时强制32 * num_warps ≤ 1024、shared_mem_size ≤ 64KB超出即被判定为 unsupported回退给TunableOp处理平台产物不同CUDA 平台产物为cubinROCm 平台产物为hsaco两者通过objcopy统一转成 ELF 目标文件后链接跨平台无需改 C 代码构建期依赖 Triton Python 包tools/ci_build/compile_triton.py直接import triton编译环境需安装与仓库 CI 兼容的 Triton 版本。十、小结一套Python 编写、构建期编译、运行期调优的算子接入范式通过本文的 softmax 示例可以归纳出 ONNX Runtime 接入 Triton kernel 的标准范式在 Python 模块中用triton.jit编写 kernel并通过get_function_table()声明dtypes/blocks/name_pattern/sig_pattern/group_pattern与kwargs构建时以--use_triton_kernel启用框架compile_triton.py 自动完成编译、二进制嵌入与triton_kernel_infos.h生成C 侧实现返回候选 Triton kernel的Get*TritonOps()函数与TunableOp无缝衔接完成自动调优运行时由 triton_kernel.cu 统一负责加载、元数据管理与cuLaunchKernel启动用 kernel_explorer 基准测试验证调优结果确认 Triton kernel 在目标场景下确实优于其他实现。这套框架让 onnxruntime 能以较低的工程成本复用 Triton 生态的高性能 kernelkernel 逻辑用 Python 表达、多配置版本自动展开、启动细节由框架统一封装最终性能表现由TunableOp以实测数据说话——这正是文档示例中softmax_fp16_2048能够脱颖而出的原因。【免费下载链接】onnxruntimeONNX Runtime: cross-platform, high performance ML inferencing and training accelerator项目地址: https://gitcode.com/GitHub_Trending/on/onnxruntime创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表