[Diffusion] Fuse LongCat residual gate updates (#36577)
This commit is contained in:
+6
-6
@@ -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`.
|
- 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.
|
- 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`
|
- Kernel: `diffusion_residual_gate_add`
|
||||||
- Locations: `diffusion/residual_gate_add.py`, `csrc/diffusion/residual_gate_add.cuh`, `runtime/models/dits/ltx_2.py`
|
- 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.
|
- 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.
|
- 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`.
|
- 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_residual_gate_add.py`.
|
- 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`.
|
- 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
|
11. MiniMax-H3 indexed AdaLN modulation and gated residual fusion
|
||||||
- Kernels: `indexed_scale_shift_bf16_`, `indexed_gate_bf16_`
|
- Kernels: `indexed_scale_shift_bf16_`, `indexed_gate_bf16_`
|
||||||
|
|||||||
@@ -39,6 +39,7 @@ from sglang.kernels.ops.diffusion import (
|
|||||||
fused_gelu_active,
|
fused_gelu_active,
|
||||||
fused_linear_gelu_tanh,
|
fused_linear_gelu_tanh,
|
||||||
mark_fused_gelu_site,
|
mark_fused_gelu_site,
|
||||||
|
residual_gate_add,
|
||||||
tensors_equal,
|
tensors_equal,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.distributed import get_tp_world_size
|
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)
|
hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2)
|
||||||
gate = gate.unsqueeze(1)
|
gate = gate.unsqueeze(1)
|
||||||
hidden_states, _ = self.proj_out(hidden_states)
|
hidden_states, _ = self.proj_out(hidden_states)
|
||||||
hidden_states = gate * hidden_states
|
hidden_states = residual_gate_add(residual, hidden_states, gate)
|
||||||
hidden_states = residual + hidden_states
|
|
||||||
if hidden_states.dtype == torch.float16:
|
if hidden_states.dtype == torch.float16:
|
||||||
hidden_states = hidden_states.clip(-65504, 65504)
|
hidden_states = hidden_states.clip(-65504, 65504)
|
||||||
|
|
||||||
@@ -662,19 +662,22 @@ class _TransformerBlock(nn.Module):
|
|||||||
positions=positions,
|
positions=positions,
|
||||||
)
|
)
|
||||||
|
|
||||||
attn_output = gate_msa.unsqueeze(1) * attn_output
|
hidden_states = residual_gate_add(
|
||||||
hidden_states = hidden_states + attn_output
|
hidden_states, attn_output, gate_msa.unsqueeze(1)
|
||||||
|
)
|
||||||
|
|
||||||
norm_hidden_states = self.norm2(hidden_states)
|
norm_hidden_states = self.norm2(hidden_states)
|
||||||
norm_hidden_states = (
|
norm_hidden_states = (
|
||||||
norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None]
|
norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None]
|
||||||
)
|
)
|
||||||
ff_output = self.ff(norm_hidden_states)
|
ff_output = self.ff(norm_hidden_states)
|
||||||
ff_output = gate_mlp.unsqueeze(1) * ff_output
|
hidden_states = residual_gate_add(
|
||||||
hidden_states = hidden_states + ff_output
|
hidden_states, ff_output, gate_mlp.unsqueeze(1)
|
||||||
|
)
|
||||||
|
|
||||||
context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output
|
encoder_hidden_states = residual_gate_add(
|
||||||
encoder_hidden_states = encoder_hidden_states + context_attn_output
|
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 = self.norm2_context(encoder_hidden_states)
|
||||||
norm_encoder_hidden_states = (
|
norm_encoder_hidden_states = (
|
||||||
@@ -682,8 +685,10 @@ class _TransformerBlock(nn.Module):
|
|||||||
+ c_shift_mlp[:, None]
|
+ c_shift_mlp[:, None]
|
||||||
)
|
)
|
||||||
context_ff_output = self.ff_context(norm_encoder_hidden_states)
|
context_ff_output = self.ff_context(norm_encoder_hidden_states)
|
||||||
encoder_hidden_states = (
|
encoder_hidden_states = residual_gate_add(
|
||||||
encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output
|
encoder_hidden_states,
|
||||||
|
context_ff_output,
|
||||||
|
c_gate_mlp.unsqueeze(1),
|
||||||
)
|
)
|
||||||
if encoder_hidden_states.dtype == torch.float16:
|
if encoder_hidden_states.dtype == torch.float16:
|
||||||
encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504)
|
encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504)
|
||||||
|
|||||||
Reference in New Issue
Block a user