(test_data_loader, generator, loss_fn, embed_space_evaluator, args)
| 287 | |
| 288 | |
| 289 | def evaluate_testset(test_data_loader, generator, loss_fn, embed_space_evaluator, args): |
| 290 | # to evaluation mode |
| 291 | generator.train(False) |
| 292 | |
| 293 | if embed_space_evaluator: |
| 294 | embed_space_evaluator.reset() |
| 295 | losses = AverageMeter('loss') |
| 296 | joint_mae = AverageMeter('mae_on_joint') |
| 297 | accel = AverageMeter('accel') |
| 298 | start = time.time() |
| 299 | |
| 300 | with torch.no_grad(): |
| 301 | for iter_idx, data in enumerate(test_data_loader): |
| 302 | time_seq, in_text, text_lengths, in_text_padded, _, target_vec, in_audio, in_spec, aux_info = data |
| 303 | batch_size = target_vec.size(0) |
| 304 | |
| 305 | in_text = in_text.to(device) |
| 306 | time_seq = time_seq.to(device) |
| 307 | in_text_padded = in_text_padded.to(device) |
| 308 | in_audio = in_audio.to(device) |
| 309 | in_spec = in_spec.to(device) |
| 310 | target = target_vec.to(device) |
| 311 | |
| 312 | # speaker input |
| 313 | speaker_model = utils.train_utils.get_speaker_model(generator) |
| 314 | if speaker_model: |
| 315 | vid_indices = [random.choice(list(speaker_model.word2index.values())) for _ in range(batch_size)] |
| 316 | vid_indices = torch.LongTensor(vid_indices).to(device) |
| 317 | else: |
| 318 | vid_indices = None |
| 319 | |
| 320 | pre_seq = target.new_zeros((target.shape[0], target.shape[1], target.shape[2] + 1)) |
| 321 | pre_seq[:, 0:args.n_pre_poses, :-1] = target[:, 0:args.n_pre_poses] |
| 322 | pre_seq[:, 0:args.n_pre_poses, -1] = 1 # indicating bit for constraints |
| 323 | pre_seq_partial = pre_seq[:, 0:args.n_pre_poses, :-1] |
| 324 | |
| 325 | if args.model == 'joint_embedding': |
| 326 | loss, out_dir_vec = eval_embed(in_text_padded, in_audio, pre_seq_partial, |
| 327 | target, generator, mode='speech') |
| 328 | elif args.model == 'gesture_autoencoder': |
| 329 | loss, _ = eval_embed(in_text_padded, in_audio, pre_seq_partial, target, generator) |
| 330 | elif args.model == 'seq2seq': |
| 331 | out_dir_vec = generator(in_text, text_lengths, target, None) |
| 332 | loss = loss_fn(out_dir_vec, target) |
| 333 | elif args.model == 'speech2gesture': |
| 334 | out_dir_vec = generator(in_spec, pre_seq_partial) |
| 335 | loss = loss_fn(out_dir_vec, target) |
| 336 | elif args.model == 'multimodal_context': |
| 337 | out_dir_vec, *_ = generator(pre_seq, time_seq, in_text_padded, in_audio, vid_indices) |
| 338 | loss = F.l1_loss(out_dir_vec, target) |
| 339 | else: |
| 340 | assert False |
| 341 | |
| 342 | losses.update(loss.item(), batch_size) |
| 343 | |
| 344 | if args.model != 'gesture_autoencoder': |
| 345 | if embed_space_evaluator: |
| 346 | embed_space_evaluator.push_samples(in_text_padded, in_audio, out_dir_vec, target) |
no test coverage detected