1111import torch .nn as nn
1212from diffusers .models .attention import FeedForward
1313from diffusers .models .attention_processor import Attention
14+ from diffusers .models .normalization import RMSNorm
1415from diffusers .models .transformers .transformer_z_image import FeedForward as ZImageFeedForward
1516from diffusers .models .transformers .transformer_z_image import ZImageTransformer2DModel , ZImageTransformerBlock
1617from huggingface_hub import utils
1718
1819from nunchaku .models .unets .unet_sdxl import NunchakuSDXLFeedForward
1920
21+ from ...ops .gemm import svdq_gemm_w4a4_cuda
22+ from ...ops .quantize import svdq_quantize_w4a4_act_fuse_lora_cuda
2023from ...utils import get_precision , pad_tensor
2124from ..attention import NunchakuBaseAttention
2225from ..attention_processors .zimage import NunchakuZSingleStreamAttnProcessor
2326from ..embeddings import pack_rotemb
2427from ..linear import SVDQW4A4Linear
2528from ..utils import fuse_linears
26- from .utils import NunchakuModelLoaderMixin , patch_scale_key
29+ from .utils import NunchakuModelLoaderMixin , convert_fp16 , patch_scale_key
2730
2831
2932class NunchakuZImageRopeHook :
@@ -50,6 +53,77 @@ def __call__(self, module: nn.Module, input_args: tuple, input_kwargs: dict):
5053 return input_args , new_input_kwargs
5154
5255
56+ class NunchakuZImageFusedModule (nn .Module ):
57+ """
58+ Fused module for quantized QKV projection, RMS normalization, and rotary embedding for ZImage attention.
59+
60+ Parameters
61+ ----------
62+ qkv : SVDQW4A4Linear
63+ Quantized QKV projection layer.
64+ norm_q : RMSNorm
65+ RMSNorm for query.
66+ norm_k : RMSNorm
67+ RMSNorm for key.
68+ """
69+
70+ def __init__ (self , qkv : SVDQW4A4Linear , norm_q : RMSNorm , norm_k : RMSNorm ):
71+ super ().__init__ ()
72+ for name , param in qkv .named_parameters (prefix = "qkv_" ):
73+ setattr (self , name .replace ("." , "" ), param )
74+ self .qkv_precision = qkv .precision
75+ self .qkv_out_features = qkv .out_features
76+ for name , param in norm_q .named_parameters (prefix = "norm_q_" ):
77+ setattr (self , name .replace ("." , "" ), param )
78+ for name , param in norm_k .named_parameters (prefix = "norm_k_" ):
79+ setattr (self , name .replace ("." , "" ), param )
80+
81+ def forward (self , x : torch .Tensor , freqs_cis : Optional [torch .Tensor ] = None ):
82+ """
83+ Fuse QKV projection, RMS normalizaion and rotary embedding.
84+
85+ Parameters
86+ ----------
87+ x : torch.Tensor
88+ The hidden states tensor
89+ freqs_cis : torch.Tensor, optional
90+ The rotary embedding tensor
91+
92+ Returns
93+ -------
94+ The projection results of q, k, v. q result and k result are RMS-normalized and applied RoPE.
95+ """
96+ batch_size , seq_len , channels = x .shape
97+ x = x .view (batch_size * seq_len , channels )
98+ quantized_x , ascales , lora_act_out = svdq_quantize_w4a4_act_fuse_lora_cuda (
99+ x ,
100+ lora_down = self .qkv_proj_down ,
101+ smooth = self .qkv_smooth_factor ,
102+ fp4 = self .qkv_precision == "nvfp4" ,
103+ pad_size = 256 ,
104+ )
105+ output = torch .empty (batch_size * seq_len , self .qkv_out_features , dtype = x .dtype , device = x .device )
106+ svdq_gemm_w4a4_cuda (
107+ act = quantized_x ,
108+ wgt = self .qkv_qweight ,
109+ out = output ,
110+ ascales = ascales ,
111+ wscales = self .qkv_wscales ,
112+ lora_act_in = lora_act_out ,
113+ lora_up = self .qkv_proj_up ,
114+ bias = getattr (self , "qkv_bias" , None ),
115+ fp4 = self .qkv_precision == "nvfp4" ,
116+ alpha = 1.0 if self .qkv_precision == "nvfp4" else None ,
117+ wcscales = self .qkv_wcscales if self .qkv_precision == "nvfp4" else None ,
118+ norm_q = self .norm_q_weight ,
119+ norm_k = self .norm_k_weight ,
120+ rotary_emb = freqs_cis ,
121+ )
122+
123+ output = output .view (batch_size , seq_len , - 1 )
124+ return output
125+
126+
53127class NunchakuZImageAttention (NunchakuBaseAttention ):
54128 """
55129 Nunchaku-optimized Attention module for ZImage with quantized and fused QKV projections.
@@ -172,6 +246,14 @@ def _convert_z_image_ff(z_ff: ZImageFeedForward) -> FeedForward:
172246 return converted_ff
173247
174248
249+ def replace_fused_module (module , incompatible_keys ):
250+ assert isinstance (module , NunchakuZImageAttention )
251+ module .fused_module = NunchakuZImageFusedModule (module .to_qkv , module .norm_q , module .norm_k )
252+ del module .to_qkv
253+ del module .norm_q
254+ del module .norm_k
255+
256+
175257class NunchakuZImageFeedForward (NunchakuSDXLFeedForward ):
176258 """
177259 Quantized feed-forward block for :class:`NunchakuZImageTransformerBlock`.
@@ -218,6 +300,7 @@ def _patch_model(self, skip_refiners: bool = False, **kwargs):
218300 def _patch_transformer_block (block_list : List [ZImageTransformerBlock ]):
219301 for _ , block in enumerate (block_list ):
220302 block .attention = NunchakuZImageAttention (block .attention , ** kwargs )
303+ block .attention .register_load_state_dict_post_hook (replace_fused_module )
221304 block .feed_forward = NunchakuZImageFeedForward (block .feed_forward , ** kwargs )
222305
223306 def _convert_feed_forward (block_list : List [ZImageTransformerBlock ]):
@@ -323,10 +406,12 @@ def from_pretrained(cls, pretrained_model_name_or_path: str | os.PathLike[str],
323406
324407 print (f"quantization_config: { quantization_config } , rank={ rank } , skip_refiners={ skip_refiners } " )
325408
326- transformer ._patch_model (skip_refiners = skip_refiners , precision = precision , rank = rank )
409+ transformer ._patch_model (skip_refiners = skip_refiners , precision = precision , rank = rank , ** kwargs )
327410 transformer = transformer .to_empty (device = device )
328411
329412 patch_scale_key (transformer , model_state_dict )
413+ if torch_dtype == torch .float16 :
414+ convert_fp16 (transformer , model_state_dict )
330415
331416 transformer .load_state_dict (model_state_dict )
332417
0 commit comments