From a12edc1e907768125d7b509823a8027563f5fc7e Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 15 Sep 2025 20:22:35 +0300 Subject: [PATCH] SDNQ use nan_to_num_ with fp8 quantization in case of zeros --- modules/sdnq/__init__.py | 2 ++ modules/sdnq/dequantizer.py | 6 ++++++ modules/sdnq/layers/linear/linear_fp8.py | 4 ++-- modules/sdnq/layers/linear/linear_fp8_tensorwise.py | 5 ++--- modules/sdnq/layers/linear/linear_int8.py | 1 + 5 files changed, 13 insertions(+), 5 deletions(-) diff --git a/modules/sdnq/__init__.py b/modules/sdnq/__init__.py index 8b95fab0c..a7829d44d 100644 --- a/modules/sdnq/__init__.py +++ b/modules/sdnq/__init__.py @@ -42,6 +42,8 @@ def quantize_weight(weight: torch.FloatTensor, reduction_axes: Union[int, List[i zero_point = None if dtype_dict[weights_dtype]["is_integer"]: quantized_weight.round_() + else: + quantized_weight.nan_to_num_() quantized_weight = quantized_weight.clamp_(dtype_dict[weights_dtype]["min"], dtype_dict[weights_dtype]["max"]).to(dtype_dict[weights_dtype]["torch_dtype"]) return quantized_weight, scale, zero_point diff --git a/modules/sdnq/dequantizer.py b/modules/sdnq/dequantizer.py index 9000f4475..9ac902c57 100644 --- a/modules/sdnq/dequantizer.py +++ b/modules/sdnq/dequantizer.py @@ -42,6 +42,12 @@ def quantize_int8(input: torch.FloatTensor, dim: int = -1) -> Tuple[torch.CharTe return input, scale +def quantize_fp8(input: torch.FloatTensor, dim: int = -1) -> Tuple[torch.Tensor, torch.FloatTensor]: + scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(448) + input = torch.div(input, scale).nan_to_num_().clamp_(-448, 448).to(dtype=torch.float8_e4m3fn) + return input, scale + + def re_quantize_matmul_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, result_shape: torch.Size) -> Tuple[torch.CharTensor, torch.FloatTensor]: result = dequantize_asymmetric(weight, scale, zero_point, scale.dtype, result_shape) if result.ndim > 2: # convs diff --git a/modules/sdnq/layers/linear/linear_fp8.py b/modules/sdnq/layers/linear/linear_fp8.py index 9b22f4696..5707ad225 100644 --- a/modules/sdnq/layers/linear/linear_fp8.py +++ b/modules/sdnq/layers/linear/linear_fp8.py @@ -5,12 +5,12 @@ from typing import Tuple import torch from ...common import use_torch_compile # noqa: TID252 +from ...dequantizer import quantize_fp8 # noqa: TID252 def quantize_fp8_matmul_input(input: torch.FloatTensor) -> Tuple[torch.Tensor, torch.FloatTensor]: input = input.flatten(0,-2).to(dtype=torch.float32) - input_scale = torch.amax(input.abs(), dim=-1, keepdims=True).div_(448) - input = torch.div(input, input_scale).clamp_(-448, 448).to(dtype=torch.float8_e4m3fn) + input, input_scale = quantize_fp8(input, dim=-1) return input, input_scale diff --git a/modules/sdnq/layers/linear/linear_fp8_tensorwise.py b/modules/sdnq/layers/linear/linear_fp8_tensorwise.py index 778194e3d..aa80251ae 100644 --- a/modules/sdnq/layers/linear/linear_fp8_tensorwise.py +++ b/modules/sdnq/layers/linear/linear_fp8_tensorwise.py @@ -5,13 +5,12 @@ from typing import Tuple import torch from ...common import use_torch_compile # noqa: TID252 -from ...dequantizer import dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252 +from ...dequantizer import quantize_fp8, dequantize_symmetric, dequantize_symmetric_with_bias # noqa: TID252 def quantize_fp8_matmul_input_tensorwise(input: torch.FloatTensor, scale: torch.FloatTensor) -> Tuple[torch.Tensor, torch.FloatTensor]: input = input.flatten(0,-2).to(dtype=scale.dtype) - input_scale = torch.amax(input.abs(), dim=-1, keepdims=True).div_(448) - input = torch.div(input, input_scale).clamp_(-448, 448).to(dtype=torch.float8_e4m3fn) + input, input_scale = quantize_fp8(input, dim=-1) scale = torch.mul(input_scale, scale) if scale.dtype == torch.float16: # fp16 will overflow scale = scale.to(dtype=torch.float32) diff --git a/modules/sdnq/layers/linear/linear_int8.py b/modules/sdnq/layers/linear/linear_int8.py index ca202ea6e..3a1b82c1f 100644 --- a/modules/sdnq/layers/linear/linear_int8.py +++ b/modules/sdnq/layers/linear/linear_int8.py @@ -1,6 +1,7 @@ # pylint: disable=relative-beyond-top-level,redefined-builtin,protected-access from typing import Tuple + import torch from ...common import use_torch_compile # noqa: TID252