已关闭
[Requirement|需求建议]: MulNoNan/FusedMulAdd/FusedMulAddAdd tiling补充输入维度数上限8校验 #2413
TangPC创建于  7月28日关闭于  7月28日
TangPC
TangPC成员
7月28日 创建

一、背景信息

MulNoNanFusedMulAddFusedMulAddAdd 三个 arch35 算子的 tiling(op_host/arch35/*_tiling_arch35.cpp)目前在 DoOpTiling()只做了 dtype 一致性校验CheckDtype),对输入 shape 的维度数(rank)没有任何约束。

昇腾张量维度上限为 8。当前这三个算子在收到超过 8 维的输入时,host 侧不会拦截,shape 被直接透传给 Ops::Base::BroadcastBaseTiling<OpDag>,问题在更下游才暴露,且缺少明确指向具体输入名与实际维度数的错误信息,定位成本高。

本仓库其他算子已普遍在 host 侧显式拦截该场景,例如:

算子 位置 常量
IsPosInf math/is_pos_inf/op_host/arch35/is_pos_inf_tiling_arch35.cpp:30 constexpr int64_t MAX_DIM_NUM = 8;
IsNegInf math/is_neg_inf/op_host/arch35/is_neg_inf_tiling_arch35.cpp:29 constexpr int64_t MAX_DIM_NUM = 8;
IsInf math/is_inf/op_host/is_inf_tiling.cpp:32 constexpr uint32_t MAX_DIM = 8;
Cross math/cross/op_kernel/arch35/cross_struct.h:21 constexpr int64_t MAX_DIM = 8;

因此该需求是补齐这三个算子缺失的入参校验,对齐仓库既有实践

二、价值/作用

  1. 错误前移到 host 侧 tiling 阶段:非法维度在图编译期即被拦截并返回 ge::GRAPH_FAILED,不再下沉到 broadcast tiling 之后才失败。
  2. 报错信息可定位:复用仓库统一宏 OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON,日志中明确给出出错的输入名x1/x2/x3/x4)、实际维度数原因The dim num must be no more than 8.),与 IsPosInfCast 等算子的报错格式保持一致。
  3. 覆盖全部输入而非仅首个输入FusedMulAdd(3 输入)、FusedMulAddAdd(4 输入)中,非首位输入单独超维的场景同样会被拦截,消除漏检。
  4. 约束显式化:三个算子 README 的「约束说明」章节同步补充维度上限说明,使用者查文档即可获知限制,无需读代码。

应用场景:图模式(GE IR 构图)下调用上述三个算子,输入张量 rank 超过 8 时。

三、设计方案

3.1 使能方式(涉及哪些框架:如Aclnn直调、Pytorch训练等)

仅图模式(通过算子 IR 构图调用)。

依据(均可在代码中直接核实):

  • 三个算子目录下均无 op_api/ 目录,即未提供 aclnn 直调接口;
  • 三个算子 README 的「调用说明」章节均只列出图模式一种调用方式,样例分别为 math/mul_no_nan/examples/arch35/test_geir_mul_no_nan.cppmath/fused_mul_add/examples/arch35/test_geir_fused_mul_add.cppmath/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 分支):

算子 输入 支持数据类型 数据格式
MulNoNan x1, x2 FLOAT16, FLOAT, INT32, BFLOAT16 ND
FusedMulAdd x1, x2, x3 FLOAT16, FLOAT, INT32 ND
FusedMulAddAdd x1, x2, x3, x4 FLOAT16, FLOAT, INT32 ND

各算子要求全部输入与输出为同一种数据类型(由已有的 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

算子 维度上限常量 输入个数常量 输入名数组
MulNoNan MUL_NO_NAN_MAX_DIM_NUM = 8 MUL_NO_NAN_INPUT_NUM = 2 {"x1", "x2"}
FusedMulAdd FUSED_MUL_ADD_MAX_DIM_NUM = 8 FUSED_MUL_ADD_INPUT_NUM = 3 {"x1", "x2", "x3"}
FusedMulAddAdd FUSED_MUL_ADD_ADD_MAX_DIM_NUM = 8 FUSED_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:45math/fused_mul_add/op_host/fused_mul_add_def.cpp:50math/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 算子约束限制

本需求新增的约束:

  • MulNoNanx1x2 的维度数(rank)均不能超过 8,超过时在 tiling 阶段校验失败。
  • FusedMulAddx1x2x3 的维度数(rank)均不能超过 8,超过时在 tiling 阶段校验失败。
  • FusedMulAddAddx1x2x3x4 的维度数(rank)均不能超过 8,超过时在 tiling 阶段校验失败。

上述约束已同步写入三个算子 README 的「约束说明」章节。

各算子原有约束保持不变,摘录如下(来自 README):

  • MulNoNan:输入输出须同一数据类型;仅判 x2 一侧,x2 = -0 同样进入零臂;支持任意 NumPy 广播形态与动态 shape / 动态 rank。
  • FusedMulAdd:输入输出须同一数据类型;支持任意 NumPy 广播形态与动态 shape / 动态 rank。
  • FusedMulAddAdd:输入输出须同一数据类型;计算顺序为 ((x1 * x2) + x3) + x4,不可交换;x1 必须为完整输出 shape,当前 runtime 不支持 x1 自身向上广播;不支持含 -1 的动态 shape。

验收用例(已在 PR 中随实现一并提交,每个算子 3 个,共 9 个):

用例 场景 预期
max_dim_num_8d_fp32 全部输入为 8 维 {2,1,1,1,1,1,1,2} GRAPH_SUCCESS(上边界不误杀)
x1_dim_num_over_8_failed 输入为 9 维 GRAPH_FAILED
x2/x3/x4``_dim_num_over_8_failed 仅末位输入为 9 维,其余为 {2} GRAPH_FAILED(验证遍历全部输入)

💡 备注

关联 PR:https://gitcode.com/cann/ops-math/pull/4294

likedislike
TangPCTangPC成员
7月28日 添加了label:requirement
TangPC
TangPC成员
7月28日 评论:
likedislike
CANN-robotCANN-robot成员
7月28日 将 pingchuantang 设为负责人
CANN-robotCANN-robot成员
7月28日 关闭了 issue
CANN-robotCANN-robot成员
7月28日 添加了label:resolved
TangPC
TangPC成员
8月1日 评论:

/open

likedislike