mingyuan/MotionDiffuse
69
1from datetime import datetime
2import numpy as np
3import torch
4from datasets import get_dataset_motion_loader, get_motion_loader
5from models import MotionTransformer
6from utils.get_opt import get_opt
7from utils.metrics import *
8from datasets import EvaluatorModelWrapper
9from collections import OrderedDict
10from utils.plot_script import *
11from utils import paramUtil
12from utils.utils import *
13from trainers import DDPMTrainer
14
15from os.path import join as pjoin
16import sys
17
18
19def build_models(opt, dim_pose):
20 encoder = MotionTransformer(
21 input_feats=dim_pose,
22 num_frames=opt.max_motion_length,
23 num_layers=opt.num_layers,
24 latent_dim=opt.latent_dim,
25 no_clip=opt.no_clip,
26 no_eff=opt.no_eff)
27 return encoder
28
29
30torch.multiprocessing.set_sharing_strategy('file_system')
31
32
33def evaluate_matching_score(motion_loaders, file):
34 match_score_dict = OrderedDict({})
35 R_precision_dict = OrderedDict({})
36 activation_dict = OrderedDict({})
37 # print(motion_loaders.keys())
38 print('========== Evaluating Matching Score ==========')
39 for motion_loader_name, motion_loader in motion_loaders.items():
40 all_motion_embeddings = []
41 score_list = []
42 all_size = 0
43 matching_score_sum = 0
44 top_k_count = 0
45 # print(motion_loader_name)
46 with torch.no_grad():
47 for idx, batch in enumerate(motion_loader):
48 word_embeddings, pos_one_hots, _, sent_lens, motions, m_lens, _ = batch
49 text_embeddings, motion_embeddings = eval_wrapper.get_co_embeddings(
50 word_embs=word_embeddings,
51 pos_ohot=pos_one_hots,
52 cap_lens=sent_lens,
53 motions=motions,
54 m_lens=m_lens
55 )
56 dist_mat = euclidean_distance_matrix(text_embeddings.cpu().numpy(),
57 motion_embeddings.cpu().numpy())
58 matching_score_sum += dist_mat.trace()
59
60 argsmax = np.argsort(dist_mat, axis=1)
61 top_k_mat = calculate_top_k(argsmax, top_k=3)
62 top_k_count += top_k_mat.sum(axis=0)
63
64 all_size += text_embeddings.shape[0]
65
66 all_motion_embeddings.append(motion_embeddings.cpu().numpy())
67
68 all_motion_embeddings = np.concatenate(all_motion_embeddings, axis=0)
69 matching_score = matching_score_sum / all_size
70 R_precision = top_k_count / all_size
71 match_score_dict[motion_loader_name] = matching_score
72 R_precision_dict[motion_loader_name] = R_precision
73 activation_dict[motion_loader_name] = all_motion_embeddings
74
75 print(f'---> [{motion_loader_name}] Matching Score: {matching_score:.4f}')
76 print(f'---> [{motion_loader_name}] Matching Score: {matching_score:.4f}', file=file, flush=True)
77
78 line = f'---> [{motion_loader_name}] R_precision: '
79 for i in range(len(R_precision)):
80 line += '(top %d): %.4f ' % (i+1, R_precision[i])
81 print(line)
82 print(line, file=file, flush=True)
83
84 return match_score_dict, R_precision_dict, activation_dict
85
86
87def evaluate_fid(groundtruth_loader, activation_dict, file):
88 eval_dict = OrderedDict({})
89 gt_motion_embeddings = []
90 print('========== Evaluating FID ==========')
91 with torch.no_grad():
92 for idx, batch in enumerate(groundtruth_loader):
93 _, _, _, sent_lens, motions, m_lens, _ = batch
94 motion_embeddings = eval_wrapper.get_motion_embeddings(
95 motions=motions,
96 m_lens=m_lens
97 )
98 gt_motion_embeddings.append(motion_embeddings.cpu().numpy())
99 gt_motion_embeddings = np.concatenate(gt_motion_embeddings, axis=0)
100 gt_mu, gt_cov = calculate_activation_statistics(gt_motion_embeddings)
101
102 # print(gt_mu)
103 for model_name, motion_embeddings in activation_dict.items():
104 mu, cov = calculate_activation_statistics(motion_embeddings)
105 # print(mu)
106 fid = calculate_frechet_distance(gt_mu, gt_cov, mu, cov)
107 print(f'---> [{model_name}] FID: {fid:.4f}')
108 print(f'---> [{model_name}] FID: {fid:.4f}', file=file, flush=True)
109 eval_dict[model_name] = fid
110 return eval_dict
111
112
113def evaluate_diversity(activation_dict, file):
114 eval_dict = OrderedDict({})
115 print('========== Evaluating Diversity ==========')
116 for model_name, motion_embeddings in activation_dict.items():
117 diversity = calculate_diversity(motion_embeddings, diversity_times)
118 eval_dict[model_name] = diversity
119 print(f'---> [{model_name}] Diversity: {diversity:.4f}')
120 print(f'---> [{model_name}] Diversity: {diversity:.4f}', file=file, flush=True)
121 return eval_dict
122
123
124def evaluate_multimodality(mm_motion_loaders, file):
125 eval_dict = OrderedDict({})
126 print('========== Evaluating MultiModality ==========')
127 for model_name, mm_motion_loader in mm_motion_loaders.items():
128 mm_motion_embeddings = []
129 with torch.no_grad():
130 for idx, batch in enumerate(mm_motion_loader):
131 # (1, mm_replications, dim_pos)
132 motions, m_lens = batch
133 motion_embedings = eval_wrapper.get_motion_embeddings(motions[0], m_lens[0])
134 mm_motion_embeddings.append(motion_embedings.unsqueeze(0))
135 if len(mm_motion_embeddings) == 0:
136 multimodality = 0
137 else:
138 mm_motion_embeddings = torch.cat(mm_motion_embeddings, dim=0).cpu().numpy()
139 multimodality = calculate_multimodality(mm_motion_embeddings, mm_num_times)
140 print(f'---> [{model_name}] Multimodality: {multimodality:.4f}')
141 print(f'---> [{model_name}] Multimodality: {multimodality:.4f}', file=file, flush=True)
142 eval_dict[model_name] = multimodality
143 return eval_dict
144
145
146def get_metric_statistics(values):
147 mean = np.mean(values, axis=0)
148 std = np.std(values, axis=0)
149 conf_interval = 1.96 * std / np.sqrt(replication_times)
150 return mean, conf_interval
151
152
153def evaluation(log_file):
154 with open(log_file, 'w') as f:
155 all_metrics = OrderedDict({'Matching Score': OrderedDict({}),
156 'R_precision': OrderedDict({}),
157 'FID': OrderedDict({}),
158 'Diversity': OrderedDict({}),
159 'MultiModality': OrderedDict({})})
160 for replication in range(replication_times):
161 motion_loaders = {}
162 mm_motion_loaders = {}
163 motion_loaders['ground truth'] = gt_loader
164 for motion_loader_name, motion_loader_getter in eval_motion_loaders.items():
165 motion_loader, mm_motion_loader = motion_loader_getter()
166 motion_loaders[motion_loader_name] = motion_loader
167 mm_motion_loaders[motion_loader_name] = mm_motion_loader
168
169 print(f'==================== Replication {replication} ====================')
170 print(f'==================== Replication {replication} ====================', file=f, flush=True)
171 print(f'Time: {datetime.now()}')
172 print(f'Time: {datetime.now()}', file=f, flush=True)
173 mat_score_dict, R_precision_dict, acti_dict = evaluate_matching_score(motion_loaders, f)
174
175 print(f'Time: {datetime.now()}')
176 print(f'Time: {datetime.now()}', file=f, flush=True)
177 fid_score_dict = evaluate_fid(gt_loader, acti_dict, f)
178
179 print(f'Time: {datetime.now()}')
180 print(f'Time: {datetime.now()}', file=f, flush=True)
181 div_score_dict = evaluate_diversity(acti_dict, f)
182
183 print(f'Time: {datetime.now()}')
184 print(f'Time: {datetime.now()}', file=f, flush=True)
185 mm_score_dict = evaluate_multimodality(mm_motion_loaders, f)
186
187 print(f'!!! DONE !!!')
188 print(f'!!! DONE !!!', file=f, flush=True)
189
190 for key, item in mat_score_dict.items():
191 if key not in all_metrics['Matching Score']:
192 all_metrics['Matching Score'][key] = [item]
193 else:
194 all_metrics['Matching Score'][key] += [item]
195
196 for key, item in R_precision_dict.items():
197 if key not in all_metrics['R_precision']:
198 all_metrics['R_precision'][key] = [item]
199 else:
200 all_metrics['R_precision'][key] += [item]
201
202 for key, item in fid_score_dict.items():
203 if key not in all_metrics['FID']:
204 all_metrics['FID'][key] = [item]
205 else:
206 all_metrics['FID'][key] += [item]
207
208 for key, item in div_score_dict.items():
209 if key not in all_metrics['Diversity']:
210 all_metrics['Diversity'][key] = [item]
211 else:
212 all_metrics['Diversity'][key] += [item]
213
214 for key, item in mm_score_dict.items():
215 if key not in all_metrics['MultiModality']:
216 all_metrics['MultiModality'][key] = [item]
217 else:
218 all_metrics['MultiModality'][key] += [item]
219
220
221 # print(all_metrics['Diversity'])
222 for metric_name, metric_dict in all_metrics.items():
223 print('========== %s Summary ==========' % metric_name)
224 print('========== %s Summary ==========' % metric_name, file=f, flush=True)
225
226 for model_name, values in metric_dict.items():
227 # print(metric_name, model_name)
228 mean, conf_interval = get_metric_statistics(np.array(values))
229 # print(mean, mean.dtype)
230 if isinstance(mean, np.float64) or isinstance(mean, np.float32):
231 print(f'---> [{model_name}] Mean: {mean:.4f} CInterval: {conf_interval:.4f}')
232 print(f'---> [{model_name}] Mean: {mean:.4f} CInterval: {conf_interval:.4f}', file=f, flush=True)
233 elif isinstance(mean, np.ndarray):
234 line = f'---> [{model_name}]'
235 for i in range(len(mean)):
236 line += '(top %d) Mean: %.4f CInt: %.4f;' % (i+1, mean[i], conf_interval[i])
237 print(line)
238 print(line, file=f, flush=True)
239
240
241if __name__ == '__main__':
242 mm_num_samples = 100
243 mm_num_repeats = 30
244 mm_num_times = 10
245
246 diversity_times = 300
247 replication_times = 1
248 batch_size = 32
249 opt_path = sys.argv[1]
250 dataset_opt_path = opt_path
251
252 try:
253 device_id = int(sys.argv[2])
254 except:
255 device_id = 0
256 device = torch.device('cuda:%d' % device_id if torch.cuda.is_available() else 'cpu')
257 torch.cuda.set_device(device_id)
258
259 gt_loader, gt_dataset = get_dataset_motion_loader(dataset_opt_path, batch_size, device)
260 wrapper_opt = get_opt(dataset_opt_path, device)
261 eval_wrapper = EvaluatorModelWrapper(wrapper_opt)
262
263 opt = get_opt(opt_path, device)
264 encoder = build_models(opt, opt.dim_pose)
265 trainer = DDPMTrainer(opt, encoder)
266 eval_motion_loaders = {
267 'text2motion': lambda: get_motion_loader(
268 opt,
269 batch_size,
270 trainer,
271 gt_dataset,
272 mm_num_samples,
273 mm_num_repeats
274 )
275 }
276
277 log_file = './t2m_evaluation.log'
278 evaluation(log_file)
279 