Prhub

#22786 [AMD][diffusion] Add FlyDSL fused normalization kernels for ROCm diffusion models optimization

原始 PR 作者 yctseng0211 合并时间 2026-06-08 17:42 文件变更 3 提交数 16 评论 12 代码增减 +1099 / -8

执行摘要

为 AMD ROCm 添加 FlyDSL 融合归一化核,优化扩散模型

为AMD ROCm扩散模型优化归一化层性能。现有CUDA/CUTLASS融合核在ROCm上不可用,需要原生ROCm融合核以利用硬件潜力。PR body明确指出"Add AMD FlyDSL fused normalization kernels optimized for ROCm diffusion pipelines"。

建议精读此PR,尤其是fused_residual_norm.py中FlyDSL核的实现和寄存器缓存优化技术,以及layernorm.py中多级fallback的设计模式,值得在跨平台多后端开发中参考。

讨论亮点

Review中主要有以下讨论:

  • Dockerfile依赖(HaiShaw):要求添加rocm.Dockerfile安装FlyDSL。作者通过兼容性修复(适配flydsl>=0.1.5)解决了问题,最终无需修改Dockerfile。
  • 性能优化建议(gemini-code-assist[bot]):建议将batch循环移至GPU网格以减少CPU开销、移除forward_hip中的冗余contiguous调用和调试打印、避免在gate为None时重复分配dummy张量、以及放宽维度对齐断言。这些建议并未全部被采纳,但最终HaiShaw批准了PR。

实现拆解

  1. 核心核实现:在python/sglang/jit_kernel/diffusion/flydsl/fused_residual_norm.py中使用FlyDSL DSL编写了两个核函数:flydsl_fused_residual_norm_scale_shift(融合残差加、门控乘、归一化、scale·shift)和flydsl_norm_scale_shift(融合归一化、scale·shift)。核函数利用寄存器缓存优化,Phase 1计算得到的f32中间值保留在寄存器中,Phase 2直接重用,避免重新读取HBM。
  2. 编译缓存与Fallback:通过_get_or_compile缓存编译结果,通过launch_fused_norm处理参数和降级。当维度不满足5120对齐或flydsl包未导入时,优雅降级到native PyTorch实现。
  3. 集成到现有层:在python/sglang/multimodal_gen/runtime/layers/layernorm.py中,修改RMSNorm和FusedScaleResidualNormScaleShift的forward_hip方法,添加环境变量SGLANG_USE_ROCM_FLYDSL控制开关、导入保护和维度对齐检查。
  4. 单元测试:新增python/sglang/jit_kernel/tests/diffusion/test_flydsl_fused_norm.py,包含两个参数化测试用例,覆盖多种norm_type、batch和seq_len组合,对比FlyDSL输出与PyTorch参考输出,使用宽松容忍度(atol=5e-2, rtol=5e-2)。测试注册为AMD CI。
文件 模块 状态 重要度
python/sglang/jit_kernel/diffusion/flydsl/fused_residual_norm.py 融合核 added 9.08
python/sglang/multimodal_gen/runtime/layers/layernorm.py 扩散层 modified 7.25
python/sglang/jit_kernel/tests/diffusion/test_flydsl_fused_norm.py 测试 added 7.26

关键符号

_build_fused_norm_module flydsl_fused_residual_norm_ss_kernel launch_fused_norm _get_or_compile flydsl_fused_residual_norm_scale_shift flydsl_norm_scale_shift forward_hip (RMSNorm) forward_hip (FusedScaleResidualNormScaleShift)

关键源码片段

python/sglang/jit_kernel/diffusion/flydsl/fused_residual_norm.py core-logic

核心实现,定义了 FlyDSL 融合归一化核和包装函数

"""FlyDSL fused normalization kernels for AMD ROCm (gfx950)."""WARP_SIZE = 64
_VEC = 8
_NUM_WAVES = 10
FLYDSL_NORM_MIN_ALIGNED_DIM = WARP_SIZE * _NUM_WAVES * _VEC # 5120
​
​
def _build_fused_norm_module(D: int, is_rms: bool, has_gate: bool, has_weight: bool):
    # 设置 tiling 参数,要求 D 是 5120 的倍数
    VEC = _VEC
    NUM_WAVES = _NUM_WAVES
    BLOCK = NUM_WAVES * WARP_SIZE # 640
    assert D % FLYDSL_NORM_MIN_ALIGNED_DIM == 0, f"D must be multiple of {FLYDSL_NORM_MIN_ALIGNED_DIM}"
    NUM_ITERS = D // (BLOCK * VEC)
​
    @flyc.kernel(known_block_size=[BLOCK, 1, 1])
    def flydsl_fused_residual_norm_ss_kernel(
        y_ptr: fx.Tensor, res_out_ptr: fx.Tensor, res_ptr: fx.Tensor,
        x_ptr: fx.Tensor, gate_ptr: fx.Tensor, weight_ptr: fx.Tensor,
        bias_ptr: fx.Tensor, scale_ptr: fx.Tensor, shift_ptr: fx.Tensor,
        total_rows: Int32, gate_stride: Int32, scale_stride: Int32, shift_stride: Int32,
    ):
        row = fx.block_idx.x
        tid = fx.thread_idx.x
        # 创建 buffer 资源
        y_rsrc = buffer_ops.create_buffer_resource(y_ptr, max_size=True)
        # ... 省略中间代码 ...
        # Phase 1: 计算 residual + gate * x,累积部分和 / 平方和,保留 f32 值在寄存器
        # Phase 2: 使用寄存器中的 f32 值计算 RMSNorm/LayerNorm 和 scale·shift,
        # 避免重新读取 HBM,节省约 20% 带宽
    return flydsl_fused_residual_norm_ss_kernel
python/sglang/multimodal_gen/runtime/layers/layernorm.py dependency-wiring

修改 forward_hip 以集成 FlyDSL 核,包含环境变量开关、导入保护和维度降级

def forward_hip(self, residual, x, gate, shift, scale):
    # 环境变量开关,默认不使用 FlyDSL
    if not _use_rocm_flydsl:
        return self.forward_native(residual, x, gate, shift, scale)
    # 导入保护:若 flydsl 包未安装则降级
    try:
        from sglang.jit_kernel.diffusion.flydsl.fused_residual_norm import (
            FLYDSL_NORM_MIN_ALIGNED_DIM,
            flydsl_fused_residual_norm_scale_shift,
        )
    except ImportError:
        return self.forward_native(residual, x, gate, shift, scale)
    # 维度对齐检查:仅当 hidden_size 是 5120 的倍数时才使用 FlyDSL
    if x.shape[-1] % FLYDSL_NORM_MIN_ALIGNED_DIM != 0:
        return self.forward_native(residual, x, gate, shift, scale)
    # 调用 FlyDSL 融合核
    return flydsl_fused_residual_norm_scale_shift(
        residual.contiguous(),
        x.contiguous(),
        gate.contiguous() if isinstance(gate, torch.Tensor) else None,
        _ensure_contiguous(self.norm.weight),
        _ensure_contiguous(self.norm.bias),
        scale.contiguous(),
        shift.contiguous(),
        self.norm_type,
        self.eps,
    )

评论区精华

Dockerfile 安装 FlyDSL 依赖 infra

HaiShaw 要求添加 rocm.Dockerfile 安装适当版本 FlyDSL。作者回复兼容性问题已在 0.1.5 版本解决,无需修改 Dockerfile。

结论:无需修改 Dockerfile,兼容性已解决。 · 已解决

batch 循环 CPU 开销 性能

gemini-code-assist[bot] 建议将 kernel launch 的 batch 循环移到 GPU 网格中以减少 CPU 开销。

结论:未明确是否采纳,但 PR 最终被批准。 · unresolved

连续调用和调试打印清理 style

gemini-code-assist[bot] 建议移除 forward_hip 中的冗余 contiguous 调用和调试打印计数。

结论:未明确采纳,最终代码中仍保留 contiguous 调用。 · unresolved

风险与影响

  1. 维度对齐限制:核要求hidden_size是5120的倍数,不满足时会fallback到native实现。若用户模型hidden_size较小或不是5120倍数,无法获得加速。
  2. 集成路径未在CI中测试:PR-CI未设置SGLANG_USE_ROCM_FLYDSL环境变量,从layernorm.py到FlyDSL核的集成分支从未被执行,可能引入未发现的兼容性问题。
  3. 可选依赖flydsl:环境需要安装flydsl包(>=0.1.5),增加了部署复杂度。

用户影响:AMD ROCm用户启用环境变量后,在扩散模型(如Wan2.2 T2V)推理的Denoising阶段可获得约3-10%的加速(总时间约1-3%)。未启用或非AMD平台无影响。系统影响:新增约1k行Python代码,引入可选依赖flydsl。团队影响:为AMD扩散优化建立了基础设施,未来可以类似方式添加更多融合核。

维度对齐限制 集成路径未测试 可选依赖 flydsl

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论