diff --git a/modules/processing_vae.py b/modules/processing_vae.py index 71ee84d45..5d50f0411 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -33,10 +33,8 @@ def create_latents(image, p, dtype=None, device=None): def full_vae_decode(latents, model): t0 = time.time() - if shared.opts.diffusers_move_unet and not getattr(model, 'has_accelerate', False) and hasattr(model, 'unet'): - shared.log.debug('Moving to CPU: model=UNet') - unet_device = model.unet.device - sd_models.move_model(model.unet, devices.cpu) + if shared.opts.diffusers_move_unet and not getattr(model, 'has_accelerate', False): + base_device = sd_models.move_base(model, devices.cpu) if not shared.cmd_opts.lowvram and not shared.opts.diffusers_seq_cpu_offload and hasattr(model, 'vae'): sd_models.move_model(model.vae, devices.device) latents.to(model.vae.device) @@ -69,8 +67,8 @@ def full_vae_decode(latents, model): model.vae.apply(sd_models.convert_to_faketensors) devices.torch_gc(force=True) - if shared.opts.diffusers_move_unet and not getattr(model, 'has_accelerate', False) and hasattr(model, 'unet'): - sd_models.move_model(model.unet, unet_device) + if shared.opts.diffusers_move_unet and not getattr(model, 'has_accelerate', False) and base_device is not None: + sd_models.move_base(model, base_device) t1 = time.time() debug(f'VAE decode: name={sd_vae.loaded_vae_file if sd_vae.loaded_vae_file is not None else "baked"} dtype={model.vae.dtype} upcast={upcast} images={latents.shape[0]} latents={latents.shape} time={round(t1-t0, 3)}') return decoded diff --git a/modules/sd_models.py b/modules/sd_models.py index 8eec61834..e9e6beae9 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -819,6 +819,19 @@ def move_model(model, device=None, force=False): devices.torch_gc() +def move_base(model, device): + key = 'unet' + if isinstance(model, diffusers.FluxPipeline): + key = 'transformer' + if not hasattr(model, key): + return None + shared.log.debug(f'Moving to CPU: model={key}') + model = getattr(model, key) + R = model.device + move_model(model, device) + return R + + def get_load_config(model_file, model_type, config_type='yaml'): if config_type == 'yaml': yaml = os.path.splitext(model_file)[0] + '.yaml'