CANNJudge-Add算子实现

Ascend C Add 算子### 摘要

基于 PyTorch torch.add(x1, x2) 语义,在昇腾 NPU 上用 Ascend C 实现的高性能逐元素加法算子。

一、题目到底要求什么

题目内容

Add算子
一、赛题背景

Add(逐元素加法)是最基础的二元向量算子,完成两个张量按广播规则对齐后的逐元素相加。它广泛用于残差网络shortcut分支累加(ResNet的 x + F(x))、优化器参数更新(SGD/Adam 的 w += g)、多尺度特征融合、bias 累加等场景,是所有深度学习框架与算子融合体系中的公共底座。

本题要求基于PyTorch中torch.add(x1, x2)的核心业务逻辑,采用Ascend C编程语言进行算子原生开发,

在昇腾NPU硬件上实现一款高性能的Add算子。

二、算子功能描述

实现的Add算子需完成以下核心计算:

  1. 步骤1 : 读取输入张量 x1 与 x2,按广播规则对齐(shape 需可广播,输出取广播后 shape);
  2. 步骤2 : 对齐后的每个元素计算 y = x1 + x2,写出结果张量 y。

算子的核心难点在于:广播对齐时的数据搬运与访存合并、超大规模输入(千万级元素)下的带宽利用与多核切分均衡,以及 fp16/int32 混合精度路径下的类型转换流水(浮点内部 fp32 计算、整型 int64 流水回绕)。

三、核心定义与约束
3.1 参考算子

等价python实现:

python 复制代码
import torch

# x1/x2 同 dtype,shape 可广播
y = torch.add(x1, x2)
3.2 数学公式

yi=x1i+x2i y_i = x1_i + x2_i yi=x1i+x2i

其中 iii 遍历广播结果中的所有元素位置。

3.3 输入输出与属性总览
类型 参数名 类型 维度形状 支持数据类型 数据格式 备注
输入 x1 required 任意ND(与x2可广播) float16, float32, int32 ND 加数
输入 x2 required 任意ND(与x1可广播) float16, float32, int32 ND 加数,与x1同dtype
输出 y required broadcast(x1, x2) float16, float32, int32 ND 与输入同 dtype

本算子无属性。

3.4 关键输入约束
  • 数据类型: x1 与 x2 同 dtype,支持 float16/float32/int32;
  • 维度场景: 任意 ND,两输入 shape 满足 numpy 广播规则(Infershape4Broadcast),同形为最常用场景;
  • 维度取值范围(均为正整数):
  • 总元素数 N ∈ 1, 67108864(约 6.7 千万元素压测上限);
  • 单维长度 ∈ 1, 2147483647;
  • 其他约束: 广播维长度为 1 的输入沿该维复制展开;题目用例的数据范围为 [-2, 2)(int32 为 [-100000, 100000)),无溢出分支(int32 环绕语义见 3.7)。
3.5 核心属性说明

无属性。

3.6 输出严格要求
  • 输出 y 的 shape 严格等于 broadcast(x1.shape, x2.shape),dtype 与输入一致;
  • 输出元素与输入元素按广播位置一一对应。
3.7 特殊值处理规则

语义基准:PyTorch torch.add(x1, x2) 实测(torch 2.10,实测算例见生成小结)。

  • inf/nan 传播(IEEE-754) : 同号无穷相加保持同号 inf;异号无穷相加(inf + (-inf))结果为 nan;任一输入为 nan 时结果为 nan;无穷加有限值保持无穷符号。实测 [inf, inf, -inf, nan] + [-inf, 1, 1, 1] = [nan, inf, -inf, nan];
  • fp16 溢出舍入: 内部 fp32 计算后按 RINT(向最近偶舍入)回转 fp16。实测 65504+2 = 65504(不进位),65504+16 = inf(超出 fp16 表示范围为 inf);
  • fp32 大数相消: 同量级有限大数相加结果按 IEEE 规则精确表示,不引入额外 eps;
  • int32 环绕 : 整型加法按 模运算环绕(与 torch 实测一致:2147483647 + 1 = -2147483648,-2147483648 + (-1) = 2147483647);golden 以 int64 流水相加后回绕 int32,题目 int32 用例取值范围保证和无溢出的日常场景一致。
四、规则要求
  1. 一致性 : 实现结果须与参考实现 torch.add(x1, x2) 在相同输入下数值一致(满足第五节精度阈值);
  2. 广播支持: 正确处理 x1/x2 shape 不一致但可广播的场景(行/列向量广播均须支持);
  3. 类型支持: 正确覆盖 float16/float32/int32 三种数据类型的转换流水;
  4. 性能要求: 大规模用例(≥8M 元素)须充分利用多核与向量流水,带宽效率不低于同类 BLAS-1 内核水平。
五、精度判断规则

以更高精度的参考实现结果作为 golden(标杆)进行逐元素比较。

逐元素通过条件 :abs(actual - golden) <= atol + rtol * abs(golden)。

整体通过条件 :matched_ratio >= required_matched_ratio 且 max_abs_error <= max_abs_error_limit。

误差阈值表:

数据类型 float32 float16 bfloat16
atol 9.77e-4 1.95e-3 1.56e-2
rtol 1.53e-5 1.95e-3 1.56e-2
required_matched_ratio 0.99 0.99 0.99
max_abs_error_limit 1e-2 1e-1 1e-0

int32 用例要求逐元素精确相等。

六、示例说明

示例1: 同形小规模

复制代码
x1 = [1, 2, 3] (float32), x2 = [4, 5, 6] (float32)
y  = [1+4, 2+5, 3+6] = [5, 7, 9]

示例2: 广播

复制代码
x1 = [[0, 0, 0], [1, 1, 1]] (2x3, float32), x2 = [10, 20, 30] (3, float32)
y  = [[10, 20, 30], [11, 21, 31]]

示例3: int32 环绕

复制代码
x1 = [2147483647] (int32), x2 = [1] (int32)
y  = [-2147483648]   (按 2^32 环绕)

分析

1.1 数学语义

y=x1+x2y = x1 + x2y=x1+x2

两个张量按 numpy 广播规则 对齐后逐元素相加,输出 shape = broadcast(x1.shape, x2.shape),dtype 与输入一致。

  • 输入:x1、x2,同 dtype,任意 ND,shape 可广播
  • 输出:y,同 dtype,shape 为广播结果
  • 支持 dtype:float16 / float32 / int32
  • 规模:总元素数 N ∈ 1, 67,108,864(约 6.7 千万)
1.2 必须踩中的"坑"(特殊语义)

这些是和 PyTorch 对齐时最容易写错的地方:

场景 规则 代码对应
inf/nan 传播 同号 inf 相加仍是 inf;inf + (-inf) = nan;任一 nan 则 nan AscendC::Add 硬件按 IEEE-754 天然处理,无需特判
fp16 内部精度 内部按 fp32 计算,结果用 RINT(向最近偶) 舍入回 fp16。例:65504+2=65504(不进位),65504+16=inf Cast → fp32 Add → Cast(CAST_RINT) 回 fp16
int32 环绕 按 mod 2³² 环绕:INT_MAX+1 = INT_MIN,与 C 有符号溢出 UB 不同 用 uint32_t 做加法,再 cast 回 int32_t
广播维 stride=0 某维输入长度为 1 时,沿该维"复制",等价于 stride=0 host 端算 stride 时把 dim==1 的 stride 置 0
1.3 精度门槛
  • float:逐元素 abs(actual-golden) <= atol + rtol*|golden|,匹配率 ≥ 99%
  • int32:必须逐元素精确相等
  • fp16 的 atol/rtol 比 fp32 宽,正是因为中间走 fp32、末尾 RINT
1.4 性能要求

≥ 8M 元素的大规模用例要吃满 多核 + 向量流水 ,带宽效率对齐 BLAS-1。本质上 Add 是纯内存带宽受限算子------算力远大于带宽,所以一切设计都围绕"怎么把数据搬得最顺"。


二、整体设计思路:为什么要分三条路径

最朴素的写法:对每个输出元素,算出它在 x1/x2 里的偏移,读两个数、加、写回。这在小 tensor 上没问题,但在 6.7 千万元素时会完全跑不满带宽 ------因为 Ascend 的 DataCopy(向量化搬数)一次搬一整块,而 GetValue/SetValue(标量访问)一个元素一次请求,效率差几个数量级。

所以核心矛盾是:

广播会让"连续内存"被打断,但性能又要求连续搬数。

代码因此把输入分成三类,匹配三种实现:

复制代码
                 ┌─────────────────────────────────────┐
host 端分析 shape │  能否线性切片?有多长的连续尾段?         │
                 └─────────────────────────────────────┘
                                  │
        ┌─────────────────────────┼─────────────────────────┐
        ▼                         ▼                         ▼
  ① linear 路径            ② segment 路径            ③ generic 路径
 两输入完全连续             尾部一段可向量搬            真·不规则广播
 或其中一个是标量           (外层循环 + 内层块拷贝)    标量 odometer
 纯 DataCopy 流水          块级 DataCopy + 偏移计算     GetValue/SetValue
 (最快,覆盖 90% 用例)    (折中,行列广播)           (兜底,正确性优先)

优先级:linear > segment > generic。host 能走快路就绝不让设备走慢路。

题目要干嘛

就是做一个 y = x1 + x2 的加法算子,跑在华为昇腾 NPU 上。听起来简单,但有几个麻烦点:

  1. 两个张量形状可能不一样 ,要按广播规则对齐。比如 x1 是 [2,3] 的矩阵,x2 是 [3] 的一行,那 x2 要 "横着复制两遍" 再相加。
  2. 三种数据类型都要支持:fp16、fp32、int32。其中 fp16 要求内部先转成 fp32 算完再舍回来(不然精度不够),int32 加法溢出要像 C 语言那样 "环绕"(2147483647+1 突然变成负数)。
  3. 要快。动辄几千万个元素,慢了不行。
为什么代码要这么写

核心原因就一句话:在 NPU 上,加法本身不花钱,把数据从内存搬进搬出才花钱。

NPU 跟 CPU 不一样。它不能一个元素一个元素地从大内存(GM)里抠出来算 ------ 那样慢得要死。它的高效方式是:一次搬一大块连续的数据进片上缓存(UB),然后用向量指令一把加完,再整块写回去。这叫 "向量化搬数"。

但广播会把内存连续性打断。比如 x2 是一行 [a,b,c],要加到矩阵每一行上,那它在内存里其实只有一份,但逻辑上要被用很多次。这种情况下你没法简单地 "从头搬到尾"。

所以代码的思路是见什么菜下什么碟:

  • 最好的情况(两个张量形状一模一样,内存连续):那最简单,切几块,每个核搬一块、加一块、写一块,双缓冲让搬数据和算重叠起来,跑满带宽就行。如果其中一个其实是个标量(比如加 bias),那连搬都不用搬,直接拿那个值铺满一整块。
  • 稍微麻烦点(最右边那一维是连续的,只是外面套了几层广播):那就外层循环算一下 "这一块数据在内存里的起始地址是多少",里面那一段连续的还是照样整块搬。
  • 最麻烦(乱七八糟的广播):那就老实一个元素一个元素地算地址、读、加、写。慢,但保证不出错。

为什么要在 host 端(CPU)就把这些 shape 分析、选路都做完?因为 NPU 核很多,让每个核自己去分析 shape 是重复劳动,而且设备端不擅长这种分支判断。host 算好一张 "每个维度步长多少" 的表,直接发给每个核,核拿到就干活。

还有几个细节:

  • 步长设 0 就代表广播,不用真的把数据复制一份,省内存省时间。
  • 对齐。NPU 搬内存要求按 32 字节一块搬,所以开头对齐的主体走向量快路,多出来的几个零头让 0 号核单独标量处理,不影响大局。
  • 多核切匀。用整除公式保证每个核干的活差不多,不会有的核忙死有的核闲着。

一句话总结:这题表面是考加法,实际考的是 "怎么把不连续的广播数据,尽可能地拼成大块连续内存来搬"。代码做的三件事 ------host 分析步长、分快 / 慢三条路径、双缓冲向量化 ------ 全都是为了这个目标服务的。


三、Host 侧:把 shape 翻译成"步长表"

3.1 AddBuildHostMeta ------ 建元信息

这是整个算子最聪明的部分。它不真的展开广播(那会浪费内存和时间),而是只算每个维度的步长(stride)。

从最右维(最快变化维)往左扫:

cpp 复制代码
raw[axis] = {
    outputDim,
    x1Dim == 1 && outputDim != 1 ? 0 : x1Stride,   // ← 广播维 stride=0
    x2Dim == 1 && outputDim != 1 ? 0 : x2Stride,
};
x1Stride *= x1Dim;   // 累积"跨过一个元素要走多少"
x2Stride *= x2Dim;
total *= outputDim;

关键点:广播维(输入 dim=1、输出 dim>1)的 stride 直接设为 0 。

这样后面算偏移时 coordinate * 0 = 0,不管坐标怎么变,输入指针都不动------正好就是"沿该维复制"的语义。比真正复制数据省掉一次内存搬运。

3.2 合并相邻维(dim merge)

cpp 复制代码
if (x1Merge && x2Merge) {
    previous.dim *= current.dim;      // 两维合成一维
    previous.x1Stride = current.x1Stride;
    previous.x2Stride = current.x2Stride;
}

什么时候能合并?当上一维的 stride 恰好等于"下一维长度 × 下一维 stride"时,说明这两维在内存里是连续的,可以拍平成一维。例如 shape [2, 3, 4]、stride 分别 [12, 4, 1] 能合并成 [24]、stride [1]------也就是一整块连续内存。

这一步直接决定了"linear 路径"能不能命中:合并后只剩一维、且两输入 stride 都是 1,就是纯连续。

3.3 为什么限制 8 维(ADD_MAX_DIMS = 8)

设备侧不能动态分配数组,维度信息要通过函数参数一个个传下去(见后面 ADD_META_PARAMETERS 宏)。传 8 个是性能与通用性的折中。超过 8 维怎么办? host 把前导维切成"前缀",剩下后 8 维交给 generic 内核,外层 host 循环每段切片调一次 kernel(见 run_kernel 里 prefixCount 那段)。这是个工程妥协:任意高维都能正确算,只是超维部分退化为切片循环。


四、设备侧三条路径详解

4.1 公共基础设施

双缓冲队列(double buffer)
cpp 复制代码
AscendC::TQue<AscendC::QuePosition::VECIN, ADD_BUFFER_NUM> x1Queue;  // ADD_BUFFER_NUM=2

NPU 的执行分三段流水:GM→UB(搬入)→ 向量计算 → UB→GM(搬出)。用 2 个 buffer 后,第 N 块在算的时候,第 N+1 块已经在搬入,搬算重叠,藏住访存延迟。这是带宽受限算子必须做的。

复制代码
时间轴:  [搬x1块0][搬x1块1][搬x1块2]...
              [算块0  ][算块1  ][算块2  ]...
                   [写块0  ][写块1  ][写块2  ]...
核间切分
cpp 复制代码
const uint64_t block = AscendC::GetBlockIdx();        // 本核编号
begin = workUnits * block / coreNum;                  // 本核负责的区间
end   = workUnits * (block+1) / coreNum;

用 workUnits * i / coreNum 这种整除写法,保证各核任务量最多差一个 work unit,不会出现某核忙死某核闲死。这就是题目说的"多核切分均衡"。

32 字节对齐

DataCopy 按 32B 块搬,要求地址和个数对齐:

cpp 复制代码
constexpr uint32_t alignElements = 32 / sizeof(T);   // fp32=8, fp16=16
alignedTotal = total / alignElements * alignElements; // 对齐后的主体

不齐的尾巴交给 0 号核用标量 GetValue/SetValue 收尾。这样主体走向量化快路,尾巴只有几十上百个元素,不影响带宽。


4.2 路径①:AddLinearVectorImpl(连续/标量)

命中条件:两输入元素总数都 == 输出总数(完全连续),或其中一个只有 1 个元素(标量广播)。

三种 mode:

  • ADD_BOTH_CONTIGUOUS:两边都正常 DataCopy
  • ADD_X1_SCALAR:x1 是标量,不用搬,直接 Duplicate 把它填进整块 local tensor
  • ADD_X2_SCALAR:同理
cpp 复制代码
if (mode == ADD_X1_SCALAR) {
    AscendC::Duplicate(x1Local, x1Scalar, count);   // 不读 GM,广播到寄存器/UB
} else {
    AscendC::DataCopy(x1Local, x1Gm[offset], count);
}

标量用 Duplicate 而不是重复 GetValue,一次向量指令铺满整个 buffer------这就是 bias 累加(x + bias)这种最常见场景的优化。

4.3 fp16 专用路径:AddLinearFp16Impl

fp16 不能直接相加(精度不够,而且题目要求内部 fp32),所以多三个 fp32 的 TBuf:

cpp 复制代码
Cast(x1Fp32, x1Local, CAST_NONE, count);          // half -> float,无损
Cast(x2Fp32, x2Local, CAST_NONE, count);
Add(yFp32, x1Fp32, x2Fp32, count);                // fp32 里加
Cast(yLocal, yFp32, CAST_RINT, count);            // float -> half,最近偶舍入
  • CAST_NONE:输入侧,half→float 是精确的,无需舍入
  • CAST_RINT:输出侧,向最近偶数舍入,和 PyTorch 实测对齐(65504+2 不进位靠它)
  • tile 大小取 4096 个 fp16(8KB 输入),因为 fp32 buffer 要占 16KB,总 UB 预算要算着用

为什么不直接 Add(half, half, half)?因为 Ascend 向量指令做 fp16 加法内部也是 fp32 流水,但题目明确要求 RINT 回转和大数相消语义,显式 Cast 路径能确定性地复现 PyTorch 行为,避免不同固件舍入模式不一致。

4.4 路径②:AddSegmentVectorImpl(段式广播)

命中条件:尾部有一段足够长、且对齐的连续区间 (inner),外面套若干组(groups)。

典型例子:x1 = [2,3](逐行)广播加 x2 = [3](逐列)?不对------这种是列广播。真正的 segment 场景是:最内维连续、外层维有广播 。例如 x1=[1, 3](列向量)加 y=[2, 3]:最右维长度 3,两边 stride 都是 1,所以 inner=3, groups=2。

host 端 AddTrailingSpan 从右往左找最长连续尾段:

cpp 复制代码
compatible = (x1Strides[axis]==span && x2Strides[axis]==span);
if (!compatible) break;        // 碰到不连续的维就停
span *= dims[axis];

设备端把工作拆成 chunks = groups × chunksPerGroup 均匀分给各核:

cpp 复制代码
group    = chunk / chunksPerGroup;
position = (chunk % chunksPerGroup) * tileElements;
AddCalculateOffsets(group*inner, ..., x1Base, x2Base);  // 这一组在 GM 里的基址
DataCopy(x1Local, x1Gm[x1Base + position], count);       // 只向量搬连续的那段

这样外层广播用一次标量偏移计算搞定,内层连续段照样吃满带宽。比 generic 路径快得多。

4.5 路径③:AddGenericBroadcastImpl(通用广播,兜底)

真遇到不规则广播(比如行列都有广播),就用**里程表(odometer)**逐元素推坐标:

cpp 复制代码
// 初始化:把 begin 这个输出下标拆成各维坐标,累加出初始偏移
for (axis = rank-1; axis >= 0; --axis) {
    coordinates[axis] = begin % dims[axis];
    x1Offset += coordinates[axis] * x1Strides[axis];
    x2Offset += coordinates[axis] * x2Strides[axis];
}
// 每处理完一个元素,最右维 +1;进位则向左维进位
for (...) {
    yGm.SetValue(i, AddScalarValue(x1Gm.GetValue(x1Offset), x2Gm.GetValue(x2Offset)));
    ++coordinates[axis];
    x1Offset += x1Strides[axis];   // 坐标 +1,偏移线性加
    if (coordinates[axis] < dims[axis]) break;  // 没进位,继续
    // 进位了:该维归零,偏移回退,处理左一维
    coordinates[axis] = 0;
    x1Offset -= dims[axis] * x1Strides[axis];
}

这就是十进制 odometer:个位满 10 进一位。因为 stride 里广播维是 0,所以即使坐标进位,那些维的偏移也自动保持不变------逻辑统一。

这条路径是标量访问,慢,但保证任意可广播 shape 都正确,是正确性的安全网。题目里大规模用例其实走不到这里。


五、类型与精度的几个关键写法

5.1 AddScalarValue 重载

cpp 复制代码
// fp32:直接加
float AddScalarValue(float a, float b) { return a + b; }
// fp16:升 fp32 加,再回 half(和向量路径语义一致)
half  AddScalarValue(half a, half b) {
    return static_cast<half>(static_cast<float>(a) + static_cast<float>(b));
}
// int32:用 uint32 加,再解释回 int32,实现 mod 2^32 环绕
int32_t AddScalarValue(int32_t a, int32_t b) {
    return static_cast<int32_t>(static_cast<uint32_t>(a) + static_cast<uint32_t>(b));
}

注意 generic 路径里 int32 内核实例化成了 uint32_t(add_generic_int32 调 AddGenericBroadcastImpl<uint32_t>),就是为了让 GetValue/SetValue 和加法都走无符号环绕,规避 C 有符号溢出的未定义行为。

5.2 dtype 编码

host 校验里:dtype == 0 → fp32,== 1 → fp16,== 5 → int32。这是昇腾算子框架的枚举约定,不是题目定义的,照框架来即可。


六、run_kernel:host 端的调度大脑

执行流程:

  1. 入参校验:空指针、dtype 一致、维度关系合法,不合法直接 return(防御性编程)。
  2. 建 meta :AddBuildHostMeta 算步长表。
  3. 判断 linear:两输入元素数都等于总数?或一个是标量?→ 走 linear。
  4. 超 8 维切片:前导维 host 循环,每段调 generic。
  5. 找最优尾段 :对三种 mode(双连续/x1 标量/x2 标量)分别算 AddTrailingSpan,取最长且对齐的 → 走 segment。
  6. 都不行:走 generic 兜底。

选核数 AddSelectCoreCount:

cpp 复制代码
min(workItems, availableCoreNum)

工作项比核数还少时别开满核(空核白耗调度),这也是性能细节。


kernel 完整源码实现

cpp 复制代码
#include <algorithm>
#include <cstdint>
#include <vector>
#include "kernel_operator.h"

/*
 * Add Kernel - Ascend C direct invocation
 *
 * The host reduces adjacent broadcast dimensions and selects one of three
 * device paths (broadcasts above eight effective dimensions are sliced):
 *   1. linear: both inputs are contiguous, or one input is a scalar;
 *   2. segment: an aligned trailing span can be copied and added by vector;
 *   3. generic: an odometer maps every output element to both input offsets.
 *
 * The vector paths use double-buffered queues. fp16 is promoted to fp32 for
 * the addition and rounded back to fp16. int32 uses modulo-2^32 arithmetic.
 * This file is included by main.asc; do not add main() or include guards.
 */

namespace {

constexpr uint32_t ADD_MAX_DIMS = 8;
constexpr uint32_t ADD_BUFFER_NUM = 2;
constexpr uint32_t ADD_TILE_BYTES = 16 * 1024;
constexpr uint32_t ADD_FP16_TILE_ELEMENTS = 4096;

enum AddLinearMode : uint32_t {
    ADD_BOTH_CONTIGUOUS = 0,
    ADD_X1_SCALAR = 1,
    ADD_X2_SCALAR = 2,
};

struct AddAxis {
    uint64_t dim;
    uint64_t x1Stride;
    uint64_t x2Stride;
};

struct AddHostMeta {
    uint64_t total = 0;
    uint64_t x1Elements = 0;
    uint64_t x2Elements = 0;
    uint32_t rank = 0;
    uint64_t dims[ADD_MAX_DIMS] = {};
    uint64_t x1Strides[ADD_MAX_DIMS] = {};
    uint64_t x2Strides[ADD_MAX_DIMS] = {};
    std::vector<AddAxis> axes;
};

inline void AddBuildHostMeta(const TensorInfo &x1, const TensorInfo &x2,
                             const TensorInfo &y, AddHostMeta &meta)
{
    meta.axes.clear();
    std::vector<AddAxis> raw(static_cast<size_t>(y.numDims));
    uint64_t x1Stride = 1;
    uint64_t x2Stride = 1;
    uint64_t total = 1;
    const int64_t x1Leading = y.numDims - x1.numDims;
    const int64_t x2Leading = y.numDims - x2.numDims;
    for (int64_t reverse = y.numDims; reverse > 0; --reverse) {
        const int64_t axis = reverse - 1;
        const int64_t x1Axis = axis - x1Leading;
        const int64_t x2Axis = axis - x2Leading;
        const int64_t outputDim = y.shape[axis];
        const int64_t x1Dim = x1Axis < 0 ? 1 : x1.shape[x1Axis];
        const int64_t x2Dim = x2Axis < 0 ? 1 : x2.shape[x2Axis];
        raw[axis] = {
            static_cast<uint64_t>(outputDim),
            x1Dim == 1 && outputDim != 1 ? 0 : x1Stride,
            x2Dim == 1 && outputDim != 1 ? 0 : x2Stride,
        };
        x1Stride *= static_cast<uint64_t>(x1Dim);
        x2Stride *= static_cast<uint64_t>(x2Dim);
        total *= static_cast<uint64_t>(outputDim);
    }
    meta.x1Elements = x1Stride;
    meta.x2Elements = x2Stride;
    meta.total = total;

    // Remove unit axes and merge adjacent axes whose input offsets stay linear.
    for (const AddAxis &current : raw) {
        if (current.dim == 1) {
            continue;
        }
        if (!meta.axes.empty()) {
            AddAxis &previous = meta.axes.back();
            const bool x1Merge = previous.x1Stride == current.dim * current.x1Stride;
            const bool x2Merge = previous.x2Stride == current.dim * current.x2Stride;
            if (x1Merge && x2Merge) {
                previous.dim *= current.dim;
                previous.x1Stride = current.x1Stride;
                previous.x2Stride = current.x2Stride;
                continue;
            }
        }
        meta.axes.push_back(current);
    }

    if (meta.axes.empty()) {
        meta.axes.push_back({1, 0, 0});
    }
    meta.rank = static_cast<uint32_t>(
        std::min(meta.axes.size(), static_cast<size_t>(ADD_MAX_DIMS)));
    const size_t firstStoredAxis = meta.axes.size() - meta.rank;
    for (uint32_t axis = 0; axis < meta.rank; ++axis) {
        const AddAxis &source = meta.axes[firstStoredAxis + axis];
        meta.dims[axis] = source.dim;
        meta.x1Strides[axis] = source.x1Stride;
        meta.x2Strides[axis] = source.x2Stride;
    }
}

inline uint32_t AddSelectCoreCount(uint64_t workItems, int64_t availableCoreNum)
{
    if (workItems == 0 || availableCoreNum <= 0) {
        return 0;
    }
    return static_cast<uint32_t>(std::min(workItems,
                                          static_cast<uint64_t>(availableCoreNum)));
}

inline uint64_t AddTrailingSpan(const AddHostMeta &meta, AddLinearMode mode)
{
    uint64_t span = 1;
    for (uint32_t reverse = meta.rank; reverse > 0; --reverse) {
        const uint32_t axis = reverse - 1;
        bool compatible = false;
        if (mode == ADD_BOTH_CONTIGUOUS) {
            compatible = meta.x1Strides[axis] == span && meta.x2Strides[axis] == span;
        } else if (mode == ADD_X1_SCALAR) {
            compatible = meta.x1Strides[axis] == 0 && meta.x2Strides[axis] == span;
        } else {
            compatible = meta.x2Strides[axis] == 0 && meta.x1Strides[axis] == span;
        }
        if (!compatible) {
            break;
        }
        span *= meta.dims[axis];
    }
    return span;
}

inline GM_ADDR AddByteOffset(GM_ADDR address, uint64_t bytes)
{
    return reinterpret_cast<GM_ADDR>(reinterpret_cast<uintptr_t>(address) + bytes);
}

} // namespace

__aicore__ inline float AddScalarValue(float lhs, float rhs)
{
    return lhs + rhs;
}

__aicore__ inline half AddScalarValue(half lhs, half rhs)
{
    return static_cast<half>(static_cast<float>(lhs) + static_cast<float>(rhs));
}

__aicore__ inline int32_t AddScalarValue(int32_t lhs, int32_t rhs)
{
    return static_cast<int32_t>(static_cast<uint32_t>(lhs) + static_cast<uint32_t>(rhs));
}

__aicore__ inline uint32_t AddScalarValue(uint32_t lhs, uint32_t rhs)
{
    return lhs + rhs;
}

template <typename T>
__aicore__ inline void AddLinearVectorImpl(GM_ADDR x1, GM_ADDR x2, GM_ADDR y,
                                           uint64_t total, uint32_t coreNum,
                                           uint32_t mode)
{
    constexpr uint32_t alignElements = 32 / sizeof(T);
    constexpr uint32_t tileElements = ADD_TILE_BYTES / sizeof(T);
    const uint64_t alignedTotal = total / alignElements * alignElements;
    const uint64_t workUnits = alignedTotal / alignElements;
    const uint64_t block = static_cast<uint64_t>(AscendC::GetBlockIdx());
    const uint64_t begin = workUnits * block / coreNum * alignElements;
    const uint64_t end = workUnits * (block + 1) / coreNum * alignElements;

    AscendC::GlobalTensor<T> x1Gm;
    AscendC::GlobalTensor<T> x2Gm;
    AscendC::GlobalTensor<T> yGm;
    x1Gm.SetGlobalBuffer((__gm__ T *)x1, total);
    x2Gm.SetGlobalBuffer((__gm__ T *)x2, total);
    yGm.SetGlobalBuffer((__gm__ T *)y, total);

    if (begin < end) {
        AscendC::TPipe pipe;
        AscendC::TQue<AscendC::QuePosition::VECIN, ADD_BUFFER_NUM> x1Queue;
        AscendC::TQue<AscendC::QuePosition::VECIN, ADD_BUFFER_NUM> x2Queue;
        AscendC::TQue<AscendC::QuePosition::VECOUT, ADD_BUFFER_NUM> yQueue;
        pipe.InitBuffer(x1Queue, ADD_BUFFER_NUM, ADD_TILE_BYTES);
        pipe.InitBuffer(x2Queue, ADD_BUFFER_NUM, ADD_TILE_BYTES);
        pipe.InitBuffer(yQueue, ADD_BUFFER_NUM, ADD_TILE_BYTES);

        const T x1Scalar = mode == ADD_X1_SCALAR ? x1Gm.GetValue(0) : static_cast<T>(0);
        const T x2Scalar = mode == ADD_X2_SCALAR ? x2Gm.GetValue(0) : static_cast<T>(0);
        for (uint64_t offset = begin; offset < end;) {
            const uint32_t count = static_cast<uint32_t>(
                end - offset > tileElements ? tileElements : end - offset);
            AscendC::LocalTensor<T> x1Local = x1Queue.AllocTensor<T>();
            AscendC::LocalTensor<T> x2Local = x2Queue.AllocTensor<T>();
            if (mode == ADD_X1_SCALAR) {
                AscendC::Duplicate(x1Local, x1Scalar, count);
            } else {
                AscendC::DataCopy(x1Local, x1Gm[offset], count);
            }
            if (mode == ADD_X2_SCALAR) {
                AscendC::Duplicate(x2Local, x2Scalar, count);
            } else {
                AscendC::DataCopy(x2Local, x2Gm[offset], count);
            }
            x1Queue.EnQue(x1Local);
            x2Queue.EnQue(x2Local);

            x1Local = x1Queue.DeQue<T>();
            x2Local = x2Queue.DeQue<T>();
            AscendC::LocalTensor<T> yLocal = yQueue.AllocTensor<T>();
            AscendC::Add(yLocal, x1Local, x2Local, count);
            yQueue.EnQue(yLocal);
            x1Queue.FreeTensor(x1Local);
            x2Queue.FreeTensor(x2Local);

            yLocal = yQueue.DeQue<T>();
            AscendC::DataCopy(yGm[offset], yLocal, count);
            yQueue.FreeTensor(yLocal);
            offset += count;
        }
    }

    // Basic DataCopy operates on 32-byte units. One core owns the short tail.
    if (block == 0) {
        for (uint64_t index = alignedTotal; index < total; ++index) {
            const T lhs = mode == ADD_X1_SCALAR ? x1Gm.GetValue(0) : x1Gm.GetValue(index);
            const T rhs = mode == ADD_X2_SCALAR ? x2Gm.GetValue(0) : x2Gm.GetValue(index);
            yGm.SetValue(index, AddScalarValue(lhs, rhs));
        }
    }
}

__aicore__ inline void AddLinearFp16Impl(GM_ADDR x1, GM_ADDR x2, GM_ADDR y,
                                         uint64_t total, uint32_t coreNum,
                                         uint32_t mode)
{
    constexpr uint32_t alignElements = 16;
    constexpr uint32_t inputTileBytes = ADD_FP16_TILE_ELEMENTS * sizeof(half);
    constexpr uint32_t fp32TileBytes = ADD_FP16_TILE_ELEMENTS * sizeof(float);
    const uint64_t alignedTotal = total / alignElements * alignElements;
    const uint64_t workUnits = alignedTotal / alignElements;
    const uint64_t block = static_cast<uint64_t>(AscendC::GetBlockIdx());
    const uint64_t begin = workUnits * block / coreNum * alignElements;
    const uint64_t end = workUnits * (block + 1) / coreNum * alignElements;

    AscendC::GlobalTensor<half> x1Gm;
    AscendC::GlobalTensor<half> x2Gm;
    AscendC::GlobalTensor<half> yGm;
    x1Gm.SetGlobalBuffer((__gm__ half *)x1, total);
    x2Gm.SetGlobalBuffer((__gm__ half *)x2, total);
    yGm.SetGlobalBuffer((__gm__ half *)y, total);

    if (begin < end) {
        AscendC::TPipe pipe;
        AscendC::TQue<AscendC::QuePosition::VECIN, ADD_BUFFER_NUM> x1Queue;
        AscendC::TQue<AscendC::QuePosition::VECIN, ADD_BUFFER_NUM> x2Queue;
        AscendC::TQue<AscendC::QuePosition::VECOUT, ADD_BUFFER_NUM> yQueue;
        AscendC::TBuf<AscendC::QuePosition::VECCALC> x1Fp32Buffer;
        AscendC::TBuf<AscendC::QuePosition::VECCALC> x2Fp32Buffer;
        AscendC::TBuf<AscendC::QuePosition::VECCALC> yFp32Buffer;
        pipe.InitBuffer(x1Queue, ADD_BUFFER_NUM, inputTileBytes);
        pipe.InitBuffer(x2Queue, ADD_BUFFER_NUM, inputTileBytes);
        pipe.InitBuffer(yQueue, ADD_BUFFER_NUM, inputTileBytes);
        pipe.InitBuffer(x1Fp32Buffer, fp32TileBytes);
        pipe.InitBuffer(x2Fp32Buffer, fp32TileBytes);
        pipe.InitBuffer(yFp32Buffer, fp32TileBytes);

        const half x1Scalar = mode == ADD_X1_SCALAR ? x1Gm.GetValue(0) : static_cast<half>(0);
        const half x2Scalar = mode == ADD_X2_SCALAR ? x2Gm.GetValue(0) : static_cast<half>(0);
        for (uint64_t offset = begin; offset < end;) {
            const uint32_t count = static_cast<uint32_t>(
                end - offset > ADD_FP16_TILE_ELEMENTS ? ADD_FP16_TILE_ELEMENTS : end - offset);
            AscendC::LocalTensor<half> x1Local = x1Queue.AllocTensor<half>();
            AscendC::LocalTensor<half> x2Local = x2Queue.AllocTensor<half>();
            if (mode == ADD_X1_SCALAR) {
                AscendC::Duplicate(x1Local, x1Scalar, count);
            } else {
                AscendC::DataCopy(x1Local, x1Gm[offset], count);
            }
            if (mode == ADD_X2_SCALAR) {
                AscendC::Duplicate(x2Local, x2Scalar, count);
            } else {
                AscendC::DataCopy(x2Local, x2Gm[offset], count);
            }
            x1Queue.EnQue(x1Local);
            x2Queue.EnQue(x2Local);

            x1Local = x1Queue.DeQue<half>();
            x2Local = x2Queue.DeQue<half>();
            AscendC::LocalTensor<half> yLocal = yQueue.AllocTensor<half>();
            AscendC::LocalTensor<float> x1Fp32 = x1Fp32Buffer.Get<float>();
            AscendC::LocalTensor<float> x2Fp32 = x2Fp32Buffer.Get<float>();
            AscendC::LocalTensor<float> yFp32 = yFp32Buffer.Get<float>();
            AscendC::Cast(x1Fp32, x1Local, AscendC::RoundMode::CAST_NONE, count);
            AscendC::Cast(x2Fp32, x2Local, AscendC::RoundMode::CAST_NONE, count);
            AscendC::Add(yFp32, x1Fp32, x2Fp32, count);
            AscendC::Cast(yLocal, yFp32, AscendC::RoundMode::CAST_RINT, count);
            yQueue.EnQue(yLocal);
            x1Queue.FreeTensor(x1Local);
            x2Queue.FreeTensor(x2Local);

            yLocal = yQueue.DeQue<half>();
            AscendC::DataCopy(yGm[offset], yLocal, count);
            yQueue.FreeTensor(yLocal);
            offset += count;
        }
    }

    if (block == 0) {
        for (uint64_t index = alignedTotal; index < total; ++index) {
            const half lhs = mode == ADD_X1_SCALAR ? x1Gm.GetValue(0) : x1Gm.GetValue(index);
            const half rhs = mode == ADD_X2_SCALAR ? x2Gm.GetValue(0) : x2Gm.GetValue(index);
            yGm.SetValue(index, AddScalarValue(lhs, rhs));
        }
    }
}

#define ADD_META_PARAMETERS \
    uint32_t rank, uint64_t dim0, uint64_t dim1, uint64_t dim2, uint64_t dim3, \
    uint64_t dim4, uint64_t dim5, uint64_t dim6, uint64_t dim7, \
    uint64_t x1Stride0, uint64_t x1Stride1, uint64_t x1Stride2, uint64_t x1Stride3, \
    uint64_t x1Stride4, uint64_t x1Stride5, uint64_t x1Stride6, uint64_t x1Stride7, \
    uint64_t x2Stride0, uint64_t x2Stride1, uint64_t x2Stride2, uint64_t x2Stride3, \
    uint64_t x2Stride4, uint64_t x2Stride5, uint64_t x2Stride6, uint64_t x2Stride7

#define ADD_META_ARGUMENTS \
    rank, dim0, dim1, dim2, dim3, dim4, dim5, dim6, dim7, \
    x1Stride0, x1Stride1, x1Stride2, x1Stride3, x1Stride4, x1Stride5, x1Stride6, x1Stride7, \
    x2Stride0, x2Stride1, x2Stride2, x2Stride3, x2Stride4, x2Stride5, x2Stride6, x2Stride7

__aicore__ inline void AddLoadMeta(uint64_t *dims,
                                   uint64_t *x1Strides,
                                   uint64_t *x2Strides,
                                   ADD_META_PARAMETERS)
{
    dims[0] = dim0;
    dims[1] = dim1;
    dims[2] = dim2;
    dims[3] = dim3;
    dims[4] = dim4;
    dims[5] = dim5;
    dims[6] = dim6;
    dims[7] = dim7;
    x1Strides[0] = x1Stride0;
    x1Strides[1] = x1Stride1;
    x1Strides[2] = x1Stride2;
    x1Strides[3] = x1Stride3;
    x1Strides[4] = x1Stride4;
    x1Strides[5] = x1Stride5;
    x1Strides[6] = x1Stride6;
    x1Strides[7] = x1Stride7;
    x2Strides[0] = x2Stride0;
    x2Strides[1] = x2Stride1;
    x2Strides[2] = x2Stride2;
    x2Strides[3] = x2Stride3;
    x2Strides[4] = x2Stride4;
    x2Strides[5] = x2Stride5;
    x2Strides[6] = x2Stride6;
    x2Strides[7] = x2Stride7;
    (void)rank;
}

__aicore__ inline void AddCalculateOffsets(uint64_t outputIndex, uint32_t rank,
                                           const uint64_t *dims,
                                           const uint64_t *x1Strides,
                                           const uint64_t *x2Strides,
                                           uint64_t &x10Offset, uint64_t &x20Offset)
{
    x10Offset = 0;
    x20Offset = 0;
    for (int32_t axis = static_cast<int32_t>(rank) - 1; axis >= 0; --axis) {
        const uint64_t coordinate = outputIndex % dims[axis];
        outputIndex /= dims[axis];
        x10Offset += coordinate * x1Strides[axis];
        x20Offset += coordinate * x2Strides[axis];
    }
}

template <typename T>
__aicore__ inline void AddSegmentVectorImpl(GM_ADDR x1, GM_ADDR x2, GM_ADDR y,
                                            uint64_t inner, uint64_t groups,
                                            uint32_t coreNum, uint32_t mode,
                                            ADD_META_PARAMETERS)
{
    constexpr uint32_t tileElements = ADD_TILE_BYTES / sizeof(T);
    uint64_t dims[ADD_MAX_DIMS];
    uint64_t x1Strides[ADD_MAX_DIMS];
    uint64_t x2Strides[ADD_MAX_DIMS];
    AddLoadMeta(dims, x1Strides, x2Strides, ADD_META_ARGUMENTS);

    const uint64_t chunksPerGroup = (inner + tileElements - 1) / tileElements;
    const uint64_t chunks = groups * chunksPerGroup;
    const uint64_t block = static_cast<uint64_t>(AscendC::GetBlockIdx());
    const uint64_t chunkBegin = chunks * block / coreNum;
    const uint64_t chunkEnd = chunks * (block + 1) / coreNum;
    const uint64_t total = groups * inner;
    AscendC::GlobalTensor<T> x1Gm;
    AscendC::GlobalTensor<T> x2Gm;
    AscendC::GlobalTensor<T> yGm;
    x1Gm.SetGlobalBuffer((__gm__ T *)x1, total);
    x2Gm.SetGlobalBuffer((__gm__ T *)x2, total);
    yGm.SetGlobalBuffer((__gm__ T *)y, total);

    AscendC::TPipe pipe;
    AscendC::TQue<AscendC::QuePosition::VECIN, ADD_BUFFER_NUM> x1Queue;
    AscendC::TQue<AscendC::QuePosition::VECIN, ADD_BUFFER_NUM> x2Queue;
    AscendC::TQue<AscendC::QuePosition::VECOUT, ADD_BUFFER_NUM> yQueue;
    pipe.InitBuffer(x1Queue, ADD_BUFFER_NUM, ADD_TILE_BYTES);
    pipe.InitBuffer(x2Queue, ADD_BUFFER_NUM, ADD_TILE_BYTES);
    pipe.InitBuffer(yQueue, ADD_BUFFER_NUM, ADD_TILE_BYTES);

    for (uint64_t chunk = chunkBegin; chunk < chunkEnd; ++chunk) {
        const uint64_t group = chunk / chunksPerGroup;
        const uint64_t position = (chunk % chunksPerGroup) * tileElements;
        const uint32_t count = static_cast<uint32_t>(
            inner - position > tileElements ? tileElements : inner - position);
        const uint64_t outputBase = group * inner;
        uint64_t x1Base = 0;
        uint64_t x2Base = 0;
        AddCalculateOffsets(outputBase, rank, dims, x1Strides, x2Strides, x1Base, x2Base);

        AscendC::LocalTensor<T> x1Local = x1Queue.AllocTensor<T>();
        AscendC::LocalTensor<T> x2Local = x2Queue.AllocTensor<T>();
        if (mode == ADD_X1_SCALAR) {
            AscendC::Duplicate(x1Local, x1Gm.GetValue(x1Base), count);
        } else {
            AscendC::DataCopy(x1Local, x1Gm[x1Base + position], count);
        }
        if (mode == ADD_X2_SCALAR) {
            AscendC::Duplicate(x2Local, x2Gm.GetValue(x2Base), count);
        } else {
            AscendC::DataCopy(x2Local, x2Gm[x2Base + position], count);
        }
        x1Queue.EnQue(x1Local);
        x2Queue.EnQue(x2Local);

        x1Local = x1Queue.DeQue<T>();
        x2Local = x2Queue.DeQue<T>();
        AscendC::LocalTensor<T> yLocal = yQueue.AllocTensor<T>();
        AscendC::Add(yLocal, x1Local, x2Local, count);
        yQueue.EnQue(yLocal);
        x1Queue.FreeTensor(x1Local);
        x2Queue.FreeTensor(x2Local);

        yLocal = yQueue.DeQue<T>();
        AscendC::DataCopy(yGm[outputBase + position], yLocal, count);
        yQueue.FreeTensor(yLocal);
    }
}

__aicore__ inline void AddSegmentFp16Impl(GM_ADDR x1, GM_ADDR x2, GM_ADDR y,
                                          uint64_t inner, uint64_t groups,
                                          uint32_t coreNum, uint32_t mode,
                                          ADD_META_PARAMETERS)
{
    constexpr uint32_t inputTileBytes = ADD_FP16_TILE_ELEMENTS * sizeof(half);
    constexpr uint32_t fp32TileBytes = ADD_FP16_TILE_ELEMENTS * sizeof(float);
    uint64_t dims[ADD_MAX_DIMS];
    uint64_t x1Strides[ADD_MAX_DIMS];
    uint64_t x2Strides[ADD_MAX_DIMS];
    AddLoadMeta(dims, x1Strides, x2Strides, ADD_META_ARGUMENTS);

    const uint64_t chunksPerGroup =
        (inner + ADD_FP16_TILE_ELEMENTS - 1) / ADD_FP16_TILE_ELEMENTS;
    const uint64_t chunks = groups * chunksPerGroup;
    const uint64_t block = static_cast<uint64_t>(AscendC::GetBlockIdx());
    const uint64_t chunkBegin = chunks * block / coreNum;
    const uint64_t chunkEnd = chunks * (block + 1) / coreNum;
    const uint64_t total = groups * inner;
    AscendC::GlobalTensor<half> x1Gm;
    AscendC::GlobalTensor<half> x2Gm;
    AscendC::GlobalTensor<half> yGm;
    x1Gm.SetGlobalBuffer((__gm__ half *)x1, total);
    x2Gm.SetGlobalBuffer((__gm__ half *)x2, total);
    yGm.SetGlobalBuffer((__gm__ half *)y, total);

    AscendC::TPipe pipe;
    AscendC::TQue<AscendC::QuePosition::VECIN, ADD_BUFFER_NUM> x1Queue;
    AscendC::TQue<AscendC::QuePosition::VECIN, ADD_BUFFER_NUM> x2Queue;
    AscendC::TQue<AscendC::QuePosition::VECOUT, ADD_BUFFER_NUM> yQueue;
    AscendC::TBuf<AscendC::QuePosition::VECCALC> x1Fp32Buffer;
    AscendC::TBuf<AscendC::QuePosition::VECCALC> x2Fp32Buffer;
    AscendC::TBuf<AscendC::QuePosition::VECCALC> yFp32Buffer;
    pipe.InitBuffer(x1Queue, ADD_BUFFER_NUM, inputTileBytes);
    pipe.InitBuffer(x2Queue, ADD_BUFFER_NUM, inputTileBytes);
    pipe.InitBuffer(yQueue, ADD_BUFFER_NUM, inputTileBytes);
    pipe.InitBuffer(x1Fp32Buffer, fp32TileBytes);
    pipe.InitBuffer(x2Fp32Buffer, fp32TileBytes);
    pipe.InitBuffer(yFp32Buffer, fp32TileBytes);

    for (uint64_t chunk = chunkBegin; chunk < chunkEnd; ++chunk) {
        const uint64_t group = chunk / chunksPerGroup;
        const uint64_t position = (chunk % chunksPerGroup) * ADD_FP16_TILE_ELEMENTS;
        const uint32_t count = static_cast<uint32_t>(
            inner - position > ADD_FP16_TILE_ELEMENTS ? ADD_FP16_TILE_ELEMENTS : inner - position);
        const uint64_t outputBase = group * inner;
        uint64_t x1Base = 0;
        uint64_t x2Base = 0;
        AddCalculateOffsets(outputBase, rank, dims, x1Strides, x2Strides, x1Base, x2Base);

        AscendC::LocalTensor<half> x1Local = x1Queue.AllocTensor<half>();
        AscendC::LocalTensor<half> x2Local = x2Queue.AllocTensor<half>();
        if (mode == ADD_X1_SCALAR) {
            AscendC::Duplicate(x1Local, x1Gm.GetValue(x1Base), count);
        } else {
            AscendC::DataCopy(x1Local, x1Gm[x1Base + position], count);
        }
        if (mode == ADD_X2_SCALAR) {
            AscendC::Duplicate(x2Local, x2Gm.GetValue(x2Base), count);
        } else {
            AscendC::DataCopy(x2Local, x2Gm[x2Base + position], count);
        }
        x1Queue.EnQue(x1Local);
        x2Queue.EnQue(x2Local);

        x1Local = x1Queue.DeQue<half>();
        x2Local = x2Queue.DeQue<half>();
        AscendC::LocalTensor<half> yLocal = yQueue.AllocTensor<half>();
        AscendC::LocalTensor<float> x1Fp32 = x1Fp32Buffer.Get<float>();
        AscendC::LocalTensor<float> x2Fp32 = x2Fp32Buffer.Get<float>();
        AscendC::LocalTensor<float> yFp32 = yFp32Buffer.Get<float>();
        AscendC::Cast(x1Fp32, x1Local, AscendC::RoundMode::CAST_NONE, count);
        AscendC::Cast(x2Fp32, x2Local, AscendC::RoundMode::CAST_NONE, count);
        AscendC::Add(yFp32, x1Fp32, x2Fp32, count);
        AscendC::Cast(yLocal, yFp32, AscendC::RoundMode::CAST_RINT, count);
        yQueue.EnQue(yLocal);
        x1Queue.FreeTensor(x1Local);
        x2Queue.FreeTensor(x2Local);

        yLocal = yQueue.DeQue<half>();
        AscendC::DataCopy(yGm[outputBase + position], yLocal, count);
        yQueue.FreeTensor(yLocal);
    }
}

template <typename T>
__aicore__ inline void AddGenericBroadcastImpl(GM_ADDR x1, GM_ADDR x2, GM_ADDR y,
                                               uint64_t total, uint32_t coreNum,
                                               ADD_META_PARAMETERS)
{
    uint64_t dims[ADD_MAX_DIMS];
    uint64_t x1Strides[ADD_MAX_DIMS];
    uint64_t x2Strides[ADD_MAX_DIMS];
    AddLoadMeta(dims, x1Strides, x2Strides, ADD_META_ARGUMENTS);

    const uint64_t block = static_cast<uint64_t>(AscendC::GetBlockIdx());
    const uint64_t begin = total * block / coreNum;
    const uint64_t end = total * (block + 1) / coreNum;
    AscendC::GlobalTensor<T> x1Gm;
    AscendC::GlobalTensor<T> x2Gm;
    AscendC::GlobalTensor<T> yGm;
    x1Gm.SetGlobalBuffer((__gm__ T *)x1, total);
    x2Gm.SetGlobalBuffer((__gm__ T *)x2, total);
    yGm.SetGlobalBuffer((__gm__ T *)y, total);

    uint64_t coordinates[ADD_MAX_DIMS] = {};
    uint64_t remainder = begin;
    uint64_t x10Offset = 0;
    uint64_t x20Offset = 0;
    for (int32_t axis = static_cast<int32_t>(rank) - 1; axis >= 0; --axis) {
        coordinates[axis] = remainder % dims[axis];
        remainder /= dims[axis];
        x10Offset += coordinates[axis] * x1Strides[axis];
        x20Offset += coordinates[axis] * x2Strides[axis];
    }

    for (uint64_t outputIndex = begin; outputIndex < end; ++outputIndex) {
        yGm.SetValue(outputIndex,
                     AddScalarValue(x1Gm.GetValue(x10Offset), x2Gm.GetValue(x20Offset)));
        for (int32_t axis = static_cast<int32_t>(rank) - 1; axis >= 0; --axis) {
            ++coordinates[axis];
            x10Offset += x1Strides[axis];
            x20Offset += x2Strides[axis];
            if (coordinates[axis] < dims[axis]) {
                break;
            }
            coordinates[axis] = 0;
            x10Offset -= dims[axis] * x1Strides[axis];
            x20Offset -= dims[axis] * x2Strides[axis];
        }
    }
}

extern "C" __global__ __vector__ void add_linear_fp32(
    GM_ADDR x1, GM_ADDR x2, GM_ADDR y, uint64_t total, uint32_t coreNum, uint32_t mode)
{
    AddLinearVectorImpl<float>(x1, x2, y, total, coreNum, mode);
}

extern "C" __global__ __vector__ void add_linear_fp16(
    GM_ADDR x1, GM_ADDR x2, GM_ADDR y, uint64_t total, uint32_t coreNum, uint32_t mode)
{
    AddLinearFp16Impl(x1, x2, y, total, coreNum, mode);
}

extern "C" __global__ __vector__ void add_linear_int32(
    GM_ADDR x1, GM_ADDR x2, GM_ADDR y, uint64_t total, uint32_t coreNum, uint32_t mode)
{
    AddLinearVectorImpl<int32_t>(x1, x2, y, total, coreNum, mode);
}

extern "C" __global__ __vector__ void add_segment_fp32(
    GM_ADDR x1, GM_ADDR x2, GM_ADDR y, uint64_t inner, uint64_t groups,
    uint32_t coreNum, uint32_t mode, ADD_META_PARAMETERS)
{
    AddSegmentVectorImpl<float>(x1, x2, y, inner, groups, coreNum, mode, ADD_META_ARGUMENTS);
}

extern "C" __global__ __vector__ void add_segment_fp16(
    GM_ADDR x1, GM_ADDR x2, GM_ADDR y, uint64_t inner, uint64_t groups,
    uint32_t coreNum, uint32_t mode, ADD_META_PARAMETERS)
{
    AddSegmentFp16Impl(x1, x2, y, inner, groups, coreNum, mode, ADD_META_ARGUMENTS);
}

extern "C" __global__ __vector__ void add_segment_int32(
    GM_ADDR x1, GM_ADDR x2, GM_ADDR y, uint64_t inner, uint64_t groups,
    uint32_t coreNum, uint32_t mode, ADD_META_PARAMETERS)
{
    AddSegmentVectorImpl<int32_t>(x1, x2, y, inner, groups, coreNum, mode, ADD_META_ARGUMENTS);
}

extern "C" __global__ __vector__ void add_generic_fp32(
    GM_ADDR x1, GM_ADDR x2, GM_ADDR y, uint64_t total, uint32_t coreNum,
    ADD_META_PARAMETERS)
{
    AddGenericBroadcastImpl<float>(x1, x2, y, total, coreNum, ADD_META_ARGUMENTS);
}

extern "C" __global__ __vector__ void add_generic_fp16(
    GM_ADDR x1, GM_ADDR x2, GM_ADDR y, uint64_t total, uint32_t coreNum,
    ADD_META_PARAMETERS)
{
    AddGenericBroadcastImpl<half>(x1, x2, y, total, coreNum, ADD_META_ARGUMENTS);
}

extern "C" __global__ __vector__ void add_generic_int32(
    GM_ADDR x1, GM_ADDR x2, GM_ADDR y, uint64_t total, uint32_t coreNum,
    ADD_META_PARAMETERS)
{
    AddGenericBroadcastImpl<uint32_t>(x1, x2, y, total, coreNum, ADD_META_ARGUMENTS);
}

#define ADD_HOST_META_ARGUMENTS \
    meta.rank, meta.dims[0], meta.dims[1], meta.dims[2], meta.dims[3], \
    meta.dims[4], meta.dims[5], meta.dims[6], meta.dims[7], \
    meta.x1Strides[0], meta.x1Strides[1], meta.x1Strides[2], meta.x1Strides[3], \
    meta.x1Strides[4], meta.x1Strides[5], meta.x1Strides[6], meta.x1Strides[7], \
    meta.x2Strides[0], meta.x2Strides[1], meta.x2Strides[2], meta.x2Strides[3], \
    meta.x2Strides[4], meta.x2Strides[5], meta.x2Strides[6], meta.x2Strides[7]

extern "C" void run_kernel(GM_ADDR x1, const TensorGroupInfo &info_x1,
                           GM_ADDR x2, const TensorGroupInfo &info_x2,
                           GM_ADDR y, const TensorGroupInfo &info_y,
                           int64_t availableCoreNum, aclrtStream stream)
{
    (void)stream;
    if (x1 == nullptr || x2 == nullptr || y == nullptr || availableCoreNum <= 0 ||
        info_x1.numTensors < 1 || info_x2.numTensors < 1 || info_y.numTensors < 1 ||
        info_x1.tensors == nullptr || info_x2.tensors == nullptr || info_y.tensors == nullptr) {
        return;
    }

    const TensorInfo &x1Info = info_x1.tensors[0];
    const TensorInfo &x2Info = info_x2.tensors[0];
    const TensorInfo &yInfo = info_y.tensors[0];
    if (x1Info.dtype != x2Info.dtype || x1Info.dtype != yInfo.dtype ||
        (x1Info.dtype != 0 && x1Info.dtype != 1 && x1Info.dtype != 5) ||
        x1Info.numDims < 0 || x2Info.numDims < 0 || yInfo.numDims < 0 ||
        x1Info.numDims > yInfo.numDims || x2Info.numDims > yInfo.numDims ||
        yInfo.numDims != std::max(x1Info.numDims, x2Info.numDims) ||
        (x1Info.numDims && x1Info.shape == nullptr) ||
        (x2Info.numDims && x2Info.shape == nullptr) ||
        (yInfo.numDims && yInfo.shape == nullptr)) {
        return;
    }

    AddHostMeta meta;
    AddBuildHostMeta(x1Info, x2Info, yInfo, meta);
    if (meta.total == 0) {
        return;
    }

    AddLinearMode linearMode = ADD_BOTH_CONTIGUOUS;
    bool useLinear = false;
    if (meta.x1Elements == meta.total && meta.x2Elements == meta.total) {
        useLinear = true;
    } else if (meta.x1Elements == 1 && meta.x2Elements == meta.total) {
        linearMode = ADD_X1_SCALAR;
        useLinear = true;
    } else if (meta.x2Elements == 1 && meta.x1Elements == meta.total) {
        linearMode = ADD_X2_SCALAR;
        useLinear = true;
    }

    const uint32_t elementBytes = x1Info.dtype == 1 ? 2U : 4U;
    const uint64_t alignElements = 32U / elementBytes;
    if (useLinear) {
        const uint32_t coreNum = AddSelectCoreCount(
            std::max<uint64_t>(meta.total / alignElements, 1), availableCoreNum);
        if (x1Info.dtype == 0) {
            add_linear_fp32<<<coreNum, nullptr, stream>>>(
                x1, x2, y, meta.total, coreNum, static_cast<uint32_t>(linearMode));
        } else if (x1Info.dtype == 1) {
            add_linear_fp16<<<coreNum, nullptr, stream>>>(
                x1, x2, y, meta.total, coreNum, static_cast<uint32_t>(linearMode));
        } else {
            add_linear_int32<<<coreNum, nullptr, stream>>>(
                x1, x2, y, meta.total, coreNum, static_cast<uint32_t>(linearMode));
        }
        return;
    }

    if (meta.axes.size() > ADD_MAX_DIMS) {
        uint64_t sliceElements = 1;
        for (uint32_t axis = 0; axis < meta.rank; ++axis) {
            sliceElements *= meta.dims[axis];
        }
        const uint64_t prefixCount = meta.total / sliceElements;
        const size_t prefixRank = meta.axes.size() - meta.rank;
        const uint32_t coreNum = AddSelectCoreCount(sliceElements, availableCoreNum);
        for (uint64_t prefix = 0; prefix < prefixCount; ++prefix) {
            uint64_t remainder = prefix;
            uint64_t x1Base = 0;
            uint64_t x2Base = 0;
            for (size_t reverse = prefixRank; reverse > 0; --reverse) {
                const size_t axis = reverse - 1;
                const uint64_t coordinate = remainder % meta.axes[axis].dim;
                remainder /= meta.axes[axis].dim;
                x1Base += coordinate * meta.axes[axis].x1Stride;
                x2Base += coordinate * meta.axes[axis].x2Stride;
            }

            GM_ADDR sliceX1 = AddByteOffset(x1, x1Base * elementBytes);
            GM_ADDR sliceX2 = AddByteOffset(x2, x2Base * elementBytes);
            GM_ADDR sliceY = AddByteOffset(y, prefix * sliceElements * elementBytes);
            if (x1Info.dtype == 0) {
                add_generic_fp32<<<coreNum, nullptr, stream>>>(
                    sliceX1, sliceX2, sliceY, sliceElements, coreNum,
                    ADD_HOST_META_ARGUMENTS);
            } else if (x1Info.dtype == 1) {
                add_generic_fp16<<<coreNum, nullptr, stream>>>(
                    sliceX1, sliceX2, sliceY, sliceElements, coreNum,
                    ADD_HOST_META_ARGUMENTS);
            } else {
                add_generic_int32<<<coreNum, nullptr, stream>>>(
                    sliceX1, sliceX2, sliceY, sliceElements, coreNum,
                    ADD_HOST_META_ARGUMENTS);
            }
        }
        return;
    }

    uint64_t bestSpan = 1;
    AddLinearMode segmentMode = ADD_BOTH_CONTIGUOUS;
    const AddLinearMode candidates[3] = {
        ADD_BOTH_CONTIGUOUS, ADD_X1_SCALAR, ADD_X2_SCALAR,
    };
    for (AddLinearMode candidate : candidates) {
        const uint64_t span = AddTrailingSpan(meta, candidate);
        if (span >= alignElements && span % alignElements == 0 && span > bestSpan) {
            bestSpan = span;
            segmentMode = candidate;
        }
    }

    if (bestSpan > 1) {
        const uint64_t groups = meta.total / bestSpan;
        const uint64_t tileElements = x1Info.dtype == 1
            ? ADD_FP16_TILE_ELEMENTS
            : ADD_TILE_BYTES / elementBytes;
        const uint64_t chunks = groups * ((bestSpan + tileElements - 1) / tileElements);
        const uint32_t coreNum = AddSelectCoreCount(chunks, availableCoreNum);
        if (x1Info.dtype == 0) {
            add_segment_fp32<<<coreNum, nullptr, stream>>>(
                x1, x2, y, bestSpan, groups, coreNum, static_cast<uint32_t>(segmentMode),
                ADD_HOST_META_ARGUMENTS);
        } else if (x1Info.dtype == 1) {
            add_segment_fp16<<<coreNum, nullptr, stream>>>(
                x1, x2, y, bestSpan, groups, coreNum, static_cast<uint32_t>(segmentMode),
                ADD_HOST_META_ARGUMENTS);
        } else {
            add_segment_int32<<<coreNum, nullptr, stream>>>(
                x1, x2, y, bestSpan, groups, coreNum, static_cast<uint32_t>(segmentMode),
                ADD_HOST_META_ARGUMENTS);
        }
        return;
    }

    const uint32_t coreNum = AddSelectCoreCount(meta.total, availableCoreNum);
    if (x1Info.dtype == 0) {
        add_generic_fp32<<<coreNum, nullptr, stream>>>(
            x1, x2, y, meta.total, coreNum, ADD_HOST_META_ARGUMENTS);
    } else if (x1Info.dtype == 1) {
        add_generic_fp16<<<coreNum, nullptr, stream>>>(
            x1, x2, y, meta.total, coreNum, ADD_HOST_META_ARGUMENTS);
    } else {
        add_generic_int32<<<coreNum, nullptr, stream>>>(
            x1, x2, y, meta.total, coreNum, ADD_HOST_META_ARGUMENTS);
    }
}

#undef ADD_HOST_META_ARGUMENTS
#undef ADD_META_ARGUMENTS
#undef ADD_META_PARAMETERS

附:路径选择速查

shape 情况 命中路径
x1==x2 同形连续 linear(BOTH_CONTIGUOUS)
x1 标量 + x2 大张量 linear(X1_SCALAR,Duplicate)
x=2,3 + x=3(行广播,最右连续) segment(inner=3, groups=2)
x=2,1 + x=2,3(列广播) segment 或 generic(看尾段是否对齐)
任意不规则高维广播 generic(odometer)
合并后维数 > 8 host 切前缀 × generic 后缀
相关推荐
java资料站2 小时前
八、SpringAl 会话记忆,历史对话,隔离记忆
ai
七夜zippoe3 小时前
短期记忆 vs 长期记忆:Agent 记忆系统的分层架构设计
ai·agent·长期记忆·短期记忆·记忆系统
三声三视3 小时前
周日晚九点半,两份 JD 喂进 tri-jobhunt:全中的那份和缺一项的,都拿了 70 分
人工智能·ai·skillhub·tri-skill·tri-jobhunt
葡萄城技术团队3 小时前
Jev 爆火之后,给企业应用配一个 AI 决策模型
ai
解决小子3 小时前
2026年中国就业情况报告
ai·职场和发展·创业创新·业界资讯·就业
运维开发王义杰5 小时前
Gemini 3.8 语音大模型 GA:当 TTS 学会“演戏”,家庭英语学习与视频创作价值拆解
ai
XLYcmy5 小时前
AI 时代,MOM(制造运营管理系统)该如何演进? 上
ai·llm·agent·模型·mom·harness·工业系统
启雀AI5 小时前
培训管理系统的 AI 智能陪练完整功能逻辑,以家电门店销售为例的剧本框架
人工智能·ai·软件需求·培训系统·培训平台