pyasc Matmul.set_tensor_b 详解:设置矩阵乘右矩阵 B 的三种方式与底层实现
2026/9/18 8:13:27 网站建设 项目流程

pyasc Matmul.set_tensor_b 详解:设置矩阵乘右矩阵 B 的三种方式与底层实现

【免费下载链接】pyasc本项目为Python用户提供算子编程接口,支持在昇腾AI处理器上加速计算,接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc

导读

本文围绕 CANN pyasc 项目(面向昇腾 AI 处理器的 Python 算子编程接口)中的高阶矩阵乘 APIasc.language.adv.Matmul.set_tensor_b展开,完整讲解其两种函数重载、与 Ascend C 原生SetTensorB的对应关系、参数语义、约束条件及调用流程。读完本文,你将掌握在 pyasc 中为 Matmul 对象设置右矩阵 B 的正确姿势,理解标量/GlobalTensor/LocalTensor 三种形态的适用场景,并能结合 matmul_mix.py 示例完成一个可运行的 Matmul 算子编写。

接口定位:Matmul 高阶 API 中的右矩阵设置入口

pyasc 为 Python 用户提供与 Ascend C 一一对应的算子编程接口。在 asc.language.adv.Matmul 这一组高阶 API 中,矩阵乘的计算公式为C = A * B + Bias,其中右矩阵 B 的传入正是由set_tensor_b负责。它属于 Matmul 对象在正式迭代计算(iterate_all/iterate/iterate_batch等)之前必须完成的"输入装配"环节,与左矩阵设置接口set_tensor_a对称存在。

该接口的完整函数原型定义在 matmul.py 中,通过@overload声明了两种调用形态:

def set_tensor_b(self, scalar: int) -> None: ... def set_tensor_b(self, tensor: BaseTensor, transpose: bool = False) -> None: ...
  • 标量形态:当 B 矩阵是单值标量时使用,将一个整数常量直接作为右矩阵参与计算;
  • 张量形态:当 B 矩阵是真实数据时使用,可传入GlobalTensor(全局内存)或LocalTensor(本地内存)类型,并通过transpose控制是否转置。

对应的 Ascend C 函数原型

set_tensor_b与 Ascend C 原生接口SetTensorB一一对应,包含三个重载版本,分别覆盖全局内存张量、本地内存张量和标量三种场景:

__aicore__ inline void SetTensorB(const GlobalTensor<SrcBT>& gm, bool isTransposeB = false)
__aicore__ inline void SetTensorB(const LocalTensor<SrcBT>& leftMatrix, bool isTransposeB = false)
__aicore__ inline void SetTensorB(SrcBT bScalar)

对照可见,Python 层的两个重载正是对这三个 C++ 重载的封装:tensor参数在运行时根据实际传入的是GlobalTensor还是LocalTensor分别映射到前两个原型,scalar参数映射到第三个原型,transpose则对应 C++ 侧的isTransposeB(默认均为false)。

参数说明

参数类型必填说明
scalarint二选一B 矩阵中设置的值,为标量。仅用于右矩阵为常量的场景
tensorBaseTensor二选一B 矩阵,类型为GlobalTensorLocalTensor
transposeboolB 矩阵是否需要转置,默认False

两种重载二选一使用:要么传scalar,要么传tensor(可搭配transpose),不能同时传入。

约束说明

  • 传入的 TensorB 地址空间大小需要保证不小于single_k * single_n(以元素个数计)。其中single_ksingle_n是 Matmul 单核计算的 K、N 方向分片大小,B 矩阵需要为每个核的分片计算提供完整的数据。
  • 数据类型受源码级校验约束。查看 matmul.py 的实现可知,张量形态下check_type仅允许以下类型:halffloat(即 float32)、int8以及对应的float16/float32;标量形态下则仅支持halffloatfloat16float32(matmul.py),传入其他类型会抛出ValueError("Tensor type is not supported in set_tensor_b")。这也意味着 int8 量化场景下 B 矩阵必须走张量形态,无法用标量形态表达。

底层实现:重载分发与 IR 构建

从源码结构看,set_tensor_b的 Python 实现采用OverloadDispatcher机制完成运行时重载分发(matmul.py):

@require_jit @set_matmul_docstring(api_name="set_tensor_b") def set_tensor_b(self, *args, **kwargs) -> None: dispatcher = OverloadDispatcher(__name__) builder = global_builder.get_ir_builder() @dispatcher.register(scalar=RuntimeInt) def _(scalar: RuntimeInt): check_type(self.b_dtype, [KT.half, KT.float_, KT.float16, KT.float32], "Scalar type is not supported in set_tensor_b") builder.create_asc_MatmulSetTensorBScalarOp(self.to_ir(), _mat(scalar, self.b_dtype).to_ir()) @dispatcher.register(tensor=BaseTensor, transpose=DefaultValued(RuntimeBool, False)) def _(tensor: BaseTensor, transpose: RuntimeBool = False): check_type(tensor.dtype, [KT.half, KT.float_, KT.int8, KT.float16, KT.float32], "Tensor type is not supported in set_tensor_b") transpose = _mat(transpose, KT.bit) builder.create_asc_MatmulSetTensorBOp(self.to_ir(), tensor.to_ir(), transpose.to_ir()) dispatcher(*args, **kwargs)

关键点解读:

  1. @require_jit装饰器:确保该方法仅在 JIT 编译上下文中被调用,由编译器在运行时将 Python 调用降级为 Asc IR 指令。
  2. 重载匹配OverloadDispatcher根据实参形态(scalartensor)自动选择对应的内部函数,transpose参数通过DefaultValued(RuntimeBool, False)提供默认值。
  3. 类型检查前置:在构建 IR 之前先做check_type校验,失败立即抛异常,将错误拦截在编译期之前。
  4. IR 生成:标量形态生成asc_MatmulSetTensorBScalarOp,张量形态生成asc_MatmulSetTensorBOptranspose被转换为 bit 类型的 IR 值),这些 Op 最终由后端发射为对应的 Ascend C 内核代码。

调用示例与完整使用流程

原文档给出的最小调用示例:

asc.adv.register_matmul(pipe, workspace, mm, tiling) mm.set_tensor_a(gm_a) mm.set_tensor_b(gm_b) # 设置右矩阵B mm.set_bias(gm_bias) mm.iterate_all(gm_c)

在真实算子中,set_tensor_b通常与多核切分逻辑配合。仓库示例 matmul_mix.py 给出了一个完整的 Matmul 内核写法,其中 B 矩阵的设置流程如下:

@asc.jit(always_compile=True) def matmul_kernel(a: asc.GlobalAddress, b: asc.GlobalAddress, c: asc.GlobalAddress, tiling: asc.adv.TCubeTiling, workspace: asc.GlobalAddress): offset_a, offset_b, offset_c, tail_m, tail_n = calc_offsets(tiling, IS_TRANS_A, IS_TRANS_B) a_global = asc.GlobalTensor() b_global = asc.GlobalTensor() c_global = asc.GlobalTensor() a_global.set_global_buffer(a + offset_a) b_global.set_global_buffer(b + offset_b) # 为 B 矩阵绑定全局内存基址 + 片内偏移 c_global.set_global_buffer(c + offset_c) pipe = asc.TPipe() matmul = asc.adv.Matmul( a=asc.adv.MatmulType(asc.TPosition.GM, asc.CubeFormat.ND, a_global.dtype, IS_TRANS_A), b=asc.adv.MatmulType(asc.TPosition.GM, asc.CubeFormat.ND, b_global.dtype, IS_TRANS_B), c=asc.adv.MatmulType(asc.TPosition.GM, asc.CubeFormat.ND, c_global.dtype), ) asc.adv.register_matmul(pipe, workspace, matmul, tiling) if asc.get_block_idx() < tiling.used_core_num: matmul.set_tensor_a(a_global, IS_TRANS_A) matmul.set_tensor_b(b_global, IS_TRANS_B) # 传入转置标志 matmul.set_tail(tail_m, tail_n) # 尾核调整分片 matmul.iterate_all(c_global) matmul.end() asc.pipe_barrier(asc.PipeID.PIPE_ALL)

其中IS_TRANS_B即为set_tensor_btranspose参数。若 B 需要转置,还需在 Tiling 侧同步设置,见 generate_tiling 中matmul_tiling.set_b_type(host.TPosition.GM, host.CubeFormat.ND, host.DataType.DT_FLOAT16, False)的最后一个布尔参数。

与相邻接口的配合关系

set_tensor_b不是一个孤立接口,它在 Matmul 生命周期中的位置如下:

  1. 初始化:先通过 register_matmul 完成 Matmul 对象初始化(分离模式下需在init_buffer之前调用,最多支持 4 个 Matmul 对象)。
  2. 装配输入:依次调用set_tensor_a(左矩阵)、set_tensor_b(右矩阵)、set_bias(可选偏置)。
  3. 调整分片:如果当前是尾核,还需调用 set_tail 重新设置single_core_m/single_core_n/single_core_k——注意 B 矩阵的地址空间约束正是基于single_k * single_n计算的,因此set_tail调整 N/K 分片后,需要确保 B 的内存范围仍然满足该约束。
  4. 迭代计算:调用 iterate_all(一次计算single_core_m * single_core_n的 C 矩阵)或iterate/iterate_batch等接口。
  5. 资源释放:多个 Matmul 对象切换时调用end()释放计算资源。

常见错误与避坑提示

  • 类型不匹配:向set_tensor_b传入int8以外的非支持类型,或标量形态传入不支持的类型,会直接抛出ValueError。请先确认MatmulType中 b 的 dtype 声明与实际传入张量一致。
  • 内存越界:B 张量基址加偏移后的可读范围若小于single_k * single_n个元素,将产生越界访问。多核场景下务必像 calc_offsets 那样按n_index * single_core_n(转置时按n_index * k_b * single_core_n)计算 B 的片内偏移。
  • 转置标志不一致set_tensor_btranspose必须与 Tiling 侧set_b_type中的转置配置、以及 Matmul 构造时MatmulType(..., IS_TRANS_B)保持一致,三者不一致会导致计算错误。

总结

asc.language.adv.Matmul.set_tensor_b是 pyasc 矩阵乘编程中装配右矩阵 B 的唯一入口,支持标量、GlobalTensorLocalTensor三种数据形态与可选的转置开关,底层通过OverloadDispatcher分发并经类型校验后生成对应的 Asc IR 指令,最终发射为 Ascend C 的SetTensorB内核代码。结合register_matmulset_tensor_aset_tensor_bset_tailiterate_all的标准调用链,即可在昇腾 AI 处理器上完成一个完整、可验证的 Matmul 算子。

【免费下载链接】pyasc本项目为Python用户提供算子编程接口,支持在昇腾AI处理器上加速计算,接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询