用几行 Python 换一整份 C++ 核函数:catlass_cppgen 的 CATLASS GEMM 代码生成实战

发布时间:2026/9/20 16:41:56
用几行 Python 换一整份 C++ 核函数:catlass_cppgen 的 CATLASS GEMM 代码生成实战
用几行 Python 换一整份 C 核函数catlass_cppgen 的 CATLASS GEMM 代码生成实战【免费下载链接】YiA series of large language models trained from scratch by developers 01-ai项目地址: https://gitcode.com/GitHub_Trending/yi/Yicatlass_cppgen 是一个 Python 驱动的算子代码生成框架声明张量的形状、步长、类型与目标架构后它替你完成 CATLASS 算子生成产出可直接编译的 C 核函数省去手写模板与参数绑定的功夫。三步跑通安装 catlass_cppgen 并生成第一个 GEMM 核函数装包、描述张量、出码三步走完。装包有三条路可挑先pip install build再执行python -m builddist/目录里会产出.whl和.tar.gz分发包接着pip install其中任意一个即可如果就在源码目录里干活pip install -e .直接以开发模式装好也行。装好之后用 OpTensor 声明两张纸面张量只描述形状、步长与类型不绑定真实内存交给 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_archArch.Ascend950, elementDataType.FLOAT, layoutRowMajor, Aa, Bb) print(gemm.get_kernels()[0].gen_kernel_template()) # 输出 C 核函数模板 小贴士OpTensor 不绑定底层数据这段代码在任何机器上都能直接跑通真实张量由下游推理框架提供你可以放心用它做纸面调优。整条生成链路的节奏是Gemm或GroupGemm先做算子规划get_kernels()返回一组候选 Kernel 对象再按需tune()或to_evg()最终得到一个配置完成的 Kernel。按场景挑 Kernel标准矩阵乘、批量乘与 Split-K 该选哪个拿到 Kernel 列表后不必纠结你的数据形态基本决定了该用哪一个特化类。两张二维矩阵相乘Basic 与 Batched 两个入口标准场景对应BasicMatmulKernelA、B 均为 2 维alpha 固定 1.0、beta 固定 0.0可以额外挂一个可选 Bias。如果你的数据按批次组织——A 形状为 (batchCount, M, K)、B 形状为 (batchCount, K, N) 的 3 维张量——那就换BatchedMatmulKernel它要求所有批次共享同一套矩阵维度。K 维偏大时Split-K 与 Stream-K 多核切分核函数K 方向足够长、单核串行太慢时框架备了三种沿 K 轴切分的变体MultiCoreSplitkMatmulKernel直接多核切 KTailMultiCoreSplitkMatmulKernel是尾块优化版本专治切不尽的剩余部分StreamkMatmulKernel采用 Stream-K 调度策略动态分配工作量。三者输入都是 2 维张量也都支持可选 Bias。一组矩阵维度不一Group GEMM 的 groupList 玩法一批矩阵的 M 维度互不相同时用按 M 轴切分的GroupedMatmulSliceMKernel。构造GroupGemm时多传一个groupList——INT64 类型、VectorLayout(4)布局、形状(4,)的 OpTensor用来描述分组信息from catlass_cppgen.op.group_gemm import GroupGemm from catlass_cppgen.catlass.layout.layout import VectorLayout groupList OpTensor(dtypeDataType.INT64, layoutVectorLayout(4), shape(4,)) group_gemm GroupGemm(atlas_archArch.Ascend950, Aa, Bb_3d, groupListgroupList) kernels group_gemm.get_kernels()进阶调优Tile 形状、DispatchPolicy 与 EVG 后处理扩展如果你要控制 Tile 与调度策略Kernel 默认配置就能出码想进一步榨性能时tune()让你显式指定两组 Tile 形状和调度策略from catlass_cppgen.catlass.gemm_coord import GemmShape from catlass_cppgen.catlass.gemm.dispatch_policy import MmadPingpong kernel.tune( GemmShape(128, 256, 64), GemmShape(128, 256, 64), dispatch_policyMmadPingpong(arch_tagArch.Ascend950), ) 小贴士atlas_arch决定代码生成的目标硬件AtlasA2/A3 与 Ascend950 各有一族匹配的 dispatch_policyMmadPingpong这类策略必须通过arch_tag指明架构调优意图才能落地。调优完成后gen_kernel_template()负责核函数主体gen_params_device()负责参数绑定的代码生成两者配合即是完整可编译产物。如果你想在乘法后串联 Bias 与激活函数EVG 后处理EVGEpilogue Visitor Graph是 CATLASS 的后处理框架在 catlass_cppgen 里被封装成一个evg_config字典fn_src写 epilogue 函数源码example_inputs声明后处理涉及的所有张量。可用算子面很宽——二元运算 add、sub、mul、div激活函数 relu、leakyRelu、Prelu、sigmoid、silu比较选择 max、min类型转换 cast还有 constant 常量支持多节点串联组合行方向广播计算也在能力范围内。启用时换用BasicMatmulTlaVisitorKernel构造 Gemm 时传入evg_config再用is_support_evg确认特性开关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), }, } gemm Gemm(atlas_archArch.Ascend950, evg_configevg_config, Aa, Bb) kernel gemm.get_kernels()[0] assert kernel.is_support_evg # 确认该 Kernel 支持 EVG工程导航五个目录看懂算子代码生成的源码组织catlass_cppgen/op/— 算子规划入口gemm.py与group_gemm.py在这里接收你的张量描述吐出候选 Kernel 集合catlass_cppgen/kernel/— 各算子特化实现gemm/、group_gemm/两个子目录装下全部矩阵乘变体kernel_base.py是它们的公共骨架catlass_cppgen/catlass/— CATLASS 特性层arch声明硬件代际evg承载后处理访问图layout与gemm_coord提供布局、形状抽象catlass_cppgen/common/— 通用地基OpTensor、DataType类型定义与共享工具函数docs/— 三份 API 文档 kernel_api.md、evg_api.md、optensor_api.md 各对应一类对象tests/下按 catlass、common、op 分组的单测则是最直接的用法示范谁该用这套 GEMM 代码生成工具在 AtlasA2/A3、Ascend950 这类平台上做矩阵乘、分组矩阵乘算子的性能开发时catlass_cppgen 能把描述参数 → 产出核函数这条链路压缩到十几行 Python深入用法直接翻 docs/ 的 API 文档tests/ 里的单测就是现成的参照实现。【免费下载链接】YiA series of large language models trained from scratch by developers 01-ai项目地址: https://gitcode.com/GitHub_Trending/yi/Yi创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考