[Diffusion] Fuse LTX2 Ada values (#29390)

This commit is contained in:
Xiaoyu Zhang
2026-06-26 23:13:15 +08:00
committed by GitHub
parent b91348071e
commit 18b0e5757e
3 changed files with 339 additions and 20 deletions
@@ -0,0 +1,184 @@
# Adapted from NVlabs/Sana sol-engine LTX2 Ada-value fusion.
#
# SPDX-License-Identifier: Apache-2.0
import torch
import triton
import triton.language as tl
@triton.jit
def _ltx2_ada_values9_kernel(
temb_ptr,
table_ptr,
out0_ptr,
out1_ptr,
out2_ptr,
out3_ptr,
out4_ptr,
out5_ptr,
out6_ptr,
out7_ptr,
out8_ptr,
rows: tl.constexpr,
hidden: tl.constexpr,
total_params: tl.constexpr,
table_stride_p: tl.constexpr,
table_stride_d: tl.constexpr,
BLOCK_N: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
cols = tl.arange(0, BLOCK_N)
mask = cols < hidden
temb_row = temb_ptr + row * total_params * hidden
base = row * hidden + cols
table0 = tl.load(
table_ptr + 0 * table_stride_p + cols * table_stride_d,
mask=mask,
other=0.0,
).to(tl.bfloat16)
temb0 = tl.load(
temb_row + (0 * hidden + cols),
mask=mask,
other=0.0,
).to(tl.bfloat16)
table1 = tl.load(
table_ptr + 1 * table_stride_p + cols * table_stride_d,
mask=mask,
other=0.0,
).to(tl.bfloat16)
temb1 = tl.load(
temb_row + (1 * hidden + cols),
mask=mask,
other=0.0,
).to(tl.bfloat16)
table2 = tl.load(
table_ptr + 2 * table_stride_p + cols * table_stride_d,
mask=mask,
other=0.0,
).to(tl.bfloat16)
temb2 = tl.load(
temb_row + (2 * hidden + cols),
mask=mask,
other=0.0,
).to(tl.bfloat16)
table3 = tl.load(
table_ptr + 3 * table_stride_p + cols * table_stride_d,
mask=mask,
other=0.0,
).to(tl.bfloat16)
temb3 = tl.load(
temb_row + (3 * hidden + cols),
mask=mask,
other=0.0,
).to(tl.bfloat16)
table4 = tl.load(
table_ptr + 4 * table_stride_p + cols * table_stride_d,
mask=mask,
other=0.0,
).to(tl.bfloat16)
temb4 = tl.load(
temb_row + (4 * hidden + cols),
mask=mask,
other=0.0,
).to(tl.bfloat16)
table5 = tl.load(
table_ptr + 5 * table_stride_p + cols * table_stride_d,
mask=mask,
other=0.0,
).to(tl.bfloat16)
temb5 = tl.load(
temb_row + (5 * hidden + cols),
mask=mask,
other=0.0,
).to(tl.bfloat16)
table6 = tl.load(
table_ptr + 6 * table_stride_p + cols * table_stride_d,
mask=mask,
other=0.0,
).to(tl.bfloat16)
temb6 = tl.load(
temb_row + (6 * hidden + cols),
mask=mask,
other=0.0,
).to(tl.bfloat16)
table7 = tl.load(
table_ptr + 7 * table_stride_p + cols * table_stride_d,
mask=mask,
other=0.0,
).to(tl.bfloat16)
temb7 = tl.load(
temb_row + (7 * hidden + cols),
mask=mask,
other=0.0,
).to(tl.bfloat16)
table8 = tl.load(
table_ptr + 8 * table_stride_p + cols * table_stride_d,
mask=mask,
other=0.0,
).to(tl.bfloat16)
temb8 = tl.load(
temb_row + (8 * hidden + cols),
mask=mask,
other=0.0,
).to(tl.bfloat16)
tl.store(out0_ptr + base, (table0 + temb0).to(tl.bfloat16), mask=mask)
tl.store(out1_ptr + base, (table1 + temb1).to(tl.bfloat16), mask=mask)
tl.store(out2_ptr + base, (table2 + temb2).to(tl.bfloat16), mask=mask)
tl.store(out3_ptr + base, (table3 + temb3).to(tl.bfloat16), mask=mask)
tl.store(out4_ptr + base, (table4 + temb4).to(tl.bfloat16), mask=mask)
tl.store(out5_ptr + base, (table5 + temb5).to(tl.bfloat16), mask=mask)
tl.store(out6_ptr + base, (table6 + temb6).to(tl.bfloat16), mask=mask)
tl.store(out7_ptr + base, (table7 + temb7).to(tl.bfloat16), mask=mask)
tl.store(out8_ptr + base, (table8 + temb8).to(tl.bfloat16), mask=mask)
def ltx2_ada_values9(
scale_shift_table: torch.Tensor,
timestep: torch.Tensor,
) -> tuple[torch.Tensor, ...]:
if timestep.ndim != 3:
raise ValueError("timestep must have shape [B, S, 9 * D]")
if not timestep.is_cuda or timestep.dtype != torch.bfloat16:
raise ValueError("timestep must be a CUDA bfloat16 tensor")
if not timestep.is_contiguous():
raise ValueError("timestep must be contiguous")
if scale_shift_table.ndim != 2 or scale_shift_table.shape[0] != 9:
raise ValueError("scale_shift_table must have shape [9, D]")
if (
not scale_shift_table.is_cuda
or scale_shift_table.dtype not in (torch.bfloat16, torch.float32)
or scale_shift_table.stride(-1) != 1
):
raise ValueError(
"scale_shift_table must be CUDA, bf16/fp32, last-dim contiguous"
)
total_params = int(scale_shift_table.shape[0])
hidden = int(scale_shift_table.shape[1])
if hidden <= 0 or timestep.shape[-1] != total_params * hidden:
raise ValueError("timestep last dim must equal 9 * hidden")
if hidden % 256 != 0 or hidden > 8192:
raise ValueError("hidden size is outside the supported LTX2 fast-path range")
batch, seq, _ = timestep.shape
rows = int(batch * seq)
outs = tuple(
torch.empty((batch, seq, hidden), device=timestep.device, dtype=timestep.dtype)
for _ in range(9)
)
_ltx2_ada_values9_kernel[(rows,)](
timestep,
scale_shift_table,
*outs,
rows,
hidden,
total_params,
scale_shift_table.stride(0),
scale_shift_table.stride(1),
BLOCK_N=triton.next_power_of_2(hidden),
num_warps=4 if hidden >= 4096 else 8,
)
return outs
@@ -49,6 +49,8 @@ logger = init_logger(__name__)
ADALN_NUM_BASE_PARAMS = 6
ADALN_NUM_CROSS_ATTN_PARAMS = 3
_LTX2_FUSED_ADA_VALUES_RUNTIME_DISABLED = False
def adaln_embedding_coefficient(cross_attention_adaln: bool) -> int:
return ADALN_NUM_BASE_PARAMS + (
@@ -56,6 +58,48 @@ def adaln_embedding_coefficient(cross_attention_adaln: bool) -> int:
)
def _ltx2_disable_fused_ada_values(exc: Exception) -> None:
global _LTX2_FUSED_ADA_VALUES_RUNTIME_DISABLED
_LTX2_FUSED_ADA_VALUES_RUNTIME_DISABLED = True
logger.warning_once(f"Disabling LTX2 fused Ada values fast path: {exc}")
def _ltx2_try_fused_ada_values9(
scale_shift_table: torch.Tensor,
batch_size: int,
timestep: torch.Tensor,
) -> tuple[torch.Tensor, ...] | None:
if (
_LTX2_FUSED_ADA_VALUES_RUNTIME_DISABLED
or get_tp_world_size() != 1
or not timestep.is_cuda
or timestep.dtype != torch.bfloat16
or timestep.ndim != 3
or int(timestep.shape[0]) != int(batch_size)
or not timestep.is_contiguous()
or not scale_shift_table.is_cuda
or scale_shift_table.dtype not in (torch.bfloat16, torch.float32)
or scale_shift_table.ndim != 2
or int(scale_shift_table.shape[0]) != 9
or scale_shift_table.stride(-1) != 1
):
return None
hidden = int(scale_shift_table.shape[1])
if hidden % 256 != 0 or hidden > 8192 or timestep.shape[-1] != 9 * hidden:
return None
try:
from sglang.jit_kernel.diffusion.triton.ltx2_ada_values import (
ltx2_ada_values9,
)
return ltx2_ada_values9(scale_shift_table, timestep)
except Exception as exc:
_ltx2_disable_fused_ada_values(exc)
return None
def _ltx2_is_perturbed(
perturbation_config: dict[str, object],
key: str,
@@ -1037,13 +1081,21 @@ class LTX2TransformerBlock(nn.Module):
audio_replicated_for_sp: bool = False,
video_memory_prefix_len: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
batch_size = hidden_states.size(0)
video_ada_values = _ltx2_try_fused_ada_values9(
self.scale_shift_table, batch_size, temb
)
audio_ada_values = _ltx2_try_fused_ada_values9(
self.audio_scale_shift_table, batch_size, temb_audio
)
# 1. Video and Audio Self-Attention
vshift_msa, vscale_msa, vgate_msa = self.get_ada_values(
self.scale_shift_table, batch_size, temb, slice(0, 3)
)
if video_ada_values is None:
vshift_msa, vscale_msa, vgate_msa = self.get_ada_values(
self.scale_shift_table, batch_size, temb, slice(0, 3)
)
else:
vshift_msa, vscale_msa, vgate_msa = video_ada_values[0:3]
norm_hidden_states = (
self.rms_norm(hidden_states, self.norm_eps) * (1 + vscale_msa) + vshift_msa
)
@@ -1058,9 +1110,12 @@ class LTX2TransformerBlock(nn.Module):
)
hidden_states = hidden_states + attn_hidden_states * vgate_msa
ashift_msa, ascale_msa, agate_msa = self.get_ada_values(
self.audio_scale_shift_table, batch_size, temb_audio, slice(0, 3)
)
if audio_ada_values is None:
ashift_msa, ascale_msa, agate_msa = self.get_ada_values(
self.audio_scale_shift_table, batch_size, temb_audio, slice(0, 3)
)
else:
ashift_msa, ascale_msa, agate_msa = audio_ada_values[0:3]
norm_audio_hidden_states = (
self.rms_norm(audio_hidden_states, self.norm_eps) * (1 + ascale_msa)
+ ashift_msa
@@ -1081,9 +1136,12 @@ class LTX2TransformerBlock(nn.Module):
raise ValueError(
"cross_attention_adaln requires prompt modulation tensors."
)
vshift_q, vscale_q, vgate_q = self.get_ada_values(
self.scale_shift_table, batch_size, temb, slice(6, 9)
)
if video_ada_values is None:
vshift_q, vscale_q, vgate_q = self.get_ada_values(
self.scale_shift_table, batch_size, temb, slice(6, 9)
)
else:
vshift_q, vscale_q, vgate_q = video_ada_values[6:9]
v_prompt_shift, v_prompt_scale = self.get_ada_values(
self.prompt_scale_shift_table, batch_size, temb_prompt, slice(None)
)
@@ -1100,9 +1158,12 @@ class LTX2TransformerBlock(nn.Module):
)
hidden_states = hidden_states + attn_hidden_states * vgate_q
ashift_q, ascale_q, agate_q = self.get_ada_values(
self.audio_scale_shift_table, batch_size, temb_audio, slice(6, 9)
)
if audio_ada_values is None:
ashift_q, ascale_q, agate_q = self.get_ada_values(
self.audio_scale_shift_table, batch_size, temb_audio, slice(6, 9)
)
else:
ashift_q, ascale_q, agate_q = audio_ada_values[6:9]
a_prompt_shift, a_prompt_scale = self.get_ada_values(
self.audio_prompt_scale_shift_table,
batch_size,
@@ -1251,18 +1312,24 @@ class LTX2TransformerBlock(nn.Module):
audio_hidden_states + v2a_gate * v2a_attn_hidden_states
)
# 4. Feedforward
vshift_mlp, vscale_mlp, vgate_mlp = self.get_ada_values(
self.scale_shift_table, batch_size, temb, slice(3, 6)
)
if video_ada_values is None:
vshift_mlp, vscale_mlp, vgate_mlp = self.get_ada_values(
self.scale_shift_table, batch_size, temb, slice(3, 6)
)
else:
vshift_mlp, vscale_mlp, vgate_mlp = video_ada_values[3:6]
norm_hidden_states = (
self.rms_norm(hidden_states, self.norm_eps) * (1 + vscale_mlp) + vshift_mlp
)
ff_output = self.ff(norm_hidden_states)
hidden_states = hidden_states + ff_output * vgate_mlp
ashift_mlp, ascale_mlp, agate_mlp = self.get_ada_values(
self.audio_scale_shift_table, batch_size, temb_audio, slice(3, 6)
)
if audio_ada_values is None:
ashift_mlp, ascale_mlp, agate_mlp = self.get_ada_values(
self.audio_scale_shift_table, batch_size, temb_audio, slice(3, 6)
)
else:
ashift_mlp, ascale_mlp, agate_mlp = audio_ada_values[3:6]
norm_audio_hidden_states = (
self.rms_norm(audio_hidden_states, self.norm_eps) * (1 + ascale_mlp)
+ ashift_mlp
@@ -1673,7 +1740,6 @@ class LTX2VideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
late_audio_self_attention_mask: Optional[torch.Tensor] = None,
**kwargs,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
batch_size = hidden_states.size(0)
audio_timestep = audio_timestep if audio_timestep is not None else timestep