diff --git a/modules/lora/networks.py b/modules/lora/networks.py index 38b6f979b..da606a16c 100644 --- a/modules/lora/networks.py +++ b/modules/lora/networks.py @@ -19,6 +19,64 @@ default_components = ['text_encoder', 'text_encoder_2', 'text_encoder_3', 'text_ deactivate_components = ['text_encoder', 'text_encoder_2', 'text_encoder_3', 'unet', 'transformer', 'llm_adapter'] +class ActivationPass: + """State of one activation walk, built before the walk so a pass that raises still has it. + + `wanted_names` is built once here and reaches every layer as + `component_wanted`, either this tuple or the empty one. The factor cache + keys its pass entry on that object's identity, so an equal tuple rebuilt + per component would send every lookup back to disk. + """ + + def __init__(self, fuse): + self.sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) + self.fuse = fuse + self.elimit = None # the error limiter, bound for the duration of the walk + self.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 () + self.stack_sig = lora_stack.signature() + lora_blocks.signature() + lora_sdnq.signature() # tracked beside network_current_names so stack-setting, block-weight and mechanism changes re-apply + self.select_active = len(l.loaded_networks) > 0 and lora_stack.active_select(len(l.loaded_networks)) # restore-only walks have nothing to stack; the count warning would fire on every network-free generation + self.component_wanted = () + self.device = None + self.group_offload = shared.opts.diffusers_offload_mode == "group" + self.group_stripped = {} + self.pbar = nullcontext() + self.task = None + self.total = 0 + self.active_components = [] + self.applied_weight = 0 + self.applied_bias = 0 + self.refused = 0 + self.backup_size = 0 + + def stamp(self, module): + """Mark the layer as carrying this set under these settings; the pair is the skip key.""" + module.network_current_names = self.component_wanted + module.network_current_stack = self.stack_sig + + def tick(self, description=None): + if self.task is None: + return + if description is None: + self.pbar.update(self.task, advance=1) + else: + self.pbar.update(self.task, advance=1, description=description) + + def claim(self, module, network_layer_name, changed): + """Accept a layer one of the mechanisms took; only a layer whose weights changed counts as applied.""" + if changed and self.component_wanted: + applied_layers.append(network_layer_name) + self.applied_weight += 1 + self.stamp(module) + self.tick() + + def keep_selected(self, module, network_layer_name, sel_backup): + """Hold a scheduled weight-kind layer on its pristine tensor until the schedule applies the winner.""" + self.backup_size += sel_backup # counted only where this branch keeps the layer; the weight path below re-enters the shared backup call, which counts it then + network_apply_weights(module, None, None, device=self.device) + self.claim(module, network_layer_name, True) + return True + + def group_will_mutate(module, network_layer_name: str, loaded) -> bool: """True when the pass will write to this module: a loaded network covers its layer, a tensor backup awaits restore, or an svd factor stash awaits removal.""" @@ -116,6 +174,108 @@ def should_skip(module, network_layer_name, wanted, stack_sig): return getattr(module, 'network_current_names', ()) == wanted and getattr(module, 'network_current_stack', 'sum') == stack_sig +def try_select(ctx, module, network_layer_name): + """Put the layer under a selection schedule; True when it took the layer. + + The three arms are mutually exclusive and their warnings are keyed, so a + layer that cannot be scheduled reports one reason and falls through. + """ + if not ctx.select_active or not ctx.component_wanted or network_layer_name.startswith('lora_te'): + return False + if lora_sdnq.select_candidate(module, network_layer_name, ctx.component_wanted): # SDNQ pairs ride the channel as separate segments at any bit width; weight rewrites cannot flip a quantized layer + restore_pristine(module, ctx.device) + applied = lora_sdnq.apply_select_cached(module, network_layer_name, ctx.component_wanted) # a stored score record and factor pair serve before the deltas are assembled + if applied is None: + per_net, sel_bias = network_calc_weights(module, network_layer_name, elimit=ctx.elimit, per_net=True) + if sel_bias is None: + applied = lora_sdnq.apply_select(module, network_layer_name, per_net, ctx.component_wanted) + if applied is not None: + ctx.claim(module, network_layer_name, applied) + return True + lora_stack.warn_once('select-unridable', f'Network stack: mode={lora_stack.mode()} layer="{network_layer_name}" fallback=sum') # a pair the channel cannot carry (bias delta or malformed member) sums like any unsupported set + elif getattr(module, 'sdnq_dequantizer', None) is not None: # hosting disabled: quantized layers have no side-channel to carry segments and packed backups cannot flip, so the sum paths below take the layer + if any(net.modules.get(network_layer_name, None) is not None for net in l.loaded_networks): + lora_stack.warn_once('select-host-disabled', f'Network stack: mode={lora_stack.mode()} quant=sdnq host=disabled fallback=sum') + else: # other layers select by recomputing the winner from the pristine backup at schedule time + sel_backup = network_backup_weights(module, network_layer_name, ctx.component_wanted, ctx.fuse) + if tensor_backup(module) is not None: # a flip recomputes the winner from the pristine tensor, which fuse mode does not keep + if lora_stack.register_weight_pair_cached(network_layer_name, module, ctx.component_wanted): # a stored score record registers without assembling the pair + return ctx.keep_selected(module, network_layer_name, sel_backup) + per_net, sel_bias = network_calc_weights(module, network_layer_name, elimit=ctx.elimit, per_net=True) + if sel_bias is None and lora_stack.register_weight_pair(network_layer_name, module, per_net, ctx.component_wanted): + return ctx.keep_selected(module, network_layer_name, sel_backup) + return False + + +def try_factors(ctx, module, network_layer_name): + """Attach the set to the quantized side channel as exact factors; True when it took the layer.""" + if not lora_sdnq.factor_candidate(module, network_layer_name, ctx.component_wanted): + return False + restore_pristine(module, ctx.device) # an earlier non-factorable set may have requantized this layer + applied = lora_sdnq.apply_factors(module, network_layer_name, ctx.component_wanted) + if applied is None: # the exact path declined; hosting or the weight path takes the layer + return False + ctx.claim(module, network_layer_name, applied) + return True + + +def try_hosted(ctx, module, network_layer_name): + """Host the combined delta on the side channel as truncated factors. + + Returns whether it took the layer and, when it declined after assembling + the delta, that delta, so the weight path applies it without a second + calc. A returned pair of Nones still counts as assembled. + """ + if not lora_sdnq.host_candidate(module, network_layer_name, ctx.component_wanted): + return False, None + restore_pristine(module, ctx.device) # the hosted delta is measured against the pristine base + batch = None + hosted = lora_sdnq.apply_cached(module, network_layer_name, ctx.component_wanted) # a stored entry serves the layer before the delta is assembled + if hosted is None: + batch_updown, batch_ex_bias = network_calc_weights(module, network_layer_name, elimit=ctx.elimit) + batch = (batch_updown, batch_ex_bias) + if batch_ex_bias is None: # bias deltas need the plain path; weight-only sets ride the side-channel without a weight backup + hosted = lora_sdnq.apply_hosted(module, network_layer_name, batch_updown, ctx.component_wanted) + if hosted is not None: + batch = None # hosting took the delta + if hosted is None: + return False, batch + ctx.claim(module, network_layer_name, hosted) + return True, None + + +def apply_generic(ctx, module, network_layer_name, batch): + """The weight path, which takes any layer the mechanisms above declined.""" + stripped = lora_sdnq.remove_factors(module) # the mechanism gate can decline a layer still carrying attached factors; the weight path must start from the pristine channel + if stripped and not ctx.component_wanted: # factor-mode layers have no tensor backup, dropping the factors is the whole restore + ctx.stamp(module) + ctx.tick() + return + ctx.backup_size += network_backup_weights(module, network_layer_name, ctx.component_wanted, ctx.fuse) + if not ctx.component_wanted: + lora_stack.drop(network_layer_name) # a restored layer must leave the selection schedule + if tensor_backup(module) is None: # fuse mode has no tensor backup, restore stays with network_deactivate + ctx.tick() + return + batch_updown, batch_ex_bias = None, None # restore-only pass, apply with no weights reverts to backup + else: + batch_updown, batch_ex_bias = batch if batch is not None else network_calc_weights(module, network_layer_name, elimit=ctx.elimit) + if batch_updown is not None: + lora_sdnq.note_fallback(module, network_layer_name) # only layers whose quantized weight actually takes a delta + if ctx.fuse: + weight_written, bias_written = network_apply_direct(module, batch_updown, batch_ex_bias, device=ctx.device) + else: + weight_written, bias_written = network_apply_weights(module, batch_updown, batch_ex_bias, device=ctx.device) + if batch_updown is not None or batch_ex_bias is not None: + applied_layers.append(network_layer_name) + ctx.applied_weight += 1 if weight_written else 0 + ctx.applied_bias += 1 if bias_written else 0 + ctx.refused += (batch_updown is not None and not weight_written) + (batch_ex_bias is not None and not bias_written) # a delta the module would not take leaves that layer on its base value + ctx.stamp(module) + bs = round(ctx.backup_size/1024/1024/1024, 2) if ctx.backup_size > 0 else None + ctx.tick(f'networks={len(l.loaded_networks)} modules={ctx.active_components} layers={ctx.total} weights={ctx.applied_weight} bias={ctx.applied_bias} backup={bs} device={ctx.device}') + + def network_activate(include=None, exclude=None): if exclude is None: exclude = [] @@ -123,169 +283,52 @@ def network_activate(include=None, exclude=None): include = [] promote_pending() t0 = time.time() - fuse = lora_overrides.fuse_native() # resolve once: backup and apply passes must agree + ctx = ActivationPass(lora_overrides.fuse_native()) # fuse resolved once: the backup, apply and restore paths must agree with limit_errors("network_activate") as elimit: - sd_model = prepare_model_for_write(getattr(shared.sd_model, "pipe", shared.sd_model)) - group_offload = shared.opts.diffusers_offload_mode == "group" - group_stripped = {} - device = None - modules, components, active_components, total = collect_components(sd_model, include, exclude, default_components, restore_filtered=True) - pbar, task = pass_progress('activate', total, len(l.loaded_networks) > 0) - applied_weight = 0 - applied_bias = 0 - refused = 0 - 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_blocks.signature() + lora_sdnq.signature() # tracked beside network_current_names so stack-setting, block-weight and mechanism changes re-apply - select_active = len(l.loaded_networks) > 0 and lora_stack.active_select(len(l.loaded_networks)) # restore-only walks have nothing to stack; the count warning would fire on every network-free generation + ctx.elimit = elimit + ctx.sd_model = prepare_model_for_write(ctx.sd_model) + modules, components, ctx.active_components, ctx.total = collect_components(ctx.sd_model, include, exclude, default_components, restore_filtered=True) + ctx.pbar, ctx.task = pass_progress('activate', ctx.total, len(l.loaded_networks) > 0) + with devices.inference_context(), ctx.pbar: applied_layers.clear() lora_sdnq.fallback_layers.clear() # a raise mid-pass leaves stale entries behind lora_sdnq.hosted_layers.clear() lora_sdnq.factor_layers.clear() lora_sdnq.select_layers.clear() - backup_size = 0 for component in modules.keys(): - component_wanted = wanted_names if component in components else () - device = getattr(sd_model, component, None).device + ctx.component_wanted = ctx.wanted_names if component in components else () # the pass tuple itself, never a copy + ctx.device = getattr(ctx.sd_model, component, None).device for _, module in modules[component]: network_layer_name = getattr(module, 'network_layer_name', None) - if should_skip(module, network_layer_name, component_wanted, stack_sig): - if task is not None: - pbar.update(task, advance=1) + if should_skip(module, network_layer_name, ctx.component_wanted, ctx.stack_sig): + ctx.tick() continue lora_stack.drop(network_layer_name) # re-application invalidates any live selection schedule; the select branch re-registers - 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.select_candidate(module, network_layer_name, component_wanted): # SDNQ pairs ride the channel as separate segments at any bit width; weight rewrites cannot flip a quantized layer - restore_pristine(module, device) - applied = lora_sdnq.apply_select_cached(module, network_layer_name, component_wanted) # a stored score record and factor pair serve before the deltas are assembled - if applied is None: - 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 - lora_stack.warn_once('select-unridable', f'Network stack: mode={lora_stack.mode()} layer="{network_layer_name}" fallback=sum') # a pair the channel cannot carry (bias delta or malformed member) sums like any unsupported set - elif getattr(module, 'sdnq_dequantizer', None) is not None: # hosting disabled: quantized layers have no side-channel to carry segments and packed backups cannot flip, so the sum paths below take the layer - if any(net.modules.get(network_layer_name, None) is not None for net in l.loaded_networks): - lora_stack.warn_once('select-host-disabled', f'Network stack: mode={lora_stack.mode()} quant=sdnq host=disabled fallback=sum') - else: # other layers select by recomputing the winner from the pristine backup at schedule time - sel_backup = network_backup_weights(module, network_layer_name, component_wanted, fuse) - if tensor_backup(module) is not None: # a flip recomputes the winner from the pristine tensor, which fuse mode does not keep - if lora_stack.register_weight_pair_cached(network_layer_name, module, component_wanted): # a stored score record registers without assembling the pair - backup_size += sel_backup - 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 - 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, component_wanted): - backup_size += sel_backup # counted only when this branch keeps the layer; the fallthrough re-enters the shared backup call below, which counts it then - 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): - restore_pristine(module, device) # an earlier non-factorable set may have requantized this layer - applied = lora_sdnq.apply_factors(module, network_layer_name, component_wanted) - if applied is not None: # exact path took the layer; None falls through to hosting or requantize - 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 - if lora_sdnq.host_candidate(module, network_layer_name, component_wanted): - restore_pristine(module, device) # the hosted delta is measured against the pristine base - hosted = lora_sdnq.apply_cached(module, network_layer_name, component_wanted) # a stored entry serves the layer before the delta is assembled - if hosted is None: - batch_updown, batch_ex_bias = network_calc_weights(module, network_layer_name, elimit=elimit) - calced = True - if batch_ex_bias is None: # bias deltas need the plain path; weight-only sets ride the side-channel without a weight backup - hosted = lora_sdnq.apply_hosted(module, network_layer_name, batch_updown, component_wanted) - if hosted is not None: - batch_updown, batch_ex_bias = None, None - del batch_updown, batch_ex_bias - if hosted is not None: - if hosted 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 - stripped = lora_sdnq.remove_factors(module) # the mechanism gate can decline a layer still carrying attached factors; the weight path must start from the pristine channel - if stripped and not component_wanted: # factor-mode layers have no tensor backup, dropping the factors is the whole restore - module.network_current_names = () - module.network_current_stack = stack_sig - if task is not None: - pbar.update(task, advance=1) + if ctx.group_offload and component not in ctx.group_stripped and group_will_mutate(module, network_layer_name, l.loaded_networks): + ctx.device = group_offload_strip(ctx.sd_model, component, ctx.group_stripped) + if try_select(ctx, module, network_layer_name): 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 - if tensor_backup(module) is None: # fuse mode has no tensor backup, restore stays with network_deactivate - if task is not None: - pbar.update(task, advance=1) - continue - batch_updown, batch_ex_bias = None, None # restore-only pass, apply with no weights reverts to backup - else: - if not calced: # the host branch may have assembled the delta already; a declined layer reuses it - batch_updown, batch_ex_bias = network_calc_weights(module, network_layer_name, elimit=elimit) - if batch_updown is not None: - lora_sdnq.note_fallback(module, network_layer_name) # only layers whose quantized weight actually takes a delta - if fuse: - weight_written, bias_written = network_apply_direct(module, batch_updown, batch_ex_bias, device=device) - else: - weight_written, bias_written = network_apply_weights(module, batch_updown, batch_ex_bias, device=device) - if batch_updown is not None or batch_ex_bias is not None: - applied_layers.append(network_layer_name) - applied_weight += 1 if weight_written else 0 - applied_bias += 1 if bias_written else 0 - refused += (batch_updown is not None and not weight_written) + (batch_ex_bias is not None and not bias_written) # a delta the module would not take leaves that layer on its base value - batch_updown, batch_ex_bias = None, None - del batch_updown, batch_ex_bias - module.network_current_names = component_wanted - module.network_current_stack = stack_sig - if task is not None: - bs = round(backup_size/1024/1024/1024, 2) if backup_size > 0 else None - pbar.update(task, advance=1, description=f'networks={len(l.loaded_networks)} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} backup={bs} device={device}') - - if task is not None and len(applied_layers) == 0: - pbar.remove_task(task) # hide progress bar for no action + if try_factors(ctx, module, network_layer_name): + continue + hosted, batch = try_hosted(ctx, module, network_layer_name) + if hosted: + continue + apply_generic(ctx, module, network_layer_name, batch) + if ctx.task is not None and len(applied_layers) == 0: + ctx.pbar.remove_task(ctx.task) # hide progress bar for no action global native_active, refused_writes # pylint: disable=global-statement lora_sdnq.report_fallbacks() native_active = len(l.loaded_networks) > 0 - refused_writes = refused - l.last_backup_size = backup_size + refused_writes = ctx.refused + l.last_backup_size = ctx.backup_size l.timer.activate += time.time() - t0 - if refused > 0: - log.error(f'Network load: type=LoRA networks={[n.name for n in l.loaded_networks]} weights={applied_weight} bias={applied_bias} refused={refused} network partially applied') + if ctx.refused > 0: + log.error(f'Network load: type=LoRA networks={[n.name for n in l.loaded_networks]} weights={ctx.applied_weight} bias={ctx.applied_bias} refused={ctx.refused} network partially applied') if l.debug and len(l.loaded_networks) > 0: - log.debug(f'Network load: type=LoRA networks={[n.name for n in l.loaded_networks]} modules={active_components} layers={total} weights={applied_weight} bias={applied_bias} refused={refused} backup={round(backup_size/1024/1024/1024, 2)} fuse={fuse}:{shared.opts.lora_fuse_diffusers} device={device} time={l.timer.summary}') + log.debug(f'Network load: type=LoRA networks={[n.name for n in l.loaded_networks]} modules={ctx.active_components} layers={ctx.total} weights={ctx.applied_weight} bias={ctx.applied_bias} refused={ctx.refused} backup={round(ctx.backup_size/1024/1024/1024, 2)} fuse={ctx.fuse}:{shared.opts.lora_fuse_diffusers} device={ctx.device} time={l.timer.summary}') modules.clear() - if len(applied_layers) > 0 or shared.opts.diffusers_offload_mode == "sequential" or len(group_stripped) > 0: - sd_models.set_diffuser_offload(sd_model, op="model") + if len(applied_layers) > 0 or shared.opts.diffusers_offload_mode == "sequential" or len(ctx.group_stripped) > 0: + sd_models.set_diffuser_offload(ctx.sd_model, op="model") def effective_mode(): diff --git a/test/test-sdnq-lora-factors.py b/test/test-sdnq-lora-factors.py index 191429890..dbfd89482 100644 --- a/test/test-sdnq-lora-factors.py +++ b/test/test-sdnq-lora-factors.py @@ -900,6 +900,41 @@ def test_route_fat_dense_delta_requantizes(): return True +def test_declined_host_delta_is_not_recomputed(): + layer = build_layer('uint4') + torch.manual_seed(21) + D = torch.randn(OUT_F, IN_F, device=DEVICE) * 1e-2 # fat and full-rank: hosting assembles the delta and then routes it to the grid + net = make_dense_net('recompute', layer, D) + with host_rank(256), mock_model(lin=layer), counting_calc() as calls: + activate(net) + assert not hasattr(layer, 'sdnq_lora_svd_stash'), 'the fixture must reach the weight path, not the side channel' + assert calls['n'] == 1, f'a declined host hands its delta on instead of assembling it twice, got {calls["n"]}' + return True + + +def test_pass_presents_one_wanted_names_tuple(): + from modules.lora import lora_factor_cache as fc + layer = build_layer('uint4') + _A, _B, D = make_delta(sigma=3e-3) + net = make_dense_net('identity', layer, D) + seen = [] # holds the objects, so a freed tuple cannot lend its address to the next one + real = fc.begin_pass + + def recording(wanted_names): + seen.append(wanted_names) + return real(wanted_names) + + fc.begin_pass = recording + try: + with host_rank(64), mock_model(lin=layer): + activate(net) + finally: + fc.begin_pass = real + assert len(seen) >= 2, f'the hosted path must consult the cache more than once for this to prove anything, got {len(seen)}' + assert all(x is seen[0] for x in seen), 'one walk must present one tuple: the cache memoizes its entry on identity, and an equal rebuild rereads it from disk' + return True + + def test_route_rule_terms_gate_both_ways(): layer = build_layer('uint4') torch.manual_seed(23) @@ -2914,7 +2949,8 @@ def run_tests(): log.warning('=== Hosting ===') 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_fat_dense_delta_requantizes, test_declined_host_delta_is_not_recomputed, test_pass_presents_one_wanted_names_tuple, + 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_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)