Intro-ops/course/scripts/03_reduce_sum_script.md

4.6 KiB
Raw Permalink Blame History

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

讲解要点:

  1. 目标:每行 N 个元素求和为 1 个值
  2. Step 1: 各线程先独立累加自己负责的列 → 存入 shared memory
  3. Step 2: 树形归约——stride = blockDim.x/2, /4, ..., 1,每次配对相加
  4. 为什么叫"树形"——每步参与线程减半log₂(blockDim.x) 步完成
  5. 关键:每一步之后必须 __syncthreads()

可视化辅助: 手动画出 8→4→2→1 的归约树

过渡语: "听起来简单,但写起来有两个坑——我们直接在 IDE 里看。"


[10min] 代码实操

CUDA Kernel6min

打开文件: 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 Kernel4min

打开文件: 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__ 时尤其注意。"

错误 3blockDim.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%+ 带宽利用率。"