云计算百科
云计算领域专业知识百科平台

逐元素算子链融合

概述

如果一串算子都是"逐元素"的(对每个元素独立做同样的运算),那么可以把它们合并到一个循环里,每个元素从头到尾算完,中间结果只活在寄存器里。

形如 y = max(0.1, tanh(x)) * 2 + 1 的逐元素链(激活、归一化、残差加、缩放……)


先理解什么是"逐元素算子"

逐元素算子(Elementwise Operator):对张量里的每个元素独立做同样的运算,输出形状和输入形状相同,且输出坐标 = 输入坐标。

比如:

算子公式输入输出
tanh y[i]=tanh⁡(x[i])y[i] = \\tanh(x[i])y[i]=tanh(x[i]) x[0],x[1],…x[0], x[1], \\dotsx[0],x[1], y[0],y[1],…y[0], y[1], \\dotsy[0],y[1],
缩放 y[i]=x[i]×2y[i] = x[i] \\times 2y[i]=x[i]×2 同上 同上
平移 y[i]=x[i]+1y[i] = x[i] + 1y[i]=x[i]+1 同上 同上
max y[i]=max⁡(x[i],0.1)y[i] = \\max(x[i], 0.1)y[i]=max(x[i],0.1) 同上 同上

关键特征:第 i 个输出元素只依赖第 i 个输入元素,不依赖其他位置的元素。

这就是"逐元素"的含义


问题在哪里?

假设有一个 5 个算子的链:

y = max(0.1, tanh(x) * 2 + 1)

它拆成 5 个算子:

t1 = tanh(x) # ① 双曲正切
t2 = t1 * 2 # ② 缩放
t3 = t2 + 1 # ③ 平移
t4 = max(t3, 0.1) # ④ 下限截断
y = t4 # ⑤ 赋值(可能还有更多)

如果每个算子单独跑一个 kernel:

kernel1: 读 x[N],写 t1[N] → DRAM 流量:读 N + 写 N
kernel2: 读 t1[N],写 t2[N] → DRAM 流量:读 N + 写 N
kernel3: 读 t2[N],写 t3[N] → DRAM 流量:读 N + 写 N
kernel4: 读 t3[N],写 t4[N] → DRAM 流量:读 N + 写 N
kernel5: 读 t4[N],写 y[N] → DRAM 流量:读 N + 写 N

假设每个张量大小是 S 字节(N 个元素,每个元素若干字节),那么不融合的总 DRAM 流量是:
5×(2S)=10S
5×(2S)=10S
5×(2S)=10S

即:每个中间张量都要被写一次、读一次,共 5 对读写。


具体数值例子

为了简单,假设:

  • 张量大小 N = 4(4 个元素)
  • 输入 x = [0.5, 1.0, -0.5, 2.0]
  • 每个元素是 f32(4 字节),所以 S = 4×4 = 16 字节

步骤 1:不融合,逐个算子算

t1 = tanh(x)

tanh⁡(0.5)=0.462,tanh⁡(1.0)=0.762,tanh⁡(−0.5)=−0.462,tanh⁡(2.0)=0.964
\\tanh(0.5) = 0.462, \\quad \\tanh(1.0) = 0.762, \\quad \\tanh(-0.5) = -0.462, \\quad \\tanh(2.0) = 0.964
tanh(0.5)=0.462,tanh(1.0)=0.762,tanh(0.5)=0.462,tanh(2.0)=0.964

t1=[0.462,0.762,−0.462,0.964]
t1 = [0.462, 0.762, -0.462, 0.964]
t1=[0.462,0.762,0.462,0.964]

t2 = t1 × 2

t2=[0.924,1.524,−0.924,1.928]
t2 = [0.924, 1.524, -0.924, 1.928]
t2=[0.924,1.524,0.924,1.928]

t3 = t2 + 1

t2=[1.924,2.524,0.076,2.928]
t2 = [1.924,2.524,0.076,2.928]
t2=[1.924,2.524,0.076,2.928]

t4 = max(t3, 0.1)

t4=[1.924,2.524,0.1,2.928]
t4=[1.924,2.524,0.1,2.928]
t4=[1.924,2.524,0.1,2.928]

y = t4

y=[1.924,2.524,0.1,2.928]
y=[1.924,2.524,0.1,2.928]
y=[1.924,2.524,0.1,2.928]

不融合的最终结果:[1.924, 2.524, 0.1, 2.928]

DRAM 流量:

算子读写
tanh 16 B (x) 16 B (t1)
×2 16 B (t1) 16 B (t2)
+1 16 B (t2) 16 B (t3)
max 16 B (t3) 16 B (t4)
赋值 16 B (t4) 16 B (y)

步骤 2:融合后,一个 kernel 算完

融合后的 kernel 只有一个循环,遍历所有元素:

for i in 0..N:
t1 = tanh(x[i])
t2 = t1 * 2
t3 = t2 + 1
y[i] = max(t3, 0.1)

对于 i=0:x[0]=0.5

  • t1 = tanh(0.5) = 0.462
  • t2 = 0.462 × 2 = 0.924
  • t3 = 0.924 + 1 = 1.924
  • y[0] = max(1.924, 0.1) = 1.924

对于 i=1:x[1]=1.0

  • t1 = tanh(1.0) = 0.762
  • t2 = 0.762 × 2 = 1.524
  • t3 = 1.524 + 1 = 2.524
  • y[1] = max(2.524, 0.1) = 2.524

对于 i=2:x[2]=-0.5

  • t1 = tanh(-0.5) = -0.462
  • t2 = -0.462 × 2 = -0.924
  • t3 = -0.924 + 1 = 0.076
  • y[2] = max(0.076, 0.1) = 0.1

对于 i=3:x[3]=2.0

  • t1 = tanh(2.0) = 0.964
  • t2 = 0.964 × 2 = 1.928
  • t3 = 1.928 + 1 = 2.928
  • y[3] = max(2.928, 0.1) = 2.928

融合后的最终结果:[1.924, 2.524, 0.1, 2.928] ✅ 与不融合完全一致。

DRAM 流量:

操作读写
融合 kernel 16 B (x) 16 B (y)

总计:32 字节(2S = 2×16 = 32)。


融合前后对比

维度不融合融合后
Kernel 数量 5 个 1 个
DRAM 流量 10S = 160 B 2S = 32 B
中间张量 t1~t4 都写回 DRAM,再读出来 只活在寄存器里
Kernel 启动开销 5 次 1 次
最终结果 相同 相同

流量减少到原来的 1/5。如果链上有 50 个逐元素算子,流量减少到 1/50。


为什么逐元素链可以"白送"地融合?

关键原因:输出坐标 = 输入坐标(恒等映射)。

  • 对于第 i 个输出元素 y[i],它只依赖 x[i],不需要访问 x 的其他位置。
  • 所以一个循环遍历 i 从 0 到 N-1,每个 i 独立地把整条链算完,不存在坐标换算负担。

对比卷积:卷积的输出 (n, co, ho, wo) 需要访问输入的窗口 (n, ci, ho * s+kh, wo * s+kw),涉及 stride、padding 等复杂坐标映射。所以卷积融合的索引推导更难,而逐元素链融合是"最简单、最安全"的融合。


什么情况下不能融合?

逐元素链融合的前提是:

所有算子都是逐元素的:不能有 reduce(如 sum、mean)、reshape、transpose、matmul 等跨元素操作。

没有副作用:不能有 print、文件写入、随机数生成(每次结果不同)。

没有分支改变形状:不能有 if 导致输出形状变化。

如果链中混入了一个非逐元素算子(比如中间夹了一个 softmax),融合就会被打断,需要分段处理。


总结

逐元素算子链融合:因为每个输出元素只依赖对应的输入元素(恒等坐标映射),所以可以把整条链塞进一个循环,每个元素从头算到尾,中间结果只活在寄存器里。

DRAM 流量从 2 × 算子数 × S 降到 2S(只读输入一次、写输出一次),是最简单、收益最直接的融合模式。

这也是为什么激活函数、缩放平移、逐元素加等算子在几乎所有推理引擎里都是"免费"融合的。

赞(0)
未经允许不得转载:网硕互联帮助中心 » 逐元素算子链融合
分享到: 更多 (0)

评论 抢沙发

评论前必须登录!