From 1cbc313fa57050a1b98e75d72ba1adf5c45a6b99 Mon Sep 17 00:00:00 2001 From: CalamitousFelicitousness Date: Fri, 14 Aug 2026 23:58:24 +0100 Subject: [PATCH] fix(offload): derive group offload roles structurally and unify the dispatch Group offload placed components by name matching, which sent every modular text encoder through a second apply with the denoiser profile and missed vae-class components not named vae. Placement now derives from the component: an encode or decode entry bridge and an upstream group-offload opt-out both take whole-module on-demand placement, an undecorated encode or decode entry stays resident with a warning, denoiser slot names take the per-step profile, and everything else takes the leaf profile. One pass serves classic and modular pipelines, enumerating modular components from their specs, with hooks landing on text encoder wrappers rather than their inner models. A signature sentinel marks a component mid-apply; a failed apply is torn down on the next pass. The pass logs the placement of every component. --- CHANGELOG.md | 2 + modules/sd_offload.py | 164 +++++++++++++++++++++--------------------- 2 files changed, 84 insertions(+), 82 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 62778560f..15658f323 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -48,6 +48,8 @@ - improve handling of hf auth - improve pipeline detection for non-cached models - cleanup alt offload codepaths + - offload: text encoders no longer take the denoiser profile on modular pipelines + - offload: components entered through encode or decode are detected by structure rather than by name - hf progress bars - ltx: send the guidance stack and cross-timestep on every 2.x call path - ltx: distilled variants no longer force dynamic shifting on, which remapped their sigma schedule diff --git a/modules/sd_offload.py b/modules/sd_offload.py index f04b785e0..0368af422 100644 --- a/modules/sd_offload.py +++ b/modules/sd_offload.py @@ -20,6 +20,10 @@ offload_allow_none = ['sd', 'sdxl'] offload_post = ['h1'] offload_hook_instance = None balanced_offload_exclude = ['CogView4Pipeline', 'MeissonicPipeline'] +group_offload_main = [ # component names entered once per denoising step + "unet", "transformer", "transformer_2", "transformer_ref", "unconditional_transformer", + "prior", "prior_prior", "decoder", "dit_model", "model", "controlnet", +] # a denoiser registered under any other name takes the aux profile until listed here no_split_module_classes = [ "Linear", "Conv1d", "Conv2d", "Conv3d", "ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d", "Embedding", "SDNQLinear", "SDNQConv1d", "SDNQConv2d", "SDNQConv3d", "SDNQConvTranspose1d", "SDNQConvTranspose2d", "SDNQConvTranspose3d", "SDNQEmbedding", @@ -27,6 +31,7 @@ no_split_module_classes = [ "MiniMaxH3TransformerBlock", "MiniMaxH3TokenRefinerBlock", ] accelerate_dtype_byte_size = None +group_stats_reported = set() move_stream = None @@ -161,30 +166,44 @@ def apply_group_offload_component(module, module_name: str, main: bool) -> bool: return False if hasattr(module, '_hf_hook'): # leftover accelerate hooks from a previous offload mode abort the group apply upstream module = accelerate.hooks.remove_hook_from_module(module, recurse=True) + module.sdnext_ondemand = False # group placement replaces any on-demand hook remove_group_offload_component(module) module.requires_grad_(False) log.debug(f'Offload: type=group op=apply type={shared.opts.group_offload_type} module={module_name} pin={cfg["use_stream"] and not cfg["low_cpu_mem_usage"]}') # before the apply: pinning large components takes a while and would otherwise run silently + module.sdnext_group_offload_sig = 'partial' # a raise below leaves hooks that only a non-empty signature will remove apply_group_offloading(module, onload_device=devices.device, offload_device=devices.cpu, **cfg) module.sdnext_group_offload_sig = sig return True -def set_group_resident(module): - """VAE-class components never take group hooks: the hooks are forward-scoped, while - pipelines enter through encode/decode, and tiled calls re-enter per tile.""" +def set_group_resident(module) -> bool: + """Keep a component on the accelerator with no hooks of any kind.""" + changed = False if hasattr(module, '_hf_hook'): module = accelerate.hooks.remove_hook_from_module(module, recurse=True) - remove_group_offload_component(module) + changed = True + if remove_group_offload_component(module): + changed = True + module.sdnext_ondemand = False module.requires_grad_(False) - module.to(devices.device) + if any(not devices.same_device(t.device, devices.device) for t in itertools.chain(module.parameters(), module.buffers())): # an interrupted generation can leave a group-hooked module split across devices + module.to(devices.device) + changed = True + return changed -def group_offload_role(module_name) -> str: - if any(m in module_name for m in ['vae']): - return 'vae' - if any(m in module_name for m in ['text_encoder', 'image_encoder', 'safety_checker']): - return 'aux' - return 'main' +def group_offload_role(module_name: str, module) -> str: + """Placement role for one component: resident to stay put, ondemand for whole-module onload, main for per-step denoisers, aux for the rest.""" + if has_entry_bridge(module): + return 'ondemand' # encode and decode bypass the forward that group hooks scope to + if callable(getattr(module, 'encode', None)) or callable(getattr(module, 'decode', None)): + log.warning(f'Offload: type=group module={module_name} cls={module.__class__.__name__} bridge=missing role=resident') # decorate the entry points with apply_forward_hook to make the component offloadable + return 'resident' # nothing fires an onload for an undecorated entry point, so any hook placement strands the weights on cpu + if not getattr(module, '_supports_group_offloading', True): + return 'ondemand' # upstream marks modules that read submodule weights outside those submodules' forward + if module_name in group_offload_main: + return 'main' + return 'aux' def has_entry_bridge(module) -> bool: @@ -212,27 +231,17 @@ class OnDemandHook(accelerate.hooks.ModelHook): return args, kwargs -def apply_group_offload_vae(sd_model, module, module_name: str) -> str: - """Placement policy for vae-class components, which never take group hooks. Components onload whole when their decode or encode entry point fires.""" - if not has_entry_bridge(module): - log.warning(f'Offload: type=group module={module_name} class={module.__class__.__name__} no entry bridge') - set_group_resident(module) # TODO group offload: this will fail as vae will end up on cpu and there will be nothing to pull it back on gpu - module.sdnext_ondemand = False # a lingering stamp would let the seams offload a component with no onload hook - names = getattr(sd_model, 'sdnext_ondemand_modules', None) or [] - if module_name in names: - sd_model.sdnext_ondemand_modules = [n for n in names if n != module_name] +def apply_group_offload_ondemand(module) -> bool: + """Placement for components that never take group hooks: they onload whole at their entry point.""" + if getattr(module, 'sdnext_ondemand', False) and hasattr(module, '_hf_hook'): return False - if not getattr(module, 'sdnext_ondemand', False) or not hasattr(module, '_hf_hook'): - if hasattr(module, '_hf_hook'): - module = accelerate.hooks.remove_hook_from_module(module, recurse=True) - remove_group_offload_component(module) - module.requires_grad_(False) - accelerate.hooks.add_hook_to_module(module, OnDemandHook(), append=False) - module.sdnext_ondemand = True - module.to(devices.cpu) - names = getattr(sd_model, 'sdnext_ondemand_modules', None) or [] - if module_name not in names: - sd_model.sdnext_ondemand_modules = names + [module_name] + if hasattr(module, '_hf_hook'): + module = accelerate.hooks.remove_hook_from_module(module, recurse=True) + remove_group_offload_component(module) + module.requires_grad_(False) + accelerate.hooks.add_hook_to_module(module, OnDemandHook(), append=False) + module.sdnext_ondemand = True + module.to(devices.cpu) return True @@ -265,9 +274,10 @@ def offload_ondemand(sd_model, include=[], exclude=[], reason='', force=False): def report_group_stats(sd_model, module_names): """Per-component stats block once per loaded model; balanced mode prints its own from the hook map.""" - if getattr(sd_model, 'sdnext_group_stats_reported', False): + checkpoint_name = sd_model.sd_checkpoint_info.name if getattr(sd_model, "sd_checkpoint_info", None) is not None else sd_model.__class__.__name__ + if checkpoint_name in group_stats_reported: # keyed by checkpoint since a task switch rebuilds the pipe object return - sd_model.sdnext_group_stats_reported = True + group_stats_reported.add(checkpoint_name) total = 0.0 counted = [] for module_name in module_names: @@ -279,51 +289,35 @@ def report_group_stats(sd_model, module_names): log.info(f'Model class={sd_model.__class__.__name__} modules={len(counted)} size={total:.3f}') -def apply_modular_group_offload(sd_model): - """Per-component group offload for modular pipelines, which lack the pipeline-level - enable_*_offload entry points. The model and sequential modes also route here.""" - if shared.opts.diffusers_offload_mode != 'group' and not getattr(sd_model, 'sdnext_modular_offload_warned', False): - sd_model.sdnext_modular_offload_warned = True - log.warning(f'Offload: desired={shared.opts.diffusers_offload_mode} override=group reason="modular pipeline"') - applied = [] - loaded = [name for name, component in sd_model.components.items() if isinstance(component, torch.nn.Module)] - for name in loaded: - module = getattr(sd_model, name, None) - if 'text_encoder' in name: - if apply_group_offload_component(getattr(module, 'model', module), name, main=False): - applied.append(name) - if 'vae' in name: - if apply_group_offload_vae(sd_model, module, name): - applied.append(name) - else: - if apply_group_offload_component(module, name, main=True): - applied.append(name) - # has_accelerate stays unset: group hooks are not accelerate hooks, and the modular pipeline's own to() skips group-offloaded components when move_model runs - log.info(f'Offload: type=group type={shared.opts.group_offload_type} modules={applied}') - report_group_stats(sd_model, ('transformer', 'transformer_ref', 'text_encoder', 'vae', 'audio_vae')) - - def apply_group_offload(sd_model): - applied, resident, ondemand = [], [], [] - for module_name in get_module_names(sd_model): + """Per-component group offload for classic and modular pipelines.""" + changed = False + placements = [] + module_names = get_module_names(sd_model) + for module_name in module_names: module = getattr(sd_model, module_name, None) if not isinstance(module, torch.nn.Module): continue try: - role = group_offload_role(module_name) - if role == 'vae': - if apply_group_offload_vae(sd_model, module, module_name) == 'ondemand': - ondemand.append(module_name) - else: - resident.append(module_name) - elif apply_group_offload_component(module, module_name, main=role == 'main'): - applied.append(module_name) + role = group_offload_role(module_name, module) + placements.append(f'{module_name}:{role}') + if role == 'resident': + applied = set_group_resident(module) + elif role == 'ondemand': + applied = apply_group_offload_ondemand(module) + else: + applied = apply_group_offload_component(module, module_name, main=role == 'main') + changed = changed or applied except Exception as e: log.error(f'Offload: type=group module={module_name} {e}') - set_accelerate(sd_model) - if applied: - log.info(f'Offload: type=group type={shared.opts.group_offload_type} modules={applied} resident={resident} ondemand={ondemand}') - report_group_stats(sd_model, get_module_names(sd_model)) + sd_model.sdnext_ondemand_modules = [name for name in module_names if getattr(getattr(sd_model, name, None), 'sdnext_ondemand', False)] + if sd_models.get_diffusers_task(sd_model) != sd_models.DiffusersTaskType.MODULAR: # group hooks are not accelerate hooks, so modular pipelines stay unstamped + set_accelerate(sd_model) + if changed: + log.info(f'Offload: type=group modules={placements}') + else: + log.debug(f'Offload: type=group modules={placements}') + report_group_stats(sd_model, module_names) return sd_model @@ -387,7 +381,10 @@ def set_diffuser_offload(sd_model, op:str='model', quiet:bool=False, force:bool= accelerate.utils.modeling.dtype_byte_size = dtype_byte_size if sd_models.get_diffusers_task(sd_model) == sd_models.DiffusersTaskType.MODULAR and shared.opts.diffusers_offload_mode in {'model', 'sequential', 'group'}: - apply_modular_group_offload(sd_model) + if shared.opts.diffusers_offload_mode != 'group' and not getattr(sd_model, 'sdnext_modular_offload_warned', False): + sd_model.sdnext_modular_offload_warned = True + log.warning(f'Offload: desired={shared.opts.diffusers_offload_mode} override=group reason="modular pipeline"') + apply_group_offload(sd_model) process_timer.add('offload', time.time() - t0) return @@ -583,16 +580,19 @@ def get_module_names(pipe=None, exclude=None): else: return [] modules_names = [] - try: - dict_keys = pipe._internal_dict.keys() # pylint: disable=protected-access - modules_names.extend(dict_keys) - except Exception: - pass - try: - dict_keys = get_signature(pipe).keys() - modules_names.extend(dict_keys) - except Exception: - pass + if hasattr(pipe, '_component_specs'): # modular pipelines name their components in specs; the config dict also carries scalars + modules_names.extend(pipe.components) + else: + try: + dict_keys = pipe._internal_dict.keys() # pylint: disable=protected-access + modules_names.extend(dict_keys) + except Exception: + pass + try: + dict_keys = get_signature(pipe).keys() + modules_names.extend(dict_keys) + except Exception: + pass modules_names = [m for m in modules_names if m not in exclude and not m.startswith('_')] modules_names = [m for m in modules_names if is_valid(m)] modules_names = sorted(set(modules_names))