NVIDIA CUTLASS 中的 EVT 功能扩展:实现 Sqrt 操作支持
2025-05-31 17:33:09作者:邬祺芯Juliet
背景介绍
NVIDIA CUTLASS 是一个高性能 CUDA C++ 模板库,用于实现矩阵乘法和其他线性代数运算。其中的 Epilogue Visitor Tree (EVT) 提供了一种灵活的方式来定义和组合核函数的尾端操作。在 Python 接口中,用户可以通过 cutlass.epilogue.trace 方法来定义自定义的尾端操作。
问题发现
在尝试使用 Python 接口定义 Adam 优化器的尾端操作时,开发者发现当前 EVT 实现缺少对 sqrt (平方根) 运算的支持。具体场景是在 Adam 优化器的实现中,需要计算梯度归一化项时使用了 torch.sqrt 操作,但 EVT 系统无法识别这一操作。
技术分析
通过分析 CUTLASS 的代码结构,我们发现 EVT 系统的操作映射主要在以下几个部分实现:
- C++ 核心功能:在
include/cutlass/functional.h中定义了各种基础数学运算 - Python 绑定:在 Python 接口中通过
ast_op_to_bindings方法将 Python AST 节点映射到 CUTLASS 的功能操作
当前系统已经支持了基本的算术运算(加、减、乘、除)和一些常用函数(如最大值、最小值),但缺少对平方根运算的支持。
解决方案
要实现 sqrt 操作的支持,需要在以下几个层面进行修改:
-
C++ 核心层:
- 在
functional.h中添加sqrt的函数实现 - 确保实现支持 CUDA 设备代码和模板参数
- 在
-
Python 接口层:
- 扩展
ast_op_to_bindings的映射表,添加对sqrt函数的支持 - 处理 Python AST 中
Call节点的特殊处理逻辑
- 扩展
-
类型系统:
- 确保新操作支持 CUTLASS 支持的各种数据类型(float16, float32, bfloat16 等)
- 实现类型推导规则
实现建议
对于想要贡献此功能的开发者,建议按照以下步骤进行:
- 首先在
functional.h中添加sqrt的模板函数实现 - 在 Python 接口中添加对应的操作映射
- 添加单元测试验证功能正确性
- 考虑性能优化(如使用 CUDA 内置函数)
- 文档更新,说明新支持的操作
扩展思考
这个问题反映了 EVT 系统的一个通用扩展模式。类似的数学函数(如指数、对数等)也可以通过相同的方式添加。CUTLASS 团队可以考虑建立一个更系统化的机制来支持常见数学函数的添加,而不是逐个硬编码。
总结
通过为 CUTLASS EVT 添加 sqrt 操作支持,可以显著增强其在机器学习优化算法(如 Adam)中的应用能力。这一改进不仅解决了眼前的问题,也为未来扩展更多数学函数提供了参考模式。对于深度学习框架开发者来说,这样的扩展意味着能够更灵活地在高性能核函数中实现复杂的数学运算组合。
登录后查看全文
热门项目推荐
相关项目推荐
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0214
cann-learning-hubCANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。Jupyter Notebook0138
uni-appA cross-platform framework using Vue.jsJavaScript08
GLM-5.2智谱开源 GLM-5.2,这是针对长文本任务的最新旗舰模型。相较于前代产品 GLM-5.1,它在长文本任务处理能力上实现了显著飞跃,并且首次在稳定的 100 万 token 上下文中提供这一能力。Jinja00
SwanLab⚡️SwanLab - an open-source, modern-design AI training tracking and visualization tool. Supports Cloud / Self-hosted use. Integrated with PyTorch / Transformers / LLaMA Factory / veRL/ Swift / Ultralytics / MMEngine / Keras etc.Python00
tiny-universe《大模型白盒子构建指南》:一个全手搓的Tiny-UniverseJupyter Notebook03
热门内容推荐
最新内容推荐
项目优选
收起
deepin linux kernel
C
32
16
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
469
465
暂无描述
Dockerfile
778
5.08 K
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
877
2.03 K
Ascend Extension for PyTorch
Python
758
968
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
697
1.4 K
昇腾LLM分布式训练框架
Python
185
231
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.1 K
1.14 K
本仓库是 Flutter SDK 与 Flutter Engine 的 OpenHarmony 适配版本,由 CPF-Flutter 团队维护。开发者可使用熟悉的 Flutter 技术栈开发 OpenHarmony 应用,3.35.7 及以后的适配版本可基于本仓库源码构建支持 OpenHarmony 的 Flutter Engine。
Dart
1.04 K
271
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
2.25 K
677