(opt, ref_pcs, logger)
| 403 | |
| 404 | |
| 405 | def 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 |
no test coverage detected