mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
NNCF experimental direct INT8 MatMul support
This commit is contained in:
+217
-109
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user