CANN pyasc 算子开发:set_mm_layout_transform 接口详解——控制 Mmad 矩阵乘加的 M/N 方向遍历顺序
【免费下载链接】pyasc本项目为Python用户提供算子编程接口,支持在昇腾AI处理器上加速计算,接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc
导读
asc.language.basic.set_mm_layout_transform是 CANN pyasc(Python 版 Ascend C 编程接口)中用于控制矩阵乘加计算(Mmad)遍历方向的基础算子接口。本文以 接口文档 为主体,结合仓库中的 Python 源码实现、IR 定义与 MLIR 测试用例,系统讲解该接口的参数语义、底层调用链、完整使用方式与适用场景,帮助算子开发者理解并正确使用这一影响 CUBE 计算结果产出顺序的关键开关。
接口概览:函数签名与返回值
set_mm_layout_transform位于asc.language.basic模块,可直接通过顶层asc命名空间导入使用。其完整签名如下:
asc.language.basic.set_mm_layout_transform(mm_layout_mode: bool) → None- 所属模块:
asc.language.basic(基础算子接口),对应源码 python/asc/language/basic/common.py - 参数:
mm_layout_mode,bool 类型 - 返回值:无(
None)
该接口对应的 Ascend C 原生函数原型为:
__aicore__ inline void SetMMLayoutTransform(bool mmLayoutMode);从接口命名可以看出,pyasc 采用了与 Ascend C 一一对应的设计原则:Python 接口名set_mm_layout_transform与 C++ 接口名SetMMLayoutTransform保持严格映射,开发者已有的 Ascend C 算子编写经验可以直接迁移到 pyasc 中。
参数语义:mm_layout_mode 控制什么
mm_layout_mode用于设置 Mmad(Matrix Multiply and Add,矩阵乘加单元,即 CUBE)的M/N 方向优先顺序,决定矩阵乘加计算产生结果时按何种维度顺序遍历:
| 参数取值 | 行为语义 |
|---|---|
True | CUBE 将先按 N 方向、再按 M 方向产生结果 |
False | CUBE 将先按 M 方向、再按 N 方向产生结果 |
约束说明
接口文档明确标注:本接口无约束说明。即没有类型、调用时机或硬件配置方面的额外限制,可在算子内核函数中按需调用。
底层实现:从 Python 到 IR 的调用链
了解该接口在 pyasc 内部的完整调用链,有助于理解其工作机制。在 python/asc/language/basic/common.py 中,接口的 JIT 实现如下:
@overload def set_mm_layout_transform(mm_layout_mode: bool) -> None: ... @require_jit @set_common_docstring(api_name="set_mm_layout_transform") def set_mm_layout_transform(mm_layout_mode: RuntimeBool) -> None: builder = global_builder.get_ir_builder() mm_layout_mode = _mat(mm_layout_mode) builder.create_asc_SetMMLayoutTransformOp(mm_layout_mode.to_ir())实现要点:
@overload声明:对外暴露类型化签名,便于类型检查与 IDE 提示;@require_jit装饰器:标记该函数只能在asc.jit编译的内核函数内调用,保证参数在编译期被正确捕获;_mat转换:将 Python bool 值转换为 IR 侧的运行时布尔值;create_asc_SetMMLayoutTransformOp:通过 IR Builder 创建对应的ascendc.set_mm_layout_transform指令节点,将高层 Python 调用下沉到 Asc IR。
IR 指令定义
对应的 IR 操作定义位于 include/ascir/Dialect/Asc/IR/Basic/Common.td:
def AscendC_SetMMLayoutTransformOp : APIOp<"set_mm_layout_transform", "SetMMLayoutTransform", [AscFunc]> { let summary = "Set MM layout transform mode"; let arguments = (ins I1:$mm_layout_mode); let assemblyFormat = "$mm_layout_mode attr-dict `:` type($mm_layout_mode)"; }- 该 Op 使用
APIOp<...>基类,声明为AscFunc特性,说明它作为函数级 API 被注册进 Asc 方言; - 参数
mm_layout_mode类型为I1(1-bit 整数),与 Python 侧 bool 类型一一对应。
编译验证:MLIR 测试中的降级行为
仓库的 Lit 测试文件 test/Target/AscendC/basic/common.mlir 验证了该 Op 到 Ascend C 代码的降级输出:
// CHECK-LABEL:void emit_set_mm_layout_transform(bool v1) { // CHECK-NEXT: AscendC::SetMMLayoutTransform(v1); // CHECK-NEXT: return; // CHECK-NEXT:} func.func @emit_set_mm_layout_transform(%mode: i1) { ascendc.set_mm_layout_transform %mode : i1 return }即ascendc.set_mm_layout_transform %mode : i1会被降级为 C++ 调用AscendC::SetMMLayoutTransform(v1);,印证了 pyasc 从 Python → Asc IR → Ascend C 代码的完整编译链路。
完整使用示例
最小调用示例
在asc.jit修饰的内核函数中直接调用即可,无需额外导入(asc顶层命名空间已导出该接口):
import asc @asc.jit def kernel_set_mm_layout_transform() -> None: # 先按 N 方向、再按 M 方向产生结果 asc.set_mm_layout_transform(True) # 先按 M 方向、再按 N 方向产生结果 asc.set_mm_layout_transform(False) kernel_set_mm_layout_transform[1]()单元测试中的标准写法
仓库单元测试 python/test/unit/language/basic/test_common_api.py 给出了标准测试模式:
def test_set_mm_layout_transform(mock_launcher_run): @asc.jit def kernel_set_mm_layout_transform() -> None: asc.set_mm_layout_transform(True) asc.set_mm_layout_transform(False) kernel_set_mm_layout_transform[1]() assert mock_launcher_run.call_count == 1该测试验证了:两个开关在同一个内核中连续调用是合法的,内核能够正常编译并启动一次。
在矩阵乘算子中的典型用法
set_mm_layout_transform通常与 Matmul 相关算子配合使用。一个典型场景是在矩阵乘内核中按需切换结果遍历方向:
import asc @asc.jit def matmul_with_layout() -> None: # 让 CUBE 先按 N 方向产生结果,便于后续按行优先做输出搬运 asc.set_mm_layout_transform(True) # 后续矩阵乘加计算将遵循该方向顺序 # ... 其余 Matmul 计算与数据搬运逻辑 ...接口文档位置与相关资源
- 接口详解文档:docs/python-api/language/generated/asc.language.basic.set_mm_layout_transform.md
- API 索引:docs/python-api/language/basic.md(basic 模块 API 总览表,可检索全部基础算子接口)
- Python 源码实现:python/asc/language/basic/common.py
- docstring 定义:python/asc/language/basic/utils.py(生成接口文档的元数据来源)
- IR Op 定义:include/ascir/Dialect/Asc/IR/Basic/Common.td
- MLIR 降级测试:test/Target/AscendC/basic/common.mlir
- 单元测试:python/test/unit/language/basic/test_common_api.py
总结
asc.language.basic.set_mm_layout_transform是一个轻量但影响矩阵乘加计算遍历顺序的关键开关:
- 参数简单:仅一个 bool 参数,
True表示先 N 后 M,False表示先 M 后 N; - 无约束:可在内核中自由、多次调用;
- 链路清晰:Python 调用 →
create_asc_SetMMLayoutTransformOp创建 IR → 降级为AscendC::SetMMLayoutTransform,全链路在仓库源码与测试中均有对应证据; - 实践要点:需在
@asc.jit内核函数内使用,通过顶层asc命名空间即可调用。
当你的矩阵乘算子对结果张量的遍历顺序有特定偏好(例如与后续向量运算或数据搬运的内存布局对齐)时,合理设置该开关可以显著简化数据处理逻辑。
【免费下载链接】pyasc本项目为Python用户提供算子编程接口,支持在昇腾AI处理器上加速计算,接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考