accelerate for model loading
This commit is contained in:
+26
-38
@@ -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 = []
|
||||
|
||||
Reference in New Issue
Block a user