4.6 KiB
Reduce Sum 算子 — 教学视频脚本
时长:15-20 分钟 | 难度:进阶
[5min] 概念讲解
开场(30s)
"前两个算子 copy 和 vector_add,线程之间是零通信的——各自算各自的。今天 reduce_sum 完全不同——线程需要互相通信,通过 shared memory 把部分结果合并为最终结果。"
算子在深度学习中的用途(1min)
"归约(reduction)操作在深度学习中非常常见:LayerNorm 里的 mean/std、attention 里的 softmax 分母 sum、loss 函数的最终求和——都是归约。理解归约是 GPU 编程的第一个分水岭。"
算法推导(3.5min)
关键画面: 展示树形归约流程图(docs/diagrams/tree-reduction.md)
讲解要点:
- 目标:每行 N 个元素求和为 1 个值
- Step 1: 各线程先独立累加自己负责的列 → 存入 shared memory
- Step 2: 树形归约——
stride = blockDim.x/2, /4, ..., 1,每次配对相加 - 为什么叫"树形"——每步参与线程减半,log₂(blockDim.x) 步完成
- 关键:每一步之后必须
__syncthreads()
可视化辅助: 手动画出 8→4→2→1 的归约树
过渡语: "听起来简单,但写起来有两个坑——我们直接在 IDE 里看。"
[10min] 代码实操
CUDA Kernel(6min)
打开文件: ops/reduce_sum/nvidia/kernel.cuh
Step 1: 线程各自累加(2min)
"每个线程用 grid-stride loop 跨步累加自己负责的列——这个和 copy 一样。但这次结果不是写到全局内存,而是写到 shared memory。"
Step 2: shared memory 写入 + 同步(1min)
smem[threadIdx.x] = sum;
__syncthreads(); // 关键!
关键决策点: "为什么这里一定要 __syncthreads()?因为 tree reduction 下一步要读 smem[tid+s]——那是别的线程写的。不同步的话可能读到旧数据。"
Step 3: 树形归约(2min)
for (int s = blockDim.x / 2; s > 0; s >>= 1) {
if (threadIdx.x < s) {
smem[threadIdx.x] += smem[threadIdx.x + s];
}
__syncthreads(); // 在 if 外面!
}
"注意 __syncthreads() 在 if 外面——这是最常见的 bug。如果放进 if,只有一半线程执行同步,另一半跳过,整个 block 死锁。"
Step 4: 输出(1min)
"只有 thread 0 写回全局内存——因为 smem[0] 已经是整行的和。"
TileLang Kernel(4min)
打开文件: ops/reduce_sum/tilelang/kernel.py
边写边讲:
- "外层
T.Parallel分发到行——不同行之间独立" - "内层
T.Serial顺序遍历列分块——因为累加器有状态依赖" - "
T.reduce_sum(tile)替代手写树形归约——编译器为你生成最优代码" - "注意
T.alloc_fragment分配累加器——这是寄存器级别的存储"
对比: "CUDA 版 20+ 行,TileLang 版 8 行。T.reduce_sum 内部帮你处理了同步、bank conflict、warp divergence——这些在 CUDA 里都是你要手写的。"
[3min] 测试验证
跑测试
PYTHONPATH=python:. CAMP_BUILD_DIR=build-nvidia \
pytest tests/op_tests/test_reduce_sum.py -v --backend nvidia
重点: "注意测试里的容差——归约涉及很多加法,浮点累加误差比 copy 大。这是正常的,只要在容差内就行。"
Benchmark 解读
"reduce_sum 的计算量和内存访问量之比(算术强度)比 copy 高——这意味着它从 memory-bound 向 compute-bound 靠近。优化方向也变了:从追求带宽利用率转向减少 bank conflict。"
[2min] 常见错误演示
错误 1:__syncthreads() 在 if 内(30s)
故意写错并运行 → 死锁/挂起
"这是最容易踩的坑。症状是程序 hang 住不动。记住:__syncthreads() 永远放在条件分支外面。"
错误 2:忘记 shared memory 初始化(30s)
"如果把部分和写进 shared memory 之前没有清零——你存的是上一次 kernel launch 的垃圾数据。用 extern __shared__ 时尤其注意。"
错误 3:blockDim.x 不是 2 的幂(30s)
"树形归约假设 block size 是 2 的幂。如果用 300 个线程——stride 从 150 开始,配对就会错位。建议 block size = 128/256/512。"
错误 4:写了 shared memory 但忘记声明(30s)
"extern __shared__ float smem[] 在 kernel 参数里声明还不够——launch 时 <<<grid, block, shared_mem_size>>> 第三个参数必须传。报错 uses too much shared data 时检查这里。"
课后挑战
"消除 reduce_sum 的 bank conflict:给 shared memory 加 padding,对比优化前后的 bandwidth。目标:提升 30%+ 带宽利用率。"