[diffusion] LTX-2: mount the bit-exact fused modulate at the 8 bare adaLN sites (ltx23-one-stage denoise -2.8% H100 / -2.6% H200) (#34315)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
dd8c5849af
commit
546965fc72
@@ -10,6 +10,7 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from sglang.kernels.ops.diffusion.bitexact_gate import BitExactFusionGate
|
||||||
from sglang.kernels.ops.diffusion.fused_linear_gelu import (
|
from sglang.kernels.ops.diffusion.fused_linear_gelu import (
|
||||||
can_fuse_linear_gelu,
|
can_fuse_linear_gelu,
|
||||||
fused_gelu_active,
|
fused_gelu_active,
|
||||||
@@ -26,6 +27,10 @@ from sglang.kernels.ops.diffusion.ltx2_rmsnorm_modulate import (
|
|||||||
ltx2_rms_norm_modulate_active,
|
ltx2_rms_norm_modulate_active,
|
||||||
mark_ltx2_rms_norm_modulate_site,
|
mark_ltx2_rms_norm_modulate_site,
|
||||||
)
|
)
|
||||||
|
from sglang.kernels.ops.diffusion.modulate_scale_shift import (
|
||||||
|
can_use_modulate_scale_shift_cuda,
|
||||||
|
modulate_scale_shift_cuda,
|
||||||
|
)
|
||||||
from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add
|
from sglang.kernels.ops.diffusion.residual_gate_add import residual_gate_add
|
||||||
from sglang.multimodal_gen.configs.models.dits.ltx_2 import LTX2ArchConfig, LTX2Config
|
from sglang.multimodal_gen.configs.models.dits.ltx_2 import LTX2ArchConfig, LTX2Config
|
||||||
from sglang.multimodal_gen.configs.models.fsdp import (
|
from sglang.multimodal_gen.configs.models.fsdp import (
|
||||||
@@ -140,6 +145,54 @@ def adaln_embedding_coefficient(cross_attention_adaln: bool) -> int:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_LTX2_MODULATE = BitExactFusionGate("LTX-2 fused modulate")
|
||||||
|
|
||||||
|
|
||||||
|
def _ltx2_modulate(
|
||||||
|
x: torch.Tensor, scale: torch.Tensor, shift: torch.Tensor
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""``x * (1 + scale) + shift`` via the bit-exact fused CUDA modulate.
|
||||||
|
|
||||||
|
The adaLN scale/shift rows here are ``(B, 1, D)`` views (from
|
||||||
|
``unbind``/``squeeze`` of the combined tables), so they are made dense and
|
||||||
|
squeezed to the ``(B, D)`` layout the kernel expects — a bit-exact copy of
|
||||||
|
tiny tensors. Per-token modulation rows (``dim 1 > 1``) and any other
|
||||||
|
unsupported layout keep the eager chain; the first fused call
|
||||||
|
self-verifies ``torch.equal`` against the eager chain and falls back
|
||||||
|
permanently on any mismatch.
|
||||||
|
"""
|
||||||
|
verified = _LTX2_MODULATE.verified
|
||||||
|
if (
|
||||||
|
not _LTX2_MODULATE.disabled
|
||||||
|
and x.dim() == 3
|
||||||
|
and x.is_contiguous()
|
||||||
|
and scale.dim() == 3
|
||||||
|
and scale.shape == (x.shape[0], 1, x.shape[-1])
|
||||||
|
and shift.shape == scale.shape
|
||||||
|
and (verified or _LTX2_MODULATE.can_attempt_once())
|
||||||
|
):
|
||||||
|
scale_rows = scale.squeeze(1).contiguous()
|
||||||
|
shift_rows = shift.squeeze(1).contiguous()
|
||||||
|
if can_use_modulate_scale_shift_cuda(x, scale_rows, shift_rows):
|
||||||
|
try:
|
||||||
|
out = modulate_scale_shift_cuda(x, scale_rows, shift_rows)
|
||||||
|
except Exception as exc:
|
||||||
|
_LTX2_MODULATE.on_exception(exc, logger=logger)
|
||||||
|
else:
|
||||||
|
if verified:
|
||||||
|
return out
|
||||||
|
return _LTX2_MODULATE.accept_or_fallback(
|
||||||
|
out,
|
||||||
|
x * (1 + scale) + shift,
|
||||||
|
logger=logger,
|
||||||
|
mismatch_msg=(
|
||||||
|
"LTX-2 fused modulate is not bit-exact on this "
|
||||||
|
"platform; falling back to eager"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
return x * (1 + scale) + shift
|
||||||
|
|
||||||
|
|
||||||
def _ltx2_rms_norm_modulate(
|
def _ltx2_rms_norm_modulate(
|
||||||
block: nn.Module,
|
block: nn.Module,
|
||||||
rms_norm: nn.Module,
|
rms_norm: nn.Module,
|
||||||
@@ -1301,8 +1354,8 @@ class LTX2TransformerBlock(nn.Module):
|
|||||||
norm_hidden_states = _ltx2_rms_norm_modulate(
|
norm_hidden_states = _ltx2_rms_norm_modulate(
|
||||||
self, self.rms_norm, hidden_states, vscale_q, vshift_q, self.norm_eps
|
self, self.rms_norm, hidden_states, vscale_q, vshift_q, self.norm_eps
|
||||||
)
|
)
|
||||||
mod_encoder_hidden_states = (
|
mod_encoder_hidden_states = _ltx2_modulate(
|
||||||
encoder_hidden_states * (1 + v_prompt_scale) + v_prompt_shift
|
encoder_hidden_states, v_prompt_scale, v_prompt_shift
|
||||||
)
|
)
|
||||||
attn_hidden_states = self.attn2(
|
attn_hidden_states = self.attn2(
|
||||||
norm_hidden_states,
|
norm_hidden_states,
|
||||||
@@ -1333,8 +1386,8 @@ class LTX2TransformerBlock(nn.Module):
|
|||||||
ashift_q,
|
ashift_q,
|
||||||
self.norm_eps,
|
self.norm_eps,
|
||||||
)
|
)
|
||||||
mod_audio_encoder_hidden_states = (
|
mod_audio_encoder_hidden_states = _ltx2_modulate(
|
||||||
audio_encoder_hidden_states * (1 + a_prompt_scale) + a_prompt_shift
|
audio_encoder_hidden_states, a_prompt_scale, a_prompt_shift
|
||||||
)
|
)
|
||||||
attn_audio_hidden_states = self.audio_attn2(
|
attn_audio_hidden_states = self.audio_attn2(
|
||||||
norm_audio_hidden_states,
|
norm_audio_hidden_states,
|
||||||
@@ -1423,11 +1476,11 @@ class LTX2TransformerBlock(nn.Module):
|
|||||||
v2a_gate = audio_ca_gate[0].squeeze(2)
|
v2a_gate = audio_ca_gate[0].squeeze(2)
|
||||||
|
|
||||||
# A2V
|
# A2V
|
||||||
mod_norm_hidden_states = (
|
mod_norm_hidden_states = _ltx2_modulate(
|
||||||
norm_hidden_states * (1 + video_a2v_ca_scale) + video_a2v_ca_shift
|
norm_hidden_states, video_a2v_ca_scale, video_a2v_ca_shift
|
||||||
)
|
)
|
||||||
mod_norm_audio_hidden_states = (
|
mod_norm_audio_hidden_states = _ltx2_modulate(
|
||||||
norm_audio_hidden_states * (1 + audio_a2v_ca_scale) + audio_a2v_ca_shift
|
norm_audio_hidden_states, audio_a2v_ca_scale, audio_a2v_ca_shift
|
||||||
)
|
)
|
||||||
|
|
||||||
if not skip_a2v_cross_attn:
|
if not skip_a2v_cross_attn:
|
||||||
@@ -1448,11 +1501,11 @@ class LTX2TransformerBlock(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# V2A
|
# V2A
|
||||||
mod_norm_hidden_states = (
|
mod_norm_hidden_states = _ltx2_modulate(
|
||||||
norm_hidden_states * (1 + video_v2a_ca_scale) + video_v2a_ca_shift
|
norm_hidden_states, video_v2a_ca_scale, video_v2a_ca_shift
|
||||||
)
|
)
|
||||||
mod_norm_audio_hidden_states = (
|
mod_norm_audio_hidden_states = _ltx2_modulate(
|
||||||
norm_audio_hidden_states * (1 + audio_v2a_ca_scale) + audio_v2a_ca_shift
|
norm_audio_hidden_states, audio_v2a_ca_scale, audio_v2a_ca_shift
|
||||||
)
|
)
|
||||||
|
|
||||||
if not skip_v2a_cross_attn:
|
if not skip_v2a_cross_attn:
|
||||||
@@ -2220,7 +2273,7 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
shift, scale = scale_shift_values[:, :, 0], scale_shift_values[:, :, 1]
|
shift, scale = scale_shift_values[:, :, 0], scale_shift_values[:, :, 1]
|
||||||
with torch.autocast(device_type=hidden_states.device.type, enabled=False):
|
with torch.autocast(device_type=hidden_states.device.type, enabled=False):
|
||||||
hidden_states = self.norm_out(hidden_states)
|
hidden_states = self.norm_out(hidden_states)
|
||||||
hidden_states = hidden_states * (1 + scale) + shift
|
hidden_states = _ltx2_modulate(hidden_states, scale, shift)
|
||||||
hidden_states, _ = self.proj_out(hidden_states)
|
hidden_states, _ = self.proj_out(hidden_states)
|
||||||
|
|
||||||
# Audio
|
# Audio
|
||||||
@@ -2233,7 +2286,9 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
)
|
)
|
||||||
with torch.autocast(device_type=audio_hidden_states.device.type, enabled=False):
|
with torch.autocast(device_type=audio_hidden_states.device.type, enabled=False):
|
||||||
audio_hidden_states = self.audio_norm_out(audio_hidden_states)
|
audio_hidden_states = self.audio_norm_out(audio_hidden_states)
|
||||||
audio_hidden_states = audio_hidden_states * (1 + audio_scale) + audio_shift
|
audio_hidden_states = _ltx2_modulate(
|
||||||
|
audio_hidden_states, audio_scale, audio_shift
|
||||||
|
)
|
||||||
audio_hidden_states, _ = self.audio_proj_out(audio_hidden_states)
|
audio_hidden_states, _ = self.audio_proj_out(audio_hidden_states)
|
||||||
# Unpatchify if requested (default True for pipeline compatibility)
|
# Unpatchify if requested (default True for pipeline compatibility)
|
||||||
return_latents = kwargs.get("return_latents", True)
|
return_latents = kwargs.get("return_latents", True)
|
||||||
|
|||||||
@@ -0,0 +1,48 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.models.dits.ltx_2 import _ltx2_modulate
|
||||||
|
|
||||||
|
|
||||||
|
class TestLtx2ModulateMount(unittest.TestCase):
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
|
||||||
|
def test_row_broadcast_matches_eager(self):
|
||||||
|
torch.manual_seed(0)
|
||||||
|
x = torch.randn(2, 517, 4096, device="cuda", dtype=torch.bfloat16)
|
||||||
|
scale = torch.randn(2, 1, 4096, device="cuda", dtype=torch.bfloat16)
|
||||||
|
shift = torch.randn(2, 1, 4096, device="cuda", dtype=torch.bfloat16)
|
||||||
|
reference = x * (1 + scale) + shift
|
||||||
|
self.assertTrue(torch.equal(reference, _ltx2_modulate(x, scale, shift)))
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
|
||||||
|
def test_non_contiguous_rows_match_eager(self):
|
||||||
|
# unbind()/squeeze() of the combined adaLN tables produces strided
|
||||||
|
# (B, 1, D) views; the helper densifies them before the kernel.
|
||||||
|
torch.manual_seed(1)
|
||||||
|
x = torch.randn(2, 33, 512, device="cuda", dtype=torch.bfloat16)
|
||||||
|
table = torch.randn(2, 1, 4, 512, device="cuda", dtype=torch.bfloat16)
|
||||||
|
scale, shift = table.unbind(dim=2)[:2]
|
||||||
|
reference = x * (1 + scale) + shift
|
||||||
|
self.assertTrue(torch.equal(reference, _ltx2_modulate(x, scale, shift)))
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
|
||||||
|
def test_per_token_rows_fall_back(self):
|
||||||
|
torch.manual_seed(2)
|
||||||
|
x = torch.randn(2, 16, 128, device="cuda", dtype=torch.bfloat16)
|
||||||
|
scale = torch.randn(2, 16, 128, device="cuda", dtype=torch.bfloat16)
|
||||||
|
shift = torch.randn(2, 16, 128, device="cuda", dtype=torch.bfloat16)
|
||||||
|
reference = x * (1 + scale) + shift
|
||||||
|
self.assertTrue(torch.equal(reference, _ltx2_modulate(x, scale, shift)))
|
||||||
|
|
||||||
|
def test_cpu_falls_back(self):
|
||||||
|
torch.manual_seed(3)
|
||||||
|
x = torch.randn(1, 9, 64, dtype=torch.float32)
|
||||||
|
scale = torch.randn(1, 1, 64, dtype=torch.float32)
|
||||||
|
shift = torch.randn(1, 1, 64, dtype=torch.float32)
|
||||||
|
reference = x * (1 + scale) + shift
|
||||||
|
self.assertTrue(torch.equal(reference, _ltx2_modulate(x, scale, shift)))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user