From f324b7c0e592463b15f347ef8e5d6514f1637e1e Mon Sep 17 00:00:00 2001 From: Disty0 Date: Thu, 21 Aug 2025 02:21:05 +0300 Subject: [PATCH] SDNQ remove unnecessary .contiguous() --- modules/sdnq/__init__.py | 4 ++-- modules/sdnq/layers/linear/linear_fp8.py | 2 +- modules/sdnq/layers/linear/linear_fp8_tensorwise.py | 2 +- modules/sdnq/layers/linear/linear_int8.py | 2 +- 4 files changed, 5 insertions(+), 5 deletions(-) diff --git a/modules/sdnq/__init__.py b/modules/sdnq/__init__.py index 642506842..548a3bd34 100644 --- a/modules/sdnq/__init__.py +++ b/modules/sdnq/__init__.py @@ -163,8 +163,8 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz scale.transpose_(0,1) layer.weight.transpose_(0,1) if not dtype_dict[weights_dtype]["is_integer"]: - stride = layer.weight.stride() - if stride[0] > stride[1] and stride[1] == 1: + weight_stride = layer.weight.stride() + if not (weight_stride[0] == 1 and weight_stride[1] > 1): layer.weight.data = layer.weight.t().contiguous().t() if not use_tensorwise_fp8_matmul: scale = scale.to(torch.float32) diff --git a/modules/sdnq/layers/linear/linear_fp8.py b/modules/sdnq/layers/linear/linear_fp8.py index ba261eb50..7613a5616 100644 --- a/modules/sdnq/layers/linear/linear_fp8.py +++ b/modules/sdnq/layers/linear/linear_fp8.py @@ -8,7 +8,7 @@ from ...common import use_torch_compile # noqa: TID252 def quantize_fp8_matmul_input(input: torch.FloatTensor) -> Tuple[torch.Tensor, torch.FloatTensor]: - input = input.flatten(0,-2).contiguous().to(dtype=torch.float32) + 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) 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 f07154aab..778194e3d 100644 --- a/modules/sdnq/layers/linear/linear_fp8_tensorwise.py +++ b/modules/sdnq/layers/linear/linear_fp8_tensorwise.py @@ -9,7 +9,7 @@ from ...dequantizer import dequantize_symmetric, dequantize_symmetric_with_bias def quantize_fp8_matmul_input_tensorwise(input: torch.FloatTensor, scale: torch.FloatTensor) -> Tuple[torch.Tensor, torch.FloatTensor]: - input = input.flatten(0,-2).contiguous().to(dtype=scale.dtype) + 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) scale = torch.mul(input_scale, scale) diff --git a/modules/sdnq/layers/linear/linear_int8.py b/modules/sdnq/layers/linear/linear_int8.py index 6d94052a2..fd1189c2e 100644 --- a/modules/sdnq/layers/linear/linear_int8.py +++ b/modules/sdnq/layers/linear/linear_int8.py @@ -10,7 +10,7 @@ from ...dequantizer import dequantize_symmetric, dequantize_symmetric_with_bias def quantize_int8_matmul_input(input: torch.FloatTensor, scale: torch.FloatTensor) -> Tuple[torch.CharTensor, torch.FloatTensor]: - input = input.flatten(0,-2).contiguous().to(dtype=scale.dtype) + input = input.flatten(0,-2).to(dtype=scale.dtype) input_scale = torch.amax(input.abs(), dim=-1, keepdims=True).div_(127) input = torch.div(input, input_scale).round_().clamp_(-128, 127).to(dtype=torch.int8) scale = torch.mul(input_scale, scale)