From a1491a660c6edcaf0ad71240bf34868cc5793d32 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Fri, 9 May 2025 23:36:50 +0300 Subject: [PATCH] Cleanup --- modules/model_quant_nncf.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/modules/model_quant_nncf.py b/modules/model_quant_nncf.py index 2fc41bcf4..3d18caf23 100644 --- a/modules/model_quant_nncf.py +++ b/modules/model_quant_nncf.py @@ -387,8 +387,8 @@ def decompress_int4_asymmetric(input: torch.Tensor, scale: torch.Tensor, zero_po return decompress_asymmetric(unpack_uint4(input, shape), scale, zero_point, dtype) -def decompress_int4_symmetric(input: torch.Tensor, scale: torch.Tensor, shape: torch.Size, dtype: torch.dtype, result_dtype: torch.dtype) -> torch.Tensor: - return decompress_symmetric(unpack_int4(input, shape, dtype=dtype), scale, result_dtype) +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) if shared.opts.nncf_decompress_compile: @@ -522,7 +522,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.scale.dtype, self.result_dtype) + result = decompress_int4_symmetric(x.weight, self.scale, self.compressed_weight_shape, self.result_dtype) if return_decompressed_only: return result else: