From 5993f91f843c3429b0e3bdf21bf05a9f1236f731 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Tue, 1 Sep 2026 22:51:10 +0800 Subject: [PATCH] [Kernel] Register merged diffusion agent kernels with KDA backend (#37385) --- python/sglang/kernels/fused_op.py | 5 +++ .../sglang/kernels/ops/diffusion/__init__.py | 43 ++++++++++++++----- .../diffusion/norm/norm_scale_shift_jit.py | 27 ++++++++++++ python/sglang/kernels/spec.py | 1 + .../kernels/ops/layernorm/test_fused_op.py | 10 ++++- .../ops/layernorm/test_kernels_namespace.py | 37 ++++++++++++++++ 6 files changed, 111 insertions(+), 12 deletions(-) diff --git a/python/sglang/kernels/fused_op.py b/python/sglang/kernels/fused_op.py index 44708945c..fb1d15bb5 100644 --- a/python/sglang/kernels/fused_op.py +++ b/python/sglang/kernels/fused_op.py @@ -108,6 +108,7 @@ BACKEND_METHODS: Dict[KernelBackend, str] = { KernelBackend.AOT: "forward_aot", KernelBackend.CUTE_DSL: "forward_cute_dsl", KernelBackend.FLYDSL: "forward_flydsl", + KernelBackend.KDA: "forward_kda", KernelBackend.FLASHINFER: "forward_flashinfer", KernelBackend.DEEPGEMM: "forward_deepgemm", KernelBackend.AITER: "forward_aiter", @@ -122,6 +123,7 @@ _METHOD_BACKEND_LABELS: Dict[str, str] = { # must never trigger a surprise compilation in a serving process; force it # explicitly when wanted. Per-op priority overrides this (see BaseFusedOp). DEFAULT_PRIORITY: Tuple[KernelBackend, ...] = ( + KernelBackend.KDA, KernelBackend.AOT, KernelBackend.JIT, KernelBackend.FLASHINFER, @@ -434,6 +436,9 @@ class BaseFusedOp(nn.Module, ABC): def forward_cute_dsl(self, *args, **kwargs): raise NotImplementedError(f"{self._op_label()}: no cute_dsl backend") + def forward_kda(self, *args, **kwargs): + raise NotImplementedError(f"{self._op_label()}: no KDA backend") + def forward_flashinfer(self, *args, **kwargs): raise NotImplementedError(f"{self._op_label()}: no flashinfer backend") diff --git a/python/sglang/kernels/ops/diffusion/__init__.py b/python/sglang/kernels/ops/diffusion/__init__.py index dd20a1522..16d9ec32f 100644 --- a/python/sglang/kernels/ops/diffusion/__init__.py +++ b/python/sglang/kernels/ops/diffusion/__init__.py @@ -82,6 +82,13 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = ( _CUDA, "Bit-exact RMSNorm + adaLN scale/shift.", ), + ( + "diffusion.scale_residual_norm_scale_shift", + KernelBackend.KDA, + "norm.norm_scale_shift_jit:kda_scale_residual_norm_scale_shift", + _CUDA_SM100_PLUS, + "KDA B200 native CUDA residual + LayerNorm + scale/shift (#27392).", + ), ( "diffusion.scale_residual_norm_scale_shift", KernelBackend.TRITON, @@ -110,6 +117,13 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = ( _CUDA, "Qwen residual LayerNorm/modulation + NVFP4 quantization.", ), + ( + "diffusion.norm_scale_shift", + KernelBackend.KDA, + "norm.norm_scale_shift_jit:kda_norm_scale_shift", + _CUDA_SM100_PLUS, + "KDA B200 native CUDA LayerNorm + scale/shift (#27392).", + ), ( "diffusion.norm_scale_shift", KernelBackend.CUTE_DSL, @@ -168,10 +182,10 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = ( ), ( "diffusion.residual_gate_add", - KernelBackend.JIT, + KernelBackend.KDA, "modulate.residual_gate_add_jit:residual_gate_add", _CUDA, - "Fused residual + gate * update.", + "KDA native CUDA residual + gate * update (#29361).", ), ( "diffusion.timestep_embedding", @@ -201,19 +215,26 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = ( _CUDA, "Fused in-place QK RMS-norm + RoPE.", ), + ( + "diffusion.flux2_layernorm_modulate_fp8_quant", + KernelBackend.KDA, + "norm.layernorm_modulate_triton:fused_layernorm_modulate_fp8_quant_raw", + _CUDA, + "KDA-generated FLUX.2 LayerNorm + adaLN modulation + static FP8 quantization.", + ), ( "diffusion.flux2_qkv_epilogue", - KernelBackend.JIT, + KernelBackend.KDA, "rope.flux2_qkv_epilogue_jit:try_fused_flux2_qkv_epilogue", _CUDA, - "FLUX.2 QK RMS-norm + RoPE + joint QKV packing.", + "KDA-generated FLUX.2 QK RMS-norm + RoPE + joint QKV packing.", ), ( "diffusion.flux2_token_cat_fp8", - KernelBackend.TRITON, + KernelBackend.KDA, "layout.flux2_token_cat_fp8_triton:try_flux2_token_cat_fp8", _CUDA, - "FLUX.2 single-block token concatenation + static FP8 quantization.", + "KDA-generated FLUX.2 token concatenation + static FP8 quantization.", ), ( "diffusion.qwen_qkv_epilogue", @@ -224,10 +245,10 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = ( ), ( "diffusion.ltx2_qknorm_split_rope", - KernelBackend.JIT, + KernelBackend.KDA, "rope.ltx2_qknorm_split_rope_jit:ltx2_qknorm_split_rope_cuda", - _CUDA, - "LTX-2 QK-norm + split RoPE.", + _CUDA_SM100_PLUS, + "KDA native CUDA LTX-2 QK-norm + split RoPE (#29708).", ), ( "diffusion.ltx25_decoder_rope", @@ -343,10 +364,10 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = ( ), ( "diffusion.causal_conv3d_cat_pad", - KernelBackend.JIT, + KernelBackend.KDA, "layout.causal_conv3d_cat_pad_jit:fused_causal_conv3d_cat_pad_cuda", _CUDA, - "Causal Conv3d cat + pad.", + "KDA native CUDA causal Conv3d cat + pad (#29281).", ), ( "diffusion.causal_conv3d_cat_pad", diff --git a/python/sglang/kernels/ops/diffusion/norm/norm_scale_shift_jit.py b/python/sglang/kernels/ops/diffusion/norm/norm_scale_shift_jit.py index 94b08a387..0726cc0a3 100644 --- a/python/sglang/kernels/ops/diffusion/norm/norm_scale_shift_jit.py +++ b/python/sglang/kernels/ops/diffusion/norm/norm_scale_shift_jit.py @@ -222,6 +222,33 @@ def try_fused_scale_residual_norm_scale_shift( return y, residual_out +def kda_norm_scale_shift(x, weight, bias, scale, shift, norm_type, eps): + """Run the KDA B200 native CUDA path introduced by PR #27392. + + Unlike ``try_fused_norm_scale_shift``, this explicit backend entry point + fails on unsupported inputs instead of silently returning ``None`` for a + caller-owned fallback. + """ + out = try_fused_norm_scale_shift(x, weight, bias, scale, shift, norm_type, eps) + if out is None: + raise RuntimeError("unsupported input for KDA norm-scale-shift CUDA") + return out + + +def kda_scale_residual_norm_scale_shift( + residual, x, gate, weight, bias, scale, shift, norm_type, eps +): + """Run the KDA B200 residual + norm + scale/shift path from PR #27392.""" + out = try_fused_scale_residual_norm_scale_shift( + residual, x, gate, weight, bias, scale, shift, norm_type, eps + ) + if out is None: + raise RuntimeError( + "unsupported input for KDA scale-residual-norm-scale-shift CUDA" + ) + return out + + def try_fused_norm_scale_shift_fp8( x, weight, bias, scale, shift, input_scale, norm_type, eps ): diff --git a/python/sglang/kernels/spec.py b/python/sglang/kernels/spec.py index d1220a35e..913534a26 100644 --- a/python/sglang/kernels/spec.py +++ b/python/sglang/kernels/spec.py @@ -44,6 +44,7 @@ class KernelBackend(str, Enum): AOT = "aot" # sgl_kernel wheel (CUDA / ROCm builds) CUTE_DSL = "cute_dsl" FLYDSL = "flydsl" # FlyDSL MLIR compiler (device=HIP, gfx950) + KDA = "KDA" # Kernel Design Agents generated implementation FLASHINFER = "flashinfer" DEEPGEMM = "deepgemm" AITER = "aiter" # AMD aiter library (device=HIP) diff --git a/test/registered/kernels/ops/layernorm/test_fused_op.py b/test/registered/kernels/ops/layernorm/test_fused_op.py index b23485d06..849d86776 100644 --- a/test/registered/kernels/ops/layernorm/test_fused_op.py +++ b/test/registered/kernels/ops/layernorm/test_fused_op.py @@ -10,7 +10,7 @@ import pytest import torch import sglang.kernels as K -from sglang.kernels.fused_op import BaseFusedOp +from sglang.kernels.fused_op import BACKEND_METHODS, DEFAULT_PRIORITY, BaseFusedOp from sglang.kernels.registry import KernelRegistry from sglang.kernels.spec import CapabilityRequirement as Cap from sglang.kernels.spec import KernelBackend, KernelSpec @@ -64,6 +64,14 @@ def test_available_backends(): } +def test_kda_backend_contract(): + assert KernelBackend.KDA.value == "KDA" + assert BACKEND_METHODS[KernelBackend.KDA] == "forward_kda" + assert DEFAULT_PRIORITY[0] is KernelBackend.KDA + with pytest.raises(NotImplementedError, match="no KDA backend"): + _ToyAdd().forward(_t(1.0), _t(2.0), backend=KernelBackend.KDA) + + def test_priority_dispatch(): # TRITON is first in priority and always eligible (no capability). assert _ToyAdd()(_t(1.0), _t(2.0)).item() == 1003.0 diff --git a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py index d8483c4bf..ce68445ec 100644 --- a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py +++ b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py @@ -30,6 +30,19 @@ EXPECTED = { "quantization.nvfp4_gemm_swiglu_nvfp4_quant": {"cute_dsl"}, "kvcache.reshape_and_cache_flash": {"triton"}, "diffusion.apply_group_norm_silu": {"triton"}, + "diffusion.norm_scale_shift": {"KDA", "cute_dsl", "flydsl"}, + "diffusion.scale_residual_norm_scale_shift": { + "KDA", + "triton", + "cute_dsl", + "flydsl", + }, + "diffusion.residual_gate_add": {"KDA"}, + "diffusion.ltx2_qknorm_split_rope": {"KDA"}, + "diffusion.causal_conv3d_cat_pad": {"KDA", "triton"}, + "diffusion.flux2_layernorm_modulate_fp8_quant": {"KDA"}, + "diffusion.flux2_qkv_epilogue": {"KDA"}, + "diffusion.flux2_token_cat_fp8": {"KDA"}, } _CPU = PlatformInfo(device_type="cpu") @@ -83,6 +96,30 @@ def test_sparse_linear_attention_registry_targets_forward_kernel(): assert spec.target.endswith(":_attn_fwd") +@pytest.mark.parametrize( + "op, target_suffix", + ( + ("diffusion.norm_scale_shift", ":kda_norm_scale_shift"), + ( + "diffusion.scale_residual_norm_scale_shift", + ":kda_scale_residual_norm_scale_shift", + ), + ("diffusion.residual_gate_add", ":residual_gate_add"), + ( + "diffusion.ltx2_qknorm_split_rope", + ":ltx2_qknorm_split_rope_cuda", + ), + ( + "diffusion.causal_conv3d_cat_pad", + ":fused_causal_conv3d_cat_pad_cuda", + ), + ), +) +def test_merged_diffusion_kda_provenance_backend(op, target_suffix): + spec = K.registry.get_backend(op, KernelBackend.KDA) + assert spec.target.endswith(target_suffix) + + def test_single_backend_resolves_without_backend(): assert ( K.select_kernel("kvcache.reshape_and_cache_flash").backend