SDNQ add quantized matmul support for Conv1d and Conv3d too

This commit is contained in:
Disty0
2025-06-06 00:19:54 +03:00
parent 976f0ba61f
commit 06fcc3cf85
2 changed files with 128 additions and 98 deletions
+127 -97
View File
@@ -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
+1 -1
Submodule wiki updated: d6fce6bde6...0663976403