diff --git a/modules/model_quant_nncf.py b/modules/model_quant_nncf.py index 3d18caf23..5a4ad53aa 100644 --- a/modules/model_quant_nncf.py +++ b/modules/model_quant_nncf.py @@ -50,6 +50,7 @@ def nncf_compress_layer(layer, num_bits, is_asym_mode, torch_dtype=None, quant_c if layer.__class__.__name__ in allowed_types: if torch_dtype is None: torch_dtype = devices.dtype + result_shape = None if layer.__class__.__name__ in conv_types: if is_asym_mode or not quant_conv: # don't quant convs with asym mode @@ -60,6 +61,22 @@ def nncf_compress_layer(layer, num_bits, is_asym_mode, torch_dtype=None, quant_c return layer reduction_axes = [i for i in range(layer.weight.ndim) if i != 1] else: + if shared.opts.nncf_compress_weights_num_of_groups > 1: + num_of_groups = shared.opts.nncf_compress_weights_num_of_groups + reduction_axes = layer.weight.ndim - 1 + channel_size = layer.weight.shape[reduction_axes] + + group_size = channel_size // num_of_groups + while channel_size % group_size != 0: # find something divisible + num_of_groups -= 1 + group_size = channel_size // num_of_groups + + if num_of_groups > 1: + result_shape = layer.weight.shape + new_shape = list(result_shape) + new_shape[reduction_axes : reduction_axes + 1] = (num_of_groups, group_size) + layer.weight.data = layer.weight.reshape(new_shape) + reduction_axes = [layer.weight.ndim - 1] if shared.opts.diffusers_offload_mode != "none": @@ -116,12 +133,14 @@ def nncf_compress_layer(layer, num_bits, is_asym_mode, torch_dtype=None, quant_c zero_point=zero_point.data, compressed_weight_shape=compressed_weight.shape, result_dtype=torch_dtype, + result_shape=result_shape, ) else: decompressor = INT4SymmetricWeightsDecompressor( scale=scale.data, compressed_weight_shape=compressed_weight.shape, result_dtype=torch_dtype, + result_shape=result_shape, ) else: if is_asym_mode: @@ -129,11 +148,13 @@ def nncf_compress_layer(layer, num_bits, is_asym_mode, torch_dtype=None, quant_c scale=scale.data, zero_point=zero_point.data, result_dtype=torch_dtype, + result_shape=result_shape, ) else: decompressor = INT8SymmetricWeightsDecompressor( scale=scale.data, result_dtype=torch_dtype, + result_shape=result_shape, ) compressed_weight = decompressor.pack_weight(compressed_weight) @@ -375,20 +396,26 @@ def pack_int4(tensor: torch.Tensor) -> torch.Tensor: return pack_uint4(tensor.to(dtype=torch.uint8)) -def decompress_asymmetric(input: torch.Tensor, scale: torch.Tensor, zero_point: torch.Tensor, dtype: torch.dtype) -> torch.Tensor: - return torch.mul(torch.sub(input.to(dtype=scale.dtype), zero_point), scale).to(dtype=dtype) +def decompress_asymmetric(input: torch.Tensor, scale: torch.Tensor, zero_point: torch.Tensor, dtype: torch.dtype, result_shape: torch.Size) -> torch.Tensor: + result = torch.mul(torch.sub(input.to(dtype=scale.dtype), zero_point), scale).to(dtype=dtype) + if result_shape is not None: + result = result.reshape(result_shape) + return result -def decompress_symmetric(input: torch.Tensor, scale: torch.Tensor, dtype: torch.dtype) -> torch.Tensor: - return torch.mul(input.to(dtype=scale.dtype), scale).to(dtype=dtype) +def decompress_symmetric(input: torch.Tensor, scale: torch.Tensor, dtype: torch.dtype, result_shape: torch.Size) -> torch.Tensor: + result = torch.mul(input.to(dtype=scale.dtype), scale).to(dtype=dtype) + if result_shape is not None: + result = result.reshape(result_shape) + return result -def decompress_int4_asymmetric(input: torch.Tensor, scale: torch.Tensor, zero_point: torch.Tensor, shape: torch.Size, dtype: torch.dtype) -> torch.Tensor: - return decompress_asymmetric(unpack_uint4(input, shape), scale, zero_point, dtype) +def decompress_int4_asymmetric(input: torch.Tensor, scale: torch.Tensor, zero_point: torch.Tensor, shape: torch.Size, dtype: torch.dtype, result_shape: torch.Size) -> torch.Tensor: + return decompress_asymmetric(unpack_uint4(input, shape), scale, zero_point, dtype, result_shape) -def decompress_int4_symmetric(input: torch.Tensor, scale: torch.Tensor, shape: torch.Size, dtype: torch.dtype) -> torch.Tensor: - return decompress_symmetric(unpack_int4(input, shape, dtype=scale.dtype), scale, dtype) +def decompress_int4_symmetric(input: torch.Tensor, scale: torch.Tensor, shape: torch.Size, dtype: torch.dtype, result_shape: torch.Size) -> torch.Tensor: + return decompress_symmetric(unpack_int4(input, shape, dtype=scale.dtype), scale, dtype, result_shape) if shared.opts.nncf_decompress_compile: @@ -399,15 +426,22 @@ if shared.opts.nncf_decompress_compile: 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.logs.warning(f"Quantization: type=nncf Decompress using torch.compile is not available: {e}") + shared.log.warning(f"Quantization: type=nncf Decompress using torch.compile is not available: {e}") class INT8AsymmetricWeightsDecompressor(torch.nn.Module): - def __init__(self, scale: torch.Tensor, zero_point: torch.Tensor, result_dtype: torch.dtype): + def __init__( + self, + scale: torch.Tensor, + zero_point: torch.Tensor, + result_dtype: torch.dtype, + result_shape: torch.Size, + ): super().__init__() self.scale = scale self.zero_point = zero_point self.result_dtype = result_dtype + self.result_shape = result_shape @property def num_bits(self): @@ -424,7 +458,7 @@ class INT8AsymmetricWeightsDecompressor(torch.nn.Module): 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) + result = decompress_asymmetric(x.weight, self.scale, self.zero_point, self.result_dtype, self.result_shape) if return_decompressed_only: return result else: @@ -432,10 +466,16 @@ class INT8AsymmetricWeightsDecompressor(torch.nn.Module): class INT8SymmetricWeightsDecompressor(torch.nn.Module): - def __init__(self, scale: torch.Tensor, result_dtype: torch.dtype): + def __init__( + self, + scale: torch.Tensor, + result_dtype: torch.dtype, + result_shape: torch.Size, + ): super().__init__() self.scale = scale self.result_dtype = result_dtype + self.result_shape = result_shape @property def num_bits(self): @@ -452,7 +492,7 @@ class INT8SymmetricWeightsDecompressor(torch.nn.Module): 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) + result = decompress_symmetric(x.weight, self.scale, self.result_dtype, self.result_shape) if return_decompressed_only: return result else: @@ -466,12 +506,14 @@ class INT4AsymmetricWeightsDecompressor(torch.nn.Module): zero_point: torch.Tensor, compressed_weight_shape: torch.Size, result_dtype: torch.dtype, + result_shape: torch.Size, ): super().__init__() 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): @@ -488,7 +530,7 @@ class INT4AsymmetricWeightsDecompressor(torch.nn.Module): 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) + result = decompress_int4_asymmetric(x.weight, self.scale, self.zero_point, self.compressed_weight_shape, self.result_dtype, self.result_shape) if return_decompressed_only: return result else: @@ -501,11 +543,13 @@ class INT4SymmetricWeightsDecompressor(torch.nn.Module): scale: torch.Tensor, compressed_weight_shape: torch.Size, result_dtype: torch.dtype, + result_shape: torch.Size, ): super().__init__() 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): @@ -522,7 +566,7 @@ class INT4SymmetricWeightsDecompressor(torch.nn.Module): 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) + result = decompress_int4_symmetric(x.weight, self.scale, self.compressed_weight_shape, self.result_dtype, self.result_shape) if return_decompressed_only: return result else: diff --git a/modules/shared.py b/modules/shared.py index 1b90c6573..cbe5f0830 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -542,16 +542,17 @@ options_templates.update(options_section(('quantization', "Quantization Settings "nncf_compress_sep": OptionInfo("