Rename NNCF to SDNQ and rename quant schemes

This commit is contained in:
Disty0
2025-05-26 02:39:51 +03:00
parent 9c2e15433e
commit 4453efee76
7 changed files with 114 additions and 122 deletions
+35 -35
View File
@@ -103,35 +103,35 @@ def create_quanto_config(kwargs = None, allow_quanto: bool = True, module: str =
return kwargs
def create_nncf_config(kwargs = None, allow_nncf: bool = True, module: str = 'Model'):
def create_sdnq_config(kwargs = None, allow_sdnq: bool = True, module: str = 'Model', weights_dtype: str = None):
from modules import shared
if len(shared.opts.nncf_compress_weights) > 0 and (shared.opts.nncf_compress_mode == 'pre') and allow_nncf:
if 'Model' in shared.opts.nncf_compress_weights or (module is not None and module in shared.opts.nncf_compress_weights) or module == 'any':
from modules.model_quant_nncf import NNCFQuantizer, NNCFConfig
diffusers.quantizers.auto.AUTO_QUANTIZER_MAPPING["nncf"] = NNCFQuantizer
transformers.quantizers.auto.AUTO_QUANTIZER_MAPPING["nncf"] = NNCFQuantizer
diffusers.quantizers.auto.AUTO_QUANTIZATION_CONFIG_MAPPING["nncf"] = NNCFConfig
transformers.quantizers.auto.AUTO_QUANTIZATION_CONFIG_MAPPING["nncf"] = NNCFConfig
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':
from modules.model_quant_sdnq import SDNQQuantizer, SDNQConfig
diffusers.quantizers.auto.AUTO_QUANTIZER_MAPPING["sdnq"] = SDNQQuantizer
transformers.quantizers.auto.AUTO_QUANTIZER_MAPPING["sdnq"] = SDNQQuantizer
diffusers.quantizers.auto.AUTO_QUANTIZATION_CONFIG_MAPPING["sdnq"] = SDNQConfig
transformers.quantizers.auto.AUTO_QUANTIZATION_CONFIG_MAPPING["sdnq"] = SDNQConfig
nncf_config = NNCFConfig(
weights_dtype=shared.opts.nncf_compress_weights_mode.lower(),
group_size=shared.opts.nncf_compress_weights_group_size,
use_int8_matmul=shared.opts.nncf_decompress_int8_matmul,
sdnq_config = SDNQConfig(
weights_dtype=weights_dtype if weights_dtype is not None else shared.opts.sdnq_quantize_weights_mode,
group_size=shared.opts.sdnq_quantize_weights_group_size,
use_int8_matmul=shared.opts.sdnq_decompress_int8_matmul,
)
log.debug(f'Quantization: module="{module}" type=nncf dtype={shared.opts.nncf_compress_weights_mode}')
log.debug(f'Quantization: module="{module}" type=sdnq dtype={shared.opts.sdnq_quantize_weights_mode}')
if kwargs is None:
return nncf_config
return sdnq_config
else:
kwargs['quantization_config'] = nncf_config
kwargs['quantization_config'] = sdnq_config
return kwargs
return kwargs
def check_quant(module: str = ''):
from modules import shared
if 'Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization or 'Model' in shared.opts.nncf_compress_weights:
if 'Model' in shared.opts.bnb_quantization or 'Model' in shared.opts.torchao_quantization or 'Model' in shared.opts.quanto_quantization or 'Model' in shared.opts.sdnq_quantize_weights:
return True
if module in shared.opts.bnb_quantization or module in shared.opts.torchao_quantization or module in shared.opts.quanto_quantization or module in shared.opts.nncf_compress_weights:
if module in shared.opts.bnb_quantization or module in shared.opts.torchao_quantization or module in shared.opts.quanto_quantization or module in shared.opts.sdnq_quantize_weights:
return True
return False
@@ -165,10 +165,10 @@ def create_config(kwargs = None, allow: bool = True, module: str = 'Model'):
if debug:
log.trace(f'Quantization: type=quanto config={kwargs.get("quantization_config", None)}')
return kwargs
kwargs = create_nncf_config(kwargs, allow_nncf=allow, module=module)
kwargs = create_sdnq_config(kwargs, allow_sdnq=allow, module=module)
if kwargs is not None and 'quantization_config' in kwargs:
if debug:
log.trace(f'Quantization: type=nncf config={kwargs.get("quantization_config", None)}')
log.trace(f'Quantization: type=sdnq config={kwargs.get("quantization_config", None)}')
return kwargs
return kwargs
@@ -298,18 +298,18 @@ def apply_layerwise(sd_model, quiet:bool=False):
log.error(f'Quantization: type=layerwise {e}')
def nncf_compress_model(model, op=None, sd_model=None, do_gc=True):
def sdnq_quantize_model(model, op=None, sd_model=None, do_gc=True):
global quant_last_model_name, quant_last_model_device # pylint: disable=global-statement
from modules import devices, shared
from modules.model_quant_nncf import apply_nncf_to_module
from modules.model_quant_sdnq import apply_sdnq_to_module
model.eval()
if model.__class__.__name__ in {"T5EncoderModel", "UMT5EncoderModel"}:
import torch
from modules.model_quant_nncf import NNCF_T5DenseGatedActDense # T5DenseGatedActDense uses fp32
from modules.model_quant_sdnq import SDNQ_T5DenseGatedActDense # T5DenseGatedActDense uses fp32
for i in range(len(model.encoder.block)):
model.encoder.block[i].layer[1].DenseReluDense = NNCF_T5DenseGatedActDense(
model.encoder.block[i].layer[1].DenseReluDense = SDNQ_T5DenseGatedActDense(
model.encoder.block[i].layer[1].DenseReluDense,
dtype=torch.float32 if devices.dtype != torch.bfloat16 else torch.bfloat16
)
@@ -318,15 +318,15 @@ def nncf_compress_model(model, op=None, sd_model=None, do_gc=True):
if hasattr(model, "get_input_embeddings"):
backup_embeddings = copy.deepcopy(model.get_input_embeddings())
num_bits = 8 if shared.opts.nncf_compress_weights_mode in {"INT8", "INT8_SYM", "INT8_ASYM"} else 4
is_asym_mode = shared.opts.nncf_compress_weights_mode in {"INT8", "INT4", "INT8_ASYM", "INT4_ASYM"}
model = apply_nncf_to_module(model, num_bits, is_asym_mode, quant_conv=shared.opts.nncf_quantize_conv_layers)
model.quantization_method = 'NNCF'
num_bits = 8 if shared.opts.sdnq_quantize_weights_mode in {"int8", "uint8"} else 4
is_asym_mode = shared.opts.sdnq_quantize_weights_mode in {"uint8", "uint4"}
model = apply_sdnq_to_module(model, num_bits, is_asym_mode, quant_conv=shared.opts.sdnq_quantize_conv_layers)
model.quantization_method = 'SDNQ'
if hasattr(model, "set_input_embeddings") and backup_embeddings is not None:
model.set_input_embeddings(backup_embeddings)
if op is not None and shared.opts.nncf_quantize_shuffle_weights:
if op is not None and shared.opts.sdnq_quantize_shuffle_weights:
if quant_last_model_name is not None:
if "." in quant_last_model_name:
last_model_names = quant_last_model_name.split(".")
@@ -347,14 +347,14 @@ def nncf_compress_model(model, op=None, sd_model=None, do_gc=True):
return model
def nncf_compress_weights(sd_model):
def sdnq_quantize_weights(sd_model):
try:
t0 = time.time()
from modules import shared, devices, sd_models
log.info(f"Quantization: type=NNCF modules={shared.opts.nncf_compress_weights}")
log.info(f"Quantization: type=SDNQ modules={shared.opts.sdnq_quantize_weights}")
global quant_last_model_name, quant_last_model_device # pylint: disable=global-statement
sd_model = sd_models.apply_function_to_model(sd_model, nncf_compress_model, shared.opts.nncf_compress_weights, op="nncf")
sd_model = sd_models.apply_function_to_model(sd_model, sdnq_quantize_model, shared.opts.sdnq_quantize_weights, op="sdnq")
if quant_last_model_name is not None:
if "." in quant_last_model_name:
last_model_names = quant_last_model_name.split(".")
@@ -366,9 +366,9 @@ def nncf_compress_weights(sd_model):
quant_last_model_device = None
t1 = time.time()
log.info(f"Quantization: type=NNCF time={t1-t0:.2f}")
log.info(f"Quantization: type=SDNQ time={t1-t0:.2f}")
except Exception as e:
log.warning(f"Quantization: type=NNCF {e}")
log.warning(f"Quantization: type=SDNQ {e}")
return sd_model
@@ -529,8 +529,8 @@ def get_dit_args(load_config:dict={}, module:str=None, device_map:bool=False, al
def do_post_load_quant(sd_model):
from modules import shared
if shared.opts.nncf_compress_weights and shared.opts.nncf_compress_mode == 'post' and not (shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx"):
sd_model = nncf_compress_weights(sd_model)
if shared.opts.sdnq_quantize_weights and shared.opts.sdnq_quantize_mode == 'post':
sd_model = sdnq_quantize_weights(sd_model)
if shared.opts.optimum_quanto_weights:
sd_model = optimum_quanto_weights(sd_model)
if shared.opts.torchao_quantization and shared.opts.torchao_quantization_mode == 'post':