forked from ccf-ai-infra/Intro-ops
115 lines
4.6 KiB
Markdown
115 lines
4.6 KiB
Markdown
# 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 + 累加 sum(reduce sum)
|
||
- Pass 3: `out[i] /= sum`
|
||
|
||
4. **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](../../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:忘记减 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 版一样)。分析性能提升幅度,写一段注释解释为什么快。"
|