Intro-ops/course/scripts/04_softmax_script.md

4.6 KiB
Raw Blame History

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

讲解要点:

  1. 朴素公式的问题: exp(100) 溢出 → NaN。演示一个小 demotorch.tensor([100.0, 200.0, 300.0]) 的朴素 exp 直接炸

  2. 稳定的公式:

    • softmax(x_i) = exp(x_i - max) / Σ exp(x_j - max)
    • 为什么等价?分子分母同除 exp(max)
    • 为什么稳定?x_i - max ≤ 0,所以 exp() ≤ 1,永不溢出
  3. 三趟扫描流程:

    • Pass 1: 求 max_valreduce max
    • Pass 2: 写 exp(x - max) 到 out + 累加 sumreduce sum
    • Pass 3: out[i] /= sum
  4. Online SoftmaxTileLang 版):

    • 用 log-sum-exp (lse) 滚动更新,两趟完成
    • exp2/log2exp/log 在硬件上更快

过渡语: "理论比代码复杂——但代码本身并不可怕。来写。"


[10min] 代码实操

CUDA Kernel5min

打开文件: ops/softmax/nvidia/kernel.cuh

Step 1: Pass 1 — 求行最大值1.5min

"和 reduce_sum 的归约一模一样——只是把 += 改成 max()。注意初始值max 初始化为 in[row * cols] 而不是 0——因为输入可能是全负数。"

Step 2: Pass 2 — exp + 累加 sum2min

"第二趟扫描做了两件事:计算 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 Kernel5min

打开文件: ops/softmax/tilelang/kernel.py

关键概念:

  • "log2_e = 1.44269504exp(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忘记减 max30s

"最经典的 bug。输入 [100, 200, 300]——朴素 exp 输出全是 NaN。减了 max 后正常输出 [0, 0, 1]。"

错误 2Pass 2 和 Pass 3 之间没同步30s

"Pass 2 的 sum 归约完成后thread 0 有正确的 sum——但 thread 1 可能还在写 smem。不 sync 的话 thread 1 在 Pass 3 读到的是旧 sum。"

错误 3max 初始化为 0 而非第一个元素30s

"如果输入全是负数——max=0 比真实值大。exp(x - 0) 没问题(仍然 ≤ 1但 exp(x - real_max) 的精度更好。不影响正确性但影响精度。"

错误 4TileLang 里 Pass 1 用 T.Parallel30s

"lse 有跨 tile 的状态依赖——必须 T.Serial。用 T.Parallel 会读到未初始化的 lse 值。"


课后挑战

"把 CUDA 的三趟扫描改成两趟 online softmax像 TileLang 版一样)。分析性能提升幅度,写一段注释解释为什么快。"