diff --git a/modules/model_quant.py b/modules/model_quant.py index 307707708..9f2bd6a0b 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -37,7 +37,7 @@ def get_quant(name): return 'none' -def create_bnb_config(kwargs = None, allow_bnb: bool = True, module: str = 'Model'): +def create_bnb_config(kwargs = None, allow_bnb: bool = True, module: str = 'Model', modules_to_not_convert: list = []): from modules import shared, devices if len(shared.opts.bnb_quantization) > 0 and allow_bnb: if 'Model' in shared.opts.bnb_quantization or (module is not None and module in shared.opts.bnb_quantization) or module == 'any': @@ -49,7 +49,8 @@ def create_bnb_config(kwargs = None, allow_bnb: bool = True, module: str = 'Mode load_in_4bit=shared.opts.bnb_quantization_type in ['nf4', 'fp4'], bnb_4bit_quant_storage=shared.opts.bnb_quantization_storage, bnb_4bit_quant_type=shared.opts.bnb_quantization_type, - bnb_4bit_compute_dtype=devices.dtype + bnb_4bit_compute_dtype=devices.dtype, + #modules_to_not_convert=modules_to_not_convert, # ignored by bnb ) log.debug(f'Quantization: module={module} type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') if kwargs is None: @@ -60,7 +61,7 @@ def create_bnb_config(kwargs = None, allow_bnb: bool = True, module: str = 'Mode return kwargs -def create_ao_config(kwargs = None, allow_ao: bool = True, module: str = 'Model'): +def create_ao_config(kwargs = None, allow_ao: bool = True, module: str = 'Model', modules_to_not_convert: list = []): from modules import shared if len(shared.opts.torchao_quantization) > 0 and (shared.opts.torchao_quantization_mode == 'pre') and allow_ao: if 'Model' in shared.opts.torchao_quantization or (module is not None and module in shared.opts.torchao_quantization) or module == 'any': @@ -68,9 +69,9 @@ def create_ao_config(kwargs = None, allow_ao: bool = True, module: str = 'Model' if torchao is None: return kwargs if module in {'TE', 'LLM'}: - ao_config = transformers.TorchAoConfig(quant_type=shared.opts.torchao_quantization_type) + ao_config = transformers.TorchAoConfig(quant_type=shared.opts.torchao_quantization_type, modules_to_not_convert=modules_to_not_convert) else: - ao_config = diffusers.TorchAoConfig(shared.opts.torchao_quantization_type) + ao_config = diffusers.TorchAoConfig(shared.opts.torchao_quantization_type, modules_to_not_convert=modules_to_not_convert) log.debug(f'Quantization: module={module} type=torchao dtype={shared.opts.torchao_quantization_type}') if kwargs is None: return ao_config @@ -80,7 +81,7 @@ def create_ao_config(kwargs = None, allow_ao: bool = True, module: str = 'Model' return kwargs -def create_quanto_config(kwargs = None, allow_quanto: bool = True, module: str = 'Model'): +def create_quanto_config(kwargs = None, allow_quanto: bool = True, module: str = 'Model', modules_to_not_convert: list = []): from modules import shared if len(shared.opts.quanto_quantization) > 0 and allow_quanto: if 'Model' in shared.opts.quanto_quantization or (module is not None and module in shared.opts.quanto_quantization) or module == 'any': @@ -88,10 +89,10 @@ def create_quanto_config(kwargs = None, allow_quanto: bool = True, module: str = if optimum_quanto is None: return kwargs if module in {'TE', 'LLM'}: - quanto_config = transformers.QuantoConfig(weights=shared.opts.quanto_quantization_type) + quanto_config = transformers.QuantoConfig(weights=shared.opts.quanto_quantization_type, modules_to_not_convert=modules_to_not_convert) quanto_config.weights_dtype = quanto_config.weights else: - quanto_config = diffusers.QuantoConfig(weights_dtype=shared.opts.quanto_quantization_type) + quanto_config = diffusers.QuantoConfig(weights_dtype=shared.opts.quanto_quantization_type, modules_to_not_convert=modules_to_not_convert) quanto_config.activations = None # patch so it works with transformers quanto_config.weights = quanto_config.weights_dtype log.debug(f'Quantization: module={module} type=quanto dtype={shared.opts.quanto_quantization_type}') @@ -103,7 +104,7 @@ def create_quanto_config(kwargs = None, allow_quanto: bool = True, module: str = return kwargs -def create_sdnq_config(kwargs = None, allow_sdnq: bool = True, module: str = 'Model', weights_dtype: str = None): +def create_sdnq_config(kwargs = None, allow_sdnq: bool = True, module: str = 'Model', weights_dtype: str = None, modules_to_not_convert: list = []): from modules import devices, shared if len(shared.opts.sdnq_quantize_weights) > 0 and (shared.opts.sdnq_quantize_mode == 'pre') and allow_sdnq: if 'Model' in shared.opts.sdnq_quantize_weights or (module is not None and module in shared.opts.sdnq_quantize_weights) or module == 'any': @@ -114,15 +115,8 @@ def create_sdnq_config(kwargs = None, allow_sdnq: bool = True, module: str = 'Mo transformers.quantizers.auto.AUTO_QUANTIZATION_CONFIG_MAPPING["sdnq"] = SDNQConfig if weights_dtype is None: - if module in {"TE", "LLM"}: - if shared.opts.sdnq_quantize_weights_mode_te == "none": - return kwargs - elif shared.opts.sdnq_quantize_weights_mode_te in {"same as model", "default"}: - weights_dtype = shared.opts.sdnq_quantize_weights_mode - else: - weights_dtype = shared.opts.sdnq_quantize_weights_mode_te - elif shared.opts.sdnq_quantize_weights_mode == "none": - return kwargs + if module in {"TE", "LLM"} and shared.opts.sdnq_quantize_weights_mode_te not in {"same as model", "default"}: + weights_dtype = shared.opts.sdnq_quantize_weights_mode_te else: weights_dtype = shared.opts.sdnq_quantize_weights_mode if weights_dtype is None or weights_dtype == 'none': @@ -150,6 +144,7 @@ def create_sdnq_config(kwargs = None, allow_sdnq: bool = True, module: str = 'Mo dequantize_fp32=shared.opts.sdnq_dequantize_fp32, quantization_device=quantization_device, return_device=return_device, + modules_to_not_convert=modules_to_not_convert, ) log.debug(f'Quantization: module="{module}" type=sdnq dtype={weights_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} group_size={shared.opts.sdnq_quantize_weights_group_size} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} dequantize_fp32={shared.opts.sdnq_dequantize_fp32} quantize_with_gpu={shared.opts.sdnq_quantize_with_gpu} quantization_device={quantization_device} return_device={return_device}') if kwargs is None: @@ -180,25 +175,25 @@ def check_nunchaku(module: str = ''): return True -def create_config(kwargs = None, allow: bool = True, module: str = 'Model'): +def create_config(kwargs = None, allow: bool = True, module: str = 'Model', modules_to_not_convert = []): if kwargs is None: kwargs = {} - kwargs = create_sdnq_config(kwargs, allow_sdnq=allow, module=module) + kwargs = create_sdnq_config(kwargs, allow_sdnq=allow, module=module, modules_to_not_convert=modules_to_not_convert) if kwargs is not None and 'quantization_config' in kwargs: if debug: log.trace(f'Quantization: type=sdnq config={kwargs.get("quantization_config", None)}') return kwargs - kwargs = create_bnb_config(kwargs, allow_bnb=allow, module=module) + kwargs = create_bnb_config(kwargs, allow_bnb=allow, module=module, modules_to_not_convert=modules_to_not_convert) if kwargs is not None and 'quantization_config' in kwargs: if debug: log.trace(f'Quantization: type=bnb config={kwargs.get("quantization_config", None)}') return kwargs - kwargs = create_quanto_config(kwargs, allow_quanto=allow, module=module) + kwargs = create_quanto_config(kwargs, allow_quanto=allow, module=module, modules_to_not_convert=modules_to_not_convert) if kwargs is not None and 'quantization_config' in kwargs: if debug: log.trace(f'Quantization: type=quanto config={kwargs.get("quantization_config", None)}') return kwargs - kwargs = create_ao_config(kwargs, allow_ao=allow, module=module) + kwargs = create_ao_config(kwargs, allow_ao=allow, module=module, modules_to_not_convert=modules_to_not_convert) if kwargs is not None and 'quantization_config' in kwargs: if debug: log.trace(f'Quantization: type=torchao config={kwargs.get("quantization_config", None)}') @@ -331,20 +326,16 @@ def apply_layerwise(sd_model, quiet:bool=False): log.error(f'Quantization: type=layerwise {e}') -def sdnq_quantize_model(model, op=None, sd_model=None, do_gc=True): +def sdnq_quantize_model(model, op=None, sd_model=None, do_gc: bool = True, weights_dtype: str = None, modules_to_not_convert: list = []): global quant_last_model_name, quant_last_model_device # pylint: disable=global-statement from modules import devices, shared from modules.sdnq import apply_sdnq_to_module - model.eval() - backup_embeddings = None - if hasattr(model, "get_input_embeddings"): - backup_embeddings = copy.deepcopy(model.get_input_embeddings()) - - if shared.opts.sdnq_quantize_weights_mode_te != "default" and op is not None and "text_encoder" in op: - weights_dtype = shared.opts.sdnq_quantize_weights_mode_te - else: - weights_dtype = shared.opts.sdnq_quantize_weights_mode + if weights_dtype is None: + if op is not None and ("text_encoder" in op or op in {"TE", "LLM"}) and shared.opts.sdnq_quantize_weights_mode_te not in {"same as model", "default"}: + weights_dtype = shared.opts.sdnq_quantize_weights_mode_te + else: + weights_dtype = shared.opts.sdnq_quantize_weights_mode if weights_dtype is None or weights_dtype == 'none': return model @@ -361,9 +352,15 @@ def sdnq_quantize_model(model, op=None, sd_model=None, do_gc=True): quantization_device = None return_device = None - modules_to_not_convert = getattr(model, "_keep_in_fp32_modules", []) - if modules_to_not_convert is None: - modules_to_not_convert = [] + if getattr(model, "_keep_in_fp32_modules", None) is not None: + modules_to_not_convert.extend(model._keep_in_fp32_modules) + if model.__class__.__name__ == "ChromaTransformer2DModel": + modules_to_not_convert.append("distilled_guidance_layer") + + model.eval() + backup_embeddings = None + if hasattr(model, "get_input_embeddings"): + backup_embeddings = copy.deepcopy(model.get_input_embeddings()) model = apply_sdnq_to_module( model, @@ -558,7 +555,7 @@ def torchao_quantization(sd_model): return sd_model -def get_dit_args(load_config:dict={}, module:str=None, device_map:bool=False, allow_quant:bool=True): +def get_dit_args(load_config:dict={}, module:str=None, device_map:bool=False, allow_quant:bool=True, modules_to_not_convert: list = []): from modules import shared, devices config = load_config.copy() if 'torch_dtype' not in config: @@ -581,7 +578,7 @@ def get_dit_args(load_config:dict={}, module:str=None, device_map:bool=False, al elif shared.opts.device_map == 'gpu': config['device_map'] = devices.device if allow_quant: - quant_args = create_config(module=module) + quant_args = create_config(module=module, modules_to_not_convert=modules_to_not_convert) else: quant_args = {} return config, quant_args diff --git a/modules/sdnq/__init__.py b/modules/sdnq/__init__.py index c76d1e9f9..b4e251601 100644 --- a/modules/sdnq/__init__.py +++ b/modules/sdnq/__init__.py @@ -441,5 +441,8 @@ class SDNQConfig(QuantizationConfigMixin): accepted_weights = ["int8", "int7", "int6", "int5", "int4", "int3", "int2", "uint8", "uint7", "uint6", "uint5", "uint4", "uint3", "uint2", "uint1", "bool", "float8_e4m3fn", "float8_e4m3fnuz", "float8_e5m2", "float8_e5m2fnuz"] if self.weights_dtype not in accepted_weights: raise ValueError(f"Only support weights in {accepted_weights} but found {self.weights_dtype}") - if not isinstance(self.modules_to_not_convert, list): + + if self.modules_to_not_convert is None: + self.modules_to_not_convert = [] + elif not isinstance(self.modules_to_not_convert, list): self.modules_to_not_convert = [self.modules_to_not_convert]