Team Ai
Apppublic

mingyuan/MotionDiffuse

sourceHugging Facemitupdated 3y agoView on Hugging Face
69likes
evaluation.py279 linesDownload Raw Back to tools
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