[diffusion] Accelerate lossless Ideogram norm post-processing (#34931)

This commit is contained in:
Xiaoyu Zhang
2026-08-16 17:20:09 +08:00
committed by GitHub
parent b752f1e533
commit 0761d3f3a4
2 changed files with 47 additions and 1 deletions
@@ -17,6 +17,8 @@ from sglang.kernels.ops.diffusion.fused_gate_rmsnorm import (
fused_rmsnorm_tanh_residual,
mark_fused_gate_rmsnorm_site,
)
from sglang.kernels.ops.diffusion.modulate_scale_shift import modulate_scale_shift
from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add
from sglang.kernels.ops.diffusion.triton.rope_rotate_half_bitexact import (
fused_rope_rotate_half_bitexact,
)
@@ -68,6 +70,7 @@ LLM_TOKEN_INDICATOR = 3
_IDEOGRAM_ROPE = BitExactFusionGate("Ideogram fused RoPE")
_IDEOGRAM_SWIGLU = BitExactFusionGate("Ideogram fused SiLU-mul")
_IDEOGRAM_ZERO_SHIFTS: dict[tuple[torch.device, torch.dtype, int], torch.Tensor] = {}
def _can_use_fused_rope(
@@ -434,6 +437,20 @@ def _norm_scale(
)
if y is not None:
return y
if (
not torch.compiler.is_compiling()
and x.is_cuda
and not torch.cuda.is_current_stream_capturing()
and x.dim() == 3
and scale.shape == (x.shape[0], 1, x.shape[-1])
and x.shape[0] == 1
):
key = (x.device, x.dtype, x.shape[-1])
zero_shift = _IDEOGRAM_ZERO_SHIFTS.get(key)
if zero_shift is None:
zero_shift = torch.zeros(1, x.shape[-1], device=x.device, dtype=x.dtype)
_IDEOGRAM_ZERO_SHIFTS[key] = zero_shift
return modulate_scale_shift(norm(x), scale.squeeze(1), zero_shift)
return norm(x) * (1.0 + scale)
@@ -455,7 +472,15 @@ def _gate_residual(
)
if y is not None:
return y
return residual + torch.tanh(gate) * norm(x)
normed = norm(x)
tanh_gate = torch.tanh(gate)
if (
not torch.compiler.is_compiling()
and x.is_cuda
and not torch.cuda.is_current_stream_capturing()
):
return residual_gate_add(residual, normed, tanh_gate)
return residual + tanh_gate * normed
class Ideogram4TransformerBlock(nn.Module):
@@ -68,8 +68,11 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload im
from sglang.multimodal_gen.runtime.models.dits.ideogram import (
Ideogram4ColumnParallelLinear,
Ideogram4MergedColumnParallelLinear,
Ideogram4RMSNorm,
Ideogram4RowParallelLinear,
Ideogram4Transformer2DModel,
_gate_residual,
_norm_scale,
)
from sglang.multimodal_gen.runtime.models.encoders.ideogram import (
IdeogramQwen3VLTextEncoder,
@@ -192,6 +195,24 @@ def _fake_ideogram_pipeline(transformer, unconditional_transformer):
class TestIdeogram4(unittest.TestCase):
def test_lossless_norm_postprocess_preserves_cpu_reference(self):
norm = Ideogram4RMSNorm(16, eps=1e-5)
x = torch.randn(1, 7, 16)
update = torch.randn_like(x)
scale = torch.randn(1, 1, 16)
gate = torch.randn_like(scale)
self.assertTrue(
torch.equal(_norm_scale(x, scale, norm, False), norm(x) * (1 + scale))
)
self.assertTrue(
torch.equal(
_gate_residual(update, gate, x, norm, False),
x + torch.tanh(gate) * norm(update),
)
)
self.assertEqual(set(norm.state_dict()), {"weight"})
def test_ideogram_dit_supports_layerwise_offload(self):
self.assertTrue(
issubclass(Ideogram4Transformer2DModel, LayerwiseOffloadableModuleMixin)