[diffusion] fix: fix corrupted image editing outputs in Multi-GPU SP mode for FLUX.2-klein models (#19454)
This commit is contained in:
@@ -530,10 +530,6 @@ class Flux2PipelineConfig(FluxPipelineConfig):
|
|||||||
txt_ids = _prepare_text_ids(prompt_embeds).to(device=device)
|
txt_ids = _prepare_text_ids(prompt_embeds).to(device=device)
|
||||||
|
|
||||||
img_ids = batch.latent_ids
|
img_ids = batch.latent_ids
|
||||||
if batch.image_latent is not None:
|
|
||||||
image_latent_ids = batch.condition_image_latent_ids
|
|
||||||
img_ids = torch.cat([img_ids, image_latent_ids], dim=1).to(device=device)
|
|
||||||
|
|
||||||
if img_ids.ndim == 3:
|
if img_ids.ndim == 3:
|
||||||
img_ids = img_ids[0]
|
img_ids = img_ids[0]
|
||||||
if txt_ids.ndim == 3:
|
if txt_ids.ndim == 3:
|
||||||
@@ -544,6 +540,16 @@ class Flux2PipelineConfig(FluxPipelineConfig):
|
|||||||
img_cos = shard_rotary_emb_for_sp(img_cos)
|
img_cos = shard_rotary_emb_for_sp(img_cos)
|
||||||
img_sin = shard_rotary_emb_for_sp(img_sin)
|
img_sin = shard_rotary_emb_for_sp(img_sin)
|
||||||
|
|
||||||
|
if batch.image_latent is not None:
|
||||||
|
cond_ids = batch.condition_image_latent_ids
|
||||||
|
if cond_ids.ndim == 3:
|
||||||
|
cond_ids = cond_ids[0]
|
||||||
|
cond_cos, cond_sin = rotary_emb.forward(cond_ids)
|
||||||
|
cond_cos = shard_rotary_emb_for_sp(cond_cos)
|
||||||
|
cond_sin = shard_rotary_emb_for_sp(cond_sin)
|
||||||
|
img_cos = torch.cat([img_cos, cond_cos], dim=0)
|
||||||
|
img_sin = torch.cat([img_sin, cond_sin], dim=0)
|
||||||
|
|
||||||
txt_cos, txt_sin = rotary_emb.forward(txt_ids)
|
txt_cos, txt_sin = rotary_emb.forward(txt_ids)
|
||||||
|
|
||||||
cos = torch.cat([txt_cos, img_cos], dim=0).to(device=device)
|
cos = torch.cat([txt_cos, img_cos], dim=0).to(device=device)
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ from sglang.multimodal_gen.runtime.distributed.communication_op import (
|
|||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
get_ring_parallel_world_size,
|
get_ring_parallel_world_size,
|
||||||
get_sequence_parallel_world_size,
|
get_sequence_parallel_world_size,
|
||||||
|
get_sp_group,
|
||||||
get_sp_parallel_rank,
|
get_sp_parallel_rank,
|
||||||
get_sp_world_size,
|
get_sp_world_size,
|
||||||
get_ulysses_parallel_world_size,
|
get_ulysses_parallel_world_size,
|
||||||
@@ -352,23 +353,23 @@ class USPAttention(nn.Module):
|
|||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
k: torch.Tensor,
|
k: torch.Tensor,
|
||||||
v: torch.Tensor,
|
v: torch.Tensor,
|
||||||
replicated_q: torch.Tensor | None = None,
|
num_replicated_prefix: int = 0,
|
||||||
replicated_k: torch.Tensor | None = None,
|
|
||||||
replicated_v: torch.Tensor | None = None,
|
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
Forward pass for USPAttention.
|
Forward pass for USPAttention.
|
||||||
|
|
||||||
q, k, v: [B, S_local, H, D]
|
q, k, v: [B, S_local, H, D]
|
||||||
|
num_replicated_prefix: number of leading tokens in q/k/v that are
|
||||||
|
replicated (identical) across all SP ranks, e.g. text tokens
|
||||||
|
in FLUX joint attention. These tokens are excluded from the
|
||||||
|
Ulysses all-to-all so they appear exactly once in the gathered
|
||||||
|
sequence, preserving correct attention weights.
|
||||||
|
|
||||||
Note: Replicated tensors are not supported in this implementation.
|
Note: Replicated tensors are not supported in this implementation.
|
||||||
When skip_sequence_parallel=True (set at construction time), all SP
|
When skip_sequence_parallel=True (set at construction time), all SP
|
||||||
communication is bypassed — use this for cross-attention where KV
|
communication is bypassed — use this for cross-attention where KV
|
||||||
content is replicated across ranks (distinct from replicated_k/v args).
|
content is replicated across ranks (distinct from replicated_k/v args).
|
||||||
"""
|
"""
|
||||||
assert (
|
|
||||||
replicated_q is None and replicated_k is None and replicated_v is None
|
|
||||||
), "USPAttention does not support replicated_qkv."
|
|
||||||
forward_context: ForwardContext = get_forward_context()
|
forward_context: ForwardContext = get_forward_context()
|
||||||
ctx_attn_metadata = forward_context.attn_metadata
|
ctx_attn_metadata = forward_context.attn_metadata
|
||||||
if self.skip_sequence_parallel or get_sequence_parallel_world_size() == 1:
|
if self.skip_sequence_parallel or get_sequence_parallel_world_size() == 1:
|
||||||
@@ -376,8 +377,14 @@ class USPAttention(nn.Module):
|
|||||||
out = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
out = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
sp_size = get_ulysses_parallel_world_size()
|
||||||
|
if sp_size > 1 and num_replicated_prefix > 0:
|
||||||
|
return self._forward_with_replicated_prefix(
|
||||||
|
q, k, v, ctx_attn_metadata, num_replicated_prefix
|
||||||
|
)
|
||||||
|
|
||||||
# Ulysses-style All-to-All for sequence/head sharding
|
# Ulysses-style All-to-All for sequence/head sharding
|
||||||
if get_ulysses_parallel_world_size() > 1:
|
if sp_size > 1:
|
||||||
# -> [B, S, H_local, D]
|
# -> [B, S, H_local, D]
|
||||||
q = _usp_input_all_to_all(q, head_dim=2)
|
q = _usp_input_all_to_all(q, head_dim=2)
|
||||||
k = _usp_input_all_to_all(k, head_dim=2)
|
k = _usp_input_all_to_all(k, head_dim=2)
|
||||||
@@ -398,8 +405,66 @@ class USPAttention(nn.Module):
|
|||||||
out = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
out = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||||
|
|
||||||
# Ulysses-style All-to-All to restore original sharding
|
# Ulysses-style All-to-All to restore original sharding
|
||||||
if get_ulysses_parallel_world_size() > 1:
|
if sp_size > 1:
|
||||||
# -> [B, S_local, H, D]
|
# -> [B, S_local, H, D]
|
||||||
out = _usp_output_all_to_all(out, head_dim=2)
|
out = _usp_output_all_to_all(out, head_dim=2)
|
||||||
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
def _forward_with_replicated_prefix(
|
||||||
|
self,
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
v: torch.Tensor,
|
||||||
|
ctx_attn_metadata,
|
||||||
|
num_rep: int,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Ulysses attention where the first *num_rep* tokens are replicated
|
||||||
|
across SP ranks (e.g. text tokens) and should NOT be duplicated by the
|
||||||
|
all-to-all.
|
||||||
|
|
||||||
|
Strategy:
|
||||||
|
1. Split q/k/v into replicated prefix and SP-sharded suffix.
|
||||||
|
2. All-to-all only the sharded suffix (gathers sequence, shards heads).
|
||||||
|
3. Locally slice the replicated prefix to the same head shard.
|
||||||
|
4. Concatenate [prefix_h_local, gathered_suffix] and run attention.
|
||||||
|
5. Split output, all-to-all back the suffix, all-gather prefix heads.
|
||||||
|
"""
|
||||||
|
sp_size = get_ulysses_parallel_world_size()
|
||||||
|
sp_rank = get_sp_parallel_rank()
|
||||||
|
|
||||||
|
q_rep, q_shard = q[:, :num_rep], q[:, num_rep:]
|
||||||
|
k_rep, k_shard = k[:, :num_rep], k[:, num_rep:]
|
||||||
|
v_rep, v_shard = v[:, :num_rep], v[:, num_rep:]
|
||||||
|
|
||||||
|
q_shard = _usp_input_all_to_all(q_shard, head_dim=2)
|
||||||
|
k_shard = _usp_input_all_to_all(k_shard, head_dim=2)
|
||||||
|
v_shard = _usp_input_all_to_all(v_shard, head_dim=2)
|
||||||
|
|
||||||
|
h_local = q_shard.shape[2]
|
||||||
|
h_start = sp_rank * h_local
|
||||||
|
h_end = h_start + h_local
|
||||||
|
q_rep = q_rep[:, :, h_start:h_end, :].contiguous()
|
||||||
|
k_rep = k_rep[:, :, h_start:h_end, :].contiguous()
|
||||||
|
v_rep = v_rep[:, :, h_start:h_end, :].contiguous()
|
||||||
|
|
||||||
|
q = torch.cat([q_rep, q_shard], dim=1)
|
||||||
|
k = torch.cat([k_rep, k_shard], dim=1)
|
||||||
|
v = torch.cat([v_rep, v_shard], dim=1)
|
||||||
|
|
||||||
|
out = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||||
|
|
||||||
|
out_rep = out[:, :num_rep]
|
||||||
|
out_shard = out[:, num_rep:]
|
||||||
|
|
||||||
|
out_shard = _usp_output_all_to_all(out_shard, head_dim=2)
|
||||||
|
|
||||||
|
gathered = [torch.empty_like(out_rep) for _ in range(sp_size)]
|
||||||
|
torch.distributed.all_gather(
|
||||||
|
gathered,
|
||||||
|
out_rep.contiguous(),
|
||||||
|
group=get_sp_group().ulysses_group,
|
||||||
|
)
|
||||||
|
out_rep = torch.cat(gathered, dim=2)
|
||||||
|
|
||||||
|
return torch.cat([out_rep, out_shard], dim=1)
|
||||||
|
|||||||
@@ -264,7 +264,10 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
query, key, cos_sin_cache, is_neox=False
|
query, key, cos_sin_cache, is_neox=False
|
||||||
)
|
)
|
||||||
|
|
||||||
hidden_states = self.attn(query, key, value)
|
num_rep = (
|
||||||
|
encoder_hidden_states.shape[1] if encoder_hidden_states is not None else 0
|
||||||
|
)
|
||||||
|
hidden_states = self.attn(query, key, value, num_replicated_prefix=num_rep)
|
||||||
|
|
||||||
hidden_states = hidden_states.flatten(2, 3)
|
hidden_states = hidden_states.flatten(2, 3)
|
||||||
hidden_states = hidden_states.to(query.dtype)
|
hidden_states = hidden_states.to(query.dtype)
|
||||||
@@ -366,6 +369,7 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
attention_mask: Optional[torch.Tensor] = None,
|
attention_mask: Optional[torch.Tensor] = None,
|
||||||
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||||
|
num_replicated_prefix: int = 0,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
# Parallel in (QKV + MLP in) projection
|
# Parallel in (QKV + MLP in) projection
|
||||||
@@ -398,7 +402,9 @@ class Flux2ParallelSelfAttention(torch.nn.Module, AttentionModuleMixin):
|
|||||||
query, key = apply_flashinfer_rope_qk_inplace(
|
query, key = apply_flashinfer_rope_qk_inplace(
|
||||||
query, key, cos_sin_cache, is_neox=False
|
query, key, cos_sin_cache, is_neox=False
|
||||||
)
|
)
|
||||||
hidden_states = self.attn(query, key, value)
|
hidden_states = self.attn(
|
||||||
|
query, key, value, num_replicated_prefix=num_replicated_prefix
|
||||||
|
)
|
||||||
hidden_states = hidden_states.flatten(2, 3)
|
hidden_states = hidden_states.flatten(2, 3)
|
||||||
hidden_states = hidden_states.to(query.dtype)
|
hidden_states = hidden_states.to(query.dtype)
|
||||||
|
|
||||||
@@ -468,6 +474,7 @@ class Flux2SingleTransformerBlock(nn.Module):
|
|||||||
attn_output = self.attn(
|
attn_output = self.attn(
|
||||||
hidden_states=norm_hidden_states,
|
hidden_states=norm_hidden_states,
|
||||||
freqs_cis=freqs_cis,
|
freqs_cis=freqs_cis,
|
||||||
|
num_replicated_prefix=text_seq_len or 0,
|
||||||
**joint_attention_kwargs,
|
**joint_attention_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -879,6 +886,7 @@ class Flux2Transformer2DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
temb_mod_params=single_stream_mod,
|
temb_mod_params=single_stream_mod,
|
||||||
freqs_cis=freqs_cis,
|
freqs_cis=freqs_cis,
|
||||||
joint_attention_kwargs=joint_attention_kwargs,
|
joint_attention_kwargs=joint_attention_kwargs,
|
||||||
|
text_seq_len=num_txt_tokens,
|
||||||
)
|
)
|
||||||
# Remove text tokens from concatenated stream
|
# Remove text tokens from concatenated stream
|
||||||
hidden_states = hidden_states[:, num_txt_tokens:, ...]
|
hidden_states = hidden_states[:, num_txt_tokens:, ...]
|
||||||
|
|||||||
@@ -23,7 +23,6 @@ from sglang.multimodal_gen import envs
|
|||||||
from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType, STA_Mode
|
from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType, STA_Mode
|
||||||
from sglang.multimodal_gen.configs.pipeline_configs.wan import (
|
from sglang.multimodal_gen.configs.pipeline_configs.wan import (
|
||||||
Wan2_2_TI2V_5B_Config,
|
Wan2_2_TI2V_5B_Config,
|
||||||
WanI2V480PConfig,
|
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.cache.cache_dit_integration import (
|
from sglang.multimodal_gen.runtime.cache.cache_dit_integration import (
|
||||||
CacheDitConfig,
|
CacheDitConfig,
|
||||||
@@ -794,13 +793,9 @@ class DenoisingStage(PipelineStage):
|
|||||||
else:
|
else:
|
||||||
batch.did_sp_shard_latents = False
|
batch.did_sp_shard_latents = False
|
||||||
|
|
||||||
# For I2I tasks like QwenImageEdit, where the image latents is provided as condition, the image_latent (input image) should be
|
# image_latent must be sharded consistently with latents when it is
|
||||||
# replicated on all SP ranks, not sharded, as it provides global context.
|
# concatenated along the sequence dimension in the denoising loop.
|
||||||
# For Wan2_2_TI2V_5B_Config, it has very special settings
|
if batch.image_latent is not None:
|
||||||
if (
|
|
||||||
isinstance(server_args.pipeline_config, WanI2V480PConfig)
|
|
||||||
and batch.image_latent is not None
|
|
||||||
):
|
|
||||||
batch.image_latent, _ = server_args.pipeline_config.shard_latents_for_sp(
|
batch.image_latent, _ = server_args.pipeline_config.shard_latents_for_sp(
|
||||||
batch, batch.image_latent
|
batch, batch.image_latent
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2051,6 +2051,26 @@
|
|||||||
"expected_e2e_ms": 6359.87,
|
"expected_e2e_ms": 6359.87,
|
||||||
"expected_avg_denoise_ms": 99.47,
|
"expected_avg_denoise_ms": 99.47,
|
||||||
"expected_median_denoise_ms": 100.27
|
"expected_median_denoise_ms": 100.27
|
||||||
|
},
|
||||||
|
"flux_2_klein_ti2i_2_gpus": {
|
||||||
|
"stages_ms": {
|
||||||
|
"InputValidationStage": 40.19,
|
||||||
|
"TextEncodingStage": 88.84,
|
||||||
|
"ImageVAEEncodingStage": 80.81,
|
||||||
|
"LatentPreparationStage": 1.05,
|
||||||
|
"TimestepPreparationStage": 28.64,
|
||||||
|
"DenoisingStage": 354.04,
|
||||||
|
"DecodingStage": 11.11
|
||||||
|
},
|
||||||
|
"denoise_step_ms": {
|
||||||
|
"0": 33.54,
|
||||||
|
"1": 61.3,
|
||||||
|
"2": 86.9,
|
||||||
|
"3": 87.55
|
||||||
|
},
|
||||||
|
"expected_e2e_ms": 716.81,
|
||||||
|
"expected_avg_denoise_ms": 67.32,
|
||||||
|
"expected_median_denoise_ms": 74.1
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -932,8 +932,9 @@ def get_generate_fn(
|
|||||||
if is_image_url(image_path):
|
if is_image_url(image_path):
|
||||||
new_image_paths.append(download_image_from_url(str(image_path)))
|
new_image_paths.append(download_image_from_url(str(image_path)))
|
||||||
else:
|
else:
|
||||||
new_image_paths.append(Path(image_path))
|
local_path = Path(image_path)
|
||||||
if not image_path.exists():
|
new_image_paths.append(local_path)
|
||||||
|
if not local_path.exists():
|
||||||
pytest.skip(f"{case_id}: file missing: {image_path}")
|
pytest.skip(f"{case_id}: file missing: {image_path}")
|
||||||
|
|
||||||
image_paths = new_image_paths
|
image_paths = new_image_paths
|
||||||
|
|||||||
@@ -790,6 +790,15 @@ TWO_GPU_CASES_B = [
|
|||||||
),
|
),
|
||||||
T2I_sampling_params,
|
T2I_sampling_params,
|
||||||
),
|
),
|
||||||
|
DiffusionTestCase(
|
||||||
|
"flux_2_klein_ti2i_2_gpus",
|
||||||
|
DiffusionServerArgs(
|
||||||
|
model_path="black-forest-labs/FLUX.2-klein-4B",
|
||||||
|
modality="image",
|
||||||
|
num_gpus=2,
|
||||||
|
),
|
||||||
|
TI2I_sampling_params,
|
||||||
|
),
|
||||||
]
|
]
|
||||||
|
|
||||||
# Load global configuration
|
# Load global configuration
|
||||||
|
|||||||
Reference in New Issue
Block a user