-
Notifications
You must be signed in to change notification settings - Fork 10
Expand file tree
/
Copy pathtest_denoiser.py
More file actions
72 lines (58 loc) · 2.8 KB
/
Copy pathtest_denoiser.py
File metadata and controls
72 lines (58 loc) · 2.8 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
import torch
from os.path import join as pjoin
from diffusers import DDIMScheduler
from models.vae.model import VAE
from models.denoiser.model import Denoiser
from models.denoiser.trainer import DenoiserTrainer
from options.denoiser_option import arg_parse
from utils.get_opt import get_opt
from utils.fixseed import fixseed
from motion_loaders.dataset_motion_loader import get_dataset_motion_loader
from models.t2m_eval_wrapper import EvaluatorModelWrapper
def load_vae(vae_opt):
print(f'Loading VAE Model {vae_opt.name}')
model = VAE(vae_opt)
ckpt = torch.load(pjoin(vae_opt.checkpoints_dir, vae_opt.dataset_name, vae_opt.name, 'model', 'net_best_fid.tar'),
map_location='cpu')
model.load_state_dict(ckpt["vae"])
model.freeze()
return model
def load_denoiser(opt, vae_dim):
print(f'Loading Denoiser Model {opt.name}')
denoiser = Denoiser(opt, vae_dim)
ckpt = torch.load(pjoin(opt.checkpoints_dir, opt.dataset_name, opt.name, 'model', 'net_best_fid.tar'),
map_location='cpu')
missing_keys, unexpected_keys = denoiser.load_state_dict(ckpt["denoiser"], strict=False)
assert len(unexpected_keys) == 0
assert all([k.startswith('clip_model.') for k in missing_keys])
return denoiser
if __name__ == '__main__':
opt = arg_parse(False)
vae_name = get_opt(pjoin(opt.checkpoints_dir, opt.dataset_name, opt.name, 'opt.txt'), opt.device).vae_name
vae_opt = get_opt(pjoin(opt.checkpoints_dir, opt.dataset_name, vae_name, 'opt.txt'), opt.device)
cond_scale = opt.cond_scale
num_inference_timesteps = opt.num_inference_timesteps
opt = get_opt(pjoin(opt.checkpoints_dir, opt.dataset_name, opt.name, 'opt.txt'), opt.device)
opt.cond_scale = cond_scale
opt.num_inference_timesteps = num_inference_timesteps
fixseed(opt.seed)
# evaluation setup
dataset_opt_path = f"checkpoints/{opt.dataset_name}/Comp_v6_KLD005/opt.txt"
wrapper_opt = get_opt(dataset_opt_path, torch.device('cuda'))
eval_wrapper = EvaluatorModelWrapper(wrapper_opt)
eval_val_loader, _ = get_dataset_motion_loader(dataset_opt_path, 32, 'test', device=opt.device)
# models & noise scheduler
vae_model = load_vae(vae_opt).to(opt.device)
denoiser = load_denoiser(opt, vae_opt.latent_dim).to(opt.device)
scheduler = DDIMScheduler(
num_train_timesteps=opt.num_train_timesteps,
beta_start=opt.beta_start,
beta_end=opt.beta_end,
beta_schedule=opt.beta_schedule,
prediction_type=opt.prediction_type,
clip_sample=False,
)
# train
trainer = DenoiserTrainer(opt, denoiser, vae_model, scheduler)
trainer.test(eval_wrapper, eval_val_loader, 20,
save_dir=pjoin(opt.checkpoints_dir, opt.dataset_name, opt.name, 'eval'), cal_mm=False, save_motion=False)