diff --git a/modules/sdnq/__init__.py b/modules/sdnq/__init__.py index 6eb2e3553..bca210c1b 100644 --- a/modules/sdnq/__init__.py +++ b/modules/sdnq/__init__.py @@ -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, diff --git a/pipelines/flux/flux_lora.py b/pipelines/flux/flux_lora.py index 50f138df5..d7c80dc81 100644 --- a/pipelines/flux/flux_lora.py +++ b/pipelines/flux/flux_lora.py @@ -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.")