84 lines
3.1 KiB
Python
84 lines
3.1 KiB
Python
#import sys
|
|
from collections import OrderedDict
|
|
import torch
|
|
#sys.path.insert(1, os.path.join(sys.path[0], '..', '..'))
|
|
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"]
|
|
|
|
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"
|
|
try:
|
|
model.load_state_dict(filtered_state_dict, strict=full_strict)
|
|
except:
|
|
## rename the keys for 256x256 model
|
|
new_pl_sd = OrderedDict()
|
|
for k,v in state_dict.items():
|
|
new_pl_sd[k] = v
|
|
|
|
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)
|
|
|
|
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 = []
|
|
for idx, line in enumerate(f.readlines()):
|
|
l = line.strip()
|
|
if len(l) != 0:
|
|
prompt_list.append(l)
|
|
f.close()
|
|
return prompt_list
|
|
|
|
def get_latent_z(model, videos):
|
|
b, c, t, h, w = videos.shape
|
|
x = rearrange(videos, 'b c t h w -> (b t) c h w')
|
|
z = model.encode_first_stage(x)
|
|
z = rearrange(z, '(b t) c h w -> b c t h w', b=b, t=t)
|
|
return z
|
|
|
|
def get_latent_z_with_hidden_states(model, videos):
|
|
b, c, t, h, w = videos.shape
|
|
x = rearrange(videos, 'b c t h w -> (b t) c h w')
|
|
encoder_posterior, hidden_states = model.first_stage_model.encode(x, return_hidden_states=True)
|
|
|
|
hidden_states_first_last = []
|
|
### use only the first and last hidden states
|
|
for hid in hidden_states:
|
|
hid = rearrange(hid, '(b t) c h w -> b c t h w', t=t)
|
|
hid_new = torch.cat([hid[:, :, 0:1], hid[:, :, -1:]], dim=2)
|
|
hidden_states_first_last.append(hid_new)
|
|
|
|
z = model.get_first_stage_encoding(encoder_posterior).detach()
|
|
z = rearrange(z, '(b t) c h w -> b c t h w', b=b, t=t)
|
|
return z, hidden_states_first_last |