ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

PyPTO 张量运算与转置实战指南:pypto.matmul 维度约束、广播、reshape 与 `.T` 陷阱全解析

PyPTO 张量运算与转置实战指南:pypto.matmul 维度约束、广播、reshape 与 `.T` 陷阱全解析 PyPTO 张量运算与转置实战指南pypto.matmul 维度约束、广播、reshape 与.T陷阱全解析【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym本篇技术指南基于 tensor-ops.mdDEBUG_GUIDEBOOK §9.7 张量运算与 §9.16 转置运算并结合 matmul.md、python-operators.md、pypto-view.md 等姊妹文档与仓库测试代码系统梳理 PyPTO JIT kernel 编写中与张量运算、转置相关的维度约束、广播规则、初始化和常见错误。读者读完后将掌握如何在pypto.frontend.jit中正确处理 1D/2D 张量的 matmul、如何用 reshape 完成广播、为什么.T在 PyPTO 中间张量上不可用、以及 matmul 转置标志a_trans/b_trans的正确用法。1. 背景这些模式从哪里来本文档记录的模式源自仓库中 GDRGated Delta Rule类 kernel 的真实开发排错经验属于 Agent 在 NPU kernel 开发中沉淀的可复用规律。这些经验与仓库 pypto-general-debug 技能包中的其他排错文档互为补充遇到 matmul 语法问题先看 matmul.md遇到 Python 运算符问题看 python-operators.md遇到 view/assemble 维度问题看 pypto-view.md。仓库中真实使用这些 API 的佐证可见于 test_glm_attention_pre_quant.py其中同时出现了pypto.set_vec_tile_shapes(...)、pypto.set_cube_tile_shapes([32, 32], [256, 512], [256, 256])、pypto.matmul(x_int8, weight, pypto.DT_INT32)以及大量pypto.reshape(..., inplaceTrue)的调用印证了本文所述 API 的实际用法。2. §9.7 张量运算Tensor Operations2.1 pypto.matmul 要求 2D 及以上的张量典型报错RuntimeError: Tensor dimension mismatch. Expect input_dim mat2_dim and both in [2, 3, 4], got input_dim: 2, mat2_dim: 1.原因pypto.matmul要求两个输入张量都至少是 2 维的。1D 张量必须先 reshape 成 2D 才能参与矩阵乘。最常见的场景——向量与矩阵相乘假设c_cum形状为[bt, bt]2Dgc_raw形状为[bt]1D直接传入会报错# c_cum 是 [bt, bt]2Dgc_raw 是 [bt]1D # 错误写法 g_cum pypto.matmul(c_cum, gc_raw, ...) # 正确写法——先 reshape 成 2D结果再 reshape 回 [bt] gc_raw_2d gc_raw.reshape([bt, 1]) g_cum pypto.matmul(c_cum, gc_raw_2d, ...).reshape([bt])从 matmul 得到 1D 结果的通用模式当期望结果形状是[bt]而 matmul 天然产出[bt, 1]时追加一次 reshape 即可result pypto.matmul(matrix, vector_2d, ...).reshape([bt])完整的pypto.matmulAPI转置标志、tile shapes、enable_split_k等参见 matmul.md §9.19。尤其是65535 ND 内轴上限规则operand.shape[-1] 65535该限制作用于每个 ND 操作数的物理内轴而非逻辑 K当逻辑 K 过大时应优先通过布局 /pypto.view形状 /a_trans/b_trans让 K 成为物理外轴再调用完整 matmul。2.2 广播用 pypto.reshape 增删维度PyPTO 中广播依赖显式的 reshape 来对齐维度标量参与逐元素运算时需要先显式扩展成与张量匹配的形状# 错误写法——标量没有显式 reshape result tensor * scalar # 正确写法——把标量扩展为 [bt, 1] 再逐元素相乘 scalar_reshaped pypto.reshape(scalar, [bt, 1]) result pypto.mul(tensor, scalar_reshaped)在 JIT 函数内Python 运算符*、、.exp()等是被支持的更简洁写法详见 python-operators.md§9.14。例如result a * b * (c d)与pypto.mul(pypto.mul(a, b), pypto.add(c, d))等价。2.3 zeros 初始化分配全零张量是 kernel 中常见的初始化操作显式指定形状与数据类型zeros pypto.zeros([M, N], pypto.DT_FP32)pypto.DT_FP32是 PyPTO 的数据类型枚举之一仓库中pypto.matmul(x_int8, weight, pypto.DT_INT32)见 test_glm_attention_pre_quant.py展示了显式传 dtype 的同类用法。2.4 逐元素运算对张量取负可以借助乘以 -1.0 实现negated pypto.mul(tensor, -1.0) # 乘以负一同样地在 JIT 内使用 Python 运算符与链式方法tensor.neg()、tensor.abs()、tensor.rsqrt()、tensor.sqrt()、tensor.exp()可以写出更干净的代码。3. §9.16 转置运算Transpose Operations3.1 两种可用的转置方式PyPTO 提供两种转置途径二者均可工作# 方式一.T 属性更简洁优先推荐 transposed tensor.T # 方式二显式转置函数 transposed pypto.transpose(tensor, dim0, dim1)3.2 关键陷阱.T对 PyPTO 中间张量不可用重要限制.T在pypto.frontend.jit内对PyTorch 后端的张量有效但对由 view / matmul 产生的 PyPTO 中间张量无效会抛出AttributeError: Tensor object has no attribute T结论与建议对 PyPTO 张量执行转置时优先使用pypto.transpose(tensor, dim0, dim1)对 matmul 中的转置需求直接使用a_trans/b_trans标志而不是先对操作数.T——详见 matmul.md 中 .Tattribute doesnt exist on PyPTO tensors 一节。3.3 matmul 转置标志的正确用法由于.T不可靠matmul 中的转置应通过a_trans/b_trans表达四种组合覆盖全部场景# a b.T → a_transFalse, b_transTrue result pypto.matmul(a, b, dtype, a_transFalse, b_transTrue) # a.T b → a_transTrue, b_transFalse result pypto.matmul(a, b, dtype, a_transTrue, b_transFalse) # a b → 两者默认 False result pypto.matmul(a, b, dtype) # a.T b.T → 两者均为 True result pypto.matmul(a, b, dtype, a_transTrue, b_transTrue)注意pypto.matmul(a, b.T, out_dtypepypto.DT_FP32)这类传out_dtype关键字的写法是错误语法正确的完整调用形如pypto.matmul(a, b, pypto.DT_FP32, a_transFalse, b_transTrue)dtype 作为位置参数转置走标志位。3.4 2D 张量显式转置的额外限制在 tiling 系统下pypto.transpose(tensor, 0, 1)对2D 张量可能并不生效会触发 TileShape dim num should same to input 错误。处理 2D 矩阵转置时优先考虑# 不要这样2D 转置在 tiling 系统下可能失败 kc_t pypto.transpose(kc, 0, 1) qk pypto.matmul(qc, kc_t, dtype) # 应该这样——用 b_trans 标志 qk pypto.matmul(qc, kc, dtype, a_transFalse, b_transTrue)对称和m_mat m_mat.T的等价实现由于显式转置不可用可拆成两次 matmul 累加dk_c dk_c pypto.matmul(m_mat, kc, dtype) # m_mat kc dk_c dk_c pypto.matmul(m_mat, kc, dtype, a_transTrue, b_transFalse) # m_mat.T kc4. 张量运算相关的高频排错速查与张量运算强相关、在同目录姊妹文档中沉淀的排错规律如下可直接对照定位问题4.1 pypto.assemble 形状不匹配报错CHECK FAILED: dest.GetShape().size() tensor.GetShape().size() Assemble: src and dest requires same shape根因pypto.assemble(tensor, offsets, dest)要求 src 张量的维度数必须与目标位置 view 的维度数一致src 形状不匹配目标 view 维度时会触发该错误。解决reshape 输出张量以匹配目标 view 形状# 错误——out_chunk 是 [bt, V]但 assemble 期望 [1, bt, 1, V] pypto.assemble(out_chunk, [b, t0, h, 0], out_out) # 正确——先 reshape 成与 view 维度一致 pypto.assemble(out_chunk.reshape([1, bt, 1, V]), [b, t0, h, 0], out_out)4.2 reshape([1]) 应用于多元素张量报错CHECK FAILED: capacity 1 Shape size not match, func CheckAndInferShape根因试图把容量大于 1 的张量 reshape 成单元素形状。解决用pypto.view先按偏移取出单个元素再 reshape# 错误——bt4 的张量容量为 4不能 reshape 成 [1] gl g_cum.reshape([1])[bt - 1:bt].reshape([1]) # 正确——先 view 出最后一个元素再 reshape gl pypto.view(g_cum, [1], [bt - 1]).reshape([1])取 1D 张量最后一个元素的通用模式last_elem pypto.view(tensor, [1], [bt - 1]).reshape([1])。4.3 规约轴需要 32 字节对齐报错Reduce op: the tileShape of last axis need to 32Byte align!根因PyPTO 的规约运算.sum()、.mean()要求参与规约的维度按 32 字节对齐。FP32 下即dim * 4必须能被 32 整除等价于维度取 8 的倍数维度字节数FP32是否对齐bt416 B❌ 不对齐bt832 B✅ 对齐V1664 B✅ 对齐V32 / K32128 B✅ 对齐# 错误——bt4 导致对齐错误 bt 4 result tensor.sum(-1) # 正确——bt8 满足 32 字节对齐 bt 8 result tensor.sum(-1)规则任何参与.sum()/.mean()等规约的维度须满足(dim * bytes_per_element) % 32 0FP32 下使用 8 的倍数维度。4.4 对齐维度下 sum 仍失败 → matmul 规约兜底即使维度理论上对齐如 V32由于 PyPTO 内部 tiling 方式.sum(-1)仍可能失败。此时可用预计算 ones 向量 matmul替代规约# 在 host 端创建 ones 向量任意维度均可不要求对齐 ones_v torch.ones(V, 1, devicedevice, dtypetorch.float32) # [V, 1] ones_k torch.ones(K, 1, devicedevice, dtypetorch.float32) # [K, 1] # 在 kernel 签名中把 ones 向量作为参数传入 pypto.frontend.jit(...) def kernel(..., ones_v: pypto.Tensor([], pypto.DT_FP32), ones_k: pypto.Tensor([], pypto.DT_FP32), ...): # 替换 .sum(-1) db_c (dvb * vc).sum(-1) # 错误写法可能失败 db_c pypto.matmul(dvb * vc, ones_v, pypto.DT_FP32).reshape([bt]) # 正确 # 替换 .sum(0) 与 .sum(1) d_l_l_mat_0 pypto.matmul(d_l * l_mat, ones_k, pypto.DT_FP32, a_transTrue, b_transFalse).reshape([bt]) # 行和 d_l_l_mat_1 pypto.matmul(d_l * l_mat, ones_k, pypto.DT_FP32).reshape([bt]) # 列和关键洞察matmul 走 cube 运算不依赖向量单元的 32 字节对齐约束而.sum()走向量运算要求严格对齐。ones 向量第一维必须等于被规约维度否则会触发 K-dimension valid shape mismatch 错误——因此对不同维度V / K / bt应分别准备对应的 ones 向量。4.5 5D 视图 4D vec tile shapes 触发 Run pass failed报错Errcode: FFFFFF! Run pass failed., func CompileFunction根因使用 5D 张量视图如pypto.view(A_in, [1, 1, 1, bt, bt], [b, h, c, 0, 0])时pypto.set_vec_tile_shapes仅配置了 4D tile shapes框架无法处理 5D 运算。解决在 host_wrapper 中把 5D 张量 reshape 成 2D 再传入 kernel并用 2D view 配合 2 个偏移# host 端5D/4D/3D → 2D A_2d A_5d.reshape([B * H * nt, bt * bt]) w_2d w_4d.reshape([B * H * nt, bt * K]) g_cum_2d g_cum_3d.reshape([B * H * nt, bt]) # kernel 内2D viewlen(shape) len(offsets) 必须成立 for c in pypto.loop(0, nt, 1): cache_idx session_base c a pypto.view(A_2d, [bt, bt], [cache_idx, 0]).reshape([bt, bt]) g_cum pypto.view(g_cum_2d, [1, bt], [cache_idx, 0]).reshape([bt])规则pypto.view强制len(shape) len(offsets)kernel 内所有张量尽量保持 ≤4D1D 视图从 2D 张量切出时用shape[1, size]配合 2 个偏移。更完整的 view 指南1 填充、offset 匹配、常见错误表见 pypto-view.md。5. 与 Python 运算符写法的取舍tensor-ops.md 中pypto.mul的显式写法与 python-operators.md§9.14的发现并不冲突PyPTO JIT 内Python 运算符与链式方法是被支持的推荐用于更简洁的代码例如# 冗余写法 result pypto.mul(pypto.mul(a, b), pypto.add(c, d)) result pypto.sum(x, dim-1, keepdimTrue) result pypto.rsqrt(pypto.add(sum_sq, eps)) # 推荐写法等价且更清晰 result a * b * (c d) result x.sum(-1, keepdimTrue) result (sum_sq eps).rsqrt()两条纪律始终成立.sum()的规约轴须 32 字节对齐不满足时走 matmul 规约兜底.T只在 PyTorch 后端张量上可用PyPTO 中间张量请改用pypto.transpose或 matmul 的a_trans/b_trans。6. 核心要点总结matmul 维度约束pypto.matmul要求两个操作数均为 2D1D 向量先.reshape([N, 1])结果再.reshape([N])。广播靠显式 reshape标量参与逐元素运算时先扩展维度如[bt, 1]或用 Python 运算符与链式方法简化。初始化与逐元素pypto.zeros([M, N], pypto.DT_FP32)取负用pypto.mul(tensor, -1.0)或tensor.neg()。转置三条路.T仅限 PyTorch 后端张量、pypto.transpose(tensor, dim0, dim1)注意 2D 张量在 tiling 系统下的限制、matmul 的a_trans/b_transmatmul 场景的首选。关联排错assemble 维度匹配、reshape([1])容量限制、规约 32 字节对齐、matmul-ones 规约兜底、5D 视图 4D tile 冲突逐一对应 matmul.md 与 pypto-view.md 中的完整展开。如需将上述规则落到完整的 kernel 排错流程可配合 DEBUG_GUIDEBOOK.md 的章节索引§9.7、§9.14、§9.16、§9.19按需取用对应叶子文档。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进