mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
Merge pull request #5070 from vladmandic/refactor/lora-activate-breakdown
refactor(lora): break down network_activate
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user