mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
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:
@@ -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'):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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"},
|
||||
|
||||
Reference in New Issue
Block a user