[diffusion] perf: add unified SP shard helpers and zero-copy tail-pad attention (#30107)
This commit is contained in:
@@ -145,37 +145,15 @@ def pad_text_embeddings_with_mask(
|
|||||||
|
|
||||||
|
|
||||||
def shard_rotary_emb_for_sp(emb):
|
def shard_rotary_emb_for_sp(emb):
|
||||||
"""
|
"""Shard rotary embeddings [S, D] along the sequence for SP; non-divisible
|
||||||
Shard rotary embeddings [S, D] along sequence for SP.
|
lengths pad by repeating the last row (position labels, never attention
|
||||||
If S is not divisible by SP degree, pad by repeating the last row.
|
K/V, so the pad value only needs to stay finite)."""
|
||||||
"""
|
|
||||||
# Sequence Parallelism: slice image RoPE to local shard if enabled
|
|
||||||
try:
|
try:
|
||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import shard_seq
|
||||||
get_sp_parallel_rank,
|
|
||||||
get_sp_world_size,
|
|
||||||
)
|
|
||||||
|
|
||||||
sp_world_size = get_sp_world_size()
|
return shard_seq(emb, dim=0, pad_mode="repeat_last")[0]
|
||||||
except Exception:
|
except Exception:
|
||||||
sp_world_size = 1
|
# Distributed state not initialized (single-process utilities).
|
||||||
seq_len = emb.shape[0]
|
|
||||||
if seq_len % sp_world_size != 0:
|
|
||||||
pad_len = sp_world_size - (seq_len % sp_world_size)
|
|
||||||
pad = emb[-1:].repeat(pad_len, 1)
|
|
||||||
emb = torch.cat([emb, pad], dim=0)
|
|
||||||
if sp_world_size > 1:
|
|
||||||
try:
|
|
||||||
rank = get_sp_parallel_rank()
|
|
||||||
except Exception:
|
|
||||||
rank = 0
|
|
||||||
seq_len = emb.shape[0]
|
|
||||||
local_len = seq_len // sp_world_size
|
|
||||||
start = rank * local_len
|
|
||||||
end = start + local_len
|
|
||||||
emb = emb[start:end]
|
|
||||||
return emb
|
|
||||||
else:
|
|
||||||
return emb
|
return emb
|
||||||
|
|
||||||
|
|
||||||
@@ -543,18 +521,15 @@ class PipelineConfig:
|
|||||||
return latents, False
|
return latents, False
|
||||||
time_dim = latents.shape[2]
|
time_dim = latents.shape[2]
|
||||||
|
|
||||||
# Pad to next multiple of SP degree if needed
|
# Zero-padding a non-divisible time dim would enter self-attention
|
||||||
|
# unmasked (video models pass no attn_mask) and corrupt real tokens;
|
||||||
|
# keep such shapes unsharded until models consume the sp_shard meta.
|
||||||
if time_dim > 0 and time_dim % sp_world_size != 0:
|
if time_dim > 0 and time_dim % sp_world_size != 0:
|
||||||
logger.debug(
|
logger.warning_once(
|
||||||
"Padding latents to next multiple of SP degree, performance is sub-optimal"
|
f"Latent time dim {time_dim} is not divisible by SP degree "
|
||||||
|
f"{sp_world_size}; skipping sequence shard for correctness."
|
||||||
)
|
)
|
||||||
pad_len = sp_world_size - (time_dim % sp_world_size)
|
return latents, False
|
||||||
pad = torch.zeros(
|
|
||||||
(*latents.shape[:2], pad_len, *latents.shape[3:]),
|
|
||||||
dtype=latents.dtype,
|
|
||||||
device=latents.device,
|
|
||||||
)
|
|
||||||
latents = torch.cat([latents, pad], dim=2)
|
|
||||||
|
|
||||||
assert latents.shape[2] % sp_world_size == 0
|
assert latents.shape[2] % sp_world_size == 0
|
||||||
sharded_tensor = rearrange(
|
sharded_tensor = rearrange(
|
||||||
|
|||||||
@@ -0,0 +1,225 @@
|
|||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
"""Unified sequence-parallel shard / pad / gather helpers.
|
||||||
|
|
||||||
|
Layout invariant: padding always sits at the end of the LAST rank's local
|
||||||
|
chunk, so the ulysses-gathered sequence carries one contiguous pad block at its
|
||||||
|
global tail. `tail_attn_meta` then lets attention skip that block for free
|
||||||
|
(the pad becomes its own varlen segment - no repacking, no mask compute).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.distributed.communication_op import (
|
||||||
|
sequence_model_parallel_all_gather,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
|
get_ring_parallel_world_size,
|
||||||
|
get_sp_parallel_rank,
|
||||||
|
get_sp_world_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Text shorter than this stays replicated instead of SP-sharded (see
|
||||||
|
# plan_text_strategy). 0 = always shard when legal; H100 bench showed sharding
|
||||||
|
# wins from trivial lengths on, so the knob exists only as an escape hatch.
|
||||||
|
_TEXT_SHARD_MIN = int(os.environ.get("SGLANG_SP_TEXT_SHARD_MIN", "0"))
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class SpShard:
|
||||||
|
"""Facts of one tail-padded even shard, shared by tensors of that stream."""
|
||||||
|
|
||||||
|
orig_len: int # real tokens (global)
|
||||||
|
local_len: int # per-rank chunk length (equal on every rank)
|
||||||
|
num_pad: int # pad tokens, all at the last rank's local tail
|
||||||
|
sp_size: int
|
||||||
|
sp_rank: int
|
||||||
|
|
||||||
|
@property
|
||||||
|
def local_pad(self) -> int:
|
||||||
|
"""Pad rows inside THIS rank's chunk (tail rows of the last rank)."""
|
||||||
|
return self.num_pad if self.sp_rank == self.sp_size - 1 else 0
|
||||||
|
|
||||||
|
@property
|
||||||
|
def local_real_len(self) -> int:
|
||||||
|
return self.local_len - self.local_pad
|
||||||
|
|
||||||
|
|
||||||
|
def build_shard_plan(seq_len: int) -> SpShard:
|
||||||
|
"""Shard math only; tensors are sliced separately via `shard_like`."""
|
||||||
|
sp_size = get_sp_world_size()
|
||||||
|
if sp_size <= 1:
|
||||||
|
return SpShard(seq_len, seq_len, 0, 1, 0)
|
||||||
|
local_len = (seq_len + sp_size - 1) // sp_size
|
||||||
|
return SpShard(
|
||||||
|
orig_len=seq_len,
|
||||||
|
local_len=local_len,
|
||||||
|
num_pad=local_len * sp_size - seq_len,
|
||||||
|
sp_size=sp_size,
|
||||||
|
sp_rank=get_sp_parallel_rank(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def shard_like(
|
||||||
|
x: torch.Tensor, shard: SpShard, dim: int = 1, pad_mode: str = "zeros"
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Apply a planned shard to one tensor (RoPE caches use the same plan as
|
||||||
|
hidden states so their chunks stay aligned)."""
|
||||||
|
if shard.sp_size <= 1:
|
||||||
|
return x
|
||||||
|
if shard.num_pad > 0:
|
||||||
|
if pad_mode == "repeat_last":
|
||||||
|
pad = x.narrow(dim, x.shape[dim] - 1, 1)
|
||||||
|
pad = pad.expand(
|
||||||
|
*[shard.num_pad if i == dim else -1 for i in range(x.dim())]
|
||||||
|
)
|
||||||
|
x = torch.cat([x, pad], dim=dim)
|
||||||
|
else:
|
||||||
|
# F.pad pads dims last-to-first: (left, right) pairs from dim -1.
|
||||||
|
pads = [0, 0] * (x.dim() - 1 - dim) + [0, shard.num_pad]
|
||||||
|
x = F.pad(x, pads)
|
||||||
|
return x.narrow(dim, shard.sp_rank * shard.local_len, shard.local_len)
|
||||||
|
|
||||||
|
|
||||||
|
def shard_seq(
|
||||||
|
x: torch.Tensor, dim: int = 1, pad_mode: str = "zeros"
|
||||||
|
) -> tuple[torch.Tensor, SpShard]:
|
||||||
|
"""
|
||||||
|
mode:
|
||||||
|
zeroes: pad with zeroes at tail
|
||||||
|
repeat_last: repeat the last token, only for rotary embedding
|
||||||
|
"""
|
||||||
|
shard = build_shard_plan(x.shape[dim])
|
||||||
|
return shard_like(x, shard, dim=dim, pad_mode=pad_mode), shard
|
||||||
|
|
||||||
|
|
||||||
|
def gather_seq(local: torch.Tensor, orig_len: int, dim: int = 1) -> torch.Tensor:
|
||||||
|
"""All-gather an SP-sharded sequence and trim the tail padding"""
|
||||||
|
if get_sp_world_size() <= 1:
|
||||||
|
return local
|
||||||
|
full = sequence_model_parallel_all_gather(local.contiguous(), dim=dim)
|
||||||
|
if full.shape[dim] > orig_len:
|
||||||
|
full = full.narrow(dim, 0, orig_len)
|
||||||
|
return full
|
||||||
|
|
||||||
|
|
||||||
|
def shard_seq_prefix(
|
||||||
|
x: torch.Tensor, prefix_len: int, shard: SpShard, dim: int = 0
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Shard only the leading ``prefix_len`` rows (e.g. the text segment of a
|
||||||
|
joint RoPE cache) with an existing plan; the remainder is kept as-is."""
|
||||||
|
rest = x.shape[dim] - prefix_len
|
||||||
|
return torch.cat(
|
||||||
|
[
|
||||||
|
shard_like(x.narrow(dim, 0, prefix_len), shard, dim=dim),
|
||||||
|
x.narrow(dim, prefix_len, rest),
|
||||||
|
],
|
||||||
|
dim=dim,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def join_seqs(
|
||||||
|
prefix: torch.Tensor, body: torch.Tensor, local_pad: int, dim: int = 1
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Concatenate local sharded ``[prefix (txt tokens, padding tokens), body (img tokens)]`` for joint attention, while relocating the
|
||||||
|
prefix's ``local_pad`` tail rows behind the body.
|
||||||
|
|
||||||
|
Why leave the padding at tail: the shard pads the *text* chunk, but the local joint layout is
|
||||||
|
[text, image].
|
||||||
|
|
||||||
|
In naive implementation, after the ulysses all-to-all, that pad would sit mid-sequence (of last rank)
|
||||||
|
([... txt_last, PAD, img_last]), which required further mem copy (for the padding tokens), inefficient in this case
|
||||||
|
|
||||||
|
With the pad relocated behind the image, the padding forms one global-tail block that the zero-copy varlen
|
||||||
|
path (tail_attn_meta, implemented in USPAttention.forward) skips for free
|
||||||
|
"""
|
||||||
|
if local_pad > 0:
|
||||||
|
real = prefix.shape[dim] - local_pad
|
||||||
|
return torch.cat(
|
||||||
|
[
|
||||||
|
# txt tokens
|
||||||
|
prefix.narrow(dim, 0, real),
|
||||||
|
body,
|
||||||
|
# leave the padding at global-tail
|
||||||
|
prefix.narrow(dim, real, local_pad),
|
||||||
|
],
|
||||||
|
dim=dim,
|
||||||
|
)
|
||||||
|
return torch.cat([prefix, body], dim=dim)
|
||||||
|
|
||||||
|
|
||||||
|
def split_seqs(
|
||||||
|
joint: torch.Tensor, prefix_len: int, local_pad: int, dim: int = 1
|
||||||
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Inverse of ``join_seqs``: recover ``(prefix, body)`` from the joint output, with the pad rows rejoining the prefix tail so the residual text
|
||||||
|
stream keeps its per-rank shape.
|
||||||
|
|
||||||
|
([... txt_last, PAD, img_last]) -> prefix (txt + pad), body (img)
|
||||||
|
"""
|
||||||
|
total = joint.shape[dim]
|
||||||
|
if local_pad > 0:
|
||||||
|
real = prefix_len - local_pad
|
||||||
|
body_end = total - local_pad
|
||||||
|
prefix = torch.cat(
|
||||||
|
[joint.narrow(dim, 0, real), joint.narrow(dim, body_end, local_pad)],
|
||||||
|
dim=dim,
|
||||||
|
)
|
||||||
|
return prefix, joint.narrow(dim, real, body_end - real)
|
||||||
|
return (
|
||||||
|
joint.narrow(dim, 0, prefix_len),
|
||||||
|
joint.narrow(dim, prefix_len, total - prefix_len),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def should_shard_text(txt_len: int) -> bool:
|
||||||
|
"""True when the joint-attention text stream should be SP-sharded here
|
||||||
|
(see plan_text_strategy for the policy)."""
|
||||||
|
return get_sp_world_size() > 1 and plan_text_strategy(txt_len) == "shard"
|
||||||
|
|
||||||
|
|
||||||
|
def tail_attn_meta(
|
||||||
|
shard: SpShard,
|
||||||
|
batch_size: int,
|
||||||
|
device: torch.device,
|
||||||
|
image_seq_len: int = 0,
|
||||||
|
) -> dict | None:
|
||||||
|
"""Per-request attention meta for a tail-padded shard: `cu_seqlens_tail`
|
||||||
|
splits each batch row into [valid | pad] varlen segments over the gathered
|
||||||
|
layout, so USPAttention runs varlen FA on the padded q/k/v with zero
|
||||||
|
repacking. Built once per request, reused by every block."""
|
||||||
|
if shard.sp_size <= 1 or shard.num_pad == 0:
|
||||||
|
return None
|
||||||
|
seq = shard.sp_size * (shard.local_len + image_seq_len)
|
||||||
|
valid = seq - shard.num_pad
|
||||||
|
row = torch.tensor([valid, shard.num_pad], dtype=torch.int32, device=device)
|
||||||
|
seglens = row.repeat(batch_size)
|
||||||
|
cu_seqlens = torch.zeros(2 * batch_size + 1, dtype=torch.int32, device=device)
|
||||||
|
cu_seqlens[1:] = torch.cumsum(seglens, dim=0)
|
||||||
|
return {
|
||||||
|
"pad_start": valid,
|
||||||
|
"pad_end": seq,
|
||||||
|
"local_pad": shard.local_pad,
|
||||||
|
"cu_seqlens_tail": cu_seqlens,
|
||||||
|
"max_seqlen_tail": max(valid, shard.num_pad),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def plan_text_strategy(txt_len: int) -> str:
|
||||||
|
"""Choose "shard" or "replicate" for the joint-attention text stream.
|
||||||
|
|
||||||
|
Prefer "shard" by default. for small sequence (shorter than SGLANG_SP_TEXT_SHARD_MIN), choose "replicate" for better performance
|
||||||
|
|
||||||
|
"""
|
||||||
|
sp_size = get_sp_world_size()
|
||||||
|
if sp_size <= 1:
|
||||||
|
return "replicate"
|
||||||
|
if txt_len % sp_size != 0 and get_ring_parallel_world_size() > 1:
|
||||||
|
return "replicate"
|
||||||
|
if txt_len < _TEXT_SHARD_MIN:
|
||||||
|
return "replicate"
|
||||||
|
return "shard"
|
||||||
@@ -533,7 +533,21 @@ class USPAttention(nn.Module):
|
|||||||
effective_skip_sp = (
|
effective_skip_sp = (
|
||||||
self.skip_sequence_parallel or skip_sequence_parallel_override
|
self.skip_sequence_parallel or skip_sequence_parallel_override
|
||||||
)
|
)
|
||||||
if attn_mask is not None:
|
# Tail-pad meta alone (sp_shard.tail_attn_meta; mask derivable from the
|
||||||
|
# pad span) also opts into the masked SP branch. gap_* = legacy alias.
|
||||||
|
meta_pad_start = meta_pad_end = None
|
||||||
|
if attn_mask_meta is not None:
|
||||||
|
meta_pad_start = attn_mask_meta.get(
|
||||||
|
"pad_start", attn_mask_meta.get("gap_start")
|
||||||
|
)
|
||||||
|
meta_pad_end = attn_mask_meta.get("pad_end", attn_mask_meta.get("gap_end"))
|
||||||
|
meta_only_pad = (
|
||||||
|
attn_mask is None
|
||||||
|
and meta_pad_start is not None
|
||||||
|
and not effective_skip_sp
|
||||||
|
and get_sequence_parallel_world_size() > 1
|
||||||
|
)
|
||||||
|
if attn_mask is not None or meta_only_pad:
|
||||||
|
|
||||||
def _prepare_sdpa_mask(
|
def _prepare_sdpa_mask(
|
||||||
mask: torch.Tensor, *, dtype: torch.dtype, device: torch.device
|
mask: torch.Tensor, *, dtype: torch.dtype, device: torch.device
|
||||||
@@ -626,7 +640,7 @@ class USPAttention(nn.Module):
|
|||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"USPAttention masked path does not support ring parallelism yet."
|
"USPAttention masked path does not support ring parallelism yet."
|
||||||
)
|
)
|
||||||
if attn_mask.dim() != 2:
|
if attn_mask is not None and attn_mask.dim() != 2:
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
"USPAttention masked SP path currently expects a [B, S_local] key mask."
|
"USPAttention masked SP path currently expects a [B, S_local] key mask."
|
||||||
)
|
)
|
||||||
@@ -637,26 +651,46 @@ class USPAttention(nn.Module):
|
|||||||
k = _usp_input_all_to_all(k, head_dim=2)
|
k = _usp_input_all_to_all(k, head_dim=2)
|
||||||
v = _usp_input_all_to_all(v, head_dim=2)
|
v = _usp_input_all_to_all(v, head_dim=2)
|
||||||
|
|
||||||
gap_start = None
|
|
||||||
gap_end = None
|
|
||||||
if attn_mask_meta is not None:
|
|
||||||
gap_start = attn_mask_meta.get("gap_start")
|
|
||||||
gap_end = attn_mask_meta.get("gap_end")
|
|
||||||
if (
|
if (
|
||||||
_VARLEN_FA_ENABLED
|
_VARLEN_FA_ENABLED
|
||||||
and self.backend == AttentionBackendEnum.FA
|
and self.backend == AttentionBackendEnum.FA
|
||||||
and gap_start is not None
|
and meta_pad_start is not None
|
||||||
and gap_end is not None
|
and meta_pad_end is not None
|
||||||
and gap_end > gap_start
|
and meta_pad_end > meta_pad_start
|
||||||
and q.device.type == "cuda"
|
and q.device.type == "cuda"
|
||||||
and q.dtype in (torch.float16, torch.bfloat16)
|
and q.dtype in (torch.float16, torch.bfloat16)
|
||||||
):
|
):
|
||||||
bs, seq = q.shape[0], q.shape[1]
|
bs, seq = q.shape[0], q.shape[1]
|
||||||
assert 0 <= gap_start < gap_end <= seq
|
assert 0 <= meta_pad_start < meta_pad_end <= seq
|
||||||
valid_seq = seq - (gap_end - gap_start)
|
cu_tail = attn_mask_meta.get("cu_seqlens_tail")
|
||||||
q_dense = torch.cat([q[:, :gap_start], q[:, gap_end:]], dim=1)
|
if cu_tail is not None and meta_pad_end == seq:
|
||||||
k_dense = torch.cat([k[:, :gap_start], k[:, gap_end:]], dim=1)
|
# Zero-copy tail path: run varlen FA straight over the
|
||||||
v_dense = torch.cat([v[:, :gap_start], v[:, gap_end:]], dim=1)
|
# padded layout, each row split into [valid | pad] segments
|
||||||
|
# (contiguous reshapes only, no repacking).
|
||||||
|
assert (
|
||||||
|
cu_tail.numel() == 2 * bs + 1
|
||||||
|
), "cu_seqlens_tail does not match the batch size"
|
||||||
|
out = flash_attn_varlen_func(
|
||||||
|
q=q.reshape(bs * seq, *q.shape[2:]),
|
||||||
|
k=k.reshape(bs * seq, *k.shape[2:]),
|
||||||
|
v=v.reshape(bs * seq, *v.shape[2:]),
|
||||||
|
cu_seqlens_q=cu_tail,
|
||||||
|
cu_seqlens_k=cu_tail,
|
||||||
|
max_seqlen_q=attn_mask_meta["max_seqlen_tail"],
|
||||||
|
max_seqlen_k=attn_mask_meta["max_seqlen_tail"],
|
||||||
|
softmax_scale=self.softmax_scale,
|
||||||
|
causal=False,
|
||||||
|
ver=_fa_backend.fa_ver,
|
||||||
|
).reshape(bs, seq, *q.shape[2:])
|
||||||
|
# Match the packed paths: masked query rows read as zeros.
|
||||||
|
out[:, meta_pad_start:].zero_()
|
||||||
|
if sp_size > 1:
|
||||||
|
out = _usp_output_all_to_all(out, head_dim=2)
|
||||||
|
return out
|
||||||
|
valid_seq = seq - (meta_pad_end - meta_pad_start)
|
||||||
|
q_dense = torch.cat([q[:, :meta_pad_start], q[:, meta_pad_end:]], dim=1)
|
||||||
|
k_dense = torch.cat([k[:, :meta_pad_start], k[:, meta_pad_end:]], dim=1)
|
||||||
|
v_dense = torch.cat([v[:, :meta_pad_start], v[:, meta_pad_end:]], dim=1)
|
||||||
cu_seqlens = torch.arange(
|
cu_seqlens = torch.arange(
|
||||||
0,
|
0,
|
||||||
(bs + 1) * valid_seq,
|
(bs + 1) * valid_seq,
|
||||||
@@ -677,10 +711,17 @@ class USPAttention(nn.Module):
|
|||||||
ver=_fa_backend.fa_ver,
|
ver=_fa_backend.fa_ver,
|
||||||
).reshape(bs, valid_seq, *q.shape[2:])
|
).reshape(bs, valid_seq, *q.shape[2:])
|
||||||
gap_out = out_dense.new_zeros(
|
gap_out = out_dense.new_zeros(
|
||||||
bs, gap_end - gap_start, out_dense.shape[2], out_dense.shape[3]
|
bs,
|
||||||
|
meta_pad_end - meta_pad_start,
|
||||||
|
out_dense.shape[2],
|
||||||
|
out_dense.shape[3],
|
||||||
)
|
)
|
||||||
out = torch.cat(
|
out = torch.cat(
|
||||||
[out_dense[:, :gap_start], gap_out, out_dense[:, gap_start:]],
|
[
|
||||||
|
out_dense[:, :meta_pad_start],
|
||||||
|
gap_out,
|
||||||
|
out_dense[:, meta_pad_start:],
|
||||||
|
],
|
||||||
dim=1,
|
dim=1,
|
||||||
)
|
)
|
||||||
if sp_size > 1:
|
if sp_size > 1:
|
||||||
@@ -691,9 +732,17 @@ class USPAttention(nn.Module):
|
|||||||
# attn_mask is inconsistent across SP ranks (None on some, Tensor on
|
# attn_mask is inconsistent across SP ranks (None on some, Tensor on
|
||||||
# others), which causes all_gather participant mismatch. Upstream
|
# others), which causes all_gather participant mismatch. Upstream
|
||||||
# mask builders must ensure all ranks produce the same mask type.
|
# mask builders must ensure all ranks produce the same mask type.
|
||||||
gathered_mask = sequence_model_parallel_all_gather(
|
if attn_mask is None:
|
||||||
attn_mask.contiguous(), dim=1
|
# Meta-only tail-pad caller on a non-FA fallback: the gathered
|
||||||
)
|
# mask is fully determined by the pad span, no collective needed.
|
||||||
|
gathered_mask = torch.ones(
|
||||||
|
q.shape[0], q.shape[1], dtype=torch.bool, device=q.device
|
||||||
|
)
|
||||||
|
gathered_mask[:, meta_pad_start:meta_pad_end] = False
|
||||||
|
else:
|
||||||
|
gathered_mask = sequence_model_parallel_all_gather(
|
||||||
|
attn_mask.contiguous(), dim=1
|
||||||
|
)
|
||||||
if (
|
if (
|
||||||
_VARLEN_FA_ENABLED
|
_VARLEN_FA_ENABLED
|
||||||
and self.backend == AttentionBackendEnum.FA
|
and self.backend == AttentionBackendEnum.FA
|
||||||
|
|||||||
@@ -134,10 +134,14 @@ class ErnieImageSelfAttention(nn.Module):
|
|||||||
self.norm_q = RMSNorm(head_dim, eps=eps)
|
self.norm_q = RMSNorm(head_dim, eps=eps)
|
||||||
self.norm_k = RMSNorm(head_dim, eps=eps)
|
self.norm_k = RMSNorm(head_dim, eps=eps)
|
||||||
|
|
||||||
|
# The joint [image, text] stream is fully replicated, so the ulysses
|
||||||
|
# all-to-all would wrongly treat it as sharded and duplicate it. Skip
|
||||||
|
# SP until the stream is sharded (sp_shard + num_replicated_suffix).
|
||||||
self.attn = USPAttention(
|
self.attn = USPAttention(
|
||||||
num_heads=self.num_local_heads,
|
num_heads=self.num_local_heads,
|
||||||
head_size=head_dim,
|
head_size=head_dim,
|
||||||
prefix=f"{prefix}.attn",
|
prefix=f"{prefix}.attn",
|
||||||
|
skip_sequence_parallel=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
|
|||||||
@@ -30,10 +30,17 @@ from torch.nn import LayerNorm as LayerNorm
|
|||||||
from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
|
from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
|
||||||
from sglang.multimodal_gen.runtime.distributed import (
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
divide,
|
divide,
|
||||||
get_sp_parallel_rank,
|
|
||||||
get_sp_world_size,
|
|
||||||
get_tp_world_size,
|
get_tp_world_size,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import (
|
||||||
|
build_shard_plan,
|
||||||
|
join_seqs,
|
||||||
|
shard_like,
|
||||||
|
shard_seq_prefix,
|
||||||
|
should_shard_text,
|
||||||
|
split_seqs,
|
||||||
|
tail_attn_meta,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
||||||
from sglang.multimodal_gen.runtime.layers.layernorm import (
|
from sglang.multimodal_gen.runtime.layers.layernorm import (
|
||||||
RMSNorm,
|
RMSNorm,
|
||||||
@@ -69,99 +76,6 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
|||||||
logger = init_logger(__name__) # pylint: disable=invalid-name
|
logger = init_logger(__name__) # pylint: disable=invalid-name
|
||||||
|
|
||||||
|
|
||||||
def _shard_text_for_sp(
|
|
||||||
encoder_hidden_states: torch.Tensor,
|
|
||||||
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
|
||||||
image_seq_len: int,
|
|
||||||
num_txt_tokens: int,
|
|
||||||
) -> Tuple[
|
|
||||||
torch.Tensor,
|
|
||||||
Optional[Tuple[torch.Tensor, torch.Tensor]],
|
|
||||||
int,
|
|
||||||
Optional[torch.Tensor],
|
|
||||||
Optional[Dict[str, int]],
|
|
||||||
]:
|
|
||||||
sp_size = get_sp_world_size()
|
|
||||||
num_replicated_prefix = num_txt_tokens
|
|
||||||
if sp_size == 1:
|
|
||||||
return encoder_hidden_states, freqs_cis, num_replicated_prefix, None, None
|
|
||||||
|
|
||||||
sp_rank = get_sp_parallel_rank()
|
|
||||||
local_txt_tokens = (num_txt_tokens + sp_size - 1) // sp_size
|
|
||||||
padded_txt_tokens = local_txt_tokens * sp_size
|
|
||||||
num_pad_tokens = padded_txt_tokens - num_txt_tokens
|
|
||||||
|
|
||||||
if num_pad_tokens > 0:
|
|
||||||
pad_hidden_states = encoder_hidden_states.new_zeros(
|
|
||||||
encoder_hidden_states.shape[0],
|
|
||||||
num_pad_tokens,
|
|
||||||
encoder_hidden_states.shape[2],
|
|
||||||
)
|
|
||||||
encoder_hidden_states = torch.cat(
|
|
||||||
[encoder_hidden_states, pad_hidden_states], dim=1
|
|
||||||
)
|
|
||||||
|
|
||||||
encoder_hidden_states = torch.chunk(encoder_hidden_states, sp_size, dim=1)[sp_rank]
|
|
||||||
if freqs_cis is not None:
|
|
||||||
cos, sin = freqs_cis
|
|
||||||
txt_cos = cos[:num_txt_tokens]
|
|
||||||
txt_sin = sin[:num_txt_tokens]
|
|
||||||
if num_pad_tokens > 0:
|
|
||||||
pad_cos = txt_cos.new_ones(num_pad_tokens, txt_cos.shape[1])
|
|
||||||
pad_sin = txt_sin.new_zeros(num_pad_tokens, txt_sin.shape[1])
|
|
||||||
txt_cos = torch.cat([txt_cos, pad_cos], dim=0)
|
|
||||||
txt_sin = torch.cat([txt_sin, pad_sin], dim=0)
|
|
||||||
freqs_cis = (
|
|
||||||
torch.cat(
|
|
||||||
[
|
|
||||||
torch.chunk(txt_cos, sp_size, dim=0)[sp_rank],
|
|
||||||
cos[num_txt_tokens:],
|
|
||||||
],
|
|
||||||
dim=0,
|
|
||||||
),
|
|
||||||
torch.cat(
|
|
||||||
[
|
|
||||||
torch.chunk(txt_sin, sp_size, dim=0)[sp_rank],
|
|
||||||
sin[num_txt_tokens:],
|
|
||||||
],
|
|
||||||
dim=0,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
num_replicated_prefix = 0
|
|
||||||
if num_pad_tokens == 0:
|
|
||||||
return encoder_hidden_states, freqs_cis, num_replicated_prefix, None, None
|
|
||||||
|
|
||||||
txt_start = sp_rank * local_txt_tokens
|
|
||||||
valid_txt_tokens = min(local_txt_tokens, max(num_txt_tokens - txt_start, 0))
|
|
||||||
text_mask = torch.zeros(
|
|
||||||
encoder_hidden_states.shape[0],
|
|
||||||
local_txt_tokens,
|
|
||||||
dtype=torch.bool,
|
|
||||||
device=encoder_hidden_states.device,
|
|
||||||
)
|
|
||||||
text_mask[:, :valid_txt_tokens] = True
|
|
||||||
image_mask = torch.ones(
|
|
||||||
encoder_hidden_states.shape[0],
|
|
||||||
image_seq_len,
|
|
||||||
dtype=torch.bool,
|
|
||||||
device=encoder_hidden_states.device,
|
|
||||||
)
|
|
||||||
return (
|
|
||||||
encoder_hidden_states,
|
|
||||||
freqs_cis,
|
|
||||||
num_replicated_prefix,
|
|
||||||
torch.cat([text_mask, image_mask], dim=1),
|
|
||||||
{
|
|
||||||
"gap_start": (sp_size - 1) * (local_txt_tokens + image_seq_len)
|
|
||||||
+ local_txt_tokens
|
|
||||||
- num_pad_tokens,
|
|
||||||
"gap_end": (sp_size - 1) * (local_txt_tokens + image_seq_len)
|
|
||||||
+ local_txt_tokens,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import]
|
from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import]
|
||||||
from nunchaku.models.normalization import ( # type: ignore[import]
|
from nunchaku.models.normalization import ( # type: ignore[import]
|
||||||
@@ -600,9 +514,12 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
allow_inplace=True,
|
allow_inplace=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
query = torch.cat([encoder_query, query], dim=1)
|
# join_seqs relocates any SP text tail-pad behind the image (see
|
||||||
key = torch.cat([encoder_key, key], dim=1)
|
# sp_shard.join_seqs for why).
|
||||||
value = torch.cat([encoder_value, value], dim=1)
|
sp_txt_pad = (attn_mask_meta or {}).get("local_pad", 0)
|
||||||
|
query = join_seqs(encoder_query, query, sp_txt_pad)
|
||||||
|
key = join_seqs(encoder_key, key, sp_txt_pad)
|
||||||
|
value = join_seqs(encoder_value, value, sp_txt_pad)
|
||||||
else:
|
else:
|
||||||
query, key = apply_qk_norm_with_optional_rope(
|
query, key = apply_qk_norm_with_optional_rope(
|
||||||
q=query,
|
q=query,
|
||||||
@@ -627,12 +544,8 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
x = x.to(query.dtype)
|
x = x.to(query.dtype)
|
||||||
|
|
||||||
if encoder_hidden_states is not None:
|
if encoder_hidden_states is not None:
|
||||||
encoder_hidden_states, x = x.split_with_sizes(
|
encoder_hidden_states, x = split_seqs(
|
||||||
[
|
x, encoder_hidden_states.shape[1], sp_txt_pad
|
||||||
encoder_hidden_states.shape[1],
|
|
||||||
x.shape[1] - encoder_hidden_states.shape[1],
|
|
||||||
],
|
|
||||||
dim=1,
|
|
||||||
)
|
)
|
||||||
if not self.pre_only:
|
if not self.pre_only:
|
||||||
x, _ = self.to_out[0](x)
|
x, _ = self.to_out[0](x)
|
||||||
@@ -775,11 +688,16 @@ class FluxSingleTransformerBlock(nn.Module):
|
|||||||
num_replicated_prefix: int = 0,
|
num_replicated_prefix: int = 0,
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
text_seq_len = encoder_hidden_states.shape[1]
|
text_seq_len = encoder_hidden_states.shape[1]
|
||||||
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
|
joint_attention_kwargs = joint_attention_kwargs or {}
|
||||||
|
# join_seqs relocates any SP text tail-pad behind the image; the caller
|
||||||
|
# hands single blocks a RoPE cache reordered the same way.
|
||||||
|
sp_txt_pad = (joint_attention_kwargs.get("attn_mask_meta") or {}).get(
|
||||||
|
"local_pad", 0
|
||||||
|
)
|
||||||
|
hidden_states = join_seqs(encoder_hidden_states, hidden_states, sp_txt_pad)
|
||||||
|
|
||||||
residual = hidden_states
|
residual = hidden_states
|
||||||
norm_hidden_states, gate = self.norm(hidden_states, emb=temb)
|
norm_hidden_states, gate = self.norm(hidden_states, emb=temb)
|
||||||
joint_attention_kwargs = joint_attention_kwargs or {}
|
|
||||||
|
|
||||||
if self.use_nunchaku_structure:
|
if self.use_nunchaku_structure:
|
||||||
if _nunchaku_fused_ops_available:
|
if _nunchaku_fused_ops_available:
|
||||||
@@ -824,9 +742,8 @@ class FluxSingleTransformerBlock(nn.Module):
|
|||||||
if hidden_states.dtype == torch.float16:
|
if hidden_states.dtype == torch.float16:
|
||||||
hidden_states = hidden_states.clip(-65504, 65504)
|
hidden_states = hidden_states.clip(-65504, 65504)
|
||||||
|
|
||||||
encoder_hidden_states, hidden_states = (
|
encoder_hidden_states, hidden_states = split_seqs(
|
||||||
hidden_states[:, :text_seq_len],
|
hidden_states, text_seq_len, sp_txt_pad
|
||||||
hidden_states[:, text_seq_len:],
|
|
||||||
)
|
)
|
||||||
return encoder_hidden_states, hidden_states
|
return encoder_hidden_states, hidden_states
|
||||||
|
|
||||||
@@ -1186,24 +1103,41 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
num_txt_tokens = encoder_hidden_states.shape[1]
|
num_txt_tokens = encoder_hidden_states.shape[1]
|
||||||
encoder_hidden_states, _ = self.context_embedder(encoder_hidden_states)
|
encoder_hidden_states, _ = self.context_embedder(encoder_hidden_states)
|
||||||
|
|
||||||
(
|
# Shard the replicated text stream across SP ranks (image latents are
|
||||||
encoder_hidden_states,
|
# already sharded); non-divisible lengths tail-pad the last rank and the
|
||||||
freqs_cis,
|
# per-request tail meta lets attention skip the pad for free.
|
||||||
num_replicated_prefix,
|
num_replicated_prefix = num_txt_tokens
|
||||||
attn_mask,
|
singles_freqs_cis = freqs_cis
|
||||||
attn_mask_meta,
|
if should_shard_text(num_txt_tokens):
|
||||||
) = _shard_text_for_sp(
|
txt_shard = build_shard_plan(num_txt_tokens)
|
||||||
encoder_hidden_states,
|
encoder_hidden_states = shard_like(encoder_hidden_states, txt_shard)
|
||||||
freqs_cis,
|
if freqs_cis is not None:
|
||||||
hidden_states.shape[1],
|
cos, sin = freqs_cis
|
||||||
num_txt_tokens,
|
cos = shard_seq_prefix(cos, num_txt_tokens, txt_shard)
|
||||||
)
|
sin = shard_seq_prefix(sin, num_txt_tokens, txt_shard)
|
||||||
if attn_mask is not None:
|
freqs_cis = (cos, sin)
|
||||||
joint_attention_kwargs = (
|
singles_freqs_cis = freqs_cis
|
||||||
joint_attention_kwargs.copy() if joint_attention_kwargs else {}
|
num_replicated_prefix = 0
|
||||||
|
tail_meta = tail_attn_meta(
|
||||||
|
txt_shard,
|
||||||
|
encoder_hidden_states.shape[0],
|
||||||
|
hidden_states.device,
|
||||||
|
image_seq_len=hidden_states.shape[1],
|
||||||
)
|
)
|
||||||
joint_attention_kwargs["attn_mask"] = attn_mask
|
if tail_meta is not None:
|
||||||
joint_attention_kwargs["attn_mask_meta"] = attn_mask_meta
|
joint_attention_kwargs = (
|
||||||
|
joint_attention_kwargs.copy() if joint_attention_kwargs else {}
|
||||||
|
)
|
||||||
|
joint_attention_kwargs["attn_mask_meta"] = tail_meta
|
||||||
|
# Single blocks apply RoPE on the relocated [txt_real, img, pad]
|
||||||
|
# layout, so hand them a cache reordered the same way.
|
||||||
|
if freqs_cis is not None:
|
||||||
|
t_loc = txt_shard.local_len
|
||||||
|
pad = txt_shard.local_pad
|
||||||
|
singles_freqs_cis = (
|
||||||
|
join_seqs(cos[:t_loc], cos[t_loc:], pad, dim=0),
|
||||||
|
join_seqs(sin[:t_loc], sin[t_loc:], pad, dim=0),
|
||||||
|
)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
joint_attention_kwargs is not None
|
joint_attention_kwargs is not None
|
||||||
@@ -1229,7 +1163,7 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
encoder_hidden_states=encoder_hidden_states,
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
temb=temb,
|
temb=temb,
|
||||||
freqs_cis=freqs_cis,
|
freqs_cis=singles_freqs_cis,
|
||||||
joint_attention_kwargs=joint_attention_kwargs,
|
joint_attention_kwargs=joint_attention_kwargs,
|
||||||
num_replicated_prefix=num_replicated_prefix,
|
num_replicated_prefix=num_replicated_prefix,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -23,10 +23,17 @@ from diffusers.models.normalization import AdaLayerNormContinuous
|
|||||||
from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
|
from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
|
||||||
from sglang.multimodal_gen.runtime.distributed import (
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
divide,
|
divide,
|
||||||
get_sp_parallel_rank,
|
|
||||||
get_sp_world_size,
|
|
||||||
get_tp_world_size,
|
get_tp_world_size,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import (
|
||||||
|
build_shard_plan,
|
||||||
|
join_seqs,
|
||||||
|
shard_like,
|
||||||
|
shard_seq_prefix,
|
||||||
|
should_shard_text,
|
||||||
|
split_seqs,
|
||||||
|
tail_attn_meta,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
||||||
from sglang.multimodal_gen.runtime.layers.layernorm import (
|
from sglang.multimodal_gen.runtime.layers.layernorm import (
|
||||||
RMSNorm,
|
RMSNorm,
|
||||||
@@ -60,115 +67,6 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
|||||||
logger = init_logger(__name__) # pylint: disable=invalid-name
|
logger = init_logger(__name__) # pylint: disable=invalid-name
|
||||||
|
|
||||||
|
|
||||||
def _shard_text_for_sp(
|
|
||||||
encoder_hidden_states: torch.Tensor,
|
|
||||||
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
|
||||||
image_seq_len: int,
|
|
||||||
num_txt_tokens: int,
|
|
||||||
) -> Tuple[
|
|
||||||
torch.Tensor,
|
|
||||||
Optional[Tuple[torch.Tensor, torch.Tensor]],
|
|
||||||
int,
|
|
||||||
int,
|
|
||||||
Optional[torch.Tensor],
|
|
||||||
Optional[Dict[str, int]],
|
|
||||||
]:
|
|
||||||
sp_size = get_sp_world_size()
|
|
||||||
num_replicated_prefix = num_txt_tokens
|
|
||||||
if sp_size == 1:
|
|
||||||
return (
|
|
||||||
encoder_hidden_states,
|
|
||||||
freqs_cis,
|
|
||||||
num_replicated_prefix,
|
|
||||||
num_txt_tokens,
|
|
||||||
None,
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
|
|
||||||
sp_rank = get_sp_parallel_rank()
|
|
||||||
local_txt_tokens = (num_txt_tokens + sp_size - 1) // sp_size
|
|
||||||
padded_txt_tokens = local_txt_tokens * sp_size
|
|
||||||
num_pad_tokens = padded_txt_tokens - num_txt_tokens
|
|
||||||
|
|
||||||
if num_pad_tokens > 0:
|
|
||||||
pad_hidden_states = encoder_hidden_states.new_zeros(
|
|
||||||
encoder_hidden_states.shape[0],
|
|
||||||
num_pad_tokens,
|
|
||||||
encoder_hidden_states.shape[2],
|
|
||||||
)
|
|
||||||
encoder_hidden_states = torch.cat(
|
|
||||||
[encoder_hidden_states, pad_hidden_states], dim=1
|
|
||||||
)
|
|
||||||
|
|
||||||
encoder_hidden_states = torch.chunk(encoder_hidden_states, sp_size, dim=1)[sp_rank]
|
|
||||||
if freqs_cis is not None:
|
|
||||||
cos, sin = freqs_cis
|
|
||||||
txt_cos = cos[:num_txt_tokens]
|
|
||||||
txt_sin = sin[:num_txt_tokens]
|
|
||||||
if num_pad_tokens > 0:
|
|
||||||
pad_cos = txt_cos.new_ones(num_pad_tokens, txt_cos.shape[1])
|
|
||||||
pad_sin = txt_sin.new_zeros(num_pad_tokens, txt_sin.shape[1])
|
|
||||||
txt_cos = torch.cat([txt_cos, pad_cos], dim=0)
|
|
||||||
txt_sin = torch.cat([txt_sin, pad_sin], dim=0)
|
|
||||||
freqs_cis = (
|
|
||||||
torch.cat(
|
|
||||||
[
|
|
||||||
torch.chunk(txt_cos, sp_size, dim=0)[sp_rank],
|
|
||||||
cos[num_txt_tokens:],
|
|
||||||
],
|
|
||||||
dim=0,
|
|
||||||
),
|
|
||||||
torch.cat(
|
|
||||||
[
|
|
||||||
torch.chunk(txt_sin, sp_size, dim=0)[sp_rank],
|
|
||||||
sin[num_txt_tokens:],
|
|
||||||
],
|
|
||||||
dim=0,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
num_replicated_prefix = 0
|
|
||||||
if num_pad_tokens == 0:
|
|
||||||
return (
|
|
||||||
encoder_hidden_states,
|
|
||||||
freqs_cis,
|
|
||||||
num_replicated_prefix,
|
|
||||||
local_txt_tokens,
|
|
||||||
None,
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
|
|
||||||
txt_start = sp_rank * local_txt_tokens
|
|
||||||
valid_txt_tokens = min(local_txt_tokens, max(num_txt_tokens - txt_start, 0))
|
|
||||||
text_mask = torch.zeros(
|
|
||||||
encoder_hidden_states.shape[0],
|
|
||||||
local_txt_tokens,
|
|
||||||
dtype=torch.bool,
|
|
||||||
device=encoder_hidden_states.device,
|
|
||||||
)
|
|
||||||
text_mask[:, :valid_txt_tokens] = True
|
|
||||||
image_mask = torch.ones(
|
|
||||||
encoder_hidden_states.shape[0],
|
|
||||||
image_seq_len,
|
|
||||||
dtype=torch.bool,
|
|
||||||
device=encoder_hidden_states.device,
|
|
||||||
)
|
|
||||||
return (
|
|
||||||
encoder_hidden_states,
|
|
||||||
freqs_cis,
|
|
||||||
num_replicated_prefix,
|
|
||||||
local_txt_tokens,
|
|
||||||
torch.cat([text_mask, image_mask], dim=1),
|
|
||||||
{
|
|
||||||
"gap_start": (sp_size - 1) * (local_txt_tokens + image_seq_len)
|
|
||||||
+ local_txt_tokens
|
|
||||||
- num_pad_tokens,
|
|
||||||
"gap_end": (sp_size - 1) * (local_txt_tokens + image_seq_len)
|
|
||||||
+ local_txt_tokens,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _get_qkv_projections(
|
def _get_qkv_projections(
|
||||||
attn: "Flux2Attention", hidden_states, encoder_hidden_states=None
|
attn: "Flux2Attention", hidden_states, encoder_hidden_states=None
|
||||||
):
|
):
|
||||||
@@ -465,9 +363,12 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
allow_inplace=True,
|
allow_inplace=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
query = torch.cat([encoder_query, query], dim=1)
|
# join_seqs relocates any SP text tail-pad behind the image (see
|
||||||
key = torch.cat([encoder_key, key], dim=1)
|
# sp_shard.join_seqs for why).
|
||||||
value = torch.cat([encoder_value, value], dim=1)
|
sp_txt_pad = (attn_mask_meta or {}).get("local_pad", 0)
|
||||||
|
query = join_seqs(encoder_query, query, sp_txt_pad)
|
||||||
|
key = join_seqs(encoder_key, key, sp_txt_pad)
|
||||||
|
value = join_seqs(encoder_value, value, sp_txt_pad)
|
||||||
else:
|
else:
|
||||||
query, key = apply_qk_norm_with_optional_rope(
|
query, key = apply_qk_norm_with_optional_rope(
|
||||||
q=query,
|
q=query,
|
||||||
@@ -493,12 +394,8 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
hidden_states = hidden_states.to(query.dtype)
|
hidden_states = hidden_states.to(query.dtype)
|
||||||
|
|
||||||
if encoder_hidden_states is not None:
|
if encoder_hidden_states is not None:
|
||||||
encoder_hidden_states, hidden_states = hidden_states.split_with_sizes(
|
encoder_hidden_states, hidden_states = split_seqs(
|
||||||
[
|
hidden_states, encoder_hidden_states.shape[1], sp_txt_pad
|
||||||
encoder_hidden_states.shape[1],
|
|
||||||
hidden_states.shape[1] - encoder_hidden_states.shape[1],
|
|
||||||
],
|
|
||||||
dim=1,
|
|
||||||
)
|
)
|
||||||
encoder_hidden_states, _ = self.to_add_out(encoder_hidden_states)
|
encoder_hidden_states, _ = self.to_add_out(encoder_hidden_states)
|
||||||
|
|
||||||
@@ -1196,25 +1093,43 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
hidden_states, _ = self.x_embedder(hidden_states)
|
hidden_states, _ = self.x_embedder(hidden_states)
|
||||||
encoder_hidden_states, _ = self.context_embedder(encoder_hidden_states)
|
encoder_hidden_states, _ = self.context_embedder(encoder_hidden_states)
|
||||||
|
|
||||||
(
|
# Shard the replicated text stream across SP ranks (image latents are
|
||||||
encoder_hidden_states,
|
# already sharded); non-divisible lengths tail-pad the last rank and the
|
||||||
freqs_cis,
|
# per-request tail meta lets attention skip the pad for free.
|
||||||
num_replicated_prefix,
|
num_replicated_prefix = num_txt_tokens
|
||||||
num_txt_tokens,
|
sp_txt_pad = 0
|
||||||
attn_mask,
|
singles_freqs_cis = freqs_cis
|
||||||
attn_mask_meta,
|
if should_shard_text(num_txt_tokens):
|
||||||
) = _shard_text_for_sp(
|
txt_shard = build_shard_plan(num_txt_tokens)
|
||||||
encoder_hidden_states,
|
encoder_hidden_states = shard_like(encoder_hidden_states, txt_shard)
|
||||||
freqs_cis,
|
if freqs_cis is not None:
|
||||||
hidden_states.shape[1],
|
cos, sin = freqs_cis
|
||||||
num_txt_tokens,
|
cos = shard_seq_prefix(cos, num_txt_tokens, txt_shard)
|
||||||
)
|
sin = shard_seq_prefix(sin, num_txt_tokens, txt_shard)
|
||||||
if attn_mask is not None:
|
freqs_cis = (cos, sin)
|
||||||
joint_attention_kwargs = (
|
singles_freqs_cis = freqs_cis
|
||||||
joint_attention_kwargs.copy() if joint_attention_kwargs else {}
|
num_replicated_prefix = 0
|
||||||
|
num_txt_tokens = txt_shard.local_len
|
||||||
|
tail_meta = tail_attn_meta(
|
||||||
|
txt_shard,
|
||||||
|
encoder_hidden_states.shape[0],
|
||||||
|
hidden_states.device,
|
||||||
|
image_seq_len=hidden_states.shape[1],
|
||||||
)
|
)
|
||||||
joint_attention_kwargs["attn_mask"] = attn_mask
|
if tail_meta is not None:
|
||||||
joint_attention_kwargs["attn_mask_meta"] = attn_mask_meta
|
joint_attention_kwargs = (
|
||||||
|
joint_attention_kwargs.copy() if joint_attention_kwargs else {}
|
||||||
|
)
|
||||||
|
joint_attention_kwargs["attn_mask_meta"] = tail_meta
|
||||||
|
sp_txt_pad = txt_shard.local_pad
|
||||||
|
# The single-stream trunk applies RoPE on the relocated
|
||||||
|
# [txt_real, img, pad] layout; reorder its cache to match.
|
||||||
|
if freqs_cis is not None:
|
||||||
|
t_loc = txt_shard.local_len
|
||||||
|
singles_freqs_cis = (
|
||||||
|
join_seqs(cos[:t_loc], cos[t_loc:], sp_txt_pad, dim=0),
|
||||||
|
join_seqs(sin[:t_loc], sin[t_loc:], sp_txt_pad, dim=0),
|
||||||
|
)
|
||||||
|
|
||||||
# 4. Double Stream Transformer Blocks
|
# 4. Double Stream Transformer Blocks
|
||||||
for index_block, block in enumerate(self.transformer_blocks):
|
for index_block, block in enumerate(self.transformer_blocks):
|
||||||
@@ -1227,8 +1142,11 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
joint_attention_kwargs=joint_attention_kwargs,
|
joint_attention_kwargs=joint_attention_kwargs,
|
||||||
num_replicated_prefix=num_replicated_prefix,
|
num_replicated_prefix=num_replicated_prefix,
|
||||||
)
|
)
|
||||||
# Concatenate text and image streams for single-block inference
|
# Concatenate text and image streams for single-block inference;
|
||||||
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
|
# join_seqs relocates any SP text tail-pad behind the image once for
|
||||||
|
# the whole trunk (see sp_shard.join_seqs for why).
|
||||||
|
txt_real = num_txt_tokens - sp_txt_pad
|
||||||
|
hidden_states = join_seqs(encoder_hidden_states, hidden_states, sp_txt_pad)
|
||||||
|
|
||||||
# 5. Single Stream Transformer Blocks
|
# 5. Single Stream Transformer Blocks
|
||||||
for index_block, block in enumerate(self.single_transformer_blocks):
|
for index_block, block in enumerate(self.single_transformer_blocks):
|
||||||
@@ -1236,13 +1154,14 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
encoder_hidden_states=None,
|
encoder_hidden_states=None,
|
||||||
temb_mod_params=single_stream_mod,
|
temb_mod_params=single_stream_mod,
|
||||||
freqs_cis=freqs_cis,
|
freqs_cis=singles_freqs_cis,
|
||||||
joint_attention_kwargs=joint_attention_kwargs,
|
joint_attention_kwargs=joint_attention_kwargs,
|
||||||
text_seq_len=num_txt_tokens,
|
text_seq_len=txt_real,
|
||||||
num_replicated_prefix=num_replicated_prefix,
|
num_replicated_prefix=num_replicated_prefix,
|
||||||
)
|
)
|
||||||
# Remove text tokens from concatenated stream
|
# Remove text (and any tail pad) from the concatenated stream
|
||||||
hidden_states = hidden_states[:, num_txt_tokens:, ...]
|
img_end = hidden_states.shape[1] - sp_txt_pad
|
||||||
|
hidden_states = hidden_states[:, txt_real:img_end, ...]
|
||||||
|
|
||||||
# 6. Output layers
|
# 6. Output layers
|
||||||
hidden_states = self.norm_out(hidden_states, temb)
|
hidden_states = self.norm_out(hidden_states, temb)
|
||||||
|
|||||||
@@ -153,13 +153,14 @@ class SelfAttention(nn.Module):
|
|||||||
softmax_scale=None,
|
softmax_scale=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x, freqs):
|
def forward(self, x, freqs, attn_mask_meta=None):
|
||||||
"""
|
"""
|
||||||
Forward pass for self-attention.
|
Forward pass for self-attention.
|
||||||
|
|
||||||
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, 1, head_dim] - should match x's sequence length
|
||||||
|
attn_mask_meta: sp_shard tail-pad meta; excludes SP padding from attention
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Output tensor [B, S_local, D]
|
Output tensor [B, S_local, D]
|
||||||
@@ -189,8 +190,9 @@ class SelfAttention(nn.Module):
|
|||||||
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)
|
||||||
|
|
||||||
# USPAttention handles SP communication internally
|
# USPAttention handles SP communication internally; the tail meta keeps
|
||||||
out = self.attn(q, k, v)
|
# SP padding out of the softmax.
|
||||||
|
out = self.attn(q, k, v, attn_mask_meta=attn_mask_meta)
|
||||||
out = rearrange(out, "b s n d -> b s (n d)")
|
out = rearrange(out, "b s n d -> b s (n d)")
|
||||||
|
|
||||||
out, _ = self.o(out)
|
out, _ = self.o(out)
|
||||||
@@ -326,7 +328,7 @@ class DiTBlock(nn.Module):
|
|||||||
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||||
self.mlp_residual = MulAdd()
|
self.mlp_residual = MulAdd()
|
||||||
|
|
||||||
def forward(self, x, context, t_mod, freqs):
|
def forward(self, x, context, t_mod, freqs, attn_mask_meta=None):
|
||||||
has_seq = len(t_mod.shape) == 4
|
has_seq = len(t_mod.shape) == 4
|
||||||
chunk_dim = 2 if has_seq else 1
|
chunk_dim = 2 if has_seq else 1
|
||||||
# msa: multi-head self-attention mlp: multi-layer perceptron
|
# msa: multi-head self-attention mlp: multi-layer perceptron
|
||||||
@@ -347,7 +349,9 @@ class DiTBlock(nn.Module):
|
|||||||
# - layernorm(x) * (1 + scale_msa) + shift_msa
|
# - layernorm(x) * (1 + scale_msa) + shift_msa
|
||||||
input_x = self.norm1(x, shift_msa, scale_msa)
|
input_x = self.norm1(x, shift_msa, scale_msa)
|
||||||
# 2. torch.compile may fuse mlp_residual and self_attn_norm
|
# 2. torch.compile may fuse mlp_residual and self_attn_norm
|
||||||
x = self.mlp_residual(self.self_attn(input_x, freqs), gate_msa, x)
|
x = self.mlp_residual(
|
||||||
|
self.self_attn(input_x, freqs, attn_mask_meta=attn_mask_meta), gate_msa, x
|
||||||
|
)
|
||||||
norm_x = self.self_attn_norm(x)
|
norm_x = self.self_attn_norm(x)
|
||||||
# 3. Cross-attention, fuse:
|
# 3. Cross-attention, fuse:
|
||||||
# - x = x + 1 * cross_output
|
# - x = x + 1 * cross_output
|
||||||
|
|||||||
@@ -20,10 +20,16 @@ from sglang.multimodal_gen.runtime.distributed import (
|
|||||||
get_tp_world_size,
|
get_tp_world_size,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
get_ring_parallel_world_size,
|
|
||||||
get_sp_parallel_rank,
|
|
||||||
get_sp_world_size,
|
get_sp_world_size,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import (
|
||||||
|
build_shard_plan,
|
||||||
|
join_seqs,
|
||||||
|
shard_like,
|
||||||
|
should_shard_text,
|
||||||
|
split_seqs,
|
||||||
|
tail_attn_meta,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.attention import (
|
from sglang.multimodal_gen.runtime.layers.attention import (
|
||||||
USPAttention,
|
USPAttention,
|
||||||
build_varlen_mask_meta,
|
build_varlen_mask_meta,
|
||||||
@@ -80,99 +86,6 @@ def _local_seq_len(seq_len: int, sp_world_size: int) -> int:
|
|||||||
return padded_len // sp_world_size
|
return padded_len // sp_world_size
|
||||||
|
|
||||||
|
|
||||||
def _shard_text_for_sp(
|
|
||||||
encoder_hidden_states: torch.Tensor,
|
|
||||||
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
|
||||||
) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]:
|
|
||||||
"""Shard the replicated text stream evenly across SP ranks.
|
|
||||||
|
|
||||||
The image latents are already sharded by the pipeline while the text stream
|
|
||||||
is replicated. This splits the text embeddings (and their RoPE cache) so each
|
|
||||||
rank keeps ``1/sp_size`` of the text tokens, making the joint attention fully
|
|
||||||
sequence-sharded (``num_replicated_prefix=0``). Callers must ensure the text
|
|
||||||
length divides evenly across SP ranks.
|
|
||||||
"""
|
|
||||||
sp_size = get_sp_world_size()
|
|
||||||
if sp_size == 1:
|
|
||||||
return encoder_hidden_states, freqs_cis
|
|
||||||
|
|
||||||
sp_rank = get_sp_parallel_rank()
|
|
||||||
encoder_hidden_states = torch.chunk(encoder_hidden_states, sp_size, dim=1)[sp_rank]
|
|
||||||
|
|
||||||
if freqs_cis is not None:
|
|
||||||
img_cache, txt_cache = freqs_cis
|
|
||||||
txt_cache = torch.chunk(txt_cache, sp_size, dim=0)[sp_rank]
|
|
||||||
freqs_cis = (img_cache, txt_cache)
|
|
||||||
|
|
||||||
return encoder_hidden_states, freqs_cis
|
|
||||||
|
|
||||||
|
|
||||||
def _pad_shard_text_for_sp_varlen(
|
|
||||||
encoder_hidden_states: torch.Tensor,
|
|
||||||
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
|
||||||
image_seq_len: int,
|
|
||||||
) -> Tuple[
|
|
||||||
torch.Tensor,
|
|
||||||
Optional[Tuple[torch.Tensor, torch.Tensor]],
|
|
||||||
torch.Tensor,
|
|
||||||
Dict[str, int],
|
|
||||||
]:
|
|
||||||
"""Right-pad a non-divisible replicated text stream to a multiple of the SP
|
|
||||||
world size and shard it evenly across ranks, so the joint ``[text, image]``
|
|
||||||
sequence is fully sequence-parallel.
|
|
||||||
|
|
||||||
The pad tokens occupy a single contiguous block at the tail of the last
|
|
||||||
rank's text chunk. The returned ``attn_mask`` (joint ``[text, image]``
|
|
||||||
validity mask) and ``attn_mask_meta`` (``gap_start`` / ``gap_end``) describe
|
|
||||||
that block so ``USPAttention`` excludes it from attention via the varlen
|
|
||||||
kernel.
|
|
||||||
|
|
||||||
Callers must ensure the text length is NOT divisible by the SP world size;
|
|
||||||
the evenly-divisible case is a plain ``_shard_text_for_sp`` with no mask.
|
|
||||||
|
|
||||||
Returns ``(encoder_hidden_states, freqs_cis, attn_mask, attn_mask_meta)``.
|
|
||||||
"""
|
|
||||||
sp_size = get_sp_world_size()
|
|
||||||
t_real = encoder_hidden_states.shape[1]
|
|
||||||
num_pad = sp_size - t_real % sp_size
|
|
||||||
|
|
||||||
encoder_hidden_states = F.pad(encoder_hidden_states, (0, 0, 0, num_pad))
|
|
||||||
if freqs_cis is not None:
|
|
||||||
img_cache, txt_cache = freqs_cis
|
|
||||||
txt_cache = F.pad(txt_cache, (0, 0, 0, num_pad))
|
|
||||||
freqs_cis = (img_cache, txt_cache)
|
|
||||||
|
|
||||||
local_txt = (t_real + num_pad) // sp_size
|
|
||||||
encoder_hidden_states, freqs_cis = _shard_text_for_sp(
|
|
||||||
encoder_hidden_states, freqs_cis
|
|
||||||
)
|
|
||||||
|
|
||||||
sp_rank = get_sp_parallel_rank()
|
|
||||||
txt_start = sp_rank * local_txt
|
|
||||||
valid_txt = min(local_txt, max(t_real - txt_start, 0))
|
|
||||||
text_mask = torch.zeros(
|
|
||||||
encoder_hidden_states.shape[0],
|
|
||||||
local_txt,
|
|
||||||
dtype=torch.bool,
|
|
||||||
device=encoder_hidden_states.device,
|
|
||||||
)
|
|
||||||
text_mask[:, :valid_txt] = True
|
|
||||||
image_mask = torch.ones(
|
|
||||||
encoder_hidden_states.shape[0],
|
|
||||||
image_seq_len,
|
|
||||||
dtype=torch.bool,
|
|
||||||
device=encoder_hidden_states.device,
|
|
||||||
)
|
|
||||||
joint_mask = torch.cat([text_mask, image_mask], dim=1)
|
|
||||||
# Gathered joint layout is rank-major [txt_0, img_0, ..., txt_{sp-1},
|
|
||||||
# img_{sp-1}]; the pad block is the tail of the last rank's text chunk.
|
|
||||||
gap_meta = {
|
|
||||||
"gap_start": (sp_size - 1) * (local_txt + image_seq_len) + local_txt - num_pad,
|
|
||||||
"gap_end": (sp_size - 1) * (local_txt + image_seq_len) + local_txt,
|
|
||||||
}
|
|
||||||
return encoder_hidden_states, freqs_cis, joint_mask, gap_meta
|
|
||||||
|
|
||||||
|
|
||||||
def _get_qkv_projections(
|
def _get_qkv_projections(
|
||||||
attn: "QwenImageCrossAttention", hidden_states, encoder_hidden_states=None
|
attn: "QwenImageCrossAttention", hidden_states, encoder_hidden_states=None
|
||||||
):
|
):
|
||||||
@@ -776,6 +689,8 @@ class QwenImageCrossAttention(nn.Module):
|
|||||||
# When the text stream is sharded across SP ranks the joint sequence is
|
# When the text stream is sharded across SP ranks the joint sequence is
|
||||||
# fully sequence-parallel, so no leading tokens are replicated.
|
# fully sequence-parallel, so no leading tokens are replicated.
|
||||||
sp_text_sharded = cross_attention_kwargs.get("sp_text_sharded", False)
|
sp_text_sharded = cross_attention_kwargs.get("sp_text_sharded", False)
|
||||||
|
# Rows of tail padding inside THIS rank's text chunk (sp_shard meta).
|
||||||
|
sp_txt_pad = (attn_mask_meta or {}).get("local_pad", 0)
|
||||||
|
|
||||||
(
|
(
|
||||||
img_query,
|
img_query,
|
||||||
@@ -834,11 +749,11 @@ class QwenImageCrossAttention(nn.Module):
|
|||||||
txt_query, txt_key, txt_cache, is_neox=False
|
txt_query, txt_key, txt_cache, is_neox=False
|
||||||
)
|
)
|
||||||
|
|
||||||
# Concatenate for joint attention
|
# Joint order [text, image]; join_seqs relocates any SP text tail-pad
|
||||||
# Order: [text, image]
|
# behind the image (see sp_shard.join_seqs for why).
|
||||||
joint_query = torch.cat([txt_query, img_query], dim=1)
|
joint_query = join_seqs(txt_query, img_query, sp_txt_pad)
|
||||||
joint_key = torch.cat([txt_key, img_key], dim=1)
|
joint_key = join_seqs(txt_key, img_key, sp_txt_pad)
|
||||||
joint_value = torch.cat([txt_value, img_value], dim=1)
|
joint_value = join_seqs(txt_value, img_value, sp_txt_pad)
|
||||||
if attn_mask is None and encoder_hidden_states_mask is not None:
|
if attn_mask is None and encoder_hidden_states_mask is not None:
|
||||||
image_mask = torch.ones(
|
image_mask = torch.ones(
|
||||||
(hidden_states.shape[0], img_query.shape[1]),
|
(hidden_states.shape[0], img_query.shape[1]),
|
||||||
@@ -865,8 +780,9 @@ class QwenImageCrossAttention(nn.Module):
|
|||||||
joint_hidden_states = joint_hidden_states.to(joint_query.dtype)
|
joint_hidden_states = joint_hidden_states.to(joint_query.dtype)
|
||||||
|
|
||||||
# Split attention outputs back
|
# Split attention outputs back
|
||||||
txt_attn_output = joint_hidden_states[:, :seq_len_txt, :] # Text part
|
txt_attn_output, img_attn_output = split_seqs(
|
||||||
img_attn_output = joint_hidden_states[:, seq_len_txt:, :] # Image part
|
joint_hidden_states, seq_len_txt, sp_txt_pad
|
||||||
|
)
|
||||||
|
|
||||||
# Apply output projections
|
# Apply output projections
|
||||||
img_attn_output, _ = self.to_out[0](img_attn_output)
|
img_attn_output, _ = self.to_out[0](img_attn_output)
|
||||||
@@ -1502,7 +1418,6 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
|
|
||||||
block_attention_kwargs = attention_kwargs.copy() if attention_kwargs else {}
|
block_attention_kwargs = attention_kwargs.copy() if attention_kwargs else {}
|
||||||
sp_text_sharded = False
|
sp_text_sharded = False
|
||||||
sp_size = get_sp_world_size()
|
|
||||||
if encoder_hidden_states_mask is not None:
|
if encoder_hidden_states_mask is not None:
|
||||||
encoder_hidden_states_mask = encoder_hidden_states_mask.to(
|
encoder_hidden_states_mask = encoder_hidden_states_mask.to(
|
||||||
device=hidden_states.device, dtype=torch.bool
|
device=hidden_states.device, dtype=torch.bool
|
||||||
@@ -1520,29 +1435,24 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
block_attention_kwargs["attn_mask_meta"] = build_varlen_mask_meta(
|
block_attention_kwargs["attn_mask_meta"] = build_varlen_mask_meta(
|
||||||
joint_mask
|
joint_mask
|
||||||
)
|
)
|
||||||
elif sp_size > 1 and encoder_hidden_states.shape[1] % sp_size == 0:
|
elif should_shard_text(encoder_hidden_states.shape[1]):
|
||||||
# Text divides evenly across SP ranks: plain even shard, no mask.
|
# Shard the replicated text stream across SP ranks; non-divisible
|
||||||
encoder_hidden_states, freqs_cis = _shard_text_for_sp(
|
# lengths tail-pad the last rank and attention skips the pad via the
|
||||||
encoder_hidden_states, freqs_cis
|
# per-request tail meta. Otherwise fall through to replicated text.
|
||||||
|
txt_shard = build_shard_plan(encoder_hidden_states.shape[1])
|
||||||
|
encoder_hidden_states = shard_like(encoder_hidden_states, txt_shard)
|
||||||
|
if freqs_cis is not None:
|
||||||
|
img_cache, txt_cache = freqs_cis
|
||||||
|
freqs_cis = (img_cache, shard_like(txt_cache, txt_shard, dim=0))
|
||||||
|
tail_meta = tail_attn_meta(
|
||||||
|
txt_shard,
|
||||||
|
encoder_hidden_states.shape[0],
|
||||||
|
hidden_states.device,
|
||||||
|
image_seq_len=hidden_states.shape[1],
|
||||||
)
|
)
|
||||||
|
if tail_meta is not None:
|
||||||
|
block_attention_kwargs["attn_mask_meta"] = tail_meta
|
||||||
sp_text_sharded = True
|
sp_text_sharded = True
|
||||||
elif sp_size > 1 and get_ring_parallel_world_size() == 1:
|
|
||||||
# Text does not divide evenly: pad to an SP multiple and shard, with a
|
|
||||||
# pad-gap mask so USPAttention excludes the padding via the varlen
|
|
||||||
# kernel. The varlen masked path does not support ring parallelism, so
|
|
||||||
# uneven text under ring>1 instead falls through to the replicated
|
|
||||||
# path below.
|
|
||||||
(
|
|
||||||
encoder_hidden_states,
|
|
||||||
freqs_cis,
|
|
||||||
pad_mask,
|
|
||||||
pad_meta,
|
|
||||||
) = _pad_shard_text_for_sp_varlen(
|
|
||||||
encoder_hidden_states, freqs_cis, hidden_states.shape[1]
|
|
||||||
)
|
|
||||||
sp_text_sharded = True
|
|
||||||
block_attention_kwargs["attn_mask"] = pad_mask
|
|
||||||
block_attention_kwargs["attn_mask_meta"] = pad_meta
|
|
||||||
block_attention_kwargs["sp_text_sharded"] = sp_text_sharded
|
block_attention_kwargs["sp_text_sharded"] = sp_text_sharded
|
||||||
|
|
||||||
temb = self.time_text_embed(timestep, hidden_states, additional_t_cond)
|
temb = self.time_text_embed(timestep, hidden_states, additional_t_cond)
|
||||||
|
|||||||
+44
-65
@@ -25,13 +25,16 @@ from sglang.multimodal_gen.runtime.distributed import (
|
|||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.distributed.communication_op import (
|
from sglang.multimodal_gen.runtime.distributed.communication_op import (
|
||||||
cfg_model_parallel_all_reduce,
|
cfg_model_parallel_all_reduce,
|
||||||
sequence_model_parallel_all_gather,
|
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
get_cfg_group,
|
get_cfg_group,
|
||||||
get_classifier_free_guidance_rank,
|
get_classifier_free_guidance_rank,
|
||||||
get_sp_parallel_rank,
|
)
|
||||||
get_sp_world_size,
|
from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import (
|
||||||
|
SpShard,
|
||||||
|
gather_seq,
|
||||||
|
shard_seq,
|
||||||
|
tail_attn_meta,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
|
from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context
|
||||||
|
|
||||||
@@ -651,64 +654,15 @@ class MOVADenoisingStage(PipelineStage):
|
|||||||
|
|
||||||
def _shard_sequence_for_sp(
|
def _shard_sequence_for_sp(
|
||||||
self, x: torch.Tensor, dim: int = 1
|
self, x: torch.Tensor, dim: int = 1
|
||||||
) -> tuple[torch.Tensor, int]:
|
) -> tuple[torch.Tensor, SpShard]:
|
||||||
"""
|
"""Tail-padded even shard along the sequence dim (sp_shard.shard_seq)."""
|
||||||
Shard tensor along sequence dimension for Sequence Parallelism.
|
return shard_seq(x, dim=dim)
|
||||||
|
|
||||||
Args:
|
|
||||||
x: Input tensor
|
|
||||||
dim: Dimension to shard along
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
(sharded_tensor, pad_len)
|
|
||||||
"""
|
|
||||||
sp_size = get_sp_world_size()
|
|
||||||
if sp_size <= 1:
|
|
||||||
return x, 0
|
|
||||||
|
|
||||||
sp_rank = get_sp_parallel_rank()
|
|
||||||
seq_len = x.shape[dim]
|
|
||||||
|
|
||||||
# Pad if needed
|
|
||||||
pad_len = (sp_size - (seq_len % sp_size)) % sp_size
|
|
||||||
if pad_len > 0:
|
|
||||||
pad_shape = list(x.shape)
|
|
||||||
pad_shape[dim] = pad_len
|
|
||||||
pad = torch.zeros(pad_shape, dtype=x.dtype, device=x.device)
|
|
||||||
x = torch.cat([x, pad], dim=dim)
|
|
||||||
|
|
||||||
# Shard
|
|
||||||
chunk_size = x.shape[dim] // sp_size
|
|
||||||
start = sp_rank * chunk_size
|
|
||||||
end = start + chunk_size
|
|
||||||
idx = [slice(None)] * x.dim()
|
|
||||||
idx[dim] = slice(start, end)
|
|
||||||
return x[tuple(idx)], pad_len
|
|
||||||
|
|
||||||
def _gather_sequence_from_sp(
|
def _gather_sequence_from_sp(
|
||||||
self, x: torch.Tensor, pad_len: int, dim: int = 1
|
self, x: torch.Tensor, shard: SpShard, dim: int = 1
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""Gather an SP-sharded tensor and trim the tail padding."""
|
||||||
Gather tensor along sequence dimension after Sequence Parallelism.
|
return gather_seq(x, shard.orig_len, dim=dim)
|
||||||
|
|
||||||
Args:
|
|
||||||
x: Sharded tensor
|
|
||||||
pad_len: Padding length that was added during sharding
|
|
||||||
dim: Dimension to gather along
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Gathered tensor with padding removed
|
|
||||||
"""
|
|
||||||
sp_size = get_sp_world_size()
|
|
||||||
if sp_size <= 1:
|
|
||||||
return x
|
|
||||||
|
|
||||||
gathered = sequence_model_parallel_all_gather(x, dim=dim)
|
|
||||||
if pad_len > 0:
|
|
||||||
idx = [slice(None)] * gathered.dim()
|
|
||||||
idx[dim] = slice(0, gathered.shape[dim] - pad_len)
|
|
||||||
gathered = gathered[tuple(idx)]
|
|
||||||
return gathered
|
|
||||||
|
|
||||||
def inference_single_step(
|
def inference_single_step(
|
||||||
self,
|
self,
|
||||||
@@ -817,13 +771,20 @@ class MOVADenoisingStage(PipelineStage):
|
|||||||
).reshape(full_audio_seq_len, 1, -1)
|
).reshape(full_audio_seq_len, 1, -1)
|
||||||
|
|
||||||
# Shard sequences for SP
|
# Shard sequences for SP
|
||||||
visual_x, visual_pad_len = self._shard_sequence_for_sp(visual_x, dim=1)
|
visual_x, visual_shard = self._shard_sequence_for_sp(visual_x, dim=1)
|
||||||
audio_x, audio_pad_len = self._shard_sequence_for_sp(audio_x, dim=1)
|
audio_x, audio_shard = self._shard_sequence_for_sp(audio_x, dim=1)
|
||||||
|
|
||||||
# Shard freqs to match local sequence length
|
# Shard freqs to match local sequence length
|
||||||
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)
|
||||||
|
|
||||||
|
# Tail-pad meta so self-attention excludes SP padding (built once per
|
||||||
|
# step, shared by every block).
|
||||||
|
visual_attn_meta = tail_attn_meta(
|
||||||
|
visual_shard, visual_x.shape[0], visual_x.device
|
||||||
|
)
|
||||||
|
audio_attn_meta = tail_attn_meta(audio_shard, audio_x.shape[0], audio_x.device)
|
||||||
|
|
||||||
# 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,
|
||||||
@@ -839,11 +800,13 @@ class MOVADenoisingStage(PipelineStage):
|
|||||||
video_fps=video_fps,
|
video_fps=video_fps,
|
||||||
full_visual_seq_len=full_visual_seq_len,
|
full_visual_seq_len=full_visual_seq_len,
|
||||||
full_audio_seq_len=full_audio_seq_len,
|
full_audio_seq_len=full_audio_seq_len,
|
||||||
|
visual_attn_meta=visual_attn_meta,
|
||||||
|
audio_attn_meta=audio_attn_meta,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Gather sequences back from SP before head/unpatchify
|
# Gather sequences back from SP before head/unpatchify
|
||||||
visual_x = self._gather_sequence_from_sp(visual_x, visual_pad_len, dim=1)
|
visual_x = self._gather_sequence_from_sp(visual_x, visual_shard, dim=1)
|
||||||
audio_x = self._gather_sequence_from_sp(audio_x, audio_pad_len, dim=1)
|
audio_x = self._gather_sequence_from_sp(audio_x, audio_shard, dim=1)
|
||||||
|
|
||||||
visual_output = visual_dit.head(visual_x, visual_t)
|
visual_output = visual_dit.head(visual_x, visual_t)
|
||||||
visual_output = visual_dit.unpatchify(visual_output, grid_size)
|
visual_output = visual_dit.unpatchify(visual_output, grid_size)
|
||||||
@@ -871,6 +834,8 @@ class MOVADenoisingStage(PipelineStage):
|
|||||||
condition_scale: float | None = 1.0,
|
condition_scale: float | None = 1.0,
|
||||||
a2v_condition_scale: float | None = None,
|
a2v_condition_scale: float | None = None,
|
||||||
v2a_condition_scale: float | None = None,
|
v2a_condition_scale: float | None = None,
|
||||||
|
visual_attn_meta: dict | None = None,
|
||||||
|
audio_attn_meta: dict | None = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Forward pass through dual-tower DiT with cross-modal interaction.
|
Forward pass through dual-tower DiT with cross-modal interaction.
|
||||||
@@ -933,15 +898,29 @@ class MOVADenoisingStage(PipelineStage):
|
|||||||
|
|
||||||
# Self-attention and FFN in DiT blocks
|
# Self-attention and FFN in DiT blocks
|
||||||
visual_x = visual_block(
|
visual_x = visual_block(
|
||||||
visual_x, visual_context, visual_t_mod, visual_freqs
|
visual_x,
|
||||||
|
visual_context,
|
||||||
|
visual_t_mod,
|
||||||
|
visual_freqs,
|
||||||
|
attn_mask_meta=visual_attn_meta,
|
||||||
|
)
|
||||||
|
audio_x = audio_block(
|
||||||
|
audio_x,
|
||||||
|
audio_context,
|
||||||
|
audio_t_mod,
|
||||||
|
audio_freqs,
|
||||||
|
attn_mask_meta=audio_attn_meta,
|
||||||
)
|
)
|
||||||
audio_x = audio_block(audio_x, audio_context, audio_t_mod, audio_freqs)
|
|
||||||
|
|
||||||
# Process remaining visual layers (if visual has more layers than audio)
|
# Process remaining visual layers (if visual has more layers than audio)
|
||||||
for layer_idx in range(min_layers, visual_layers):
|
for layer_idx in range(min_layers, visual_layers):
|
||||||
visual_block = visual_dit.blocks[layer_idx]
|
visual_block = visual_dit.blocks[layer_idx]
|
||||||
visual_x = visual_block(
|
visual_x = visual_block(
|
||||||
visual_x, visual_context, visual_t_mod, visual_freqs
|
visual_x,
|
||||||
|
visual_context,
|
||||||
|
visual_t_mod,
|
||||||
|
visual_freqs,
|
||||||
|
attn_mask_meta=visual_attn_meta,
|
||||||
)
|
)
|
||||||
|
|
||||||
return visual_x, audio_x
|
return visual_x, audio_x
|
||||||
|
|||||||
@@ -229,12 +229,6 @@
|
|||||||
"psnr_threshold": 30.0,
|
"psnr_threshold": 30.0,
|
||||||
"mean_abs_diff_threshold": 5.6
|
"mean_abs_diff_threshold": 5.6
|
||||||
},
|
},
|
||||||
"mova_360p_ring1_uly2": {
|
|
||||||
"clip_threshold": 0.96,
|
|
||||||
"ssim_threshold": 0.91,
|
|
||||||
"psnr_threshold": 30.0,
|
|
||||||
"mean_abs_diff_threshold": 6.2
|
|
||||||
},
|
|
||||||
"wan2_1_i2v_14b_lora_2gpu": {
|
"wan2_1_i2v_14b_lora_2gpu": {
|
||||||
"clip_threshold": 0.97,
|
"clip_threshold": 0.97,
|
||||||
"ssim_threshold": 0.90,
|
"ssim_threshold": 0.90,
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
SGL_TEST_FILES_CI_DATA_REVISION = "c28c11c24125b009f9e602fd5ae0d0ddcdc11d36"
|
SGL_TEST_FILES_CI_DATA_REVISION = "46b9b53a429606cb6739c861f275c1277c314a10"
|
||||||
|
|
||||||
if current_platform.is_npu():
|
if current_platform.is_npu():
|
||||||
SGL_TEST_FILES_CI_DATA_REVISION = "670d66a8a290b62c0c3c077b3e9b0f4a4d9a44e7"
|
SGL_TEST_FILES_CI_DATA_REVISION = "670d66a8a290b62c0c3c077b3e9b0f4a4d9a44e7"
|
||||||
|
|||||||
@@ -0,0 +1,213 @@
|
|||||||
|
"""Unit tests for the unified SP shard helpers (pure logic, no distributed)."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.distributed import sp_shard_utils as sps
|
||||||
|
from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import (
|
||||||
|
SpShard,
|
||||||
|
shard_like,
|
||||||
|
tail_attn_meta,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _fake_sp(monkeypatch, sp_size, sp_rank=0, ring=1):
|
||||||
|
monkeypatch.setattr(sps, "get_sp_world_size", lambda: sp_size)
|
||||||
|
monkeypatch.setattr(sps, "get_sp_parallel_rank", lambda: sp_rank)
|
||||||
|
monkeypatch.setattr(sps, "get_ring_parallel_world_size", lambda: ring)
|
||||||
|
|
||||||
|
|
||||||
|
# --- build_shard_plan math --------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_plan_shard_divisible(monkeypatch):
|
||||||
|
_fake_sp(monkeypatch, 2, 1)
|
||||||
|
s = sps.build_shard_plan(16)
|
||||||
|
assert (s.local_len, s.num_pad, s.local_pad) == (8, 0, 0)
|
||||||
|
|
||||||
|
|
||||||
|
def test_plan_shard_padded_last_rank(monkeypatch):
|
||||||
|
_fake_sp(monkeypatch, 4, 3)
|
||||||
|
s = sps.build_shard_plan(14)
|
||||||
|
assert (s.local_len, s.num_pad) == (4, 2)
|
||||||
|
assert s.local_pad == 2 and s.local_real_len == 2
|
||||||
|
|
||||||
|
|
||||||
|
def test_plan_shard_pad_only_on_last_rank(monkeypatch):
|
||||||
|
_fake_sp(monkeypatch, 4, 0)
|
||||||
|
s = sps.build_shard_plan(14)
|
||||||
|
assert s.local_pad == 0 and s.local_real_len == 4
|
||||||
|
|
||||||
|
|
||||||
|
def test_plan_shard_sp1_noop(monkeypatch):
|
||||||
|
_fake_sp(monkeypatch, 1)
|
||||||
|
s = sps.build_shard_plan(15)
|
||||||
|
assert (s.local_len, s.num_pad, s.sp_size) == (15, 0, 1)
|
||||||
|
|
||||||
|
|
||||||
|
# --- shard_like -------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_shard_like_zero_pads_tail():
|
||||||
|
shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=1)
|
||||||
|
x = torch.arange(15, dtype=torch.float32).unsqueeze(0).unsqueeze(-1)
|
||||||
|
local = shard_like(x, shard, dim=1)
|
||||||
|
assert local.shape[1] == 8
|
||||||
|
assert local[0, -1, 0].item() == 0.0 # tail pad
|
||||||
|
assert local[0, 0, 0].item() == 8.0 # rank1 starts at token 8
|
||||||
|
|
||||||
|
|
||||||
|
def test_shard_like_repeat_last():
|
||||||
|
shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=1)
|
||||||
|
x = torch.arange(15, dtype=torch.float32).unsqueeze(-1)
|
||||||
|
local = shard_like(x, shard, dim=0, pad_mode="repeat_last")
|
||||||
|
assert local[-1, 0].item() == 14.0 # repeated last row, not zero
|
||||||
|
|
||||||
|
|
||||||
|
def test_shard_like_chunks_align_across_tensors():
|
||||||
|
# RoPE cache sharded with the same plan stays aligned with hidden states.
|
||||||
|
shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=0)
|
||||||
|
x = torch.arange(15).unsqueeze(0).unsqueeze(-1).float()
|
||||||
|
rope = torch.arange(15).unsqueeze(-1).float()
|
||||||
|
assert torch.equal(
|
||||||
|
shard_like(x, shard, dim=1)[0, :, 0], shard_like(rope, shard, dim=0)[:, 0]
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# --- tail_attn_meta ---------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_tail_meta_none_when_divisible():
|
||||||
|
shard = SpShard(orig_len=16, local_len=8, num_pad=0, sp_size=2, sp_rank=0)
|
||||||
|
assert tail_attn_meta(shard, 1, torch.device("cpu")) is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_tail_meta_single_stream():
|
||||||
|
shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=1)
|
||||||
|
meta = tail_attn_meta(shard, 1, torch.device("cpu"))
|
||||||
|
assert meta["pad_start"] == 15 and meta["pad_end"] == 16
|
||||||
|
assert meta["local_pad"] == 1
|
||||||
|
assert meta["cu_seqlens_tail"].tolist() == [0, 15, 16]
|
||||||
|
assert meta["max_seqlen_tail"] == 15
|
||||||
|
|
||||||
|
|
||||||
|
def test_tail_meta_joint_layout_and_batch():
|
||||||
|
# sp=2, local_txt=8 (1 pad), img=100 per rank -> S = 2*(8+100) = 216.
|
||||||
|
shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=1)
|
||||||
|
meta = tail_attn_meta(shard, 2, torch.device("cpu"), image_seq_len=100)
|
||||||
|
assert meta["pad_start"] == 215 and meta["pad_end"] == 216
|
||||||
|
assert meta["cu_seqlens_tail"].tolist() == [0, 215, 216, 431, 432]
|
||||||
|
|
||||||
|
|
||||||
|
def test_tail_meta_max_seqlen_covers_pad_segment():
|
||||||
|
# Degenerate short sequence: num_pad (3) > valid (1). FA requires
|
||||||
|
# max_seqlen >= the longest segment, i.e. the pad block here.
|
||||||
|
shard = SpShard(orig_len=1, local_len=1, num_pad=3, sp_size=4, sp_rank=3)
|
||||||
|
meta = tail_attn_meta(shard, 1, torch.device("cpu"))
|
||||||
|
assert meta["max_seqlen_tail"] == 3
|
||||||
|
|
||||||
|
|
||||||
|
def test_tail_meta_matches_legacy_gap_formula():
|
||||||
|
# The tail layout puts the pad exactly where the legacy per-model gap
|
||||||
|
# formula pointed, minus the relocation: end == S (global tail).
|
||||||
|
sp, local_txt, img, num_pad = 3, 5, 40, 2
|
||||||
|
shard = SpShard(
|
||||||
|
orig_len=sp * local_txt - num_pad,
|
||||||
|
local_len=local_txt,
|
||||||
|
num_pad=num_pad,
|
||||||
|
sp_size=sp,
|
||||||
|
sp_rank=sp - 1,
|
||||||
|
)
|
||||||
|
meta = tail_attn_meta(shard, 1, torch.device("cpu"), image_seq_len=img)
|
||||||
|
seq = sp * (local_txt + img)
|
||||||
|
assert meta["pad_end"] == seq
|
||||||
|
assert meta["pad_start"] == seq - num_pad
|
||||||
|
|
||||||
|
|
||||||
|
# --- plan_text_strategy -----------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_strategy_sp1_replicates(monkeypatch):
|
||||||
|
_fake_sp(monkeypatch, 1)
|
||||||
|
assert sps.plan_text_strategy(100) == "replicate"
|
||||||
|
|
||||||
|
|
||||||
|
def test_strategy_shard_when_legal(monkeypatch):
|
||||||
|
_fake_sp(monkeypatch, 2)
|
||||||
|
assert sps.plan_text_strategy(15) == "shard"
|
||||||
|
assert sps.plan_text_strategy(16) == "shard"
|
||||||
|
|
||||||
|
|
||||||
|
def test_strategy_ring_blocks_padded_shard(monkeypatch):
|
||||||
|
_fake_sp(monkeypatch, 2, ring=2)
|
||||||
|
assert sps.plan_text_strategy(15) == "replicate" # padded shard needs mask
|
||||||
|
assert sps.plan_text_strategy(16) == "shard" # divisible: no mask needed
|
||||||
|
|
||||||
|
|
||||||
|
def test_strategy_min_len_threshold(monkeypatch):
|
||||||
|
_fake_sp(monkeypatch, 2)
|
||||||
|
monkeypatch.setattr(sps, "_TEXT_SHARD_MIN", 64)
|
||||||
|
assert sps.plan_text_strategy(32) == "replicate"
|
||||||
|
assert sps.plan_text_strategy(64) == "shard"
|
||||||
|
|
||||||
|
|
||||||
|
# --- join_seqs / split_seqs / shard_seq_prefix ------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_join_split_roundtrip_with_pad():
|
||||||
|
# Joint [text, image] with 2 tail-pad rows relocated behind the image.
|
||||||
|
txt = torch.arange(6, dtype=torch.float32).view(1, 6, 1) # rows 4,5 are pad
|
||||||
|
img = (torch.arange(3, dtype=torch.float32) + 100).view(1, 3, 1)
|
||||||
|
joint = sps.join_seqs(txt, img, local_pad=2)
|
||||||
|
assert joint[0, :, 0].tolist() == [0, 1, 2, 3, 100, 101, 102, 4, 5]
|
||||||
|
txt_back, img_back = sps.split_seqs(joint, prefix_len=6, local_pad=2)
|
||||||
|
assert torch.equal(txt_back, txt) and torch.equal(img_back, img)
|
||||||
|
|
||||||
|
|
||||||
|
def test_join_split_roundtrip_no_pad():
|
||||||
|
txt = torch.randn(1, 4, 2)
|
||||||
|
img = torch.randn(1, 3, 2)
|
||||||
|
joint = sps.join_seqs(txt, img, local_pad=0)
|
||||||
|
assert torch.equal(joint, torch.cat([txt, img], dim=1))
|
||||||
|
txt_back, img_back = sps.split_seqs(joint, prefix_len=4, local_pad=0)
|
||||||
|
assert torch.equal(txt_back, txt) and torch.equal(img_back, img)
|
||||||
|
|
||||||
|
|
||||||
|
def test_shard_seq_prefix_only_touches_prefix():
|
||||||
|
# Joint RoPE cache [txt(15); img(4)]: text segment shards, image stays.
|
||||||
|
shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=1)
|
||||||
|
cache = torch.arange(19, dtype=torch.float32).unsqueeze(-1)
|
||||||
|
out = sps.shard_seq_prefix(cache, 15, shard, dim=0)
|
||||||
|
assert out.shape[0] == 8 + 4
|
||||||
|
assert out[0, 0].item() == 8.0 # rank1 text chunk starts at token 8
|
||||||
|
assert out[-4:, 0].flatten().tolist() == [15, 16, 17, 18] # image untouched
|
||||||
|
|
||||||
|
|
||||||
|
def test_should_shard_text_gate(monkeypatch):
|
||||||
|
_fake_sp(monkeypatch, 2)
|
||||||
|
assert sps.should_shard_text(15) is True
|
||||||
|
_fake_sp(monkeypatch, 1)
|
||||||
|
assert sps.should_shard_text(15) is False
|
||||||
|
|
||||||
|
|
||||||
|
# --- gather_seq -------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_gather_seq_sp1_noop(monkeypatch):
|
||||||
|
_fake_sp(monkeypatch, 1)
|
||||||
|
x = torch.randn(1, 5, 2)
|
||||||
|
assert sps.gather_seq(x, 5, dim=1) is x
|
||||||
|
|
||||||
|
|
||||||
|
def test_gather_seq_trims(monkeypatch):
|
||||||
|
_fake_sp(monkeypatch, 2)
|
||||||
|
monkeypatch.setattr(
|
||||||
|
sps, "sequence_model_parallel_all_gather", lambda t, dim: torch.cat([t, t], dim)
|
||||||
|
)
|
||||||
|
local = torch.randn(1, 8, 2)
|
||||||
|
out = sps.gather_seq(local, 15, dim=1)
|
||||||
|
assert out.shape[1] == 15
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
pytest.main([__file__, "-q"])
|
||||||
Reference in New Issue
Block a user