mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
register sdnq classes
This commit is contained in:
@@ -27,8 +27,8 @@ class NetworkModuleLora(network.NetworkModule):
|
||||
return None
|
||||
linear_modules = [torch.nn.Linear, torch.nn.modules.linear.NonDynamicallyQuantizableLinear, torch.nn.MultiheadAttention, diffusers_lora.LoRACompatibleLinear]
|
||||
typ = type(self.sd_module)
|
||||
is_linear = typ in linear_modules or self.sd_module.__class__.__name__ in ["NNCFLinear", "QLinear", "Linear4bit"]
|
||||
is_conv = (typ in [torch.nn.Conv2d, diffusers_lora.LoRACompatibleConv]) or (self.sd_module.__class__.__name__ in ["NNCFConv2d", "QConv2d"]) or (typ.__name__ in ['downsampler_block', 'upsampler_block'])
|
||||
is_linear = typ in linear_modules or self.sd_module.__class__.__name__ in ["SDNQLinear", "QLinear", "Linear4bit"]
|
||||
is_conv = (typ in [torch.nn.Conv2d, diffusers_lora.LoRACompatibleConv]) or (self.sd_module.__class__.__name__ in ["SDNQConv2d", "QConv2d"]) or (typ.__name__ in ['downsampler_block', 'upsampler_block'])
|
||||
if is_linear:
|
||||
weight = weight.reshape(weight.shape[0], -1)
|
||||
module = torch.nn.Linear(weight.shape[1], weight.shape[0], bias=False)
|
||||
|
||||
@@ -18,7 +18,11 @@ offload_allow_none = ['sd', 'sdxl']
|
||||
offload_post = ['h1']
|
||||
offload_hook_instance = None
|
||||
balanced_offload_exclude = ['CogView4Pipeline', 'MeissonicPipeline']
|
||||
no_split_module_classes = ["Linear", "Conv1d", "Conv2d", "Conv3d", "ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d", "WanTransformerBlock"]
|
||||
no_split_module_classes = [
|
||||
"Linear", "Conv1d", "Conv2d", "Conv3d", "ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d",
|
||||
"SDNQLinear", "SDNQConv1d", "SDNQConv2d", "SDNQConv3d", "SDNQConvTranspose1d", "SDNQConvTranspose2d", "SDNQConvTranspose3d",
|
||||
"WanTransformerBlock",
|
||||
]
|
||||
accelerate_dtype_byte_size = None
|
||||
move_stream = None
|
||||
|
||||
|
||||
Reference in New Issue
Block a user