[diffusion] operator: unify rotary embedding impl (#18164)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
gemini-code-assist[bot]
parent
b21390f8f3
commit
26b2c63d03
@@ -28,143 +28,167 @@
|
|||||||
|
|
||||||
import functools
|
import functools
|
||||||
from collections import OrderedDict
|
from collections import OrderedDict
|
||||||
from typing import Any, Optional, Tuple
|
from typing import Any, Optional, Tuple, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import get_sp_group
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import get_sp_group
|
||||||
from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp
|
from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp
|
||||||
from sglang.multimodal_gen.runtime.layers.triton_ops import apply_rotary_embedding
|
from sglang.multimodal_gen.runtime.layers.triton_ops import (
|
||||||
|
apply_rotary_embedding,
|
||||||
|
apply_rotary_embedding_qk,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
try:
|
||||||
|
from flashinfer.rope import apply_rope_with_cos_sin_cache_inplace
|
||||||
|
except ImportError:
|
||||||
|
apply_rope_with_cos_sin_cache_inplace = None
|
||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
_is_flashinfer_available = (
|
||||||
|
current_platform.is_cuda() and apply_rope_with_cos_sin_cache_inplace is not None
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def apply_flashinfer_rope_qk_inplace(
|
def _rope_impl_naive(
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
k: torch.Tensor,
|
k: Optional[torch.Tensor],
|
||||||
cos_sin_cache: torch.Tensor,
|
cos: torch.Tensor,
|
||||||
*,
|
sin: torch.Tensor,
|
||||||
head_size: Optional[int] = None,
|
is_neox_style: bool = False,
|
||||||
is_neox: bool = False,
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
positions: Optional[torch.Tensor] = None,
|
cos = cos.float().unsqueeze(1)
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
sin = sin.float().unsqueeze(1)
|
||||||
if q.dim() != 4 or k.dim() != 4:
|
|
||||||
raise ValueError(
|
|
||||||
f"Expected q/k to be 4D [bsz, seqlen, nheads, head_size], "
|
|
||||||
f"got q:{tuple(q.shape)} k:{tuple(k.shape)}"
|
|
||||||
)
|
|
||||||
if q.shape != k.shape:
|
|
||||||
raise ValueError(
|
|
||||||
f"q and k must have the same shape, got {q.shape} vs {k.shape}"
|
|
||||||
)
|
|
||||||
|
|
||||||
if not (isinstance(cos_sin_cache, torch.Tensor) and cos_sin_cache.dim() == 2):
|
def _rope(x):
|
||||||
raise ValueError("cos_sin_cache must be a 2D torch.Tensor")
|
if is_neox_style:
|
||||||
|
x1, x2 = torch.chunk(x, 2, dim=-1)
|
||||||
bsz, seqlen, nheads, d = q.shape
|
|
||||||
if head_size is None:
|
|
||||||
head_size = d
|
|
||||||
if head_size != d:
|
|
||||||
raise ValueError(f"head_size mismatch: inferred {d}, but head_size={head_size}")
|
|
||||||
|
|
||||||
try:
|
|
||||||
from flashinfer.rope import apply_rope_with_cos_sin_cache_inplace
|
|
||||||
except ImportError:
|
|
||||||
# Triton fallback for AMD/ROCm where FlashInfer is not available
|
|
||||||
import warnings
|
|
||||||
|
|
||||||
warnings.warn(
|
|
||||||
"FlashInfer not available, using Triton fallback for RoPE",
|
|
||||||
stacklevel=2,
|
|
||||||
)
|
|
||||||
half_size = cos_sin_cache.shape[-1] // 2
|
|
||||||
if positions is None:
|
|
||||||
cos = cos_sin_cache[:seqlen, :half_size].to(q.dtype)
|
|
||||||
sin = cos_sin_cache[:seqlen, half_size:].to(q.dtype)
|
|
||||||
cos = cos.unsqueeze(0).expand(bsz, -1, -1).reshape(bsz * seqlen, -1)
|
|
||||||
sin = sin.unsqueeze(0).expand(bsz, -1, -1).reshape(bsz * seqlen, -1)
|
|
||||||
else:
|
else:
|
||||||
positions = positions.to(cos_sin_cache.device).view(-1)
|
x1, x2 = x[..., 0::2], x[..., 1::2]
|
||||||
cos = cos_sin_cache[positions, :half_size].to(q.dtype)
|
|
||||||
sin = cos_sin_cache[positions, half_size:].to(q.dtype)
|
|
||||||
q_flat = q.reshape(bsz * seqlen, nheads, d)
|
|
||||||
k_flat = k.reshape(bsz * seqlen, nheads, d)
|
|
||||||
q_rot = apply_rotary_embedding(q_flat, cos, sin, interleaved=not is_neox)
|
|
||||||
k_rot = apply_rotary_embedding(k_flat, cos, sin, interleaved=not is_neox)
|
|
||||||
return q_rot.view(bsz, seqlen, nheads, d), k_rot.view(bsz, seqlen, nheads, d)
|
|
||||||
|
|
||||||
if positions is None:
|
o1 = (x1.float() * cos - x2.float() * sin).to(dtype=x.dtype)
|
||||||
pos_1d = torch.arange(seqlen, device=q.device, dtype=torch.long)
|
o2 = (x2.float() * cos + x1.float() * sin).to(dtype=x.dtype)
|
||||||
positions = pos_1d if bsz == 1 else pos_1d.repeat(bsz)
|
|
||||||
|
if is_neox_style:
|
||||||
|
return torch.cat((o1, o2), dim=-1)
|
||||||
|
else:
|
||||||
|
return torch.stack((o1, o2), dim=-1).flatten(-2)
|
||||||
|
|
||||||
|
q_out = _rope(q)
|
||||||
|
if k is not None:
|
||||||
|
k_out = _rope(k)
|
||||||
|
return q_out, k_out
|
||||||
|
return q_out
|
||||||
|
|
||||||
|
|
||||||
|
def _rope_impl_triton(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: Optional[torch.Tensor],
|
||||||
|
cos: torch.Tensor,
|
||||||
|
sin: torch.Tensor,
|
||||||
|
is_neox_style: bool = False,
|
||||||
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
|
cos, sin = cos.float(), sin.float()
|
||||||
|
|
||||||
|
if k is not None:
|
||||||
|
return apply_rotary_embedding_qk(q, k, cos, sin, is_neox_style)
|
||||||
else:
|
else:
|
||||||
if not (
|
return apply_rotary_embedding(q, cos, sin, is_neox_style)
|
||||||
isinstance(positions, torch.Tensor)
|
|
||||||
and positions.dtype == torch.long
|
|
||||||
and positions.dim() == 1
|
def _rope_impl_flashinfer(
|
||||||
):
|
q: torch.Tensor,
|
||||||
raise ValueError("positions must be a 1D torch.long Tensor")
|
k: Optional[torch.Tensor],
|
||||||
if positions.numel() != bsz * seqlen:
|
cos: torch.Tensor,
|
||||||
raise ValueError(
|
sin: torch.Tensor,
|
||||||
f"positions length must be bsz*seqlen={bsz*seqlen}, got {positions.numel()}"
|
is_neox_style: bool = False,
|
||||||
)
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
|
|
||||||
|
q = q.contiguous()
|
||||||
|
if k is not None:
|
||||||
|
_k = k.contiguous()
|
||||||
|
else:
|
||||||
|
_k = torch.empty_like(q)
|
||||||
|
|
||||||
|
cos_sin_cache = torch.cat([cos, sin], dim=-1).float()
|
||||||
|
|
||||||
|
if q.dim() == 3:
|
||||||
|
seq_len, num_q_heads, head_dim = q.shape
|
||||||
|
bsz = 1
|
||||||
|
else:
|
||||||
|
bsz, seq_len, num_q_heads, head_dim = q.shape
|
||||||
|
|
||||||
|
num_kv_heads = _k.size(-2)
|
||||||
|
|
||||||
|
pos_1d = torch.arange(seq_len, device=q.device, dtype=torch.long)
|
||||||
|
positions = pos_1d if bsz == 1 else pos_1d.repeat(bsz)
|
||||||
|
|
||||||
q_flat = q.reshape(bsz * seqlen, nheads * d).contiguous()
|
|
||||||
k_flat = k.reshape(bsz * seqlen, nheads * d).contiguous()
|
|
||||||
apply_rope_with_cos_sin_cache_inplace(
|
apply_rope_with_cos_sin_cache_inplace(
|
||||||
positions=positions,
|
positions=positions,
|
||||||
query=q_flat,
|
query=q.view(-1, num_q_heads * head_dim),
|
||||||
key=k_flat,
|
key=_k.view(-1, num_kv_heads * head_dim),
|
||||||
head_size=d,
|
head_size=head_dim,
|
||||||
cos_sin_cache=cos_sin_cache,
|
cos_sin_cache=cos_sin_cache,
|
||||||
is_neox=is_neox,
|
is_neox=is_neox_style,
|
||||||
)
|
)
|
||||||
return q_flat.view(bsz, seqlen, nheads, d), k_flat.view(bsz, seqlen, nheads, d)
|
|
||||||
|
|
||||||
|
if k is not None:
|
||||||
def _rotate_neox(x: torch.Tensor) -> torch.Tensor:
|
return q, _k
|
||||||
x1 = x[..., : x.shape[-1] // 2]
|
return q
|
||||||
x2 = x[..., x.shape[-1] // 2 :]
|
|
||||||
return torch.cat((-x2, x1), dim=-1)
|
|
||||||
|
|
||||||
|
|
||||||
def _rotate_gptj(x: torch.Tensor) -> torch.Tensor:
|
|
||||||
x1 = x[..., ::2]
|
|
||||||
x2 = x[..., 1::2]
|
|
||||||
x = torch.stack((-x2, x1), dim=-1)
|
|
||||||
return x.flatten(-2)
|
|
||||||
|
|
||||||
|
|
||||||
def _apply_rotary_emb(
|
def _apply_rotary_emb(
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
cos: torch.Tensor,
|
cos: torch.Tensor,
|
||||||
sin: torch.Tensor,
|
sin: torch.Tensor,
|
||||||
is_neox_style: bool,
|
is_neox_style: bool = False,
|
||||||
interleaved: bool = False,
|
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
Args:
|
Args:
|
||||||
x: [num_tokens, num_heads, head_size] or [num_tokens, head_size]
|
x: [batch_size, seq_len, num_heads, head_dim] or [seq_len, num_heads, head_dim]
|
||||||
cos: [num_tokens, head_size // 2]
|
cos: [seq_len, head_dim // 2]
|
||||||
sin: [num_tokens, head_size // 2]
|
sin: [seq_len, head_dim // 2]
|
||||||
is_neox_style: Whether to use the Neox-style or GPT-J-style rotary
|
is_neox_style: Whether to use the Neox-style or GPT-J-style rotary
|
||||||
positional embeddings.
|
positional embeddings.
|
||||||
"""
|
"""
|
||||||
# cos = cos.unsqueeze(-2).to(x.dtype)
|
if _is_flashinfer_available and x.dtype in {torch.bfloat16, torch.float16}:
|
||||||
# sin = sin.unsqueeze(-2).to(x.dtype)
|
return _rope_impl_flashinfer(x, None, cos, sin, is_neox_style)
|
||||||
if is_neox_style:
|
|
||||||
cos = cos.unsqueeze(-2)
|
|
||||||
sin = sin.unsqueeze(-2)
|
|
||||||
if is_neox_style:
|
|
||||||
x1, x2 = torch.chunk(x, 2, dim=-1)
|
|
||||||
else:
|
|
||||||
x1 = x[..., ::2]
|
|
||||||
x2 = x[..., 1::2]
|
|
||||||
o1 = (x1.float() * cos - x2.float() * sin).type_as(x)
|
|
||||||
o2 = (x2.float() * cos + x1.float() * sin).type_as(x)
|
|
||||||
return torch.cat((o1, o2), dim=-1)
|
|
||||||
else:
|
else:
|
||||||
return apply_rotary_embedding(x, cos, sin, interleaved)
|
try:
|
||||||
|
return _rope_impl_triton(x, None, cos, sin, is_neox_style)
|
||||||
|
except Exception:
|
||||||
|
return _rope_impl_naive(x, None, cos, sin, is_neox_style)
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_rotary_emb_qk(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
cos: torch.Tensor,
|
||||||
|
sin: torch.Tensor,
|
||||||
|
is_neox_style: bool = False,
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
q: [batch_size, seq_len, num_heads, head_dim] or [seq_len, num_heads, head_dim]
|
||||||
|
k: [batch_size, seq_len, num_heads, head_dim] or [seq_len, num_heads, head_dim]
|
||||||
|
cos: [seq_len, head_dim // 2]
|
||||||
|
sin: [seq_len, head_dim // 2]
|
||||||
|
is_neox_style: Whether to use the Neox-style or GPT-J-style rotary
|
||||||
|
positional embeddings.
|
||||||
|
"""
|
||||||
|
if (
|
||||||
|
_is_flashinfer_available
|
||||||
|
and q.dtype in {torch.bfloat16, torch.float16}
|
||||||
|
and k.dtype in {torch.bfloat16, torch.float16}
|
||||||
|
):
|
||||||
|
return _rope_impl_flashinfer(q, k, cos, sin, is_neox_style)
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
return _rope_impl_triton(q, k, cos, sin, is_neox_style)
|
||||||
|
except Exception:
|
||||||
|
return _rope_impl_naive(q, k, cos, sin, is_neox_style)
|
||||||
|
|
||||||
|
|
||||||
@CustomOp.register("rotary_embedding")
|
@CustomOp.register("rotary_embedding")
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||||
|
|
||||||
# TODO: for temporary usage, expecting a refactor
|
# TODO: for temporary usage, expecting a refactor
|
||||||
from typing import Optional
|
from typing import Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import triton # type: ignore
|
import triton # type: ignore
|
||||||
@@ -450,8 +450,12 @@ def _rotary_embedding_kernel(
|
|||||||
cos_vals = tl.load(cos_row_ptr + offsets_half, mask=mask, other=0.0)
|
cos_vals = tl.load(cos_row_ptr + offsets_half, mask=mask, other=0.0)
|
||||||
sin_vals = tl.load(sin_row_ptr + offsets_half, mask=mask, other=0.0)
|
sin_vals = tl.load(sin_row_ptr + offsets_half, mask=mask, other=0.0)
|
||||||
|
|
||||||
offsets_x1 = 2 * offsets_half
|
if interleaved:
|
||||||
offsets_x2 = 2 * offsets_half + 1
|
offsets_x1 = 2 * offsets_half
|
||||||
|
offsets_x2 = 2 * offsets_half + 1
|
||||||
|
else:
|
||||||
|
offsets_x1 = offsets_half
|
||||||
|
offsets_x2 = offsets_half + head_size_half
|
||||||
|
|
||||||
x1_vals = tl.load(x_row_ptr + offsets_x1, mask=mask, other=0.0)
|
x1_vals = tl.load(x_row_ptr + offsets_x1, mask=mask, other=0.0)
|
||||||
x2_vals = tl.load(x_row_ptr + offsets_x2, mask=mask, other=0.0)
|
x2_vals = tl.load(x_row_ptr + offsets_x2, mask=mask, other=0.0)
|
||||||
@@ -467,6 +471,85 @@ def _rotary_embedding_kernel(
|
|||||||
tl.store(output_row_ptr + offsets_x2, o2_vals.to(x2_vals.dtype), mask=mask)
|
tl.store(output_row_ptr + offsets_x2, o2_vals.to(x2_vals.dtype), mask=mask)
|
||||||
|
|
||||||
|
|
||||||
|
@triton.autotune(
|
||||||
|
configs=[
|
||||||
|
triton.Config({"BLOCK_HS_HALF": 32}, num_warps=2),
|
||||||
|
triton.Config({"BLOCK_HS_HALF": 64}, num_warps=4),
|
||||||
|
triton.Config({"BLOCK_HS_HALF": 128}, num_warps=4),
|
||||||
|
triton.Config({"BLOCK_HS_HALF": 256}, num_warps=8),
|
||||||
|
],
|
||||||
|
key=["head_size", "interleaved"],
|
||||||
|
)
|
||||||
|
@triton.jit
|
||||||
|
def _rotary_embedding_qk_kernel(
|
||||||
|
output_q_ptr,
|
||||||
|
output_k_ptr,
|
||||||
|
q_ptr,
|
||||||
|
k_ptr,
|
||||||
|
cos_ptr,
|
||||||
|
sin_ptr,
|
||||||
|
num_heads,
|
||||||
|
head_size,
|
||||||
|
num_tokens,
|
||||||
|
stride_q_row,
|
||||||
|
stride_k_row,
|
||||||
|
stride_cos_row,
|
||||||
|
stride_sin_row,
|
||||||
|
interleaved: tl.constexpr,
|
||||||
|
BLOCK_HS_HALF: tl.constexpr,
|
||||||
|
):
|
||||||
|
row_idx = tl.program_id(0)
|
||||||
|
token_idx = (row_idx // num_heads) % num_tokens
|
||||||
|
|
||||||
|
q_row_ptr = q_ptr + row_idx * stride_q_row
|
||||||
|
k_row_ptr = k_ptr + row_idx * stride_k_row
|
||||||
|
cos_row_ptr = cos_ptr + token_idx * stride_cos_row
|
||||||
|
sin_row_ptr = sin_ptr + token_idx * stride_sin_row
|
||||||
|
output_q_row_ptr = output_q_ptr + row_idx * stride_q_row
|
||||||
|
output_k_row_ptr = output_k_ptr + row_idx * stride_k_row
|
||||||
|
|
||||||
|
# half size for x1 and x2
|
||||||
|
head_size_half = head_size // 2
|
||||||
|
|
||||||
|
for block_start in range(0, head_size_half, BLOCK_HS_HALF):
|
||||||
|
offsets_half = block_start + tl.arange(0, BLOCK_HS_HALF)
|
||||||
|
mask = offsets_half < head_size_half
|
||||||
|
|
||||||
|
cos_vals = tl.load(cos_row_ptr + offsets_half, mask=mask, other=0.0)
|
||||||
|
sin_vals = tl.load(sin_row_ptr + offsets_half, mask=mask, other=0.0)
|
||||||
|
|
||||||
|
if interleaved:
|
||||||
|
offsets_x1 = 2 * offsets_half
|
||||||
|
offsets_x2 = 2 * offsets_half + 1
|
||||||
|
else:
|
||||||
|
offsets_x1 = offsets_half
|
||||||
|
offsets_x2 = offsets_half + head_size_half
|
||||||
|
|
||||||
|
q1_vals = tl.load(q_row_ptr + offsets_x1, mask=mask, other=0.0)
|
||||||
|
q2_vals = tl.load(q_row_ptr + offsets_x2, mask=mask, other=0.0)
|
||||||
|
k1_vals = tl.load(k_row_ptr + offsets_x1, mask=mask, other=0.0)
|
||||||
|
k2_vals = tl.load(k_row_ptr + offsets_x2, mask=mask, other=0.0)
|
||||||
|
|
||||||
|
q1_fp32 = q1_vals.to(tl.float32)
|
||||||
|
q2_fp32 = q2_vals.to(tl.float32)
|
||||||
|
k1_fp32 = k1_vals.to(tl.float32)
|
||||||
|
k2_fp32 = k2_vals.to(tl.float32)
|
||||||
|
|
||||||
|
cos_fp32 = cos_vals.to(tl.float32)
|
||||||
|
sin_fp32 = sin_vals.to(tl.float32)
|
||||||
|
|
||||||
|
qo1_vals = tl.fma(-q2_fp32, sin_fp32, q1_fp32 * cos_fp32)
|
||||||
|
qo2_vals = tl.fma(q1_fp32, sin_fp32, q2_fp32 * cos_fp32)
|
||||||
|
ko1_vals = tl.fma(-k2_fp32, sin_fp32, k1_fp32 * cos_fp32)
|
||||||
|
ko2_vals = tl.fma(k1_fp32, sin_fp32, k2_fp32 * cos_fp32)
|
||||||
|
|
||||||
|
tl.store(output_q_row_ptr + offsets_x1, qo1_vals.to(q1_vals.dtype), mask=mask)
|
||||||
|
tl.store(output_q_row_ptr + offsets_x2, qo2_vals.to(q2_vals.dtype), mask=mask)
|
||||||
|
|
||||||
|
tl.store(output_k_row_ptr + offsets_x1, ko1_vals.to(k1_vals.dtype), mask=mask)
|
||||||
|
tl.store(output_k_row_ptr + offsets_x2, ko2_vals.to(k2_vals.dtype), mask=mask)
|
||||||
|
|
||||||
|
|
||||||
def apply_rotary_embedding(
|
def apply_rotary_embedding(
|
||||||
x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, interleaved: bool = False
|
x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, interleaved: bool = False
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
@@ -486,13 +569,6 @@ def apply_rotary_embedding(
|
|||||||
# num_tokens per head, 1 token per block
|
# num_tokens per head, 1 token per block
|
||||||
grid = (bsz * num_tokens * num_heads,)
|
grid = (bsz * num_tokens * num_heads,)
|
||||||
|
|
||||||
if interleaved and cos.shape[-1] == head_size:
|
|
||||||
cos = cos[..., ::2].contiguous()
|
|
||||||
sin = sin[..., ::2].contiguous()
|
|
||||||
else:
|
|
||||||
cos = cos.contiguous()
|
|
||||||
sin = sin.contiguous()
|
|
||||||
|
|
||||||
_rotary_embedding_kernel[grid](
|
_rotary_embedding_kernel[grid](
|
||||||
output_reshaped,
|
output_reshaped,
|
||||||
x_reshaped,
|
x_reshaped,
|
||||||
@@ -510,6 +586,54 @@ def apply_rotary_embedding(
|
|||||||
return output
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
def apply_rotary_embedding_qk(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
cos: torch.Tensor,
|
||||||
|
sin: torch.Tensor,
|
||||||
|
interleaved: bool = False,
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
output_q = torch.empty_like(q)
|
||||||
|
output_k = torch.empty_like(k)
|
||||||
|
cos = cos.contiguous()
|
||||||
|
sin = sin.contiguous()
|
||||||
|
|
||||||
|
if q.dim() > 3:
|
||||||
|
bsz, num_tokens, num_heads, head_size = q.shape
|
||||||
|
else:
|
||||||
|
num_tokens, num_heads, head_size = q.shape
|
||||||
|
bsz = 1
|
||||||
|
|
||||||
|
assert head_size % 2 == 0, "head_size must be divisible by 2"
|
||||||
|
|
||||||
|
q_reshaped = q.view(-1, head_size)
|
||||||
|
k_reshaped = k.view(-1, head_size)
|
||||||
|
output_q_reshaped = output_q.view(-1, head_size)
|
||||||
|
output_k_reshaped = output_k.view(-1, head_size)
|
||||||
|
|
||||||
|
# num_tokens per head, 1 token per block
|
||||||
|
grid = (bsz * num_tokens * num_heads,)
|
||||||
|
|
||||||
|
_rotary_embedding_qk_kernel[grid](
|
||||||
|
output_q_reshaped,
|
||||||
|
output_k_reshaped,
|
||||||
|
q_reshaped,
|
||||||
|
k_reshaped,
|
||||||
|
cos,
|
||||||
|
sin,
|
||||||
|
num_heads,
|
||||||
|
head_size,
|
||||||
|
num_tokens,
|
||||||
|
q_reshaped.stride(0),
|
||||||
|
k_reshaped.stride(0),
|
||||||
|
cos.stride(0),
|
||||||
|
sin.stride(0),
|
||||||
|
interleaved,
|
||||||
|
)
|
||||||
|
|
||||||
|
return output_q, output_k
|
||||||
|
|
||||||
|
|
||||||
# RMSNorm-fp32
|
# RMSNorm-fp32
|
||||||
def maybe_contiguous_lastdim(x):
|
def maybe_contiguous_lastdim(x):
|
||||||
return x.contiguous() if x is not None and x.stride(-1) != 1 else x
|
return x.contiguous() if x is not None and x.stride(-1) != 1 else x
|
||||||
|
|||||||
@@ -23,9 +23,7 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
|||||||
ReplicatedLinear,
|
ReplicatedLinear,
|
||||||
RowParallelLinear,
|
RowParallelLinear,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import _apply_rotary_emb
|
||||||
apply_flashinfer_rope_qk_inplace,
|
|
||||||
)
|
|
||||||
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||||
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
|
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
@@ -48,14 +46,14 @@ def compute_rope_cos_sin(
|
|||||||
making it compatible with FSDP meta device initialization.
|
making it compatible with FSDP meta device initialization.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
position_ids: Position IDs tensor [B, L] or [1, L]
|
position_ids: Position IDs tensor [L]
|
||||||
head_dim: Dimension of each attention head
|
head_dim: Dimension of each attention head
|
||||||
base: RoPE base frequency (default: 10000.0)
|
base: RoPE base frequency (default: 10000.0)
|
||||||
device: Target device
|
device: Target device
|
||||||
dtype: Output dtype
|
dtype: Output dtype
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(cos, sin): Each with shape [B, L, head_dim]
|
(cos, sin): Each with shape [L, head_dim]
|
||||||
"""
|
"""
|
||||||
device = device or position_ids.device
|
device = device or position_ids.device
|
||||||
dtype = dtype or torch.float32
|
dtype = dtype or torch.float32
|
||||||
@@ -66,18 +64,10 @@ def compute_rope_cos_sin(
|
|||||||
** (torch.arange(0, head_dim, 2, dtype=torch.float32, device=device) / head_dim)
|
** (torch.arange(0, head_dim, 2, dtype=torch.float32, device=device) / head_dim)
|
||||||
)
|
)
|
||||||
|
|
||||||
# Expand for batch computation: [B, L] -> [B, 1, L] @ [1, head_dim/2, 1] -> [B, head_dim/2, L]
|
freqs = torch.outer(inv_freq.float(), position_ids.float()).transpose(0, 1)
|
||||||
inv_freq_expanded = inv_freq[None, :, None].expand(position_ids.shape[0], -1, 1)
|
|
||||||
position_ids_expanded = position_ids[:, None, :].float()
|
|
||||||
|
|
||||||
# Compute frequencies: [B, head_dim/2, L] -> [B, L, head_dim/2]
|
cos = freqs.cos().to(dtype=dtype)
|
||||||
freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2)
|
sin = freqs.sin().to(dtype=dtype)
|
||||||
|
|
||||||
# Double the frequencies for full head_dim: [B, L, head_dim]
|
|
||||||
emb = torch.cat((freqs, freqs), dim=-1)
|
|
||||||
|
|
||||||
cos = emb.cos().to(dtype=dtype)
|
|
||||||
sin = emb.sin().to(dtype=dtype)
|
|
||||||
|
|
||||||
return cos, sin
|
return cos, sin
|
||||||
|
|
||||||
@@ -245,45 +235,13 @@ class ConditionalCrossAttention(nn.Module):
|
|||||||
if x_freqs is not None:
|
if x_freqs is not None:
|
||||||
x_cos, x_sin = x_freqs
|
x_cos, x_sin = x_freqs
|
||||||
q_view = rearrange(q, "b l (h d) -> b l h d", d=self.head_dim)
|
q_view = rearrange(q, "b l (h d) -> b l h d", d=self.head_dim)
|
||||||
x_cos = x_cos.to(q_view.dtype).to(q_view.device).squeeze(0)
|
q_view = _apply_rotary_emb(q_view, x_cos, x_sin, is_neox_style=True)
|
||||||
x_sin = x_sin.to(q_view.dtype).to(q_view.device).squeeze(0)
|
|
||||||
# FlashInfer expects cos_sin_cache with shape [seqlen, head_dim],
|
|
||||||
# where the first half is cos and the second half is sin, each with
|
|
||||||
# head_dim//2 elements. Since compute_rope_cos_sin duplicates the
|
|
||||||
# frequencies (cat((freqs, freqs))), we only take the first half.
|
|
||||||
half_dim = self.head_dim // 2
|
|
||||||
cos_sin_cache = torch.cat(
|
|
||||||
[
|
|
||||||
x_cos[:, :half_dim].to(dtype=torch.float32).contiguous(),
|
|
||||||
x_sin[:, :half_dim].to(dtype=torch.float32).contiguous(),
|
|
||||||
],
|
|
||||||
dim=-1,
|
|
||||||
)
|
|
||||||
q_view, _ = apply_flashinfer_rope_qk_inplace(
|
|
||||||
q_view, q_view.clone(), cos_sin_cache, is_neox=True
|
|
||||||
)
|
|
||||||
q = rearrange(q_view, "b l h d -> b l (h d)")
|
q = rearrange(q_view, "b l h d -> b l (h d)")
|
||||||
|
|
||||||
if y_freqs is not None:
|
if y_freqs is not None:
|
||||||
y_cos, y_sin = y_freqs
|
y_cos, y_sin = y_freqs
|
||||||
k_view = rearrange(k, "b l (h d) -> b l h d", d=self.head_dim)
|
k_view = rearrange(k, "b l (h d) -> b l h d", d=self.head_dim)
|
||||||
y_cos = y_cos.to(k_view.dtype).to(k_view.device).squeeze(0)
|
k_view = _apply_rotary_emb(k_view, y_cos, y_sin, is_neox_style=True)
|
||||||
y_sin = y_sin.to(k_view.dtype).to(k_view.device).squeeze(0)
|
|
||||||
# FlashInfer expects cos_sin_cache with shape [seqlen, head_dim],
|
|
||||||
# where the first half is cos and the second half is sin, each with
|
|
||||||
# head_dim//2 elements. Since compute_rope_cos_sin duplicates the
|
|
||||||
# frequencies (cat((freqs, freqs))), we only take the first half.
|
|
||||||
half_dim = self.head_dim // 2
|
|
||||||
cos_sin_cache = torch.cat(
|
|
||||||
[
|
|
||||||
y_cos[:, :half_dim].to(dtype=torch.float32).contiguous(),
|
|
||||||
y_sin[:, :half_dim].to(dtype=torch.float32).contiguous(),
|
|
||||||
],
|
|
||||||
dim=-1,
|
|
||||||
)
|
|
||||||
k_view, _ = apply_flashinfer_rope_qk_inplace(
|
|
||||||
k_view, k_view.clone(), cos_sin_cache, is_neox=True
|
|
||||||
)
|
|
||||||
k = rearrange(k_view, "b l h d -> b l (h d)")
|
k = rearrange(k_view, "b l h d -> b l (h d)")
|
||||||
|
|
||||||
q = rearrange(q, "b l (h d) -> b l h d", h=self.num_heads_per_rank)
|
q = rearrange(q, "b l (h d) -> b l h d", h=self.num_heads_per_rank)
|
||||||
@@ -538,7 +496,7 @@ class DualTowerConditionalBridge(
|
|||||||
dtype = dtype or torch.float32
|
dtype = dtype or torch.float32
|
||||||
|
|
||||||
# Audio positions: 0, 1, 2, ..., L_a-1
|
# Audio positions: 0, 1, 2, ..., L_a-1
|
||||||
audio_pos = torch.arange(L_a, device=device, dtype=torch.float32).unsqueeze(0)
|
audio_pos = torch.arange(L_a, device=device, dtype=torch.float32)
|
||||||
|
|
||||||
# Video positions: Align video frames to audio step units
|
# Video positions: Align video frames to audio step units
|
||||||
if self.apply_first_frame_bias_in_rope:
|
if self.apply_first_frame_bias_in_rope:
|
||||||
@@ -558,7 +516,7 @@ class DualTowerConditionalBridge(
|
|||||||
torch.arange(f_v, device=device, dtype=torch.float32) * scale
|
torch.arange(f_v, device=device, dtype=torch.float32) * scale
|
||||||
)
|
)
|
||||||
|
|
||||||
video_pos = video_pos_per_frame.repeat_interleave(h * w).unsqueeze(0)
|
video_pos = video_pos_per_frame.repeat_interleave(h * w)
|
||||||
|
|
||||||
# Use functional RoPE to compute cos/sin
|
# Use functional RoPE to compute cos/sin
|
||||||
cos_v, sin_v = compute_rope_cos_sin(
|
cos_v, sin_v = compute_rope_cos_sin(
|
||||||
|
|||||||
@@ -36,7 +36,7 @@ from sglang.multimodal_gen.runtime.layers.layernorm import (
|
|||||||
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
||||||
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
||||||
_apply_rotary_emb,
|
_apply_rotary_emb_qk,
|
||||||
get_rotary_pos_embed,
|
get_rotary_pos_embed,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.visual_embedding import PatchEmbed
|
from sglang.multimodal_gen.runtime.layers.visual_embedding import PatchEmbed
|
||||||
@@ -116,8 +116,9 @@ class CausalWanSelfAttention(nn.Module):
|
|||||||
cache_start = current_start
|
cache_start = current_start
|
||||||
|
|
||||||
cos, sin = freqs_cis
|
cos, sin = freqs_cis
|
||||||
roped_query = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v)
|
roped_query, roped_key = _apply_rotary_emb_qk(
|
||||||
roped_key = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
|
q, k, cos, sin, is_neox_style=False
|
||||||
|
)
|
||||||
|
|
||||||
if kv_cache is None:
|
if kv_cache is None:
|
||||||
# Padding for flex attention
|
# Padding for flex attention
|
||||||
|
|||||||
@@ -36,7 +36,7 @@ from sglang.multimodal_gen.runtime.layers.linear import ColumnParallelLinear
|
|||||||
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
||||||
NDRotaryEmbedding,
|
NDRotaryEmbedding,
|
||||||
apply_flashinfer_rope_qk_inplace,
|
_apply_rotary_emb_qk,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.visual_embedding import (
|
from sglang.multimodal_gen.runtime.layers.visual_embedding import (
|
||||||
CombinedTimestepGuidanceTextProjEmbeddings,
|
CombinedTimestepGuidanceTextProjEmbeddings,
|
||||||
@@ -195,15 +195,12 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
|
|
||||||
if freqs_cis is not None:
|
if freqs_cis is not None:
|
||||||
cos, sin = freqs_cis
|
cos, sin = freqs_cis
|
||||||
cos_sin_cache = torch.cat(
|
query, key = _apply_rotary_emb_qk(
|
||||||
[
|
query,
|
||||||
cos.to(dtype=torch.float32).contiguous(),
|
key,
|
||||||
sin.to(dtype=torch.float32).contiguous(),
|
cos,
|
||||||
],
|
sin,
|
||||||
dim=-1,
|
is_neox_style=False,
|
||||||
)
|
|
||||||
query, key = apply_flashinfer_rope_qk_inplace(
|
|
||||||
query, key, cos_sin_cache, is_neox=False
|
|
||||||
)
|
)
|
||||||
|
|
||||||
x = self.attn(query, key, value)
|
x = self.attn(query, key, value)
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, apply_qk_nor
|
|||||||
from sglang.multimodal_gen.runtime.layers.linear import ColumnParallelLinear
|
from sglang.multimodal_gen.runtime.layers.linear import ColumnParallelLinear
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
||||||
NDRotaryEmbedding,
|
NDRotaryEmbedding,
|
||||||
apply_flashinfer_rope_qk_inplace,
|
_apply_rotary_emb_qk,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
@@ -225,16 +225,7 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
|
|
||||||
if freqs_cis is not None:
|
if freqs_cis is not None:
|
||||||
cos, sin = freqs_cis
|
cos, sin = freqs_cis
|
||||||
cos_sin_cache = torch.cat(
|
query, key = _apply_rotary_emb_qk(query, key, cos, sin, is_neox_style=False)
|
||||||
[
|
|
||||||
cos.to(dtype=torch.float32).contiguous(),
|
|
||||||
sin.to(dtype=torch.float32).contiguous(),
|
|
||||||
],
|
|
||||||
dim=-1,
|
|
||||||
)
|
|
||||||
query, key = apply_flashinfer_rope_qk_inplace(
|
|
||||||
query, key, cos_sin_cache, is_neox=False
|
|
||||||
)
|
|
||||||
|
|
||||||
hidden_states = self.attn(query, key, value)
|
hidden_states = self.attn(query, key, value)
|
||||||
|
|
||||||
@@ -357,16 +348,8 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
|
|
||||||
if freqs_cis is not None:
|
if freqs_cis is not None:
|
||||||
cos, sin = freqs_cis
|
cos, sin = freqs_cis
|
||||||
cos_sin_cache = torch.cat(
|
query, key = _apply_rotary_emb_qk(query, key, cos, sin, is_neox_style=False)
|
||||||
[
|
|
||||||
cos.to(dtype=torch.float32).contiguous(),
|
|
||||||
sin.to(dtype=torch.float32).contiguous(),
|
|
||||||
],
|
|
||||||
dim=-1,
|
|
||||||
)
|
|
||||||
query, key = apply_flashinfer_rope_qk_inplace(
|
|
||||||
query, key, cos_sin_cache, is_neox=False
|
|
||||||
)
|
|
||||||
hidden_states = self.attn(query, key, value)
|
hidden_states = self.attn(query, key, value)
|
||||||
hidden_states = hidden_states.flatten(2, 3)
|
hidden_states = hidden_states.flatten(2, 3)
|
||||||
hidden_states = hidden_states.to(query.dtype)
|
hidden_states = hidden_states.to(query.dtype)
|
||||||
|
|||||||
@@ -25,10 +25,7 @@ from sglang.multimodal_gen.runtime.layers.layernorm import (
|
|||||||
ScaleResidualLayerNormScaleShift,
|
ScaleResidualLayerNormScaleShift,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import _apply_rotary_emb_qk
|
||||||
_apply_rotary_emb,
|
|
||||||
apply_flashinfer_rope_qk_inplace,
|
|
||||||
)
|
|
||||||
from sglang.multimodal_gen.runtime.layers.visual_embedding import Timesteps
|
from sglang.multimodal_gen.runtime.layers.visual_embedding import Timesteps
|
||||||
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||||
from sglang.multimodal_gen.runtime.platforms import (
|
from sglang.multimodal_gen.runtime.platforms import (
|
||||||
@@ -390,28 +387,15 @@ class GlmImageAttention(torch.nn.Module):
|
|||||||
# 3. Rotational positional embeddings applied to latent stream
|
# 3. Rotational positional embeddings applied to latent stream
|
||||||
if image_rotary_emb is not None:
|
if image_rotary_emb is not None:
|
||||||
cos, sin = image_rotary_emb
|
cos, sin = image_rotary_emb
|
||||||
|
query[:, text_seq_length:, :, :], key[:, text_seq_length:, :, :] = (
|
||||||
if _is_cuda and cos.dim() == 2:
|
_apply_rotary_emb_qk(
|
||||||
q_img = query[:, text_seq_length:, :, :]
|
query[:, text_seq_length:, :, :],
|
||||||
k_img = key[:, text_seq_length:, :, :]
|
key[:, text_seq_length:, :, :],
|
||||||
cos_sin_cache = torch.cat(
|
cos,
|
||||||
[
|
sin,
|
||||||
cos.to(dtype=torch.float32).contiguous(),
|
is_neox_style=True,
|
||||||
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
|
|
||||||
)
|
)
|
||||||
|
)
|
||||||
|
|
||||||
if kv_cache is not None:
|
if kv_cache is not None:
|
||||||
if kv_cache.mode == "write":
|
if kv_cache.mode == "write":
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ from sglang.multimodal_gen.runtime.layers.layernorm import (
|
|||||||
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
||||||
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
||||||
_apply_rotary_emb,
|
_apply_rotary_emb_qk,
|
||||||
get_rotary_pos_embed,
|
get_rotary_pos_embed,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.visual_embedding import (
|
from sglang.multimodal_gen.runtime.layers.visual_embedding import (
|
||||||
@@ -194,14 +194,13 @@ class MMDoubleStreamBlock(nn.Module):
|
|||||||
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :, 2]
|
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :, 2]
|
||||||
|
|
||||||
# Apply QK-Norm if needed
|
# Apply QK-Norm if needed
|
||||||
|
|
||||||
img_q = self.img_attn_q_norm(img_q.contiguous()).to(img_v)
|
img_q = self.img_attn_q_norm(img_q.contiguous()).to(img_v)
|
||||||
img_k = self.img_attn_k_norm(img_k.contiguous()).to(img_v)
|
img_k = self.img_attn_k_norm(img_k.contiguous()).to(img_v)
|
||||||
|
|
||||||
# Apply rotary embeddings
|
# Apply rotary embeddings
|
||||||
cos, sin = freqs_cis
|
cos, sin = freqs_cis
|
||||||
img_q, img_k = _apply_rotary_emb(
|
img_q, img_k = _apply_rotary_emb_qk(img_q, img_k, cos, sin, is_neox_style=False)
|
||||||
img_q, cos, sin, is_neox_style=False
|
|
||||||
), _apply_rotary_emb(img_k, cos, sin, is_neox_style=False)
|
|
||||||
# Prepare text for attention using fused operation
|
# Prepare text for attention using fused operation
|
||||||
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
|
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
|
||||||
|
|
||||||
@@ -360,11 +359,10 @@ class MMSingleStreamBlock(nn.Module):
|
|||||||
img_q, txt_q = q[:, :-txt_len], q[:, -txt_len:]
|
img_q, txt_q = q[:, :-txt_len], q[:, -txt_len:]
|
||||||
img_k, txt_k = k[:, :-txt_len], k[:, -txt_len:]
|
img_k, txt_k = k[:, :-txt_len], k[:, -txt_len:]
|
||||||
img_v, txt_v = v[:, :-txt_len], v[:, -txt_len:]
|
img_v, txt_v = v[:, :-txt_len], v[:, -txt_len:]
|
||||||
|
|
||||||
# Apply rotary embeddings to image parts
|
# Apply rotary embeddings to image parts
|
||||||
cos, sin = freqs_cis
|
cos, sin = freqs_cis
|
||||||
img_q, img_k = _apply_rotary_emb(
|
img_q, img_k = _apply_rotary_emb_qk(img_q, img_k, cos, sin, is_neox_style=False)
|
||||||
img_q, cos, sin, is_neox_style=False
|
|
||||||
), _apply_rotary_emb(img_k, cos, sin, is_neox_style=False)
|
|
||||||
|
|
||||||
# Run distributed attention
|
# Run distributed attention
|
||||||
img_attn_output, txt_attn_output = self.attn(
|
img_attn_output, txt_attn_output = self.attn(
|
||||||
|
|||||||
@@ -233,9 +233,10 @@ class WanAudioModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
],
|
],
|
||||||
dim=-1,
|
dim=-1,
|
||||||
)
|
)
|
||||||
.reshape(f, 1, -1)
|
.reshape(f, -1)
|
||||||
.to(x.device)
|
.to(x.device)
|
||||||
)
|
)
|
||||||
|
freqs = (freqs.real.contiguous().float(), freqs.imag.contiguous().float())
|
||||||
|
|
||||||
for block in self.blocks:
|
for block in self.blocks:
|
||||||
x = block(x, context, t_mod, freqs)
|
x = block(x, context, t_mod, freqs)
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
|||||||
RowParallelLinear,
|
RowParallelLinear,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
||||||
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import _apply_rotary_emb_qk
|
||||||
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||||
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
|
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
@@ -76,25 +77,6 @@ def precompute_freqs_cis(
|
|||||||
return freqs_cis
|
return freqs_cis
|
||||||
|
|
||||||
|
|
||||||
def rope_apply(x, freqs, num_heads):
|
|
||||||
x = rearrange(x, "b s (n d) -> b s n d", n=num_heads)
|
|
||||||
x_out = torch.view_as_complex(
|
|
||||||
x.to(torch.float64).reshape(x.shape[0], x.shape[1], x.shape[2], -1, 2)
|
|
||||||
)
|
|
||||||
x_out = torch.view_as_real(x_out * freqs).flatten(2)
|
|
||||||
return x_out.to(x.dtype)
|
|
||||||
|
|
||||||
|
|
||||||
def rope_apply_head_dim(x, freqs, head_dim):
|
|
||||||
x = rearrange(x, "b s (n d) -> b s n d", d=head_dim)
|
|
||||||
x_out = torch.view_as_complex(
|
|
||||||
x.to(torch.float64).reshape(x.shape[0], x.shape[1], x.shape[2], -1, 2)
|
|
||||||
)
|
|
||||||
# print(f"{x_out.shape = }, {freqs.shape = }")
|
|
||||||
x_out = torch.view_as_real(x_out * freqs).flatten(2)
|
|
||||||
return x_out.to(x.dtype)
|
|
||||||
|
|
||||||
|
|
||||||
class SelfAttention(nn.Module):
|
class SelfAttention(nn.Module):
|
||||||
"""
|
"""
|
||||||
Self-Attention module for MOVA DiT with Sequence Parallelism support.
|
Self-Attention module for MOVA DiT with Sequence Parallelism support.
|
||||||
@@ -139,14 +121,11 @@ class SelfAttention(nn.Module):
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
x: Input tensor [B, S_local, D] - already sharded by SP when SP > 1
|
x: Input tensor [B, S_local, D] - already sharded by SP when SP > 1
|
||||||
freqs: RoPE frequencies [S_local, 1, head_dim] - should match x's sequence length
|
freqs: RoPE frequencies [S_local, head_dim] - should match x's sequence length
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Output tensor [B, S_local, D]
|
Output tensor [B, S_local, D]
|
||||||
"""
|
"""
|
||||||
if isinstance(freqs, DTensor):
|
|
||||||
freqs = freqs.to_local()
|
|
||||||
|
|
||||||
# Compute Q, K, V on local sequence
|
# Compute Q, K, V on local sequence
|
||||||
q, _ = self.q(x)
|
q, _ = self.q(x)
|
||||||
k, _ = self.k(x)
|
k, _ = self.k(x)
|
||||||
@@ -160,15 +139,15 @@ class SelfAttention(nn.Module):
|
|||||||
q = self.norm_q(q)
|
q = self.norm_q(q)
|
||||||
k = self.norm_k(k)
|
k = self.norm_k(k)
|
||||||
|
|
||||||
# Apply RoPE
|
|
||||||
q = rope_apply_head_dim(q, freqs, self.head_dim)
|
|
||||||
k = rope_apply_head_dim(k, freqs, self.head_dim)
|
|
||||||
|
|
||||||
# USPAttention expects [B, S_local, H, D] format
|
# USPAttention expects [B, S_local, H, D] format
|
||||||
q = rearrange(q, "b s (n d) -> b s n d", n=self.num_heads_per_rank)
|
q = rearrange(q, "b s (n d) -> b s n d", n=self.num_heads_per_rank)
|
||||||
k = rearrange(k, "b s (n d) -> b s n d", n=self.num_heads_per_rank)
|
k = rearrange(k, "b s (n d) -> b s n d", n=self.num_heads_per_rank)
|
||||||
v = rearrange(v, "b s (n d) -> b s n d", n=self.num_heads_per_rank)
|
v = rearrange(v, "b s (n d) -> b s n d", n=self.num_heads_per_rank)
|
||||||
|
|
||||||
|
# Apply RoPE
|
||||||
|
cos, sin = freqs
|
||||||
|
q, k = _apply_rotary_emb_qk(q, k, cos, sin, is_neox_style=False)
|
||||||
|
|
||||||
# USPAttention handles SP communication internally
|
# USPAttention handles SP communication internally
|
||||||
out = self.attn(q, k, v)
|
out = self.attn(q, k, v)
|
||||||
out = rearrange(out, "b s n d -> b s (n d)")
|
out = rearrange(out, "b s n d -> b s (n d)")
|
||||||
@@ -515,9 +494,10 @@ class WanModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
],
|
],
|
||||||
dim=-1,
|
dim=-1,
|
||||||
)
|
)
|
||||||
.reshape(f * h * w, 1, -1)
|
.reshape(f * h * w, -1)
|
||||||
.to(x.device)
|
.to(x.device)
|
||||||
)
|
)
|
||||||
|
freqs = (freqs.real.contiguous().float(), freqs.imag.contiguous().float())
|
||||||
|
|
||||||
for block in self.blocks:
|
for block in self.blocks:
|
||||||
x = block(x, context, t_mod, freqs)
|
x = block(x, context, t_mod, freqs)
|
||||||
|
|||||||
@@ -38,9 +38,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.configs.nunchaku_config i
|
|||||||
NunchakuConfig,
|
NunchakuConfig,
|
||||||
is_nunchaku_available,
|
is_nunchaku_available,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import _apply_rotary_emb_qk
|
||||||
apply_flashinfer_rope_qk_inplace,
|
|
||||||
)
|
|
||||||
from sglang.multimodal_gen.runtime.layers.triton_ops import (
|
from sglang.multimodal_gen.runtime.layers.triton_ops import (
|
||||||
fuse_scale_shift_gate_select01_kernel,
|
fuse_scale_shift_gate_select01_kernel,
|
||||||
)
|
)
|
||||||
@@ -763,12 +761,14 @@ class QwenImageCrossAttention(nn.Module):
|
|||||||
raise RuntimeError("image_rotary_emb must be cos_sin_cache tensors")
|
raise RuntimeError("image_rotary_emb must be cos_sin_cache tensors")
|
||||||
|
|
||||||
img_cache, txt_cache = image_rotary_emb
|
img_cache, txt_cache = image_rotary_emb
|
||||||
|
img_cos, img_sin = img_cache.chunk(2, dim=-1)
|
||||||
|
txt_cos, txt_sin = txt_cache.chunk(2, dim=-1)
|
||||||
|
|
||||||
img_query, img_key = apply_flashinfer_rope_qk_inplace(
|
img_query, img_key = _apply_rotary_emb_qk(
|
||||||
img_query, img_key, img_cache, is_neox=False
|
img_query, img_key, img_cos, img_sin, is_neox_style=False
|
||||||
)
|
)
|
||||||
txt_query, txt_key = apply_flashinfer_rope_qk_inplace(
|
txt_query, txt_key = _apply_rotary_emb_qk(
|
||||||
txt_query, txt_key, txt_cache, is_neox=False
|
txt_query, txt_key, txt_cos, txt_sin, is_neox_style=False
|
||||||
)
|
)
|
||||||
|
|
||||||
# Concatenate for joint attention
|
# Concatenate for joint attention
|
||||||
|
|||||||
@@ -38,8 +38,7 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
|||||||
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
from sglang.multimodal_gen.runtime.layers.mlp import MLP
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
||||||
NDRotaryEmbedding,
|
NDRotaryEmbedding,
|
||||||
_apply_rotary_emb,
|
_apply_rotary_emb_qk,
|
||||||
apply_flashinfer_rope_qk_inplace,
|
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.visual_embedding import (
|
from sglang.multimodal_gen.runtime.layers.visual_embedding import (
|
||||||
ModulateProjection,
|
ModulateProjection,
|
||||||
@@ -458,21 +457,8 @@ class WanTransformerBlock(nn.Module):
|
|||||||
|
|
||||||
# Apply rotary embeddings
|
# Apply rotary embeddings
|
||||||
cos, sin = freqs_cis
|
cos, sin = freqs_cis
|
||||||
if _is_cuda and query.shape == key.shape:
|
query, key = _apply_rotary_emb_qk(query, key, cos, sin, is_neox_style=False)
|
||||||
cos_sin_cache = torch.cat(
|
|
||||||
[
|
|
||||||
cos.to(dtype=torch.float32).contiguous(),
|
|
||||||
sin.to(dtype=torch.float32).contiguous(),
|
|
||||||
],
|
|
||||||
dim=-1,
|
|
||||||
)
|
|
||||||
query, key = apply_flashinfer_rope_qk_inplace(
|
|
||||||
query, key, cos_sin_cache, is_neox=False
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
query, key = _apply_rotary_emb(
|
|
||||||
query, cos, sin, is_neox_style=False
|
|
||||||
), _apply_rotary_emb(key, cos, sin, is_neox_style=False)
|
|
||||||
attn_output = self.attn1(query, key, value)
|
attn_output = self.attn1(query, key, value)
|
||||||
attn_output = attn_output.flatten(2)
|
attn_output = attn_output.flatten(2)
|
||||||
attn_output, _ = self.to_out(attn_output)
|
attn_output, _ = self.to_out(attn_output)
|
||||||
@@ -640,21 +626,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
|||||||
|
|
||||||
# Apply rotary embeddings
|
# Apply rotary embeddings
|
||||||
cos, sin = freqs_cis
|
cos, sin = freqs_cis
|
||||||
if _is_cuda and query.shape == key.shape:
|
query, key = _apply_rotary_emb_qk(query, key, cos, sin, is_neox_style=False)
|
||||||
cos_sin_cache = torch.cat(
|
|
||||||
[
|
|
||||||
cos.to(dtype=torch.float32).contiguous(),
|
|
||||||
sin.to(dtype=torch.float32).contiguous(),
|
|
||||||
],
|
|
||||||
dim=-1,
|
|
||||||
)
|
|
||||||
query, key = apply_flashinfer_rope_qk_inplace(
|
|
||||||
query, key, cos_sin_cache, is_neox=False
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
query, key = _apply_rotary_emb(
|
|
||||||
query, cos, sin, is_neox_style=False
|
|
||||||
), _apply_rotary_emb(key, cos, sin, is_neox_style=False)
|
|
||||||
|
|
||||||
attn_output = self.attn1(query, key, value, gate_compress=gate_compress)
|
attn_output = self.attn1(query, key, value, gate_compress=gate_compress)
|
||||||
attn_output = attn_output.flatten(2)
|
attn_output = attn_output.flatten(2)
|
||||||
|
|||||||
@@ -15,10 +15,7 @@ from sglang.multimodal_gen.runtime.layers.linear import (
|
|||||||
ReplicatedLinear,
|
ReplicatedLinear,
|
||||||
RowParallelLinear,
|
RowParallelLinear,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import _apply_rotary_emb_qk
|
||||||
_apply_rotary_emb,
|
|
||||||
apply_flashinfer_rope_qk_inplace,
|
|
||||||
)
|
|
||||||
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
|
||||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||||
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
|
from sglang.multimodal_gen.runtime.utils.layerwise_offload import OffloadableDiTMixin
|
||||||
@@ -181,20 +178,7 @@ class ZImageAttention(nn.Module):
|
|||||||
|
|
||||||
if freqs_cis is not None:
|
if freqs_cis is not None:
|
||||||
cos, sin = freqs_cis
|
cos, sin = freqs_cis
|
||||||
if _is_cuda and q.shape == k.shape:
|
q, k = _apply_rotary_emb_qk(q, k, cos, sin, is_neox_style=False)
|
||||||
cos_sin_cache = torch.cat(
|
|
||||||
[
|
|
||||||
cos.to(dtype=torch.float32).contiguous(),
|
|
||||||
sin.to(dtype=torch.float32).contiguous(),
|
|
||||||
],
|
|
||||||
dim=-1,
|
|
||||||
)
|
|
||||||
q, k = apply_flashinfer_rope_qk_inplace(
|
|
||||||
q, k, cos_sin_cache, is_neox=False
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
q = _apply_rotary_emb(q, cos, sin, is_neox_style=False)
|
|
||||||
k = _apply_rotary_emb(k, cos, sin, is_neox_style=False)
|
|
||||||
|
|
||||||
hidden_states = self.attn(q, k, v)
|
hidden_states = self.attn(q, k, v)
|
||||||
hidden_states = hidden_states.flatten(2)
|
hidden_states = hidden_states.flatten(2)
|
||||||
|
|||||||
+13
-4
@@ -701,7 +701,7 @@ class MOVADenoisingStage(PipelineStage):
|
|||||||
],
|
],
|
||||||
dim=-1,
|
dim=-1,
|
||||||
)
|
)
|
||||||
.reshape(full_visual_seq_len, 1, -1)
|
.reshape(full_visual_seq_len, -1)
|
||||||
.to(visual_x.device)
|
.to(visual_x.device)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -720,7 +720,7 @@ class MOVADenoisingStage(PipelineStage):
|
|||||||
],
|
],
|
||||||
dim=-1,
|
dim=-1,
|
||||||
)
|
)
|
||||||
.reshape(full_audio_seq_len, 1, -1)
|
.reshape(full_audio_seq_len, -1)
|
||||||
.to(audio_x.device)
|
.to(audio_x.device)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -732,6 +732,15 @@ class MOVADenoisingStage(PipelineStage):
|
|||||||
visual_freqs, _ = self._shard_sequence_for_sp(visual_freqs, dim=0)
|
visual_freqs, _ = self._shard_sequence_for_sp(visual_freqs, dim=0)
|
||||||
audio_freqs, _ = self._shard_sequence_for_sp(audio_freqs, dim=0)
|
audio_freqs, _ = self._shard_sequence_for_sp(audio_freqs, dim=0)
|
||||||
|
|
||||||
|
visual_freqs = (
|
||||||
|
visual_freqs.real.contiguous().float(),
|
||||||
|
visual_freqs.imag.contiguous().float(),
|
||||||
|
)
|
||||||
|
audio_freqs = (
|
||||||
|
audio_freqs.real.contiguous().float(),
|
||||||
|
audio_freqs.imag.contiguous().float(),
|
||||||
|
)
|
||||||
|
|
||||||
# Forward through dual-tower DiT
|
# Forward through dual-tower DiT
|
||||||
visual_x, audio_x = self.forward_dual_tower_dit(
|
visual_x, audio_x = self.forward_dual_tower_dit(
|
||||||
visual_dit=visual_dit,
|
visual_dit=visual_dit,
|
||||||
@@ -770,8 +779,8 @@ class MOVADenoisingStage(PipelineStage):
|
|||||||
audio_context: torch.Tensor,
|
audio_context: torch.Tensor,
|
||||||
visual_t_mod: torch.Tensor,
|
visual_t_mod: torch.Tensor,
|
||||||
audio_t_mod: torch.Tensor,
|
audio_t_mod: torch.Tensor,
|
||||||
visual_freqs: torch.Tensor,
|
visual_freqs: tuple[torch.Tensor, torch.Tensor],
|
||||||
audio_freqs: torch.Tensor,
|
audio_freqs: tuple[torch.Tensor, torch.Tensor],
|
||||||
grid_size: tuple[int, int, int],
|
grid_size: tuple[int, int, int],
|
||||||
video_fps: float,
|
video_fps: float,
|
||||||
full_visual_seq_len: int,
|
full_visual_seq_len: int,
|
||||||
|
|||||||
Reference in New Issue
Block a user