[diffusion] improve: improve qwen-image-edit performance to align with LightX2V (#15812)
Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
@@ -22,27 +22,10 @@ class QwenImageArchConfig(DiTArchConfig):
|
||||
axes_dims_rope: Tuple[int, int, int] = (16, 56, 56)
|
||||
zero_cond_t: bool = False
|
||||
|
||||
stacked_params_mapping: list[tuple[str, str, str]] = field(
|
||||
default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".to_qkv", ".to_q", "q"),
|
||||
(".to_qkv", ".to_k", "k"),
|
||||
(".to_qkv", ".to_v", "v"),
|
||||
(".to_added_qkv", ".add_q_proj", "q"),
|
||||
(".to_added_qkv", ".add_k_proj", "k"),
|
||||
(".to_added_qkv", ".add_v_proj", "v"),
|
||||
]
|
||||
)
|
||||
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
# QKV fusion mappings
|
||||
r"(.*)\.to_q\.(weight|bias)$": (r"\1.to_qkv.\2", 0, 3),
|
||||
r"(.*)\.to_k\.(weight|bias)$": (r"\1.to_qkv.\2", 1, 3),
|
||||
r"(.*)\.to_v\.(weight|bias)$": (r"\1.to_qkv.\2", 2, 3),
|
||||
r"(.*)\.add_q_proj\.(weight|bias)$": (r"\1.to_added_qkv.\2", 0, 3),
|
||||
r"(.*)\.add_k_proj\.(weight|bias)$": (r"\1.to_added_qkv.\2", 1, 3),
|
||||
r"(.*)\.add_v_proj\.(weight|bias)$": (r"\1.to_added_qkv.\2", 2, 3),
|
||||
# LoRA mappings
|
||||
r"^(transformer_blocks\.\d+\.attn\..*\.lora_[AB])\.default$": r"\1",
|
||||
}
|
||||
|
||||
@@ -156,16 +156,14 @@ class QwenImagePipelineConfig(ImagePipelineConfig):
|
||||
# img_shapes: for global entire image
|
||||
img_freqs, txt_freqs = rotary_emb(img_shapes, txt_seq_lens, device=device)
|
||||
|
||||
img_cos, img_sin = (
|
||||
img_freqs.real.to(dtype=dtype),
|
||||
img_freqs.imag.to(dtype=dtype),
|
||||
)
|
||||
txt_cos, txt_sin = (
|
||||
txt_freqs.real.to(dtype=dtype),
|
||||
txt_freqs.imag.to(dtype=dtype),
|
||||
)
|
||||
img_cos_half = img_freqs.real.to(dtype=torch.float32).contiguous()
|
||||
img_sin_half = img_freqs.imag.to(dtype=torch.float32).contiguous()
|
||||
txt_cos_half = txt_freqs.real.to(dtype=torch.float32).contiguous()
|
||||
txt_sin_half = txt_freqs.imag.to(dtype=torch.float32).contiguous()
|
||||
|
||||
return (img_cos, img_sin), (txt_cos, txt_sin)
|
||||
img_cos_sin_cache = torch.cat([img_cos_half, img_sin_half], dim=-1)
|
||||
txt_cos_sin_cache = torch.cat([txt_cos_half, txt_sin_half], dim=-1)
|
||||
return img_cos_sin_cache, txt_cos_sin_cache
|
||||
|
||||
def _prepare_cond_kwargs(self, batch, prompt_embeds, rotary_emb, device, dtype):
|
||||
batch_size = prompt_embeds[0].shape[0]
|
||||
@@ -184,15 +182,15 @@ class QwenImagePipelineConfig(ImagePipelineConfig):
|
||||
] * batch_size
|
||||
txt_seq_lens = [prompt_embeds[0].shape[1]]
|
||||
|
||||
(img_cos, img_sin), (txt_cos, txt_sin) = self.get_freqs_cis(
|
||||
freqs_cis = self.get_freqs_cis(
|
||||
img_shapes, txt_seq_lens, rotary_emb, device, dtype
|
||||
)
|
||||
|
||||
img_cos = shard_rotary_emb_for_sp(img_cos)
|
||||
img_sin = shard_rotary_emb_for_sp(img_sin)
|
||||
img_cache, txt_cache = freqs_cis
|
||||
img_cache = shard_rotary_emb_for_sp(img_cache)
|
||||
return {
|
||||
"txt_seq_lens": txt_seq_lens,
|
||||
"freqs_cis": ((img_cos, img_sin), (txt_cos, txt_sin)),
|
||||
"freqs_cis": (img_cache, txt_cache),
|
||||
}
|
||||
|
||||
def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype):
|
||||
@@ -252,7 +250,7 @@ class QwenImageEditPipelineConfig(QwenImagePipelineConfig):
|
||||
],
|
||||
] * batch_size
|
||||
txt_seq_lens = [prompt_embeds[0].shape[1]]
|
||||
(img_cos, img_sin), (txt_cos, txt_sin) = QwenImagePipelineConfig.get_freqs_cis(
|
||||
freqs_cis = QwenImagePipelineConfig.get_freqs_cis(
|
||||
img_shapes, txt_seq_lens, rotary_emb, device, dtype
|
||||
)
|
||||
|
||||
@@ -261,20 +259,14 @@ class QwenImageEditPipelineConfig(QwenImagePipelineConfig):
|
||||
1 * (height // vae_scale_factor // 2) * (width // vae_scale_factor // 2)
|
||||
)
|
||||
|
||||
noisy_img_cos = shard_rotary_emb_for_sp(img_cos[:noisy_img_seq_len, :])
|
||||
noisy_img_sin = shard_rotary_emb_for_sp(img_sin[:noisy_img_seq_len, :])
|
||||
|
||||
# concat back the img_cos for input image (since it is not sp-shared later)
|
||||
img_cos = torch.cat([noisy_img_cos, img_cos[noisy_img_seq_len:, :]], dim=0).to(
|
||||
device=device
|
||||
)
|
||||
img_sin = torch.cat([noisy_img_sin, img_sin[noisy_img_seq_len:, :]], dim=0).to(
|
||||
device=device
|
||||
)
|
||||
|
||||
img_cache, txt_cache = freqs_cis
|
||||
noisy_img_cache = shard_rotary_emb_for_sp(img_cache[:noisy_img_seq_len, :])
|
||||
img_cache = torch.cat(
|
||||
[noisy_img_cache, img_cache[noisy_img_seq_len:, :]], dim=0
|
||||
).to(device=device)
|
||||
return {
|
||||
"txt_seq_lens": txt_seq_lens,
|
||||
"freqs_cis": ((img_cos, img_sin), (txt_cos, txt_sin)),
|
||||
"freqs_cis": (img_cache, txt_cache),
|
||||
}
|
||||
|
||||
def preprocess_condition_image(
|
||||
@@ -441,16 +433,28 @@ class QwenImageEditPlusPipelineConfig(QwenImageEditPipelineConfig):
|
||||
] * batch_size
|
||||
txt_seq_lens = [prompt_embeds[0].shape[1]]
|
||||
|
||||
(img_cos, img_sin), (txt_cos, txt_sin) = (
|
||||
QwenImageEditPlusPipelineConfig.get_freqs_cis(
|
||||
img_shapes, txt_seq_lens, rotary_emb, device, dtype
|
||||
)
|
||||
freqs_cis = QwenImageEditPlusPipelineConfig.get_freqs_cis(
|
||||
img_shapes, txt_seq_lens, rotary_emb, device, dtype
|
||||
)
|
||||
|
||||
# perform sp shard on noisy image tokens
|
||||
noisy_img_seq_len = (
|
||||
1 * (height // vae_scale_factor // 2) * (width // vae_scale_factor // 2)
|
||||
)
|
||||
|
||||
if isinstance(freqs_cis[0], torch.Tensor) and freqs_cis[0].dim() == 2:
|
||||
img_cache, txt_cache = freqs_cis
|
||||
noisy_img_cache = shard_rotary_emb_for_sp(img_cache[:noisy_img_seq_len, :])
|
||||
img_cache = torch.cat(
|
||||
[noisy_img_cache, img_cache[noisy_img_seq_len:, :]], dim=0
|
||||
).to(device=device)
|
||||
return {
|
||||
"txt_seq_lens": txt_seq_lens,
|
||||
"freqs_cis": (img_cache, txt_cache),
|
||||
"img_shapes": img_shapes,
|
||||
}
|
||||
|
||||
(img_cos, img_sin), (txt_cos, txt_sin) = freqs_cis
|
||||
noisy_img_cos = shard_rotary_emb_for_sp(img_cos[:noisy_img_seq_len, :])
|
||||
noisy_img_sin = shard_rotary_emb_for_sp(img_sin[:noisy_img_seq_len, :])
|
||||
|
||||
|
||||
Reference in New Issue
Block a user