MindSpore 2.3.1 + Ascend 910B:自定义算子在动态Shape场景下 grad 反向传播出现 Shape 推断不一致,报错 InferShape 失败
收藏回复举报
MindSpore 2.3.1 + Ascend 910B:自定义算子在动态Shape场景下 grad 反向传播出现 Shape 推断不一致,报错 InferShape 失败
t('forum.solved') 已解决
发表于2026-03-17 21:37:02
0 查看

最近在复现一个基于稀疏注意力机制的研究项目时遇到了一个比较棘手的底层问题,查阅了官方文档和过往的 Issue,暂时没找到解决方案,特来求助。

【环境信息】

  • OS: EulerOS 2.10 (ARM64)
  • NPU: Ascend 910B (驱动版本 24.1.RC1, CANN 8.0.RC1)
  • MindSpore: 2.3.1 (pip install 方式)
  • Python: 3.9.18
  • 运行模式: context.set_context(mode=context.PYNATIVE_MODE, device_target="Ascend")

【问题描述】
我编写了一个自定义算子 SparseMaskAttention,用于处理变长序列的掩码计算。该算子在 静态Shape 下前向推理(Forward)和反向梯度(Backward)均正常。

但在开启动态Shape功能后:

1dynamic_shape = [None, 128, 64] # 第一维 Sequence Length 动态
2input_x = Tensor(np.random.randn(32, 128, 64).astype(np.float32))
3# 注册动态Shape逻辑...
4output = ops.SparseMaskAttention()(input_x)
5loss = output.sum()
6grad_fn = mindspore.value_and_grad(network, None)(input_x)

执行反向传播 grad_fn 时,程序崩溃,报错信息如下:

1[ERROR] ME(12345:7f8b0c,MainProcess):2026-03-17-21:30:15.892.321 [mindspore/ccsrc/backend/common/infershape.cc:450] InferShapeImpl] 
2The shape inferred by custom op 'SparseMaskAttention' in backward pass does not match the forward output shape. 
3Expected: (None, 128, 64), Got: (32, 128, 64). 
4Please check the 'bprop' implementation or the 'infer_shape' function for dynamic axis handling.

【已尝试的排查步骤】

  1. 检查 infer_shape 函数: 我在自定义算子的 Python 侧重写了 infer_shape,对于动态轴使用了 -1 或 AbstractScalar 进行占位,但似乎只在 Forward 生效,Backward 阶段似乎丢失了动态标记,被固化成了具体数值(32)。
  2. 对比静态图: 切换回 GRAPH_MODE 并关闭动态Shape,一切正常。
  3. 查看 CANN 日志: 开启 export MS_ENABLE_LBK=1,发现 TBE 编译生成的算子二进制文件中,反向算子的输入Shape确实被硬编码了。
  4. 文档查阅: 参考了《自定义算子开发指南》中关于动态Shape的部分,提到需要在 bprop 中显式传递 dyn_shapes 参数,但我按照示例修改后,报错变成了 TypeError: got an unexpected keyword argument 'dyn_shapes',怀疑是 API 在 2.3 版本有变动但未更新文档。

【核心疑问】

  1. 在 MindSpore 2.3+ 版本中,自定义算子支持 PYNATIVE 模式下的动态Shape反向传播,是否需要在 @cu_op 或 TBE 接口定义中添加特殊的属性标记(如 dynamic_rank 或 dynamic_axes)?
  2. 如果是在 Python 侧实现 bprop,如何正确获取当前 Input 的抽象形状(AbstractShape)而不是具体数值,以确保输出的 Gradient Shape 也是动态的?

附件中贴出了简化后的算子代码片段(省略了具体的数学计算逻辑,只保留了 Shape 推断部分):

1# 简化的 infer_shape 逻辑
2def infer_shape(self, input_shape):
3    # 尝试保留动态维
4    if input_shape[0] is None: 
5        return (None, input_shape[1], input_shape[2])
6    return input_shape
7
8# bprop 逻辑
9def bprop(self, input_x, output, dout):
10    # 这里直接返回了 dout,是否导致了 Shape 固化?
11    return (dout,)

是否有大佬遇到过类似 InferShape 在反向阶段动态性丢失的问题?或者是否有最新的 Demo 工程可以参考?

非常感谢!

我要发帖子