[diffusion] fix: fix Z-Image Cache-DiT sequence-parallel override (#25305)
This commit is contained in:
@@ -256,6 +256,7 @@ class ZImageAttention(nn.Module):
|
|||||||
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||||
num_replicated_prefix: int = 0,
|
num_replicated_prefix: int = 0,
|
||||||
num_replicated_suffix: int = 0,
|
num_replicated_suffix: int = 0,
|
||||||
|
skip_sequence_parallel_override: bool = False,
|
||||||
):
|
):
|
||||||
if self.use_fused_qkv:
|
if self.use_fused_qkv:
|
||||||
qkv, _ = self.to_qkv(hidden_states)
|
qkv, _ = self.to_qkv(hidden_states)
|
||||||
@@ -360,6 +361,7 @@ class ZImageAttention(nn.Module):
|
|||||||
v,
|
v,
|
||||||
num_replicated_prefix=num_replicated_prefix,
|
num_replicated_prefix=num_replicated_prefix,
|
||||||
num_replicated_suffix=num_replicated_suffix,
|
num_replicated_suffix=num_replicated_suffix,
|
||||||
|
skip_sequence_parallel_override=skip_sequence_parallel_override,
|
||||||
)
|
)
|
||||||
hidden_states = hidden_states.flatten(2)
|
hidden_states = hidden_states.flatten(2)
|
||||||
|
|
||||||
@@ -452,6 +454,7 @@ class ZImageTransformerBlock(nn.Module):
|
|||||||
adaln_input: Optional[torch.Tensor] = None,
|
adaln_input: Optional[torch.Tensor] = None,
|
||||||
num_replicated_prefix: int = 0,
|
num_replicated_prefix: int = 0,
|
||||||
num_replicated_suffix: int = 0,
|
num_replicated_suffix: int = 0,
|
||||||
|
skip_sequence_parallel_override: bool = False,
|
||||||
):
|
):
|
||||||
if self.modulation:
|
if self.modulation:
|
||||||
assert adaln_input is not None
|
assert adaln_input is not None
|
||||||
@@ -467,6 +470,7 @@ class ZImageTransformerBlock(nn.Module):
|
|||||||
freqs_cis=freqs_cis,
|
freqs_cis=freqs_cis,
|
||||||
num_replicated_prefix=num_replicated_prefix,
|
num_replicated_prefix=num_replicated_prefix,
|
||||||
num_replicated_suffix=num_replicated_suffix,
|
num_replicated_suffix=num_replicated_suffix,
|
||||||
|
skip_sequence_parallel_override=skip_sequence_parallel_override,
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
_is_cuda
|
_is_cuda
|
||||||
@@ -509,6 +513,7 @@ class ZImageTransformerBlock(nn.Module):
|
|||||||
freqs_cis=freqs_cis,
|
freqs_cis=freqs_cis,
|
||||||
num_replicated_prefix=num_replicated_prefix,
|
num_replicated_prefix=num_replicated_prefix,
|
||||||
num_replicated_suffix=num_replicated_suffix,
|
num_replicated_suffix=num_replicated_suffix,
|
||||||
|
skip_sequence_parallel_override=skip_sequence_parallel_override,
|
||||||
)
|
)
|
||||||
x = x + self.attention_norm2(attn_out)
|
x = x + self.attention_norm2(attn_out)
|
||||||
|
|
||||||
@@ -1002,13 +1007,13 @@ class ZImageTransformer2DModel(CachableDiT, OffloadableDiTMixin):
|
|||||||
)
|
)
|
||||||
num_replicated_suffix = cap_seq_len if not use_full_unified_sequence else 0
|
num_replicated_suffix = cap_seq_len if not use_full_unified_sequence else 0
|
||||||
|
|
||||||
for layer_id, layer in enumerate(self.layers):
|
for layer in self.layers:
|
||||||
layer.attention.attn.skip_sequence_parallel = use_full_unified_sequence
|
|
||||||
unified = layer(
|
unified = layer(
|
||||||
unified,
|
unified,
|
||||||
unified_freqs_cis,
|
unified_freqs_cis,
|
||||||
adaln_input,
|
adaln_input,
|
||||||
num_replicated_suffix=num_replicated_suffix,
|
num_replicated_suffix=num_replicated_suffix,
|
||||||
|
skip_sequence_parallel_override=use_full_unified_sequence,
|
||||||
)
|
)
|
||||||
|
|
||||||
unified = self.all_final_layer[f"{patch_size}-{f_patch_size}"](
|
unified = self.all_final_layer[f"{patch_size}-{f_patch_size}"](
|
||||||
|
|||||||
Reference in New Issue
Block a user