From 06fcc3cf85ac3e7992b9574a4f67767314927712 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Fri, 6 Jun 2025 00:19:54 +0300 Subject: [PATCH] SDNQ add quantized matmul support for Conv1d and Conv3d too --- modules/model_quant_sdnq.py | 224 ++++++++++++++++++++---------------- wiki | 2 +- 2 files changed, 128 insertions(+), 98 deletions(-) diff --git a/modules/model_quant_sdnq.py b/modules/model_quant_sdnq.py index 2f741ca91..b6d61c9bf 100644 --- a/modules/model_quant_sdnq.py +++ b/modules/model_quant_sdnq.py @@ -202,12 +202,12 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz elif is_conv_type: if use_quantized_matmul: if dtype_dict[weights_dtype]["is_integer"]: - layer.forward = quantized_conv2d_forward_int8_matmul + layer.forward = quantized_conv_forward_int8_matmul else: if use_tensorwise_fp8_matmul: - layer.forward = quantized_conv2d_forward_fp8_matmul_tensorwise + layer.forward = quantized_conv_forward_fp8_matmul_tensorwise else: - layer.forward = quantized_conv2d_forward_fp8_matmul + layer.forward = quantized_conv_forward_fp8_matmul else: layer.forward = quantized_conv_forward elif is_conv_transpose_type: @@ -472,7 +472,56 @@ def int8_matmul( return result -def conv2d_fp8_matmul( +def process_conv_input(conv_type, input, reversed_padding_repeated_twice, padding_mode, result_shape, stride, padding, dilation): + if conv_type == 1: + batch_size, _, L_in = input.shape + C_out, _, K_l = result_shape + L_out = (L_in + 2 * padding[1] - dilation[1] * (K_l - 1) - 1) // stride[1] + 1 + mm_output_shape = (batch_size, L_out, C_out) + kernel_size = (1, K_l) + if conv_type == 2: + batch_size, _, H_in, W_in = input.shape + C_out, _, K_h, K_w = result_shape + H_out = (H_in + 2 * padding[0] - dilation[0] * (K_h - 1) - 1) // stride[0] + 1 + W_out = (W_in + 2 * padding[1] - dilation[1] * (K_w - 1) - 1) // stride[1] + 1 + mm_output_shape = (batch_size, H_out, W_out, C_out) + kernel_size = (K_h, K_w) + elif conv_type == 3: + batch_size, _, D_in, H_in, W_in = input.shape + C_out, _, K_d, K_h, K_w = result_shape + D_out = (D_in + 2 * padding[0] - dilation[0] * (K_d - 1) - 1) // stride[0] + 1 + H_out = (H_in + 2 * padding[1] - dilation[1] * (K_h - 1) - 1) // stride[1] + 1 + W_out = (W_in + 2 * padding[2] - dilation[2] * (K_w - 1) - 1) // stride[2] + 1 + mm_output_shape = (batch_size, D_out, H_out, W_out, C_out) + kernel_size = (K_d, K_h, K_w) + + if padding_mode != "zeros": + input = torch.nn.functional.pad(input, reversed_padding_repeated_twice, mode=padding_mode) + padding = (0,) * (conv_type if conv_type != 1 else 2) + elif conv_type == 3: + input = torch.nn.functional.pad(input, reversed_padding_repeated_twice) + + if conv_type == 1: + input = input.unsqueeze(2) + + if conv_type == 3: + K_D_eff = kernel_size[0] + (kernel_size[0] - 1) * (dilation[0] - 1) + K_H_eff = kernel_size[1] + (kernel_size[1] - 1) * (dilation[0] - 1) + K_W_eff = kernel_size[2] + (kernel_size[2] - 1) * (dilation[0] - 1) + input = input.unfold(2, K_D_eff, stride[0]).unfold(3, K_H_eff, stride[1]).unfold(4, K_W_eff, stride[2]) + if dilation[0] > 1: + input = input[..., ::dilation[0], :, :] + if dilation[1] > 1: + input = input[..., ::dilation[1], :] + if dilation[2] > 1: + input = input[..., ::dilation[2]] + input = input.permute(0, 2, 3, 4, 1, 5, 6, 7).reshape(mm_output_shape[0], mm_output_shape[1] * mm_output_shape[2] * mm_output_shape[3], -1) + else: + input = torch.nn.functional.unfold(input, kernel_size=kernel_size, padding=padding, stride=stride, dilation=dilation).transpose(1,2) + return input, mm_output_shape + + +def conv_fp8_matmul( input: torch.FloatTensor, weight: torch.ByteTensor, bias: torch.FloatTensor, @@ -480,25 +529,16 @@ def conv2d_fp8_matmul( result_shape: torch.Size, weights_dtype: str, reversed_padding_repeated_twice: List[int], - padding_mode: str, groups: int, - stride_h: int, stride_w: int, - padding_h: int, padding_w: int, - dilation_h: int, dilation_w: int, + padding_mode: str, conv_type: int, + groups: int, stride: List[int], + padding: List[int], dilation: List[int], ) -> torch.FloatTensor: return_dtype = input.dtype - mm_output_shape, K_h, K_w = get_conv2d_shapes(input.shape, result_shape, stride_h, stride_w, padding_h, padding_w, dilation_h, dilation_w) - if padding_mode != "zeros": - input = torch.nn.functional.pad(input, reversed_padding_repeated_twice, mode=padding_mode) - padding_h = padding_w = 0 - - input, input_scale = quantize_fp8_matmul_input( - torch.nn.functional.unfold( - input, kernel_size=(K_h, K_w), padding=(padding_h, padding_w), stride=(stride_h, stride_w), dilation=(dilation_h, dilation_w) - ).transpose(1,2), - ) + input, mm_output_shape = process_conv_input(conv_type, input, reversed_padding_repeated_twice, padding_mode, result_shape, stride, padding, dilation) + input, input_scale = quantize_fp8_matmul_input(input) if groups == 1: - result = torch._scaled_mm(input, weight, scale_a=input_scale, scale_b=scale, bias=bias, out_dtype=return_dtype).reshape(mm_output_shape).permute(0,3,1,2) + result = torch._scaled_mm(input, weight, scale_a=input_scale, scale_b=scale, bias=bias, out_dtype=return_dtype).reshape(mm_output_shape) else: scale = scale.reshape(groups, 1, scale.shape[1] // groups) input_scale = input_scale.reshape(groups, input_scale.shape[0] // groups, 1) @@ -510,11 +550,17 @@ def conv2d_fp8_matmul( result = torch.cat(result, dim=-1).reshape(mm_output_shape) if bias is not None: result.add_(bias) + + if conv_type == 1: + result = result.transpose(1,2) + elif conv_type == 2: result = result.permute(0,3,1,2) + elif conv_type == 3: + result = result.permute(0,4,1,2,3) return result -def conv2d_fp8_matmul_tensorwise( +def conv_fp8_matmul_tensorwise( input: torch.FloatTensor, weight: torch.ByteTensor, bias: torch.FloatTensor, @@ -522,25 +568,15 @@ def conv2d_fp8_matmul_tensorwise( result_shape: torch.Size, weights_dtype: str, reversed_padding_repeated_twice: List[int], - padding_mode: str, groups: int, - stride_h: int, stride_w: int, - padding_h: int, padding_w: int, - dilation_h: int, dilation_w: int, + padding_mode: str, conv_type: int, + groups: int, stride: List[int], + padding: List[int], dilation: List[int], ) -> torch.FloatTensor: return_dtype = input.dtype - mm_output_shape, K_h, K_w = get_conv2d_shapes(input.shape, result_shape, stride_h, stride_w, padding_h, padding_w, dilation_h, dilation_w) - if padding_mode != "zeros": - input = torch.nn.functional.pad(input, reversed_padding_repeated_twice, mode=padding_mode) - padding_h = padding_w = 0 - - input, scale = quantize_fp8_matmul_input_tensorwise( - torch.nn.functional.unfold( - input, kernel_size=(K_h, K_w), padding=(padding_h, padding_w), stride=(stride_h, stride_w), dilation=(dilation_h, dilation_w) - ).transpose(1,2), - scale, - ) - + input, mm_output_shape = process_conv_input(conv_type, input, reversed_padding_repeated_twice, padding_mode, result_shape, stride, padding, dilation) + input, scale = quantize_fp8_matmul_input_tensorwise(input, scale) dummy_input_scale = torch.ones(1, device=input.device, dtype=torch.float32) + if groups == 1: result = decompress_symmetric(torch._scaled_mm(input, weight, scale_a=dummy_input_scale, scale_b=dummy_input_scale, bias=None, out_dtype=scale.dtype), scale, return_dtype, mm_output_shape) else: @@ -552,10 +588,17 @@ def conv2d_fp8_matmul_tensorwise( result = decompress_symmetric(torch.cat(result, dim=-1), scale, return_dtype, mm_output_shape) if bias is not None: result.add_(bias) - return result.permute(0,3,1,2) + + if conv_type == 1: + result = result.transpose(1,2) + elif conv_type == 2: + result = result.permute(0,3,1,2) + elif conv_type == 3: + result = result.permute(0,4,1,2,3) + return result -def conv2d_int8_matmul( +def conv_int8_matmul( input: torch.FloatTensor, weight: torch.ByteTensor, bias: torch.FloatTensor, @@ -564,26 +607,16 @@ def conv2d_int8_matmul( compressed_weight_shape: torch.Size, weights_dtype: str, reversed_padding_repeated_twice: List[int], - padding_mode: str, groups: int, - stride_h: int, stride_w: int, - padding_h: int, padding_w: int, - dilation_h: int, dilation_w: int, + padding_mode: str, conv_type: int, + groups: int, stride: List[int], + padding: List[int], dilation: List[int], ) -> torch.FloatTensor: return_dtype = input.dtype - mm_output_shape, K_h, K_w = get_conv2d_shapes(input.shape, result_shape, stride_h, stride_w, padding_h, padding_w, dilation_h, dilation_w) - if padding_mode != "zeros": - input = torch.nn.functional.pad(input, reversed_padding_repeated_twice, mode=padding_mode) - padding_h = padding_w = 0 + input, mm_output_shape = process_conv_input(conv_type, input, reversed_padding_repeated_twice, padding_mode, result_shape, stride, padding, dilation) + input, scale = quantize_int8_matmul_input(input, scale) if compressed_weight_shape is not None: weight = unpack_int_symetric(weight, compressed_weight_shape, weights_dtype, dtype=torch.int8, transpose=True) - input, scale = quantize_int8_matmul_input( - torch.nn.functional.unfold( - input, kernel_size=(K_h, K_w), padding=(padding_h, padding_w), stride=(stride_h, stride_w), dilation=(dilation_h, dilation_w) - ).transpose(1,2), - scale, - ) - if groups == 1: result = decompress_symmetric(torch._int_mm(input, weight), scale, return_dtype, mm_output_shape) else: @@ -595,7 +628,14 @@ def conv2d_int8_matmul( result = decompress_symmetric(torch.cat(result, dim=-1), scale, return_dtype, mm_output_shape) if bias is not None: result.add_(bias) - return result.permute(0,3,1,2) + + if conv_type == 1: + result = result.transpose(1,2) + elif conv_type == 2: + result = result.permute(0,3,1,2) + elif conv_type == 3: + result = result.permute(0,4,1,2,3) + return result def quantized_linear_forward_fp8_matmul(self, input: torch.FloatTensor) -> torch.FloatTensor: @@ -620,73 +660,63 @@ def quantized_linear_forward(self, input: torch.FloatTensor) -> torch.FloatTenso return torch.nn.functional.linear(input, self.sdnq_decompressor(self.weight), self.bias) -def get_conv2d_args(stride, padding, dilation): +def get_conv_args(input_ndim, stride, padding, dilation): + if input_ndim == 3: + conv_type = 1 + elif input_ndim == 4: + conv_type = 2 + elif input_ndim == 5: + conv_type = 3 if isinstance(stride, int): - stride_h = stride_w = stride - else: - stride_h, stride_w = stride + stride = (stride,) * conv_type if isinstance(padding, int): - padding_h = padding_w = padding - else: - padding_h, padding_w = padding + padding = (padding,) * conv_type if isinstance(dilation, int): - dilation_h = dilation_w = dilation - else: - dilation_h, dilation_w = dilation - return stride_h, stride_w, padding_h, padding_w, dilation_h, dilation_w + dilation = (dilation,) * conv_type + if conv_type == 1: + stride = (1, stride[0]) + padding = (0, padding[0]) + dilation = (1, dilation[0]) + return conv_type, stride, padding, dilation -def get_conv2d_shapes(input_shape, result_shape, stride_h, stride_w, padding_h, padding_w, dilation_h, dilation_w): - batch_size, _, H_in, W_in = input_shape - C_out, _, K_h, K_w = result_shape - W_out = (W_in + 2 * padding_w - dilation_w * (K_w - 1) - 1) // stride_w + 1 - H_out = (H_in + 2 * padding_h - dilation_h * (K_h - 1) - 1) // stride_h + 1 - return (batch_size, H_out, W_out, C_out), K_h, K_w - - -def quantized_conv2d_forward_fp8_matmul(self, input) -> torch.FloatTensor: - stride_h, stride_w, padding_h, padding_w, dilation_h, dilation_w = get_conv2d_args(self.stride, self.padding, self.dilation) - return conv2d_fp8_matmul( +def quantized_conv_forward_fp8_matmul(self, input) -> torch.FloatTensor: + conv_type, stride, padding, dilation = get_conv_args(input.ndim, self.stride, self.padding, self.dilation) + return conv_fp8_matmul( input, self.weight, self.bias, self.sdnq_decompressor.scale, self.sdnq_decompressor.result_shape, self.sdnq_decompressor.weights_dtype, self._reversed_padding_repeated_twice, - self.padding_mode, self.groups, - stride_h, stride_w, - padding_h, padding_w, - dilation_h, dilation_w, + self.padding_mode, conv_type, + self.groups, stride, padding, dilation, ) -def quantized_conv2d_forward_fp8_matmul_tensorwise(self, input) -> torch.FloatTensor: - stride_h, stride_w, padding_h, padding_w, dilation_h, dilation_w = get_conv2d_args(self.stride, self.padding, self.dilation) - return conv2d_fp8_matmul_tensorwise( +def quantized_conv_forward_fp8_matmul_tensorwise(self, input) -> torch.FloatTensor: + conv_type, stride, padding, dilation = get_conv_args(input.ndim, self.stride, self.padding, self.dilation) + return conv_fp8_matmul_tensorwise( input, self.weight, self.bias, self.sdnq_decompressor.scale, self.sdnq_decompressor.result_shape, self.sdnq_decompressor.weights_dtype, self._reversed_padding_repeated_twice, - self.padding_mode, self.groups, - stride_h, stride_w, - padding_h, padding_w, - dilation_h, dilation_w, + self.padding_mode, conv_type, + self.groups, stride, padding, dilation, ) -def quantized_conv2d_forward_int8_matmul(self, input) -> torch.FloatTensor: - stride_h, stride_w, padding_h, padding_w, dilation_h, dilation_w = get_conv2d_args(self.stride, self.padding, self.dilation) - return conv2d_int8_matmul( +def quantized_conv_forward_int8_matmul(self, input) -> torch.FloatTensor: + conv_type, stride, padding, dilation = get_conv_args(input.ndim, self.stride, self.padding, self.dilation) + return conv_int8_matmul( input, self.weight, self.bias, self.sdnq_decompressor.scale, self.sdnq_decompressor.result_shape, getattr(self.sdnq_decompressor, "compressed_weight_shape", None), self.sdnq_decompressor.weights_dtype, self._reversed_padding_repeated_twice, - self.padding_mode, self.groups, - stride_h, stride_w, - padding_h, padding_w, - dilation_h, dilation_w, + self.padding_mode, conv_type, + self.groups, stride, padding, dilation, ) @@ -1057,9 +1087,9 @@ if shared.opts.sdnq_decompress_compile: int8_matmul = torch.compile(int8_matmul, fullgraph=True) fp8_matmul = torch.compile(fp8_matmul, fullgraph=True) fp8_matmul_tensorwise = torch.compile(fp8_matmul_tensorwise, fullgraph=True) - conv2d_int8_matmul = torch.compile(conv2d_int8_matmul, fullgraph=True) - conv2d_fp8_matmul = torch.compile(conv2d_fp8_matmul, fullgraph=True) - conv2d_fp8_matmul_tensorwise = torch.compile(conv2d_fp8_matmul_tensorwise, fullgraph=True) + conv_int8_matmul = torch.compile(conv_int8_matmul, fullgraph=True) + conv_fp8_matmul = torch.compile(conv_fp8_matmul, fullgraph=True) + conv_fp8_matmul_tensorwise = torch.compile(conv_fp8_matmul_tensorwise, fullgraph=True) except Exception as e: shared.log.warning(f"Quantization: type=sdnq Decompress using torch.compile is not available: {e}") decompress_asymmetric_compiled = decompress_asymmetric diff --git a/wiki b/wiki index d6fce6bde..066397640 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit d6fce6bde637f71a5703dcd6dff348d32fd2992a +Subproject commit 0663976403988e9c8c5ff5a6e4c9da6f56e9b65a