概述
如果一串算子都是"逐元素"的(对每个元素独立做同样的运算),那么可以把它们合并到一个循环里,每个元素从头到尾算完,中间结果只活在寄存器里。
形如 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(只读输入一次、写输出一次),是最简单、收益最直接的融合模式。
这也是为什么激活函数、缩放平移、逐元素加等算子在几乎所有推理引擎里都是"免费"融合的。
网硕互联帮助中心







评论前必须登录!
注册