H100 · Stream-K 调度
H100 · Stream-K 调度
本文档基于课程讲义《8.1 Stream-K.pdf》(Lesson 8.1,10 页)整理, 深入讲解 Stream-K 调度器的机制:fixup、split、三重角色、锁、分组与 L2 局部性, 以及与 HyTiS 的对比。前置阅读:
H100-Kernel-Design.md(§10.3 Stream-K 概述)。
目录
- 第一性原理:fixup 从何而来
- 工作单元(Unit of Work)
- Split(一次拆分的贡献)
- 跟踪三元组:K_idx / k_tile_count / is_final_split
- 三种角色
- 锁(The Lock)
- Groups(分组)与 L2 局部性
- HyTiS:另一种解决波量化的思路
- 总结与学习衔接
1. 第一性原理:fixup 从何而来
对单个输出 tile C_tile,GEMM 计算:
C_tile = Σ_{k-tiles} (A_tile_k · B_tile_k)
- 一个 CTA 算完该输出 tile 的所有 K-tile → 无需跨 CTA 归约。
- 多个 CTA 各自只算 K-tile 的一个子集 → 每个 CTA 产出一个部分和,这些部分和必须合并。
在 Stream-K 调度里,这个合并步骤就叫 fixup。
2. 工作单元(Unit of Work)
每个工作单元包含:
- tile 坐标:
(m_idx, n_idx, l_idx)(l_idx用于 batch/分组问题)。 k_idx:该工作单元在输出 tile 内起始的 k-tile。k_tile_count:该工作单元为这个输出 tile 计算了多少个 k-tile。
是否需要归约由下面判断(这是 requires_fixup(...) 背后的关键谓词):
| 情况 | 判定 | 归约 |
|---|---|---|
| K 上的完整 tile | k_tile_count == k_tiles_per_output_tile | 不需要 |
| K 上的部分 tile | k_tile_count != k_tiles_per_output_tile | 需要 |
3. Split(一次拆分的贡献)
一个 CTA 被分配的 k-tile 迭代区间,可能正好落在一个输出 tile 的中间。
- 例:3 个输出 tile(每个 90 个 k-tile),4 个 CTA 单元。
- “split” = 一个 CTA 对单个输出 tile 的贡献。
- 上面的 Unit 1 有两个 split:一个是 tile 0 的尾部,一个是 tile 1 的头部。
- 代码一次处理一个 split(这就是
advance_to_next_work里的k_tile_remaining循环)。
4. 跟踪三元组:K_idx / k_tile_count / is_final_split
对单个 split(一个 CTA 在某个输出 tile 上的工作):
K_idx:该 split 在输出 tile 的 K 维内从哪里开始。K_idx = tile_iter_start - output_tile_iter_start- Unit 0 对 tile 0:
K_idx = 0; - Unit 1 对 tile 0:
K_idx = 67。
- Unit 0 对 tile 0:
k_tile_count:该 split 处理多少个 k-tile。- Unit 0 对 tile 0:67;
- Unit 1 对 tile 0:23(
= 90 − 67)。
is_final_split():(K_idx + k_tile_count) == k_tiles_per_output_tile。- 为真 → 该 split 覆盖到 K 维的末尾。Unit 1 对 tile 0 就是 final split(
67 + 23 = 90)。
- 为真 → 该 split 覆盖到 K 维的末尾。Unit 1 对 tile 0 就是 final split(
5. 三种角色
给定一个输出 tile 可能有 2–4 个 CTA 各算 K 的一段,角色直接由三元组决定:
| 条件 | 角色 | 行为 |
|---|---|---|
K_idx == 0 | first split(首个) | 算了 [0, N),在你之前没人写 workspace → 直接 store |
is_final_split() == true 且非”分离归约” | final split(末尾)+ epilogue 拥有者 | 算了 [X, 90),等前面所有人、载入他们的累加结果、加上自己的、跑 epilogue |
| 其它 | middle split(中间) | 算了 [A, B)(0 < A < B < 90),把自己的部分和归约进 workspace 已有的值 |
6. 锁(The Lock)
- 物理上:global memory 里一个连续的 int 数组,每个输出 tile 一个(多 warpgroup kernel 再乘
num_barriers)。 - 该数组的指针就放在归约数据缓冲区之后的同一块分配里;kernel 启动时每个锁都从 0 开始。
- 锁是一个整数,编码”这个输出 tile 已完成并写进 workspace 的 K 维工作量”。普通(非分离归约)模式下, 它计数累计已处理的 k-tile 数,且只增不减。
- 锁在 K-tile 空间里编码进度:
- 确定性模式(deterministic):每个 split 等待恰好等于它起始位置的累计 k-tile 数, 强制严格的从左到右归约顺序。
- 非确定性模式(non-deterministic):middle split 只需知道 workspace 已被初始化(
lock >= 1), 然后竞争式地原子归约进去。
7. Groups(分组)与 L2 局部性
Groups 是 L2 局部性优化:把 stream-K 单元划分成 G 个独立子组,每组只协作处理自己那部分 stream-K tile。
- 这是 stream-K 特有的优化:因为 stream-K 破坏了”波式光栅化”——一个 CTA 可能横跨输出网格不同区域的 tile,破坏局部性。
- 无分组(G=1):所有 stream-K 单元共享一个跨所有输出 tile 的大 K-tile 池。Unit 0 可能处理 tile 0 和 tile 1, 而 Unit 7 处理 tile 5 和 tile 6——空间位置完全不同,L2 里的数据毫不重叠。
- 有分组:每组里的 unit 会计算”与数据并行公式中、按光栅化顺序属于同一波的那些 tile”的相同 K 区间。
7.1 分组层级
Groups(最多 8 个,为了 L2 局部性)
└── 每个 group 含多个 cluster-tile
└── 每个 cluster 含多个 CTA(线程块)
└── 每个 CTA 处理 K-tile
- 分组沿光栅化维度确定。例:沿 M 光栅化、
problem_blocks_m / cluster_m = 4→ 得 4 个 group。 - group 在输出空间里交错,最终 tile id 计算:
output_tile_id = (output_tile_id_in_group * num_groups) + group_idx
8. HyTiS:另一种解决波量化的思路
- HyTiS 通过让部分波(partial wave)用更细粒度的 tile 来解决波量化,让更多 SM 保持忙碌。
- 它是纯空间分解(M×N)+ 异构 tile 大小,对比 Stream-K 的 K 维分解 + 同构 tile 大小。
- HyTiS 在一次 kernel 启动里用两种 tile 大小:
- 大 tile(如 128×256)给完整波 → 最大吞吐;
- 小 tile(如 64×64)给部分波 → 最小延迟。
- 无归约、无 workspace、无 barrier、无 fixup。
- 代价:当问题在 M、N 上很小而 K 很大时,HyTiS 无能为力(因为它只能分解 M×N 空间)。
8.1 Stream-K vs HyTiS
| Stream-K | HyTiS | |
|---|---|---|
| 分解维度 | K(同构 tile) | M×N 空间(异构 tile) |
| 归约/workspace/fixup | 有(fixup + 锁) | 无 |
| 适用 | 各种形状,含”小 M/N 大 K” | M、N 足够大、可分空间 |
| 权衡 | 引入跨 CTA 归约开销 | 小 M/N、大 K 时帮不上 |
9. 总结与学习衔接
9.1 核心脉络速记
| 概念 | 一句话 |
|---|---|
| fixup | 多 CTA 各算一段 K 后,把部分和合并的步骤 |
| split | 一个 CTA 对单个输出 tile 的贡献 |
| 三元组 | K_idx(起始)、k_tile_count(数量)、is_final_split()(是否覆盖 K 末尾) |
| 三角色 | first(直接 store)/ final(epilogue 拥有者)/ middle(归约进 workspace) |
| 锁 | global 里每 tile 一个 int,编码累计 K 进度;确定/非确定两种归约模式 |
| Groups | 把单元分组以恢复 L2 局部性(最多 8 组) |
| HyTiS | 用异构 tile 大小(大 tile 整波 + 小 tile 尾波)免归约地解决波量化 |
9.2 与本课程其他内容的衔接
| 本文概念 | 对应后续专题 |
|---|---|
| Stream-K 的 kernel 启动与协同 | 《8.2 Kernel Launch》 |
| 跨 CTA/跨 SM 归约与多 GPU | 《9. Multi GPU》《10. Multi GPU Part 2》 |
9.3 一句话记忆
Stream-K = 把”输出 tile”沿 K 维切成工作单元,让所有 SM 平分数学量、一起同时干完; 代价是跨 CTA 的 fixup 归约,用”锁”(累计 K 进度)协调首/中/尾三种角色, 用”Groups”找回被 1D 带破坏的 L2 局部性;HyTiS 则走另一条路——用异构 tile 免归约地打散尾波。
参考来源:
8.1 Stream-K.pdf(Lesson 8.1,10 页)。
