SDNQ Flux lora, use shape from sdnq_dequantizer

This commit is contained in:
Disty0
2025-08-18 19:53:57 +03:00
parent 0f72024999
commit 47154db8b1
2 changed files with 9 additions and 14 deletions
+1 -10
View File
@@ -245,15 +245,6 @@ def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_si
return model
class SDNQParameter(torch.nn.Parameter):
def __new__(cls, data=None, requires_grad=False):
return super().__new__(cls, data, requires_grad)
def __init__(self, data=None, requires_grad=False): # pylint: disable=unused-argument
self.original_shape = data.shape
super().__init__()
class SDNQQuantizer(DiffusersQuantizer):
r"""
Diffusers Quantizer for SDNQ
@@ -342,7 +333,7 @@ class SDNQQuantizer(DiffusersQuantizer):
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.weight = SDNQParameter(param_value, requires_grad=False)
layer.weight = torch.nn.Parameter(param_value, requires_grad=False)
layer = sdnq_quantize_layer(
layer,
weights_dtype=weights_dtype,
+8 -4
View File
@@ -4,20 +4,24 @@ def calculate_module_shape(model, base_module=None, base_weight_param_name=None)
return weight.quant_state.shape
elif weight.__class__.__name__ == "GGUFParameter":
return weight.quant_shape
elif weight.__class__.__name__ == "SDNQParameter":
return weight.original_shape
else:
return weight.shape
if base_module is not None:
return _get_weight_shape(base_module.weight)
if hasattr(base_module, "sdnq_dequantizer"):
return base_module.sdnq_dequantizer.original_shape
else:
return _get_weight_shape(base_module.weight)
elif base_weight_param_name is not None:
from diffusers.utils import get_submodule_by_name
if not base_weight_param_name.endswith(".weight"):
raise ValueError(f"Invalid `base_weight_param_name` passed as it does not end with '.weight' {base_weight_param_name=}.")
module_path = base_weight_param_name.rsplit(".weight", 1)[0]
submodule = get_submodule_by_name(model, module_path)
return _get_weight_shape(submodule.weight)
if hasattr(submodule, "sdnq_dequantizer"):
return submodule.sdnq_dequantizer.original_shape
else:
return _get_weight_shape(submodule.weight)
raise ValueError("Either `base_module` or `base_weight_param_name` must be provided.")