[Diffusion][Refactor] Refactor and unify RoPE execution for DiT models using RotaryEmbedding based on CustomOp (#33555)

Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
Alexandr
2026-09-11 08:31:05 +08:00
committed by GitHub
co-authored by ronnie_zheng
parent 41da06adca
commit d0035da34e
13 changed files with 763 additions and 112 deletions
@@ -330,7 +330,7 @@ class QwenImagePipelineConfig(QwenImageRolloutPipelineMixin, ImagePipelineConfig
img_cos_sin_cache = torch.cat([img_cos_half, img_sin_half], dim=-1)
txt_cos_sin_cache = torch.cat([txt_cos_half, txt_sin_half], dim=-1)
return img_cos_sin_cache, txt_cos_sin_cache
return (img_cos_sin_cache, txt_cos_sin_cache), (img_freqs, txt_freqs)
def _prepare_cond_kwargs(
self, batch, prompt_embeds, rotary_emb, device, dtype, *, negative=False
@@ -364,19 +364,24 @@ class QwenImagePipelineConfig(QwenImageRolloutPipelineMixin, ImagePipelineConfig
"img_shapes": img_shapes,
"txt_seq_lens": txt_seq_lens,
"freqs_cis": None,
"freqs_complex": None,
"encoder_hidden_states_mask": encoder_hidden_states_mask,
}
return cond_kwargs
freqs_cis = self.get_freqs_cis(
freqs_cis, freqs_complex = self.get_freqs_cis(
img_shapes, txt_seq_lens, rotary_emb, device, dtype
)
img_cache, txt_cache = freqs_cis
img_cache = shard_rotary_emb_for_sp(img_cache)
img_complex, txt_complex = freqs_complex
img_complex = shard_rotary_emb_for_sp(img_complex)
cond_kwargs = {
"txt_seq_lens": txt_seq_lens,
"freqs_cis": (img_cache, txt_cache),
"freqs_complex": (img_complex, txt_complex),
"img_shapes": img_shapes,
"encoder_hidden_states_mask": encoder_hidden_states_mask,
}
@@ -534,11 +539,12 @@ class QwenImageEditPipelineConfig(QwenImagePipelineConfig):
"img_shapes": img_shapes,
"txt_seq_lens": txt_seq_lens,
"freqs_cis": None,
"freqs_complex": None,
"encoder_hidden_states_mask": encoder_hidden_states_mask,
}
return cond_kwargs
freqs_cis = QwenImagePipelineConfig.get_freqs_cis(
freqs_cis, freqs_complex = QwenImagePipelineConfig.get_freqs_cis(
img_shapes, txt_seq_lens, rotary_emb, device, dtype
)
@@ -550,9 +556,13 @@ class QwenImageEditPipelineConfig(QwenImagePipelineConfig):
img_cache, txt_cache = _shard_qwen_edit_freqs_cis_for_sp(
freqs_cis, noisy_img_seq_len, device
)
img_complex, txt_complex = _shard_qwen_edit_freqs_cis_for_sp(
freqs_complex, noisy_img_seq_len, device
)
cond_kwargs = {
"txt_seq_lens": txt_seq_lens,
"freqs_cis": (img_cache, txt_cache),
"freqs_complex": (img_complex, txt_complex),
"img_shapes": img_shapes,
"encoder_hidden_states_mask": encoder_hidden_states_mask,
}
@@ -736,7 +746,7 @@ class QwenImageEditPlusPipelineConfig(QwenImageEditPipelineConfig):
batch, 0, text_seq_len, batch_size, negative=negative
)
freqs_cis = QwenImageEditPlusPipelineConfig.get_freqs_cis(
freqs_cis, freqs_complex = QwenImageEditPlusPipelineConfig.get_freqs_cis(
img_shapes, txt_seq_lens, rotary_emb, device, dtype
)
@@ -750,6 +760,9 @@ class QwenImageEditPlusPipelineConfig(QwenImageEditPipelineConfig):
"freqs_cis": _shard_qwen_edit_freqs_cis_for_sp(
freqs_cis, noisy_img_seq_len, device
),
"freqs_complex": _shard_qwen_edit_freqs_cis_for_sp(
freqs_complex, noisy_img_seq_len, device
),
"img_shapes": img_shapes,
"encoder_hidden_states_mask": encoder_hidden_states_mask,
}
@@ -800,7 +813,7 @@ class QwenImageLayeredPipelineConfig(QwenImageEditPipelineConfig):
batch, 0, text_seq_len, batch_size, negative=negative
)
freqs_cis = QwenImageEditPlusPipelineConfig.get_freqs_cis(
freqs_cis, freqs_complex = QwenImageEditPlusPipelineConfig.get_freqs_cis(
img_shapes, txt_seq_lens, rotary_emb, device, dtype
)
@@ -815,10 +828,17 @@ class QwenImageLayeredPipelineConfig(QwenImageEditPipelineConfig):
[noisy_img_cache, img_cache[noisy_img_seq_len:, :]], dim=0
).to(device=device)
img_complex, txt_complex = freqs_complex
noisy_img_complex = shard_rotary_emb_for_sp(img_complex[:noisy_img_seq_len, :])
img_complex = torch.cat(
[noisy_img_complex, img_complex[noisy_img_seq_len:, :]], dim=0
).to(device=device)
cond_kwargs = {
"txt_seq_lens": txt_seq_lens,
"img_shapes": img_shapes,
"freqs_cis": (img_cache, txt_cache),
"freqs_complex": (img_complex, txt_complex),
"additional_t_cond": torch.tensor([0], device=device, dtype=torch.long),
"encoder_hidden_states_mask": encoder_hidden_states_mask,
}
@@ -27,6 +27,9 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_tp_group,
)
from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
RotaryEmbedding,
)
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.platforms.aiter import USE_AITER
from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var
@@ -75,6 +78,8 @@ if _is_xpu:
if not _is_cpu:
from sglang.kernels.ops.diffusion import norm_infer, rms_norm_fn
_QK_NORM_ROPE_DICT: dict[tuple[int, bool], RotaryEmbedding] = {}
# Copied and adapted from sglang
@CustomOp.register("rms_norm")
@@ -925,6 +930,7 @@ def apply_qk_norm_with_optional_rope(
k_norm: "RMSNorm",
head_dim: int,
cos_sin_cache: Optional[torch.Tensor] = None,
freqs_complex: Optional[torch.Tensor] = None,
*,
is_neox: bool = False,
positions: Optional[torch.Tensor] = None,
@@ -950,6 +956,7 @@ def apply_qk_norm_with_optional_rope(
k_norm=k_norm,
head_dim=head_dim,
cos_sin_cache=cos_sin_cache,
freqs_complex=freqs_complex,
is_neox=is_neox,
positions=positions,
position_offset=position_offset,
@@ -965,6 +972,7 @@ def apply_qk_norm_rope(
head_dim: int,
cos_sin_cache: torch.Tensor,
*,
freqs_complex: Optional[torch.Tensor] = None,
is_neox: bool = False,
positions: Optional[torch.Tensor] = None,
position_offset: int = 0,
@@ -981,10 +989,6 @@ def apply_qk_norm_rope(
requires the fused CUDA path; the ordinary cache stores half-width cos/sin.
"""
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
apply_flashinfer_rope_qk_inplace,
)
if q.dim() != 4 or k.dim() != 4:
raise ValueError(
f"apply_qk_norm_rope expects 4D q/k tensors, got q:{tuple(q.shape)} k:{tuple(k.shape)}"
@@ -1121,13 +1125,25 @@ def apply_qk_norm_rope(
head_dim=head_dim,
allow_inplace=allow_inplace,
)
return apply_flashinfer_rope_qk_inplace(
q=q,
k=k,
cos_sin_cache=cos_sin_cache,
head_size=head_dim,
is_neox=is_neox,
rope_key = (head_dim, is_neox)
rotary_emb = _QK_NORM_ROPE_DICT.get(rope_key)
if rotary_emb is None:
rotary_emb = RotaryEmbedding(
head_size=head_dim,
rotary_dim=head_dim,
use_precomputed_cache=False,
is_neox_style=is_neox,
)
_QK_NORM_ROPE_DICT[rope_key] = rotary_emb
return rotary_emb(
query=q,
key=k,
positions=positions,
complex_freqs=(
freqs_complex.unsqueeze(-2) if freqs_complex is not None else None
),
cos_sin_cache=cos_sin_cache,
)
@@ -1,10 +1,25 @@
"""RotaryEmbedding base class and LinearScalingRotaryEmbedding variant."""
from typing import Optional, Tuple
import torch
from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp
from sglang.multimodal_gen.runtime.platforms import current_platform
from .utils import _apply_rotary_emb
from .utils import (
_apply_rotary_emb,
_apply_rotary_emb_complex,
apply_flashinfer_rope_qk_inplace,
)
if current_platform.is_npu():
import torch_npu
from sglang.kernels.ops.diffusion.common.fallback_npu import (
NPU_ROTARY_MUL_MAX_HEAD_SIZE,
NPU_ROTARY_MUL_MAX_NUM_HEADS,
)
@CustomOp.register("rotary_embedding")
@@ -15,10 +30,12 @@ class RotaryEmbedding(CustomOp):
self,
head_size: int,
rotary_dim: int,
max_position_embeddings: int,
base: int | float,
is_neox_style: bool,
dtype: torch.dtype,
max_position_embeddings: Optional[int] = 4096,
base: Optional[int | float] = 10000,
is_neox_style: bool = False,
dtype: Optional[torch.dtype] = torch.float16,
use_precomputed_cache: Optional[bool] = True,
complex_dtype: torch.dtype = torch.float32,
) -> None:
super().__init__()
self.head_size = head_size
@@ -27,11 +44,21 @@ class RotaryEmbedding(CustomOp):
self.base = base
self.is_neox_style = is_neox_style
self.dtype = dtype
self.use_precomputed_cache = use_precomputed_cache
self._complex_dtype = complex_dtype
self._is_full_rotation = rotary_dim == head_size
self._is_complex_style = not is_neox_style
self._is_npu_rotary_mul = (
current_platform.is_npu()
and is_neox_style
and rotary_dim < NPU_ROTARY_MUL_MAX_HEAD_SIZE
)
cache = self._compute_cos_sin_cache()
cache = cache.to(dtype)
self.cos_sin_cache: torch.Tensor
self.register_buffer("cos_sin_cache", cache, persistent=False)
if self.use_precomputed_cache:
cache = self._compute_cos_sin_cache()
cache = cache.to(dtype)
self.cos_sin_cache: torch.Tensor
self.register_buffer("cos_sin_cache", cache, persistent=False)
def _compute_inv_freq(self, base: int | float) -> torch.Tensor:
"""Compute the inverse frequency."""
@@ -58,41 +85,409 @@ class RotaryEmbedding(CustomOp):
cache = torch.cat((cos, sin), dim=-1)
return cache
def forward_cuda(self, *args, **kwargs):
return self.forward_native(*args, **kwargs)
def _combine_rotated_and_pass(
self,
q_rotated: torch.Tensor,
k_rotated: torch.Tensor,
q_pass: torch.Tensor,
k_pass: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Reattach the untouched tail (rotary_dim < head_size), if any.
torch.cat against an empty q_pass/k_pass (rotary_dim == head_size)
still allocates + copies, so skip it in that common case.
"""
if self._is_full_rotation:
return q_rotated, k_rotated
return (
torch.cat((q_rotated, q_pass), dim=-1),
torch.cat((k_rotated, k_pass), dim=-1),
)
def forward_npu(
self,
positions: Optional[torch.Tensor] = None,
query: Optional[torch.Tensor] = None,
key: Optional[torch.Tensor] = None,
position_offset: int = 0,
cos: Optional[torch.Tensor] = None,
sin: Optional[torch.Tensor] = None,
complex_freqs: Optional[torch.Tensor] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
offsets: Optional[torch.Tensor] = None,
**kwargs,
) -> Tuple[torch.Tensor, torch.Tensor]:
if self.use_precomputed_cache or query.dim() == 3 or key.dim() == 3:
return self.forward_native(
query=query,
key=key,
positions=positions,
position_offset=position_offset,
cos=cos,
sin=sin,
complex_freqs=complex_freqs,
cos_sin_cache=cos_sin_cache,
offsets=offsets,
**kwargs,
)
if query.dim() != 4 or key.dim() != 4:
raise ValueError(
f"query and key must be [batch_size, seq_len, num_heads, head_dim],"
f"got query: {tuple(query.shape)}, key: {tuple(key.shape)}"
)
seq_len = query.shape[1]
support_complex_style = (
complex_freqs is not None
and complex_freqs.dim() == 3
and self._is_complex_style
and self._is_full_rotation
)
if support_complex_style:
return (
_apply_rotary_emb_complex(
query, complex_freqs, dtype=self._complex_dtype
),
_apply_rotary_emb_complex(
key, complex_freqs, dtype=self._complex_dtype
),
)
is_complex_derivable = (
complex_freqs is None
and cos is not None
and sin is not None
and self._is_complex_style
and self._is_full_rotation
and cos.shape[0] == seq_len
)
if is_complex_derivable:
# No fused kernel for interleaved rotation here; complex-multiply
# is equivalent and needs fewer kernel launches.
derived_complex_freqs = torch.complex(
cos.to(torch.float32), sin.to(torch.float32)
).unsqueeze(-2)
return (
_apply_rotary_emb_complex(
query, derived_complex_freqs, dtype=self._complex_dtype
),
_apply_rotary_emb_complex(
key, derived_complex_freqs, dtype=self._complex_dtype
),
)
if cos is not None and sin is not None:
num_heads = query.shape[2]
support_npu_rotary_mul = (
self._is_npu_rotary_mul
and cos.shape[0] == seq_len
and num_heads < NPU_ROTARY_MUL_MAX_NUM_HEADS
)
if support_npu_rotary_mul:
# Called directly on the BSND [batch, seq, heads, rotary_dim]
# layout (no batch*seq flatten): cos/sin get a batch dim and
# a heads dim of 1 and broadcast against query/key (the
# documented "1S1D" pattern), avoiding a per-call
# expand+copy of cos/sin across the batch. The size gate
# above mirrors apply_rotary_embedding_native's own gate —
# if that gate disagreed, that function would silently take
# its always-interleaved fallback, which is wrong here.
q_rot = query[..., : self.rotary_dim]
q_pass = query[..., self.rotary_dim :]
k_rot = key[..., : self.rotary_dim]
k_pass = key[..., self.rotary_dim :]
cos_prepared = cos.reshape(1, seq_len, 1, -1).to(query.dtype)
sin_prepared = sin.reshape(1, seq_len, 1, -1).to(query.dtype)
if cos_prepared.size(-1) * 2 == self.rotary_dim:
cos_prepared = torch.cat((cos_prepared, cos_prepared), dim=-1)
sin_prepared = torch.cat((sin_prepared, sin_prepared), dim=-1)
q_rotated = torch_npu.npu_rotary_mul(q_rot, cos_prepared, sin_prepared)
k_rotated = torch_npu.npu_rotary_mul(k_rot, cos_prepared, sin_prepared)
return self._combine_rotated_and_pass(
q_rotated, k_rotated, q_pass, k_pass
)
# No [batch*seq, ...] flatten: cos/sin are [seq_len,
# rotary_dim // 2], shared across the batch, and only
# broadcast correctly this way for batch_size > 1.
q_rot = query[..., : self.rotary_dim]
q_pass = query[..., self.rotary_dim :]
k_rot = key[..., : self.rotary_dim]
k_pass = key[..., self.rotary_dim :]
q_rotated = _apply_rotary_emb(
q_rot,
cos,
sin,
is_neox_style=self.is_neox_style,
interleaved=not self.is_neox_style,
)
k_rotated = _apply_rotary_emb(
k_rot,
cos,
sin,
is_neox_style=self.is_neox_style,
interleaved=not self.is_neox_style,
)
return self._combine_rotated_and_pass(q_rotated, k_rotated, q_pass, k_pass)
if cos_sin_cache is not None:
return self.forward_native(
query=query,
key=key,
positions=positions,
position_offset=position_offset,
cos=cos,
sin=sin,
complex_freqs=complex_freqs,
cos_sin_cache=cos_sin_cache,
offsets=offsets,
**kwargs,
)
raise ValueError(
"No valid inputs (complex_freqs, cos/sin, or cos_sin_cache) for interleaved RoPE."
)
def forward_cuda(
self,
positions: Optional[torch.Tensor] = None,
query: Optional[torch.Tensor] = None,
key: Optional[torch.Tensor] = None,
position_offset: int = 0,
cos: Optional[torch.Tensor] = None,
sin: Optional[torch.Tensor] = None,
complex_freqs: Optional[torch.Tensor] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
offsets: Optional[torch.Tensor] = None,
**kwargs,
) -> Tuple[torch.Tensor, torch.Tensor]:
support_cuda_style = (
(cos_sin_cache is not None or cos is not None and sin is not None)
and not self.use_precomputed_cache
and query.dim() == 4
and key.dim() == 4
)
if not support_cuda_style:
return self.forward_native(
query=query,
key=key,
positions=positions,
position_offset=position_offset,
cos=cos,
sin=sin,
complex_freqs=complex_freqs,
cos_sin_cache=cos_sin_cache,
offsets=offsets,
**kwargs,
)
if cos_sin_cache is None:
cos_sin_cache = torch.cat(
[
cos.to(dtype=torch.float32).contiguous(),
sin.to(dtype=torch.float32).contiguous(),
],
dim=-1,
)
batch_size, seq_len, _, head_dim = query.shape
if positions is None:
pos_1d = torch.arange(
position_offset,
position_offset + seq_len,
device=query.device,
dtype=torch.int64,
)
positions = pos_1d if batch_size == 1 else pos_1d.repeat(batch_size)
else:
positions = positions.to(device=query.device, dtype=torch.long)
return apply_flashinfer_rope_qk_inplace(
q=query,
k=key,
cos_sin_cache=cos_sin_cache,
head_size=head_dim,
is_neox=self.is_neox_style,
positions=positions,
)
def forward_xpu(self, *args, **kwargs):
return self.forward_native(*args, **kwargs)
def forward_native(
self,
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
offsets: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
positions: Optional[torch.Tensor] = None,
query: Optional[torch.Tensor] = None,
key: Optional[torch.Tensor] = None,
position_offset: int = 0,
cos: Optional[torch.Tensor] = None,
sin: Optional[torch.Tensor] = None,
complex_freqs: Optional[torch.Tensor] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
offsets: Optional[torch.Tensor] = None,
**kwargs,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""A PyTorch-native implementation of forward()."""
if offsets is not None:
positions = positions + offsets
positions = positions.flatten()
num_tokens = positions.shape[0]
cos_sin = self.cos_sin_cache.index_select(0, positions)
cos, sin = cos_sin.chunk(2, dim=-1)
query_shape = query.shape
query = query.reshape(num_tokens, -1, self.head_size)
query_rot = query[..., : self.rotary_dim]
query_pass = query[..., self.rotary_dim :]
query_rot = _apply_rotary_emb(query_rot, cos, sin, self.is_neox_style)
query = torch.cat((query_rot, query_pass), dim=-1).reshape(query_shape)
use_precomputed_cache = self.use_precomputed_cache
if use_precomputed_cache:
if offsets is not None:
positions = positions + offsets
positions = positions.flatten()
num_tokens = positions.shape[0]
cos_sin = self.cos_sin_cache.index_select(0, positions)
cos, sin = cos_sin.chunk(2, dim=-1)
key_shape = key.shape
key = key.reshape(num_tokens, -1, self.head_size)
key_rot = key[..., : self.rotary_dim]
key_pass = key[..., self.rotary_dim :]
key_rot = _apply_rotary_emb(key_rot, cos, sin, self.is_neox_style)
key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape)
return query, key
is_complex_derivable = (
not use_precomputed_cache
and complex_freqs is None
and cos is not None
and sin is not None
and self._is_complex_style
and query.dim() == 4
and key.dim() == 4
and self._is_full_rotation
and cos.shape[0] == query.shape[1]
)
if is_complex_derivable:
# No fused kernel for interleaved rotation on native backends;
# complex-multiply is equivalent and needs fewer kernel
# launches. CUDA has its own fused path (forward_cuda) and
# never reaches here for this case.
derived_complex_freqs = torch.complex(
cos.to(torch.float32), sin.to(torch.float32)
).unsqueeze(-2)
return (
_apply_rotary_emb_complex(
query, derived_complex_freqs, dtype=self._complex_dtype
),
_apply_rotary_emb_complex(
key, derived_complex_freqs, dtype=self._complex_dtype
),
)
if cos is not None and sin is not None:
if use_precomputed_cache:
# Legacy callers (llama/qwen3/gemma2/gemma3 via get_rope())
# pass 3D [batch, seq, hidden]; cos/sin were already
# index_select'd per-token above, so they already match
# num_tokens = batch*seq row-for-row.
q_shape = query.shape
q_flat = query.reshape(num_tokens, -1, self.head_size)
q_rot = q_flat[..., : self.rotary_dim]
q_pass = q_flat[..., self.rotary_dim :]
k_shape = key.shape
k_flat = key.reshape(num_tokens, -1, self.head_size)
k_rot = k_flat[..., : self.rotary_dim]
k_pass = k_flat[..., self.rotary_dim :]
q_rotated = _apply_rotary_emb(
q_rot,
cos,
sin,
is_neox_style=self.is_neox_style,
interleaved=not self.is_neox_style,
)
k_rotated = _apply_rotary_emb(
k_rot,
cos,
sin,
is_neox_style=self.is_neox_style,
interleaved=not self.is_neox_style,
)
q, k = self._combine_rotated_and_pass(
q_rotated, k_rotated, q_pass, k_pass
)
return q.reshape(q_shape), k.reshape(k_shape)
# Direct DiT-style call: same batch/seq broadcast reasoning as
# forward_npu's cos/sin path.
q_rot = query[..., : self.rotary_dim]
q_pass = query[..., self.rotary_dim :]
k_rot = key[..., : self.rotary_dim]
k_pass = key[..., self.rotary_dim :]
q_rotated = _apply_rotary_emb(
q_rot,
cos,
sin,
is_neox_style=self.is_neox_style,
interleaved=not self.is_neox_style,
)
k_rotated = _apply_rotary_emb(
k_rot,
cos,
sin,
is_neox_style=self.is_neox_style,
interleaved=not self.is_neox_style,
)
return self._combine_rotated_and_pass(q_rotated, k_rotated, q_pass, k_pass)
if query.dim() != 4 or key.dim() != 4:
raise ValueError(
f"query and key must be [batch_size, seq_len, num_heads, head_dim],"
f"got query: {tuple(query.shape)}, key: {tuple(key.shape)}"
)
support_complex_style = (
complex_freqs is not None
and complex_freqs.dim() == 3
and self._is_complex_style
and self._is_full_rotation
)
if support_complex_style:
return (
_apply_rotary_emb_complex(
query, complex_freqs, dtype=self._complex_dtype
),
_apply_rotary_emb_complex(
key, complex_freqs, dtype=self._complex_dtype
),
)
if cos_sin_cache is not None:
batch_size, seq_len, _, _ = query.shape
num_tokens = batch_size * seq_len
if positions is None:
pos_1d = torch.arange(
position_offset,
position_offset + seq_len,
device=query.device,
dtype=torch.int64,
)
positions = pos_1d if batch_size == 1 else pos_1d.repeat(batch_size)
else:
if positions.dim() != 1 or positions.numel() != num_tokens:
raise ValueError(
f"positions must be 1D of length {num_tokens}, got shape={tuple(positions.shape)}"
)
positions = positions.to(device=query.device, dtype=torch.long)
return apply_flashinfer_rope_qk_inplace(
q=query,
k=key,
cos_sin_cache=cos_sin_cache,
head_size=self.head_size,
is_neox=self.is_neox_style,
positions=positions,
)
raise ValueError(
"No valid inputs (complex_freqs, cos/sin, or cos_sin_cache) for RoPE."
)
def extra_repr(self) -> str:
s = f"head_size={self.head_size}, rotary_dim={self.rotary_dim}"
@@ -68,6 +68,7 @@ def _apply_rotary_emb(
def _apply_rotary_emb_complex(
x: torch.Tensor, # [b, s, h, d]
freqs: torch.Tensor, # [s, 1, d // 2]
dtype: torch.dtype = torch.float32,
) -> torch.Tensor: # [b, s, h, d]
"""
Apply complex rotary positional embeddings designed for interleaved=True, neox_style=False.
@@ -77,16 +78,16 @@ def _apply_rotary_emb_complex(
Args:
x: Input activation tensor in bf16/fp16.
Shape: [batch, num_tokens, num_heads, head_size]
freqs: Complex-valued frequency tensor in complex64 format.
freqs: Complex-valued frequency tensor, real/imag parts in `dtype`.
Shape: [num_tokens, 1, head_size // 2]
dtype: Intermediate real dtype for the complex multiply.
Returns:
torch.Tensor: The same shape and dtype as x.
"""
b, s, h, d = x.shape
dtype_c = torch.float64
x_complex = torch.view_as_complex(x.to(dtype_c).reshape(b, s, h, d // 2, 2))
x_complex = torch.view_as_complex(x.to(dtype).reshape(b, s, h, d // 2, 2))
x_out = torch.view_as_real(x_complex * freqs)
x_out = x_out.view(b, s, h, d)
return x_out.to(x.dtype)
@@ -247,6 +247,23 @@ def _rope_cos_sin_cache(
)
def _rope_complex_freqs(
freqs_cis: Union[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, None],
) -> Optional[torch.Tensor]:
"""Complex-valued sibling of ``_rope_cos_sin_cache``: build [seq, dim//2]
complex64 freqs for the (is_neox=False) NPU fast path in
apply_qk_norm_with_optional_rope. Accepts the same inputs as
_rope_cos_sin_cache — a raw (cos, sin) tuple, or its already-hoisted
cat([cos, sin], dim=-1) cache tensor, split back in half."""
if freqs_cis is None:
return None
if isinstance(freqs_cis, torch.Tensor):
cos, sin = freqs_cis.chunk(2, dim=-1)
else:
cos, sin = freqs_cis
return torch.complex(cos.to(torch.float32), sin.to(torch.float32))
try:
from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import]
from nunchaku.models.normalization import ( # type: ignore[import]
@@ -633,6 +650,7 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
x: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
freqs_cis=None,
complex_freqs: Optional[torch.Tensor] = None,
num_replicated_prefix: int = 0,
attn_mask: Optional[torch.Tensor] = None,
attn_mask_meta: Optional[Dict[str, int]] = None,
@@ -659,6 +677,13 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
encoder_value = encoder_value.unflatten(-1, (num_heads, -1))
text_seq_len = encoder_query.shape[1]
# complex_freqs covers [text, image] positions in order (same
# table cos_sin_cache/positions index into); slice per call the
# same way position_offset selects rows below — the class's
# complex_freqs path does not do positional indexing itself.
text_freqs_complex = (
complex_freqs[:text_seq_len] if complex_freqs is not None else None
)
encoder_query, encoder_key = apply_qk_norm_with_optional_rope(
q=encoder_query,
k=encoder_key,
@@ -666,9 +691,16 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
k_norm=self.norm_added_k,
head_dim=self.head_dim,
cos_sin_cache=cos_sin_cache,
freqs_complex=text_freqs_complex,
is_neox=False,
allow_inplace=True,
)
img_seq_len = query.shape[1]
img_freqs_complex = (
complex_freqs[text_seq_len : text_seq_len + img_seq_len]
if complex_freqs is not None
else None
)
query, key = apply_qk_norm_with_optional_rope(
q=query,
k=key,
@@ -676,6 +708,7 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
k_norm=self.norm_k,
head_dim=self.head_dim,
cos_sin_cache=cos_sin_cache,
freqs_complex=img_freqs_complex,
is_neox=False,
position_offset=text_seq_len,
allow_inplace=True,
@@ -688,6 +721,10 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
key = join_seqs(encoder_key, key, sp_txt_pad)
value = join_seqs(encoder_value, value, sp_txt_pad)
else:
seq_len = query.shape[1]
joint_freqs_complex = (
complex_freqs[:seq_len] if complex_freqs is not None else None
)
query, key = apply_qk_norm_with_optional_rope(
q=query,
k=key,
@@ -695,6 +732,7 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
k_norm=self.norm_k,
head_dim=self.head_dim,
cos_sin_cache=cos_sin_cache,
freqs_complex=joint_freqs_complex,
is_neox=False,
allow_inplace=True,
)
@@ -854,6 +892,7 @@ class FluxSingleTransformerBlock(nn.Module):
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: Union[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, None] = None,
complex_freqs: Optional[torch.Tensor] = None,
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
num_replicated_prefix: int = 0,
) -> Tuple[torch.Tensor, torch.Tensor]:
@@ -882,6 +921,7 @@ class FluxSingleTransformerBlock(nn.Module):
attn_output = self.attn(
x=norm_hidden_states,
freqs_cis=freqs_cis,
complex_freqs=complex_freqs,
num_replicated_prefix=num_replicated_prefix,
**joint_attention_kwargs,
)
@@ -906,6 +946,7 @@ class FluxSingleTransformerBlock(nn.Module):
attn_output = self.attn(
x=norm_hidden_states,
freqs_cis=freqs_cis,
complex_freqs=complex_freqs,
num_replicated_prefix=num_replicated_prefix,
**joint_attention_kwargs,
)
@@ -1020,6 +1061,7 @@ class FluxTransformerBlock(nn.Module):
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: Union[Tuple[torch.Tensor, torch.Tensor], torch.Tensor, None] = None,
complex_freqs: Optional[torch.Tensor] = None,
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
num_replicated_prefix: int = 0,
) -> Tuple[torch.Tensor, torch.Tensor]:
@@ -1041,6 +1083,7 @@ class FluxTransformerBlock(nn.Module):
x=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
freqs_cis=freqs_cis,
complex_freqs=complex_freqs,
num_replicated_prefix=num_replicated_prefix,
**joint_attention_kwargs,
)
@@ -1326,8 +1369,14 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
join_seqs(sin[:t_loc], sin[t_loc:], pad, dim=0),
)
# Build the RoPE cos/sin cache once per step; every attention call
# below reuses the same tensor.
# Build the RoPE cos/sin cache and complex_freqs once per step; every
# attention call below reuses the same tensors.
complex_freqs = _rope_complex_freqs(freqs_cis)
singles_complex_freqs = (
complex_freqs
if singles_freqs_cis is freqs_cis
else _rope_complex_freqs(singles_freqs_cis)
)
hoisted_freqs_cis = _rope_cos_sin_cache(freqs_cis)
singles_freqs_cis = (
hoisted_freqs_cis
@@ -1358,6 +1407,7 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
encoder_hidden_states=encoder_hidden_states,
temb=temb,
freqs_cis=freqs_cis,
complex_freqs=complex_freqs,
joint_attention_kwargs=joint_attention_kwargs,
num_replicated_prefix=num_replicated_prefix,
)
@@ -1367,6 +1417,7 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
encoder_hidden_states=encoder_hidden_states,
temb=temb,
freqs_cis=singles_freqs_cis,
complex_freqs=singles_complex_freqs,
joint_attention_kwargs=joint_attention_kwargs,
num_replicated_prefix=num_replicated_prefix,
)
@@ -202,6 +202,33 @@ def _defer_gated_residual(
return residual_gate_add(residual, update, gate)
def _flux2_derive_rope_tensors(
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]],
) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]:
"""(cos_sin_cache, complex_freqs) from one (cos, sin) pair.
Called once per Flux2Transformer2DModel.forward() instead of once per
block: freqs_cis is identical across every block in a forward pass, so
deriving it per-attention-call recomputed the same tensors up to 56x
per denoising step.
"""
if freqs_cis is None:
return None, None
cos, sin = freqs_cis
cos_sin_cache = torch.cat(
[
cos.to(dtype=torch.float32).contiguous(),
sin.to(dtype=torch.float32).contiguous(),
],
dim=-1,
)
# is_neox=False here, so this can hit the NPU _apply_rotary_emb_complex
# fast path in RotaryEmbedding instead of the interleaved fallback (no
# fused NPU kernel for it).
complex_freqs = torch.complex(cos.to(torch.float32), sin.to(torch.float32))
return cos_sin_cache, complex_freqs
def _flux2_gated_resnorm(
norm: nn.Module,
residual: torch.Tensor,
@@ -611,7 +638,8 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
self,
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
complex_freqs: Optional[torch.Tensor] = None,
num_replicated_prefix: int = 0,
attn_mask: Optional[torch.Tensor] = None,
attn_mask_meta: Optional[Dict[str, int]] = None,
@@ -634,17 +662,6 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
key = key.unflatten(-1, (self.local_heads, -1))
value = value.unflatten(-1, (self.local_heads, -1))
cos_sin_cache = None
if freqs_cis is not None:
cos, sin = freqs_cis
cos_sin_cache = torch.cat(
[
cos.to(dtype=torch.float32).contiguous(),
sin.to(dtype=torch.float32).contiguous(),
],
dim=-1,
)
joint_qkv = None
sp_txt_pad = 0
if self.added_kv_proj_dim is not None:
@@ -686,6 +703,13 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
tensor.contiguous()
for tensor in (encoder_query, encoder_key, encoder_value)
]
# complex_freqs covers [text, image] positions in order (same
# table cos_sin_cache/positions index into); slice per call the
# same way position_offset selects rows below — the class's
# complex_freqs path does not do positional indexing itself.
text_freqs_complex = (
complex_freqs[:text_seq_len] if complex_freqs is not None else None
)
encoder_query, encoder_key = apply_qk_norm_with_optional_rope(
q=encoder_query,
k=encoder_key,
@@ -693,9 +717,16 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
k_norm=self.norm_added_k,
head_dim=self.head_dim,
cos_sin_cache=cos_sin_cache,
freqs_complex=text_freqs_complex,
is_neox=False,
allow_inplace=True,
)
img_seq_len = query.shape[1]
img_freqs_complex = (
complex_freqs[text_seq_len : text_seq_len + img_seq_len]
if complex_freqs is not None
else None
)
query, key = apply_qk_norm_with_optional_rope(
q=query,
k=key,
@@ -703,6 +734,7 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
k_norm=self.norm_k,
head_dim=self.head_dim,
cos_sin_cache=cos_sin_cache,
freqs_complex=img_freqs_complex,
is_neox=False,
position_offset=text_seq_len,
allow_inplace=True,
@@ -717,6 +749,10 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
query, key, value = [
tensor.contiguous() for tensor in (query, key, value)
]
seq_len = query.shape[1]
joint_freqs_complex = (
complex_freqs[:seq_len] if complex_freqs is not None else None
)
query, key = apply_qk_norm_with_optional_rope(
q=query,
k=key,
@@ -724,6 +760,7 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
k_norm=self.norm_k,
head_dim=self.head_dim,
cos_sin_cache=cos_sin_cache,
freqs_complex=joint_freqs_complex,
is_neox=False,
allow_inplace=True,
)
@@ -880,7 +917,8 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
self,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
complex_freqs: Optional[torch.Tensor] = None,
num_replicated_prefix: int = 0,
**kwargs,
) -> torch.Tensor:
@@ -907,16 +945,8 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
key = key.unflatten(-1, (self.local_heads, -1))
value = value.unflatten(-1, (self.local_heads, -1))
cos_sin_cache = None
if freqs_cis is not None:
cos, sin = freqs_cis
cos_sin_cache = torch.cat(
[
cos.to(dtype=torch.float32).contiguous(),
sin.to(dtype=torch.float32).contiguous(),
],
dim=-1,
)
if complex_freqs is not None:
complex_freqs = complex_freqs[: query.shape[1]]
# QK-norm (+ RoPE) via the shared helper so the fused kernel path is used
# here too — the single-stream block previously ran norm and RoPE as separate ops.
@@ -927,6 +957,7 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
k_norm=self.norm_k,
head_dim=self.head_dim,
cos_sin_cache=cos_sin_cache,
freqs_complex=complex_freqs,
is_neox=False,
allow_inplace=True,
)
@@ -1019,7 +1050,8 @@ class Flux2SingleTransformerBlock(nn.Module):
hidden_states: torch.Tensor | PendingGatedResidual,
encoder_hidden_states: Optional[torch.Tensor],
temb_mod_params: Tuple[torch.Tensor, torch.Tensor, torch.Tensor],
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
complex_freqs: Optional[torch.Tensor] = None,
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
split_hidden_states: bool = False,
text_seq_len: Optional[int] = None,
@@ -1046,7 +1078,8 @@ class Flux2SingleTransformerBlock(nn.Module):
joint_attention_kwargs = joint_attention_kwargs or {}
attn_output = self.attn(
hidden_states=norm_hidden_states,
freqs_cis=freqs_cis,
cos_sin_cache=cos_sin_cache,
complex_freqs=complex_freqs,
num_replicated_prefix=num_replicated_prefix,
**joint_attention_kwargs,
)
@@ -1164,7 +1197,8 @@ class Flux2TransformerBlock(nn.Module):
temb_mod_params_txt: Tuple[
Tuple[torch.Tensor, torch.Tensor, torch.Tensor], ...
],
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
complex_freqs: Optional[torch.Tensor] = None,
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
num_replicated_prefix: int = 0,
) -> Tuple[
@@ -1230,7 +1264,8 @@ class Flux2TransformerBlock(nn.Module):
attention_outputs = self.attn(
hidden_states=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
freqs_cis=freqs_cis,
cos_sin_cache=cos_sin_cache,
complex_freqs=complex_freqs,
num_replicated_prefix=num_replicated_prefix,
**joint_attention_kwargs,
)
@@ -1676,6 +1711,14 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
join_seqs(sin[:t_loc], sin[t_loc:], sp_txt_pad, dim=0),
)
# freqs_cis/singles_freqs_cis are fixed for the rest of this forward
# pass, so derive cos_sin_cache/complex_freqs once here instead of
# once per block (56x per full denoising step).
cos_sin_cache, complex_freqs = _flux2_derive_rope_tensors(freqs_cis)
singles_cos_sin_cache, singles_complex_freqs = _flux2_derive_rope_tensors(
singles_freqs_cis
)
# 4. Double Stream Transformer Blocks
for index_block, block in enumerate(self.transformer_blocks):
encoder_hidden_states, hidden_states = block(
@@ -1683,7 +1726,8 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
encoder_hidden_states=encoder_hidden_states,
temb_mod_params_img=double_stream_mod_img,
temb_mod_params_txt=double_stream_mod_txt,
freqs_cis=freqs_cis,
cos_sin_cache=cos_sin_cache,
complex_freqs=complex_freqs,
joint_attention_kwargs=joint_attention_kwargs,
num_replicated_prefix=num_replicated_prefix,
)
@@ -1703,7 +1747,8 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
hidden_states=hidden_states,
encoder_hidden_states=None,
temb_mod_params=single_stream_mod,
freqs_cis=singles_freqs_cis,
cos_sin_cache=singles_cos_sin_cache,
complex_freqs=singles_complex_freqs,
joint_attention_kwargs=joint_attention_kwargs,
text_seq_len=txt_real,
num_replicated_prefix=num_replicated_prefix,
@@ -52,8 +52,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor
QuantizationConfig,
)
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
_apply_rotary_emb,
apply_flashinfer_rope_qk_inplace,
RotaryEmbedding,
)
from sglang.multimodal_gen.runtime.layers.visual_embedding import Timesteps
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
@@ -590,6 +589,12 @@ class GlmImageAttention(torch.nn.Module):
raise ValueError(
f"unknown qk_norm: {qk_norm}. Should be one of None, 'layer_norm', 'fp32_layer_norm', 'layer_norm_across_heads', 'rms_norm', 'rms_norm_across_heads', 'l2'."
)
self.rotary_emb = RotaryEmbedding(
head_size=dim_head,
rotary_dim=dim_head,
use_precomputed_cache=False,
is_neox_style=True,
)
self.attn = USPAttention(
num_heads=self.num_local_heads,
@@ -635,28 +640,17 @@ class GlmImageAttention(torch.nn.Module):
# 3. Rotational positional embeddings applied to latent stream
if image_rotary_emb is not None:
cos, sin = image_rotary_emb
q_img = query[:, text_seq_length:, :, :]
k_img = key[:, text_seq_length:, :, :]
if _is_cuda and cos.dim() == 2:
q_img = query[:, text_seq_length:, :, :]
k_img = key[:, text_seq_length:, :, :]
cos_sin_cache = torch.cat(
[
cos.to(dtype=torch.float32).contiguous(),
sin.to(dtype=torch.float32).contiguous(),
],
dim=-1,
)
# apply_flashinfer_rope_qk_inplace is inplace kernel and q_img/k_img are views of query/key, so we need not copy back
q_out, k_out = apply_flashinfer_rope_qk_inplace(
q_img, k_img, cos_sin_cache, is_neox=True
)
else:
query[:, text_seq_length:, :, :] = _apply_rotary_emb(
query[:, text_seq_length:, :, :], cos, sin, is_neox_style=True
)
key[:, text_seq_length:, :, :] = _apply_rotary_emb(
key[:, text_seq_length:, :, :], cos, sin, is_neox_style=True
)
q_img, k_img = self.rotary_emb(
query=q_img,
key=k_img,
cos=cos,
sin=sin,
)
query[:, text_seq_length:, :, :] = q_img
key[:, text_seq_length:, :, :] = k_img
if kv_cache is not None:
if kv_cache.mode == "write":
@@ -70,6 +70,16 @@ def fused_add_gate(
return torch.addcmul(residual, x, gate.unsqueeze(1))
def _joy_complex_freqs(freqs_cis: Optional[torch.Tensor]) -> Optional[torch.Tensor]:
"""Complex-valued RoPE table from a hoisted cat([cos, sin], dim=-1)
cos_sin_cache tensor, split back in half.
"""
if freqs_cis is None:
return None
cos, sin = freqs_cis.chunk(2, dim=-1)
return torch.complex(cos.to(torch.float32), sin.to(torch.float32))
class ModulateWan(nn.Module):
"""Modulation layer for WanX."""
@@ -220,6 +230,8 @@ class MMDoubleStreamBlock(nn.Module):
vec: torch.Tensor,
vis_freqs_cis: Optional[torch.Tensor] = None,
txt_freqs_cis: Optional[torch.Tensor] = None,
vis_complex_freqs: Optional[torch.Tensor] = None,
txt_complex_freqs: Optional[torch.Tensor] = None,
num_replicated_suffix: int = 0,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Forward pass through multimodal double stream block."""
@@ -268,6 +280,7 @@ class MMDoubleStreamBlock(nn.Module):
k_norm=self.img_attn_k_norm,
head_dim=img_q.shape[-1],
cos_sin_cache=vis_freqs_cis,
freqs_complex=vis_complex_freqs,
is_neox=False,
allow_inplace=True,
)
@@ -295,6 +308,7 @@ class MMDoubleStreamBlock(nn.Module):
k_norm=self.txt_attn_k_norm,
head_dim=txt_q.shape[-1],
cos_sin_cache=txt_freqs_cis,
freqs_complex=txt_complex_freqs,
is_neox=False,
allow_inplace=True,
)
@@ -555,6 +569,9 @@ class JoyTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
txt_suffix_len = txt.shape[1] if sequence_shard_enabled else 0
vis_complex_freqs = _joy_complex_freqs(vis_freqs_cis)
txt_complex_freqs = _joy_complex_freqs(txt_freqs_cis)
# Pass through DiT blocks
for block in self.double_blocks:
img, txt = block(
@@ -563,6 +580,8 @@ class JoyTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
vec,
vis_freqs_cis,
txt_freqs_cis,
vis_complex_freqs=vis_complex_freqs,
txt_complex_freqs=txt_complex_freqs,
num_replicated_suffix=txt_suffix_len,
)
@@ -35,7 +35,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config impor
QuantizationConfig,
)
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
_apply_rotary_emb_complex,
RotaryEmbedding,
)
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
LayerwiseOffloadableModuleMixin,
@@ -130,6 +130,14 @@ class SelfAttention(nn.Module):
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
self.rotary_emb = RotaryEmbedding(
head_size=self.head_dim,
rotary_dim=self.head_dim,
use_precomputed_cache=False,
is_neox_style=False,
complex_dtype=torch.float64,
)
self.attn = USPAttention(
# Local heads per TP rank.
num_heads=self.num_heads_per_rank,
@@ -172,8 +180,11 @@ class SelfAttention(nn.Module):
v = v.view(b, s, self.num_heads_per_rank, self.head_dim)
# Apply RoPE
q = _apply_rotary_emb_complex(q, freqs)
k = _apply_rotary_emb_complex(k, freqs)
q, k = self.rotary_emb(
query=q,
key=k,
complex_freqs=freqs,
)
# USPAttention expects [B, S_local, H, D] format
# USPAttention handles SP communication internally; the tail meta keeps
@@ -977,6 +977,12 @@ class QwenImageCrossAttention(nn.Module):
make_contiguous=not self.use_fused_qkv_epilogue,
)
freqs_complex = cross_attention_kwargs.get("freqs_complex")
if freqs_complex is not None:
img_complex, txt_complex = freqs_complex
else:
img_complex = txt_complex = None
# Reshape for multi-head attention
img_query = img_query.unflatten(-1, (self.local_num_heads, self.head_dim))
img_key = img_key.unflatten(-1, (self.local_num_heads, self.head_dim))
@@ -1040,6 +1046,7 @@ class QwenImageCrossAttention(nn.Module):
k_norm=self.norm_k,
head_dim=self.head_dim,
cos_sin_cache=img_cache,
freqs_complex=img_complex,
is_neox=False,
allow_inplace=True,
)
@@ -1050,6 +1057,7 @@ class QwenImageCrossAttention(nn.Module):
k_norm=self.norm_added_k,
head_dim=self.head_dim,
cos_sin_cache=txt_cache,
freqs_complex=txt_complex,
is_neox=False,
allow_inplace=True,
)
@@ -2272,6 +2280,7 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
img_shapes: Optional[List[Tuple[int, int, int]]] = None,
txt_seq_lens: Optional[List[int]] = None,
freqs_cis: tuple[torch.Tensor, torch.Tensor] = None,
freqs_complex: tuple[torch.Tensor, torch.Tensor] = None,
additional_t_cond: Optional[torch.Tensor] = None,
guidance: torch.Tensor = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
@@ -2390,6 +2399,9 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
if freqs_cis is not None:
img_cache, txt_cache = freqs_cis
freqs_cis = (img_cache, shard_like(txt_cache, txt_shard, dim=0))
if freqs_complex is not None:
img_complex, txt_complex = freqs_complex
freqs_complex = (img_complex, shard_like(txt_complex, txt_shard, dim=0))
tail_meta = tail_attn_meta(
txt_shard,
encoder_hidden_states.shape[0],
@@ -2412,6 +2424,9 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
temb_txt_silu = temb_img_silu
image_rotary_emb = freqs_cis
if freqs_complex is not None:
block_attention_kwargs["freqs_complex"] = freqs_complex
for index_block, block in enumerate(self.transformer_blocks):
encoder_hidden_states, hidden_states = block(
hidden_states=hidden_states,
@@ -207,7 +207,7 @@
"wan2_2_t2v_14b_w8a8_2npu": {
"stages_ms": {
"InputValidationStage": 0.1,
"TextEncodingStage": 247.21,
"TextEncodingStage": 663.31,
"LatentPreparationStage": 0.37,
"TimestepPreparationStage": 5.99,
"DenoisingStage": 185484.28,
@@ -45,7 +45,7 @@ SGL_TEST_FILES_CI_DATA_REVISION = "4ce5eeb9606e378478b2d0964d83e960af4e88cf"
# The NPU pin is kept as a separate branch so ascend GT can be bumped independently
# when it's regenerated on its own cadence.
if current_platform.is_npu():
SGL_TEST_FILES_CI_DATA_REVISION = "7df858ead07940ff4d9489230fa9f040dd186789"
SGL_TEST_FILES_CI_DATA_REVISION = "cbeaa640573a4bb08dc1c41af92221bfae51b7f6"
SGL_TEST_FILES_CONSISTENCY_GT_ROOT = (
"https://raw.githubusercontent.com/"
@@ -0,0 +1,84 @@
"""Guards against complex_freqs silently mixing RoPE styles or crashing on a
confusing broadcast error instead of failing at the actual precondition.
"""
import unittest
import torch
from sglang.multimodal_gen.runtime.layers.layernorm import apply_qk_norm_rope
from sglang.multimodal_gen.runtime.layers.rotary_embedding import RotaryEmbedding
def _complex_freqs(seq_len: int, half_dim: int) -> torch.Tensor:
cos = torch.randn(seq_len, 1, half_dim)
sin = torch.randn(seq_len, 1, half_dim)
return torch.complex(cos, sin)
class TestRotaryEmbeddingComplexFreqsPreconditions(unittest.TestCase):
def test_neox_style_with_only_complex_freqs_fails_instead_of_wrong_result(self):
head_size = 64
rope = RotaryEmbedding(
head_size=head_size,
rotary_dim=head_size,
is_neox_style=True,
use_precomputed_cache=False,
)
seq_len = 5
query = torch.randn(1, seq_len, 2, head_size, dtype=torch.bfloat16)
key = torch.randn_like(query)
complex_freqs = _complex_freqs(seq_len, head_size // 2)
with self.assertRaisesRegex(ValueError, "No valid inputs"):
rope.forward_native(query=query, key=key, complex_freqs=complex_freqs)
def test_partial_rotary_dim_with_only_complex_freqs_fails_instead_of_crashing(self):
head_size, rotary_dim = 8, 4
rope = RotaryEmbedding(
head_size=head_size,
rotary_dim=rotary_dim,
is_neox_style=False,
use_precomputed_cache=False,
)
seq_len = 3
query = torch.randn(1, seq_len, 2, head_size, dtype=torch.bfloat16)
key = torch.randn_like(query)
# Sized by rotary_dim, as a real caller's derived table would be;
# this is the shape that used to crash inside view_as_complex.
complex_freqs = _complex_freqs(seq_len, rotary_dim // 2)
with self.assertRaisesRegex(ValueError, "No valid inputs"):
rope.forward_native(query=query, key=key, complex_freqs=complex_freqs)
class TestApplyQkNormRopeRequiresCache(unittest.TestCase):
def test_raises_without_cos_sin_cache(self):
# apply_qk_norm_with_optional_rope only reaches apply_qk_norm_rope
# when cos_sin_cache is not None; this pins down that apply_qk_norm_rope
# itself still enforces that precondition for its other direct
# callers (cosmos3video.py, ernie_image.py, zimage.py, ...), which
# never go through the wrapper. Passing freqs_complex must not let a
# caller substitute it for cos_sin_cache -- same raise either way.
head_dim = 8
seq_len = 3
q = torch.randn(1, seq_len, 2, head_dim, dtype=torch.bfloat16)
k = torch.randn_like(q)
freqs_complex = _complex_freqs(seq_len, head_dim // 2).squeeze(1)
with self.assertRaisesRegex(
ValueError, "cos_sin_cache must be a 2D torch.Tensor"
):
apply_qk_norm_rope(
q=q,
k=k,
q_norm=None,
k_norm=None,
head_dim=head_dim,
cos_sin_cache=None,
freqs_complex=freqs_complex,
)
if __name__ == "__main__":
unittest.main()