diff --git a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md index 14ca25968..cfaebe848 100644 --- a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md +++ b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-benchmark-profile/existing-fast-paths.md @@ -170,15 +170,15 @@ framework-specific optimization workflow. - Constraints: `cos` and `sin` shapes must match `[B, H, S, head_dim / 2]`, and `inner_dim == H * head_dim`. - Workflow rule: if LTX-2 traces show a large split-RoPE PyTorch chain, check whether the LTX2-specific Triton path was disabled by shape or dtype before proposing a new RoPE kernel. -10. LTX2 residual-gate add fusion +10. LTX2 and LongCat-Image residual-gate add fusion - Kernel: `diffusion_residual_gate_add` -- Locations: `diffusion/residual_gate_add.py`, `csrc/diffusion/residual_gate_add.cuh`, `runtime/models/dits/ltx_2.py` -- Use case: `residual + update * gate` in LTX2 self-attention, prompt cross-attention, audio/video cross-attention, and feed-forward residual updates. +- Locations: `kernels/ops/diffusion/modulate/residual_gate_add_jit.py`, `kernels/jit/csrc/diffusion/residual_gate_add.cuh`, `runtime/models/dits/ltx_2.py`, `runtime/models/dits/longcat_image.py` +- Use case: `residual + update * gate` in LTX2 self-attention, prompt cross-attention, audio/video cross-attention, and feed-forward residual updates, plus LongCat-Image joint- and single-stream transformer residuals. - Constraints: `residual`, `update`, and `gate` must be CUDA tensors on the same device, contiguous, same dtype (`fp16`, `bf16`, or `fp32`), with `update.shape == residual.shape`; `gate` can match `residual` or be row-broadcast with the last dimension matching. -- Behavior: LTX2 calls `residual_gate_add(...)` from the kernels package directly. The CUDA custom op is used while guards pass. On a runtime exception outside `torch.compile`, it logs once, disables the fast path for the process, and falls back to `residual + update * gate`. -- Validation: `test/registered/kernels/ops/diffusion/test_residual_gate_add.py`. +- Behavior: LTX2 and LongCat-Image call `residual_gate_add(...)` from the kernels package directly. The CUDA custom op is used while guards pass. On a runtime exception outside `torch.compile`, it logs once, disables the fast path for the process, and falls back to `residual + update * gate`. +- Validation: `test/registered/kernels/ops/diffusion/test_modulate.py`, `python/sglang/multimodal_gen/test/unit/test_longcat_image_residual_gate.py`. - Microbench: `test/registered/kernels/benchmark/diffusion/bench_residual_gate_add.py`. -- Workflow rule: if LTX2 traces show repeated elementwise `mul` + `add` ladders around attention or MLP residuals, check whether this existing CUDA path was disabled by shape, dtype, contiguity, or a prior runtime failure before proposing another elementwise fusion. +- Workflow rule: if LTX2 or LongCat-Image traces show repeated elementwise `mul` + `add` ladders around attention or MLP residuals, check whether this existing CUDA path was disabled by shape, dtype, contiguity, or a prior runtime failure before proposing another elementwise fusion. 11. MiniMax-H3 indexed AdaLN modulation and gated residual fusion - Kernels: `indexed_scale_shift_bf16_`, `indexed_gate_bf16_` 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 5739d1dba..b2d5ba293 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py @@ -39,6 +39,7 @@ from sglang.kernels.ops.diffusion import ( fused_gelu_active, fused_linear_gelu_tanh, mark_fused_gelu_site, + residual_gate_add, tensors_equal, ) from sglang.multimodal_gen.runtime.distributed import get_tp_world_size @@ -597,8 +598,7 @@ class _SingleTransformerBlock(nn.Module): hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) gate = gate.unsqueeze(1) hidden_states, _ = self.proj_out(hidden_states) - hidden_states = gate * hidden_states - hidden_states = residual + hidden_states + hidden_states = residual_gate_add(residual, hidden_states, gate) if hidden_states.dtype == torch.float16: hidden_states = hidden_states.clip(-65504, 65504) @@ -662,19 +662,22 @@ class _TransformerBlock(nn.Module): positions=positions, ) - attn_output = gate_msa.unsqueeze(1) * attn_output - hidden_states = hidden_states + attn_output + hidden_states = residual_gate_add( + 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] ) ff_output = self.ff(norm_hidden_states) - ff_output = gate_mlp.unsqueeze(1) * ff_output - hidden_states = hidden_states + ff_output + hidden_states = residual_gate_add( + hidden_states, ff_output, gate_mlp.unsqueeze(1) + ) - context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output - encoder_hidden_states = encoder_hidden_states + context_attn_output + encoder_hidden_states = residual_gate_add( + 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 = ( @@ -682,8 +685,10 @@ class _TransformerBlock(nn.Module): + c_shift_mlp[:, None] ) context_ff_output = self.ff_context(norm_encoder_hidden_states) - encoder_hidden_states = ( - encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output + encoder_hidden_states = residual_gate_add( + encoder_hidden_states, + context_ff_output, + c_gate_mlp.unsqueeze(1), ) if encoder_hidden_states.dtype == torch.float16: encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504)