diff --git a/modules/lora/lora_sdnq.py b/modules/lora/lora_sdnq.py index b58b26799..45a514693 100644 --- a/modules/lora/lora_sdnq.py +++ b/modules/lora/lora_sdnq.py @@ -272,12 +272,15 @@ def select_candidate(self, network_layer_name, wanted_names): def host_candidate(self, network_layer_name, wanted_names): - """True when a non-factorable set on this layer should be hosted as a truncated svd.""" + """True when this layer's set should ride the svd channel as a truncated svd: non-factorable sets below 8 bits, dense-combined sets at any width.""" if not select_candidate(self, network_layer_name, wanted_names): return False + if lora_stack.mode() in lora_stack.DENSE_MODES and not network_layer_name.startswith('lora_te'): + if sum(1 for net in l.loaded_networks if net.modules.get(network_layer_name, None) is not None) >= 2: + return True # combined deltas host at any width: requantizing them is checkpoint-fragile, while single-adapter requantize is well retained from sdnq.common import dtype_dict if dtype_dict[self.sdnq_dequantizer.weights_dtype]['num_bits'] >= 8: - return False # requantize retains most of the delta at 8 bits and above; truncation would lose more than it saves + return False # requantize retains most of a single set's delta at 8 bits and above; truncation would lose more than it saves return True @@ -312,7 +315,7 @@ def apply_cached(self, network_layer_name, wanted_names): factors = get_module_factors(module, devices.device, dtype, original_shape=deq.original_shape) if factors is not None: members.append(factors) - if len(members) == 0 and self.svd_up is None: + if not stack_dense and len(members) == 0 and self.svd_up is None: step = float(self.scale.detach().float().mean()) if step > 0 and rms / step > REQUANT_RATIO and energy < REQUANT_ENERGY: return None # routed to the grid: the caller assembles the delta and requantizes @@ -369,9 +372,9 @@ def apply_hosted(self, network_layer_name, updown, wanted_names): # cut: both terms must agree, since a thin delta rounds away on the grid however # low its capture, and a low-rank delta hosts exactly however fat it is. Scoped to # sets the side-channel would otherwise carry whole: factorable members ride - # exactly. + # exactly and dense-combined deltas stay hosted at any magnitude. delta_rms = float(updown.detach().float().square().mean().sqrt()) - maybe_requant = len(members) == 0 and self.svd_up is None + maybe_requant = not stack_dense and len(members) == 0 and self.svd_up is None if maybe_requant: step = float(self.scale.detach().float().mean()) maybe_requant = step > 0 and delta_rms / step > REQUANT_RATIO diff --git a/test/test-sdnq-lora-factors.py b/test/test-sdnq-lora-factors.py index a94e38b68..8dcae106c 100644 --- a/test/test-sdnq-lora-factors.py +++ b/test/test-sdnq-lora-factors.py @@ -872,6 +872,19 @@ def test_route_svd_checkpoint_keeps_hosting(): return True +def test_route_dense_stack_keeps_hosting(): + layer = build_layer('uint4') + torch.manual_seed(27) + D1 = torch.randn(OUT_F, IN_F, device=DEVICE) * 1e-2 + D2 = torch.randn(OUT_F, IN_F, device=DEVICE) * 1e-2 + net1, net2 = make_dense_net('df1', layer, D1), make_dense_net('df2', layer, D2) + with host_rank(256), stack_mode('ties'), mock_model(lin=layer): + activate(net1, net2) + assert hasattr(layer, 'sdnq_lora_svd_stash'), 'dense-combined deltas host at any magnitude' + activate() + return True + + def test_route_replay_from_cache(): import tempfile layer = build_layer('uint4') @@ -1515,6 +1528,39 @@ def test_dense_two_plain_loras_hosted_not_summed(): return True +def test_dense_pair_hosts_at_int8(): + layer = build_layer('int8') + A1, B1, D1 = make_delta(seed=41, sigma=1e-2) + A2, B2, D2 = make_delta(seed=42, sigma=1e-2) + n1 = make_net('ti1', layer, A1, B1) + n2 = make_net('ti2', layer, A2, B2) + with host_rank(64), mock_model(lin=layer): + Wdq0 = dq(layer) + with stack_mode('ties', dens=0.5): + activate(n1, n2) + assert hasattr(layer, 'sdnq_lora_svd_stash'), 'dense pair at int8 must host, not requantize' + assert getattr(layer, 'network_weights_backup', None) is None, 'hosted dense pair must not take a weight backup' + eff = dq(layer) - Wdq0 + activate() + assert torch.equal(dq(layer), Wdq0), 'removal must restore bit-exact' + with stack_mode('ties', dens=0.5): + ref = lora_stack.combine([('ti1', D1), ('ti2', D2)], 'lora_transformer_test') + assert rho_of(eff, ref) > 0.8, f'hosted int8 ties delta must track the ties reference, rho={rho_of(eff, ref):.3f}' + return True + + +def test_dense_single_nonfactorable_int8_keeps_requantize(): + layer = build_layer('int8') + _A, _B, D = make_delta(sigma=3e-3) + net = make_dense_net('ti8solo', layer, D) + with host_rank(256), mock_model(lin=layer), stack_mode('ties', dens=0.5): + activate(net) + assert not hasattr(layer, 'sdnq_lora_svd_stash'), 'a single non-factorable set at int8 must keep the requantize path even under a dense mode' + assert isinstance(getattr(layer, 'network_weights_backup', None), torch.Tensor), 'the requantize fallback must take the backup' + activate() + return True + + def test_single_net_ignores_dense_mode(): layer = build_layer('uint4') A, B, D = make_delta(seed=33) @@ -2018,7 +2064,7 @@ def run_tests(): for fn in [test_hosted_low_rank_delta_is_kept, test_hosted_dense_delta_beats_requant, test_hosted_skips_int8, test_hosted_disabled_by_option, test_hosted_transitions_and_rng_isolation, test_route_fat_dense_delta_requantizes, test_route_rule_terms_gate_both_ways, test_route_low_rank_fat_delta_stays_hosted, - test_route_mixed_set_keeps_hosting, test_route_svd_checkpoint_keeps_hosting, + test_route_mixed_set_keeps_hosting, test_route_svd_checkpoint_keeps_hosting, test_route_dense_stack_keeps_hosting, test_route_replay_from_cache, test_hosted_null_tail_collapses_to_effective_rank, test_hosted_flat_spectrum_keeps_cap]: run_test(CAT_HOST, fn) log.warning('=== Calibration ===') @@ -2034,7 +2080,8 @@ def run_tests(): run_test(CAT_FCACHE, fn) log.warning('=== Stack modes: dense ===') for fn in [test_ties_sign_consensus_drops_conflicts, test_dare_mask_is_deterministic_across_calls, test_dare_rescales_by_inverse_density, - test_magnitude_prune_keeps_top_density, test_dense_two_plain_loras_hosted_not_summed, test_single_net_ignores_dense_mode, + test_magnitude_prune_keeps_top_density, test_dense_two_plain_loras_hosted_not_summed, + test_dense_pair_hosts_at_int8, test_dense_single_nonfactorable_int8_keeps_requantize, test_single_net_ignores_dense_mode, test_te_layer_stays_plain_sum, test_sum_mode_keeps_exact_stacking]: run_test(CAT_STACK, fn) log.warning('=== Stack modes: select ===')