mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 01:04:32 +02:00
enable diffusers_move_unet for Flux
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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'
|
||||
|
||||
Reference in New Issue
Block a user