From b0c28e8f448fa6877c9321cbce43e3df847ae081 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Tue, 8 Sep 2026 20:03:21 +0800 Subject: [PATCH 1/2] Optimize LongCat normalization and modulation with verified fusion --- .../runtime/models/dits/longcat_image.py | 102 +++++++++++++--- .../test_longcat_image_norm_modulate.py | 111 ++++++++++++++++++ 2 files changed, 197 insertions(+), 16 deletions(-) create mode 100644 test/registered/unit/models/test_longcat_image_norm_modulate.py diff --git a/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py b/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py index 986b793c7664..178d8f7066f7 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py @@ -35,10 +35,14 @@ from sglang.kernels.ops.diffusion import ( BitExactFusionGate, can_use_fused_inplace_qknorm_rope, + can_use_fused_layernorm_modulate, can_use_linear_gelu, fused_gelu_active, + fused_layernorm_modulate, fused_linear_gelu_tanh, + is_plain_layer_norm, mark_fused_gelu_site, + modulate_scale_shift, residual_gate_add, tensors_equal, ) @@ -62,6 +66,75 @@ logger = init_logger(__name__) _LONGCAT_QKNORM_ROPE = BitExactFusionGate("LongCat fused QKNorm+RoPE") +_LONGCAT_LN_MOD = BitExactFusionGate("LongCat fused LN+modulate", per_signature=True) + + +def _longcat_norm_modulate( + norm: nn.Module, + x: torch.Tensor, + scale: torch.Tensor, + shift: torch.Tensor, +) -> torch.Tensor: + # LayerNorm's reduction depends on the live aten dispatch. Verify each + # shape/stride before using the bit-exact fused kernel, outside capture. + if ( + not _LONGCAT_LN_MOD.disabled + and is_plain_layer_norm(norm, x.shape[-1]) + and can_use_fused_layernorm_modulate(x, scale, shift) + ): + sig = ( + x.shape, + x.stride(), + scale.shape, + scale.stride(), + shift.shape, + shift.stride(), + norm.eps, + x.dtype, + x.device, + ) + verified = _LONGCAT_LN_MOD.is_verified(sig) + if verified or not ( + torch.compiler.is_compiling() or torch.cuda.is_current_stream_capturing() + ): + try: + out = fused_layernorm_modulate(x, scale, shift, norm.eps) + except Exception as exc: + _LONGCAT_LN_MOD.on_exception(exc, logger=logger) + else: + if verified: + return out + reference = modulate_scale_shift(norm(x), scale, shift) + return _LONGCAT_LN_MOD.accept_or_fallback( + out, reference, sig=sig, logger=logger + ) + return modulate_scale_shift(norm(x), scale, shift) + + +class _LongCatAdaLayerNormZero(AdaLayerNormZero): + def forward( + self, + x: torch.Tensor, + timestep: Optional[torch.Tensor] = None, + class_labels: Optional[torch.LongTensor] = None, + hidden_dtype: Optional[torch.dtype] = None, + emb: Optional[torch.Tensor] = None, + ): + if self.emb is not None: + emb = self.emb(timestep, class_labels, hidden_dtype=hidden_dtype) + emb = self.linear(self.silu(emb)) + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = emb.chunk( + 6, dim=1 + ) + x = _longcat_norm_modulate(self.norm, x, scale_msa, shift_msa) + return x, gate_msa, shift_mlp, scale_mlp, gate_mlp + + +class _LongCatAdaLayerNormZeroSingle(AdaLayerNormZeroSingle): + def forward(self, x: torch.Tensor, emb: Optional[torch.Tensor] = None): + emb = self.linear(self.silu(emb)) + shift_msa, scale_msa, gate_msa = emb.chunk(3, dim=1) + return _longcat_norm_modulate(self.norm, x, scale_msa, shift_msa), gate_msa def _longcat_qknorm_rope_reference( @@ -268,9 +341,9 @@ def __init__( super().__init__() tp_size = get_tp_world_size() self.num_local_heads = num_attention_heads // tp_size - assert num_attention_heads % tp_size == 0, ( - f"num_attention_heads ({num_attention_heads}) must be divisible by tp_size ({tp_size})" - ) + assert ( + num_attention_heads % tp_size == 0 + ), f"num_attention_heads ({num_attention_heads}) must be divisible by tp_size ({tp_size})" self.head_dim = attention_head_dim inner_dim = num_attention_heads * attention_head_dim @@ -429,9 +502,9 @@ def __init__( super().__init__() tp_size = get_tp_world_size() self.num_local_heads = num_attention_heads // tp_size - assert num_attention_heads % tp_size == 0, ( - f"num_attention_heads ({num_attention_heads}) must be divisible by tp_size ({tp_size})" - ) + assert ( + num_attention_heads % tp_size == 0 + ), f"num_attention_heads ({num_attention_heads}) must be divisible by tp_size ({tp_size})" self.head_dim = attention_head_dim inner_dim = num_attention_heads * attention_head_dim @@ -500,7 +573,7 @@ def __init__( ): super().__init__() self.mlp_hidden_dim = int(dim * mlp_ratio) - self.norm = AdaLayerNormZeroSingle(dim) + self.norm = _LongCatAdaLayerNormZeroSingle(dim) # proj_mlp: ColumnParallelLinear with gather_output=False keeps output # head-sharded, consistent with attn_output from _LongCatSingleAttention. self.proj_mlp = ColumnParallelLinear( @@ -619,8 +692,8 @@ def __init__( prefix: str = "", ): super().__init__() - self.norm1 = AdaLayerNormZero(dim) - self.norm1_context = AdaLayerNormZero(dim) + self.norm1 = _LongCatAdaLayerNormZero(dim) + self.norm1_context = _LongCatAdaLayerNormZero(dim) self.attn = _LongCatJointAttention( dim=dim, num_attention_heads=num_attention_heads, @@ -666,9 +739,8 @@ def forward( hidden_states, attn_output, gate_msa.unsqueeze(1) ) - norm_hidden_states = self.norm2(hidden_states) - norm_hidden_states = ( - norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None] + norm_hidden_states = _longcat_norm_modulate( + self.norm2, hidden_states, scale_mlp, shift_mlp ) ff_output = self.ff(norm_hidden_states) hidden_states = residual_gate_add( @@ -679,10 +751,8 @@ def forward( encoder_hidden_states, context_attn_output, c_gate_msa.unsqueeze(1) ) - norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states) - norm_encoder_hidden_states = ( - norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) - + c_shift_mlp[:, None] + norm_encoder_hidden_states = _longcat_norm_modulate( + self.norm2_context, encoder_hidden_states, c_scale_mlp, c_shift_mlp ) context_ff_output = self.ff_context(norm_encoder_hidden_states) encoder_hidden_states = residual_gate_add( diff --git a/test/registered/unit/models/test_longcat_image_norm_modulate.py b/test/registered/unit/models/test_longcat_image_norm_modulate.py new file mode 100644 index 000000000000..78c81eb63464 --- /dev/null +++ b/test/registered/unit/models/test_longcat_image_norm_modulate.py @@ -0,0 +1,111 @@ +"""LongCat normalization parity and graph-safe fusion dispatch.""" + +import unittest +from unittest.mock import patch + +import torch +from diffusers.models.normalization import AdaLayerNormZero, AdaLayerNormZeroSingle + +import sglang.multimodal_gen.runtime.models.dits.longcat_image as longcat +from sglang.kernels.ops.diffusion import BitExactFusionGate, modulate_scale_shift +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=45, stage="base-b", runner_config="1-gpu-large") + + +class TestLongCatNormModulation(CustomTestCase): + def setUp(self): + super().setUp() + self.original_gate = longcat._LONGCAT_LN_MOD + longcat._LONGCAT_LN_MOD = BitExactFusionGate("test", per_signature=True) + torch.manual_seed(42) + + def tearDown(self): + longcat._LONGCAT_LN_MOD = self.original_gate + super().tearDown() + + def require_cuda(self): + if not torch.cuda.is_available(): + self.skipTest("CUDA required") + + @torch.inference_mode() + def test_adaln_checkpoint_and_output_parity(self): + for device, dtype, dim, seq in [ + ("cpu", torch.float32, 64, 17), + ("cuda", torch.bfloat16, 3072, 512), + ("cuda", torch.bfloat16, 3072, 4608), + ]: + if device == "cuda" and not torch.cuda.is_available(): + continue + for reference_cls, candidate_cls in [ + (AdaLayerNormZero, longcat._LongCatAdaLayerNormZero), + (AdaLayerNormZeroSingle, longcat._LongCatAdaLayerNormZeroSingle), + ]: + with self.subTest(device=device, seq=seq, cls=reference_cls.__name__): + reference = reference_cls(dim).to(device=device, dtype=dtype) + candidate = candidate_cls(dim).to(device=device, dtype=dtype) + candidate.load_state_dict(reference.state_dict(), strict=True) + x = torch.randn(1, seq, dim, device=device, dtype=dtype) + emb = torch.randn(1, dim, device=device, dtype=dtype) + expected, actual = reference(x, emb=emb), candidate(x, emb=emb) + for a, b in zip(expected, actual, strict=True): + self.assertTrue(torch.equal(a, b)) + if device == "cuda": + self.assertTrue(longcat._LONGCAT_LN_MOD.verified) + self.assertFalse(longcat._LONGCAT_LN_MOD.disabled) + + def inputs(self, seq=4096): + self.require_cuda() + x = torch.randn(1, seq, 3072, device="cuda", dtype=torch.bfloat16) + modulation = torch.randn(1, 6 * 3072, device="cuda", dtype=torch.bfloat16) + shift, scale, *_ = modulation.chunk(6, dim=-1) + norm = torch.nn.LayerNorm(3072, elementwise_affine=False, eps=1e-6).cuda() + return norm, x, scale, shift + + @torch.inference_mode() + def test_changed_inputs_are_used_by_graph_replay(self): + norm, x, scale, shift = self.inputs() + longcat._longcat_norm_modulate(norm, x, scale, shift) + self.assertTrue(longcat._LONGCAT_LN_MOD.verified) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + actual = longcat._longcat_norm_modulate(norm, x, scale, shift) + x.add_(0.25) + scale.neg_() + shift.mul_(0.5) + graph.replay() + expected = norm(x) * (1 + scale[:, None]) + shift[:, None] + self.assertTrue(torch.equal(actual, expected)) + + @torch.inference_mode() + def test_unverified_capture_uses_eager_reference(self): + norm, x, scale, shift = self.inputs(seq=17) + expected = modulate_scale_shift(norm(x), scale, shift) + with patch.object(longcat, "fused_layernorm_modulate") as fused: + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + actual = longcat._longcat_norm_modulate(norm, x, scale, shift) + graph.replay() + fused.assert_not_called() + self.assertFalse(longcat._LONGCAT_LN_MOD.verified) + self.assertTrue(torch.equal(actual, expected)) + + @torch.inference_mode() + def test_mismatch_disables_fusion_and_returns_reference(self): + norm, x, scale, shift = self.inputs(seq=17) + expected = norm(x) * (1 + scale[:, None]) + shift[:, None] + with patch.object( + longcat, "fused_layernorm_modulate", return_value=torch.zeros_like(x) + ): + actual = longcat._longcat_norm_modulate(norm, x, scale, shift) + self.assertTrue(longcat._LONGCAT_LN_MOD.disabled) + self.assertTrue(torch.equal(actual, expected)) + with patch.object(longcat, "fused_layernorm_modulate") as fused: + actual = longcat._longcat_norm_modulate(norm, x, scale, shift) + fused.assert_not_called() + self.assertTrue(torch.equal(actual, expected)) + + +if __name__ == "__main__": + unittest.main() From cafc9addfd05ccb17a99011bc49d7b9154250200 Mon Sep 17 00:00:00 2001 From: BBuf <1182563586@qq.com> Date: Tue, 8 Sep 2026 23:27:57 +0800 Subject: [PATCH 2/2] fix(diffusion): preserve LongCat reference for autograd and register GPU coverage --- .../runtime/models/dits/longcat_image.py | 34 +++++++++---------- .../test_longcat_image_norm_modulate.py | 28 ++++++++++++--- 2 files changed, 40 insertions(+), 22 deletions(-) rename test/registered/{unit/models => kernel/diffusion}/test_longcat_image_norm_modulate.py (77%) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py b/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py index 178d8f7066f7..9f8ed95c33d1 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py @@ -32,17 +32,14 @@ AdaLayerNormZeroSingle, ) +from sglang.kernels.ops import diffusion as diffusion_ops from sglang.kernels.ops.diffusion import ( BitExactFusionGate, can_use_fused_inplace_qknorm_rope, - can_use_fused_layernorm_modulate, can_use_linear_gelu, fused_gelu_active, - fused_layernorm_modulate, fused_linear_gelu_tanh, - is_plain_layer_norm, mark_fused_gelu_site, - modulate_scale_shift, residual_gate_add, tensors_equal, ) @@ -75,12 +72,15 @@ def _longcat_norm_modulate( scale: torch.Tensor, shift: torch.Tensor, ) -> torch.Tensor: + if torch.is_grad_enabled() or torch.compiler.is_compiling(): + return norm(x) * (1 + scale[:, None]) + shift[:, None] # LayerNorm's reduction depends on the live aten dispatch. Verify each # shape/stride before using the bit-exact fused kernel, outside capture. if ( not _LONGCAT_LN_MOD.disabled - and is_plain_layer_norm(norm, x.shape[-1]) - and can_use_fused_layernorm_modulate(x, scale, shift) + and x.is_cuda + and diffusion_ops.is_plain_layer_norm(norm, x.shape[-1]) + and diffusion_ops.can_use_fused_layernorm_modulate(x, scale, shift) ): sig = ( x.shape, @@ -94,21 +94,19 @@ def _longcat_norm_modulate( x.device, ) verified = _LONGCAT_LN_MOD.is_verified(sig) - if verified or not ( - torch.compiler.is_compiling() or torch.cuda.is_current_stream_capturing() - ): + if verified or not torch.cuda.is_current_stream_capturing(): try: - out = fused_layernorm_modulate(x, scale, shift, norm.eps) + out = diffusion_ops.fused_layernorm_modulate(x, scale, shift, norm.eps) except Exception as exc: _LONGCAT_LN_MOD.on_exception(exc, logger=logger) else: if verified: return out - reference = modulate_scale_shift(norm(x), scale, shift) + reference = norm(x) * (1 + scale[:, None]) + shift[:, None] return _LONGCAT_LN_MOD.accept_or_fallback( out, reference, sig=sig, logger=logger ) - return modulate_scale_shift(norm(x), scale, shift) + return norm(x) * (1 + scale[:, None]) + shift[:, None] class _LongCatAdaLayerNormZero(AdaLayerNormZero): @@ -341,9 +339,9 @@ def __init__( super().__init__() tp_size = get_tp_world_size() self.num_local_heads = num_attention_heads // tp_size - assert ( - num_attention_heads % tp_size == 0 - ), f"num_attention_heads ({num_attention_heads}) must be divisible by tp_size ({tp_size})" + assert num_attention_heads % tp_size == 0, ( + f"num_attention_heads ({num_attention_heads}) must be divisible by tp_size ({tp_size})" + ) self.head_dim = attention_head_dim inner_dim = num_attention_heads * attention_head_dim @@ -502,9 +500,9 @@ def __init__( super().__init__() tp_size = get_tp_world_size() self.num_local_heads = num_attention_heads // tp_size - assert ( - num_attention_heads % tp_size == 0 - ), f"num_attention_heads ({num_attention_heads}) must be divisible by tp_size ({tp_size})" + assert num_attention_heads % tp_size == 0, ( + f"num_attention_heads ({num_attention_heads}) must be divisible by tp_size ({tp_size})" + ) self.head_dim = attention_head_dim inner_dim = num_attention_heads * attention_head_dim diff --git a/test/registered/unit/models/test_longcat_image_norm_modulate.py b/test/registered/kernel/diffusion/test_longcat_image_norm_modulate.py similarity index 77% rename from test/registered/unit/models/test_longcat_image_norm_modulate.py rename to test/registered/kernel/diffusion/test_longcat_image_norm_modulate.py index 78c81eb63464..fd22bbd90ab2 100644 --- a/test/registered/unit/models/test_longcat_image_norm_modulate.py +++ b/test/registered/kernel/diffusion/test_longcat_image_norm_modulate.py @@ -11,7 +11,7 @@ from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import CustomTestCase -register_cuda_ci(est_time=45, stage="base-b", runner_config="1-gpu-large") +register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="1-gpu-large") class TestLongCatNormModulation(CustomTestCase): @@ -63,6 +63,24 @@ def inputs(self, seq=4096): norm = torch.nn.LayerNorm(3072, elementwise_affine=False, eps=1e-6).cuda() return norm, x, scale, shift + def test_grad_enabled_uses_differentiable_reference(self): + norm, x, scale, shift = self.inputs(seq=17) + with torch.inference_mode(): + longcat._longcat_norm_modulate(norm, x, scale, shift) + self.assertTrue(longcat._LONGCAT_LN_MOD.verified) + leaves = [t.detach().clone().requires_grad_() for t in (x, scale, shift)] + refs = [t.detach().clone().requires_grad_() for t in leaves] + with patch.object(longcat.diffusion_ops, "fused_layernorm_modulate") as fused: + actual = longcat._longcat_norm_modulate(norm, *leaves) + actual.float().sum().backward() + fused.assert_not_called() + expected = norm(refs[0]) * (1 + refs[1][:, None]) + refs[2][:, None] + expected.float().sum().backward() + self.assertTrue(torch.equal(actual, expected)) + for a, b in zip(leaves, refs, strict=True): + self.assertIsNotNone(a.grad) + self.assertTrue(torch.equal(a.grad, b.grad)) + @torch.inference_mode() def test_changed_inputs_are_used_by_graph_replay(self): norm, x, scale, shift = self.inputs() @@ -82,7 +100,7 @@ def test_changed_inputs_are_used_by_graph_replay(self): def test_unverified_capture_uses_eager_reference(self): norm, x, scale, shift = self.inputs(seq=17) expected = modulate_scale_shift(norm(x), scale, shift) - with patch.object(longcat, "fused_layernorm_modulate") as fused: + with patch.object(longcat.diffusion_ops, "fused_layernorm_modulate") as fused: graph = torch.cuda.CUDAGraph() with torch.cuda.graph(graph): actual = longcat._longcat_norm_modulate(norm, x, scale, shift) @@ -96,12 +114,14 @@ def test_mismatch_disables_fusion_and_returns_reference(self): norm, x, scale, shift = self.inputs(seq=17) expected = norm(x) * (1 + scale[:, None]) + shift[:, None] with patch.object( - longcat, "fused_layernorm_modulate", return_value=torch.zeros_like(x) + longcat.diffusion_ops, + "fused_layernorm_modulate", + return_value=torch.zeros_like(x), ): actual = longcat._longcat_norm_modulate(norm, x, scale, shift) self.assertTrue(longcat._LONGCAT_LN_MOD.disabled) self.assertTrue(torch.equal(actual, expected)) - with patch.object(longcat, "fused_layernorm_modulate") as fused: + with patch.object(longcat.diffusion_ops, "fused_layernorm_modulate") as fused: actual = longcat._longcat_norm_modulate(norm, x, scale, shift) fused.assert_not_called() self.assertTrue(torch.equal(actual, expected))