From 73cab2acf557ebf1764610601302d4b96726916f Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Sun, 2 Aug 2026 08:29:17 -0600 Subject: [PATCH] Reworked merge in out of loras with convrot weights for better roundtrip accuracy. --- toolkit/lora_special.py | 29 +++++++++++++++++++-------- toolkit/network_mixins.py | 42 +++++++++++++++++++++++++++++---------- version.py | 2 +- 3 files changed, 54 insertions(+), 19 deletions(-) diff --git a/toolkit/lora_special.py b/toolkit/lora_special.py index 86994db..bd3c717 100644 --- a/toolkit/lora_special.py +++ b/toolkit/lora_special.py @@ -235,20 +235,33 @@ class FullModule(ToolkitModuleMixin, torch.nn.Module): def merge_in(self: 'FullModule', merge_weight=1.0): if not self.can_merge_in: return - om = self.org_module[0] - if 'weight._data' in om.state_dict(): - # quanto quantized weight, can't merge + # a zero diff merges to identity: skip entirely (a quantized base would + # otherwise still get requantized, which is not lossless) + if not self.diff.any() and (self.diff_b is None or not self.diff_b.any()): return - 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) + om = self.org_module[0] + if getattr(om, "is_ostris_quantized", False): + # fp32 dequant straight from the backend; the bf16 weight property + # would resample the quant scales on every merge cycle + orig_dtype = om.ostris_orig_dtype + base_weight = om.ostris_quantizer.dequantize(om) + weight_device = base_weight.device + else: + if 'weight._data' in om.state_dict(): + # quanto quantized weight, can't merge + return + org_weight = om.weight + orig_dtype = org_weight.dtype + base_weight = _dequantize_if_needed(org_weight).float() + weight_device = org_weight.device + # fold the full precision delta in + merged_weight = base_weight + merge_weight * self.diff.float().to(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) + om.weight.data = merged_weight.to(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) diff --git a/toolkit/network_mixins.py b/toolkit/network_mixins.py index cacd69d..761edee 100644 --- a/toolkit/network_mixins.py +++ b/toolkit/network_mixins.py @@ -378,20 +378,42 @@ class ToolkitModuleMixin: up_weight = self.lora_up.weight.clone().float() down_weight = self.lora_down.weight.clone().float() - # 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 - if 'weight._data' in org_sd: - # quantized weight + # a zero delta merges to identity: skip entirely. On quantized bases a + # "merge" is dequantize -> add -> requantize, which is not lossless (the + # scales resample), so an untrained module merging zero would still + # perturb the base weights the first time. + if self.full_rank: + if not down_weight.any(): + return + elif not up_weight.any() or not down_weight.any(): return weight_key = "weight" 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() + om = self.org_module[0] + org_sd = None + if not getattr(om, "is_ostris_quantized", False): + # extract weight from org_module (also dequantizes OstrisLinear, so + # only fetched on the non-ostris paths that actually use it) + org_sd = om.state_dict() + # todo find a way to merge in weights when doing quantized model + if 'weight._data' in org_sd: + # quantized weight + return + if getattr(om, "is_ostris_quantized", False): + # fp32 dequant straight from the backend. The bf16 weight property + # would re-round the reconstruction, and requantizing that resamples + # every row scale with bf16 error — repeated merge cycles walk the + # weights (~0.1% output drift per cycle per layer) + is_ao_quantized = True + orig_dtype = om.ostris_orig_dtype + weight = om.ostris_quantizer.dequantize(om) + else: + org_weight = om.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 diff --git a/version.py b/version.py index c1830e4..15ab0ed 100644 --- a/version.py +++ b/version.py @@ -1 +1 @@ -VERSION = "0.12.0" +VERSION = "0.12.1"