2626from fastvideo .logger import init_logger
2727from fastvideo .models .dits .base import BaseDiT
2828from fastvideo .platforms import AttentionBackendEnum , current_platform
29+ from fastvideo .layers .quantization import QuantizationConfig
2930
3031from fastvideo .distributed .parallel_state import get_sp_world_size
3132
@@ -106,7 +107,9 @@ def __init__(self,
106107 window_size = (- 1 , - 1 ),
107108 qk_norm = True ,
108109 eps = 1e-6 ,
109- parallel_attention = False ) -> None :
110+ parallel_attention = False ,
111+ quant_config : QuantizationConfig | None = None ,
112+ prefix : str = "" ) -> None :
110113 assert dim % num_heads == 0
111114 super ().__init__ ()
112115 self .dim = dim
@@ -118,10 +121,10 @@ def __init__(self,
118121 self .parallel_attention = parallel_attention
119122
120123 # layers
121- self .to_q = ReplicatedLinear (dim , dim )
122- self .to_k = ReplicatedLinear (dim , dim )
123- self .to_v = ReplicatedLinear (dim , dim )
124- self .to_out = ReplicatedLinear (dim , dim )
124+ self .to_q = ReplicatedLinear (dim , dim , quant_config = quant_config , prefix = f" { prefix } .to_q" )
125+ self .to_k = ReplicatedLinear (dim , dim , quant_config = quant_config , prefix = f" { prefix } .to_k" )
126+ self .to_v = ReplicatedLinear (dim , dim , quant_config = quant_config , prefix = f" { prefix } .to_v" )
127+ self .to_out = ReplicatedLinear (dim , dim , quant_config = quant_config , prefix = f" { prefix } .to_out" )
125128 self .norm_q = RMSNorm (dim , eps = eps ) if qk_norm else nn .Identity ()
126129 self .norm_k = RMSNorm (dim , eps = eps ) if qk_norm else nn .Identity ()
127130
@@ -194,13 +197,15 @@ def __init__(
194197 qk_norm = True ,
195198 eps = 1e-6 ,
196199 supported_attention_backends : tuple [AttentionBackendEnum , ...]
197- | None = None
200+ | None = None ,
201+ quant_config : QuantizationConfig | None = None ,
202+ prefix : str = "" ,
198203 ) -> None :
199204 super ().__init__ (dim , num_heads , window_size , qk_norm , eps ,
200- supported_attention_backends )
205+ supported_attention_backends , quant_config = quant_config , prefix = prefix )
201206
202- self .add_k_proj = ReplicatedLinear (dim , dim )
203- self .add_v_proj = ReplicatedLinear (dim , dim )
207+ self .add_k_proj = ReplicatedLinear (dim , dim , quant_config = quant_config , prefix = f" { prefix } .add_k_proj" )
208+ self .add_v_proj = ReplicatedLinear (dim , dim , quant_config = quant_config , prefix = f" { prefix } .add_v_proj" )
204209 self .norm_added_k = RMSNorm (dim , eps = eps ) if qk_norm else nn .Identity ()
205210 self .norm_added_q = RMSNorm (dim , eps = eps ) if qk_norm else nn .Identity ()
206211
@@ -246,16 +251,17 @@ def __init__(self,
246251 added_kv_proj_dim : int | None = None ,
247252 supported_attention_backends : tuple [AttentionBackendEnum , ...]
248253 | None = None ,
254+ quant_config : QuantizationConfig | None = None ,
249255 prefix : str = "" ):
250256 super ().__init__ ()
251257
252258 # 1. Self-attention
253259 self .norm1 = FP32LayerNorm (dim , eps , elementwise_affine = False )
254- self .to_q = ReplicatedLinear (dim , dim , bias = True )
255- self .to_k = ReplicatedLinear (dim , dim , bias = True )
256- self .to_v = ReplicatedLinear (dim , dim , bias = True )
260+ self .to_q = ReplicatedLinear (dim , dim , bias = True , quant_config = quant_config , prefix = f" { prefix } .to_q" )
261+ self .to_k = ReplicatedLinear (dim , dim , bias = True , quant_config = quant_config , prefix = f" { prefix } .to_k" )
262+ self .to_v = ReplicatedLinear (dim , dim , bias = True , quant_config = quant_config , prefix = f" { prefix } .to_v" )
257263
258- self .to_out = ReplicatedLinear (dim , dim , bias = True )
264+ self .to_out = ReplicatedLinear (dim , dim , bias = True , quant_config = quant_config , prefix = f" { prefix } .to_out" )
259265 self .attn1 = DistributedAttention (
260266 num_heads = num_heads ,
261267 head_size = dim // num_heads ,
@@ -290,13 +296,17 @@ def __init__(self,
290296 self .attn2 = WanI2VCrossAttention (dim ,
291297 num_heads ,
292298 qk_norm = qk_norm ,
293- eps = eps )
299+ eps = eps ,
300+ quant_config = quant_config ,
301+ prefix = f"{ prefix } .attn2" )
294302 else :
295303 # T2V
296304 self .attn2 = WanT2VCrossAttention (dim ,
297305 num_heads ,
298306 qk_norm = qk_norm ,
299- eps = eps )
307+ eps = eps ,
308+ quant_config = quant_config ,
309+ prefix = f"{ prefix } .attn2" )
300310 self .cross_attn_residual_norm = ScaleResidualLayerNormScaleShift (
301311 dim ,
302312 norm_type = "layer" ,
@@ -306,7 +316,7 @@ def __init__(self,
306316 compute_dtype = torch .float32 )
307317
308318 # 3. Feed-forward
309- self .ffn = MLP (dim , ffn_dim , act_type = "gelu_pytorch_tanh" )
319+ self .ffn = MLP (dim , ffn_dim , act_type = "gelu_pytorch_tanh" , quant_config = quant_config , prefix = f" { prefix } .ffn" )
310320 self .mlp_residual = ScaleResidual ()
311321
312322 self .scale_shift_table = nn .Parameter (torch .randn (1 , 6 , dim ) / dim ** 0.5 )
@@ -406,17 +416,17 @@ def __init__(self,
406416 added_kv_proj_dim : int | None = None ,
407417 supported_attention_backends : tuple [AttentionBackendEnum , ...]
408418 | None = None ,
419+ quant_config : QuantizationConfig | None = None ,
409420 prefix : str = "" ):
410421 super ().__init__ ()
411422
412423 # 1. Self-attention
413424 self .norm1 = FP32LayerNorm (dim , eps , elementwise_affine = False )
414- self .to_q = ReplicatedLinear (dim , dim , bias = True )
415- self .to_k = ReplicatedLinear (dim , dim , bias = True )
416- self .to_v = ReplicatedLinear (dim , dim , bias = True )
417- self .to_gate_compress = ReplicatedLinear (dim , dim , bias = True )
418-
419- self .to_out = ReplicatedLinear (dim , dim , bias = True )
425+ self .to_q = ReplicatedLinear (dim , dim , bias = True , quant_config = quant_config , prefix = f"{ prefix } .to_q" )
426+ self .to_k = ReplicatedLinear (dim , dim , bias = True , quant_config = quant_config , prefix = f"{ prefix } .to_k" )
427+ self .to_v = ReplicatedLinear (dim , dim , bias = True , quant_config = quant_config , prefix = f"{ prefix } .to_v" )
428+ self .to_gate_compress = ReplicatedLinear (dim , dim , bias = True , quant_config = quant_config , prefix = f"{ prefix } .to_gate_compress" )
429+ self .to_out = ReplicatedLinear (dim , dim , bias = True , quant_config = quant_config , prefix = f"{ prefix } .to_out" )
420430 self .attn1 = DistributedAttention_VSA (
421431 num_heads = num_heads ,
422432 head_size = dim // num_heads ,
@@ -451,13 +461,17 @@ def __init__(self,
451461 self .attn2 = WanI2VCrossAttention (dim ,
452462 num_heads ,
453463 qk_norm = qk_norm ,
454- eps = eps )
464+ eps = eps ,
465+ quant_config = quant_config ,
466+ prefix = f"{ prefix } .attn2" )
455467 else :
456468 # T2V
457469 self .attn2 = WanT2VCrossAttention (dim ,
458470 num_heads ,
459471 qk_norm = qk_norm ,
460- eps = eps )
472+ eps = eps ,
473+ quant_config = quant_config ,
474+ prefix = f"{ prefix } .attn2" )
461475 self .cross_attn_residual_norm = ScaleResidualLayerNormScaleShift (
462476 dim ,
463477 norm_type = "layer" ,
@@ -467,7 +481,7 @@ def __init__(self,
467481 compute_dtype = torch .float32 )
468482
469483 # 3. Feed-forward
470- self .ffn = MLP (dim , ffn_dim , act_type = "gelu_pytorch_tanh" )
484+ self .ffn = MLP (dim , ffn_dim , act_type = "gelu_pytorch_tanh" , quant_config = quant_config , prefix = f" { prefix } .ffn" )
471485 self .mlp_residual = ScaleResidual ()
472486
473487 self .scale_shift_table = nn .Parameter (torch .randn (1 , 6 , dim ) / dim ** 0.5 )
@@ -556,6 +570,7 @@ class WanTransformer3DModel(BaseDiT):
556570 def __init__ (self , config : WanVideoConfig , hf_config : dict [str ,
557571 Any ]) -> None :
558572 super ().__init__ (config = config , hf_config = hf_config )
573+ self .quant_config = config .quant_config
559574
560575 inner_dim = config .num_attention_heads * config .attention_head_dim
561576 self .hidden_size = config .hidden_size
@@ -594,6 +609,7 @@ def __init__(self, config: WanVideoConfig, hf_config: dict[str,
594609 config .eps ,
595610 config .added_kv_proj_dim ,
596611 self ._supported_attention_backends ,
612+ quant_config = config .quant_config ,
597613 prefix = f"{ config .prefix } .blocks.{ i } " )
598614 for i in range (config .num_layers )
599615 ])
0 commit comments