update sdnq

This commit is contained in:
Disty0
2026-01-14 16:23:26 +03:00
parent 4dbdee10e3
commit 784cda80aa
13 changed files with 121 additions and 33 deletions
+11 -3
View File
@@ -154,9 +154,9 @@ if hasattr(torch, "float8_e5m2fnuz"):
dtype_dict["float8_e5m2fnuz"] = {"min": -57344.0, "max": 57344.0, "num_bits": 8, "sign": 1, "exponent": 5, "mantissa": 2, "target_dtype": "fp8", "torch_dtype": torch.float8_e5m2fnuz, "storage_dtype": torch.float8_e5m2fnuz, "is_unsigned": False, "is_integer": False, "is_packed": False}
torch_dtype_dict[torch.float8_e5m2fnuz] = "float8_e5m2fnuz"
linear_types = {"Linear"}
conv_types = {"Conv1d", "Conv2d", "Conv3d"}
conv_transpose_types = {"ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d"}
linear_types = {"Linear", "SDNQLinear"}
conv_types = {"Conv1d", "Conv2d", "Conv3d", "SDNQConv1d", "SDNQConv2d", "SDNQConv3d"}
conv_transpose_types = {"ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d", "SDNQConvTranspose1d", "SDNQConvTranspose2d", "SDNQConvTranspose3d"}
allowed_types = set.union(linear_types, conv_types, conv_transpose_types)
accepted_weight_dtypes = set(dtype_dict.keys())
@@ -314,6 +314,14 @@ module_skip_keys_dict = {
["layers.0.adaLN_modulation.0.weight", "t_embedder", "cap_embedder", "siglip_embedder", "all_x_embedder", "all_final_layer"],
{}
],
"GlmImageTransformer2DModel": [
["transformer_blocks.0.norm1.linear.weight", "image_projector", "glyph_projector", "prior_projector", "time_condition_embed", "norm_out", "proj_out"],
{}
],
"GlmImageForConditionalGeneration": [
["lm_head", "patch_embed", "embeddings", "embed_tokens", "vqmodel"],
{}
],
"HunyuanImage3ForCausalMM": [
["lm_head", "patch_embed", "time_embed", "time_embed_2", "final_layer", "wte", "ln_f", "timestep_emb", "vae", "vision_aligner", "head", "post_layernorm", "embeddings"],
{}
+10 -7
View File
@@ -95,7 +95,7 @@ def quantize_int_mm(input: torch.FloatTensor, dim: int = -1, matmul_dtype: str =
@devices.inference_context()
def quantize_int_mm_sr(input: torch.FloatTensor, dim: int = -1, matmul_dtype: str = "int8") -> Tuple[torch.Tensor, torch.FloatTensor]:
scale = torch.amax(input.abs(), dim=dim, keepdims=True).div_(dtype_dict[matmul_dtype]["max"])
input = torch.div(input, scale).add_(torch.randn_like(input), alpha=0.1).round_().clamp_(dtype_dict[matmul_dtype]["min"], dtype_dict[matmul_dtype]["max"]).to(dtype=dtype_dict[matmul_dtype]["torch_dtype"])
input = torch.div(input, scale).add_(torch.rand_like(input), alpha=0.1).round_().clamp_(dtype_dict[matmul_dtype]["min"], dtype_dict[matmul_dtype]["max"]).to(dtype=dtype_dict[matmul_dtype]["torch_dtype"])
return input, scale
@@ -308,6 +308,7 @@ class SDNQDequantizer:
def __call__(self, weight, scale, zero_point, svd_up, svd_down, skip_quantized_matmul: bool = False, skip_compile: bool = False, dtype: torch.dtype = None): # pylint: disable=unused-argument
if dtype is None:
dtype = self.result_dtype
re_quantize_for_matmul = self.re_quantize_for_matmul or self.is_packed
if self.is_packed:
if self.is_integer:
if self.is_unsigned:
@@ -317,9 +318,9 @@ class SDNQDequantizer:
return dequantize_packed_int_asymmetric_compiled(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul)
else:
if skip_compile:
return dequantize_packed_int_symmetric(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=self.re_quantize_for_matmul)
return dequantize_packed_int_symmetric(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul)
else:
return dequantize_packed_int_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=self.re_quantize_for_matmul)
return dequantize_packed_int_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul)
else:
if self.is_unsigned:
if skip_compile: # compiled training needs to be traced with the original function
@@ -328,9 +329,9 @@ class SDNQDequantizer:
return dequantize_packed_float_asymmetric_compiled(weight, scale, zero_point, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul)
else:
if skip_compile:
return dequantize_packed_float_symmetric(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=self.re_quantize_for_matmul)
return dequantize_packed_float_symmetric(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul)
else:
return dequantize_packed_float_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=self.re_quantize_for_matmul)
return dequantize_packed_float_symmetric_compiled(weight, scale, self.quantized_weight_shape, self.weights_dtype, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul)
else:
if self.is_unsigned:
if skip_compile:
@@ -339,9 +340,9 @@ class SDNQDequantizer:
return dequantize_asymmetric_compiled(weight, scale, zero_point, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul)
else:
if skip_compile:
return dequantize_symmetric(weight, scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=self.re_quantize_for_matmul)
return dequantize_symmetric(weight, scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul)
else:
return dequantize_symmetric_compiled(weight, scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=self.re_quantize_for_matmul)
return dequantize_symmetric_compiled(weight, scale, svd_up=svd_up, svd_down=svd_down, dtype=dtype, result_shape=self.result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul)
dequantize_asymmetric_compiled = compile_func(dequantize_asymmetric)
@@ -356,3 +357,5 @@ re_quantize_matmul_packed_int_asymmetric_compiled = compile_func(re_quantize_mat
re_quantize_matmul_packed_int_symmetric_compiled = compile_func(re_quantize_matmul_packed_int_symmetric)
re_quantize_matmul_packed_float_asymmetric_compiled = compile_func(re_quantize_matmul_packed_float_asymmetric)
re_quantize_matmul_packed_float_symmetric_compiled = compile_func(re_quantize_matmul_packed_float_symmetric)
torch.serialization.add_safe_globals([SDNQDequantizer])
+69
View File
@@ -0,0 +1,69 @@
import torch
class SDNQLayer(torch.nn.Module):
def __init__(self, original_layer, forward_func):
torch.nn.Module.__init__(self)
for key, value in original_layer.__dict__.items():
if key not in {"forward", "forward_func", "original_class"}:
setattr(self, key, value)
self.original_class = original_layer.__class__
self.forward_func = forward_func
def forward(self, *args, **kwargs) -> torch.Tensor:
return self.forward_func(self, *args, **kwargs)
def __repr__(self):
return f"{self.__class__.__name__}(original_class={self.original_class.__name__} forward_func={self.forward_func} sdnq_dequantizer={repr(getattr(self, 'sdnq_dequantizer', None))})"
class SDNQLinear(SDNQLayer, torch.nn.Linear):
original_class: torch.nn.Linear
class SDNQConv1d(SDNQLayer, torch.nn.Conv1d):
original_class: torch.nn.Conv1d
class SDNQConv2d(SDNQLayer, torch.nn.Conv2d):
original_class: torch.nn.Conv2d
class SDNQConv3d(SDNQLayer, torch.nn.Conv3d):
original_class: torch.nn.Conv3d
class SDNQConvTranspose1d(SDNQLayer, torch.nn.ConvTranspose1d):
original_class: torch.nn.ConvTranspose1d
class SDNQConvTranspose2d(SDNQLayer, torch.nn.ConvTranspose2d):
original_class: torch.nn.ConvTranspose2d
class SDNQConvTranspose3d(SDNQLayer, torch.nn.ConvTranspose3d):
original_class: torch.nn.ConvTranspose3d
torch.serialization.add_safe_globals([SDNQLayer])
torch.serialization.add_safe_globals([SDNQLinear])
torch.serialization.add_safe_globals([SDNQConv1d])
torch.serialization.add_safe_globals([SDNQConv2d])
torch.serialization.add_safe_globals([SDNQConv3d])
torch.serialization.add_safe_globals([SDNQConvTranspose1d])
torch.serialization.add_safe_globals([SDNQConvTranspose2d])
torch.serialization.add_safe_globals([SDNQConvTranspose3d])
def get_sdnq_wrapper_class(original_layer, forward_func):
match original_layer.__class__.__name__:
case "Linear":
return SDNQLinear(original_layer, forward_func)
case "Conv1d":
return SDNQConv1d(original_layer, forward_func)
case "Conv2d":
return SDNQConv2d(original_layer, forward_func)
case "Conv3d":
return SDNQConv3d(original_layer, forward_func)
case "ConvTranspose1d":
return SDNQConvTranspose1d(original_layer, forward_func)
case "ConvTranspose2d":
return SDNQConvTranspose2d(original_layer, forward_func)
case "ConvTranspose3d":
return SDNQConvTranspose3d(original_layer, forward_func)
case _:
return SDNQLayer(original_layer, forward_func)
+3 -2
View File
@@ -37,11 +37,12 @@ def conv_fp16_matmul(
else:
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
input, scale = quantize_fp_mm_input_tensorwise(input, scale, matmul_dtype="float16")
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float16)
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float16).t_()
scale = scale.t()
elif weight.dtype != torch.float16:
weight = weight.to(dtype=torch.float16) # fp8 weights
input, scale = quantize_fp_mm_input_tensorwise(input, scale, matmul_dtype="float16")
input, weight = check_mats(input, weight)
if groups == 1:
+3 -2
View File
@@ -33,9 +33,10 @@ def conv_fp8_matmul(
input = input.flatten(0,-2)
svd_bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
input, input_scale = quantize_fp_mm_input(input)
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn)
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_()
scale = scale.t()
input, input_scale = quantize_fp_mm_input(input)
input, weight = check_mats(input, weight)
if groups == 1:
@@ -37,9 +37,10 @@ def conv_fp8_matmul_tensorwise(
else:
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
input, scale = quantize_fp_mm_input_tensorwise(input, scale)
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn)
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_()
scale = scale.t()
input, scale = quantize_fp_mm_input_tensorwise(input, scale)
input, weight = check_mats(input, weight)
dummy_input_scale = torch.ones(1, device=input.device, dtype=torch.float32)
+3 -2
View File
@@ -37,9 +37,10 @@ def conv_int8_matmul(
else:
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
input, scale = quantize_int_mm_input(input, scale)
if quantized_weight_shape is not None:
weight = unpack_int_symetric(weight, quantized_weight_shape, weights_dtype, dtype=torch.int8)
weight = unpack_int_symetric(weight, quantized_weight_shape, weights_dtype, dtype=torch.int8).t_()
scale = scale.t()
input, scale = quantize_int_mm_input(input, scale)
input, weight = check_mats(input, weight)
if groups == 1:
+2 -1
View File
@@ -21,7 +21,8 @@ def fp16_matmul(
weights_dtype: str = None,
) -> torch.FloatTensor:
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float16)
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float16).t_()
scale = scale.t()
elif weight.dtype != torch.float16:
weight = weight.to(dtype=torch.float16) # fp8 weights
return_dtype = input.dtype
+2 -1
View File
@@ -28,7 +28,8 @@ def fp8_matmul(
weights_dtype: str = None,
) -> torch.FloatTensor:
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn)
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_()
scale = scale.t()
return_dtype = input.dtype
output_shape = (*input.shape[:-1], weight.shape[-1])
if svd_up is not None:
@@ -31,7 +31,8 @@ def fp8_matmul_tensorwise(
weights_dtype: str = None,
) -> torch.FloatTensor:
if quantized_weight_shape is not None:
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn)
weight = unpack_float(weight, quantized_weight_shape, weights_dtype).to(dtype=torch.float8_e4m3fn).t_()
scale = scale.t()
return_dtype = input.dtype
output_shape = (*input.shape[:-1], weight.shape[-1])
if svd_up is not None:
+2 -1
View File
@@ -31,7 +31,8 @@ def int8_matmul(
weights_dtype: str = None,
) -> torch.FloatTensor:
if quantized_weight_shape is not None:
weight = unpack_int_symetric(weight, quantized_weight_shape, weights_dtype, dtype=torch.int8)
weight = unpack_int_symetric(weight, quantized_weight_shape, weights_dtype, dtype=torch.int8).t_()
scale = scale.t()
return_dtype = input.dtype
output_shape = (*input.shape[:-1], weight.shape[-1])
if svd_up is not None:
+3 -4
View File
@@ -170,7 +170,7 @@ def apply_sdnq_options_to_module(model, dtype: torch.dtype = None, dequantize_fp
return model
for module_name, module in model.named_children():
if hasattr(module, "sdnq_dequantizer"):
layer_class_name = module.__class__.__name__
layer_class_name = module.original_class.__name__
current_use_quantized_matmul = use_quantized_matmul
if current_use_quantized_matmul:
if layer_class_name in conv_types:
@@ -204,7 +204,7 @@ def apply_sdnq_options_to_module(model, dtype: torch.dtype = None, dequantize_fp
module.svd_down.data = module.svd_down.to(dtype=scale_dtype)
if current_use_quantized_matmul is not None and current_use_quantized_matmul != module.sdnq_dequantizer.use_quantized_matmul:
if not module.sdnq_dequantizer.re_quantize_for_matmul:
if not module.sdnq_dequantizer.re_quantize_for_matmul and not dtype_dict[module.sdnq_dequantizer.weights_dtype]["is_packed"]:
module.scale.t_()
module.weight.t_()
if current_use_quantized_matmul:
@@ -215,8 +215,7 @@ def apply_sdnq_options_to_module(model, dtype: torch.dtype = None, dequantize_fp
if module.svd_up is not None:
module.svd_up.data, module.svd_down.data = prepare_svd_for_matmul(module.svd_up.t_(), module.svd_down.t_(), current_use_quantized_matmul)
module.sdnq_dequantizer.use_quantized_matmul = current_use_quantized_matmul
module.forward = get_forward_func(module.__class__.__name__, module.sdnq_dequantizer.quantized_matmul_dtype, current_use_quantized_matmul)
module.forward = module.forward.__get__(module, module.__class__)
module.forward_func = get_forward_func(module.original_class.__name__, module.sdnq_dequantizer.quantized_matmul_dtype, current_use_quantized_matmul)
setattr(model, module_name, module)
else:
setattr(model, module_name, apply_sdnq_options_to_module(module, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul))
+8 -7
View File
@@ -20,6 +20,7 @@ from .dequantizer import SDNQDequantizer, dequantize_sdnq_model
from .packed_int import pack_int_symetric, pack_int_asymetric
from .packed_float import pack_float
from .forward import get_forward_func
from .layers import get_sdnq_wrapper_class
class QuantizationMethod(str, Enum):
@@ -55,7 +56,7 @@ def quantize_weight(weight: torch.FloatTensor, reduction_axes: Union[int, List[i
if dtype_dict[weights_dtype]["is_integer"]:
if use_stochastic_rounding:
quantized_weight.add_(torch.randn_like(quantized_weight), alpha=0.1)
quantized_weight.add_(torch.rand_like(quantized_weight), alpha=0.1)
quantized_weight.round_()
else:
if use_stochastic_rounding:
@@ -352,7 +353,7 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
svd_down = svd_down.to(dtype=torch_dtype)
re_quantize_for_matmul = re_quantize_for_matmul or num_of_groups > 1
if use_quantized_matmul and not re_quantize_for_matmul:
if use_quantized_matmul and not re_quantize_for_matmul and not dtype_dict[weights_dtype]["is_packed"]:
scale.t_()
weight.t_()
weight = prepare_weight_for_matmul(weight)
@@ -436,7 +437,7 @@ def sdnq_quantize_layer_weight_dynamic(weight, layer_class_name=None, weights_dt
svd_down = svd_down.t_()
svd_is_transposed = True
quantization_loss = torch.nn.functional.mse_loss(weight, sdnq_dequantizer(quantized_weight, scale, zero_point, svd_up, svd_down, skip_quantized_matmul=sdnq_dequantizer.use_quantized_matmul, dtype=torch.float32)).div_(weight_std)
quantization_loss = torch.nn.functional.mse_loss(weight, sdnq_dequantizer(quantized_weight, scale, zero_point, svd_up, svd_down, skip_quantized_matmul=sdnq_dequantizer.use_quantized_matmul, dtype=torch.float32, skip_compile=True)).div_(weight_std)
if quantization_loss <= dynamic_loss_threshold:
return (quantized_weight, scale, zero_point, svd_up, svd_down, sdnq_dequantizer)
return None
@@ -499,6 +500,7 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", quantized_matmul_dtype=None
) = weight_data
del weight_data
layer = get_sdnq_wrapper_class(layer, get_forward_func(layer_class_name, layer.sdnq_dequantizer.quantized_matmul_dtype, layer.sdnq_dequantizer.use_quantized_matmul))
layer.weight = torch.nn.Parameter(layer.weight.to(return_device, non_blocking=non_blocking), requires_grad=False)
layer.scale = torch.nn.Parameter(layer.scale.to(return_device, non_blocking=non_blocking), requires_grad=False)
if layer.zero_point is not None:
@@ -506,10 +508,7 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", quantized_matmul_dtype=None
if layer.svd_up is not None:
layer.svd_up = torch.nn.Parameter(layer.svd_up.to(return_device, non_blocking=non_blocking), requires_grad=False)
layer.svd_down = torch.nn.Parameter(layer.svd_down.to(return_device, non_blocking=non_blocking), requires_grad=False)
layer = layer.to(return_device, non_blocking=non_blocking)
layer.forward = get_forward_func(layer_class_name, layer.sdnq_dequantizer.quantized_matmul_dtype, layer.sdnq_dequantizer.use_quantized_matmul)
layer.forward = layer.forward.__get__(layer, layer.__class__)
if use_dynamic_quantization:
if modules_dtype_dict is None:
@@ -814,7 +813,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
else:
param_value = param_value.to(target_device, non_blocking=self.quantization_config.non_blocking).to(dtype=torch.float32)
layer, _ = get_module_from_name(model, param_name)
layer, tensor_name = get_module_from_name(model, param_name)
layer.weight = torch.nn.Parameter(param_value, requires_grad=False)
layer, self.quantization_config.modules_to_not_convert, self.quantization_config.modules_dtype_dict = sdnq_quantize_layer(
layer,
@@ -848,6 +847,8 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
if layer.svd_up is not None:
layer.svd_up._is_hf_initialized = True # pylint: disable=protected-access
layer.svd_down._is_hf_initialized = True # pylint: disable=protected-access
parent_module, tensor_name = get_module_from_name(model, param_name.removesuffix(tensor_name).removesuffix("."))
setattr(parent_module, tensor_name, layer)
def get_quantize_ops(self):
return SDNQQuantize(self)