SDNQ INT8 matmul support for Conv2d

This commit is contained in:
Disty0
2025-06-05 13:36:49 +03:00
parent c7fb5b1690
commit e25890bb1d
2 changed files with 99 additions and 7 deletions
+97 -6
View File
@@ -57,21 +57,30 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz
if layer_class_name in conv_types:
if not quant_conv:
return layer
if dtype_dict[weights_dtype]["num_bits"] < 4:
weights_dtype = "uint4"
is_conv_type = True
reduction_axes = 1
output_channel_size, channel_size = layer.weight.shape[:2]
group_channel_size = channel_size // layer.groups
use_quantized_matmul = False
if dtype_dict[weights_dtype]["num_bits"] < 4:
weights_dtype = "uint4"
if shared.opts.sdnq_use_quantized_matmul_conv:
use_quantized_matmul = dtype_dict[weights_dtype]["is_integer"] and weights_dtype in quantized_matmul_dtypes and group_channel_size >= 32 and output_channel_size >= 32
#if use_quantized_matmul and not dtype_dict[weights_dtype]["is_integer"]:
# use_quantized_matmul = output_channel_size % 16 == 0 and group_channel_size % 16 == 0
# use_tensorwise_fp8_matmul = torch_version < 2.5 or devices.backend in {"cpu", "openvino"} or (devices.backend == "cuda" and sys.platform == "win32" and torch_version <= 2.7 and torch.cuda.get_device_capability(devices.device) == (8,9))
if use_quantized_matmul:
result_shape = layer.weight.shape
layer.weight.data = layer.weight.reshape(output_channel_size, -1)
elif layer_class_name in conv_transpose_types:
if not quant_conv:
return layer
if dtype_dict[weights_dtype]["num_bits"] < 4:
weights_dtype = "uint4"
is_conv_transpose_type = True
reduction_axes = 0
channel_size, output_channel_size = layer.weight.shape[:2]
use_quantized_matmul = False
if dtype_dict[weights_dtype]["num_bits"] < 4:
weights_dtype = "uint4"
else:
is_linear_type = True
reduction_axes = -1
@@ -193,7 +202,10 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz
else:
layer.forward = quantized_linear_forward
elif is_conv_type:
layer.forward = quantized_conv_forward
if use_quantized_matmul:
layer.forward = quantized_conv2d_forward_int8_matmul
else:
layer.forward = quantized_conv_forward
elif is_conv_transpose_type:
if layer_class_name.endswith("1d"):
layer.forward = quantized_conv_transpose_1d_forward
@@ -393,7 +405,7 @@ def quantize_fp8_matmul_input_tensorwise(input: torch.FloatTensor, scale: torch.
return input, scale
def quantize_int8_matmul_input(input: torch.FloatTensor, scale: torch.FloatTensor) -> Tuple[torch.ByteTensor, torch.FloatTensor]:
def quantize_int8_matmul_input(input: torch.FloatTensor, scale: torch.FloatTensor, flatten: bool = True) -> Tuple[torch.ByteTensor, torch.FloatTensor]:
input = input.flatten(0,-2).contiguous()
input_scale = torch.div(input.abs().amax(dim=-1, keepdims=True), 127)
input = torch.div(input, input_scale).round_().clamp_(-128, 127).to(torch.int8)
@@ -454,6 +466,54 @@ def int8_matmul(
return result
def conv2d_int8_matmul(
input: torch.FloatTensor,
weight: torch.ByteTensor,
bias: torch.FloatTensor,
scale: torch.FloatTensor,
result_shape: torch.Size,
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,
) -> torch.FloatTensor:
return_dtype = input.dtype
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
mm_output_shape = (batch_size, H_out, W_out, C_out)
if compressed_weight_shape is not None:
weight = unpack_int_symetric(weight, compressed_weight_shape, weights_dtype, dtype=torch.int8, transpose=True)
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_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:
weight = weight.reshape(weight.shape[0], groups, weight.shape[1] // groups).transpose(0,1)
input = input.reshape(input.shape[0], groups, input.shape[1] // groups).transpose(0,1)
result = []
for i in range(groups):
result.append(torch._int_mm(input[i], weight[i]))
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)
def quantized_linear_forward_fp8_matmul(self, input: torch.FloatTensor) -> torch.FloatTensor:
if torch.numel(input) / input.shape[-1] < 32:
return torch.nn.functional.linear(input, self.sdnq_decompressor(self.weight, skip_quantized_matmul=True), self.bias)
@@ -472,6 +532,36 @@ def quantized_linear_forward_int8_matmul(self, input: torch.FloatTensor) -> torc
return int8_matmul(input, self.weight, self.bias, self.sdnq_decompressor.scale, getattr(self.sdnq_decompressor, "compressed_weight_shape", None), self.sdnq_decompressor.weights_dtype)
def quantized_conv2d_forward_int8_matmul(self, input) -> torch.FloatTensor:
if isinstance(self.stride, int):
stride_h = stride_w = self.stride
else:
stride_h, stride_w = self.stride
if isinstance(self.padding, int):
padding_h = padding_w = self.padding
else:
padding_h, padding_w = self.padding
if isinstance(self.dilation, int):
dilation_h = dilation_w = self.dilation
else:
dilation_h, dilation_w = self.dilation
return conv2d_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,
)
def quantized_linear_forward(self, input: torch.FloatTensor) -> torch.FloatTensor:
return torch.nn.functional.linear(input, self.sdnq_decompressor(self.weight), self.bias)
@@ -840,6 +930,7 @@ if shared.opts.sdnq_decompress_compile:
fp8_matmul = torch.compile(fp8_matmul, fullgraph=True)
fp8_matmul_tensorwise = torch.compile(fp8_matmul_tensorwise, fullgraph=True)
int8_matmul = torch.compile(int8_matmul, fullgraph=True)
conv2d_int8_matmul = torch.compile(conv2d_int8_matmul, 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
+2 -1
View File
@@ -521,10 +521,11 @@ options_templates.update(options_section(("quantization", "Quantization Settings
"sdnq_quantize_weights_mode": OptionInfo("int8", "Quantization type", gr.Dropdown, {"choices": ["int8", "int6", "uint4", "float8_e4m3fn", "uint8", "uint6", "int4", "float8_e5m2", "float8_e4m3fnuz", "float8_e5m2fnuz", "int2", "uint2", "uint1"], "visible": native}),
"sdnq_quantize_weights_group_size": OptionInfo(0, "Group size", gr.Slider, {"minimum": -1, "maximum": 4096, "step": 1, "visible": native}),
"sdnq_quantize_conv_layers": OptionInfo(False, "Quantize the convolutional layers", gr.Checkbox, {"visible": native}),
"sdnq_decompress_fp32": OptionInfo(False, "Decompress using full precision", gr.Checkbox, {"visible": native}),
"sdnq_decompress_compile": OptionInfo(devices.has_triton(), "Decompress using torch.compile", gr.Checkbox, {"visible": native}),
"sdnq_use_quantized_matmul": OptionInfo(False, "Use Quantized MatMul", gr.Checkbox, {"visible": native}),
"sdnq_use_quantized_matmul_conv": OptionInfo(False, "Use Quantized MatMul with convolutional layers", gr.Checkbox, {"visible": native}),
"sdnq_quantize_with_gpu": OptionInfo(True, "Quantize with the GPU", gr.Checkbox, {"visible": native}),
"sdnq_decompress_fp32": OptionInfo(False, "Decompress using full precision", gr.Checkbox, {"visible": native}),
"sdnq_quantize_shuffle_weights": OptionInfo(False, "Shuffle weights in post mode", gr.Checkbox, {"visible": native}),
"bnb_quantization_sep": OptionInfo("<h2>BitsAndBytes</h2>", "", gr.HTML),