如何用 catlass_cppgen 生成 Ascend 高性能 GEMM 算子代码
【免费下载链接】YiA series of large language models trained from scratch by developers @01-ai项目地址: https://gitcode.com/GitHub_Trending/yi/Yi
catlass_cppgen 是一个 Python 代码生成框架:你只用几行代码描述张量形状、数据类型和布局,它就能自动产出基于 CATLASS 的 C++ 核函数代码,覆盖 Matmul、Group GEMM 和 EVG 后处理三类场景。
为什么需要它:手写 Ascend 算子的三个痛点
手写高性能 GEMM 核函数,难点往往不在矩阵乘法本身,而在周边细节:Tile 该怎么切、调度策略选哪个、不同代际硬件(AtlasA2/A3、Ascend950)之间参数如何对齐。稍微改一个参数,可能就要把整段 C++ 实现重读一遍。catlass_cppgen 把这些实现细节收进代码生成流程里:你在 Python 侧配置参数,框架负责输出可编译的核函数模板,并配套完整的数据类型与布局抽象,出错空间小很多。
一条命令装好 🔧
拿到源码后,开发模式下最省事的做法是直接用可编辑方式安装:
pip install -e .如果偏好分发流程,也可以先执行python -m build,在dist/下生成.whl和.tar.gz,再对产物执行pip install。
三步生成第一个 GEMM 核函数 🚀
整体流程是:描述输入 → 创建算子拿 Kernel → 打印生成代码。下面是最小示例:
from catlass_cppgen.op.gemm import Gemm from catlass_cppgen.common.op_tensor import OpTensor from catlass_cppgen.common.data_type import DataType from catlass_cppgen.catlass.layout.layout import RowMajor from catlass_cppgen.catlass.arch.arch import Arch a = OpTensor.from_shape_stride((128, 256), (256, 1), DataType.FLOAT) b = OpTensor.from_shape_stride((256, 384), (384, 1), DataType.FLOAT) gemm = Gemm(atlas_arch=Arch.Ascend950, element=DataType.FLOAT, layout=RowMajor, A=a, B=b) kernel = gemm.get_kernels()[0] print(kernel.gen_kernel_template())注意输入只描述形状、步长和类型,无需绑定真实数据。拿到 Kernel 对象后,除了gen_kernel_template()输出核函数模板,还能调用gen_params_device()生成参数绑定代码。如果不想取默认的第一个,也可以按类型从返回列表里挑选,比如显式指定BasicMatmulKernel。
能力地图:从普通 Matmul 到 Split-K 与 EVG 🗺️
当前支持的算子与对应 Kernel 类如下,基础与批处理 Matmul 均固定alpha = 1.0、beta = 0.0:
| 功能场景 | Kernel 类 | 输入形态 | 切分与调度 |
|---|---|---|---|
| 多个矩阵组各自 M 维不同 | GroupedMatmulSliceMKernel | 各组独立维度 | M 轴切分 |
| 标准 Matmul | BasicMatmulKernel | A/B 均为 2 维 | 单核,可选 Bias |
| 批处理 Matmul,各批共享维度 | BatchedMatmulKernel | 3 维(batchCount, M, K) | 无 |
| 大 K 场景多核并行 | MultiCoreSplitkMatmulKernel | A/B 均为 2 维,可选 Bias | 沿 K 轴多核切分 |
| 尾块优化变体 | TailMultiCoreSplitkMatmulKernel | A/B 均为 2 维,可选 Bias | K 轴尾块处理 |
| 负载均衡调度 | StreamkMatmulKernel | A/B 均为 2 维,可选 Bias | Stream-K 策略 |
| 带 epilogue 后处理图的 GEMM | BasicMatmulTlaVisitorKernel | A/B 均为 2 维 | 配合 EVG 图 |
除 Kernel 本体外,框架还提供 EVG(Epilogue Visitor Graph)后处理:通过evg_config(或to_evg())挂接一段 Python 函数源码,即可在 epilogue 阶段串接计算。单节点支持 add/sub/mul/div 二元运算、relu、silu、sigmoid、leakyRelu、Prelu 激活、max/min 选择、cast 类型转换与 constant 常量,节点间可自由组合,并支持行广播。写法示例:
evg_config = { "fn_src": "def epilogue(accum, bias):\n return relu(accum + bias)", "example_inputs": { "accum": OpTensor.from_shape_stride((128, 256), (256, 1), DataType.FLOAT), "bias": OpTensor.from_shape_stride((1, 256), (256, 1), DataType.FLOAT), "result": OpTensor.from_shape_stride((128, 256), (256, 1), DataType.FLOAT), }, } gemm = Gemm(atlas_arch=Arch.Ascend950, evg_config=evg_config, A=a, B=b)配置传入Gemm后,可通过kernel.is_support_evg确认当前 Kernel 是否启用了该能力。
进阶入口:调优与扩展
- Tile 与调度调优:对 Kernel 调用
tune(GemmShape, GemmShape, dispatch_policy=...),可指定两级 Tile 形状与调度策略(如MmadPingpong(arch_tag=Arch.Ascend950)),从 Kernel 对象上直接入手。 - Group GEMM:
GroupGemm配合groupList张量(DataType.INT64+VectorLayout)描述多组矩阵,随后按组做 Tile 调优。 - 模块路径:算子规划入口在
op/目录(gemm.py、group_gemm.py),Kernel 特化类在kernel/,数据类型与张量描述等通用件在common/,架构代际声明在catlass/arch。 - 测试参考:
tests/按 CATLASS 特性、通用组件(类型/排布)、算子代码生成三层组织,是新手的最佳示例库。
文档导航:三份 API 手册 📚
- docs/kernel_api.md:算子如何规划、Kernel 有哪些可查询特性与调优参数,入门先看这份。
- docs/optensor_api.md:输入张量的各种描述方式(形状、步长、类型)都在这份里定义。
- docs/evg_api.md:后处理图怎么写、支持哪些节点,做 epilogue 定制前建议通读。
【免费下载链接】YiA series of large language models trained from scratch by developers @01-ai项目地址: https://gitcode.com/GitHub_Trending/yi/Yi
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考