4.6 KiB
Softmax 算子 — 教学视频脚本
时长:15-20 分钟 | 难度:挑战
[5min] 概念讲解
开场(30s)
"softmax 是四个算子中最难的。它结合了 copy 的逐元素、reduce_sum 的归约——还要解决数值溢出问题。今天我们不只写正确的 softmax,还要理解为什么'正确'不是理所当然的。"
算子在深度学习中的用途(1min)
"softmax 是 transformer 的核心——每个 attention block 都调它。分类任务的最后一层也是 softmax。一个效率高 10% 的 softmax kernel 可以提升整个推理管线的吞吐。"
算法推导(3.5min)
关键画面: 展示 softmax 流水线图(docs/diagrams/softmax-pipeline.md)
讲解要点:
-
朴素公式的问题:
exp(100)溢出 → NaN。演示一个小 demo:torch.tensor([100.0, 200.0, 300.0])的朴素 exp 直接炸 -
稳定的公式:
softmax(x_i) = exp(x_i - max) / Σ exp(x_j - max)- 为什么等价?分子分母同除
exp(max) - 为什么稳定?
x_i - max ≤ 0,所以exp() ≤ 1,永不溢出
-
三趟扫描流程:
- Pass 1: 求
max_val(reduce max) - Pass 2: 写
exp(x - max)到 out + 累加 sum(reduce sum) - Pass 3:
out[i] /= sum
- Pass 1: 求
-
Online Softmax(TileLang 版):
- 用 log-sum-exp (lse) 滚动更新,两趟完成
exp2/log2比exp/log在硬件上更快
过渡语: "理论比代码复杂——但代码本身并不可怕。来写。"
[10min] 代码实操
CUDA Kernel(5min)
打开文件: ops/softmax/nvidia/kernel.cuh
Step 1: Pass 1 — 求行最大值(1.5min)
"和 reduce_sum 的归约一模一样——只是把 += 改成 max()。注意初始值:max 初始化为 in[row * cols] 而不是 0——因为输入可能是全负数。"
Step 2: Pass 2 — exp + 累加 sum(2min)
"第二趟扫描做了两件事:计算 exp(x - max) 写到 out,同时累加 sum。为什么要写 out?因为第三趟需要这些中间值——不写的话第三趟还得重新从全局内存读输入再算 exp,浪费带宽。"
Step 3: Pass 3 — 归一化(1.5min)
"最后一行 __syncthreads() 确保 sum 已经归约完成。除了 thread 0 知道 sum 之外——其他线程不需要知道 sum 就能做除法吗?不对——sum 还在 smem[0] 里,每个线程需要读 smem[0] 来做除法。所以第三趟之前也要 sync。"
TileLang Kernel(5min)
打开文件: ops/softmax/tilelang/kernel.py
关键概念:
- "
log2_e = 1.44269504:exp(x) = 2^(x * log2(e))的转换系数" - "Pass 1 用
T.Serial遍历列分块——和 reduce_sum 一样,lse 有状态依赖" - "
T.reduce_max+T.reduce_sum:编译器为你选择最优的归约策略" - "lse 更新公式:
m_new = max(lse, max(tile)); lse = m_new + log2(exp2(lse - m_new) + sum(exp2(tile - m_new)))"
不展开完整公式推导——引导学员去看 docs/diagrams/softmax-pipeline.md 中的 LSE 滚动更新详解。
[3min] 测试验证
跑测试
PYTHONPATH=python:. CAMP_BUILD_DIR=build-nvidia \
pytest tests/op_tests/test_softmax.py -v --backend nvidia
重点展示: "测试里会验证两件事——正确性(每行和 ≈ 1.0)和数值稳定性(大数据不出 NaN)。如果你偷懒没写减 max,大数值用例会直接炸。"
Benchmark 解读
"softmax 的算术强度比 reduce_sum 更高——exp/log 是计算密集型操作。优化方向从内存带宽转向计算吞吐。如果你的算力利用率低,可能是 exp 计算没被流水线化。"
[2min] 常见错误演示
错误 1:忘记减 max(30s)
"最经典的 bug。输入 [100, 200, 300]——朴素 exp 输出全是 NaN。减了 max 后正常输出 [0, 0, 1]。"
错误 2:Pass 2 和 Pass 3 之间没同步(30s)
"Pass 2 的 sum 归约完成后,thread 0 有正确的 sum——但 thread 1 可能还在写 smem。不 sync 的话 thread 1 在 Pass 3 读到的是旧 sum。"
错误 3:max 初始化为 0 而非第一个元素(30s)
"如果输入全是负数——max=0 比真实值大。exp(x - 0) 没问题(仍然 ≤ 1),但 exp(x - real_max) 的精度更好。不影响正确性但影响精度。"
错误 4:TileLang 里 Pass 1 用 T.Parallel(30s)
"lse 有跨 tile 的状态依赖——必须 T.Serial。用 T.Parallel 会读到未初始化的 lse 值。"
课后挑战
"把 CUDA 的三趟扫描改成两趟 online softmax(像 TileLang 版一样)。分析性能提升幅度,写一段注释解释为什么快。"