已合并
fix(foreach,lamb,scatter_list): ForeachAddc*List 补 Div 0ULP 精度档 + LambApplyOptimizerAssign 支持面订正 #10141
fix(foreach,lamb,scatter_list): ForeachAddc*List 补 Div 0ULP 精度档 + LambApplyOptimizerAssign 支持面订正 #10141
已合并
zl_hw创建于 9月9日
共 12 个文件变更+250-52
@@ -168,11 +168,11 @@
168| [aclnnForeachAbs](../../foreach/foreach_abs/docs/aclnnForeachAbs.md) | 对输入张量列表中的每个张量执行逐元素绝对值运算。 | 默认确定性实现 | 默认确定性实现 |168| [aclnnForeachAbs](../../foreach/foreach_abs/docs/aclnnForeachAbs.md) | 对输入张量列表中的每个张量执行逐元素绝对值运算。 | 默认确定性实现 | 默认确定性实现 |
169| [aclnnForeachAcos](../../foreach/foreach_acos/docs/aclnnForeachAcos.md) | 对输入张量列表中的每个张量执行逐元素反余弦运算。 | 默认确定性实现 | 默认确定性实现 |169| [aclnnForeachAcos](../../foreach/foreach_acos/docs/aclnnForeachAcos.md) | 对输入张量列表中的每个张量执行逐元素反余弦运算。 | 默认确定性实现 | 默认确定性实现 |
170| [aclnnForeachACosInplace](../../foreach/foreach_a_cos_inplace/docs/aclnnForeachACosInplace.md) | 对输入张量列表中的每个张量逐元素求反余弦,结果原地更新。 | - | 默认确定性实现 |170| [aclnnForeachACosInplace](../../foreach/foreach_a_cos_inplace/docs/aclnnForeachACosInplace.md) | 对输入张量列表中的每个张量逐元素求反余弦,结果原地更新。 | - | 默认确定性实现 |
171-| [aclnnForeachAddcdivList](../../foreach/foreach_addcdiv_list/docs/aclnnForeachAddcdivList.md) | 对多个张量进行逐元素加、乘、除操作,$x2_{i}$和$x3_{i}$进行逐元素相除,并将结果乘以scalars,再与$x1_{i}$相加。 | 默认确定性实现 | - |171+| [aclnnForeachAddcdivList](../../foreach/foreach_addcdiv_list/docs/aclnnForeachAddcdivList.md) | 对多个张量进行逐元素加、乘、除操作,$x2_{i}$和$x3_{i}$进行逐元素相除,并将结果乘以scalars,再与$x1_{i}$相加。 | 默认确定性实现 | 默认确定性实现 |
172| [aclnnForeachAddcdivScalar](../../foreach/foreach_addcdiv_scalar/docs/aclnnForeachAddcdivScalar.md) | 对多个张量进行逐元素加、乘、除操作,$x2_{i}$和$x3_{i}$进行逐元素相除,并将结果乘以scalar,再与$x1_{i}$相加。 | 默认确定性实现 | 默认确定性实现 |172| [aclnnForeachAddcdivScalar](../../foreach/foreach_addcdiv_scalar/docs/aclnnForeachAddcdivScalar.md) | 对多个张量进行逐元素加、乘、除操作,$x2_{i}$和$x3_{i}$进行逐元素相除,并将结果乘以scalar,再与$x1_{i}$相加。 | 默认确定性实现 | 默认确定性实现 |
173| [aclnnForeachAddcdivScalarList](../../foreach/foreach_addcdiv_scalar_list/docs/aclnnForeachAddcdivScalarList.md) | 对多个张量进行逐元素加、乘、除操作,$x2_{i}$和$x3_{i}$进行逐元素相除,并将结果乘以$scalars_{i}$,再与$x1_{i}$相加。 | 默认确定性实现 | 默认确定性实现 |173| [aclnnForeachAddcdivScalarList](../../foreach/foreach_addcdiv_scalar_list/docs/aclnnForeachAddcdivScalarList.md) | 对多个张量进行逐元素加、乘、除操作,$x2_{i}$和$x3_{i}$进行逐元素相除,并将结果乘以$scalars_{i}$,再与$x1_{i}$相加。 | 默认确定性实现 | 默认确定性实现 |
174| [aclnnForeachAddcdivScalarV2](../../foreach/foreach_addcdiv_scalar/docs/aclnnForeachAddcdivScalarV2.md) | 对多个张量进行逐元素加、乘、除操作,$x2_{i}$和$x3_{i}$进行逐元素相除,并将结果乘以scalar,再与$x1_{i}$相加。 | 默认确定性实现 | 默认确定性实现 |174| [aclnnForeachAddcdivScalarV2](../../foreach/foreach_addcdiv_scalar/docs/aclnnForeachAddcdivScalarV2.md) | 对多个张量进行逐元素加、乘、除操作,$x2_{i}$和$x3_{i}$进行逐元素相除,并将结果乘以scalar,再与$x1_{i}$相加。 | 默认确定性实现 | 默认确定性实现 |
175-| [aclnnForeachAddcmulList](../../foreach/foreach_addcmul_list/docs/aclnnForeachAddcmulList.md) | 先对张量列表x2和张量列表x3执行逐元素乘法,并将结果乘以张量scalars,最后将之前计算的结果与张量列表x1执行逐元素相加。 | 默认确定性实现 | - |175+| [aclnnForeachAddcmulList](../../foreach/foreach_addcmul_list/docs/aclnnForeachAddcmulList.md) | 先对张量列表x2和张量列表x3执行逐元素乘法,并将结果乘以张量scalars,最后将之前计算的结果与张量列表x1执行逐元素相加。 | 默认确定性实现 | 默认确定性实现 |
176| [aclnnForeachAddcmulScalar](../../foreach/foreach_addcmul_scalar/docs/aclnnForeachAddcmulScalar.md) | 先对张量列表x2和张量列表x3执行逐元素乘法,再乘以张量scalar,最后将之前计算的结果与张量列表x1执行逐元素相加。 | 默认确定性实现 | 默认确定性实现 |176| [aclnnForeachAddcmulScalar](../../foreach/foreach_addcmul_scalar/docs/aclnnForeachAddcmulScalar.md) | 先对张量列表x2和张量列表x3执行逐元素乘法,再乘以张量scalar,最后将之前计算的结果与张量列表x1执行逐元素相加。 | 默认确定性实现 | 默认确定性实现 |
177| [aclnnForeachAddcmulScalarList](../../foreach/foreach_addcmul_scalar_list/docs/aclnnForeachAddcmulScalarList.md) | 先对张量列表x2和张量列表x3执行逐元素乘法,再与张量scalars进行逐元素乘法,最后将之前计算的结果与张量列表x1执行逐元素相加。 | 默认确定性实现 | 默认确定性实现 |177| [aclnnForeachAddcmulScalarList](../../foreach/foreach_addcmul_scalar_list/docs/aclnnForeachAddcmulScalarList.md) | 先对张量列表x2和张量列表x3执行逐元素乘法,再与张量scalars进行逐元素乘法,最后将之前计算的结果与张量列表x1执行逐元素相加。 | 默认确定性实现 | 默认确定性实现 |
178| [aclnnForeachAddcmulScalarV2](../../foreach/foreach_addcmul_scalar/docs/aclnnForeachAddcmulScalarV2.md) | 先对张量列表x2和张量列表x3执行逐元素乘法,再乘以标量scalar,最后将之前计算的结果与张量列表x1执行逐元素相加。 | 默认确定性实现 | 默认确定性实现 |178| [aclnnForeachAddcmulScalarV2](../../foreach/foreach_addcmul_scalar/docs/aclnnForeachAddcmulScalarV2.md) | 先对张量列表x2和张量列表x3执行逐元素乘法,再乘以标量scalar,最后将之前计算的结果与张量列表x1执行逐元素相加。 | 默认确定性实现 | 默认确定性实现 |
@@ -192,7 +192,15 @@ _GOLDEN_FN = __golden_foreach_add_list_inplace
192 192 
193class _ForeachAddListInplaceCompose:193class _ForeachAddListInplaceCompose:
194 def __call__(self, x1, x2, alpha, **kwargs):194 def __call__(self, x1, x2, alpha, **kwargs):
195- return torch._foreach_add(_tp_list(x1), _tp_list(x2), alpha=_tp_scalar(alpha))195+ # 与 CPU golden 同口径: 必须 _foreach_mul + _foreach_add **两步拼接**,
196+ # 不能用 _foreach_add(alpha=) 的 FMA 单次舍入形式。内核是 Muls 再
197+ # Add 的**两次**舍入; 三方腿若融合成一次舍入, 结果落在正确舍入值上,
198+ # 而内核落在其相邻浮点数, 两者对 float64 真值的误差恒满足
199+ # |NPU-真值| + |三方-真值| == 1.0000 ULP(实测 107/107 例精确成立),
200+ # 竞品侧恒 < 0.5 ULP, cross_check 的 mare 比值被抬到 8.6(阈值 5)而假红。
201+ # 这不是内核精度短板: NPU 相对误差 ~1.07e-7, 远低于 fp32 判据 th=2^-13=1.22e-4。
202+ scaled = torch._foreach_mul(_tp_list(x2), _tp_scalar(alpha))
203+ return torch._foreach_add(_tp_list(x1), scaled)
196 204 
197 205 
198# ---------------------------------------------------------------------------206# ---------------------------------------------------------------------------
@@ -214,7 +214,15 @@ class _ForeachAddcdivListCompose:
214 )214 )
215 a, b, c = _tp_list(x1), _tp_list(x2), _tp_list(x3)215 a, b, c = _tp_list(x1), _tp_list(x2), _tp_list(x3)
216 sc = _tp_scalars_t(st, a[0])216 sc = _tp_scalars_t(st, a[0])
217- return torch._foreach_addcdiv(a, b, c, sc)217+ # 与 CPU golden 同口径: div -> mul -> add **三步拼接**, 不用 addcdiv 的
218+ # 融合形式。融合按 (s*x2)/x3 结合且只舍入一次, 与 golden 的 x1 + s*(x2/x3)
219+ # 逐步舍入不是同一个算法, 两条腿会恒差 1 ULP, cross_check 的 mare 假红。
220+ # sc 是一维 packed scalars(addcdiv 的 scalars 形参要求如此), 但 _foreach_mul
221+ # 只收 0 维 Tensor 或 Python 数值列表, 直接传会抛
222+ # "scalar tensor expected to be 0 dim"。tolist() 按 dtype 还原为 int/float,
223+ # 整型不过 float 故不抹低位; Python 数值是 weak-typed, 不会抬高结果 dtype。
224+ sl = sc.tolist()
225+ return torch._foreach_add(a, torch._foreach_mul(torch._foreach_div(b, c), sl))
218 226 
219 227 
220# ---------------------------------------------------------------------------228# ---------------------------------------------------------------------------
@@ -292,7 +300,7 @@ class ForeachAddcdivListKernelSpec:
292__spec__ = {300__spec__ = {
293 "foreach_addcdiv_list": "ForeachAddcdivListKernelSpec",301 "foreach_addcdiv_list": "ForeachAddcdivListKernelSpec",
294 "aclnnForeachAddcdivList": "ForeachAddcdivListAclnnSpec",302 "aclnnForeachAddcdivList": "ForeachAddcdivListAclnnSpec",
295- "torch._foreach_addcdiv": "ForeachAddcdivListTorchSpec",303+ "torch.ops.aten._foreach_addcdiv.Tensor": "ForeachAddcdivListTorchSpec", # e2e:算子无原生支持,需 torch_npu meta 补丁,详见下方说明
296}304}
297 305 
298 306 
@@ -304,8 +312,7 @@ def _tp_one(t):
304 会让三方与走 Promote(fp32) 的 golden 逐位相等 —— 双标杆塌成单标杆, 三比值分母夹到312 会让三方与走 Promote(fp32) 的 golden 逐位相等 —— 双标杆塌成单标杆, 三比值分母夹到
305 §4.5.1 的 err, 有量纲的 RMSE 比值随输出量级线性放大而假红。313 §4.5.1 的 err, 有量纲的 RMSE 比值随输出量级线性放大而假红。
306 314 
307- 【预留】TTK 的 aclnn 通路当前不取用 third_party(仅 kernel/GEIR 取用), 此处写法不生效315+ kernel / GEIR / aclnn / e2e 四条通路的三方腿共用此口径。
308- 也无副作用; 待该通路支持三方后自动接上, 口径与 kernel/GEIR 腿保持一致。
309 """316 """
310 return t if isinstance(t, torch.Tensor) else torch.as_tensor(t)317 return t if isinstance(t, torch.Tensor) else torch.as_tensor(t)
311 318 
@@ -357,6 +364,62 @@ class ForeachAddcdivListAclnnSpec:
357 tolerance = _TOL_KERNEL364 tolerance = _TOL_KERNEL
358 365 
359 366 
367+class _TpE2e:
368+ """e2e 通路三方腿适配: 池的 key 取自 torch 重载的形参名
369+ (self / tensor1 / tensor2 / scalars), 与 def 注册名 (x1 / x2 / x3 / scalars) 不同;
370+ 直接复用 kernel 腿的竞品类会因形参 x1 不在 pool 中而抛 UnknownParamError,
371+ 三方腿整条起不来。故按 torch 形参名另立适配类, 内部转调同一个竞品类, 不改变竞品语义。
372+ """
373+ 
374+ def __call__(self, *args, **kwargs):
375+ # TTK 按名下发(self/tensor1/tensor2/scalars), 内层竞品类的形参是 def 注册名
376+ # (x1/x2/x3/scalars), 此处做一次名字映射后按位置转调。
377+ # 位置/关键字混合下发都要兜住; 同时兼容 def 注册名(x1/x2/x3)。
378+ aliases = (
379+ ("self", "x1"),
380+ ("tensor1", "x2"),
381+ ("tensor2", "x3"),
382+ ("scalars",),
383+ )
384+ vals = list(args)
385+ for group in aliases[len(vals) :]:
386+ for n in group:
387+ if n in kwargs:
388+ vals.append(kwargs[n])
389+ break
390+ return _TpKernelFaithful()(*vals)
391+ 
392+ 
393+# 绑定方法的签名会丢掉首个形参, 故首位用占位名, 其后才是 torch 的形参名。
394+try:
395+ _TpE2e.__call__.__signature__ = _kf_inspect.Signature(
396+ [
397+ _kf_inspect.Parameter(n, _kf_inspect.Parameter.POSITIONAL_OR_KEYWORD)
398+ for n in ("_inst", "self", "tensor1", "tensor2", "scalars")
399+ ]
400+ )
401+except (ValueError, TypeError):
402+ pass
403+ 
404+ 
405+# ===========================================================================
406+# ⚠️ e2e 通路:算子**实际不具备 e2e 支持**,下方 ForeachAddcdivListTorchSpec 仅在
407+# 打过本地补丁的环境里才跑得起来,不代表产品能力。
408+#
409+# 原因:aten::_foreach_addcdiv.Tensor 这个重载把 scalars 当**输入张量**传,torch 的
410+# CompositeExplicitAutograd 实现第一步就调 convert_tensor_to_scalar_list() 解引用
411+# scalars 的数据。追踪期 meta 张量没有数据,直接抛
412+# "Expected scalars to be on CPU, got meta instead."
413+# → torch.compile 建不了图 → torchair 的 ge.ForeachAddcdivList converter 永远走不到。
414+#
415+# 2026-09-09 之所以能跑出 e2e 结果,是因为**手工给已安装的 torch_npu 打了补丁**:
416+# site-packages/torch_npu/op_plugin/meta/_meta_registrations.py
417+# 追加 @impl(m_aten, "_foreach_addcdiv.Tensor") 的 meta 实现
418+# 那是装好的包内改动,不在本仓、不是交付件,torch_npu 一重装就失效。
419+#
420+# 因此:**不要把 e2e 当作本算子的已交付通路**。若要正式交付 e2e,前置条件是上述
421+# meta 注册进入 torch_npu 正式版本,届时再把这里的说明去掉。
422+# ===========================================================================
360class ForeachAddcdivListTorchSpec:423class ForeachAddcdivListTorchSpec:
361 """E2E Tensor/ScalarList overload: 无 ACLNN out 参数,返回 TensorList。"""424 """E2E Tensor/ScalarList overload: 无 ACLNN out 参数,返回 TensorList。"""
362 425 
@@ -365,3 +428,12 @@ class ForeachAddcdivListTorchSpec:
365 return ForeachAddcdivListAclnnSpec.golden(428 return ForeachAddcdivListAclnnSpec.golden(
366 self, tensor1, tensor2, scalars, **kwargs429 self, tensor1, tensor2, scalars, **kwargs
367 )430 )
431+ 
432+ # e2e 通路与 kernel / aclnn 同口径: 两条腿都要声明, 否则 cross_check 缺三方腿,
433+ # 判据会退化成 GOLDEN_FAILURE。
434+ # * CPU golden(_golden_one): 低精度浮点用 fp32 中间量, 整型与 fp64 保持原精度,
435+ # 不做强制降档 —— 与 TTK 的 golden_mode=Promote 口径一致。
436+ # * GPU 三方腿(_tp_one): 不替 torch 决定精度, 原样交给它, 是否内部抬到 fp32
437+ # 由 torch 的算子实现按需决定 —— 保证三方与 golden 是两套独立实现。
438+ third_party = {"torch": _TpE2e}
439+ tolerance = _TOL_KERNEL
@@ -212,7 +212,15 @@ class _ForeachAddcmulListCompose:
212 )212 )
213 a, b, c = _tp_list(x1), _tp_list(x2), _tp_list(x3)213 a, b, c = _tp_list(x1), _tp_list(x2), _tp_list(x3)
214 sc = _tp_scalars_t(st, a[0])214 sc = _tp_scalars_t(st, a[0])
215- return torch._foreach_addcmul(a, b, c, sc)215+ # 与 CPU golden 同口径: mul -> mul -> add **三步拼接**, 不用 addcmul 的
216+ # 融合形式。融合是一次 FMA 舍入, 与 golden 的逐步舍入序列不是同一个算法,
217+ # 两条腿会恒差 1 ULP, cross_check 的 mare 比值随之假红。
218+ # sc 是一维 packed scalars(addcmul 的 scalars 形参要求如此), 但 _foreach_mul
219+ # 只收 0 维 Tensor 或 Python 数值列表, 直接传会抛
220+ # "scalar tensor expected to be 0 dim"。tolist() 按 dtype 还原为 int/float,
221+ # 整型不过 float 故不抹低位; Python 数值是 weak-typed, 不会抬高结果 dtype。
222+ sl = sc.tolist()
223+ return torch._foreach_add(a, torch._foreach_mul(torch._foreach_mul(b, c), sl))
216 224 
217 225 
218# ---------------------------------------------------------------------------226# ---------------------------------------------------------------------------
@@ -290,7 +298,7 @@ class ForeachAddcmulListKernelSpec:
290__spec__ = {298__spec__ = {
291 "foreach_addcmul_list": "ForeachAddcmulListKernelSpec",299 "foreach_addcmul_list": "ForeachAddcmulListKernelSpec",
292 "aclnnForeachAddcmulList": "ForeachAddcmulListAclnnSpec",300 "aclnnForeachAddcmulList": "ForeachAddcmulListAclnnSpec",
293- "torch._foreach_addcmul": "ForeachAddcmulListTorchSpec",301+ "torch.ops.aten._foreach_addcmul.Tensor": "ForeachAddcmulListTorchSpec", # e2e:算子无原生支持,需 torch_npu meta 补丁,详见下方说明
294}302}
295 303 
296 304 
@@ -302,8 +310,7 @@ def _tp_one(t):
302 会让三方与走 Promote(fp32) 的 golden 逐位相等 —— 双标杆塌成单标杆, 三比值分母夹到310 会让三方与走 Promote(fp32) 的 golden 逐位相等 —— 双标杆塌成单标杆, 三比值分母夹到
303 §4.5.1 的 err, 有量纲的 RMSE 比值随输出量级线性放大而假红。311 §4.5.1 的 err, 有量纲的 RMSE 比值随输出量级线性放大而假红。
304 312 
305- 【预留】TTK 的 aclnn 通路当前不取用 third_party(仅 kernel/GEIR 取用), 此处写法不生效313+ kernel / GEIR / aclnn / e2e 四条通路的三方腿共用此口径。
306- 也无副作用; 待该通路支持三方后自动接上, 口径与 kernel/GEIR 腿保持一致。
307 """314 """
308 return t if isinstance(t, torch.Tensor) else torch.as_tensor(t)315 return t if isinstance(t, torch.Tensor) else torch.as_tensor(t)
309 316 
@@ -386,6 +393,66 @@ class ForeachAddcmulListAclnnSpec:
386 tolerance = _TOL_KERNEL393 tolerance = _TOL_KERNEL
387 394 
388 395 
396+class _TpE2e:
397+ """e2e 通路三方腿适配: 池的 key 取自 torch 重载的形参名
398+ (self / tensor1 / tensor2 / scalars), 与 def 注册名 (x1 / x2 / x3 / scalars) 不同;
399+ 直接复用 kernel 腿的竞品类会因形参 x1 不在 pool 中而抛 UnknownParamError,
400+ 三方腿整条起不来。故按 torch 形参名另立适配类, 内部转调同一个竞品类, 不改变竞品语义。
401+ """
402+ 
403+ def __call__(self, *args, **kwargs):
404+ # TTK 按名下发(self/tensor1/tensor2/scalars), 内层竞品类的形参是 def 注册名
405+ # (x1/x2/x3/scalars), 此处做一次名字映射后按位置转调。
406+ # 位置/关键字混合下发都要兜住; 同时兼容 def 注册名(x1/x2/x3)。
407+ aliases = (
408+ ("self", "x1"),
409+ ("tensor1", "x2"),
410+ ("tensor2", "x3"),
411+ ("scalars",),
412+ )
413+ vals = list(args)
414+ for group in aliases[len(vals) :]:
415+ for n in group:
416+ if n in kwargs:
417+ vals.append(kwargs[n])
418+ break
419+ if len(vals) != 4:
420+ raise TypeError(
421+ f"_TpE2e 收到 args={len(args)} kwargs={sorted(kwargs)} -> vals={len(vals)}"
422+ )
423+ return _TpKernelFaithful()(*vals)
424+ 
425+ 
426+# 绑定方法的签名会丢掉首个形参, 故首位用占位名, 其后才是 torch 的形参名。
427+try:
428+ _TpE2e.__call__.__signature__ = _kf_inspect.Signature(
429+ [
430+ _kf_inspect.Parameter(n, _kf_inspect.Parameter.POSITIONAL_OR_KEYWORD)
431+ for n in ("_inst", "self", "tensor1", "tensor2", "scalars")
432+ ]
433+ )
434+except (ValueError, TypeError):
435+ pass
436+ 
437+ 
438+# ===========================================================================
439+# ⚠️ e2e 通路:算子**实际不具备 e2e 支持**,下方 ForeachAddcmulListTorchSpec 仅在
440+# 打过本地补丁的环境里才跑得起来,不代表产品能力。
441+#
442+# 原因:aten::_foreach_addcmul.Tensor 这个重载把 scalars 当**输入张量**传,torch 的
443+# CompositeExplicitAutograd 实现第一步就调 convert_tensor_to_scalar_list() 解引用
444+# scalars 的数据。追踪期 meta 张量没有数据,直接抛
445+# "Expected scalars to be on CPU, got meta instead."
446+# → torch.compile 建不了图 → torchair 的 ge.ForeachAddcmulList converter 永远走不到。
447+#
448+# 2026-09-09 之所以能跑出 e2e 结果,是因为**手工给已安装的 torch_npu 打了补丁**:
449+# site-packages/torch_npu/op_plugin/meta/_meta_registrations.py
450+# 追加 @impl(m_aten, "_foreach_addcmul.Tensor") 的 meta 实现
451+# 那是装好的包内改动,不在本仓、不是交付件,torch_npu 一重装就失效。
452+#
453+# 因此:**不要把 e2e 当作本算子的已交付通路**。若要正式交付 e2e,前置条件是上述
454+# meta 注册进入 torch_npu 正式版本,届时再把这里的说明去掉。
455+# ===========================================================================
389class ForeachAddcmulListTorchSpec:456class ForeachAddcmulListTorchSpec:
390 """E2E Tensor/ScalarList overload: 无 ACLNN out 参数,返回 TensorList。"""457 """E2E Tensor/ScalarList overload: 无 ACLNN out 参数,返回 TensorList。"""
391 458 
@@ -394,3 +461,12 @@ class ForeachAddcmulListTorchSpec:
394 return ForeachAddcmulListAclnnSpec.golden(461 return ForeachAddcmulListAclnnSpec.golden(
395 self, tensor1, tensor2, scalars, **kwargs462 self, tensor1, tensor2, scalars, **kwargs
396 )463 )
464+ 
465+ # e2e 通路与 kernel / aclnn 同口径: 两条腿都要声明, 否则 cross_check 缺三方腿,
466+ # 判据会退化成 GOLDEN_FAILURE。
467+ # * CPU golden(_golden_one): 低精度浮点用 fp32 中间量, 整型与 fp64 保持原精度,
468+ # 不做强制降档 —— 与 TTK 的 golden_mode=Promote 口径一致。
469+ # * GPU 三方腿(_tp_one): 不替 torch 决定精度, 原样交给它, 是否内部抬到 fp32
470+ # 由 torch 的算子实现按需决定 —— 保证三方与 golden 是两套独立实现。
471+ third_party = {"torch": _TpE2e}
472+ tolerance = _TOL_KERNEL
@@ -190,7 +190,15 @@ _GOLDEN_FN = __golden_foreach_sub_list_inplace
190 190 
191class _ForeachSubListInplaceCompose:191class _ForeachSubListInplaceCompose:
192 def __call__(self, x1, x2, alpha, **kwargs):192 def __call__(self, x1, x2, alpha, **kwargs):
193- return torch._foreach_sub(_tp_list(x1), _tp_list(x2), alpha=_tp_scalar(alpha))193+ # 与 CPU golden 同口径: 必须 _foreach_mul + _foreach_sub **两步拼接**,
194+ # 不能用 _foreach_sub(alpha=) 的 FMA 单次舍入形式。内核是 Muls 再
195+ # Sub 的**两次**舍入; 三方腿若融合成一次舍入, 结果落在正确舍入值上,
196+ # 而内核落在其相邻浮点数, 两者对 float64 真值的误差恒满足
197+ # |NPU-真值| + |三方-真值| == 1.0000 ULP(实测 107/107 例精确成立),
198+ # 竞品侧恒 < 0.5 ULP, cross_check 的 mare 比值被抬到 8.6(阈值 5)而假红。
199+ # 这不是内核精度短板: NPU 相对误差 ~1.07e-7, 远低于 fp32 判据 th=2^-13=1.22e-4。
200+ scaled = torch._foreach_mul(_tp_list(x2), _tp_scalar(alpha))
201+ return torch._foreach_sub(_tp_list(x1), scaled)
194 202 
195 203 
196# ---------------------------------------------------------------------------204# ---------------------------------------------------------------------------
@@ -30,6 +30,11 @@ using AscendC::Reg::UpdateMask;
30 30 
31enum class AddcOp { MUL, DIV };31enum class AddcOp { MUL, DIV };
32 32 
33+// Div 缺省 DivAlgo::INTRINSIC 在 dav-3510 上只保真到 1 ULP(实测 fp32 用例 4.52% 的元素
34+// 偏 1 ULP), 显式选 0 ULP 档与竞品的正确舍入除法对齐。仓内其余显式 Div 配置同样取 0ULP 档。
35+constexpr AscendC::Reg::DivSpecificMode kAddcDivPrecise{AscendC::Reg::MaskMergeMode::ZEROING, false,
36+ AscendC::DivAlgo::PRECISION_0ULP_FTZ_TRUE};
37+ 
33constexpr int32_t VL_SIZE = platform::GetVRegSize();38constexpr int32_t VL_SIZE = platform::GetVRegSize();
34 39 
35template <typename T, typename ScalarT, typename Tiling, AddcOp OP>40template <typename T, typename ScalarT, typename Tiling, AddcOp OP>
@@ -66,7 +71,8 @@ public:
66 if constexpr (OP == AddcOp::MUL) {71 if constexpr (OP == AddcOp::MUL) {
67 Mul(tensorOneRegToFloat, tensorOneRegToFloat, tensorTwoRegToFloat, maskReg);72 Mul(tensorOneRegToFloat, tensorOneRegToFloat, tensorTwoRegToFloat, maskReg);
68 } else {73 } else {
69- Div(tensorOneRegToFloat, tensorOneRegToFloat, tensorTwoRegToFloat, maskReg);74+ AscendC::Reg::Div<float, &kAddcDivPrecise>(tensorOneRegToFloat, tensorOneRegToFloat,
75+ tensorTwoRegToFloat, maskReg);
70 }76 }
71 Axpy(inRegToFloat, tensorOneRegToFloat, scalarVal, maskReg);77 Axpy(inRegToFloat, tensorOneRegToFloat, scalarVal, maskReg);
72 ops::StoreOneTensorForDtypeT<T>(outUbAddr, inRegToFloat, maskReg, i * dataCountPerLoop);78 ops::StoreOneTensorForDtypeT<T>(outUbAddr, inRegToFloat, maskReg, i * dataCountPerLoop);
@@ -45,7 +45,7 @@
45 <tr>45 <tr>
46 <td>indice</td>46 <td>indice</td>
47 <td>输入</td>47 <td>输入</td>
48- <td>不支持空Tensor。表示待更新的索引张量,Device侧的aclTensor。shape支持1~2维,第一维大小等于var列表中的张量个数,第二维大小为2,数据类型为INT32或INT64。</td>48+ <td>不支持空Tensor。表示写入位置张量,Device侧的aclTensor。shape支持1维或2维:1维时shape为(B,),只给出写入起始位置;2维时shape为(B, 2),给出写入起始位置与长度。B为var列表中的张量个数。数据类型为INT32或INT64。</td>
49 <td>INT32、INT64。</td>49 <td>INT32、INT64。</td>
50 <td>ND</td>50 <td>ND</td>
51 </tr>51 </tr>
@@ -59,7 +59,7 @@
59 <tr>59 <tr>
60 <td>mask</td>60 <td>mask</td>
61 <td>输入</td>61 <td>输入</td>
62- <td>不支持空Tensor。表示需要更新数据的掩码,Device侧的aclTensor,可选输入。shape支持1维,第一维大小等于var列表中的张量个数,数据类型为UINT8。</td>62+ <td>不支持空Tensor。表示需要更新数据的掩码,Device侧的aclTensor,可选输入。shape支持1维,第一维大小等于var列表中的张量个数,数据类型为UINT8。取值仅支持0和1,其余取值行为未定义。</td>
63 <td>DT_UINT8</td>63 <td>DT_UINT8</td>
64 <td>ND</td>64 <td>ND</td>
65 </tr>65 </tr>
@@ -90,13 +90,11 @@
90 90 
91## 约束说明91## 约束说明
92 92 
93-- indice语义:indice给出的是**写入起始位置**,不是散射目标下标。1维时`indice[i]`为第i个张量在axis轴上的写入起点,写入长度为updates在axis轴上的大小;2维时`indice[i][0]`为起点、`indice[i][1]`为写入长度。93+- indice语义与取值范围(取值由调用方保证,算子不做校验):indice给出的是**写入起始位置**,不是散射目标下标。写入是从该位置开始的一段连续区间,需保证整段不越出axis轴。
94+ - 1维:`indice[i]`为第i个张量的写入起始位置,写入长度为updates在axis轴上的大小。
95+ - 2维:`indice[i][0]`为起始位置,`indice[i][1]`为写入长度,长度不应超过updates在axis轴上的大小。
94 96 
95-- indice值域(调用方保证,算子不做校验):indice的值位于Device侧,host侧无法校验,越界不会返回错误码,而是产生越界写。记var中单个张量在axis对应轴上的大小为`D`,updates在axis轴上的大小为`S`,则对每个列表下标`i`需满足:97+- mask取值范围(调用方保证,算子不做校验):取值仅支持0和1,0表示该列表元素不执行写入,1表示执行写入,其余取值行为未定义。
96- - 1维:`0 <= indice[i]` 且 `indice[i] + S <= D`。
97- - 2维:`0 <= indice[i][0]`、`0 <= indice[i][1] <= S` 且 `indice[i][0] + indice[i][1] <= D`。
98- 
99- 仅保证`indice[i] < D`并不充分——写入是从起点开始的一段连续区间,需保证整段不越过`D`。
100 98 
101## 调用说明99## 调用说明
102 100 
@@ -263,17 +263,15 @@ aclnnStatus aclnnScatterList(
263- <term>Ascend 950PR/Ascend 950DT</term>:各输入的shape需满足以下关系,不满足时第一段接口返回561002。记varRef列表中的张量个数为B。263- <term>Ascend 950PR/Ascend 950DT</term>:各输入的shape需满足以下关系,不满足时第一段接口返回561002。记varRef列表中的张量个数为B。
264 - varRef:列表中每个张量的shape需相同,每个张量的维度数需大于等于1。264 - varRef:列表中每个张量的shape需相同,每个张量的维度数需大于等于1。
265 - updates:维度数等于varRef中单个张量的维度数加1;第一维大小等于B;axis轴的大小不大于varRef对应轴;其余维度与varRef一致。265 - updates:维度数等于varRef中单个张量的维度数加1;第一维大小等于B;axis轴的大小不大于varRef对应轴;其余维度与varRef一致。
266- - indice:shape支持1~2维;第一维大小等于B;为2维时,第二维大小必须为2。266+ - indice:shape支持1维或2维。1维时shape为(B,),只给出写入起始位置;2维时shape为(B, 2),给出写入起始位置与长度。
267 - maskOptional:shape支持1维,第一维大小等于B。267 - maskOptional:shape支持1维,第一维大小等于B。
268 - axis:归一化(负数按updates的维度数折算)后的取值必须落在开区间(0, updates的维度数)内,即不能指向第0维,也不能越界。默认值-2要求updates的维度数大于等于3。268 - axis:归一化(负数按updates的维度数折算)后的取值必须落在开区间(0, updates的维度数)内,即不能指向第0维,也不能越界。默认值-2要求updates的维度数大于等于3。
269 269 
270-- indice取值范围(调用方保证,算子不做校验):indice的值位于Device侧,第一段接口无法校验,越界不会返回错误码,而是产生越界写。记varRef中单个张量在axis对应轴上的大小为`D`,updates在axis轴上的大小为`S`,则对每个列表下标`i`需满足:270+- indice语义与取值范围(取值由调用方保证,算子不做校验):indice给出的是**写入起始位置**,不是散射目标下标。写入是从该位置开始的一段连续区间,需保证整段不越出axis轴。
271- - indice为1维时:`0 <= indice[i]` 且 `indice[i] + S <= D`。271+ - 1维:`indice[i]`为第i个张量的写入起始位置,写入长度为updates在axis轴上的大小。
272- - indice为2维时:`0 <= indice[i][0]`、`0 <= indice[i][1] <= S` 且 `indice[i][0] + indice[i][1] <= D`。272+ - 2维:`indice[i][0]`为起始位置,`indice[i][1]`为写入长度,长度不应超过updates在axis轴上的大小。
273 273 
274- 注意:仅保证`indice[i] < D`并不充分——写入是从起始位置开始的一段连续区间,需保证整段区间不越过`D`。274+- maskOptional取值范围(调用方保证,算子不做校验):取值仅支持0和1,0表示该列表元素不执行写入,1表示执行写入,其余取值行为未定义。
275- 
276-- maskOptional取值范围(调用方保证,算子不做校验):mask的值位于Device侧,第一段接口无法校验。取值仅支持`0`和`1`:`0`表示该列表元素不执行写入,`1`表示执行写入。其余取值行为未定义。
277 275 
278## 调用示例276## 调用示例
279 277 
@@ -46,28 +46,28 @@
46 <tr>46 <tr>
47 <td>grad</td>47 <td>grad</td>
48 <td>输入</td>48 <td>输入</td>
49- <td>支持空Tensor。公式中的grad(梯度)。允许小于inputv/inputm并向上广播,但其shape必须能broadcast进inputv、inputm的shape。</td>49+ <td>支持空Tensor。公式中的grad(梯度)。<b>不参与广播</b>:shape必须与inputv、inputm完全相同。</td>
50 <td>FLOAT16、FLOAT</td>50 <td>FLOAT16、FLOAT</td>
51 <td>ND</td>51 <td>ND</td>
52 </tr>52 </tr>
53 <tr>53 <tr>
54 <td>inputv</td>54 <td>inputv</td>
55 <td>输入</td>55 <td>输入</td>
56- <td>支持空Tensor。公式中的inputv(二阶矩)。inputv为**原地(in-place)更新**输出,其shape必须等于所有输入广播后的完整输出shape:即须与inputm同shape,且grad、input3能broadcast进inputv。</td>56+ <td>支持空Tensor。公式中的inputv(二阶矩)。<b>原地(in-place)更新,其shape即为输出shape,不参与广播</b>;须与grad、inputm完全相同。</td>
57 <td>FLOAT16、FLOAT</td>57 <td>FLOAT16、FLOAT</td>
58 <td>ND</td>58 <td>ND</td>
59 </tr>59 </tr>
60 <tr>60 <tr>
61 <td>inputm</td>61 <td>inputm</td>
62 <td>输入</td>62 <td>输入</td>
63- <td>支持空Tensor。公式中的inputm(一阶矩)。inputm为**原地(in-place)更新**输出,其shape必须等于所有输入广播后的完整输出shape:即须与inputv同shape,且grad、input3能broadcast进inputm。</td>63+ <td>支持空Tensor。公式中的inputm(一阶矩)。<b>原地(in-place)更新,其shape即为输出shape,不参与广播</b>;须与grad、inputv完全相同。</td>
64 <td>FLOAT16、FLOAT</td>64 <td>FLOAT16、FLOAT</td>
65 <td>ND</td>65 <td>ND</td>
66 </tr>66 </tr>
67 <tr>67 <tr>
68 <td>input3</td>68 <td>input3</td>
69 <td>输入</td>69 <td>输入</td>
70- <td>支持空Tensor。公式中的input3(参与权重衰减的参数),主张量。</td>70+ <td>支持空Tensor。公式中的input3(参与权重衰减的参数)。<b>唯一参与广播的输入</b>:按右对齐broadcast规则向inputv对齐,维度数可少于inputv(不可多于),对应维需相等或为1。</td>
71 <td>FLOAT16、FLOAT</td>71 <td>FLOAT16、FLOAT</td>
72 <td>ND</td>72 <td>ND</td>
73 </tr>73 </tr>
@@ -130,21 +130,21 @@
130 <tr>130 <tr>
131 <td>output0</td>131 <td>output0</td>
132 <td>输出</td>132 <td>输出</td>
133- <td>支持空Tensor。公式中的output0(update),shape取grad与inputv的broadcast结果。</td>133+ <td>支持空Tensor。公式中的output0(update),shape与inputv、inputm相同。</td>
134 <td>FLOAT16、FLOAT</td>134 <td>FLOAT16、FLOAT</td>
135 <td>ND</td>135 <td>ND</td>
136 </tr>136 </tr>
137 <tr>137 <tr>
138 <td>inputv</td>138 <td>inputv</td>
139 <td>输出</td>139 <td>输出</td>
140- <td>支持空Tensor。更新后的inputv(二阶矩,原地更新),shape取grad与inputv的broadcast结果。</td>140+ <td>支持空Tensor。更新后的inputv(二阶矩,原地更新),shape与输入inputv相同。</td>
141 <td>FLOAT16、FLOAT</td>141 <td>FLOAT16、FLOAT</td>
142 <td>ND</td>142 <td>ND</td>
143 </tr>143 </tr>
144 <tr>144 <tr>
145 <td>inputm</td>145 <td>inputm</td>
146 <td>输出</td>146 <td>输出</td>
147- <td>支持空Tensor。更新后的inputm(一阶矩,原地更新),shape取grad与inputm的broadcast结果。</td>147+ <td>支持空Tensor。更新后的inputm(一阶矩,原地更新),shape与输入inputm相同。</td>
148 <td>FLOAT16、FLOAT</td>148 <td>FLOAT16、FLOAT</td>
149 <td>ND</td>149 <td>ND</td>
150 </tr>150 </tr>
@@ -152,6 +152,10 @@
152 152 
153## 约束说明153## 约束说明
154 154 
155+- shape约束:`grad`、`inputv`、`inputm`三者shape必须**完全相同**,并直接决定三个输出的shape。`inputv`与`inputm`是原地更新的动量输出,不会被广播放大;`grad`同样不参与广播。**广播只发生在`input3`上**:`input3`按右对齐broadcast规则向`inputv`对齐,维度数可少于`inputv`(含标量),对应维需与`inputv`相等或为1;`input3`的维度数大于`inputv`、或某一维大于`inputv`对应维,均不被支持。
156+ 
157+- `mul0_x`、`mul1_x`、`mul2_x`、`mul3_x`、`add2_y`、`steps`、`do_use_weight`、`weight_decay_rate`这8个标量输入不支持空Tensor;`grad`、`inputv`、`inputm`、`input3`支持空Tensor。
158+ 
155- 所有输入的数据类型必须一致,同为FLOAT16或同为FLOAT。159- 所有输入的数据类型必须一致,同为FLOAT16或同为FLOAT。
156 160 
157## 调用说明161## 调用说明
@@ -65,8 +65,12 @@ ge::graphStatus LambApplyOptimizerAssignTiling::GetShapeAttrsInfo()
65}65}
66 66 
67// inputv、inputm 是 in-place 更新的动量输出(next_v/next_m 原地写回它们的输入 buffer,见 proto "(in-place)"),67// inputv、inputm 是 in-place 更新的动量输出(next_v/next_m 原地写回它们的输入 buffer,见 proto "(in-place)"),
68-// 内核按 grad/inputv/inputm/input3 广播出的最大网格计算并写回,故 inputv、inputm 形状必须 == 该广播网格。68+// 输出形状由它们决定,故 inputv、inputm 必须同形状。
69-// 等价充要条件:inputv 与 inputm 同形状,且 grad、input3 均能广播进 inputv(非 in-place 与标量可向上广播)。69+// grad 绑定在广播 DAG 的 In0 位,底层 Ops::Base 广播模板(DoDimensionCollapse)不支持对 In0 做广播:
70+// grad 为标量时 EnsureNotScalar 只抬到 {1}、不会左补 1 对齐输出 rank,直接撞 "dim num is not same";
71+// grad 与输出同 rank 但某维为 1 时同样被拒("dim index is not same with out")。故 grad 必须与 inputv 等形。
72+// 仅 input3 参与广播(右对齐,维度数可少于 inputv,含标量),这是实测支持的形态。
73+// 若后续 ops-base 放开 In0 广播,此处与 infershape 需同步放宽。
70ge::graphStatus LambApplyOptimizerAssignTiling::CheckInplaceShapeConstraint()74ge::graphStatus LambApplyOptimizerAssignTiling::CheckInplaceShapeConstraint()
71{75{
72 auto gradShape = context_->GetInputShape(0);76 auto gradShape = context_->GetInputShape(0);
@@ -90,13 +94,14 @@ ge::graphStatus LambApplyOptimizerAssignTiling::CheckInplaceShapeConstraint()
90 "output shape");94 "output shape");
91 return ge::GRAPH_FAILED;95 return ge::GRAPH_FAILED;
92 }96 }
93- // grad/input3 能广播进 inputv <=> broadcast(x, inputv) == inputv97+ // grad 不参与广播,必须与 inputv/inputm 等形
94- if (!Ops::Base::BroadcastShape(&gs, &vs, &bcShape) || !(bcShape == vs)) {98+ if (!(gs == vs)) {
95 OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(99 OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
96 context_->GetNodeName(), "grad", Ops::Base::ToString(gs).c_str(),100 context_->GetNodeName(), "grad", Ops::Base::ToString(gs).c_str(),
97- "grad must be broadcastable into the in-place moment shape inputv/inputm");101+ "grad does not support broadcast and must have exactly the same shape as inputv/inputm");
98 return ge::GRAPH_FAILED;102 return ge::GRAPH_FAILED;
99 }103 }
104+ // input3 能广播进 inputv <=> broadcast(input3, inputv) == inputv
100 if (!Ops::Base::BroadcastShape(&ps, &vs, &bcShape) || !(bcShape == vs)) {105 if (!Ops::Base::BroadcastShape(&ps, &vs, &bcShape) || !(bcShape == vs)) {
101 OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(106 OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
102 context_->GetNodeName(), "input3", Ops::Base::ToString(ps).c_str(),107 context_->GetNodeName(), "input3", Ops::Base::ToString(ps).c_str(),
@@ -45,25 +45,32 @@ static ge::graphStatus InferShape4LambApplyOptimizerAssign(gert::InferShapeConte
45 auto outputm_shape = context->GetOutputShape(OUTPUTM_IDX);45 auto outputm_shape = context->GetOutputShape(OUTPUTM_IDX);
46 OP_CHECK_NULL_WITH_CONTEXT(context, outputm_shape);46 OP_CHECK_NULL_WITH_CONTEXT(context, outputm_shape);
47 47 
48- // inputv、inputm 承载动量的更新结果,内核按广播出的完整网格计算并写回它们,48+ // grad、inputv、inputm 三者形状必须完全相同:inputv/inputm 是原地更新的动量输出,
49- // 故输出形状由 inputv/inputm 决定,grad 与 input3 只能向上广播进来。49+ // 输出形状由它们决定;grad 绑定在广播 DAG 的 In0 位,底层 Ops::Base 广播模板
50+ // (DoDimensionCollapse) 不支持对 In0 做广播——实测 grad 为标量、或任一维为 1 时
51+ // 均在 tiling 阶段被拒(“dim num is not same”/“dim index is not same with out”)。
52+ // 故此处直接按等形拒绝,避免放行后到 tiling 才抛 E90003。
53+ // 仅 input3 参与广播(右对齐,维度数可少于 inputv)。
50 // 此处的判定与 tiling 的 CheckInplaceShapeConstraint 保持一致,避免两个 host54 // 此处的判定与 tiling 的 CheckInplaceShapeConstraint 保持一致,避免两个 host
51 // 阶段对同一组合给出不同结论。55 // 阶段对同一组合给出不同结论。
52 OP_CHECK_IF(!(*inputv_shape == *inputm_shape),56 OP_CHECK_IF(!(*inputv_shape == *inputm_shape),
53- OP_LOGE(context->GetNodeName(),57+ OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(
54- "inputv %s and inputm %s must have the same shape, they carry the moment update outputs!",58+ context->GetNodeName(), "inputv and inputm",
55- ToString(*inputv_shape).c_str(), ToString(*inputm_shape).c_str()),59+ (ToString(*inputv_shape) + " and " + ToString(*inputm_shape)).c_str(),
60+ "inputv and inputm are in-place updated moments and must have the same shape"),
56 return ge::GRAPH_FAILED);61 return ge::GRAPH_FAILED);
57 62 
58 gert::Shape broadcast_shape;63 gert::Shape broadcast_shape;
59- OP_CHECK_IF(!BroadcastShape(grad_shape, inputv_shape, &broadcast_shape) || !(broadcast_shape == *inputv_shape),64+ OP_CHECK_IF(!(*grad_shape == *inputv_shape),
60- OP_LOGE(context->GetNodeName(), "grad %s must be broadcastable into the moment shape %s!",65+ OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
61- ToString(*grad_shape).c_str(), ToString(*inputv_shape).c_str()),66+ context->GetNodeName(), "grad", ToString(*grad_shape).c_str(),
67+ "grad does not support broadcast and must have exactly the same shape as inputv/inputm"),
62 return ge::GRAPH_FAILED);68 return ge::GRAPH_FAILED);
63 69 
64 OP_CHECK_IF(!BroadcastShape(input3_shape, inputv_shape, &broadcast_shape) || !(broadcast_shape == *inputv_shape),70 OP_CHECK_IF(!BroadcastShape(input3_shape, inputv_shape, &broadcast_shape) || !(broadcast_shape == *inputv_shape),
65- OP_LOGE(context->GetNodeName(), "input3 %s must be broadcastable into the moment shape %s!",71+ OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(
66- ToString(*input3_shape).c_str(), ToString(*inputv_shape).c_str()),72+ context->GetNodeName(), "input3", ToString(*input3_shape).c_str(),
73+ "input3 must be broadcastable into the in-place moment shape inputv/inputm"),
67 return ge::GRAPH_FAILED);74 return ge::GRAPH_FAILED);
68 75 
69 *output0_shape = *inputv_shape;76 *output0_shape = *inputv_shape;
@@ -148,17 +148,33 @@ ge::graphStatus RunInferShape(const gert::Shape& grad, const gert::Shape& inputv
148}148}
149} // namespace149} // namespace
150 150 
151-// inputv/inputm 承载动量更新结果,输出形状由它们决定;grad 更小时向上广播进来。151+// grad/inputv/inputm 三者同形,输出形状由它们决定。
152TEST_F(LambApplyOptimizerAssignProtoTest, moment_shape_decides_output_shape)152TEST_F(LambApplyOptimizerAssignProtoTest, moment_shape_decides_output_shape)
153{153{
154 gert::Shape o0 = {}, o1 = {}, o2 = {};154 gert::Shape o0 = {}, o1 = {}, o2 = {};
155 gert::Shape expect = {512, 1024};155 gert::Shape expect = {512, 1024};
156- ASSERT_EQ(RunInferShape({1, 1024}, {512, 1024}, {512, 1024}, {512, 1024}, &o0, &o1, &o2), ge::GRAPH_SUCCESS);156+ ASSERT_EQ(RunInferShape({512, 1024}, {512, 1024}, {512, 1024}, {512, 1024}, &o0, &o1, &o2), ge::GRAPH_SUCCESS);
157 ASSERT_EQ(Ops::Base::ToString(o0), Ops::Base::ToString(expect));157 ASSERT_EQ(Ops::Base::ToString(o0), Ops::Base::ToString(expect));
158 ASSERT_EQ(Ops::Base::ToString(o1), Ops::Base::ToString(expect));158 ASSERT_EQ(Ops::Base::ToString(o1), Ops::Base::ToString(expect));
159 ASSERT_EQ(Ops::Base::ToString(o2), Ops::Base::ToString(expect));159 ASSERT_EQ(Ops::Base::ToString(o2), Ops::Base::ToString(expect));
160}160}
161 161 
162+// grad 不参与广播:小于动量形状同样拒收(底层广播模板不支持对 In0 广播)。
163+TEST_F(LambApplyOptimizerAssignProtoTest, grad_smaller_than_moment_is_rejected)
164+{
165+ gert::Shape o0 = {}, o1 = {}, o2 = {};
166+ ASSERT_EQ(RunInferShape({1, 1024}, {512, 1024}, {512, 1024}, {512, 1024}, &o0, &o1, &o2), ge::GRAPH_FAILED);
167+}
168+ 
169+// input3 是唯一参与广播的输入:小于动量形状时按右对齐广播,须放行。
170+TEST_F(LambApplyOptimizerAssignProtoTest, input3_broadcast_into_moment_is_accepted)
171+{
172+ gert::Shape o0 = {}, o1 = {}, o2 = {};
173+ gert::Shape expect = {512, 1024};
174+ ASSERT_EQ(RunInferShape({512, 1024}, {512, 1024}, {512, 1024}, {1, 1024}, &o0, &o1, &o2), ge::GRAPH_SUCCESS);
175+ ASSERT_EQ(Ops::Base::ToString(o0), Ops::Base::ToString(expect));
176+}
177+ 
162// grad 大于动量形状:结果无处容纳,infershape 须与 tiling 一样拒收。178// grad 大于动量形状:结果无处容纳,infershape 须与 tiling 一样拒收。
163TEST_F(LambApplyOptimizerAssignProtoTest, grad_larger_than_moment_is_rejected)179TEST_F(LambApplyOptimizerAssignProtoTest, grad_larger_than_moment_is_rejected)
164{180{