已关闭
[Requirement|需求建议]: 为 FusedMulApplyMomentumExtern 算子新增 Tiling 阶段的校验逻辑:dim>8时返回FAILED, 所有参数dtype符合约束 #4630
牛康创建于  8月6日关闭于  8月6日
牛康
牛康
8月6日 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

Backgroud(背景信息)

为 FusedMulApplyMomentumExtern 算子新增了 Tiling 阶段的校验逻辑。主要在两个层面进行了扩展:一是在 fused_mul_apply_momentum_extern_tiling_arch35.cpp 中添加了输入张量的秩(rank)上限检查和各输入的数据类型(dtype)规则校验;二是新增了一个完整的异常测试文件 test_geir_fused_mul_apply_momentum_extern_exception.cpp,覆盖了不支持的数据类型、GEIR Cast 插入、9 维张量拒绝以及形状维度不匹配等场景。

主要改动

Tiling 层秩校验:在 GetShapeAttrsInfo 函数中新增 varShape->GetStorageShape().GetDimNum() > 8 的检查,当输入张量的维度超过 8 时直接返回 GRAPH_FAILED。

Tiling 层 dtype 校验:新增对全部 7 个输入的逐项数据类型检查——var 必须为 FP32;accum、lr、x1、momentum、x2 必须同为 FP32/FP16/BF16 之一且相互一致;var_copy 在 accum 为 BF16 时须为 BF16,否则须为 FP16。

Origin(信息来源)

长尾算子

Benefit / Necessity (价值/作用)

Design(设计方案)

likedislike
牛康牛康
8月6日 关联了pull request:feat: add verification on FusedMulApplyMomentumExtern
CANN-robotCANN-robot成员
8月6日 关闭了 issue
CANN-robotCANN-robot成员
8月6日 添加了label:resolved