H100 · Stream-K 调度

目录 · ← l8 · l10 →

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 概述)。


目录

  1. 第一性原理:fixup 从何而来
  2. 工作单元(Unit of Work)
  3. Split(一次拆分的贡献)
  4. 跟踪三元组:K_idx / k_tile_count / is_final_split
  5. 三种角色
  6. 锁(The Lock)
  7. Groups(分组)与 L2 局部性
  8. HyTiS:另一种解决波量化的思路
  9. 总结与学习衔接

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 上的完整 tilek_tile_count == k_tiles_per_output_tile不需要
K 上的部分 tilek_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
  • 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)。

5. 三种角色

给定一个输出 tile 可能有 2–4 个 CTA 各算 K 的一段,角色直接由三元组决定:

条件角色行为
K_idx == 0first 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-KHyTiS
分解维度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 页)。