From 2d05396b4ee9d122b08b27984d246d1898ab5ef4 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Thu, 12 Jun 2025 02:26:04 +0300 Subject: [PATCH] SDNQ simplify sym scale formula --- modules/sdnq/__init__.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/modules/sdnq/__init__.py b/modules/sdnq/__init__.py index 64134723a..717e130da 100644 --- a/modules/sdnq/__init__.py +++ b/modules/sdnq/__init__.py @@ -204,9 +204,7 @@ def get_scale_asymmetric(weight: torch.FloatTensor, reduction_axes: List[int], w def get_scale_symmetric(weight: torch.FloatTensor, reduction_axes: List[int], weights_dtype: str) -> torch.FloatTensor: - abs_min_values = torch.amin(weight, dim=reduction_axes, keepdims=True).abs_() - max_values = torch.amax(weight, dim=reduction_axes, keepdims=True) - scale = torch.where(abs_min_values >= max_values, abs_min_values, -max_values).div_(dtype_dict[weights_dtype]["max"]) + scale = torch.amax(weight.abs(), dim=reduction_axes, keepdims=True).div_(dtype_dict[weights_dtype]["max"]) eps = torch.finfo(scale.dtype).eps # prevent divison by 0 scale = torch.where(torch.abs(scale) < eps, eps, scale) return scale