Prhub

#34508 [Diffusion][LTX-2] Allocate AdaLN outputs from one contiguous slab

原始 PR 作者 BBuf 合并时间 2026-08-12 16:28 文件变更 2 提交数 1 评论 3 代码增减 +27 / -3

执行摘要

LTX-2 AdaLN 九输出合并为单 slab 分配,提速 1.63x

PR body 明确说明:ltx2_ada_values9 在每次 transformer block 调用时分配 9 个独立的 CUDA tensor,而 Triton kernel 写入的输出形状、dtype 完全相同且互不相交,这些 allocator 往返是不必要的。优化目标是在保持对外九元组 API 不变的前提下,把每次调用的分配次数从 9 次降到 1 次。

值得快速阅读。核心价值不在改动规模,而在“多输出 Triton kernel 用单 slab + unbind 视图减少 allocator 往返”这一通用优化模式,以及用 torch.compile(fullgraph=True) 做回归保护的测试思路。生产路径上如存在同类多输出内核,可直接套用该模式。

讨论亮点

该 PR 没有收到任何 review 评论;Issue 评论区仅有作者 BBuf 的 /tag-and-rerun-ci/rerun-failed-ci 指令和一条 CI 链接,属于流水线操作。PR body 自带的性能与精度数据(1.63x 提速、9 个输出逐位一致、fullgraph 编译通过)是唯一的决策依据,未出现技术争议或未解决疑虑。

实现拆解

  1. 定位分配热点python/sglang/kernels/ops/diffusion/triton/ltx2_ada_values.pyltx2_ada_values9 原实现通过生成器创建 9 个 (batch, seq, hidden) 独立张量,每次 transformer block 调用都产生 9 次 CUDA allocator 往返。
  2. 合并为单一 slab:改为分配一个 (9, batch, seq, hidden) 的连续 output_storage,再 unbind(dim=0) 得到 9 个不相交的连续视图;_ltx2_ada_values9_kernel 的调度与写入逻辑不变,返回的九元组 API 保持不变,对上层模型完全透明。
  3. 测试契约强化:在 test/registered/kernels/ops/diffusion/test_ltx2_ada_values.py 中给 test_ltx2_ada_values9 增加 is_contiguous() 断言;新增 test_ltx2_ada_values9_torch_compile_fullgraph,用 torch.compile(fullgraph=True) 验证整个调用可被单一计算图捕获、无 graph break。
  4. 验证与量化:B300 上以生产形状 B=1, S=1, D=4096、BF16 做 9 轮 × 2000 次调用,wrapper 中位延迟由 24.40us 降至 14.98us(约 1.63x),aten::empty 由 9 次降为 1 次;9 个输出与 eager scale/shift 参考实现对比 atol=0、rtol=0,逐位一致。
文件 模块 状态 重要度
python/sglang/kernels/ops/diffusion/triton/ltx2_ada_values.py 内核算子 modified 4.13
test/registered/kernels/ops/diffusion/test_ltx2_ada_values.py 内核测试 modified 4.92

关键符号

ltx2_ada_values9 test_ltx2_ada_values9 test_ltx2_ada_values9_torch_compile_fullgraph

关键源码片段

python/sglang/kernels/ops/diffusion/triton/ltx2_ada_values.py performance-optimization

核心变更文件:将 `ltx2_ada_values9` 的 9 次独立 `torch.empty` 分配合并为单次连续 slab 分配,是本 PR 性能提升的来源。

# python/sglang/kernels/ops/diffusion/triton/ltx2_ada_values.py
def ltx2_ada_values9(scale_shift_table, timestep):
    """为一个 transformer block 计算 9 组 AdaLN scale/shift 输出,返回 9 个连续视图。"""
    batch, seq, _ = timestep.shape
    hidden = scale_shift_table.shape[1]
    rows = int(batch * seq)
​
    # 旧实现为每个输出单独调用 torch.empty,一次函数调用产生 9 次
    # CUDA allocator 往返;新实现把 9 个输出合并进一个连续 slab,
    # 再用 unbind(0) 切出不相交的视图,分配次数降为 1 次。
    output_storage = torch.empty(
        (9, batch, seq, hidden),
        device=timestep.device,
        dtype=timestep.dtype,
    )
    outs = tuple(output_storage.unbind(dim=0))
​
    # Triton kernel 仍按行写入这 9 个视图;kernel 参数与调度保持不变,
    # 仅输出存储从 9 个独立张量换成 slab 视图,对上层 API 完全透明。
    _ltx2_ada_values9_kernel[(rows,)](
        timestep,
        scale_shift_table,
        # ... 其余参数与原实现一致
    )
    return outs
test/registered/kernels/ops/diffusion/test_ltx2_ada_values.py test-coverage

配套测试:为原测试补充连续性断言,并新增 `torch.compile(fullgraph=True)` 回归,确保 slab 视图路径可被完整编译。

# test/registered/kernels/ops/diffusion/test_ltx2_ada_values.py
@torch.no_grad()
def test_ltx2_ada_values9_torch_compile_fullgraph() -> None:
    hidden = 4096
    scale_shift_table = torch.randn(
        9, hidden, device=DEVICE, dtype=torch.bfloat16
    ).contiguous()
    timestep = torch.randn(
        1, 1, 9 * hidden, device=DEVICE, dtype=torch.bfloat16
    ).contiguous()
​
    # fullgraph=True 强制整个函数被捕获进单一计算图,
    # 任何隐式分支、额外分配或不可编译操作都会触发 graph break;
    # 这能确保 slab + unbind 的返回路径可被 torch.compile 完整编译。
    actual = torch.compile(ltx2_ada_values9, fullgraph=True)(
        scale_shift_table, timestep
    )
    expected = _reference(scale_shift_table, timestep)
​
    assert len(actual) == 9
    for actual_value, expected_value in zip(actual, expected):
        # 每个返回视图必须保持连续,避免上层按视图访问时发生
        # 非连续内存拷贝或性能回退
        assert actual_value.is_contiguous()
        torch.testing.assert_close(actual_value, expected_value, atol=0, rtol=0)

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

  • 内存布局语义变化:9 个输出从独立张量变为同一 slab 的视图,若 Triton kernel 对第 0 维 stride 的偏移计算有误,越界写会直接污染相邻视图;原实现中越界只影响单个张量。当前形状、dtype 一致且偏移简单,风险较低,但属于隐蔽的语义变化。
  • 视图连续性依赖unbind(dim=0) 返回的视图连续与否取决于 slab 与 hidden 维的内存排布,测试中已有 is_contiguous() 断言兜底;若未来支持非连续布局需显式处理。
  • 测试环境局限:性能数据与 fullgraph 编译验证均在 B300 上完成,且 Extra CI 曾显示失败,普通 CI 环境未必覆盖该内核,回归保护强度依赖 B300 本地验证。

影响范围仅限 LTX-2 diffusion 模型推理路径:每个 transformer block 的内存分配开销显著下降,B300 场景 wrapper 提速约 1.63x,aten::empty 从 9 次降为 1 次,对批量、序列更大的形状同样受益;数值逐位不变,无精度影响。改动集中在 1 个内核文件与 1 个测试文件,对外 API 无变化,对团队协作影响很小,且新增的 fullgraph 编译回归为后续 torch.compile 适配提供了保护。

内存布局语义变更 依赖视图连续性 回归依赖 B300 验证

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论