已关闭
[Requirement|需求建议]: MulNoNan/FusedMulAdd/FusedMulAddAdd tiling补充输入维度数上限8校验 #2413
TangPC创建于 7月28日关闭于 7月28日
7月28日 添加了label:requirement
TangPC
7月28日 评论:
7月28日 评论:
/assign @pingchuantang


7月28日 将 pingchuantang 设为负责人
7月28日 关闭了 issue
7月28日 添加了label:resolved
TangPC
8月1日 评论:
8月1日 评论:
/open


一、背景信息
MulNoNan、FusedMulAdd、FusedMulAddAdd三个 arch35 算子的 tiling(op_host/arch35/*_tiling_arch35.cpp)目前在DoOpTiling()中只做了 dtype 一致性校验(CheckDtype),对输入 shape 的维度数(rank)没有任何约束。昇腾张量维度上限为 8。当前这三个算子在收到超过 8 维的输入时,host 侧不会拦截,shape 被直接透传给
Ops::Base::BroadcastBaseTiling<OpDag>,问题在更下游才暴露,且缺少明确指向具体输入名与实际维度数的错误信息,定位成本高。本仓库其他算子已普遍在 host 侧显式拦截该场景,例如:
math/is_pos_inf/op_host/arch35/is_pos_inf_tiling_arch35.cpp:30constexpr int64_t MAX_DIM_NUM = 8;math/is_neg_inf/op_host/arch35/is_neg_inf_tiling_arch35.cpp:29constexpr int64_t MAX_DIM_NUM = 8;math/is_inf/op_host/is_inf_tiling.cpp:32constexpr uint32_t MAX_DIM = 8;math/cross/op_kernel/arch35/cross_struct.h:21constexpr int64_t MAX_DIM = 8;因此该需求是补齐这三个算子缺失的入参校验,对齐仓库既有实践。
二、价值/作用
ge::GRAPH_FAILED,不再下沉到 broadcast tiling 之后才失败。OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON,日志中明确给出出错的输入名(x1/x2/x3/x4)、实际维度数与原因(The dim num must be no more than 8.),与IsPosInf、Cast等算子的报错格式保持一致。FusedMulAdd(3 输入)、FusedMulAddAdd(4 输入)中,非首位输入单独超维的场景同样会被拦截,消除漏检。应用场景:图模式(GE IR 构图)下调用上述三个算子,输入张量 rank 超过 8 时。
三、设计方案
3.1 使能方式(涉及哪些框架:如Aclnn直调、Pytorch训练等)
仅图模式(通过算子 IR 构图调用)。
依据(均可在代码中直接核实):
op_api/目录,即未提供 aclnn 直调接口;math/mul_no_nan/examples/arch35/test_geir_mul_no_nan.cpp、math/fused_mul_add/examples/arch35/test_geir_fused_mul_add.cpp、math/fused_mul_add_add/examples/arch35/test_geir_fused_mul_add_add.cpp,通过op_graph/*_proto.h中的算子 IR 构图调用。本需求不新增任何对外接口,不改变使能方式。
3.2 总体设计
3.2.1 算子支持的数据类型
本需求不改变算子已支持的数据类型,现状如下(取自各算子 README 参数说明与 tiling 的 dtype 分支):
各算子要求全部输入与输出为同一种数据类型(由已有的
CheckDtype保证)。3.2.2 host侧设计
(1)头文件新增私有方法声明(
op_host/arch35/*_tiling_arch35.h)private: uint64_t tilingKey = 0; bool CheckDtype(...) const; bool CheckShape() const; // 新增(2)实现文件新增文件级常量(
op_host/arch35/*_tiling_arch35.cpp)MUL_NO_NAN_MAX_DIM_NUM = 8MUL_NO_NAN_INPUT_NUM = 2{"x1", "x2"}FUSED_MUL_ADD_MAX_DIM_NUM = 8FUSED_MUL_ADD_INPUT_NUM = 3{"x1", "x2", "x3"}FUSED_MUL_ADD_ADD_MAX_DIM_NUM = 8FUSED_MUL_ADD_ADD_INPUT_NUM = 4{"x1", "x2", "x3", "x4"}(3)
CheckShape()实现(以 MulNoNan 为例,另两个算子同构,仅常量名与输入个数不同)bool MulNoNanTiling::CheckShape() const { for (size_t i = 0; i < MUL_NO_NAN_INPUT_NUM; i++) { auto inputShape = context_->GetInputShape(i); OP_CHECK_IF(inputShape == nullptr, OP_LOGE(context_->GetNodeName(), "The shape of %s is nullptr.", MUL_NO_NAN_INPUT_NAMES[i]), return false); size_t dimNum = inputShape->GetStorageShape().GetDimNum(); if (dimNum > MUL_NO_NAN_MAX_DIM_NUM) { std::string reasonMsg = "The dim num must be no more than " + std::to_string(MUL_NO_NAN_MAX_DIM_NUM) + "."; OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(context_->GetNodeName(), MUL_NO_NAN_INPUT_NAMES[i], std::to_string(dimNum), reasonMsg); return false; } } return true; }(4)调用点:置于
DoOpTiling()最前,先于已有的 dtype 校验,超限即短路返回。ge::graphStatus MulNoNanTiling::DoOpTiling() { if (!CheckShape()) { return ge::GRAPH_FAILED; } // 原有 dtype 校验与 BroadcastBaseTiling 分发逻辑不变 ... }设计取舍说明:
for遍历INPUT_NUM个输入而非只校验x1,保证非首位输入单独超维时同样被拦截;inputShape->GetStorageShape().GetDimNum();OP_CHECK_IF/OP_LOGE/OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON,不引入新的报错风格;GetInputShape()返回nullptr的分支。3.2.3 kernel侧设计
无 kernel 侧改动。 校验完全在 host 侧 tiling 阶段完成,
op_kernel/下的 DAG、struct、kernel 入口均不涉及,已有正常 shape 的执行路径与性能不受影响。3.3 支持硬件
Ascend 950PR / Ascend 950DT。
依据:三个算子的
op_host/*_def.cpp中均为this->AICore().AddConfig("ascend950", aicoreConfig);(math/mul_no_nan/op_host/mul_no_nan_def.cpp:45、math/fused_mul_add/op_host/fused_mul_add_def.cpp:50、math/fused_mul_add_add/op_host/fused_mul_add_add_def.cpp:55),且实现位于arch35目录下;三个 README 的「产品支持情况」表格中,仅Ascend 950PR/Ascend 950DT一行为 √,Atlas A3 / A2 / 200I·500 A2 / 推理系列 / 训练系列均为 ×。本需求不改变硬件支持范围。
3.4 算子约束限制
本需求新增的约束:
MulNoNan:x1、x2的维度数(rank)均不能超过 8,超过时在 tiling 阶段校验失败。FusedMulAdd:x1、x2、x3的维度数(rank)均不能超过 8,超过时在 tiling 阶段校验失败。FusedMulAddAdd:x1、x2、x3、x4的维度数(rank)均不能超过 8,超过时在 tiling 阶段校验失败。上述约束已同步写入三个算子 README 的「约束说明」章节。
各算子原有约束保持不变,摘录如下(来自 README):
x2一侧,x2 = -0同样进入零臂;支持任意 NumPy 广播形态与动态 shape / 动态 rank。((x1 * x2) + x3) + x4,不可交换;x1必须为完整输出 shape,当前 runtime 不支持x1自身向上广播;不支持含-1的动态 shape。验收用例(已在 PR 中随实现一并提交,每个算子 3 个,共 9 个):
max_dim_num_8d_fp32{2,1,1,1,1,1,1,2}GRAPH_SUCCESS(上边界不误杀)x1_dim_num_over_8_failedGRAPH_FAILEDx2/x3/x4``_dim_num_over_8_failed{2}GRAPH_FAILED(验证遍历全部输入)💡 备注
关联 PR:https://gitcode.com/cann/ops-math/pull/4294