[diffusion] perf: shard text when using sp in flux.1/2 (#27066)

Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
Aleksi Vesanto
2026-06-17 21:17:39 +08:00
committed by GitHub
co-authored by Mick
parent 3fb65ebabd
commit dad890fff1
3 changed files with 395 additions and 23 deletions
@@ -70,7 +70,7 @@ def build_varlen_mask_meta(
Returns ``cu_seqlens``, ``indices``, ``inv_indices``, ``max_seqlen``.
Passing the result via ``joint_attention_kwargs`` opts the caller into
``USPAttention``'s varlen FA fast path, which zero-fills masked query
rows on output only use when those rows are dropped or ignored
rows on output; only use when those rows are dropped or ignored
downstream.
"""
assert key_mask.dim() == 2, "key_mask must be [B, S]"
@@ -500,10 +500,10 @@ class USPAttention(nn.Module):
conditioning prefix (e.g. cached text K/V) followed by a
sequence-sharded suffix (image tokens). Q has no replicated
portion and is fully sequence-sharded.
attn_mask_meta: optional varlen metadata from
``build_varlen_mask_meta(attn_mask)``. Supplying this opts
into the varlen FA fast path, in which masked query rows
are zero-filled on output (differs from SDPA semantics).
attn_mask_meta: optional metadata for the varlen FA fast path.
Callers may pass ``build_varlen_mask_meta(attn_mask)`` or a
known contiguous padding gap. Masked query rows are zero-filled
on output (differs from SDPA semantics).
Note: Replicated tensors are not supported in this implementation.
When skip_sequence_parallel=True (set at construction time), all SP
@@ -619,6 +619,56 @@ class USPAttention(nn.Module):
k = _usp_input_all_to_all(k, 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 (
_VARLEN_FA_ENABLED
and self.backend == AttentionBackendEnum.FA
and gap_start is not None
and gap_end is not None
and gap_end > gap_start
and q.device.type == "cuda"
and q.dtype in (torch.float16, torch.bfloat16)
):
bs, seq = q.shape[0], q.shape[1]
assert 0 <= gap_start < gap_end <= seq
valid_seq = seq - (gap_end - gap_start)
q_dense = torch.cat([q[:, :gap_start], q[:, gap_end:]], dim=1)
k_dense = torch.cat([k[:, :gap_start], k[:, gap_end:]], dim=1)
v_dense = torch.cat([v[:, :gap_start], v[:, gap_end:]], dim=1)
cu_seqlens = torch.arange(
0,
(bs + 1) * valid_seq,
valid_seq,
dtype=torch.int32,
device=q.device,
)
out_dense = flash_attn_varlen_func(
q=q_dense.reshape(bs * valid_seq, *q.shape[2:]),
k=k_dense.reshape(bs * valid_seq, *k.shape[2:]),
v=v_dense.reshape(bs * valid_seq, *v.shape[2:]),
cu_seqlens_q=cu_seqlens,
cu_seqlens_k=cu_seqlens,
max_seqlen_q=valid_seq,
max_seqlen_k=valid_seq,
softmax_scale=self.softmax_scale,
causal=False,
ver=_fa_backend.fa_ver,
).reshape(bs, valid_seq, *q.shape[2:])
gap_out = out_dense.new_zeros(
bs, gap_end - gap_start, out_dense.shape[2], out_dense.shape[3]
)
out = torch.cat(
[out_dense[:, :gap_start], gap_out, out_dense[:, gap_start:]],
dim=1,
)
if sp_size > 1:
out = _usp_output_all_to_all(out, head_dim=2)
return out
# If NCCL timeout/deadlock occurs here, check whether
# attn_mask is inconsistent across SP ranks (None on some, Tensor on
# others), which causes all_gather participant mismatch. Upstream
@@ -626,6 +676,42 @@ class USPAttention(nn.Module):
gathered_mask = sequence_model_parallel_all_gather(
attn_mask.contiguous(), dim=1
)
if (
_VARLEN_FA_ENABLED
and self.backend == AttentionBackendEnum.FA
and gathered_mask.dtype
in (torch.bool, torch.uint8, torch.int32, torch.int64)
and q.device.type == "cuda"
and gathered_mask.device == q.device
and q.dtype in (torch.float16, torch.bfloat16)
and q.shape[:2] == gathered_mask.shape == k.shape[:2] == v.shape[:2]
):
bs, seq = q.shape[0], q.shape[1]
gathered_mask_meta = build_varlen_mask_meta(gathered_mask)
indices = gathered_mask_meta["indices"]
inv_indices = gathered_mask_meta["inv_indices"]
assert (
inv_indices.shape[0] == bs * seq
), "gathered attn_mask shape does not match q/k/v"
if indices.shape[0] > 0:
q_unpad, k_unpad, v_unpad = fused_pack_qkv(q, k, v, indices)
out_unpad = flash_attn_varlen_func(
q=q_unpad,
k=k_unpad,
v=v_unpad,
cu_seqlens_q=gathered_mask_meta["cu_seqlens"],
cu_seqlens_k=gathered_mask_meta["cu_seqlens"],
max_seqlen_q=gathered_mask_meta["max_seqlen"],
max_seqlen_k=gathered_mask_meta["max_seqlen"],
softmax_scale=self.softmax_scale,
causal=False,
ver=_fa_backend.fa_ver,
)
out = fused_scatter_to_padded(out_unpad, inv_indices, bs, seq)
if sp_size > 1:
out = _usp_output_all_to_all(out, head_dim=2)
return out
q_ = q.transpose(1, 2)
k_ = k.transpose(1, 2)
v_ = v.transpose(1, 2)
@@ -28,7 +28,12 @@ from diffusers.models.normalization import (
from torch.nn import LayerNorm as LayerNorm
from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
from sglang.multimodal_gen.runtime.distributed import divide, get_tp_world_size
from sglang.multimodal_gen.runtime.distributed import (
divide,
get_sp_parallel_rank,
get_sp_world_size,
get_tp_world_size,
)
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
from sglang.multimodal_gen.runtime.layers.fused_linear_act import linear_gelu_tanh
from sglang.multimodal_gen.runtime.layers.layernorm import (
@@ -64,6 +69,100 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
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:
from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import]
from nunchaku.models.normalization import ( # type: ignore[import]
@@ -448,6 +547,8 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
encoder_hidden_states: Optional[torch.Tensor] = None,
freqs_cis=None,
num_replicated_prefix: int = 0,
attn_mask: Optional[torch.Tensor] = None,
attn_mask_meta: Optional[Dict[str, int]] = None,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
(
query,
@@ -504,9 +605,6 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
query = torch.cat([encoder_query, query], dim=1)
key = torch.cat([encoder_key, key], dim=1)
value = torch.cat([encoder_value, value], dim=1)
num_replicated_prefix = (
num_replicated_prefix or encoder_hidden_states.shape[1]
)
else:
query, key = apply_qk_norm_with_optional_rope(
q=query,
@@ -519,7 +617,14 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin):
allow_inplace=True,
)
x = self.attn(query, key, value, num_replicated_prefix=num_replicated_prefix)
x = self.attn(
query,
key,
value,
attn_mask=attn_mask,
attn_mask_meta=attn_mask_meta,
num_replicated_prefix=num_replicated_prefix,
)
x = x.flatten(2, 3)
x = x.to(query.dtype)
@@ -669,6 +774,7 @@ class FluxSingleTransformerBlock(nn.Module):
temb: torch.Tensor,
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
num_replicated_prefix: int = 0,
) -> Tuple[torch.Tensor, torch.Tensor]:
text_seq_len = encoder_hidden_states.shape[1]
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
@@ -676,7 +782,6 @@ class FluxSingleTransformerBlock(nn.Module):
residual = hidden_states
norm_hidden_states, gate = self.norm(hidden_states, emb=temb)
joint_attention_kwargs = joint_attention_kwargs or {}
joint_attention_kwargs.setdefault("num_replicated_prefix", text_seq_len or 0)
if self.use_nunchaku_structure:
if _nunchaku_fused_ops_available:
@@ -691,6 +796,7 @@ class FluxSingleTransformerBlock(nn.Module):
attn_output = self.attn(
x=norm_hidden_states,
freqs_cis=freqs_cis,
num_replicated_prefix=num_replicated_prefix,
**joint_attention_kwargs,
)
if isinstance(attn_output, tuple):
@@ -706,6 +812,7 @@ class FluxSingleTransformerBlock(nn.Module):
attn_output = self.attn(
x=norm_hidden_states,
freqs_cis=freqs_cis,
num_replicated_prefix=num_replicated_prefix,
**joint_attention_kwargs,
)
@@ -815,6 +922,7 @@ class FluxTransformerBlock(nn.Module):
temb: torch.Tensor,
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
num_replicated_prefix: int = 0,
) -> Tuple[torch.Tensor, torch.Tensor]:
norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(
hidden_states, emb=temb
@@ -834,6 +942,7 @@ class FluxTransformerBlock(nn.Module):
x=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
freqs_cis=freqs_cis,
num_replicated_prefix=num_replicated_prefix,
**joint_attention_kwargs,
)
@@ -1075,8 +1184,28 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
else:
temb = self.time_text_embed(timestep, pooled_projections)
num_txt_tokens = encoder_hidden_states.shape[1]
encoder_hidden_states, _ = self.context_embedder(encoder_hidden_states)
(
encoder_hidden_states,
freqs_cis,
num_replicated_prefix,
attn_mask,
attn_mask_meta,
) = _shard_text_for_sp(
encoder_hidden_states,
freqs_cis,
hidden_states.shape[1],
num_txt_tokens,
)
if attn_mask is not None:
joint_attention_kwargs = (
joint_attention_kwargs.copy() if joint_attention_kwargs else {}
)
joint_attention_kwargs["attn_mask"] = attn_mask
joint_attention_kwargs["attn_mask_meta"] = attn_mask_meta
if (
joint_attention_kwargs is not None
and "ip_adapter_image_embeds" in joint_attention_kwargs
@@ -1094,6 +1223,7 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
temb=temb,
freqs_cis=freqs_cis,
joint_attention_kwargs=joint_attention_kwargs,
num_replicated_prefix=num_replicated_prefix,
)
for block in self.single_transformer_blocks:
encoder_hidden_states, hidden_states = block(
@@ -1102,6 +1232,7 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
temb=temb,
freqs_cis=freqs_cis,
joint_attention_kwargs=joint_attention_kwargs,
num_replicated_prefix=num_replicated_prefix,
)
hidden_states = self.norm_out(hidden_states, temb)
@@ -21,7 +21,12 @@ from diffusers.models.embeddings import TimestepEmbedding, Timesteps
from diffusers.models.normalization import AdaLayerNormContinuous
from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig
from sglang.multimodal_gen.runtime.distributed import divide, get_tp_world_size
from sglang.multimodal_gen.runtime.distributed import (
divide,
get_sp_parallel_rank,
get_sp_world_size,
get_tp_world_size,
)
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
from sglang.multimodal_gen.runtime.layers.layernorm import (
RMSNorm,
@@ -55,6 +60,115 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
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(
attn: "Flux2Attention", hidden_states, encoder_hidden_states=None
):
@@ -295,6 +409,9 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
hidden_states: torch.Tensor,
encoder_hidden_states: Optional[torch.Tensor] = None,
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
num_replicated_prefix: int = 0,
attn_mask: Optional[torch.Tensor] = None,
attn_mask_meta: Optional[Dict[str, int]] = None,
) -> torch.Tensor:
(
query,
@@ -363,10 +480,14 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
allow_inplace=True,
)
num_rep = (
encoder_hidden_states.shape[1] if encoder_hidden_states is not None else 0
hidden_states = self.attn(
query,
key,
value,
attn_mask=attn_mask,
attn_mask_meta=attn_mask_meta,
num_replicated_prefix=num_replicated_prefix,
)
hidden_states = self.attn(query, key, value, num_replicated_prefix=num_rep)
hidden_states = hidden_states.flatten(2, 3)
hidden_states = hidden_states.to(query.dtype)
@@ -507,6 +628,11 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
num_replicated_prefix: int = 0,
**kwargs,
) -> torch.Tensor:
attn_mask = kwargs.get("attn_mask")
attn_mask_meta = kwargs.get("attn_mask_meta")
if attn_mask is None:
attn_mask = attention_mask
# Parallel in (QKV + MLP in) projection
hidden_states, _ = self.to_qkv_mlp_proj(hidden_states)
qkv, mlp_hidden_states = torch.split(
@@ -541,7 +667,12 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
query, key, cos_sin_cache, is_neox=False
)
hidden_states = self.attn(
query, key, value, num_replicated_prefix=num_replicated_prefix
query,
key,
value,
attn_mask=attn_mask,
attn_mask_meta=attn_mask_meta,
num_replicated_prefix=num_replicated_prefix,
)
hidden_states = hidden_states.flatten(2, 3)
hidden_states = hidden_states.to(query.dtype)
@@ -600,6 +731,7 @@ class Flux2SingleTransformerBlock(nn.Module):
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
split_hidden_states: bool = False,
text_seq_len: Optional[int] = None,
num_replicated_prefix: int = 0,
) -> Tuple[torch.Tensor, torch.Tensor]:
# If encoder_hidden_states is None, hidden_states is assumed to have encoder_hidden_states already
# concatenated
@@ -616,7 +748,7 @@ class Flux2SingleTransformerBlock(nn.Module):
attn_output = self.attn(
hidden_states=norm_hidden_states,
freqs_cis=freqs_cis,
num_replicated_prefix=text_seq_len or 0,
num_replicated_prefix=num_replicated_prefix,
**joint_attention_kwargs,
)
@@ -700,13 +832,16 @@ class Flux2TransformerBlock(nn.Module):
],
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
joint_attention_kwargs: Optional[Dict[str, Any]] = None,
num_replicated_prefix: int = 0,
) -> Tuple[torch.Tensor, torch.Tensor]:
joint_attention_kwargs = joint_attention_kwargs or {}
# Modulation parameters shape: [1, 1, self.dim]
(shift_msa, scale_msa, gate_msa), (shift_mlp, scale_mlp, gate_mlp) = (
temb_mod_params_img
)
(shift_msa, scale_msa, gate_msa), (
shift_mlp,
scale_mlp,
gate_mlp,
) = temb_mod_params_img
(c_shift_msa, c_scale_msa, c_gate_msa), (
c_shift_mlp,
c_scale_mlp,
@@ -728,6 +863,7 @@ class Flux2TransformerBlock(nn.Module):
hidden_states=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
freqs_cis=freqs_cis,
num_replicated_prefix=num_replicated_prefix,
**joint_attention_kwargs,
)
@@ -1060,9 +1196,26 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
hidden_states, _ = self.x_embedder(hidden_states)
encoder_hidden_states, _ = self.context_embedder(encoder_hidden_states)
# 3. Calculate RoPE embeddings from image and text tokens
# NOTE: the below logic means that we can't support batched inference with images of different resolutions or
# text prompts of different lengths. Is this a use case we want to support?
(
encoder_hidden_states,
freqs_cis,
num_replicated_prefix,
num_txt_tokens,
attn_mask,
attn_mask_meta,
) = _shard_text_for_sp(
encoder_hidden_states,
freqs_cis,
hidden_states.shape[1],
num_txt_tokens,
)
if attn_mask is not None:
joint_attention_kwargs = (
joint_attention_kwargs.copy() if joint_attention_kwargs else {}
)
joint_attention_kwargs["attn_mask"] = attn_mask
joint_attention_kwargs["attn_mask_meta"] = attn_mask_meta
# 4. Double Stream Transformer Blocks
for index_block, block in enumerate(self.transformer_blocks):
encoder_hidden_states, hidden_states = block(
@@ -1072,6 +1225,7 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
temb_mod_params_txt=double_stream_mod_txt,
freqs_cis=freqs_cis,
joint_attention_kwargs=joint_attention_kwargs,
num_replicated_prefix=num_replicated_prefix,
)
# Concatenate text and image streams for single-block inference
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
@@ -1085,6 +1239,7 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
freqs_cis=freqs_cis,
joint_attention_kwargs=joint_attention_kwargs,
text_seq_len=num_txt_tokens,
num_replicated_prefix=num_replicated_prefix,
)
# Remove text tokens from concatenated stream
hidden_states = hidden_states[:, num_txt_tokens:, ...]