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 @@
-**Image Diffusion implementation with advanced features**
+# SD.Next: All-in-one WebUI for AI generative image and video creation


@@ -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("