Entry for evaluation.
(self, step: int, stage: str = "val")
| 1425 | |
| 1426 | @torch.no_grad() |
| 1427 | def eval(self, step: int, stage: str = "val"): |
| 1428 | """Entry for evaluation.""" |
| 1429 | print("Running evaluation...") |
| 1430 | cfg = self.cfg |
| 1431 | device = self.device |
| 1432 | world_rank = self.world_rank |
| 1433 | world_size = self.world_size |
| 1434 | |
| 1435 | valloader = torch.utils.data.DataLoader( |
| 1436 | self.valset, batch_size=1, shuffle=False, num_workers=1 |
| 1437 | ) |
| 1438 | ellipse_time = 0 |
| 1439 | metrics = defaultdict(list) |
| 1440 | for i, data in enumerate(valloader): |
| 1441 | camtoworlds = data["camtoworld"].to(device) |
| 1442 | Ks = data["K"].to(device) |
| 1443 | pixels = data["image"].to(device) / 255.0 |
| 1444 | masks = data["mask"].to(device) if "mask" in data else None |
| 1445 | height, width = pixels.shape[1:3] |
| 1446 | |
| 1447 | torch.cuda.synchronize() |
| 1448 | tic = time.time() |
| 1449 | colors, _, _ = self.rasterize_splats( |
| 1450 | camtoworlds=camtoworlds, |
| 1451 | Ks=Ks, |
| 1452 | width=width, |
| 1453 | height=height, |
| 1454 | sh_degree=cfg.sh_degree, |
| 1455 | near_plane=cfg.near_plane, |
| 1456 | far_plane=cfg.far_plane, |
| 1457 | masks=masks, |
| 1458 | ) # [1, H, W, 3] |
| 1459 | torch.cuda.synchronize() |
| 1460 | ellipse_time += time.time() - tic |
| 1461 | |
| 1462 | colors = torch.clamp(colors, 0.0, 1.0) |
| 1463 | canvas_list = [pixels, colors] |
| 1464 | |
| 1465 | if world_rank == 0: |
| 1466 | # write images |
| 1467 | canvas = torch.cat(canvas_list, dim=2).squeeze(0).cpu().numpy() # side by side |
| 1468 | # canvas = canvas_list[1].squeeze(0).cpu().numpy() # signle image |
| 1469 | canvas = (canvas * 255).astype(np.uint8) |
| 1470 | imageio.imwrite( |
| 1471 | f"{self.render_dir}/{stage}_step{step}_{i:04d}.png", |
| 1472 | canvas, |
| 1473 | ) |
| 1474 | |
| 1475 | pixels_p = pixels.permute(0, 3, 1, 2) # [1, 3, H, W] |
| 1476 | colors_p = colors.permute(0, 3, 1, 2) # [1, 3, H, W] |
| 1477 | metrics["psnr"].append(self.psnr(colors_p, pixels_p)) |
| 1478 | metrics["ssim"].append(self.ssim(colors_p, pixels_p)) |
| 1479 | metrics["lpips"].append(self.lpips(colors_p, pixels_p)) |
| 1480 | if cfg.use_bilateral_grid: |
| 1481 | cc_colors = color_correct(colors, pixels) |
| 1482 | cc_colors_p = cc_colors.permute(0, 3, 1, 2) # [1, 3, H, W] |
| 1483 | metrics["cc_psnr"].append(self.psnr(cc_colors_p, pixels_p)) |
| 1484 |
no test coverage detected