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

115 lines
4.6 KiB
Markdown
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# 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。演示一个小 demo`torch.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_val`reduce max
- Pass 2: 写 `exp(x - max)` 到 out + 累加 sumreduce sum
- Pass 3: `out[i] /= sum`
4. **Online SoftmaxTileLang 版):**
- 用 log-sum-exp (lse) 滚动更新,两趟完成
- `exp2`/`log2` 比 `exp`/`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.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](../../docs/diagrams/softmax-pipeline.md) 中的 LSE 滚动更新详解。**
---
## [3min] 测试验证
### 跑测试
```bash
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 版一样)。分析性能提升幅度,写一段注释解释为什么快。"