Fix AuraFlow quant

This commit is contained in:
Disty0
2025-08-29 18:13:53 +03:00
parent 7e935743cd
commit 6e68dff381
2 changed files with 4 additions and 3 deletions
+2 -2
View File
@@ -305,7 +305,7 @@ def load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op='
elif model_type in ['AuraFlow']: # forced pipeline
from pipelines.model_auraflow import load_auraflow
sd_model = load_auraflow(checkpoint_info, diffusers_load_config)
allow_post_quant = True
allow_post_quant = False
elif model_type in ['FLUX']:
from pipelines.model_flux import load_flux
sd_model = load_flux(checkpoint_info, diffusers_load_config)
@@ -381,7 +381,7 @@ def load_diffuser_force(model_type, checkpoint_info, diffusers_load_config, op='
elif model_type in ['Kandinsky 2.2']:
from pipelines.model_kandinsky import load_kandinsky22
sd_model = load_kandinsky22(checkpoint_info, diffusers_load_config)
allow_post_quant = False
allow_post_quant = True
elif model_type in ['Kandinsky 3.0']:
from pipelines.model_kandinsky import load_kandinsky3
sd_model = load_kandinsky3(checkpoint_info, diffusers_load_config)
+2 -1
View File
@@ -90,7 +90,7 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz
try:
output_channel_size, channel_size = layer.weight.shape
except Exception as e:
raise ValueError(f"SDNQ: layer_class_name={layer_class_name} layer_weight_shape={layer.weight.shape} weights_dtype={weights_dtype} unsupported") from e
raise ValueError(f"SDNQ: param_name={param_name} layer_class_name={layer_class_name} layer_weight_shape={layer.weight.shape} weights_dtype={weights_dtype} unsupported") from e
if use_quantized_matmul:
use_quantized_matmul = weights_dtype in quantized_matmul_dtypes and channel_size >= 32 and output_channel_size >= 32
if use_quantized_matmul:
@@ -376,6 +376,7 @@ class SDNQQuantizer(DiffusersQuantizer):
keep_in_fp32_modules: List[str] = [],
**kwargs, # pylint: disable=unused-argument
):
print("AAAAA:", model.__class__.__name__, "|", getattr(model, "_keep_in_fp32_modules", "None"))
if keep_in_fp32_modules is not None:
self.modules_to_not_convert.extend(keep_in_fp32_modules)
elif getattr(model, "_keep_in_fp32_modules", None) is not None: