From 73c48c1b15ee4279d28eb31a52e4905fa2379394 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 14 Apr 2024 11:17:35 +0300 Subject: [PATCH] accelerate fixes --- nodes.py | 16 +++++++++++----- requirements.txt | 1 + 2 files changed, 12 insertions(+), 5 deletions(-) diff --git a/nodes.py b/nodes.py index 7238add..080fc11 100644 --- a/nodes.py +++ b/nodes.py @@ -156,14 +156,20 @@ class ella_model_loader: sd = model.model.state_dict_for_saving(clip_sd, vae.get_sd(), None) converted_vae = convert_ldm_vae_checkpoint(sd, converted_vae_config) - for key in converted_vae: - set_module_tensor_to_device(new_vae, key, device=device, dtype=dtype, value=converted_vae[key]) + if is_accelerate_available(): + for key in converted_vae: + set_module_tensor_to_device(new_vae, key, device=device, dtype=dtype, value=converted_vae[key]) + else: + new_vae.load_state_dict(converted_vae, strict=False) del converted_vae pbar.update(1) - converted_unet = convert_ldm_unet_checkpoint(sd, converted_unet_config) - for key in converted_unet: - set_module_tensor_to_device(unet, key, device=device, dtype=dtype, value=converted_unet[key]) + converted_unet = convert_ldm_unet_checkpoint(sd, converted_unet_config) + if is_accelerate_available(): + for key in converted_unet: + set_module_tensor_to_device(unet, key, device=device, dtype=dtype, value=converted_unet[key]) + else: + unet.load_state_dict(converted_unet, strict=False) del converted_unet ella = ELLA() diff --git a/requirements.txt b/requirements.txt index cefa596..44c2b44 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,3 +1,4 @@ diffusers>=0.26.0 +accelerate omegaconf sentencepiece