accelerate for model loading

This commit is contained in:
kijai
2024-06-30 13:24:22 +03:00
parent 78e4934d2a
commit 04d27efaa3
3 changed files with 61 additions and 53 deletions
+26 -38
View File
@@ -5,51 +5,39 @@ import torch
from einops import rearrange
from safetensors.torch import load_file
def load_model_checkpoint(model, ckpt):
def load_checkpoint(model, ckpt, full_strict):
if "safetensors" in ckpt:
try:
state_dict = load_file(ckpt)
except:
state_dict = torch.load(ckpt, map_location="cpu")
else:
state_dict = torch.load(ckpt, map_location="cpu")
if "state_dict" in list(state_dict.keys()):
state_dict = state_dict["state_dict"]
from contextlib import nullcontext
try:
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
is_accelerate_available = True
except:
pass
filtered_state_dict = {
k: v
for k, v in state_dict.items()
if not (k.startswith("cond_stage_model") or k.startswith("embedder"))
#if not (k.startswith("cond_stage_model"))
} # Filter out keys starting with "cond_stage_model" and "embedder"
def load_model_checkpoint(model, file_path, dtype, device):
if "safetensors" in file_path:
try:
model.load_state_dict(filtered_state_dict, strict=full_strict)
state_dict = load_file(file_path)
except:
## rename the keys for 256x256 model
new_pl_sd = OrderedDict()
for k,v in state_dict.items():
new_pl_sd[k] = v
state_dict = torch.load(file_path, map_location="cpu")
else:
state_dict = torch.load(file_path, map_location="cpu")
if "state_dict" in list(state_dict.keys()):
state_dict = state_dict["state_dict"]
for k in list(new_pl_sd.keys()):
if "framestride_embed" in k:
new_key = k.replace("framestride_embed", "fps_embedding")
new_pl_sd[new_key] = new_pl_sd[k]
del new_pl_sd[k]
model.load_state_dict(new_pl_sd, strict=full_strict)
# else:
# ## deepspeed
# new_pl_sd = OrderedDict()
# for key in state_dict['module'].keys():
# new_pl_sd[key[16:]]=state_dict['module'][key]
# model.load_state_dict(new_pl_sd, strict=full_strict)
filtered_state_dict = {
k: v
for k, v in state_dict.items()
if not (k.startswith("cond_stage_model") or k.startswith("embedder"))
#if not (k.startswith("cond_stage_model"))
} # Filter out keys starting with "cond_stage_model" and "embedder"
if is_accelerate_available:
for key in filtered_state_dict:
set_module_tensor_to_device(model, key, dtype=dtype, device=device, value=filtered_state_dict[key])
else:
model.load_state_dict(filtered_state_dict, strict=True)
return model
load_checkpoint(model, ckpt, full_strict=False)
print('>>> model checkpoint loaded.')
return model
def load_prompts(prompt_file):
f = open(prompt_file, 'r')
prompt_list = []