diff --git a/python/sglang/kernels/ops/diffusion/__init__.py b/python/sglang/kernels/ops/diffusion/__init__.py index 5e221bed6f..31768ea567 100644 --- a/python/sglang/kernels/ops/diffusion/__init__.py +++ b/python/sglang/kernels/ops/diffusion/__init__.py @@ -287,6 +287,13 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = ( _CUDA, "Paired in-place Helios transposed Q/K RoPE.", ), + ( + "diffusion.complex_rope", + KernelBackend.TRITON, + "rope.complex_rope_triton:fused_complex_rope", + _CUDA, + "Paired RoPE preserving PyTorch complex64 multiplication rounding.", + ), ( "diffusion.hunyuan_qkv_rope_pack", KernelBackend.TRITON, @@ -294,6 +301,13 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = ( _CUDA, "HunyuanVideo QKV pack + RoPE.", ), + ( + "diffusion.rmsnorm_preserve_reduction", + KernelBackend.TRITON, + "norm.rmsnorm_preserve_reduction:rmsnorm_preserve_reduction", + _CUDA, + "Cast-before-weight RMSNorm preserving the native FP32 mean reduction.", + ), ( "diffusion.silu_mul", KernelBackend.TRITON, @@ -545,6 +559,8 @@ _EXPORTS: dict[str, str] = { "try_fused_bias_mul_add": "sglang.kernels.kda_kernels.norm_scale_shift_jit", "try_fused_bias_scale_residual_norm_scale_shift": "sglang.kernels.kda_kernels.norm_scale_shift_jit", "triton_one_pass_rms_norm": "norm.rmsnorm_onepass_triton", + "can_use_rmsnorm_preserve_reduction": "norm.rmsnorm_preserve_reduction", + "rmsnorm_preserve_reduction": "norm.rmsnorm_preserve_reduction", "can_use_fused_rmsnorm_scale_shift": "norm.rmsnorm_scale_shift_bitexact", "can_use_fused_scale_residual_rmsnorm_scale_shift": "norm.rmsnorm_scale_shift_bitexact", "fused_rmsnorm_scale_shift_bitexact": "norm.rmsnorm_scale_shift_bitexact", @@ -598,6 +614,8 @@ _EXPORTS: dict[str, str] = { "can_use_helios_qk_rope": "rope.helios_qk_rope_jit", "fused_inplace_helios_qk_rope": "rope.helios_qk_rope_jit", "apply_rotary_embedding": "rope.rotary_triton", + "can_use_fused_complex_rope": "rope.complex_rope_triton", + "fused_complex_rope": "rope.complex_rope_triton", # Tensor layout transformations fused with downstream quantization "try_flux2_token_cat_fp8": "sglang.kernels.kda_kernels.flux2_token_cat_fp8_triton", # Activation-function fusions diff --git a/python/sglang/multimodal_gen/runtime/server_args/server_args.py b/python/sglang/multimodal_gen/runtime/server_args/server_args.py index bf65d948f9..f84ec70391 100644 --- a/python/sglang/multimodal_gen/runtime/server_args/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args/server_args.py @@ -216,8 +216,10 @@ BREAKABLE_CUDA_GRAPH_SUPPORTED_MODEL_IDS = frozenset( "minimaxai/minimax-h3", "qwen/qwen-image", "qwen/qwen-image-2512", + "qwen/qwen-image-2.1", "qwen-image", "qwen-image-2512", + "qwen-image-2.1", "tongyi-mai/z-image", "tongyi-mai/z-image-turbo", "zai-org/glm-image", @@ -238,6 +240,7 @@ BREAKABLE_CUDA_GRAPH_SUPPORTED_PIPELINE_CONFIGS = frozenset( "LongCatImageEditPipelineConfig", "MiniMaxH3PipelineConfig", "QwenImagePipelineConfig", + "QwenImage21PipelineConfig", "SanaPipelineConfig", "SanaVideoPipelineConfig", "ZImagePipelineConfig", @@ -776,7 +779,8 @@ class ServerArgs(DisaggServerArgsMixin): logger.warning( "[Diffusion BCG] disabled for %s: only FLUX.1-dev, Ideogram-4, " "jdopensource/JoyAI-Echo, Lightricks/LTX-2, LongCat-Image, " - "MiniMax-H3, Qwen/Qwen-Image, Qwen/Qwen-Image-2512, SANA1.5, " + "MiniMax-H3, Qwen/Qwen-Image, Qwen/Qwen-Image-2512, " + "Qwen/Qwen-Image-2.1, SANA1.5, " "SANA-Video, Tongyi-MAI/Z-Image/Z-Image-Turbo, and " "zai-org/GLM-Image are currently supported.", pipeline_config_name, diff --git a/test/registered/kernels/ops/diffusion/test_model_fast_paths.py b/test/registered/kernels/ops/diffusion/test_model_fast_paths.py index 4a70127127..2a95f8c779 100644 --- a/test/registered/kernels/ops/diffusion/test_model_fast_paths.py +++ b/test/registered/kernels/ops/diffusion/test_model_fast_paths.py @@ -37,8 +37,10 @@ import sglang.multimodal_gen.runtime.models.dits.glm_image as glm_image import sglang.multimodal_gen.runtime.models.dits.longcat_image as longcat_image import sglang.multimodal_gen.runtime.models.dits.ltx_2 as ltx2_module import sglang.multimodal_gen.runtime.models.dits.qwen_image as qwen_image +import sglang.multimodal_gen.runtime.models.dits.qwen_image21 as qwen_image21 import sglang.multimodal_gen.runtime.models.dits.sana as sana from sglang.kernels.ops.diffusion import ( + BitExactFusionGate, can_use_fused_layernorm_modulate, can_use_fused_qk_head_layernorm, can_use_fused_rmsnorm_scale_shift, @@ -1543,5 +1545,97 @@ def test_autoencoder_kl_fastpath_install(): assert torch.equal(opt.decode(z), ref) +@torch.no_grad() +def test_qwen21_qk_norm_verifies_and_preserves_native_fallback(monkeypatch): + x = torch.randn(1, 257, 8, 128, device="cuda", dtype=torch.bfloat16) + norm = qwen_image21.RMSNorm( + 128, 1e-6, cast_x_before_out_mul=True, force_native=True + ).to(device=x.device, dtype=x.dtype) + norm.weight.normal_() + expected = norm(x) + gate = BitExactFusionGate("test Q/K norm") + monkeypatch.setattr(qwen_image21, "_QK_NORM_FUSION", gate) + assert torch.equal(qwen_image21.apply_qk_norm(x, norm), expected) + assert gate.verified and not gate.disabled + x.normal_() + assert torch.equal(qwen_image21.apply_qk_norm(x, norm), norm(x)) + + gate = BitExactFusionGate("test mismatched Q/K norm") + monkeypatch.setattr(qwen_image21, "_QK_NORM_FUSION", gate) + monkeypatch.setattr( + qwen_image21, + "rmsnorm_preserve_reduction", + lambda x, weight, eps: torch.zeros_like(x), + ) + assert torch.equal(qwen_image21.apply_qk_norm(x, norm), norm(x)) + assert gate.disabled and not gate.verified + + +@torch.no_grad() +def test_qwen21_qk_norm_does_not_verify_during_capture(monkeypatch): + x = torch.randn(1, 17, 2, 128, device="cuda", dtype=torch.bfloat16) + norm = qwen_image21.RMSNorm( + 128, 1e-6, cast_x_before_out_mul=True, force_native=True + ).to(device=x.device, dtype=x.dtype) + norm(x) + gate = BitExactFusionGate("test captured Q/K norm") + monkeypatch.setattr(qwen_image21, "_QK_NORM_FUSION", gate) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + out = qwen_image21.apply_qk_norm(x, norm) + assert not gate.verified and not gate.disabled + x.normal_() + graph.replay() + assert torch.equal(out, norm(x)) + + +@torch.no_grad() +def test_qwen21_modulation_verifies_and_preserves_native_fallback(monkeypatch): + x = torch.randn(1, 257, 4096, device="cuda", dtype=torch.bfloat16) + scale = torch.randn(1, 1, 4096, device=x.device, dtype=x.dtype) + norm = torch.nn.LayerNorm(4096, eps=1e-6, elementwise_affine=False).cuda() + gate = BitExactFusionGate("test scale-only modulation") + monkeypatch.setattr(qwen_image21, "_MODULATION_FUSION", gate) + expected = norm(x) * (1 + scale) + actual = qwen_image21.apply_modulation(x, norm, scale) + assert torch.equal(actual.view(torch.int16), expected.view(torch.int16)) + assert gate.verified and not gate.disabled + x.normal_() + scale.normal_() + assert torch.equal( + qwen_image21.apply_modulation(x, norm, scale), norm(x) * (1 + scale) + ) + + gate = BitExactFusionGate("test mismatched modulation") + monkeypatch.setattr(qwen_image21, "_MODULATION_FUSION", gate) + monkeypatch.setattr( + qwen_image21, + "fused_layernorm_modulate", + lambda x, scale, shift, eps: torch.zeros_like(x), + ) + assert torch.equal( + qwen_image21.apply_modulation(x, norm, scale), norm(x) * (1 + scale) + ) + assert gate.disabled and not gate.verified + + +@torch.no_grad() +def test_qwen21_modulation_does_not_verify_during_capture(monkeypatch): + x = torch.randn(1, 17, 128, device="cuda", dtype=torch.bfloat16) + scale = torch.randn(1, 1, 128, device=x.device, dtype=x.dtype) + norm = torch.nn.LayerNorm(128, eps=1e-6, elementwise_affine=False).cuda() + norm(x) + gate = BitExactFusionGate("test captured modulation") + monkeypatch.setattr(qwen_image21, "_MODULATION_FUSION", gate) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + out = qwen_image21.apply_modulation(x, norm, scale) + assert not gate.verified and not gate.disabled + x.normal_() + scale.normal_() + graph.replay() + assert torch.equal(out, norm(x) * (1 + scale)) + + if __name__ == "__main__": sys.exit(pytest.main([__file__, "-v"]))