MCPcopy Create free account
hub / github.com/akira-l/SEEG / evaluate_testset

Function evaluate_testset

scripts/train.py:289–385  ·  view source on GitHub ↗
(test_data_loader, generator, loss_fn, embed_space_evaluator, args)

Source from the content-addressed store, hash-verified

287
288
289def 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)

Callers 5

mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90
train_epochsFunction · 0.70

Calls 8

updateMethod · 0.95
AverageMeterClass · 0.90
eval_embedFunction · 0.90
convert_dir_vec_to_poseFunction · 0.90
push_samplesMethod · 0.80
get_no_of_samplesMethod · 0.80
get_scoresMethod · 0.80
resetMethod · 0.45

Tested by

no test coverage detected