From 5f04ae7ad50e11b1e42f728908bac7b246fed7c5 Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Thu, 25 Jun 2026 13:19:38 -0600 Subject: [PATCH] Rework merge_network_on_save to handle dequantization on merging and saving for much more efficient full finetuning. --- .../hidream/hidream_o1_model.py | 4 +- jobs/process/BaseSDTrainProcess.py | 22 ++++++--- toolkit/lora_special.py | 48 ++++++++++++++----- toolkit/models/base_model.py | 6 +++ toolkit/models/lokr.py | 29 ++++++----- toolkit/network_mixins.py | 36 ++++++++++---- toolkit/stable_diffusion_model.py | 6 +++ toolkit/util/quantize.py | 37 +++++++++++++- 8 files changed, 147 insertions(+), 41 deletions(-) diff --git a/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py b/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py index 80790877..cf58ee0b 100644 --- a/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py +++ b/extensions_built_in/diffusion_models/hidream/hidream_o1_model.py @@ -489,6 +489,7 @@ class HidreamO1Model(BaseModel): return False def save_model(self, output_path, meta, save_dtype): + from toolkit.util.quantize import dequantize_if_quantized transformer: Qwen3VLForConditionalGeneration = unwrap_model(self.model) if self.is_comfy_weight: sd = transformer.state_dict() @@ -496,7 +497,8 @@ class HidreamO1Model(BaseModel): for key, value in sd.items(): if "lm_head.weight" in key: continue # comfy checkpoint doesnt have the lm head, so skip it - save_dict[key] = value.clone().to("cpu", dtype=save_dtype) + # dequantize any quantized (e.g. torchao) weights so we save plain full precision tensors + save_dict[key] = dequantize_if_quantized(value).clone().to("cpu", dtype=save_dtype) if not output_path.endswith(".safetensors"): output_path += ".safetensors" diff --git a/jobs/process/BaseSDTrainProcess.py b/jobs/process/BaseSDTrainProcess.py index e3817a38..4e7945fe 100644 --- a/jobs/process/BaseSDTrainProcess.py +++ b/jobs/process/BaseSDTrainProcess.py @@ -626,18 +626,26 @@ class BaseSDTrainProcess(BaseTrainProcess): ) else: if self.network is not None and self.train_config.merge_network_on_save: - # merge the network weights into a full model and save that - if not self.network.can_merge_in: + # merge the network weights into a full model and save that. + # torchao quantized weights can be force merged here (dequantize -> merge -> re-quantize) + # even though can_merge_in is False (kept False so sampling never merges). quanto and + # layer_offloading still cannot merge. + from toolkit.util.quantize import get_torchao_config + can_force_quantized_merge = ( + self.model_config.quantize and not self.model_config.layer_offloading + and get_torchao_config(self.model_config.qtype) is not None + ) + if not self.network.can_merge_in and not can_force_quantized_merge: raise ValueError("Network cannot merge in weights. Cannot save full model.") - + print_acc("Merging network weights into full model for saving...") - + self.network.merge_in(merge_weight=self.train_config.merge_network_on_save_strength) # reset weights to zero self.network.reset_weights() self.network.is_merged_in = False - print_acc("Done merging network weights.") + print_acc("Done merging network weights. Saving model...") if self.save_config.save_format == "diffusers": # saving as a folder path @@ -1810,7 +1818,9 @@ class BaseSDTrainProcess(BaseTrainProcess): self.train_config.train_unet ) - # we cannot merge in if quantized + # we cannot merge in if quantized or offloading. note: torchao quantized weights can + # still be force merged at save time for the merge-and-reset method (see save logic), + # but we keep can_merge_in False here so sampling never merges in/out. if self.model_config.quantize or self.model_config.layer_offloading: # todo find a way around this self.network.can_merge_in = False diff --git a/toolkit/lora_special.py b/toolkit/lora_special.py index a242d4e5..61eb3076 100644 --- a/toolkit/lora_special.py +++ b/toolkit/lora_special.py @@ -135,6 +135,16 @@ class LoRAModule(ToolkitModuleMixin, ExtractableModuleMixin, torch.nn.Module): # del self.org_module +def _is_quantized_tensor(t) -> bool: + # torchao stores quantized weights as tensor subclasses (e.g. AffineQuantizedTensor) under torchao.* + # that are still nn.Parameter instances and expose .dequantize(). (quanto is intentionally not handled.) + return 'torchao' in type(t).__module__ and hasattr(t, 'dequantize') + + +def _dequantize_if_needed(t): + return t.dequantize() if _is_quantized_tensor(t) else t + + class FullModule(ToolkitModuleMixin, torch.nn.Module): """ Full weight "lora" for layers that have no sensible low rank decomposition (norm layers, embeddings, @@ -142,6 +152,9 @@ class FullModule(ToolkitModuleMixin, torch.nn.Module): the original weight (and bias) of the wrapped module. On save it emits `.diff` (and `.diff_b` for bias) which ComfyUI applies as `weight += strength * diff`, so it merges directly into the model weights without any extra adapter. + + If the wrapped module's weight is torchao-quantized, the delta is kept in full precision and the original + weight is dequantized on the fly in the forward pass (the original quantized tensor is left untouched). """ def __init__( @@ -167,7 +180,10 @@ class FullModule(ToolkitModuleMixin, torch.nn.Module): self.is_checkpointing = False # trainable delta, zero initialized so an untrained layer is a no-op (zero diff) - self.diff = torch.nn.Parameter(torch.zeros_like(org_module.weight)) + # dequantize first so the delta is full precision and shaped like the real (unpacked) weight + self.weight_is_quantized = _is_quantized_tensor(org_module.weight) + ref_weight = _dequantize_if_needed(org_module.weight) + self.diff = torch.nn.Parameter(torch.zeros_like(ref_weight)) # some modules (e.g. Embedding) have no bias attribute at all org_bias = getattr(org_module, 'bias', None) if org_bias is not None: @@ -191,7 +207,10 @@ class FullModule(ToolkitModuleMixin, torch.nn.Module): mult = multiplier.mean() if isinstance(multiplier, torch.Tensor) else multiplier orig_weight = om._parameters['weight'] - eff_weight = orig_weight + (self.diff.to(orig_weight.device) * mult).to(orig_weight.dtype) + # dequantize quantized weights to full precision so the delta can be added (the original + # quantized tensor is restored in the finally block below) + base_weight = _dequantize_if_needed(orig_weight) + eff_weight = base_weight + (self.diff.to(base_weight.device) * mult).to(base_weight.dtype) has_bias = self.diff_b is not None and om._parameters.get('bias', None) is not None if has_bias: @@ -215,16 +234,23 @@ class FullModule(ToolkitModuleMixin, torch.nn.Module): def merge_in(self: 'FullModule', merge_weight=1.0): if not self.can_merge_in: return - org_sd = self.org_module[0].state_dict() - if 'weight._data' in org_sd: - # quantized weight, can't merge + om = self.org_module[0] + if 'weight._data' in om.state_dict(): + # quanto quantized weight, can't merge return - weight = org_sd['weight'] - org_sd['weight'] = (weight.float() + merge_weight * self.diff.float().to(weight.device)).to(weight.dtype) - if self.diff_b is not None and 'bias' in org_sd: - bias = org_sd['bias'] - org_sd['bias'] = (bias.float() + merge_weight * self.diff_b.float().to(bias.device)).to(bias.dtype) - self.org_module[0].load_state_dict(org_sd) + org_weight = om.weight + orig_dtype = org_weight.dtype + # dequantize torchao weights so we can fold the full precision delta in + merged_weight = _dequantize_if_needed(org_weight).float() + merge_weight * self.diff.float().to(org_weight.device) + if self.weight_is_quantized: + # re-quantize so the model stays quantized across continuous merge/reset cycles + from toolkit.util.quantize import get_torchao_config, requantize_module_weight + requantize_module_weight(om, merged_weight, orig_dtype, get_torchao_config(self._get_base_qtype())) + else: + om.weight.data = merged_weight.to(org_weight.device, orig_dtype) + # bias is never quantized + if self.diff_b is not None and getattr(om, 'bias', None) is not None: + om.bias.data = (om.bias.data.float() + merge_weight * self.diff_b.float().to(om.bias.device)).to(om.bias.dtype) def reset_weights(self: 'FullModule'): with torch.no_grad(): diff --git a/toolkit/models/base_model.py b/toolkit/models/base_model.py index 04e20b80..4caf3ae7 100644 --- a/toolkit/models/base_model.py +++ b/toolkit/models/base_model.py @@ -1594,6 +1594,12 @@ class BaseModel: def get_transformer_block_names(self) -> Optional[List[str]]: # override in child classes to get transformer block names for lora targeting return None + + def get_quantization_exclude_modules(self) -> Optional[List[str]]: + # override in child classes to keep sensitive modules in full precision when + # quantizing. Returns fnmatch patterns matched against the transformer's module + # names (e.g. "model.x_embedder*"). + return None def get_base_model_version(self) -> str: # override in child classes to get the base model version diff --git a/toolkit/models/lokr.py b/toolkit/models/lokr.py index 486c0121..74de6350 100644 --- a/toolkit/models/lokr.py +++ b/toolkit/models/lokr.py @@ -256,18 +256,18 @@ class LokrModule(ToolkitModuleMixin, nn.Module): # extract weight from org_module org_sd = self.org_module[0].state_dict() - # todo find a way to merge in weights when doing quantized model + # todo find a way to merge in weights when doing quanto quantized model if 'weight._data' in org_sd: - # quantized weight + # quanto quantized weight return weight_key = "weight" - if 'weight._data' in org_sd: - # quantized weight - weight_key = "weight._data" - - orig_dtype = org_sd[weight_key].dtype - weight = org_sd[weight_key].float() + from toolkit.util.quantize import is_quantized_tensor + org_weight = self.org_module[0].weight + is_ao_quantized = is_quantized_tensor(org_weight) + orig_dtype = org_weight.dtype + # dequantize torchao weights so the delta can be merged in full precision + weight = (org_weight.dequantize() if is_ao_quantized else org_weight).float() scale = self.scale # handle trainable scaler method locon does @@ -281,9 +281,16 @@ class LokrModule(ToolkitModuleMixin, nn.Module): + (lokr_weight * merge_weight).to(weight.device, dtype=weight.dtype) ) - # set weight to org_module - org_sd[weight_key] = merged_weight.to(orig_dtype) - self.org_module[0].load_state_dict(org_sd) + # write the merged weight back, re-quantizing if the original was torchao quantized so the + # model stays quantized across continuous merge/reset cycles + if is_ao_quantized: + from toolkit.util.quantize import get_torchao_config, requantize_module_weight + requantize_module_weight( + self.org_module[0], merged_weight, orig_dtype, get_torchao_config(self._get_base_qtype()) + ) + else: + org_sd[weight_key] = merged_weight.to(orig_dtype) + self.org_module[0].load_state_dict(org_sd) def get_orig_weight(self, device): weight = self.org_module[0].weight diff --git a/toolkit/network_mixins.py b/toolkit/network_mixins.py index 3b360f52..2012b5c3 100644 --- a/toolkit/network_mixins.py +++ b/toolkit/network_mixins.py @@ -337,6 +337,13 @@ class ToolkitModuleMixin: def disable_gradient_checkpointing(self: Module): self.is_checkpointing = False + def _get_base_qtype(self: Module): + # the qtype string the base model was quantized with (so we can re-quantize after merging), or None + network = self.network_ref() + base_ref = getattr(network, 'base_model_ref', None) + base = base_ref() if base_ref is not None else None + return getattr(getattr(base, 'model_config', None), 'qtype', None) + @torch.no_grad() def merge_out(self: Module, merge_out_weight=1.0): # make sure it is positive @@ -363,12 +370,12 @@ class ToolkitModuleMixin: return weight_key = "weight" - if 'weight._data' in org_sd: - # quantized weight - weight_key = "weight._data" - - orig_dtype = org_sd[weight_key].dtype - weight = org_sd[weight_key].float() + from toolkit.util.quantize import is_quantized_tensor + org_weight = self.org_module[0].weight + is_ao_quantized = is_quantized_tensor(org_weight) + orig_dtype = org_weight.dtype + # dequantize torchao weights so the delta can be merged in full precision + weight = (org_weight.dequantize() if is_ao_quantized else org_weight).float() multiplier = merge_weight scale = self.scale @@ -401,10 +408,19 @@ class ToolkitModuleMixin: # print(conved.size(), weight.size(), module.stride, module.padding) weight = weight + multiplier * conved * scale - # set weight to org_module - org_sd[weight_key] = weight.to(weight_device, orig_dtype) - self.org_module[0].load_state_dict(org_sd) - + # write the merged weight back, re-quantizing if the original was torchao quantized so the + # model stays quantized across continuous merge/reset cycles + if is_ao_quantized: + from toolkit.util.quantize import get_torchao_config, requantize_module_weight + config = get_torchao_config(self._get_base_qtype()) + if config is None: + print_once(f"Warning: merging into quantized layer {getattr(self, 'lora_name', '?')} " + f"without a known qtype; it will be left dequantized") + requantize_module_weight(self.org_module[0], weight.to(weight_device), orig_dtype, config) + else: + org_sd[weight_key] = weight.to(weight_device, orig_dtype) + self.org_module[0].load_state_dict(org_sd) + def reset_weights(self: Module): # reset the weights to zero org_sd = self.state_dict() diff --git a/toolkit/stable_diffusion_model.py b/toolkit/stable_diffusion_model.py index 7d9a93e3..a51e9b55 100644 --- a/toolkit/stable_diffusion_model.py +++ b/toolkit/stable_diffusion_model.py @@ -3138,6 +3138,12 @@ class StableDiffusion: # override in child classes to get transformer block names for lora targeting return None + def get_quantization_exclude_modules(self) -> Optional[List[str]]: + # override in child classes to keep sensitive modules in full precision when + # quantizing. Returns fnmatch patterns matched against the transformer's module + # names (e.g. "model.x_embedder*"). + return None + def get_base_model_version(self) -> str: if self.is_pixart: return 'pixart' diff --git a/toolkit/util/quantize.py b/toolkit/util/quantize.py index 3b0af71d..85a0652d 100644 --- a/toolkit/util/quantize.py +++ b/toolkit/util/quantize.py @@ -62,6 +62,36 @@ def get_qtype(qtype: Union[str, qtype]) -> qtype: return qtype +def is_quantized_tensor(t) -> bool: + # torchao stores quantized weights as tensor subclasses (e.g. AffineQuantizedTensor) under torchao.* + # that still report as nn.Parameter and expose .dequantize(). (quanto is handled separately.) + return 'torchao' in type(t).__module__ and hasattr(t, 'dequantize') + + +def dequantize_if_quantized(t): + return t.dequantize() if is_quantized_tensor(t) else t + + +def get_torchao_config(qtype): + # returns the torchao quantization config for a given qtype string, or None if it isn't torchao + if qtype is None: + return None + try: + q = get_qtype(qtype) + except Exception: + return None + return q.config if isinstance(q, aotype) else None + + +def requantize_module_weight(module, fp_weight, orig_dtype, config) -> None: + """Write a full precision weight back into module.weight, re-quantizing in place if a torchao + config is provided so the module stays quantized (used by the continuous merge/reset method). + If config is None the weight is left in full precision.""" + module.weight = torch.nn.Parameter(fp_weight.to(orig_dtype), requires_grad=False) + if config is not None: + torchao_quantize_(module, config) + + def quantize( model: torch.nn.Module, weights: Optional[Union[str, qtype, aotype]] = None, @@ -142,6 +172,9 @@ def quantize_model( # patch the state dict method patch_dequantization_on_save(model_to_quantize) + # sensitive modules to keep in full precision (fnmatch patterns) + exclude_modules = base_model.get_quantization_exclude_modules() or [] + if base_model.model_config.accuracy_recovery_adapter is not None: from toolkit.config_modules import NetworkConfig from toolkit.lora_special import LoRASpecialNetwork @@ -287,7 +320,7 @@ def quantize_model( quantize( model_to_quantize, weights=quantization_type, - exclude=lora_exclude_modules + exclude=lora_exclude_modules + exclude_modules ) else: # quantize model the original way without an accuracy recovery adapter @@ -319,5 +352,5 @@ def quantize_model( # device without having to move the transformer blocks to the device first base_model.print_and_status_update(" - quantizing extras") # model_to_quantize.to(base_model.device_torch, dtype=base_model.torch_dtype) - quantize(model_to_quantize, weights=quantization_type) + quantize(model_to_quantize, weights=quantization_type, exclude=exclude_modules) freeze(model_to_quantize)