博客

使用 TLX 启用集群启动控制

什么是集群启动控制 (CLC)?

Blackwell 引入了集群启动控制 (CLC) 以实现动态调度。此功能允许内核启动一个包含所需任意数量线程块的网格,这既镜像了非持久化内核使用的方法,又同时保留了较少线程块启动(由持久化内核提供)和负载均衡(由硬件驱动)的双重优势。

我们从一个简单的 GEMM 内核开始,它具有 32×32 的输出块和 144 个可用的 SM。

图 1. 非持久化调度

启用 CLC 后,从主机启动 32×32 网格时,最初会将 CTA 0–143 分配给 SM 0–143。

图 2. CLC 将初始 CTA 分配给 SM

例如,当 CTA 0 仍在 SM 0 上运行时,CLC 允许 SM 0 异步且原子地“窃取”下一个可用的工作(例如 CTA #200),这样 SM 0 无需新的线程块启动即可立即开始处理第 200 号块。

图 3. CLC 工作窃取

动态调度允许系统在执行过程中根据不断变化的工作负载和资源可用性进行调整。例如,如果运行时有 5 个额外的 SM 变得可用,它们也可以窃取并处理可用的工作。

https://docs.nvda.net.cn/cutlass/media/docs/cpp/blackwell-cluster-launch-control.html#blackwell-cluster-launch-control 

什么是 TLX?

TLX 是 Triton DSL 的底层扩展,专为需要对 GPU 操作进行精细控制的专家级用户设计。TLX 提供:

  • 硬件特定的内部函数(如 wgmma、async_copy 和 barrier)
  • 共享内存和本地内存管理
  • 指令级调度和控制
  • 跨线程束组 (warpgroup) 同步

这些特性通过公开底层的 GPU 原语以及内存、计算和异步控制流的显式结构,实现了高级内核开发。虽然 TLX 目前专注于 NVIDIA GPU,但它允许用户实现架构特定的优化,从而减少对编译器启发式算法的依赖。这种方法赋予了用户更多的责任和灵活性,但也可能导致不同硬件平台之间的差异化增加。

https://github.com/facebookexperimental/triton/tree/main 

TLX 中的 CLC

TLX 提供了三个 CLC API

  1. 初始化 tlx.clc_create_context(num_stages, num_consumers) 为 CLC 分配共享内存。
    1. num_stages 启用流水线工作负载窃取。

    2. num_consumers 支持多消费者。

  2. 生产者 tlx.clc_producer(context, k, p_producer) 尝试窃取一个工作负载阶段。
    1. context 是由 clc_create_context 返回的句柄。

    2. k 是阶段索引(0 到 num_stages-1)。

    3. p_producer 是 mbarrier 奇偶校验阶段。

  3. 消费者 tlx.clc_consumer(context, k, p_consumer) 用于 CTA ID 解码(如果成功)。
    1. k 同样是阶段索引。
    2. p_consumer 是消费者的 mbarrier 奇偶校验阶段。

初始化 API tlx.clc_create_context 同时支持多阶段流水线和多消费者工作流。CLC 生产者-消费者设置需要在共享内存中为每个阶段配置一对 mbarrier(mbar_emptymbar_full)以及一个 CLC 响应对象。

生产者 API 将通过等待 mbar_empty 来获取工作,并通过 mbar_full 进行 try_cancel 提交。消费者 API 将等待 mbar_full,从 CLC 响应中解码块 ID,并释放 mbar_empty

# init
clc_context = tlx.clc_create_context(NUM_CLC_STAGES, 1) # only 1 CLC consumer


# init mbar parity phases
clc_phase_producer = 1
clc_phase_consumer = 0
# cicular-buffer pipeline counter
clc_buf = 0


tile_id = start_pid
while tile_id != -1:
clc_buf = clc_buf % NUM_CLC_STAGES
# producer: steal workload
tlx.clc_producer(clc_context, clc_buf, clc_phase_producer)
clc_phase_producer = clc_phase_producer ^ (clc_buf == (NUM_CLC_STAGES - 1))
... # main

# consumer: decode CTA ID
tile_id = tlx.clc_consumer(clc_context, clc_buf, clc_phase_consumer)
clc_phase_consumer = clc_phase_consumer ^ (clc_buf == (NUM_CLC_STAGES - 1))
clc_buf += 1

案例研究

比较 WS GEMMCLC+WS GEMM,两者均使用 3 个 WS 区域(对比详情

  • 默认 WG(尾声消费者):同时调用 tlx.clc_producertlx.clc_consumer

图 4. 在 tlx.async_tasks 外部初始化上下文,并在 ws-region 中调用生产者 API

图 5. 在尾声 ws-region 中调用消费者 API

  • 非默认 WG(MMA 消费者):仅调用 tlx.clc_consumer

图 6. 在 MMA ws-region 中调用消费者 API

  • 非默认 WG(生产者,TMA 加载):仅调用 tlx.clc_consumer

图 7. 在 TMA 加载 ws-region 中调用消费者 API

图 8. 镜像非持久化内核中使用的网格大小

可视化流水线 GEMM 与 CLC GEMM 之间的差异

  • Y 轴:代表 144 个 SM,每个由其 SM ID(从 0 到 143)标识。
  • X 轴:代表时间,以时钟周期为单位,涵盖工作负载的持续时间。
  • 热力图的大部分是黄色的,这意味着在大多数时钟周期内,SM 都被线程块占用。
  • CLC 通过消除流水线 GEMM 中的空闲间隙(紫色)实现了更好的性能。

图 9. 流水线 GEMM 和 CLC GEMM 之间的 SM 占用热力图

  • 由于上述 GEMM 示例中所有线程块处理的工作负载大小相同,CLC 并未提升负载均衡。但对于线程块之间工作负载不均匀的内核,CLC 将像这样极大地增强负载均衡。

图 10. 启用 CLC 后内部内核的 SM 占用热力图

致谢

非常感谢 Bingyi Zhang (NVIDIA) 就 CLC 展开的启发性讨论,以及 Srivatsan Ramesh (Meta) 和 Yuanwei (Kevin) Fang (Meta) 在生成 SM 占用热力图方面提供的工具支持。