feat(offload): on-demand vae under group offload

Vae-class components never take group hooks, so group mode kept them
resident on the gpu; a MiniMax-class video vae holds about 10GB that
way while running only seconds per generation. Components above 1GB
now rest in system memory: the apply_forward_hook bridge on encode and
decode fires an on-demand hook that moves the whole module to the
device, so tiled calls find every weight already loaded, and the
processing seams return it to cpu once outputs are materialized. Small
vaes stay resident since the transfer would cost more than it frees.

- placement is decided per component by measured size and requires the
  entry bridge; components without it stay resident
- move_model no longer forces on-demand vaes to the gpu for
  non-txt2img tasks, and full_vae_encode onloads before binding the
  input, which otherwise lands on the resting device
- mode switches clear the stamp and hook in both directions
This commit is contained in:
CalamitousFelicitousness
2026-08-07 23:44:18 +01:00
parent b84b782ba4
commit abfb5ac3ed
5 changed files with 96 additions and 9 deletions
+1
View File
@@ -477,6 +477,7 @@ def process_decode(p: processing.StableDiffusionProcessing, output):
log.debug(f'Generated: frames={len(output.frames[0])}')
output.images = output.frames[0]
if output.images is not None and len(output.images) > 0 and isinstance(output.images[0], Image.Image):
sd_models.offload_ondemand(shared.sd_model) # in-pipe decode paths return materialized frames; the vae seam in processing_vae never runs
return attach_audio(output.images, audio)
model = shared.sd_model if not is_refiner_enabled(p) else shared.sd_refiner
if not hasattr(model, 'vae'):
+4
View File
@@ -189,6 +189,8 @@ def full_vae_encode(image, model):
sd_models.move_model(model.unet, devices.cpu)
if shared.opts.diffusers_offload_mode != "sequential" and hasattr(model, 'vae'):
sd_models.move_model(model.vae, devices.device)
if getattr(model.vae, 'sdnext_ondemand', False):
model.vae.to(devices.device) # the image placement below derives from vae.device, and the entry bridge would onload the weights only after the input is already bound
vae_name = sd_vae.loaded_vae_file if sd_vae.loaded_vae_file is not None else "default"
log_debug(f'Encode vae="{vae_name}" dtype={model.vae.dtype} upcast={model.vae.config.get("force_upcast", None)}')
@@ -369,6 +371,7 @@ def vae_decode(latents, model, output_type='np', vae_type='Full', width=None, he
if shared.cmd_opts.profile or debug:
t1 = time.time()
log.debug(f'Profile: VAE decode: {t1-t0:.2f}')
sd_models.offload_ondemand(model)
devices.torch_gc()
shared.state.end(jobid)
return images
@@ -393,6 +396,7 @@ def vae_encode(image, model, vae_type='Full'): # pylint: disable=unused-variable
else:
log.error('VAE not found in model')
latents = []
sd_models.offload_ondemand(model)
devices.torch_gc()
shared.state.end(jobid)
return latents
+2 -2
View File
@@ -16,7 +16,7 @@ from modules.memstats import memory_stats
from modules.shared_helpers import walk_files
from modules.modeldata import model_data
from modules.sd_checkpoint import CheckpointInfo, select_checkpoint, list_models, checkpoint_titles, get_closest_checkpoint_match, update_model_hashes, write_metadata, checkpoints_list # pylint: disable=unused-import
from modules.sd_offload import get_module_names, disable_offload, set_diffuser_offload, apply_balanced_offload, set_accelerate, remove_group_offload_component # pylint: disable=unused-import
from modules.sd_offload import get_module_names, disable_offload, set_diffuser_offload, apply_balanced_offload, set_accelerate, remove_group_offload_component, offload_ondemand # pylint: disable=unused-import
from modules.sd_models_utils import NoWatermark, get_signature, get_call, path_to_repo, apply_function_to_model, read_state_dict, get_state_dict_from_checkpoint # pylint: disable=unused-import
@@ -237,7 +237,7 @@ def move_model(model, device=None, force=False):
fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access
if getattr(model, 'vae', None) is not None and get_diffusers_task(model) != DiffusersTaskType.TEXT_2_IMAGE:
if device == devices.device and model.vae.device.type != "meta": # force vae back to gpu if not in txt2img mode
if device == devices.device and model.vae.device.type != "meta" and not getattr(model.vae, 'sdnext_ondemand', False): # force vae back to gpu if not in txt2img mode; on-demand vaes onload at their entry point instead
model.vae.to(device)
if hasattr(model.vae, '_hf_hook'):
debug_move(f'Model move: to={device} class={model.vae.__class__} fn={fn}') # pylint: disable=protected-access
+88 -6
View File
@@ -19,6 +19,7 @@ debug_move = log.trace if debug else lambda *args, **kwargs: None
offload_allow_none = ['sd', 'sdxl']
offload_post = ['h1']
offload_hook_instance = None
group_offload_vae_limit = 1.0 # GB; vae-class components above this rest on cpu and onload whole at encode/decode
balanced_offload_exclude = ['CogView4Pipeline', 'MeissonicPipeline']
no_split_module_classes = [
"Linear", "Conv1d", "Conv2d", "Conv3d", "ConvTranspose1d", "ConvTranspose2d", "ConvTranspose3d", "Embedding",
@@ -115,6 +116,15 @@ def remove_group_offload(sd_model):
if isinstance(module, torch.nn.Module) and getattr(module, 'sdnext_group_offload_sig', None) is not None:
remove_group_offload_component(module)
removed.append(module_name)
for module_name in getattr(sd_model, 'sdnext_ondemand_modules', None) or []:
module = getattr(sd_model, module_name, None)
if module is not None:
module.sdnext_ondemand = False
if hasattr(module, '_hf_hook'):
module = accelerate.hooks.remove_hook_from_module(module, recurse=True)
removed.append(f'{module_name}:ondemand')
if getattr(sd_model, 'sdnext_ondemand_modules', None):
sd_model.sdnext_ondemand_modules = []
if removed:
log.debug(f'Offload: type=group op=remove modules={removed}')
@@ -162,6 +172,76 @@ def group_offload_role(module_name: str, module) -> str:
return 'main'
def has_entry_bridge(module) -> bool:
"""Entry points decorated with diffusers' apply_forward_hook fire _hf_hook.pre_forward,
which is what carries the on-demand onload for encode and decode calls that bypass forward."""
for name in ('decode', 'encode'):
fn = getattr(module, name, None)
if fn is not None and getattr(fn, '__qualname__', '').startswith('apply_forward_hook'):
return True
return False
class OnDemandHook(accelerate.hooks.ModelHook):
"""Whole-module onload for components entered through decode or encode rather than forward.
Tiled calls re-enter inside one entry point, so the module is on device before the first
tile; the return to cpu happens at the processing seams once outputs are materialized."""
def pre_forward(self, module, *args, **kwargs):
param = next(module.parameters(), None)
if param is not None and not devices.same_device(param.device, devices.device):
t0 = time.time()
module.to(devices.device)
dt = time.time() - t0
process_timer.add('onload', dt)
log.debug(f'Offload: type=ondemand op=onload module={module.__class__.__name__} time={dt:.3f}')
return args, kwargs
def set_group_vae(sd_model, module, module_name: str) -> str:
"""Placement policy for vae-class components, which never take group hooks. Small
components stay resident; components above group_offload_vae_limit rest on cpu and
onload whole when their decode or encode entry point fires."""
size_gb, _params = get_module_size(module)
if size_gb < group_offload_vae_limit or not has_entry_bridge(module):
set_group_resident(module)
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]
return 'resident'
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]
return 'ondemand'
def offload_ondemand(sd_model):
"""Return on-demand components to cpu once their outputs are materialized."""
if sd_model is None:
return
names = getattr(sd_model, 'sdnext_ondemand_modules', None)
if not names and hasattr(sd_model, 'pipe'):
sd_model = sd_model.pipe
names = getattr(sd_model, 'sdnext_ondemand_modules', None)
for module_name in names or []:
module = getattr(sd_model, module_name, None)
param = next(module.parameters(), None) if module is not None else None
if param is not None and not devices.same_device(param.device, devices.cpu):
t0 = time.time()
module.to(devices.cpu)
dt = time.time() - t0
process_timer.add('offload', dt)
log.debug(f'Offload: type=ondemand op=offload module={module_name} time={dt:.3f}')
def apply_modular_group_offload(sd_model, op:str='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."""
@@ -182,8 +262,8 @@ def apply_modular_group_offload(sd_model, op:str='model'):
for name in ('vae', 'audio_vae'):
component = getattr(sd_model, name, None)
if component is not None:
set_group_resident(component)
applied.append(f'{name}:device')
placement = set_group_vae(sd_model, component, name)
applied.append(f'{name}:{placement}')
# 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
if any(':' not in name for name in applied):
@@ -191,7 +271,7 @@ def apply_modular_group_offload(sd_model, op:str='model'):
def apply_group_offload(sd_model, op:str='model'):
applied, resident = [], []
applied, resident, ondemand = [], [], []
for module_name in get_module_names(sd_model):
module = getattr(sd_model, module_name, None)
if not isinstance(module, torch.nn.Module):
@@ -199,15 +279,17 @@ def apply_group_offload(sd_model, op:str='model'):
try:
role = group_offload_role(module_name, module)
if role == 'resident':
set_group_resident(module)
resident.append(module_name)
if set_group_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', op=op):
applied.append(module_name)
except Exception as e:
log.error(f'Setting {op}: offload=group module={module_name} {e}')
set_accelerate(sd_model)
if applied:
log.info(f'Setting {op}: offload=group type={shared.opts.group_offload_type} modules={applied} resident={resident}')
log.info(f'Setting {op}: offload=group type={shared.opts.group_offload_type} modules={applied} resident={resident} ondemand={ondemand}')
return sd_model
+1 -1
View File
@@ -620,7 +620,7 @@
{"id":"","label":"Generic","localized":"","hint":"","ui":"video"},
{"id":"","label":"Google GenAI","localized":"","hint":"","ui":"settings_model_options"},
{"id":"","label":"Group Offload","localized":"","hint":"","ui":"settings_offload"},
{"id":"","label":"Group offload type","localized":"","hint":"Granularity used by <b>group</b> offload.<br>- <b>leaf_level</b>: offloads at the smallest module level; maximum memory savings, slower<br>- <b>block_level</b>: offloads groups of transformer blocks (size set by <b><i>Offload blocks</i></b>); faster with less savings<br>The VAE stays resident on the GPU in both modes, and text encoders always offload at leaf level.<br><br>Applies only when <b><i>Model offload mode</i></b> is <b>group</b>.<br><br>Default is <b>leaf_level</b>.","reload":"model","ui":"settings_offload"},
{"id":"","label":"Group offload type","localized":"","hint":"Granularity used by <b>group</b> offload.<br>- <b>leaf_level</b>: offloads at the smallest module level; maximum memory savings, slower<br>- <b>block_level</b>: offloads groups of transformer blocks (size set by <b><i>Offload blocks</i></b>); faster with less savings<br>Text encoders always offload at leaf level. Small VAEs stay resident on the GPU; VAEs above 1GB rest in system memory and load whole for each encode or decode.<br><br>Applies only when <b><i>Model offload mode</i></b> is <b>group</b>.<br><br>Default is <b>leaf_level</b>.","reload":"model","ui":"settings_offload"},
{"id":"","label":"Grid Options","localized":"","hint":"","ui":"settings_saving-images"},
{"id":"","label":"Grids","localized":"","hint":"","ui":"settings_saving-paths"},
{"id":"","label":"Guider","localized":"","hint":"","ui":"txt2img"},