三步生成 CATLASS 高性能内核:catlass_cppgen 代码生成框架上手指南
【免费下载链接】YiA series of large language models trained from scratch by developers @01-ai项目地址: https://gitcode.com/GitHub_Trending/yi/Yi
手写一个 GEMM 内核,要面对的是一长串模板参数:Tile 尺寸、调度策略、架构宏、布局描述,配错任何一处都是难以定位的编译报错,参数组合还会随矩阵规模不断膨胀。catlass_cppgen 是一个基于 Python 生成 CATLASS 高性能算子的代码生成框架——你在 Python 里声明张量形状、数据类型与目标架构,它替你产出可编译的 C++ 内核代码。
项目定位:Python 侧声明,C++ 侧落地
一句话概括:catlass_cppgen 把“描述算子”和“实现算子”拆开,你只管前者。它有三个核心能力值得记住:
- 声明式参数定义:用 OpTensor 描述输入的形状、步幅与数据类型即可,无需绑定真实数据,也不用关心底层内存;
- 一键内核生成:从算子对象拿到 Kernel 后,get_kernels()、tune()、gen_kernel_template() 三次调用就完成“规划 → 调优 → 出码”全流程;
- 后处理可扩展:借助 EVG(Epilogue Visitor Graph,后处理访问者图),把激活、Bias 相加、类型转换等计算挂到矩阵乘的尾声阶段,无需手写融合代码。
架构层面,框架通过 Arch 枚举声明目标代际,覆盖 AtlasA2/A3、Ascend950 等多种硬件架构。
能力全景:能生成哪些算子
GEMM 家族是主力,按优化策略分档,每档对应一个 Kernel 特化类:
- 基础矩阵乘(BasicMatmulKernel):A、B 为二维输入,固定 alpha=1.0、beta=0.0,支持可选 Bias,对应最常规的稠密 matmul;
- 批处理矩阵乘(BatchedMatmulKernel):A、B 为三维(batch, M, K)/(batch, K, N),所有批次共享同一组维度,适合 batch 推理;
- 多核 Split-K(MultiCoreSplitkMatmulKernel):沿 K 方向多核切分,K 很大而 M/N 偏小时用它提升效率,同样支持可选 Bias;尾块场景可用优化变体 TailMultiCoreSplitkMatmulKernel;
- Stream-K(StreamkMatmulKernel):采用 Stream-K 调度策略摊平负载,适合负载不均的矩阵形状;
- EVG Visitor 矩阵乘(BasicMatmulTlaVisitorKernel):面向 EVG 后处理框架的 matmul 变体。
分组矩阵乘方面,目前提供沿 M 轴切分的 GroupedMatmulSliceMKernel,用于在一次调用内完成多组 M 维度各异的矩阵乘。
后处理能力由 EVG 承载,写法接近一段 Python 小函数,可用构件包括:二元运算 add / sub / mul / div(如accum + bias)、激活函数 relu / silu / sigmoid / leakyRelu / Prelu、比较选择 max / min、类型转换 cast、常量 constant。多个节点可以串联成组合计算,并支持行广播——例如偏置向量沿行方向展开到整块累加器。
安装与快速上手
安装步骤
📦 拿到源码后,任选一种方式装进环境:
- 开发模式(改动即时生效):
pip install -e . - 构建分发包:先
pip install build,再执行python -m build,在dist/下得到.whl与.tar.gz,然后pip install dist/catlass_cppgen-*.whl - 或以普通方式安装:
pip install .
五行代码跑通基础 GEMM
最小流程是:描述输入张量 → 建立 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.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, A=a, B=b) kernel = gemm.get_kernels()[0] print(kernel.gen_kernel_template())get_kernels() 返回若干候选 Kernel,你可以按类型挑选(例如指定 BasicMatmulKernel),也可以直接取第一个。拿到 kernel 后,gen_kernel_template() 输出核函数模板,gen_params_device() 负责参数绑定的代码生成,两者配合即是一份可用的内核。
进阶玩法:Group GEMM 与 EVG 后处理
🚀 两组进阶用法各由一个关键对象表达。
Group GEMM:先构造一个 INT64 类型的 groupList 张量(VectorLayout(4)、shape 为 (4,))声明分组规模,再传给 GroupGemm(atlas_arch=..., A=a, B=b_3d, groupList=groupList) 建立算子;取回 kernels 后照常用 tune(GemmShape(256, 256, 256), GemmShape(256, 256, 64)) 做 tiling 调优。
EVG 后处理:用一个函数头(fn_src)加一组 name:tensor 示例输入(example_inputs)描述后处理计算,框架生成内核时把它织入尾声阶段:
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), }, } kernel = Gemm(atlas_arch=Arch.Ascend950, evg_config=evg_config, A=a, B=b).get_kernels()[0] assert kernel.is_support_evgis_support_evg 为 True 即表示该 Kernel 支持 EVG;Kernel 侧还可以用 to_evg() 将后处理配置绑定上去,再按需 tune 调整 Tile 形状。
调优要点:TileShape、DispatchPolicy 与架构标签
⚙️ tune() 是调优的统一入口,接收三类信息:
- 两级 GemmShape(TileShape):如 GemmShape(128, 256, 64),分别描述宏块与原子粒度的 Tile 形状;
- 可选的 dispatch_policy:调度策略对象,如 MmadPingpong(arch_tag=Arch.Ascend950);
- arch_tag 架构标签:经 Arch 枚举声明目标代际,AtlasA2/A3、Ascend950 均可选。
矩阵维度不同,合理的 Tile 与策略组合差异很大,这正是显式调优存在的意义。
资源导航:文档、测试与源码入口
🔍 想深入某个环节,按路径查:
- docs/kernel_api.md:Kernel API 基础文档,覆盖调优与特性查询;
- docs/optensor_api.md:OpTensor API,张量声明的完整输入方式;
- docs/evg_api.md:EVG 后处理 API 参考;
- tests/:单元测试按 catlass、common、op 三个维度组织,读测试用例是理解 API 最快的一条路;
- 源码侧,catlass_cppgen/op/ 是算子入口(gemm.py、group_gemm.py),catlass_cppgen/kernel/ 存放各 Kernel 特化类,EVG 相关实现位于 catlass_cppgen/catlass/evg。
回到完整链路:Gemm / GroupGemm 负责算子规划,get_kernels() 取出调优对象,tune() / to_evg() 完成配置,最后由生成方法输出 C++——从声明到落地,就差这一次调用链。
【免费下载链接】YiA series of large language models trained from scratch by developers @01-ai项目地址: https://gitcode.com/GitHub_Trending/yi/Yi
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考