diff --git a/CHANGELOG.md b/CHANGELOG.md index 629b7fdee..1cca4e8b8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,17 @@ # Change Log for SD.Next +## Update for 2025-05-12 + +Curious how your system is performing? +Run a built-in benchmark and compare to over 15k unique results world-wide: (Benchmark data)[https://vladmandic.github.io/sd-extension-system-info/pages/benchmark.html]! +From slowest 0.02 it/s running on 6th gen CPU without acceleration up to 275 it/s running on tuned GH100 system! + +- **Wiki** + - Updates for: *WSL, ZLUDA, ROCm* +- **Compute** + - NNCF: added experimental support for direct INT8 MatMul + + ## Update for 2025-05-12 ### Highlights for 2025-05-12 diff --git a/README.md b/README.md index 4331f85ca..c93ec9fd8 100644 --- a/README.md +++ b/README.md @@ -1,7 +1,7 @@
SD.Next -**Image Diffusion implementation with advanced features** +# SD.Next: All-in-one WebUI for AI generative image and video creation ![Last update](https://img.shields.io/github/last-commit/vladmandic/sdnext?svg=true) ![License](https://img.shields.io/github/license/vladmandic/sdnext?svg=true) @@ -29,9 +29,9 @@ All individual features are not listed here, instead check [ChangeLog](CHANGELOG - Multiple UIs! ▹ **Standard | Modern** - Multiple [diffusion models](https://vladmandic.github.io/sdnext-docs/Model-Support/)! -- Built-in Control for Text, Image, Batch and video processing! +- Built-in Control for Text, Image, Batch and Video processing! - Multiplatform! - ▹ **Windows | Linux | MacOS | nVidia | AMD | IntelArc/IPEX | DirectML | OpenVINO | ONNX+Olive | ZLUDA** + ▹ **Windows | Linux | MacOS | nVidia CUDA | AMD ROCm | IntelArc/IPEX | DirectML | OpenVINO | ONNX+Olive | ZLUDA** - Platform specific autodetection and tuning performed on install - Optimized processing with latest `torch` developments with built-in support for model compile, quantize and compress Compile backends: *Triton | StableFast | DeepCache | OneDiff | TeaCache | etc.* diff --git a/javascript/extraNetworks.js b/javascript/extraNetworks.js index 3e0838683..83f97ba57 100644 --- a/javascript/extraNetworks.js +++ b/javascript/extraNetworks.js @@ -6,11 +6,11 @@ let totalCards = -1; const getENActiveTab = () => { let tabName = ''; - if (gradioApp().getElementById('txt2img_prompt').checkVisibility()) return 'txt2img'; - if (gradioApp().getElementById('img2img_prompt').checkVisibility()) return 'img2img'; - if (gradioApp().getElementById('control_prompt').checkVisibility()) return 'control'; - if (gradioApp().getElementById('video_prompt').checkVisibility()) return 'video'; - if (gradioApp().getElementById('framepack_prompt_row').checkVisibility()) return 'framepack'; + if (gradioApp().getElementById('txt2img_prompt')?.checkVisibility()) return 'txt2img'; + if (gradioApp().getElementById('img2img_prompt')?.checkVisibility()) return 'img2img'; + if (gradioApp().getElementById('control_prompt')?.checkVisibility()) return 'control'; + if (gradioApp().getElementById('video_prompt')?.checkVisibility()) return 'video'; + if (gradioApp().getElementById('framepack_prompt_row')?.checkVisibility()) return 'framepack'; // legacy method if (gradioApp().getElementById('tab_txt2img').style.display === 'block') tabName = 'txt2img'; else if (gradioApp().getElementById('tab_img2img').style.display === 'block') tabName = 'img2img'; diff --git a/modules/lora/lora_apply.py b/modules/lora/lora_apply.py index 458bcb60d..0edef9831 100644 --- a/modules/lora/lora_apply.py +++ b/modules/lora/lora_apply.py @@ -151,7 +151,8 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G new_weight = dequant_weight.to(devices.device, dtype=torch.float32) + lora_weights.to(devices.device, dtype=torch.float32) self.weight = torch.nn.Parameter(new_weight, requires_grad=False) self.pre_ops.pop("0") - self = nncf_compress_layer(self, num_bits, is_asym_mode, torch_dtype=devices.dtype, quant_conv=shared.opts.nncf_quantize_conv_layers) + self._custom_forward_fn = None + self = nncf_compress_layer(self, num_bits, is_asym_mode, torch_dtype=devices.dtype, quant_conv=shared.opts.nncf_quantize_conv_layers, group_size=shared.opts.nncf_compress_weights_group_size, use_int8_matmul=shared.opts.nncf_decompress_int8_matmul) self = self.to(device) del dequant_weight except Exception as e: diff --git a/modules/model_quant.py b/modules/model_quant.py index 467d3cc86..8414ae91b 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -118,7 +118,11 @@ def create_nncf_config(kwargs = None, allow_nncf: bool = True, module: str = 'Mo diffusers.quantizers.auto.AUTO_QUANTIZATION_CONFIG_MAPPING["nncf"] = NNCFConfig transformers.quantizers.auto.AUTO_QUANTIZATION_CONFIG_MAPPING["nncf"] = NNCFConfig - nncf_config = NNCFConfig(weights_dtype=shared.opts.nncf_compress_weights_mode.lower()) + 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, + ) log.debug(f'Quantization: module="{module}" type=nncf dtype={shared.opts.nncf_compress_weights_mode}') if kwargs is None: return nncf_config diff --git a/modules/model_quant_nncf.py b/modules/model_quant_nncf.py index cdcf054f5..81b0b7d55 100644 --- a/modules/model_quant_nncf.py +++ b/modules/model_quant_nncf.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List, Optional, Union +from typing import Any, Dict, List, Tuple, Optional, Union from dataclasses import dataclass from enum import Enum @@ -45,8 +45,7 @@ class QuantizationMethod(str, Enum): NNCF = "nncf" -# de-abstracted and modified from the actual quant functions of nncf 2.16.0: -def nncf_compress_layer(layer, num_bits, is_asym_mode, torch_dtype=None, quant_conv=False, param_name=None): +def nncf_compress_layer(layer, num_bits, is_asym_mode, torch_dtype=None, quant_conv=False, group_size=0, use_int8_matmul=False, param_name=None): if layer.__class__.__name__ in allowed_types: if torch_dtype is None: torch_dtype = devices.dtype @@ -56,16 +55,18 @@ def nncf_compress_layer(layer, num_bits, is_asym_mode, torch_dtype=None, quant_c if is_asym_mode or not quant_conv: # don't quant convs with asym mode return layer reduction_axes = [i for i in range(layer.weight.ndim) if i != 0] + use_int8_matmul = False if layer.__class__.__name__ in conv_transpose_types: if is_asym_mode or not quant_conv: # don't quant convs with asym mode return layer reduction_axes = [i for i in range(layer.weight.ndim) if i != 1] + use_int8_matmul = False else: reduction_axes = -1 - if shared.opts.nncf_compress_weights_group_size > 0 or (num_bits == 4 and shared.opts.nncf_compress_weights_group_size != -1): - group_size = shared.opts.nncf_compress_weights_group_size - channel_size = layer.weight.shape[-1] + channel_size = layer.weight.shape[-1] + use_int8_matmul = use_int8_matmul and not is_asym_mode and channel_size >= 1024 and layer.weight.shape[0] >= 1024 + if not use_int8_matmul and (group_size > 0 or (num_bits == 4 and group_size != -1)): if group_size == 0: group_size = 64 num_of_groups = channel_size // group_size @@ -97,48 +98,24 @@ def nncf_compress_layer(layer, num_bits, is_asym_mode, torch_dtype=None, quant_c layer.weight.data = layer.weight.data.to(devices.device, dtype=torch.float32) if is_asym_mode: - level_low = 0 - level_high = 2**num_bits - 1 - min_values = torch.amin(layer.weight, dim=reduction_axes, keepdims=True) # [a1, r, a2] -> [a1, 1, a2] - max_values = torch.amax(layer.weight, dim=reduction_axes, keepdims=True) # [a1, r, a2] -> [a1, 1, a2] - - levels = level_high - level_low + 1 - scale = ((max_values - min_values) / (levels - 1)).to(dtype=torch.float32) - eps = torch.finfo(scale.dtype).eps - - scale = torch.where(torch.abs(scale) < eps, eps, scale) - zero_point = (level_low - (min_values / scale)).to(dtype=torch.float32) - - #this is for packing zero_point to int: - #zero_point = level_low - torch.round(min_values / scale) - #zero_point = torch.clip(zero_point.to(dtype=torch.int32), level_low, level_high) + scale, zero_point = get_int_scale_asymmetric(layer.weight, reduction_axes, num_bits) else: - factor = 2 ** (num_bits - 1) - - w_abs_min = torch.abs(torch.amin(layer.weight, dim=reduction_axes, keepdims=True)) - w_max = torch.amax(layer.weight, dim=reduction_axes, keepdims=True) - - scale = torch.where(w_abs_min >= w_max, w_abs_min, -w_max) - scale /= factor - - eps = torch.finfo(scale.dtype).eps - scale = torch.where(torch.abs(scale) < eps, eps, scale) + scale = get_int_scale_symmetric(layer.weight, reduction_axes, num_bits) zero_point = None + compressed_weight = quantize_int(layer.weight, scale, zero_point, is_asym_mode, num_bits) - dtype = torch.uint8 if is_asym_mode else torch.int8 - level_low = 0 if is_asym_mode else -(2 ** (num_bits - 1)) - level_high = 2**num_bits - 1 if is_asym_mode else 2 ** (num_bits - 1) - 1 - - compressed_weight = layer.weight.data / scale if not shared.opts.nncf_decompress_fp32: - scale = scale.to(torch_dtype) + scale = scale.to(torch_dtype) + if zero_point is not None: + zero_point = zero_point.to(torch_dtype) - if zero_point is not None: - compressed_weight += zero_point - zero_point = zero_point.to(scale.dtype) - - compressed_weight = torch.round(compressed_weight) - compressed_weight = torch.clip(compressed_weight, level_low, level_high).to(dtype) + if use_int8_matmul: + layer._custom_forward_fn = linear_forward_int8_matmul + scale = scale.squeeze(-1) + if num_bits == 8: + compressed_weight = compressed_weight.transpose(0,1) + else: + layer._custom_forward_fn = None if num_bits == 4: if is_asym_mode: @@ -148,6 +125,7 @@ def nncf_compress_layer(layer, num_bits, is_asym_mode, torch_dtype=None, quant_c compressed_weight_shape=compressed_weight.shape, result_dtype=torch_dtype, result_shape=result_shape, + use_int8_matmul=use_int8_matmul, ) else: decompressor = INT4SymmetricWeightsDecompressor( @@ -155,6 +133,7 @@ def nncf_compress_layer(layer, num_bits, is_asym_mode, torch_dtype=None, quant_c compressed_weight_shape=compressed_weight.shape, result_dtype=torch_dtype, result_shape=result_shape, + use_int8_matmul=use_int8_matmul, ) else: if is_asym_mode: @@ -163,12 +142,14 @@ def nncf_compress_layer(layer, num_bits, is_asym_mode, torch_dtype=None, quant_c zero_point=zero_point.data, result_dtype=torch_dtype, result_shape=result_shape, + use_int8_matmul=use_int8_matmul, ) else: decompressor = INT8SymmetricWeightsDecompressor( scale=scale.data, result_dtype=torch_dtype, result_shape=result_shape, + use_int8_matmul=use_int8_matmul, ) compressed_weight = decompressor.pack_weight(compressed_weight) @@ -188,7 +169,16 @@ def apply_nncf_to_module(model, num_bits, is_asym_mode, quant_conv=False): return model for param_name, module in model.named_children(): if module.__class__.__name__.startswith("NNCF") and hasattr(module, "weight") and module.weight is not None: - module = nncf_compress_layer(module, num_bits, is_asym_mode, torch_dtype=devices.dtype, quant_conv=quant_conv, param_name=param_name) + module = nncf_compress_layer( + module, + num_bits, + is_asym_mode, + torch_dtype=devices.dtype, + quant_conv=quant_conv, + group_size=shared.opts.nncf_compress_weights_group_size, + use_int8_matmul=shared.opts.nncf_decompress_int8_matmul, + param_name=param_name, + ) module = apply_nncf_to_module(module, num_bits, is_asym_mode, quant_conv=quant_conv) return model @@ -253,7 +243,9 @@ class NNCFQuantizer(DiffusersQuantizer): self.quantization_config.num_bits, self.quantization_config.is_asym_mode, torch_dtype=self.torch_dtype, - param_name=param_name + group_size=self.quantization_config.group_size, + use_int8_matmul=self.quantization_config.use_int8_matmul, + param_name=param_name, ) def adjust_max_memory(self, max_memory: Dict[str, Union[int, str]]) -> Dict[str, Union[int, str]]: @@ -341,6 +333,8 @@ class NNCFConfig(QuantizationConfigMixin): def __init__( self, weights_dtype: str = "int8_sym", + group_size: int = 0, + use_int8_matmul: bool = False, modules_to_not_convert: Optional[List[str]] = None, **kwargs, ): @@ -353,7 +347,8 @@ class NNCFConfig(QuantizationConfigMixin): self.num_bits = 8 if self.weights_dtype in {"int8", "uint8"} else 4 self.is_asym_mode = self.weights_dtype in {"uint8", "uint4"} self.is_integer = True - self.group_size = -1 + self.group_size = group_size + self.use_int8_matmul = use_int8_matmul def post_init(self): r""" @@ -384,29 +379,43 @@ class NNCF_T5DenseGatedActDense(torch.nn.Module): # forward can't find what self return hidden_states -# WeightsDecompressor classes and functions are modified from NNCF 2.16.0 +def get_int_scale_asymmetric(weight: torch.FloatTensor, reduction_axes: List[int], num_bits: int) -> Tuple[torch.FloatTensor, torch.FloatTensor]: + level_low = 0 + level_high = 2**num_bits -def unpack_uint4(packed_tensor: torch.Tensor, shape: torch.Size) -> torch.Tensor: - return torch.stack((torch.bitwise_and(packed_tensor, 15), torch.bitwise_right_shift(packed_tensor, 4)), dim=-1).reshape(shape) + min_values = torch.amin(weight, dim=reduction_axes, keepdims=True) + max_values = torch.amax(weight, dim=reduction_axes, keepdims=True) + scale = ((max_values - min_values) / (level_high - 1)) + + eps = torch.finfo(scale.dtype).eps # prevent divison by 0 + scale = torch.where(torch.abs(scale) < eps, eps, scale) + zero_point = (level_low - (min_values / scale)) + return scale, zero_point -def unpack_int4(packed_tensor: torch.Tensor, shape: torch.Size, dtype: Optional[torch.dtype] = torch.int8) -> torch.Tensor: - return unpack_uint4(packed_tensor, shape).to(dtype=dtype) - 8 +def get_int_scale_symmetric(weight: torch.FloatTensor, reduction_axes: List[int], num_bits: int) -> torch.FloatTensor: + w_abs_min = torch.abs(torch.amin(weight, dim=reduction_axes, keepdims=True)) + w_max = torch.amax(weight, dim=reduction_axes, keepdims=True) + scale = torch.where(w_abs_min >= w_max, w_abs_min, -w_max) / (2 ** (num_bits - 1)) + + eps = torch.finfo(scale.dtype).eps # prevent divison by 0 + scale = torch.where(torch.abs(scale) < eps, eps, scale) + return scale -def pack_uint4(tensor: torch.Tensor) -> torch.Tensor: - if tensor.dtype != torch.uint8: - raise RuntimeError(f"Invalid tensor dtype {tensor.type}. torch.uint8 type is supported.") - packed_tensor = tensor.contiguous().reshape(-1, 2) - packed_tensor = torch.bitwise_and(packed_tensor[..., ::2], 15) | packed_tensor[..., 1::2] << 4 - return packed_tensor +def quantize_int(weight: torch.FloatTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, is_asym_mode: bool, num_bits: int, flatten: Optional[bool] = False) -> torch.ByteTensor: + dtype = torch.uint8 if is_asym_mode else torch.int8 + level_low = 0 if is_asym_mode else -(2 ** (num_bits - 1)) + level_high = 2**num_bits - 1 if is_asym_mode else 2 ** (num_bits - 1) - 1 + compressed_weight = weight / scale + if zero_point is not None: + compressed_weight += zero_point -def pack_int4(tensor: torch.Tensor) -> torch.Tensor: - if tensor.dtype != torch.int8: - raise RuntimeError(f"Invalid tensor dtype {tensor.type}. torch.int8 type is supported.") - tensor = tensor + 8 - return pack_uint4(tensor.to(dtype=torch.uint8)) + compressed_weight = torch.round(compressed_weight).clamp_(level_low, level_high).to(dtype) + if flatten: + compressed_weight = compressed_weight.flatten(0,-2) + return compressed_weight def decompress_asymmetric(input: torch.Tensor, scale: torch.Tensor, zero_point: torch.Tensor, dtype: torch.dtype, result_shape: torch.Size) -> torch.Tensor: @@ -431,15 +440,75 @@ def decompress_int4_symmetric(input: torch.Tensor, scale: torch.Tensor, shape: t return decompress_symmetric(unpack_int4(input, shape, dtype=scale.dtype), scale, dtype, result_shape) -if shared.opts.nncf_decompress_compile: - try: - torch._dynamo.config.cache_size_limit = max(8192, torch._dynamo.config.cache_size_limit) # pylint: disable=protected-access - decompress_asymmetric = torch.compile(decompress_asymmetric, fullgraph=True) - decompress_symmetric = torch.compile(decompress_symmetric, fullgraph=True) - decompress_int4_asymmetric = torch.compile(decompress_int4_asymmetric, fullgraph=True) - decompress_int4_symmetric = torch.compile(decompress_int4_symmetric, fullgraph=True) - except Exception as e: - shared.log.warning(f"Quantization: type=nncf Decompress using torch.compile is not available: {e}") +def pack_uint4(tensor: torch.Tensor) -> torch.Tensor: + if tensor.dtype != torch.uint8: + raise RuntimeError(f"Invalid tensor dtype {tensor.type}. torch.uint8 type is supported.") + packed_tensor = tensor.contiguous().reshape(-1, 2) + packed_tensor = torch.bitwise_and(packed_tensor[..., ::2], 15) | packed_tensor[..., 1::2] << 4 + return packed_tensor + + +def pack_int4(tensor: torch.Tensor) -> torch.Tensor: + if tensor.dtype != torch.int8: + raise RuntimeError(f"Invalid tensor dtype {tensor.type}. torch.int8 type is supported.") + tensor = tensor + 8 + return pack_uint4(tensor.to(dtype=torch.uint8)) + + +def unpack_uint4(packed_tensor: torch.Tensor, shape: torch.Size, transpose: Optional[bool] = False) -> torch.Tensor: + result = torch.stack((torch.bitwise_and(packed_tensor, 15), torch.bitwise_right_shift(packed_tensor, 4)), dim=-1).reshape(shape) + if transpose: + result = result.transpose(0,1) + return result + + +def unpack_int4(packed_tensor: torch.Tensor, shape: torch.Size, dtype: Optional[torch.dtype] = torch.int8, transpose: Optional[bool] = False) -> torch.Tensor: + result = unpack_uint4(packed_tensor, shape).to(dtype=dtype) - 8 + if transpose: + result = result.transpose(0,1) + return result + + +def quantize_int8_matmul_input(input: torch.FloatTensor, scale: torch.FloatTensor) -> Tuple[torch.ByteTensor, torch.FloatTensor]: + input_scale = torch.div(input.abs().max(), 127) + input = torch.div(input, input_scale).round_().clamp_(-128, 127).to(torch.int8).flatten(0,-2) + + scale_dtype = torch.float32 if input.dtype == torch.float16 else torch.bfloat16 + scale = torch.mul(input_scale.to(dtype=scale_dtype), scale.to(dtype=scale_dtype)) + return input, scale + + +def int8_matmul( + input: torch.Tensor, + weight: torch.Tensor, + scale: torch.Tensor, + compressed_weight_shape: torch.Size, + num_bits: int, +): + if num_bits == 4: + weight = unpack_int4_compiled(weight, compressed_weight_shape, transpose=True) + + return_dtype = input.dtype + output_shape = list(input.shape) + output_shape[-1] = weight.shape[-1] + + input, scale = quantize_int8_matmul_input_compiled(input, scale) + return decompress_symmetric_compiled(torch._int_mm(input, weight), scale, return_dtype, output_shape) + + +class linear_forward_int8_matmul(): + def __func__(self, input) -> torch.FloatTensor: + if self.pre_ops["0"].skip_int8_matmul: + return torch.nn.Linear.forward(self, input) + + num_bits = self.pre_ops["0"].num_bits + scale = self.pre_ops["0"].scale + compressed_weight_shape = self.pre_ops["0"].compressed_weight_shape if num_bits == 4 else None + result = int8_matmul(input, self.weight, scale, compressed_weight_shape, num_bits) + + if self.bias is not None: + result = result + self.bias + return result class INT8AsymmetricWeightsDecompressor(torch.nn.Module): @@ -449,29 +518,25 @@ class INT8AsymmetricWeightsDecompressor(torch.nn.Module): zero_point: torch.Tensor, result_dtype: torch.dtype, result_shape: torch.Size, + use_int8_matmul: bool, ): super().__init__() + self.num_bits = 8 + self.quantization_mode = "asymmetric" + self.scale = scale self.zero_point = zero_point self.result_dtype = result_dtype self.result_shape = result_shape - @property - def num_bits(self): - return 8 - - @property - def quantization_mode(self): - return "asymmetric" - def pack_weight(self, weight: torch.Tensor) -> torch.Tensor: if debug: if torch.any((weight < 0) | (weight > 255)): raise ValueError("Weight values are not in [0, 255].") return weight.to(dtype=torch.uint8) - def forward(self, x, *args, return_decompressed_only=False): - result = decompress_asymmetric(x.weight, self.scale, self.zero_point, self.result_dtype, self.result_shape) + def forward(self, x, input=None, *args, return_decompressed_only=False): + result = decompress_asymmetric_compiled(x.weight, self.scale, self.zero_point, self.result_dtype, self.result_shape) if return_decompressed_only: return result else: @@ -484,19 +549,19 @@ class INT8SymmetricWeightsDecompressor(torch.nn.Module): scale: torch.Tensor, result_dtype: torch.dtype, result_shape: torch.Size, + use_int8_matmul: bool, ): super().__init__() + self.num_bits = 8 + self.quantization_mode = "symmetric" + self.scale = scale self.result_dtype = result_dtype self.result_shape = result_shape - @property - def num_bits(self): - return 8 - - @property - def quantization_mode(self): - return "symmetric" + self.use_int8_matmul = use_int8_matmul + self.skip_int8_matmul = False + self.input_scale = None def pack_weight(self, weight: torch.Tensor) -> torch.Tensor: if debug: @@ -504,8 +569,17 @@ class INT8SymmetricWeightsDecompressor(torch.nn.Module): raise ValueError("Weight values are not in [-128, 127].") return weight.to(dtype=torch.int8) - def forward(self, x, *args, return_decompressed_only=False): - result = decompress_symmetric(x.weight, self.scale, self.result_dtype, self.result_shape) + def forward(self, x, input=None, *args, return_decompressed_only=False): + if self.use_int8_matmul: + if input is not None: + if torch.numel(input[0]) / input[0].shape[-1] < 32: + self.skip_int8_matmul = True + else: + self.skip_int8_matmul = False + return + result = decompress_symmetric_compiled(x.weight.transpose(0,1), self.scale.unsqueeze(-1), self.result_dtype, self.result_shape) + else: + result = decompress_symmetric_compiled(x.weight, self.scale, self.result_dtype, self.result_shape) if return_decompressed_only: return result else: @@ -520,30 +594,26 @@ class INT4AsymmetricWeightsDecompressor(torch.nn.Module): compressed_weight_shape: torch.Size, result_dtype: torch.dtype, result_shape: torch.Size, + use_int8_matmul: bool, ): super().__init__() + self.num_bits = 4 + self.quantization_mode = "asymmetric" + self.scale = scale self.zero_point = zero_point self.compressed_weight_shape = compressed_weight_shape self.result_dtype = result_dtype self.result_shape = result_shape - @property - def num_bits(self): - return 4 - - @property - def quantization_mode(self): - return "asymmetric" - def pack_weight(self, weight: torch.Tensor) -> torch.Tensor: if debug: if torch.any((weight < 0) | (weight > 15)): raise ValueError("Weight values are not in [0, 15].") return pack_uint4(weight.to(dtype=torch.uint8)) - def forward(self, x, *args, return_decompressed_only=False): - result = decompress_int4_asymmetric(x.weight, self.scale, self.zero_point, self.compressed_weight_shape, self.result_dtype, self.result_shape) + def forward(self, x, input=None, *args, return_decompressed_only=False): + result = decompress_int4_asymmetric_compiled(x.weight, self.scale, self.zero_point, self.compressed_weight_shape, self.result_dtype, self.result_shape) if return_decompressed_only: return result else: @@ -557,20 +627,20 @@ class INT4SymmetricWeightsDecompressor(torch.nn.Module): compressed_weight_shape: torch.Size, result_dtype: torch.dtype, result_shape: torch.Size, + use_int8_matmul: bool, ): super().__init__() + self.num_bits = 4 + self.quantization_mode = "symmetric" + self.scale = scale self.compressed_weight_shape = compressed_weight_shape self.result_dtype = result_dtype self.result_shape = result_shape - @property - def num_bits(self): - return 4 - - @property - def quantization_mode(self): - return "symmetric" + self.use_int8_matmul = use_int8_matmul + self.skip_int8_matmul = False + self.input_scale = None def pack_weight(self, weight: torch.Tensor) -> torch.Tensor: if debug: @@ -578,9 +648,47 @@ class INT4SymmetricWeightsDecompressor(torch.nn.Module): raise ValueError("Tensor values are not in [-8, 7].") return pack_int4(weight.to(dtype=torch.int8)) - def forward(self, x, *arg, return_decompressed_only=False): - result = decompress_int4_symmetric(x.weight, self.scale, self.compressed_weight_shape, self.result_dtype, self.result_shape) + def forward(self, x, input=None, *arg, return_decompressed_only=False): + if self.use_int8_matmul: + if input is not None: + if torch.numel(input[0]) / input[0].shape[-1] < 32: + self.skip_int8_matmul = True + else: + self.skip_int8_matmul = False + return + result = decompress_int4_symmetric_compiled(x.weight, self.scale.unsqueeze(-1), self.compressed_weight_shape, self.result_dtype, self.result_shape) + else: + result = decompress_int4_symmetric_compiled(x.weight, self.scale, self.compressed_weight_shape, self.result_dtype, self.result_shape) if return_decompressed_only: return result else: x.weight = result + + +if shared.opts.nncf_decompress_compile: + try: + torch._dynamo.config.cache_size_limit = max(8192, torch._dynamo.config.cache_size_limit) # pylint: disable=protected-access + decompress_asymmetric_compiled = torch.compile(decompress_asymmetric, fullgraph=True) + decompress_symmetric_compiled = torch.compile(decompress_symmetric, fullgraph=True) + decompress_int4_asymmetric_compiled = torch.compile(decompress_int4_asymmetric, fullgraph=True) + decompress_int4_symmetric_compiled = torch.compile(decompress_int4_symmetric, fullgraph=True) + + quantize_int8_matmul_input_compiled = torch.compile(quantize_int8_matmul_input, fullgraph=True) + unpack_int4_compiled = torch.compile(unpack_int4, fullgraph=True) + except Exception as e: + shared.log.warning(f"Quantization: type=nncf Decompress using torch.compile is not available: {e}") + decompress_asymmetric_compiled = decompress_asymmetric + decompress_symmetric_compiled = decompress_symmetric + decompress_int4_asymmetric_compiled = decompress_int4_asymmetric + decompress_int4_symmetric_compiled = decompress_int4_symmetric + + quantize_int8_matmul_input_compiled = quantize_int8_matmul_input + unpack_int4_compiled = unpack_int4 +else: + decompress_asymmetric_compiled = decompress_asymmetric + decompress_symmetric_compiled = decompress_symmetric + decompress_int4_asymmetric_compiled = decompress_int4_asymmetric + decompress_int4_symmetric_compiled = decompress_int4_symmetric + + quantize_int8_matmul_input_compiled = quantize_int8_matmul_input + unpack_int4_compiled = unpack_int4 diff --git a/modules/shared.py b/modules/shared.py index 926df1333..57330a021 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -551,6 +551,7 @@ options_templates.update(options_section(('quantization', "Quantization Settings "nncf_quantize_conv_layers": OptionInfo(False, "Quantize the convolutional layers", gr.Checkbox, {"visible": native and not cmd_opts.use_openvino}), "nncf_decompress_fp32": OptionInfo(False, "Decompress using full precision", gr.Checkbox, {"visible": native and not cmd_opts.use_openvino}), "nncf_decompress_compile": OptionInfo(devices.has_triton(), "Decompress using torch.compile", gr.Checkbox, {"visible": native and not cmd_opts.use_openvino}), + "nncf_decompress_int8_matmul": OptionInfo(False, "Use direct INT8 MatMul", gr.Checkbox, {"visible": native and not cmd_opts.use_openvino}), "nncf_quantize_shuffle_weights": OptionInfo(False, "Shuffle weights in post mode", gr.Checkbox, {"visible": native and not cmd_opts.use_openvino}), "layerwise_quantization_sep": OptionInfo("

Layerwise Casting

", "", gr.HTML), diff --git a/wiki b/wiki index 223947ac7..12dbff5ca 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 223947ac73fa43d7b8e8440788bf65a42a01311f +Subproject commit 12dbff5ca440c62027a4a12685d5f4b73ea6532c