Prhub

#33617 [NVIDIA] Enable CuTe DSL BF16 GEMM on SM107

原始 PR 作者 YAMY1234 合并时间 2026-08-06 17:06 文件变更 4 提交数 3 评论 9 代码增减 +20 / -19

执行摘要

SM107 启用 CuTe DSL BF16 GEMM,修复 bias DLPack 导出

PR body 明确说明:CuTe DSL BF16 GEMM 后端在 SM10x GPU 上会被自动选中,但两个运行时入口 _tgv_bf16_gemm_run / _tgv_bf16_gemm_out_run 的显式 SM 号检查(只允许 100/103)会拒绝 SM107,导致启用后实际回退到 cuBLAS;同时打开架构 gate 还暴露了 bias 转换问题:DLPack 拒绝导出保留 requires_grad=True 的 bias view,即使在 torch.no_grad() 推理下也会报错。评论中 leejnau 实测确认 'I tested Kimi K3 on sm107 and the cutedsl bf16 gemm works ✅',说明该后端在 SM107 上对 Kimi-K3 等模型有实际收益。

值得精读。改动虽小(4 文件、+20/-19),但包含两个可复用经验:一是架构 gate 采用 capability 判断而非 SM 号枚举的取舍,二是 DLPack 与 autograd 元数据(requires_grad)冲突的规避方式——detach() 无拷贝无同步地解决推理路径导出问题。阅读时建议关注 review 中 mmangkad 关于 CUDA 版本下限的提醒,以及后续是否拆分 is_sm107_supported() 的演进。

讨论亮点

核心讨论围绕 gate 函数语义展开:

  • b8zhong:"Can we just change it to is_sm100_supported?",主张复用既有能力判断函数,随后说明自己另一部分评论是误输入。
  • leejnau:"maybe is_sm10x_supported? what is Supports 110 exactly?",对函数命名与覆盖范围提出疑问。
  • mmangkad:指出 sm107 needs >= 13.4 while is_sm100_supported currently checks >= 12.8 (actually, sm103 needs >= 12.9),建议重命名为 is_sm100_or_sm103_supported 并新增 is_sm107_supported,以便区分不同 CUDA 版本下限。
  • YAMY1234 回应:"I updated the PR to use is_sm100_supported()... Lee and I talked about it, and we don't think we need a separate is_sm107_supported() for now. If we add features that are only SM107-specific later, we can add it then as a follow up"。
  • mmangkad 接受该方案但预判:"at some point we'll probably have to distinguish them";b8zhong 与 leejnau 则认为无需额外加 CUDA 版本检查,现有检查已正确。

实现拆解

整体按三步推进,涉及 4 个文件,核心是 gate 语义统一与 DLPack 导出修复,GEMM kernel 数学与 tactic 选择完全未动。

  1. 架构 gate 统一为 capability 判断python/sglang/kernels/ops/gemm/cutedsl_bf16_gemm.py 中两个运行时入口 _tgv_bf16_gemm_run_tgv_bf16_gemm_out_run 的显式 get_device_sm() not in (100, 103) 判断改为 is_sm100_supported(),模块 docstring 从 "SM100/SM103 only" 更新为 "SM10x"。同时 python/sglang/srt/layers/quantization/unquant.pyinitialize_bf16_gemm_config 报错文案、python/sglang/srt/server_args.pybf16_gemm_backend CLI help 同步改为 "SM10x GPU"。原因:SM107 与 SM100/103 同属 major 10 能力族,逐点枚举 SM 号会导致新架构漏配,统一 gate 后未来新增 SM10x 架构无需再改。

  2. 修复 bias 的 DLPack 导出cutedsl_bf16_gemm.py_to_cute_swap 中,bias 通过 as_strided 生成广播 3D view 后,from_dlpack(bias_3d, ...) 改为 from_dlpack(bias_3d.detach(), ...)。原因:推理路径中 bias 常继承模型权重的 requires_grad=True,DLPack 协议拒绝导出仍挂 autograd 元数据的 view;detach() 只剥离元数据、不产生拷贝与同步,因此对读路径语义无影响。

  3. 测试与文档配套test/registered/kernels/ops/gemm/test_cutedsl_bf16_gemm.py 将形状参数从 N/K 两组笛卡尔积重构为 SHAPES 列表并新增 (2048, 4096) 以覆盖 SM107 目标工作负载;bias 在非 None 时设置 requires_grad_(True) 并在 torch.no_grad() 下调用 kernel,参考计算改用 bias.detach();skip 文案统一为 "SM10x required"。测试注册的 CI stage(base-b-kernel-unit、4-gpu-b200 runner)与精度比对容差(rtol=2e-2, atol=2.5)保持不变。

文件 模块 状态 重要度
python/sglang/kernels/ops/gemm/cutedsl_bf16_gemm.py GEMM 内核 modified 4.89
test/registered/kernels/ops/gemm/test_cutedsl_bf16_gemm.py 内核测试 modified 5.01
python/sglang/srt/layers/quantization/unquant.py BF16 调度 modified 5.03
python/sglang/srt/server_args.py 参数配置 modified 4.09

关键符号

_tgv_bf16_gemm_run _tgv_bf16_gemm_out_run _to_cute_swap initialize_bf16_gemm_config test_cutedsl_bf16_gemm test_cutedsl_bf16_gemm_empty_batch

关键源码片段

python/sglang/kernels/ops/gemm/cutedsl_bf16_gemm.py core-logic

核心变更文件:两个运行时入口的架构 gate 从显式 SM 号枚举改为 is_sm100_supported(),并修复 _to_cute_swap 中 bias 的 DLPack 导出(detach)。

# cutedsl_bf16_gemm.py 关键 gate 与 bias 处理片段(整理版,省略未变更部分)# 架构能力 gate:SM107 与 SM100/103 同属 major=10 能力族,
# 统一走 is_sm100_supported() 而非逐一枚举 SM 号,未来新增 SM10x 架构自动覆盖。
def _tgv_bf16_gemm_run(
    x: torch.Tensor, weight: torch.Tensor, bias: Optional[torch.Tensor]
) -> torch.Tensor:
    # SM10x 能力门控,kernel 实现与 tactic 选择保持不变
    if not is_sm100_supported():
        raise RuntimeError("cutedsl_bf16_gemm requires an SM10x GPU")
​
    assert x.dtype == torch.bfloat16 and weight.dtype == torch.bfloat16
    assert x.stride(-1) == 1, "x must be K-major [M, K]"
    assert weight.stride(-1) == 1, "weight must be K-major [N, K]"
    # ... 后续 TGV 启动逻辑与 tactic 选择未变更 ...
​
​
def _tgv_bf16_gemm_out_run(
    out: torch.Tensor,
    x: torch.Tensor,
    weight: torch.Tensor,
    bias: Optional[torch.Tensor],
) -> None:
    # 输出形态与设备相关约束校验
    if not is_sm100_supported():
        raise RuntimeError("cutedsl_bf16_gemm requires an SM10x GPU")
​
    assert x.dtype == torch.bfloat16 and weight.dtype == torch.bfloat16
    assert out.dtype == torch.bfloat16 and out.device == x.device
    assert x.ndim == 2 and weight.ndim == 2 and out.ndim == 2
    # ... 后续 kernel launch 逻辑未变更 ...
​
​
# _to_cute_swap 内部核心片段:将 PyTorch 布局的 GEMM 输入重排为 CuTe 期望布局,
# 其中 bias 映射为沿 L 维广播的 3D 视图,M/N 维度对调以匹配 c_swap。
def _to_cute_swap(...):
    M_ce = c_swap.shape[1] # == PyTorch N
    N_ce = c_swap.shape[2] # == PyTorch M
    bias_3d = bias_pt.as_strided(size=(L, M_ce, N_ce), stride=(0, 1, 0))
    # 关键修复:DLPack 拒绝导出仍带 requires_grad=True 的 Tensor view,
    # 即使推理处于 torch.no_grad();detach 只剥离 autograd 元数据,
    # 不引入拷贝或同步,读路径语义不变。
    bias_ = from_dlpack(bias_3d.detach(), assumed_align=2).mark_layout_dynamic(leading_dim=1)
    return a_, b_, c_, bias_, layout
test/registered/kernels/ops/gemm/test_cutedsl_bf16_gemm.py test-coverage

测试配套:形状参数重构并新增 (2048, 4096) 覆盖 SM107 工作负载,bias 增加 requires_grad=True 场景以复现 DLPack 导出问题。

# 形状参数化:在原有 N/K 组合基础上新增 (2048, 4096),
# 用于覆盖 SM107 目标工作负载形态。
SHAPES = [(n, k) for n in [1024, 2624, 6144] for k in [2048, 6144]] + [(2048, 4096)]
NUM_TOKENS = get_ci_test_range(list(range(1, 33)), [1, 15, 16, 32])
​
​
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
@pytest.mark.parametrize("has_bias", [False, True])
@pytest.mark.parametrize("n,k", SHAPES)
@pytest.mark.parametrize("num_tokens", NUM_TOKENS)
def test_cutedsl_bf16_gemm(num_tokens, k, n, has_bias):
    # SM10x 系列 arch major 均为 10,测试在 SM100/103/107 上都会实际执行
    if is_hip_runtime() or get_jit_cuda_arch().major != 10:
        pytest.skip("SM10x required")
​
    torch.manual_seed(num_tokens)
    x = torch.randn(num_tokens, k, dtype=torch.bfloat16, device="cuda")
    weight = torch.randn(n, k, dtype=torch.bfloat16, device="cuda")
    bias = torch.randn(n, dtype=torch.bfloat16, device="cuda") if has_bias else None
    if bias is not None:
        # 复现真实推理路径:权重 / 偏置常带 requires_grad=True,
        # 但推理在 torch.no_grad() 下执行;bias 需经 detach 后走 DLPack
        bias.requires_grad_(True)
​
    with torch.no_grad():
        out = cutedsl_bf16_gemm(x, weight, bias)
    assert out.shape == (num_tokens, n)
    assert out.dtype == torch.bfloat16
​
    # 参考计算使用 float 精度,bias 以 detach 后的值参与比对
    ref = x.float() @ weight.float().T
    if bias is not None:
        ref = ref + bias.detach().float()
    torch.testing.assert_close(out, ref.bfloat16(), rtol=2e-2, atol=2.5)

评论区精华

架构 gate 统一为 is_sm100_supported() 的取舍 设计

b8zhong 建议直接用 is_sm100_supported() 替代显式 SM 列表;leejnau 质疑函数名与 SM110 覆盖范围,建议更名 is_sm10x_supported;mmangkad 补充 SM107 需要 CUDA >= 13.4 而 is_sm100_supported 只查 >= 12.8(SM103 需 12.9),建议拆出 is_sm107_supported。

结论:采用 is_sm100_supported();YAMY1234 与 leejnau 商定暂不新增 is_sm107_supported(),待出现 SM107 专属特性再作为 follow-up 引入。 · 已解决

SM107 的 CUDA 版本下限是否单独校验 question

mmangkad 指出 SM107 运行需要 CUDA >= 13.4,现行 gate 不区分 CUDA 版本;b8zhong 认为已有检查应当正确、无需再查,leejnau 表示认同。

结论:暂不加 CUDA 版本检查;低 CUDA 版本 + SM107 组合行为未被 CI 覆盖,留待 SM107 专属能力落地时处理。 · unresolved

bias 在 DLPack 导出前的 detach 修复 正确性

PR body 说明启用 SM107 后暴露:DLPack 拒绝 requires_grad=True 的 bias view,即使推理在 torch.no_grad() 下;修复采用 detach() 且不产生拷贝或同步。评论线程未再展开,作者通过后续 commit 直接回应。

结论:在 from_dlpack 前调用 bias_3d.detach(),配套测试新增 requires_grad bias 场景,语义不受影响。 · 已解决

风险与影响

  1. CUDA 版本与 SM107 需求可能错配:mmangkad 指出 SM107 需要 CUDA >= 13.4,而 is_sm100_supported() 只按 SM 架构与 CUDA >= 12.8(SM103 需 12.9)判断。合并后未新增 CUDA 版本分层,低版本 CUDA + SM107 的组合行为没有 CI 覆盖,属于待跟进风险。
  2. 默认后端切换影响面:SM107 上 --bf16-gemm-backend auto 将从 cuBLAS 切到 CuTe DSL TGV kernel。kernel 本身未改,但 PR 标注上游 main 的 GPU CI 与性能确认仍在 pending(draft 状态),数值与性能结论依赖受控验证与后续 CI。
  3. bias detach 语义detach() 不复制、不同步,读取路径语义不变,风险低;但若未来引入训练/梯度回传路径,需重新审视该视图的 autograd 行为。
  4. 测试覆盖条件:测试 skip 条件仍为 get_jit_cuda_arch().major != 10,SM107 major 为 10 可正常执行,无漏测风险;但 CI runner 为 4-gpu-b200(SM100),SM107 实机路径未被 CI 直接覆盖。

对用户:SM107 GPU 用户启动时 auto 模式将默认启用 CuTe DSL BF16 GEMM,BF16 线性层不再回退到 F.linear,leejnau 已在 Kimi-K3 上实测验证可正常工作,预期带来低延迟收益。对系统:架构 gate 统一到 is_sm100_supported() 后,后续新增 major 10 架构无需修改 kernel 入口;CLI 帮助与报错文案同步收敛为 "SM10x"。对团队:该 PR 确立了一个可复用的 gate 策略——能力族判断优先于 SM 号枚举,同时留下一项 follow-up(SM107 专属特性出现后再拆分 CUDA 版本检查)。

CUDA 版本下限与 SM107 需求可能错配 上游 CI 与性能验证待完成 核心 GEMM 路径默认后端变更

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论