From 00231ab0357ab6480e80f78e2b21c4fdd816ac8b Mon Sep 17 00:00:00 2001 From: Disty0 Date: Wed, 15 Jul 2026 22:29:28 +0300 Subject: [PATCH] cleanup --- modules/sdnq/layers/linear/forward.py | 3 ++- modules/sdnq/quant_utils.py | 5 ++++- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/modules/sdnq/layers/linear/forward.py b/modules/sdnq/layers/linear/forward.py index 2ab292c1c..38ef65f26 100644 --- a/modules/sdnq/layers/linear/forward.py +++ b/modules/sdnq/layers/linear/forward.py @@ -6,7 +6,8 @@ from ...kernel_wrappers import use_contiguous_int8_mm, use_contiguous_fp16_mm, u def check_mats(input: torch.Tensor, weight: torch.Tensor, matmul_dtype: str = "int8") -> tuple[torch.Tensor, torch.Tensor]: - input = input.contiguous() + if input is not None: + input = input.contiguous() if ( (use_contiguous_int8_mm and matmul_dtype in {"int8", "uint8"}) or (use_contiguous_fp16_mm and matmul_dtype in {"fp16", "float16"}) diff --git a/modules/sdnq/quant_utils.py b/modules/sdnq/quant_utils.py index f131bb181..4e52295bd 100644 --- a/modules/sdnq/quant_utils.py +++ b/modules/sdnq/quant_utils.py @@ -3,7 +3,7 @@ import torch from modules import devices -from .common import dtype_dict, conv_types, conv_transpose_types +from .common import dtype_dict, compile_func, conv_types, conv_transpose_types from .kernel_wrappers import use_contiguous_int8_mm, use_contiguous_fp16_mm, use_contiguous_fp8_mm from .utils import is_pow2, is_pow4, next_power_of_2 @@ -229,3 +229,6 @@ def quantize_fp_mm(weight: torch.FloatTensor, dim: int = -1, hadamard: torch.Flo weight = weight.add_(torch.randint_like(weight, low=0, high=mantissa_difference, dtype=torch.int32)).bitwise_and_(-mantissa_difference).view(dtype=torch.float32) weight = torch.div(weight, scale).nan_to_num_().clamp_(dtype_dict[matmul_dtype]["min"], dtype_dict[matmul_dtype]["max"]).to(dtype=dtype_dict[matmul_dtype]["torch_dtype"]) return weight, scale + + +rotate_hadamard_compiled = compile_func(rotate_hadamard)