mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
SDNQ Flux lora, use shape from sdnq_dequantizer
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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.")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user