From 259e15fafe3f93ef61126f5ade3056b3e250b33e Mon Sep 17 00:00:00 2001 From: CalamitousFelicitousness Date: Sat, 18 Jul 2026 02:58:11 +0100 Subject: [PATCH] feat(lora): per-layer select stack modes klora and estlora Two-network subject+style sets select a winner per layer instead of summing: scores are top-K magnitude sums (klora) or Frobenius energies (estlora), and a timestep ramp shifts layers from the subject network toward the style network across sampling, reduced to at most one precomputed flip per layer per pass. On sub-8-bit SDNQ the pair rides the side-channel as separate segments flipped in place; other layers recompute the winner from the pristine backup, so select modes force backup mode. Selection resets per pass from the callback setup and is gated off under model compile. estlora's measured style-discrepancy term is exposed as an option. Adds XYZ axes for the stack settings. --- modules/lora/extra_networks_lora.py | 3 + modules/lora/lora_apply.py | 8 +- modules/lora/lora_overrides.py | 3 + modules/lora/lora_sdnq.py | 67 +++++++- modules/lora/lora_stack.py | 81 ++++++++-- modules/lora/networks.py | 33 ++++ modules/processing_callbacks.py | 6 + scripts/xyz/xyz_grid_classes.py | 16 ++ test/test-sdnq-lora-factors.py | 228 ++++++++++++++++++++++++++++ 9 files changed, 426 insertions(+), 19 deletions(-) diff --git a/modules/lora/extra_networks_lora.py b/modules/lora/extra_networks_lora.py index f5cf5647d..c81b94072 100644 --- a/modules/lora/extra_networks_lora.py +++ b/modules/lora/extra_networks_lora.py @@ -225,6 +225,9 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): reason = '' load_method, load_reason = lora_overrides.get_method() + from modules.lora import lora_stack + if load_method != 'native' and lora_stack.mode() != 'sum': + lora_stack.warn_once(f'method-{load_method}', f'Network stack: mode={lora_stack.mode()} method={load_method} unsupported, using sum') if debug: import sys fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access diff --git a/modules/lora/lora_apply.py b/modules/lora/lora_apply.py index 8367d61ce..11b91d9e9 100644 --- a/modules/lora/lora_apply.py +++ b/modules/lora/lora_apply.py @@ -68,7 +68,7 @@ def network_backup_weights(self: torch.nn.Conv2d | torch.nn.Linear | torch.nn.Gr return backup_size -def network_calc_weights(self: torch.nn.Conv2d | torch.nn.Linear | torch.nn.GroupNorm | torch.nn.LayerNorm | diffusers.models.lora.LoRACompatibleLinear | diffusers.models.lora.LoRACompatibleConv, network_layer_name: str, use_previous: bool = False, *, elimit: Callable[[], None] | None = None): +def network_calc_weights(self: torch.nn.Conv2d | torch.nn.Linear | torch.nn.GroupNorm | torch.nn.LayerNorm | diffusers.models.lora.LoRACompatibleLinear | diffusers.models.lora.LoRACompatibleConv, network_layer_name: str, use_previous: bool = False, *, elimit: Callable[[], None] | None = None, per_net: bool = False): if shared.opts.diffusers_offload_mode == "none": try: self.to(devices.device) @@ -77,8 +77,8 @@ def network_calc_weights(self: torch.nn.Conv2d | torch.nn.Linear | torch.nn.Grou batch_updown = None batch_ex_bias = None stack_deltas = None - if lora_stack.mode() in lora_stack.DENSE_MODES and network_layer_name is not None and not network_layer_name.startswith('lora_te'): - stack_deltas = [] # collect per-net deltas; combined after the loop (bias deltas stay summed) + if per_net or (lora_stack.mode() in lora_stack.DENSE_MODES and network_layer_name is not None and not network_layer_name.startswith('lora_te')): + stack_deltas = [] # collect per-net deltas; combined after the loop unless the caller wants them separate (bias deltas stay summed) loaded = l.loaded_networks if not use_previous else l.previously_loaded_networks for net in loaded: module = net.modules.get(network_layer_name, None) @@ -142,6 +142,8 @@ def network_calc_weights(self: torch.nn.Conv2d | torch.nn.Linear | torch.nn.Grou if elimit is not None: elimit() continue + if per_net: + return stack_deltas, batch_ex_bias if stack_deltas is not None and stack_deltas: if len(stack_deltas) >= 2: t0 = time.time() diff --git a/modules/lora/lora_overrides.py b/modules/lora/lora_overrides.py index eb629896e..078b285b2 100644 --- a/modules/lora/lora_overrides.py +++ b/modules/lora/lora_overrides.py @@ -112,6 +112,9 @@ def disable_fuse(): round-trips it through its storage format. On quantized weights that is a dequantize-add-requantize cycle per network swap whose error compounds. """ + from modules.lora import lora_stack + if lora_stack.mode() in lora_stack.SELECT_MODES: + return True # select stack modes flip per-layer winners against the pristine backup sd_model = getattr(shared.sd_model, 'pipe', shared.sd_model) if is_quantized(sd_model): return True diff --git a/modules/lora/lora_sdnq.py b/modules/lora/lora_sdnq.py index 7973121b5..b1ef10d45 100644 --- a/modules/lora/lora_sdnq.py +++ b/modules/lora/lora_sdnq.py @@ -174,6 +174,7 @@ def remove_factors(self): self.svd_up = svd_up self.svd_down = svd_down del self.sdnq_lora_svd_stash + lora_stack.drop(getattr(self, 'network_layer_name', None)) # a selection schedule must not outlive the segments it points into return True @@ -214,11 +215,23 @@ def apply_factors(self, network_layer_name, wanted_names): def append_factors(self, ups, downs): - """Concatenate ``[out, r]`` / ``[r, in]`` factor pairs onto the layer's svd channel and stash the originals.""" + """Concatenate ``[out, r]`` / ``[r, in]`` factor pairs onto the layer's svd channel and stash the originals. + + Returns the appended parts' rank ranges plus the transposed-layout flag; the + checkpoint's own factors occupy the range before the first entry and bucket + padding lands after the last, so the ranges stay valid on the live buffers. + """ deq = self.sdnq_dequantizer device = self.scale.device dtype = deq.result_dtype orig_up, orig_down = self.svd_up, self.svd_down + orig_rank = 0 + if orig_up is not None: + orig_rank = orig_up.shape[0] if deq.use_quantized_matmul else orig_up.shape[1] + segments, offset = [], orig_rank + for u in ups: + segments.append((offset, offset + u.shape[1])) + offset += u.shape[1] if deq.use_quantized_matmul: # matmul layout stores factors transposed: svd_up [r, out], svd_down [in, r] parts_up = ([orig_up.to(device=devices.device, dtype=dtype)] if orig_up is not None else []) + [u.t() for u in ups] @@ -240,6 +253,7 @@ def append_factors(self, ups, downs): self.sdnq_lora_svd_stash = (orig_up, orig_down) self.svd_up = torch.nn.Parameter(new_up.to(device=device), requires_grad=False) self.svd_down = torch.nn.Parameter(new_down.to(device=device), requires_grad=False) + return segments, deq.use_quantized_matmul def host_candidate(self, network_layer_name, wanted_names): @@ -434,6 +448,57 @@ def truncate_delta(self, D, dtype): return up_h, down_h, energy, rms is not None +def apply_select(self, network_layer_name, per_net, wanted_names): + """Attach two networks' contributions as separate side-channel segments for per-layer selection. + + Factorable members ride exactly; the rest host as their own truncated svd + with per-net cache entries. Segment ranges and selection scores register + with ``lora_stack``; the flip schedule executes from the step callback. + Returns None when the pair cannot ride the channel; the caller falls back. + """ + from sdnq.quant_utils import rotate_hadamard + deq = self.sdnq_dequantizer + changed = remove_factors(self) + if wanted_names == (): + return changed + if per_net is None or len(per_net) != 2: + return None + dtype = deq.result_dtype + lora_factor_cache.begin_pass(wanted_names) + pairs, ranks = [], [] + for i, (net_name, D) in enumerate(per_net): + if D is None or D.ndim != 2 or tuple(D.shape) != tuple(deq.original_shape): + return None + net = next((n for n in l.loaded_networks if n.name == net_name), None) + module = net.modules.get(network_layer_name, None) if net is not None else None + if module is None: + return None + ranks.append(int(getattr(module, 'dim', 0) or 0) or min(int(shared.opts.lora_sdnq_host_rank), *deq.original_shape)) + factors = get_module_factors(module, devices.device, dtype, original_shape=deq.original_shape) + if factors is not None: + up_i, down_i = factors + if deq.use_hadamard: + down_i = rotate_hadamard(down_i.to(dtype=torch.float32), group_size=deq.hadamard_group_size).to(dtype=dtype) + else: + key = f'{network_layer_name}#{i}' + cached = lora_factor_cache.fetch(key) + if cached is not None: + up_i, down_i = cached[0].to(device=devices.device, dtype=dtype), cached[1].to(device=devices.device, dtype=dtype) + hosted_layers.append((key, cached[2], cached[3])) + else: + up_i, down_i, energy, calibrated = truncate_delta(self, D.detach().to(devices.device, torch.float32), dtype) + up_i, down_i = lora_factor_cache.store(key, up_i, down_i, energy, calibrated, float(D.detach().float().square().mean().sqrt())) + hosted_layers.append((key, energy, calibrated)) + pairs.append((up_i, down_i)) + d0 = per_net[0][1].detach().to(devices.device, torch.float32) + d1 = per_net[1][1].detach().to(devices.device, torch.float32) + scores, abs_sums = lora_stack.score_pair(d0, d1, ranks[0], ranks[1]) + del d0, d1 + segments, transposed = append_factors(self, [pairs[0][0], pairs[1][0]], [pairs[0][1], pairs[1][1]]) + lora_stack.register(network_layer_name, self, 'factor', scores, segments=(segments[0], segments[1], transposed), abs_sums=abs_sums) + return True + + def note_fallback(self, network_layer_name): """Record a quantized layer taking the requantize path (summary-logged per pass); layers the routing rule sent there are counted apart.""" if getattr(self, 'sdnq_dequantizer', None) is not None and network_layer_name not in routed_layers: diff --git a/modules/lora/lora_stack.py b/modules/lora/lora_stack.py index 224505cab..313038b39 100644 --- a/modules/lora/lora_stack.py +++ b/modules/lora/lora_stack.py @@ -30,7 +30,7 @@ KLORA_BETA = 0.5 # the paper's fixed ramp offset; only the slope is user-tunable ROW_CHUNK = 512 # fp32 interiors run in first-dim slices; also fixes the DARE draw sequence SAMPLE_CAP = 1 << 22 # strided subsample bound for magnitude quantiles (full-size quantile exceeds torch limits) -state: dict = {'entries': {}, 'flips': {}, 'gamma': 1.0, 'total_steps': 0, 'finalized': False} +state: dict = {'entries': {}, 'flips': {}, 'gamma': 1.0, 'gamma_num': 0.0, 'gamma_den': 0.0, 'total_steps': 0, 'finalized': False} warned: set = set() @@ -134,6 +134,45 @@ def combine(named_deltas, layer_name): return result.to(out_dtype) +def score_pair(d0, d1, rank0, rank1): + """Selection scores for a dense delta pair: klora top-K sums (K = rank product) or est energies; plus abs-sums for the global balance.""" + abs_sums = (float(d0.abs().sum()), float(d1.abs().sum())) + if mode() == 'klora': + k = max(1, int(rank0) * int(rank1)) + s0 = float(torch.topk(d0.abs().flatten(), min(k, d0.numel()), sorted=False).values.sum()) + s1 = float(torch.topk(d1.abs().flatten(), min(k, d1.numel()), sorted=False).values.sum()) + else: + s0 = float(d0.float().square().sum()) + s1 = float(d1.float().square().sum()) + return (s0, s1), abs_sums + + +def register_weight_pair(layer_name, module, per_net): + """Score and register a weight-kind selection pair; True when the layer is scheduled.""" + from modules.lora import lora_common as l + if per_net is None or len(per_net) != 2: + return False + ranks, names = [], [] + for net_name, d in per_net: + if d is None: + return False + net = next((n for n in l.loaded_networks if n.name == net_name), None) + net_module = net.modules.get(layer_name, None) if net is not None else None + if net_module is None: + return False + names.append(net_name) + ranks.append(int(getattr(net_module, 'dim', 0) or 0) or 64) + scores, abs_sums = score_pair(per_net[0][1].float(), per_net[1][1].float(), ranks[0], ranks[1]) + register(layer_name, module, 'weight', scores, nets=tuple(names), abs_sums=abs_sums) + return True + + +def drop(layer_name): + """Forget a layer's selection entry (its factors were removed or restored).""" + if layer_name is not None and state['entries'].pop(layer_name, None) is not None: + state['finalized'] = False + + def score_topk(up, down, k): """K-LoRA layer score: sum of the top-K absolute delta entries (one dense materialization).""" d = (up.to(torch.float32) @ down.to(torch.float32)).abs().flatten() @@ -152,22 +191,28 @@ def clear(): state['entries'] = {} state['flips'] = {} state['gamma'] = 1.0 + state['gamma_num'] = 0.0 + state['gamma_den'] = 0.0 state['total_steps'] = 0 state['finalized'] = False -def register(layer_name, module, kind, segments, scores, factors=None): - """Record a select-mode layer: its two segments (or bf16 factor pairs) and static scores. +def register(layer_name, module, kind, scores, segments=None, nets=None, abs_sums=None): + """Record a select-mode layer for schedule finalization. - kind 'factor': segments = [(start, stop), (start, stop)] column ranges in svd_up/svd_down - with the transposed-layout flag appended; stashes both segments' values for flips. - kind 'weight': factors = [(up0, down0), (up1, down1)] kept for recompute-from-backup. + kind 'factor': segments = ((s0, s1), (t0, t1), transposed) column ranges on the svd + channel; both segments' pristine values are stashed for flips. kind 'weight': nets = + the two network names; the winner delta is recomputed from the layer backup at + selection time. abs_sums feeds the global magnitude balance (klora gamma). """ - entry = {'module': weakref.ref(module), 'kind': kind, 'segments': segments, 'scores': scores, 'factors': factors, 'stash': None} + entry = {'layer': layer_name, 'module': weakref.ref(module), 'kind': kind, 'segments': segments, 'scores': scores, 'nets': nets, 'stash': None} if kind == 'factor': (s0, s1), (t0, t1), transposed = segments up = module.svd_up.data entry['stash'] = (segment_view(up, s0, s1, transposed).clone(), segment_view(up, t0, t1, transposed).clone()) + if abs_sums is not None: + state['gamma_num'] += abs_sums[0] + state['gamma_den'] += abs_sums[1] state['entries'][layer_name] = entry state['finalized'] = False @@ -196,6 +241,7 @@ def layer_flip_step(scores, total_steps): def finalize(total_steps): """Build the inverted flip map for the pass; select-mode layers start at their step-0 winner.""" state['total_steps'] = int(total_steps) + state['gamma'] = (state['gamma_num'] / state['gamma_den']) if state['gamma_den'] > 0 else 1.0 state['flips'] = {} for layer_name, entry in state['entries'].items(): flip_at = layer_flip_step(entry['scores'], state['total_steps']) @@ -208,7 +254,7 @@ def finalize(total_steps): def reset(total_steps): """Per-pass reset from set_callbacks_p: restore initial selections and reschedule for this pass's step count.""" - if mode() not in SELECT_MODES or not state['entries']: + if mode() not in SELECT_MODES or not state['entries'] or int(total_steps) <= 0: return finalize(total_steps) @@ -231,20 +277,25 @@ def apply_selection(layer_name, entry, winner): if entry['kind'] == 'factor': (s0, s1), (t0, t1), transposed = entry['segments'] up = module.svd_up.data - keep, drop = ((t0, t1), (s0, s1)) if winner == 1 else ((s0, s1), (t0, t1)) + keep_seg, drop_seg = ((t0, t1), (s0, s1)) if winner == 1 else ((s0, s1), (t0, t1)) stash = entry['stash'][winner] - segment_view(up, keep[0], keep[1], transposed).copy_(stash.to(device=up.device, dtype=up.dtype)) - segment_view(up, drop[0], drop[1], transposed).zero_() + segment_view(up, keep_seg[0], keep_seg[1], transposed).copy_(stash.to(device=up.device, dtype=up.dtype)) + segment_view(up, drop_seg[0], drop_seg[1], transposed).zero_() else: weight_selection(module, entry, winner) def weight_selection(module, entry, winner): + from modules.lora import lora_common as l + from modules.lora.lora_apply import network_apply_weights backup = getattr(module, 'network_weights_backup', None) if not isinstance(backup, torch.Tensor): # fuse mode keeps a bool sentinel, not a pristine copy warn_once('select-nobackup', 'Network stack: select flip skipped, no weight backup') return - up, down = entry['factors'][winner] - weight = backup.to(device=module.weight.device, dtype=torch.float32) - delta = up.to(device=module.weight.device, dtype=torch.float32) @ down.to(device=module.weight.device, dtype=torch.float32) - module.weight.data.copy_((weight + delta.reshape(weight.shape)).to(module.weight.dtype)) + net = next((n for n in l.loaded_networks if n.name == entry['nets'][winner]), None) + net_module = net.modules.get(entry['layer'], None) if net is not None else None + if net_module is None: + return + device = module.weight.device + updown = net_module.calc_updown(backup.to(device))[0] + network_apply_weights(module, updown, None, device=device) # recomputes from the pristine backup, requantizing where the layer needs it diff --git a/modules/lora/networks.py b/modules/lora/networks.py index 58548b82a..59a66a6e0 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -91,6 +91,7 @@ def network_activate(include=None, exclude=None): with devices.inference_context(), pbar: wanted_names = tuple((x.name, x.te_multiplier, x.unet_multiplier, x.dyn_dim) for x in l.loaded_networks) if len(l.loaded_networks) > 0 else () stack_sig = lora_stack.signature() + lora_sdnq.signature() # tracked beside network_current_names so settings-only stack or mechanism changes re-apply + select_active = lora_stack.active_select(len(l.loaded_networks)) applied_layers.clear() lora_sdnq.fallback_layers.clear() # a raise mid-pass leaves stale entries behind lora_sdnq.hosted_layers.clear() @@ -108,6 +109,37 @@ def network_activate(include=None, exclude=None): if group_offload and component not in group_stripped and group_will_mutate(module, network_layer_name, l.loaded_networks): device = group_offload_strip(sd_model, component, group_stripped) calced = False # tracks whether this iteration assembled the delta, so the fallthrough reuses it instead of recomputing + if select_active and component_wanted and not network_layer_name.startswith('lora_te'): + if lora_sdnq.host_candidate(module, network_layer_name, component_wanted): # sub-8-bit SDNQ pairs ride the channel as separate segments + weights_backup = getattr(module, "network_weights_backup", None) + if weights_backup is not None and not isinstance(weights_backup, bool): + network_apply_weights(module, None, None, device=device) + per_net, sel_bias = network_calc_weights(module, network_layer_name, elimit=elimit, per_net=True) + if sel_bias is None: + applied = lora_sdnq.apply_select(module, network_layer_name, per_net, component_wanted) + if applied is not None: + if applied and component_wanted: + applied_layers.append(network_layer_name) + applied_weight += 1 + module.network_current_names = component_wanted + module.network_current_stack = stack_sig + if task is not None: + pbar.update(task, advance=1) + continue + else: # other layers select by recomputing the winner from the pristine backup at schedule time + backup_size += network_backup_weights(module, network_layer_name, component_wanted, fuse) + weights_backup = getattr(module, "network_weights_backup", None) + if weights_backup is not None and not isinstance(weights_backup, bool): + per_net, sel_bias = network_calc_weights(module, network_layer_name, elimit=elimit, per_net=True) + if sel_bias is None and lora_stack.register_weight_pair(network_layer_name, module, per_net): + network_apply_weights(module, None, None, device=device) # pristine until the schedule applies the winner + applied_layers.append(network_layer_name) + applied_weight += 1 + module.network_current_names = component_wanted + module.network_current_stack = stack_sig + if task is not None: + pbar.update(task, advance=1) + continue if lora_sdnq.factor_candidate(module, network_layer_name, component_wanted): weights_backup = getattr(module, "network_weights_backup", None) if weights_backup is not None and not isinstance(weights_backup, bool): @@ -153,6 +185,7 @@ def network_activate(include=None, exclude=None): continue backup_size += network_backup_weights(module, network_layer_name, component_wanted, fuse) if not component_wanted: + lora_stack.drop(network_layer_name) # a restored layer must leave the selection schedule weights_backup = getattr(module, "network_weights_backup", None) if weights_backup is None or isinstance(weights_backup, bool): # fuse mode has no tensor backup, restore stays with network_deactivate if task is not None: diff --git a/modules/processing_callbacks.py b/modules/processing_callbacks.py index 9b8728695..1cc11dde5 100644 --- a/modules/processing_callbacks.py +++ b/modules/processing_callbacks.py @@ -17,6 +17,8 @@ def set_callbacks_p(processing): global p, warned # pylint: disable=global-statement p = processing warned = False + from modules.lora import lora_stack + lora_stack.reset(int(getattr(processing, 'steps', 0) or 0)) # per-pass: restore initial selections and reschedule flips before any step runs def prompt_callback(step, kwargs): @@ -37,6 +39,8 @@ def prompt_callback(step, kwargs): def diffusers_callback_legacy(step: int, timestep: int, latents: torch.FloatTensor | np.ndarray): if p is None: return + from modules.lora import lora_stack + lora_stack.on_step(step) if isinstance(latents, np.ndarray): # latents from Onnx pipelines is ndarray. latents = torch.from_numpy(latents) shared.state.sampling_step = step @@ -56,6 +60,8 @@ def diffusers_callback(pipe, step: int = 0, timestep: int = 0, kwargs: dict | No if kwargs is None: kwargs = {} t0 = time.time() + from modules.lora import lora_stack + lora_stack.on_step(step) if shared.opts.torch_sync: if devices.backend == "ipex": diff --git a/scripts/xyz/xyz_grid_classes.py b/scripts/xyz/xyz_grid_classes.py index 2da29d7b3..198d93c91 100644 --- a/scripts/xyz/xyz_grid_classes.py +++ b/scripts/xyz/xyz_grid_classes.py @@ -104,6 +104,10 @@ class SharedSettingsStackHelper(): todo_ratio = None teacache_thresh = None extra_networks_default_multiplier = None + lora_stack_mode = None + lora_stack_density = None + lora_stack_alpha = None + lora_stack_discrepancy = None disable_apply_metadata = None disable_apply_params = None sdnq_quant_mode = None @@ -136,6 +140,10 @@ class SharedSettingsStackHelper(): self.sd_unet = shared.opts.sd_unet self.sd_text_encoder = shared.opts.sd_text_encoder self.extra_networks_default_multiplier = shared.opts.extra_networks_default_multiplier + self.lora_stack_mode = shared.opts.lora_stack_mode + self.lora_stack_density = shared.opts.lora_stack_density + self.lora_stack_alpha = shared.opts.lora_stack_alpha + self.lora_stack_discrepancy = shared.opts.lora_stack_discrepancy self.teacache_thresh = shared.opts.teacache_thresh self.disable_apply_metadata = shared.opts.disable_apply_metadata self.disable_apply_params = shared.opts.disable_apply_params @@ -148,6 +156,10 @@ class SharedSettingsStackHelper(): shared.opts.data["disable_apply_metadata"] = self.disable_apply_metadata shared.opts.data["disable_apply_params"] = self.disable_apply_params shared.opts.data["extra_networks_default_multiplier"] = self.extra_networks_default_multiplier + shared.opts.data["lora_stack_mode"] = self.lora_stack_mode + shared.opts.data["lora_stack_density"] = self.lora_stack_density + shared.opts.data["lora_stack_alpha"] = self.lora_stack_alpha + shared.opts.data["lora_stack_discrepancy"] = self.lora_stack_discrepancy shared.opts.data["prompt_attention"] = self.prompt_attention shared.opts.data["schedulers_solver_order"] = self.schedulers_solver_order shared.opts.data["schedulers_sigma_adjust"] = self.schedulers_sigma_adjust @@ -205,6 +217,10 @@ axis_options = [ AxisOption("[Prompt] Prompt parser", str, apply_setting("prompt_attention"), choices=lambda: ["native", "compel", "xhinker", "a1111", "fixed"]), AxisOption("[Network] LoRA", str, apply_lora, cost=0.5, choices=list_lora), AxisOption("[Network] LoRA strength", float, apply_lora_strength, cost=0.6), + AxisOption("[Network] LoRA stack mode", str, apply_setting("lora_stack_mode"), cost=0.6, choices=lambda: ["sum", "ties", "dare_ties", "dare_linear", "magnitude_prune", "klora", "estlora"]), + AxisOption("[Network] LoRA stack density", float, apply_setting("lora_stack_density"), cost=0.6), + AxisOption("[Network] LoRA stack ramp", float, apply_setting("lora_stack_alpha"), cost=0.6), + AxisOption("[Network] LoRA stack discrepancy", float, apply_setting("lora_stack_discrepancy"), cost=0.6), AxisOption("[Network] Styles", str, apply_styles, choices=lambda: [s.name for s in shared.prompt_styles.styles.values()]), AxisOption("[Param] Width", int, apply_field("width")), AxisOption("[Param] Height", int, apply_field("height")), diff --git a/test/test-sdnq-lora-factors.py b/test/test-sdnq-lora-factors.py index 1949fd6a3..b1e1881d4 100644 --- a/test/test-sdnq-lora-factors.py +++ b/test/test-sdnq-lora-factors.py @@ -1562,6 +1562,228 @@ def test_sum_mode_keeps_exact_stacking(): return True +CAT_SELECT = category('stack-select') + + +@contextmanager +def select_mode(name, alpha=None, disc=None): + old = {k: getattr(shared.opts, k, None) for k in ('lora_stack_mode', 'lora_stack_alpha', 'lora_stack_discrepancy')} + shared.opts.lora_stack_mode = name + if alpha is not None: + shared.opts.lora_stack_alpha = alpha + if disc is not None: + shared.opts.lora_stack_discrepancy = disc + lora_stack.clear() + lora_stack.warned.clear() + try: + yield + finally: + for k, v in old.items(): + setattr(shared.opts, k, v) + lora_stack.clear() + + +def select_pair(layer, seed0=41, seed1=42, scale1=1.0): + A1, B1, D1 = make_delta(seed=seed0, sigma=1e-2) + A2, B2, D2 = make_delta(seed=seed1, sigma=1e-2) + if scale1 != 1.0: + A2, D2 = A2 * scale1, D2 * scale1 + n1 = make_net('subject', layer, A1, B1) + n2 = make_net('style', layer, A2, B2) + return n1, n2, D1, D2 + + +def test_select_flip_schedule_end_to_end(): + layer = build_layer('uint4') + n1, n2, D1, D2 = select_pair(layer) + with mock_model(lin=layer), select_mode('klora', alpha=1.5): + Wdq0 = dq(layer) + activate(n1, n2) + entry = lora_stack.state['entries'].get('lora_transformer_test') + assert entry is not None and entry['kind'] == 'factor', 'a factorable pair must register factor segments' + assert entry['segments'][0] == (0, 8) and entry['segments'][1] == (8, 16), f'segments {entry["segments"]}' + total = 20 + lora_stack.reset(total) + flips = [s for s, layers in lora_stack.state['flips'].items() for _ in layers] + assert len(flips) <= 1, 'a monotone ramp allows at most one flip per layer' + eff0 = dq(layer) - Wdq0 + winner0 = 0 if rho_of(eff0, D1) > rho_of(eff0, D2) else 1 + for s in range(total): + lora_stack.on_step(s) + eff1 = dq(layer) - Wdq0 + if flips: + assert rho_of(eff1, D2) > 0.99, 'after the flip the style delta must be selected' + assert rho_of(eff0, D1) > 0.99, 'before the flip the subject delta must be selected' + else: + assert rho_of(eff1, [D1, D2][winner0]) > 0.99 + activate() + assert torch.equal(dq(layer), Wdq0), 'removal from an end-of-schedule state must restore bit-exact' + return True + + +def test_select_initial_style_when_ramp_starts_won(): + layer = build_layer('uint4') + n1, n2, _D1, D2 = select_pair(seed0=43, seed1=44, scale1=8.0, layer=layer) # style delta dominates + with mock_model(lin=layer), select_mode('estlora', alpha=1.0, disc=0.5): + Wdq0 = dq(layer) + activate(n1, n2) + lora_stack.reset(20) + eff = dq(layer) - Wdq0 + assert rho_of(eff, D2) > 0.99, 'a layer whose style side wins at step 0 must start style-selected' + return True + + +def test_select_flip_is_inplace_and_shape_stable(): + layer = build_layer('uint4') + n1, n2, _D1, _D2 = select_pair(layer, seed0=45, seed1=46) + with mock_model(lin=layer), select_mode('klora'): + activate(n1, n2) + param_id = id(layer.svd_up) + shape = tuple(layer.svd_up.shape) + lora_stack.reset(20) + entry = lora_stack.state['entries']['lora_transformer_test'] + (s0, s1), (t0, t1), transposed = entry['segments'] + zeroed = lora_stack.segment_view(layer.svd_up.data, t0, t1, transposed) + kept = lora_stack.segment_view(layer.svd_up.data, s0, s1, transposed) + assert float(zeroed.abs().sum()) == 0.0 or float(kept.abs().sum()) == 0.0, 'exactly one segment must be zeroed initially' + for s in range(20): + lora_stack.on_step(s) + assert id(layer.svd_up) == param_id and tuple(layer.svd_up.shape) == shape, 'flips must mutate in place, never reassign' + return True + + +def test_select_matmul_transposed_layout(): + layer = build_layer('uint4', use_quantized_matmul=True) + n1, n2, D1, D2 = select_pair(layer, seed0=47, seed1=48) + with mock_model(lin=layer), select_mode('klora'): + Wdq0 = dq(layer) + activate(n1, n2) + entry = lora_stack.state['entries']['lora_transformer_test'] + assert entry['segments'][2] is True, 'quantized-matmul layout must register as transposed' + lora_stack.reset(20) + eff = dq(layer) - Wdq0 + assert max(rho_of(eff, D1), rho_of(eff, D2)) > 0.99, 'initial selection must realize one delta exactly' + activate() + assert torch.equal(dq(layer), Wdq0) + return True + + +def test_select_per_net_hosted_pair(): + layer = build_layer('uint4') + torch.manual_seed(49) + Dd1 = (torch.randn(OUT_F, 24, device=DEVICE) @ torch.randn(24, IN_F, device=DEVICE)) * 1e-3 # rank inside the host cap so truncation is near-lossless + Dd2 = (torch.randn(OUT_F, 24, device=DEVICE) @ torch.randn(24, IN_F, device=DEVICE)) * 1e-3 + n1 = make_dense_net('lk1', layer, Dd1) + n2 = make_dense_net('lk2', layer, Dd2) + with host_rank(32), mock_model(lin=layer), select_mode('klora'): + Wdq0 = dq(layer) + activate(n1, n2) + entry = lora_stack.state['entries'].get('lora_transformer_test') + assert entry is not None, 'non-factorable pairs must register through per-net hosting' + assert entry['segments'][0] == (0, 24) and entry['segments'][1] == (24, 48), f'segments {entry["segments"]}' # hosting stores the effective rank (24), not the cap + lora_stack.reset(20) + eff = dq(layer) - Wdq0 + best = max(rho_of(eff, Dd1), rho_of(eff, Dd2)) + assert best > 0.9, f'initial selection must realize one hosted delta, rho={best:.3f}' + activate() + assert torch.equal(dq(layer), Wdq0) + return True + + +def test_select_reset_restores_initial_state(): + layer = build_layer('uint4') + n1, n2, _D1, _D2 = select_pair(layer, seed0=51, seed1=52) + with mock_model(lin=layer), select_mode('klora'): + activate(n1, n2) + lora_stack.reset(20) + initial = dq(layer) + for s in range(20): + lora_stack.on_step(s) + lora_stack.reset(20) + assert torch.equal(dq(layer), initial), 'a fresh pass must restore the initial selection without re-activation' + return True + + +def test_select_deactivate_from_midflip(): + layer = build_layer('uint4') + n1, n2, _D1, _D2 = select_pair(layer, seed0=53, seed1=54) + with mock_model(lin=layer), select_mode('klora'): + Wdq0 = dq(layer) + activate(n1, n2) + lora_stack.reset(20) + for s in range(10): + lora_stack.on_step(s) + activate() + assert torch.equal(dq(layer), Wdq0), 'removal mid-schedule must restore bit-exact' + assert not lora_stack.state['entries'], 'removal must drop the selection entry' + return True + + +def test_select_requires_exactly_two_nets(): + layer = build_layer('uint4') + A3, B3, _D3 = make_delta(seed=55) + n1, n2, _D1, _D2 = select_pair(layer, seed0=56, seed1=57) + n3 = make_net('third', layer, A3, B3) + with mock_model(lin=layer), select_mode('klora'): + activate(n1, n2, n3) + assert hasattr(layer, 'sdnq_lora_svd_stash'), 'three nets must fall back to the exact concat path' + assert not lora_stack.state['entries'], 'no selection entries outside the two-net case' + activate() + return True + + +def test_select_gated_off_when_compiled(): + layer = build_layer('uint4') + n1, n2, _D1, _D2 = select_pair(layer, seed0=58, seed1=59) + old_compile = getattr(shared.opts, 'cuda_compile', None) + try: + shared.opts.cuda_compile = ['Model'] + with mock_model(lin=layer), select_mode('klora'): + activate(n1, n2) + assert not lora_stack.state['entries'], 'select must gate off under model compile' + assert hasattr(layer, 'sdnq_lora_svd_stash'), 'gated select behaves as sum' + activate() + finally: + shared.opts.cuda_compile = old_compile + return True + + +def test_est_energy_matches_full_frobenius(): + torch.manual_seed(60) + up = torch.randn(64, 8, device=DEVICE) + down = torch.randn(8, 96, device=DEVICE) + gram = lora_stack.score_energy(up, down) + full = float((up @ down).square().sum()) + assert abs(gram - full) / full < 1e-5, f'{gram} vs {full}' + return True + + +def test_select_weight_kind_plain_layer(): + lin = torch.nn.Linear(IN_F, OUT_F, bias=False, dtype=torch.bfloat16, device=DEVICE) + with torch.no_grad(): + lin.weight.copy_(torch.randn(OUT_F, IN_F, device=DEVICE) * 0.02) + lin.network_layer_name = 'lora_transformer_plain' + lin.network_current_names = () + A1, B1, D1 = make_delta(seed=61, sigma=1e-2) + A2, B2, D2 = make_delta(seed=62, sigma=1e-2) + n1 = make_net('w1', lin, A1, B1) + n2 = make_net('w2', lin, A2, B2) + W0 = lin.weight.detach().float().clone() + with mock_model(lin=lin), select_mode('klora'): + activate(n1, n2) + entry = lora_stack.state['entries'].get('lora_transformer_plain') + assert entry is not None and entry['kind'] == 'weight', 'plain layers must register weight-kind selection' + assert torch.equal(lin.weight.detach().float(), W0), 'weights stay pristine until the schedule applies a winner' + lora_stack.reset(20) + eff = lin.weight.detach().float() - W0 + assert max(rho_of(eff, D1), rho_of(eff, D2)) > 0.95, 'initial selection must apply one delta from backup' + for s in range(20): + lora_stack.on_step(s) + activate() + assert torch.equal(lin.weight.detach().float(), W0), 'restore-only pass must return the pristine weight' + return True + + CAT_COMPILE = category('compile') @@ -1746,6 +1968,12 @@ def run_tests(): test_magnitude_prune_keeps_top_density, test_dense_two_plain_loras_hosted_not_summed, 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 ===') + for fn in [test_select_flip_schedule_end_to_end, test_select_initial_style_when_ramp_starts_won, test_select_flip_is_inplace_and_shape_stable, + test_select_matmul_transposed_layout, test_select_per_net_hosted_pair, test_select_reset_restores_initial_state, + test_select_deactivate_from_midflip, test_select_requires_exactly_two_nets, test_select_gated_off_when_compiled, + test_est_energy_matches_full_frobenius, test_select_weight_kind_plain_layer]: + run_test(CAT_SELECT, fn) log.warning('=== Compile ===') for fn in [test_factor_add_inside_compiled_graph, test_rank_bucket_graph_reuse]: run_test(CAT_COMPILE, fn)