(self, batch, batch_idx, dataloader_idx=0)
| 364 | |
| 365 | @rank_zero_only |
| 366 | def validation_step(self, batch, batch_idx, dataloader_idx=0): |
| 367 | batch: BatchedExample = self.data_shim(batch) |
| 368 | |
| 369 | if self.global_rank == 0: |
| 370 | print( |
| 371 | f"validation step {self.global_step}; " |
| 372 | f"scene = {batch['scene']}; " |
| 373 | f"context = {batch['context']['index'].tolist()}" |
| 374 | ) |
| 375 | |
| 376 | # Render Gaussians. |
| 377 | b, _, _, h, w = batch["target"]["image"].shape |
| 378 | assert b == 1 |
| 379 | visualization_dump = {} |
| 380 | gaussians = self.encoder( |
| 381 | batch["context"], |
| 382 | self.global_step, |
| 383 | visualization_dump=visualization_dump, |
| 384 | ) |
| 385 | output = self.decoder.forward( |
| 386 | gaussians, |
| 387 | batch["target"]["extrinsics"], |
| 388 | batch["target"]["intrinsics"], |
| 389 | batch["target"]["near"], |
| 390 | batch["target"]["far"], |
| 391 | (h, w), |
| 392 | "depth", |
| 393 | ) |
| 394 | rgb_pred = output.color[0] |
| 395 | depth_pred = vis_depth_map(output.depth[0]) |
| 396 | |
| 397 | # direct depth from gaussian means (used for visualization only) |
| 398 | gaussian_means = visualization_dump["depth"][0].squeeze() |
| 399 | if gaussian_means.shape[-1] == 3: |
| 400 | gaussian_means = gaussian_means.mean(dim=-1) |
| 401 | |
| 402 | # Compute validation metrics. |
| 403 | rgb_gt = batch["target"]["image"][0] |
| 404 | psnr = compute_psnr(rgb_gt, rgb_pred).mean() |
| 405 | self.log(f"val/psnr", psnr) |
| 406 | lpips = compute_lpips(rgb_gt, rgb_pred).mean() |
| 407 | self.log(f"val/lpips", lpips) |
| 408 | ssim = compute_ssim(rgb_gt, rgb_pred).mean() |
| 409 | self.log(f"val/ssim", ssim) |
| 410 | |
| 411 | # Construct comparison image. |
| 412 | context_img = inverse_normalize(batch["context"]["image"][0]) |
| 413 | context_img_depth = vis_depth_map(gaussian_means) |
| 414 | context = [] |
| 415 | for i in range(context_img.shape[0]): |
| 416 | context.append(context_img[i]) |
| 417 | context.append(context_img_depth[i]) |
| 418 | comparison = hcat( |
| 419 | add_label(vcat(*context), "Context"), |
| 420 | add_label(vcat(*rgb_gt), "Target (Ground Truth)"), |
| 421 | add_label(vcat(*rgb_pred), "Target (Prediction)"), |
| 422 | add_label(vcat(*depth_pred), "Depth (Prediction)"), |
| 423 | ) |
nothing calls this directly
no test coverage detected