MCPcopy Create free account
hub / github.com/Zhiyuan-R/Tiger-Diffusion / evaluate_gen

Function evaluate_gen

test_generation.py:405–438  ·  view source on GitHub ↗
(opt, ref_pcs, logger)

Source from the content-addressed store, hash-verified

403
404
405def evaluate_gen(opt, ref_pcs, logger):
406
407 if ref_pcs is None:
408 _, test_dataset = get_dataset(opt.dataroot, opt.npoints, opt.category, use_mask=False)
409 test_dataloader = torch.utils.data.DataLoader(test_dataset, batch_size=opt.batch_size,
410 shuffle=False, num_workers=int(opt.workers), drop_last=False)
411 ref = []
412 for data in tqdm(test_dataloader, total=len(test_dataloader), desc='Generating Samples'):
413 x = data['test_points']
414 m, s = data['mean'].float(), data['std'].float()
415
416 ref.append(x*s + m)
417
418 ref_pcs = torch.cat(ref, dim=0).contiguous()
419
420 logger.info("Loading sample path: %s"
421 % (opt.eval_path))
422 sample_pcs = torch.load(opt.eval_path).contiguous()
423
424 logger.info("Generation sample size:%s reference size: %s"
425 % (sample_pcs.size(), ref_pcs.size()))
426
427
428 # Compute metrics
429 results = compute_all_metrics(sample_pcs, ref_pcs, opt.batch_size)
430 results = {k: (v.cpu().detach().item()
431 if not isinstance(v, float) else v) for k, v in results.items()}
432
433 pprint(results)
434 logger.info(results)
435
436 jsd = JSD(sample_pcs.numpy(), ref_pcs.numpy())
437 pprint('JSD: {}'.format(jsd))
438 logger.info('JSD: {}'.format(jsd))
439
440
441

Callers 1

mainFunction · 0.85

Calls 2

compute_all_metricsFunction · 0.90
get_datasetFunction · 0.70

Tested by

no test coverage detected