已合并
fix(op_host): 10 个算子补齐缺失的 InferDataType 注册 #9616
fix(op_host): 10 个算子补齐缺失的 InferDataType 注册 #9616
已合并
zl_hw创建于 8月31日
zl_hw成员
8月31日

描述

问题

Relu6Grad、LambApplyOptimizerAssign、LambApplyWeightAssign、LambNextMV、LambNextMVWithDecay、LambNextRight、LambUpdateWithLr、LambUpdateWithLrV2、MultilabelMarginLoss、PoissonNllLoss 这 10 个算子在 op_host/*_infershape.cpp 中只注册了 InferShape,没有注册 InferDataType,GE 图通路上输出 dtype 缺少显式推导来源。这 10 个算子的 op_graph 原型、op_host def 与 op_kernel 均齐全,确属有 GE 图通路的算子。

方法

在各算子已有的 op_host/*_infershape.cpp 中新增 InferDataType4<OpType>,并将原有 IMPL_OP_INFERSHAPE 注册扩展为 .InferShape(...).InferDataType(...)。各输出的 dtype 来源按算子语义逐个确定,未套用统一模板:

算子 输出 dtype 来源
Relu6Grad backprops ← gradients
LambApplyOptimizerAssign output0 ← grad;inputv / inputm 为同名 ref 输出,各自跟随同名输入
LambApplyWeightAssign input_param 为同名 ref 输出,跟随第 5 个输入
LambNextMV y1~y4 ← input_mul3
LambNextMVWithDecay y1~y4 ← input_mul3
LambNextRight y1、y2 ← input_square
LambUpdateWithLr y ← input_greater1
LambUpdateWithLrV2 y ← x1
MultilabelMarginLoss y ← x;is_target ← target(INT32)
PoissonNllLoss loss ← input_x

其中 MultilabelMarginLoss 的 is_target 不能跟随 x,否则图上会把 INT32 的 is_target 推成浮点。各算子 def 中的 dtype 列表按位一一对应(lamb 系列为「全 fp16」或「全 fp32」两种组合,无混精),故以上取源自洽。

关于摆放位置

仓内两种写法都有:以本 PR 基线计,op_host 的 IMPL_OP_INFERSHAPE(...).InferDataType(...) 共 323 个文件,op_graph 的 IMPL_OP(...).InferDataType(...) 共 90 个文件。本次按多数写法放在各算子已有的 op_host infershape 文件内,可直接扩展原注册行,每算子只改一个文件,diff 中仅删除 10 行旧注册行,未触碰存量代码格式。

关联的Issue

关联 Issue #5417

测试

bash build.sh -u --ophost --soc=ascend950 --ops=relu6_grad,lamb_apply_optimizer_assign,lamb_apply_weight_assign,lamb_next_m_v,lamb_next_m_v_with_decay,lamb_next_right,lamb_update_with_lr,lamb_update_with_lr_v2,multilabel_margin_loss,poisson_nll_loss

结果:116 tests PASSED,0 FAILED。其中包含本次为 10 个算子各新增的 1 条 infer_datatype 用例:

[       OK ] relu6grad.relu6grad_infer_datatype
[       OK ] MultilabelMarginLossInferShapeTest.multilabelmarginloss_infer_datatype
[       OK ] PoissonNllLossInfershapeTest.poissonnllloss_infer_datatype
[       OK ] LambApplyOptimizerAssignProtoTest.lambapplyoptimizerassign_infer_datatype
[       OK ] LambApplyWeightAssignProtoTest.lambapplyweightassign_infer_datatype
[       OK ] LambNextMVProtoTest.lambnextmv_infer_datatype
[       OK ] LambNextMVWithDecayProtoTest.lambnextmvwithdecay_infer_datatype
[       OK ] LambNextRightProtoTest.lambnextright_infer_datatype
[       OK ] LambUpdateWithLrProtoTest.lambupdatewithlr_infer_datatype
[       OK ] LambUpdateWithLrV2ProtoTest.lambupdatewithlrv2_infer_datatype

用例中的 ASSERT_NE(dataTypeFunc, nullptr) 实证该注册可被 OpImplRegistry 取到;MultilabelMarginLoss 那条显式断言 is_target 推导为 INT32;Relu6Grad 那条按 def 声明逐 dtype 覆盖 fp16 / fp32 / bf16。

门禁自检:clang-format 18.1.8 --style=file 20 个文件全通过;仓内 scripts/oat_check.sh 20 个文件 All checks passed;trailing-whitespace / end-of-file / 冲突标记检查通过。

未覆盖:本次仅验证到 op_host UT 层(用例通过 OpImplRegistry 取到 infer_datatype 函数指针后直接调用),未做真机与 GE 图 e2e 验证。

文档更新

无。本次仅新增 InferDataType 注册与对应 UT,不涉及对外接口与算子资料变更。

类型标签

AI/Agent生成声明

likedislike
Pull Request已成功合入, 合并人@CANN-robot
(感谢 zl_hw 的贡献)
Zzl_hw成员
8月31日 创建了 pull request,commit 46f49b91
atomgit-bot
atomgit-bot
8月31日 评论:

变更摘要

本 PR 为 10 个 host 算子补齐缺失的 InferDataType 注册。改动为每个算子新增静态函数 InferDataType4Xxx(如 InferDataType4Relu6Grad、InferDataType4LambNextMV),在空指针入参时返回 ge::GRAPH_FAILED,否则通过 context->SetOutputDataType(...) 将各输出数据类型设置为对应输入的数据类型,并在 IMPL_OP_INFERSHAPE 注册链上追加 .InferDataType(...)。同时为每个算子新增了基于 gert::InferDataTypeContextFaker 的 UT 用例,校验 infer_datatype 函数已注册且输出数据类型推导正确。

主要改动

  • 为 10 个算子注册 InferDataType:Relu6Grad、MultilabelMarginLoss、PoissonNllLoss、LambApplyOptimizerAssign、LambApplyWeightAssign、LambNextMV、LambNextMVWithDecay、LambNextRight、LambUpdateWithLr、LambUpdateWithLrV2 均在 IMPL_OP_INFERSHAPE 链上追加 .InferDataType(InferDataType4Xxx),填补此前缺失的数据类型推断注册。
  • 输出数据类型跟随输入:各 InferDataType4Xxx 将输出数据类型设为对应输入的类型,例如 Relu6Grad 的 backprops 跟随 gradients(输入 0),PoissonNllLoss 的 loss 跟随 input_x,LambNextMV/LambNextMVWithDecay 的 y1~y4 均跟随输入 0,LambNextRight 的 y1/y2 跟随 input_square。
  • ref 输出与同名输入保持一致:LambApplyOptimizerAssign 的 3 个输出分别跟随输入 0/1/2(grad 及同名 ref 的 inputv、inputm),LambApplyWeightAssign 的输出 0 取第 5 个输入(input_param)。
  • 特殊类型映射:MultilabelMarginLoss 的输出 0(y)跟随输入 0(x),输出 1(is_target)跟随输入 1(target,INT32),不随 x 推导。
  • 补充 UT 用例:为上述每个算子新增 *_infer_datatype 测试,通过 gert::OpImplRegistry 获取 infer_datatype 函数,用 gert::InferDataTypeContextFaker 构造上下文并断言输出数据类型与预期一致(覆盖 DT_FLOAT 及 Relu6Grad 的 DT_FLOAT16/DT_FLOAT/DT_BF16)。
likedislike
不准确?
atomgit-bot
atomgit-bot
8月31日 评论:

代码审查

✅ 未发现问题

likedislike
不准确?
CANN-robotCANN-robot成员
8月31日 添加了label:cann-cla/no
CANN-robot
CANN-robot成员
8月31日 评论:

Thanks for your pull-request.
The full list of commands accepted by me can be found at here.
You can get sig-info at here.
You can self-configure the PR merge rules for this repository. For more details, please refer to Here.
For more, you also can visit HICANN.


PR Approval Progress

✅ Congratulations! All modules have met the lgtm and approve requirements.

Module Approval Details

module lgtm status approve status
activation ✅ 汤平川, 钱泽洪 (2/2) ✅ 汤平川, 钱泽洪 (2/1)
loss ✅ 汤平川, 钱泽洪 (2/2) ✅ 汤平川, 钱泽洪 (2/1)
optim ✅ 汤平川, 钱泽洪 (2/2) ✅ 汤平川, 钱泽洪 (2/1)

💡 Tip:

  • Committer can comment /approve or /lgtm
  • Commenting /approve implies both code review (lgtm) and intent to merge (approve)

CLA Signature Pass

zl_hw, thanks for your pull request. All authors of the commits have signed the CLA. 👍

likedislike
CANN-robotCANN-robot成员
8月31日 将zengjuan,chenfeng61,chaotang233,kevin_huang1234,crystalhu,yu-xinjie62,yangyang016,fanqirui,wangrui_,renruhai,su-yueming,zhajianqing123,chenqi317,wang-xing001,liubo75,tangweiwei2,chenxingyu18,liujie12345678,pingchuantang,Chen_HaoWen,qianzehong设为评审人
CANN-robotCANN-robot成员
8月31日 将chenfeng61,kevin_huang1234,yu-xinjie62,wangrui_,renruhai,su-yueming,zhajianqing123,wang-xing001,chenxingyu18,liujie12345678,pingchuantang,Chen_HaoWen,qianzehong设为审查人
zl_hw成员
8月31日 评论:

compile

likedislike
Zzl_hw成员
8月31日 预合并成功(commit_id: 783af059d65641780c681432a50392153326c40e)
CANN-robotCANN-robot成员
8月31日 添加了label:ci-pipeline-running
Zzl_hw成员
8月31日 修改了pull request 的描述
CANN-robotCANN-robot成员
8月31日 删除了label:ci-pipeline-running
CANN-robotCANN-robot成员
8月31日 添加了label:ci-pipeline-passed
yuning_chen
yuning_chen成员
8月31日 评论:

您好,pr合入已准备就绪,请尽快联系committer进行检视,谢谢!

likedislike
TangPC
TangPC成员
9月1日 评论:

/lgtm
/approve

likedislike
CANN-robotCANN-robot成员
9月1日 添加了label:approved
qianzehong成员
9月1日 评论:

/approve

likedislike
CANN-robotCANN-robot成员
9月1日 添加了label:lgtm
zl_hw成员
9月1日 评论:

/check-pr

likedislike
CANN-robot
CANN-robot成员
9月1日 评论:

The following label is not ready.

cann-cla/yes: Please sign CLA. If you have done, comment /check-cla to recheck again.

likedislike
zl_hw成员
9月1日 评论:

/check-cla

likedislike
CANN-robot
CANN-robot成员
9月1日 评论:

CLA检查未通过,详情可参考这里

likedislike
Zzl_hw成员
9月1日 预合并成功(commit_id: 05310564172520654c970e98797e6737e97c51c3)
Zzl_hw成员
9月1日 强制推送  1 个提交:3a917d8a-fix(op_host): 10 个算子补齐缺失的 InferDataType 注册
Zzl_hw成员
9月1日 预合并成功(commit_id: 111a469f35a032d649578ccfcb40e01d00482f0d)
CANN-robotCANN-robot成员
9月1日 删除了label:cann-cla/no
CANN-robotCANN-robot成员
9月1日 删除了label:ci-pipeline-passed
CANN-robot
CANN-robot成员
9月1日 评论:

Notification

This pull request has been changed(code update) or closed, so removes the following label(s): ci-pipeline-passed.

likedislike
CANN-robotCANN-robot成员
9月1日 添加了label:cann-cla/yes
CANN-robotCANN-robot成员
9月1日 删除了label:lgtmapproved
CANN-robot
CANN-robot成员
9月1日 评论:

Notice

New code changes of the pull request are detected and remove these labels lgtm, approved. 😳

likedislike
zl_hw成员
9月1日 评论:

compile

likedislike
Zzl_hw成员
9月1日 预合并成功(commit_id: 4ed7c91b95ec45a3d89d2e4a2ea9f2357c0ae149)
CANN-robotCANN-robot成员
9月1日 添加了label:ci-pipeline-running
CANN-robotCANN-robot成员
9月1日 删除了label:ci-pipeline-running
CANN-robotCANN-robot成员
9月1日 添加了label:ci-pipeline-passed
TangPC
TangPC成员
9月1日 评论:

/lgtm
/approve

likedislike
CANN-robotCANN-robot成员
9月1日 添加了label:approved
qianzehong成员
9月1日 评论:

/approve

likedislike
CANN-robotCANN-robot成员
9月1日 添加了label:lgtm
CANN-robotCANN-robot成员
9月1日 关闭了关联的issue
CANN-robotCANN-robot成员
9月1日 合入了pull request