diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen_image21_cuda.py b/python/sglang/multimodal_gen/test/unit/test_qwen_image21_cuda.py index b2746c236683..703915077760 100644 --- a/python/sglang/multimodal_gen/test/unit/test_qwen_image21_cuda.py +++ b/python/sglang/multimodal_gen/test/unit/test_qwen_image21_cuda.py @@ -455,11 +455,16 @@ def inputs_hd128(seed, edit): return kwargs +@pytest.mark.parametrize("cuda_kernels", [False, True]) @pytest.mark.parametrize("edit", [False, True]) @torch.no_grad() def test_cuda_qk_rope_pack_matches_eager_prefill_and_cached_steps( - bf16_model_hd128, edit, monkeypatch + bf16_model_hd128, edit, cuda_kernels, monkeypatch ): + if not cuda_kernels: + monkeypatch.setattr( + model_module, "can_use_qknorm_complex_rope_cuda", lambda *args: False + ) # The CUDA Q/K norm + RoPE + KV packing path must reproduce the Triton/eager # chain bit for bit on the prefill step, on cached steps and under BCG replay. actual_model = bf16_model_hd128 @@ -489,9 +494,12 @@ def test_cuda_qk_rope_pack_matches_eager_prefill_and_cached_steps( for timestep, output in zip((700, 300, 10), expected, strict=True): kwargs["timestep"].fill_(timestep) torch.testing.assert_close(actual_model(**kwargs), output, atol=0, rtol=0) - # The CUDA kernels are exact by construction and must have engaged. + # Unsupported CUDA kernels must leave the exact fallback and cache intact. for gate in (qk_gate, kv_gate): - assert gate.verified and not gate.disabled, gate.name + if cuda_kernels: + assert gate.verified and not gate.disabled, gate.name + else: + assert not gate.verified and not gate.disabled, gate.name # The packed and the direct-write projections are GEMM re-plumbings whose # first-sight compare depends on cuBLAS picking the same kernel for both # shapes; where it does not, the gate declines and the next tier takes over diff --git a/python/sglang/multimodal_gen/test/unit/test_residual_gate_add_dispatch.py b/python/sglang/multimodal_gen/test/unit/test_residual_gate_add_dispatch.py new file mode 100644 index 000000000000..417d6fc53082 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_residual_gate_add_dispatch.py @@ -0,0 +1,58 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Platform dispatch for the PTX-only diffusion residual fast path.""" + +import pytest +import torch + +from sglang.kernels.kda_kernels import residual_gate_add_jit as residual_ops + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="GPU required") + + +def _inputs(dtype, layout): + torch.manual_seed(13) + residual = torch.randn(1, 16, 32, device="cuda", dtype=dtype) + if layout == "transposed_row": + residual = residual.transpose(1, 2).contiguous().transpose(1, 2) + update = torch.randn(residual.shape, device="cuda", dtype=dtype) + shape = { + "full": residual.shape, + "row": (1, 1, 32), + "token": (1, 16, 1), + "transposed_row": (1, 1, 32), + }[layout] + gate = torch.randn(shape, device="cuda", dtype=dtype) + return residual, update, gate + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32]) +@pytest.mark.parametrize("layout", ["full", "row", "token", "transposed_row"]) +def test_hip_dispatch_never_attempts_ptx(monkeypatch, dtype, layout): + residual, update, gate = _inputs(dtype, layout) + expected = residual + update * gate + monkeypatch.setattr(torch.version, "hip", "test-hip") + + def unsupported(*args): + pytest.fail("HIP must not attempt the PTX-only kernel") + + monkeypatch.setattr(residual_ops, "_residual_gate_add_custom_op", unsupported) + assert not residual_ops.can_use_residual_gate_add_cuda(residual, update, gate) + torch.testing.assert_close( + residual_ops.residual_gate_add(residual, update, gate), expected, atol=0, rtol=0 + ) + with pytest.raises(RuntimeError, match="unsupported input"): + residual_ops.residual_gate_add_cuda(residual, update, gate) + + +@pytest.mark.skipif(torch.version.hip is not None, reason="PTX kernel requires CUDA") +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("layout", ["full", "row", "token", "transposed_row"]) +def test_cuda_fast_path_remains_exact(dtype, layout): + residual, update, gate = _inputs(dtype, layout) + assert residual_ops.can_use_residual_gate_add_cuda(residual, update, gate) + torch.testing.assert_close( + residual_ops.residual_gate_add_cuda(residual, update, gate), + residual + update * gate, + atol=0, + rtol=0, + )