Pull Request已成功合入, 合并人@CANN-robot
(感谢 zl_hw 的贡献)变更摘要
本 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)。


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
/approveor/lgtm- Commenting
/approveimplies 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. 👍


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


/lgtm
/approve


/approve


The following label is not ready.
cann-cla/yes: Please sign CLA. If you have done, comment /check-cla to recheck again.


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


/lgtm
/approve


/approve


描述
问题
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 来源按算子语义逐个确定,未套用统一模板:其中 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 用例:
用例中的
ASSERT_NE(dataTypeFunc, nullptr)实证该注册可被OpImplRegistry取到;MultilabelMarginLoss 那条显式断言 is_target 推导为 INT32;Relu6Grad 那条按 def 声明逐 dtype 覆盖 fp16 / fp32 / bf16。门禁自检:clang-format 18.1.8
--style=file20 个文件全通过;仓内scripts/oat_check.sh20 个文件 All checks passed;trailing-whitespace / end-of-file / 冲突标记检查通过。未覆盖:本次仅验证到 op_host UT 层(用例通过
OpImplRegistry取到infer_datatype函数指针后直接调用),未做真机与 GE 图 e2e 验证。文档更新
无。本次仅新增 InferDataType 注册与对应 UT,不涉及对外接口与算子资料变更。
类型标签
AI/Agent生成声明