已合并
fix(foreach,lamb,scatter_list): ForeachAddc*List 补 Div 0ULP 精度档 + LambApplyOptimizerAssign 支持面订正 #10141
zl_hw创建于 9月9日
fix(foreach,lamb,scatter_list): ForeachAddc*List 补 Div 0ULP 精度档 + LambApplyOptimizerAssign 支持面订正 #10141
已合并
共 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 | ||
| 193 | class _ForeachAddListInplaceCompose: | 193 | class _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_KERNEL | 364 | 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 | +# =========================================================================== | ||
| 360 | class ForeachAddcdivListTorchSpec: | 423 | class 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, **kwargs | 429 | 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_KERNEL | 393 | 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 | +# =========================================================================== | ||
| 389 | class ForeachAddcmulListTorchSpec: | 456 | class 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, **kwargs | 462 | 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 | ||
| 191 | class _ForeachSubListInplaceCompose: | 191 | class _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 | ||
| 31 | enum class AddcOp { MUL, DIV }; | 31 | enum 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 | + | ||
| 33 | constexpr int32_t VL_SIZE = platform::GetVRegSize(); | 38 | constexpr int32_t VL_SIZE = platform::GetVRegSize(); |
| 34 | 39 | ||
| 35 | template <typename T, typename ScalarT, typename Tiling, AddcOp OP> | 40 | template <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 | ## 调用说明 |
Moptim/lamb_apply_optimizer_assign/op_host/arch35/lamb_apply_optimizer_assign_tiling_arch35.cpp+10-5
| @@ -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 需同步放宽。 | ||
| 70 | ge::graphStatus LambApplyOptimizerAssignTiling::CheckInplaceShapeConstraint() | 74 | ge::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) == inputv | 97 | + // 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 保持一致,避免两个 host | 54 | // 此处的判定与 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; |
Moptim/lamb_apply_optimizer_assign/tests/ut/op_host/test_lamb_apply_optimizer_assign_infershape.cpp+18-2
| @@ -148,17 +148,33 @@ ge::graphStatus RunInferShape(const gert::Shape& grad, const gert::Shape& inputv | |||
| 148 | } | 148 | } |
| 149 | } // namespace | 149 | } // namespace |
| 150 | 150 | ||
| 151 | -// inputv/inputm 承载动量更新结果,输出形状由它们决定;grad 更小时向上广播进来。 | 151 | +// grad/inputv/inputm 三者同形,输出形状由它们决定。 |
| 152 | TEST_F(LambApplyOptimizerAssignProtoTest, moment_shape_decides_output_shape) | 152 | TEST_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 一样拒收。 |
| 163 | TEST_F(LambApplyOptimizerAssignProtoTest, grad_larger_than_moment_is_rejected) | 179 | TEST_F(LambApplyOptimizerAssignProtoTest, grad_larger_than_moment_is_rejected) |
| 164 | { | 180 | { |