MCPcopy Create free account
hub / github.com/JasonLSC/GSCodec_Studio / eval

Method eval

examples/simple_trainer_old.py:1427–1506  ·  view source on GitHub ↗

Entry for evaluation.

(self, step: int, stage: str = "val")

Source from the content-addressed store, hash-verified

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

Callers 3

trainMethod · 0.95
run_compressionMethod · 0.95
mainFunction · 0.95

Calls 4

rasterize_splatsMethod · 0.95
color_correctFunction · 0.90
updateMethod · 0.80
flushMethod · 0.80

Tested by

no test coverage detected