Entry for evaluation.
(self, step: int, stage: str = "val")
| 1100 | |
| 1101 | @torch.no_grad() |
| 1102 | def eval(self, step: int, stage: str = "val"): |
| 1103 | """Entry for evaluation.""" |
| 1104 | print("Running evaluation...") |
| 1105 | cfg = self.cfg |
| 1106 | device = self.device |
| 1107 | world_rank = self.world_rank |
| 1108 | world_size = self.world_size |
| 1109 | |
| 1110 | valloader = torch.utils.data.DataLoader( |
| 1111 | self.valset, batch_size=1, shuffle=False, num_workers=1 |
| 1112 | ) |
| 1113 | ellipse_time = 0 |
| 1114 | metrics = defaultdict(list) |
| 1115 | for i, data in enumerate(valloader): |
| 1116 | camtoworlds = data["camtoworld"].to(device) |
| 1117 | Ks = data["K"].to(device) |
| 1118 | pixels = data["image"].to(device) / 255.0 |
| 1119 | masks = data["mask"].to(device) if "mask" in data else None |
| 1120 | height, width = pixels.shape[1:3] |
| 1121 | |
| 1122 | torch.cuda.synchronize() |
| 1123 | tic = time.time() |
| 1124 | colors, _, _ = self.rasterize_splats( |
| 1125 | camtoworlds=camtoworlds, |
| 1126 | Ks=Ks, |
| 1127 | width=width, |
| 1128 | height=height, |
| 1129 | sh_degree=cfg.sh_degree, |
| 1130 | near_plane=cfg.near_plane, |
| 1131 | far_plane=cfg.far_plane, |
| 1132 | masks=masks, |
| 1133 | ) # [1, H, W, 3] |
| 1134 | torch.cuda.synchronize() |
| 1135 | ellipse_time += time.time() - tic |
| 1136 | |
| 1137 | colors = torch.clamp(colors, 0.0, 1.0) |
| 1138 | canvas_list = [pixels, colors] |
| 1139 | |
| 1140 | if world_rank == 0: |
| 1141 | # write images |
| 1142 | canvas = torch.cat(canvas_list, dim=2).squeeze(0).cpu().numpy() # side by side |
| 1143 | # canvas = canvas_list[1].squeeze(0).cpu().numpy() # signle image |
| 1144 | canvas = (canvas * 255).astype(np.uint8) |
| 1145 | imageio.imwrite( |
| 1146 | f"{self.render_dir}/{stage}_step{step}_{i:04d}.png", |
| 1147 | canvas, |
| 1148 | ) |
| 1149 | |
| 1150 | pixels_p = pixels.permute(0, 3, 1, 2) # [1, 3, H, W] |
| 1151 | colors_p = colors.permute(0, 3, 1, 2) # [1, 3, H, W] |
| 1152 | metrics["psnr"].append(self.psnr(colors_p, pixels_p)) |
| 1153 | metrics["ssim"].append(self.ssim(colors_p, pixels_p)) |
| 1154 | metrics["lpips"].append(self.lpips(colors_p, pixels_p)) |
| 1155 | if cfg.use_bilateral_grid: |
| 1156 | cc_colors = color_correct(colors, pixels) |
| 1157 | cc_colors_p = cc_colors.permute(0, 3, 1, 2) # [1, 3, H, W] |
| 1158 | metrics["cc_psnr"].append(self.psnr(cc_colors_p, pixels_p)) |
| 1159 |
no test coverage detected