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

Method eval

examples/simple_trainer.py:1102–1181  ·  view source on GitHub ↗

Entry for evaluation.

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

Source from the content-addressed store, hash-verified

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

Callers 3

trainMethod · 0.95
run_compressionMethod · 0.95

Calls 4

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

Tested by

no test coverage detected