Merge pull request #5070 from vladmandic/refactor/lora-activate-breakdown

refactor(lora): break down network_activate
This commit is contained in:
Vladimir Mandic
2026-08-30 11:51:42 +02:00
committed by GitHub
5 changed files with 443 additions and 265 deletions
+2
View File
@@ -19,3 +19,5 @@ module_types = [
loaded_networks: list = [] # no type due to circular import
previously_loaded_networks: list = [] # no type due to circular import
extra_network_lora = None # initialized in extra_networks.py
last_backup_size: int = 0 # bytes of weight backups the last activate pass held
last_mode: str = '' # how that pass left the weights: backup, fuse or factor
+19 -4
View File
@@ -47,7 +47,7 @@ a low-rank delta hosts exactly however fat it is.
import torch
from modules import devices, shared
from modules.lora import lora_calib, lora_factor_cache, lora_stack
from modules.lora import lora_calib, lora_factor_cache, lora_stack # lora_calib registers its model-load hook on import, so this one has to stay eager
from modules.lora import lora_common as l
from modules.logger import log
@@ -259,8 +259,8 @@ def append_factors(self, ups, downs):
return segments, deq.use_quantized_matmul
def select_candidate(self, network_layer_name, wanted_names):
"""True when this layer can carry a set on the svd channel; select pairs ride it at any bit width."""
def channel_candidate(self, network_layer_name, wanted_names):
"""True when this layer can carry a set on the svd channel: quantized, covered, and given a rank to spend."""
if not enabled():
return False
if int(getattr(shared.opts, 'lora_sdnq_host_rank', 0) or 0) <= 0:
@@ -272,9 +272,14 @@ def select_candidate(self, network_layer_name, wanted_names):
return any(net.modules.get(network_layer_name, None) is not None for net in l.loaded_networks)
def select_candidate(self, network_layer_name, wanted_names):
"""True when a select pair can ride this layer's svd channel; pairs ride it at any bit width."""
return channel_candidate(self, network_layer_name, wanted_names)
def host_candidate(self, network_layer_name, wanted_names):
"""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):
if not channel_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:
@@ -564,6 +569,16 @@ def note_fallback(self, network_layer_name):
fallback_layers.append(network_layer_name)
def reset_pass():
"""Clear every per-pass accumulator, so a pass that raised leaves nothing behind for the next one."""
fallback_layers.clear()
hosted_layers.clear()
hosted_ranks.clear()
factor_layers.clear()
select_layers.clear()
routed_layers.clear() # note_fallback reads this to suppress double counting, so a stale entry silences a real fallback
def report_fallbacks():
hits, misses = lora_factor_cache.flush()
if hits > 0 or misses > 0:
+9 -36
View File
@@ -32,7 +32,8 @@ class NetworkModuleLokr(network.NetworkModule): # pylint: disable=abstract-metho
self.dim = self.w2b.shape[0] if self.w2b is not None else self.dim
self.t2 = weights.w.get("lokr_t2")
def calc_updown(self, target):
def rebuild_operands(self, target):
"""The two Kronecker operands on the target's device and dtype, each either stored whole or rebuilt from its factors."""
if self.w1 is not None:
w1 = self.w1.to(target.device, dtype=target.dtype)
else:
@@ -50,8 +51,12 @@ class NetworkModuleLokr(network.NetworkModule): # pylint: disable=abstract-metho
w2a = self.w2a.to(target.device, dtype=target.dtype)
w2b = self.w2b.to(target.device, dtype=target.dtype)
w2 = lyco_helpers.make_weight_cp(t2, w2a, w2b)
return w1, w2
def calc_updown(self, target):
w1, w2 = self.rebuild_operands(target)
output_shape = [w1.size(0) * w2.size(0), w1.size(1) * w2.size(1)]
if len(target.shape) == 4:
if len(target.shape) == 4: # a conv target keeps its own shape; the chunk variants below only ever address 2-d fused weights
output_shape = target.shape
updown = make_kron(output_shape, w1, w2)
return self.finalize_updown(updown, target, output_shape)
@@ -70,23 +75,7 @@ class NetworkModuleLokrChunk(NetworkModuleLokr):
self.num_chunks = num_chunks
def calc_updown(self, target):
if self.w1 is not None:
w1 = self.w1.to(target.device, dtype=target.dtype)
else:
w1a = self.w1a.to(target.device, dtype=target.dtype)
w1b = self.w1b.to(target.device, dtype=target.dtype)
w1 = w1a @ w1b
if self.w2 is not None:
w2 = self.w2.to(target.device, dtype=target.dtype)
elif self.t2 is None:
w2a = self.w2a.to(target.device, dtype=target.dtype)
w2b = self.w2b.to(target.device, dtype=target.dtype)
w2 = w2a @ w2b
else:
t2 = self.t2.to(target.device, dtype=target.dtype)
w2a = self.w2a.to(target.device, dtype=target.dtype)
w2b = self.w2b.to(target.device, dtype=target.dtype)
w2 = lyco_helpers.make_weight_cp(t2, w2a, w2b)
w1, w2 = self.rebuild_operands(target)
full_shape = [w1.size(0) * w2.size(0), w1.size(1) * w2.size(1)]
updown = make_kron(full_shape, w1, w2)
updown = torch.chunk(updown, self.num_chunks, dim=0)[self.chunk_index]
@@ -109,23 +98,7 @@ class NetworkModuleLokrSliceChunk(NetworkModuleLokr):
self.end_row = end_row
def calc_updown(self, target):
if self.w1 is not None:
w1 = self.w1.to(target.device, dtype=target.dtype)
else:
w1a = self.w1a.to(target.device, dtype=target.dtype)
w1b = self.w1b.to(target.device, dtype=target.dtype)
w1 = w1a @ w1b
if self.w2 is not None:
w2 = self.w2.to(target.device, dtype=target.dtype)
elif self.t2 is None:
w2a = self.w2a.to(target.device, dtype=target.dtype)
w2b = self.w2b.to(target.device, dtype=target.dtype)
w2 = w2a @ w2b
else:
t2 = self.t2.to(target.device, dtype=target.dtype)
w2a = self.w2a.to(target.device, dtype=target.dtype)
w2b = self.w2b.to(target.device, dtype=target.dtype)
w2 = lyco_helpers.make_weight_cp(t2, w2a, w2b)
w1, w2 = self.rebuild_operands(target)
full_shape = [w1.size(0) * w2.size(0), w1.size(1) * w2.size(1)]
updown = make_kron(full_shape, w1, w2)
updown = updown[self.start_row:self.end_row]
+335 -223
View File
@@ -1,3 +1,44 @@
"""Applies the loaded networks to the model and takes them off again.
One walk visits every module of every component and offers each layer to
the mechanisms in a fixed order: a selection schedule, exact factors on the
quantized side channel, a truncated host on that channel, and the weight
path, which takes whatever the others declined. Order is semantics, not
preference: each mechanism is more faithful than the one after it, and only
the weight path can take any layer.
Contracts the walk depends on:
- The wanted-name tuple is built once per pass and handed to every layer as
the same object. The factor cache memoizes its pass entry on that
identity, so an equal tuple rebuilt per component makes every lookup
reread the entry from disk.
- Mechanism apply functions answer with three states: applied, took the
layer without changing it, or declined. Only a decline falls through.
- A layer is offered to a mechanism on its checkpoint weights, so factors
attach to a clean base and deltas are measured against one. `apply_cached`
can strip factors and still decline, which is why the weight path strips
again before it writes.
- The fuse decision is resolved once per pass and shared by backup, apply
and restore. Backup mode keeps a tensor and restores in the walk itself;
fuse mode keeps a marker and subtracts the delta in network_deactivate.
- Selection registration reads the factors it schedules, so it follows the
attach that produced them.
State the walk keeps on the model's own modules:
- network_layer_name: written by lora_convert and native_adapter.
- network_current_names and network_current_stack: written here, always
together, and read together as the skip key.
- network_weights_backup, network_bias_backup and the sdnq_*_backup set:
written by lora_apply, a tensor in backup mode and True as the fuse marker.
- sdnq_lora_svd_stash: written by lora_sdnq, holding the checkpoint's own
factors while a set is attached.
- sdnq_calib_rms: written by lora_calib.
- svd_up and svd_down: owned by sdnq, attached by lora_sdnq, restored by
lora_apply, and written in segments by lora_stack at flip time.
"""
from contextlib import nullcontext
import time
import rich.progress as rp
@@ -18,6 +59,64 @@ native_active: bool = False
default_components = ['text_encoder', 'text_encoder_2', 'text_encoder_3', 'text_encoder_4', 'unet', 'transformer', 'transformer_2', '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."""
@@ -45,221 +144,254 @@ def group_offload_strip(sd_model, component_name: str, stripped: dict):
return stripped[component_name]
def network_activate(include=None, exclude=None):
if exclude is None:
exclude = []
if include is None:
include = []
for net in l.loaded_networks: # promote staged multipliers only now: the deactivate pass ran against the previous values, which fuse-mode removal recomputes with
def promote_pending():
"""Promote staged multipliers onto the loaded networks; the deactivate pass ran against the previous values, which fuse-mode removal recomputes with."""
for net in l.loaded_networks:
pending = getattr(net, 'pending_config', None)
if pending is not None:
net.te_multiplier = pending['te']
net.unet_multiplier = pending['unet']
net.dyn_dim = pending['dyn']
net.block_spec = pending.get('blocks', None)
t0 = time.time()
fuse = lora_overrides.fuse_native() # resolve once: backup and apply passes must agree
with limit_errors("network_activate") as elimit:
sd_model = getattr(shared.sd_model, "pipe", shared.sd_model)
if shared.opts.diffusers_offload_mode == "sequential":
sd_models.disable_offload(sd_model)
sd_models.move_model(sd_model, device=devices.cpu)
elif shared.opts.diffusers_offload_mode == "balanced":
sd_model = sd_models.apply_balanced_offload(sd_model, force=True, silent=True) # dispatched modules hold meta tensors backed by the offload map; rebuild them real on cpu with hooks intact before touching weights
group_offload = shared.opts.diffusers_offload_mode == "group"
group_stripped = {}
device = None
modules = {}
components = include if len(include) > 0 else default_components
components = [x for x in components if x not in exclude]
filtered_components = [x for x in default_components if x not in components] # filtered components restore to backup so a filter means detached, not frozen with stale weights
active_components = []
for name in components + filtered_components:
component = getattr(sd_model, name, None)
if component is not None and hasattr(component, 'named_modules'):
if name in components:
active_components.append(name)
modules[name] = list(component.named_modules())
total = sum(len(x) for x in modules.values())
if len(l.loaded_networks) > 0:
pbar = rp.Progress(rp.TextColumn('[cyan]Network: type=LoRA action=activate'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=console)
task = pbar.add_task(description='' , total=total)
else:
task = None
pbar = nullcontext()
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
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
for _, module in modules[component]:
network_layer_name = getattr(module, 'network_layer_name', None)
current_names = getattr(module, "network_current_names", ())
if getattr(module, 'weight', None) is None or shared.state.interrupted or (network_layer_name is None) or (current_names == component_wanted and getattr(module, 'network_current_stack', 'sum') == stack_sig):
if task is not None:
pbar.update(task, advance=1)
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
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)
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)
weights_backup = getattr(module, "network_weights_backup", None)
if weights_backup is not None and not isinstance(weights_backup, bool):
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):
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) # an earlier non-factorable set requantized this layer, restore the pristine base before attaching factors
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):
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) # 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)
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:
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}')
net.pending_config = None # promotion is one-shot
if task is not None and len(applied_layers) == 0:
pbar.remove_task(task) # hide progress bar for no action
def prepare_model_for_write(sd_model):
"""Bring the model into a state where weight writes land; balanced offload returns a rebuilt model."""
if shared.opts.diffusers_offload_mode == "sequential":
sd_models.disable_offload(sd_model)
sd_models.move_model(sd_model, device=devices.cpu)
elif shared.opts.diffusers_offload_mode == "balanced":
sd_model = sd_models.apply_balanced_offload(sd_model, force=True, silent=True) # dispatched modules hold meta tensors backed by the offload map; rebuild them real on cpu with hooks intact before touching weights
return sd_model
def collect_components(sd_model, include, exclude, defaults, restore_filtered):
"""Modules to walk, as (modules, wanted components, walked component names, module count).
With restore_filtered the walk also covers the components a filter left
out, so they restore to backup instead of freezing with stale weights;
those names stay out of the reported list because nothing applies to them.
"""
components = include if len(include) > 0 else defaults
components = [x for x in components if x not in exclude]
filtered = [x for x in defaults if x not in components] if restore_filtered else []
modules = {}
active_components = []
for name in components + filtered:
component = getattr(sd_model, name, None)
if component is not None and hasattr(component, 'named_modules'):
if name in components:
active_components.append(name)
modules[name] = list(component.named_modules())
return modules, components, active_components, sum(len(x) for x in modules.values())
def pass_progress(action, total, show):
"""Progress bar for one pass, or a nullcontext with no task when there is nothing to show."""
if not show:
return nullcontext(), None
pbar = rp.Progress(rp.TextColumn(f'[cyan]Network: type=LoRA action={action}'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=console)
return pbar, pbar.add_task(description='', total=total)
def tensor_backup(module):
"""The module's weight backup when it holds real tensors; None in fuse mode, where the backup is a marker."""
weights_backup = getattr(module, 'network_weights_backup', None)
return None if isinstance(weights_backup, bool) else weights_backup
def restore_pristine(module, device):
"""Put a backed-up layer back on its checkpoint weights, so a mechanism sees the pristine base."""
if tensor_backup(module) is not None:
network_apply_weights(module, None, None, device=device)
def should_skip(module, network_layer_name, wanted, stack_sig):
"""True when the pass has nothing to do here: no weight, interrupted, unnamed, or already carrying this set under these settings."""
if getattr(module, 'weight', None) is None or shared.state.interrupted or network_layer_name is None:
return True
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 finish_pass(ctx, t0):
"""Publish what the pass did and put the model back under its offload mode.
Runs even when the error limiter aborts the walk: the hooks it stripped
and the offload it disabled have to come back, and the counters other
modules read have to describe this pass.
"""
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.last_mode = 'backup' if ctx.backup_size > 0 else ('fuse' if ctx.fuse else 'factor')
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}')
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")
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}')
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 network_activate(include=None, exclude=None):
if exclude is None:
exclude = []
if include is None:
include = []
promote_pending()
t0 = time.time()
ctx = ActivationPass(lora_overrides.fuse_native()) # fuse resolved once: the backup, apply and restore paths must agree
applied_layers.clear()
lora_sdnq.reset_pass()
modules = {}
try:
with limit_errors("network_activate") as elimit:
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:
for component in modules.keys():
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, 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 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
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
finally:
finish_pass(ctx, t0)
modules.clear()
def effective_mode():
"""Weight-state label for load logs: backup and fuse say how touched weights restore, factor means the whole load rode the svd channel and unload just drops factors."""
if getattr(l, 'last_backup_size', 0) > 0:
return 'backup'
if lora_overrides.fuse_native():
return 'fuse'
return 'factor'
"""Weight-state label for load logs: backup and fuse say how touched weights restore, factor means the whole load rode the svd channel and unload just drops factors.
Recorded by the pass rather than derived here, so the unload line
describes the pass being unloaded even when the settings it ran under
have since changed.
"""
if l.last_mode:
return l.last_mode
return 'fuse' if lora_overrides.fuse_native() else 'factor'
def network_deactivate(include=None, exclude=None):
@@ -274,31 +406,11 @@ def network_deactivate(include=None, exclude=None):
return
t0 = time.time()
with limit_errors("network_deactivate") as elimit:
sd_model = getattr(shared.sd_model, "pipe", shared.sd_model)
if shared.opts.diffusers_offload_mode == "sequential":
sd_models.disable_offload(sd_model)
sd_models.move_model(sd_model, device=devices.cpu)
elif shared.opts.diffusers_offload_mode == "balanced":
sd_model = sd_models.apply_balanced_offload(sd_model, force=True, silent=True) # dispatched modules hold meta tensors backed by the offload map; rebuild them real on cpu with hooks intact before touching weights
sd_model = prepare_model_for_write(getattr(shared.sd_model, "pipe", shared.sd_model))
group_offload = shared.opts.diffusers_offload_mode == "group"
group_stripped = {}
modules = {}
components = include if len(include) > 0 else ['text_encoder', 'text_encoder_2', 'text_encoder_3', 'unet', 'transformer', 'llm_adapter']
components = [x for x in components if x not in exclude]
active_components = []
for name in components:
component = getattr(sd_model, name, None)
if component is not None and hasattr(component, 'named_modules'):
modules[name] = list(component.named_modules())
active_components.append(name)
total = sum(len(x) for x in modules.values())
if len(l.previously_loaded_networks) > 0 and l.debug:
pbar = rp.Progress(rp.TextColumn('[cyan]Network: type=LoRA action=deactivate'), rp.BarColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=console)
task = pbar.add_task(description='', total=total)
else:
task = None
pbar = nullcontext()
modules, _components, active_components, total = collect_components(sd_model, include, exclude, default_components, restore_filtered=False)
pbar, task = pass_progress('deactivate', total, len(l.previously_loaded_networks) > 0 and l.debug)
refused = 0
with devices.inference_context(), pbar:
applied_layers.clear()
+78 -2
View File
@@ -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)
@@ -2480,6 +2515,45 @@ def test_nunchaku_entries_carry_the_network_interface():
return True
def test_native_dispatch_archs_are_native_eligible():
from modules.lora import lora_load, lora_overrides
missing = sorted(set(lora_load.NATIVE_DISPATCH) - set(lora_overrides.allow_native))
assert not missing, f'an arch with a native loader that the method choice sends elsewhere never reaches it: {missing}'
return True
def test_aborted_pass_still_publishes_its_state():
layer = build_layer('uint4')
_A, _B, D = make_delta()
net = make_dense_net('aborted', layer, D)
reported = {'n': 0}
real_prepare = networks.prepare_model_for_write
real_report = lora_sdnq.report_fallbacks
def exploding(_sd_model):
raise RuntimeError('offload rebuild failed') # a raise before the walk binds anything the epilogue reads
def counting_report():
reported['n'] += 1
real_report()
networks.prepare_model_for_write = exploding
lora_sdnq.report_fallbacks = counting_report
try:
with mock_model(lin=layer):
raised = None
try:
activate(net)
except RuntimeError as e:
raised = e
assert raised is not None and 'offload rebuild failed' in str(raised), f'the original failure must reach the caller, got {raised!r}'
assert reported['n'] == 1, 'an aborted pass must still publish its counters and put the model back under its offload mode'
finally:
networks.prepare_model_for_write = real_prepare
lora_sdnq.report_fallbacks = real_report
return True
def test_stacked_shape_mismatch_falls_back():
from types import SimpleNamespace
layer = build_layer('uint4')
@@ -2914,7 +2988,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)
@@ -2950,7 +3025,8 @@ def run_tests():
run_test(CAT_COMPILE, fn)
log.warning('=== Robustness ===')
for fn in [test_remove_factors_after_device_move, test_stacked_shape_mismatch_falls_back, test_nunchaku_entries_carry_the_network_interface,
test_four_dim_oft_blocks_load_as_boft]:
test_four_dim_oft_blocks_load_as_boft, test_aborted_pass_still_publishes_its_state,
test_native_dispatch_archs_are_native_eligible]:
run_test(CAT_ROBUST, fn)
log.warning('=== Block weights ===')
for fn in [test_block_index_sd_unet_layout, test_block_index_sdxl_unet_layout, test_block_index_flux_chains_concatenate,