[diffusion] fix: fix qwen-image with nunchaku (#21415)

This commit is contained in:
Mick
2026-03-26 16:31:44 +08:00
committed by GitHub
parent f289d173aa
commit 35720d9969
@@ -560,14 +560,26 @@ class QwenImageCrossAttention(nn.Module):
) )
if context_pre_only is not None and not context_pre_only: if context_pre_only is not None and not context_pre_only:
self.to_add_out = ReplicatedLinear(self.inner_dim, self.dim, bias=out_bias) self.to_add_out = ReplicatedLinear(
self.inner_dim,
self.dim,
bias=out_bias,
quant_config=quant_config,
prefix=f"{prefix}.to_add_out",
)
else: else:
self.to_add_out = None self.to_add_out = None
if not pre_only: if not pre_only:
self.to_out = nn.ModuleList([]) self.to_out = nn.ModuleList([])
self.to_out.append( self.to_out.append(
ReplicatedLinear(self.inner_dim, self.dim, bias=out_bias) ReplicatedLinear(
self.inner_dim,
self.dim,
bias=out_bias,
quant_config=quant_config,
prefix=f"{prefix}.to_out.0",
)
) )
else: else:
self.to_out = None self.to_out = None
@@ -706,8 +718,12 @@ class QwenImageTransformerBlock(nn.Module):
# Image processing modules # Image processing modules
self.img_mod = nn.Sequential( self.img_mod = nn.Sequential(
nn.SiLU(), nn.SiLU(),
nn.Linear( ReplicatedLinear(
dim, 6 * dim, bias=True dim,
6 * dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.img_mod",
), # For scale, shift, gate for norm1 and norm2 ), # For scale, shift, gate for norm1 and norm2
) )
self.img_norm1 = LayerNormScaleShift( self.img_norm1 = LayerNormScaleShift(
@@ -730,8 +746,12 @@ class QwenImageTransformerBlock(nn.Module):
# Text processing modules # Text processing modules
self.txt_mod = nn.Sequential( self.txt_mod = nn.Sequential(
nn.SiLU(), nn.SiLU(),
nn.Linear( ReplicatedLinear(
dim, 6 * dim, bias=True dim,
6 * dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.txt_mod",
), # For scale, shift, gate for norm1 and norm2 ), # For scale, shift, gate for norm1 and norm2
) )
self.txt_norm1 = LayerNormScaleShift( self.txt_norm1 = LayerNormScaleShift(
@@ -884,8 +904,8 @@ class QwenImageTransformerBlock(nn.Module):
modulate_index: Optional[torch.Tensor] = None, modulate_index: Optional[torch.Tensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> Tuple[torch.Tensor, torch.Tensor]:
# Get modulation parameters for both streams # Get modulation parameters for both streams
img_mod_params = self.img_mod[1](temb_img_silu) # [B, 6*dim] img_mod_params, _ = self.img_mod[1](temb_img_silu) # [B, 6*dim]
txt_mod_params = self.txt_mod[1](temb_txt_silu) # [B, 6*dim] txt_mod_params, _ = self.txt_mod[1](temb_txt_silu) # [B, 6*dim]
if ( if (
self.quant_config is not None self.quant_config is not None