MCPcopy Create free account
hub / github.com/DSL-Lab/StreamSplat / compute_losses

Method compute_losses

model/splat_model.py:69–166  ·  view source on GitHub ↗
(self, input_frames, target_depth, supv_masks, render_pkg)

Source from the content-addressed store, hash-verified

67 return decoder_out
68
69 def compute_losses(self, input_frames, target_depth, supv_masks, render_pkg):
70 output_frames = render_pkg["render"] # [B, V, C, H, W]
71 pred_depths = render_pkg["depth"]
72 depth_mask = render_pkg["alpha"] > 0.1 # [B, V, 1, H, W]
73 metrics = {}
74 with torch.no_grad():
75 B, V, C, H, W = input_frames.shape
76
77 # All frames
78 input_frames_all_256 = F.interpolate(input_frames.reshape(-1, 3, H, W), (256, 256), mode='bilinear', align_corners=False)
79 output_frames_all_256 = F.interpolate(output_frames.reshape(-1, 3, H, W), (256, 256), mode='bilinear', align_corners=False)
80 metrics['psnr'] = compute_psnr(input_frames_all_256, output_frames_all_256).mean()
81 metrics['ssim'] = compute_ssim(input_frames_all_256, output_frames_all_256).mean()
82 metrics['lpips'] = compute_lpips(input_frames_all_256 * 2 - 1, output_frames_all_256 * 2 - 1).mean()
83
84 metrics['input_frames'] = input_frames
85 metrics['pred_frames'] = output_frames
86
87 # Middle frames [1:-1]
88 if V > 2:
89 input_frames_middle = input_frames[:, 1:-1]
90 output_frames_middle = output_frames[:, 1:-1]
91
92 input_frames_middle_256 = F.interpolate(input_frames_middle.reshape(-1, 3, H, W), (256, 256), mode='bilinear', align_corners=False)
93 output_frames_middle_256 = F.interpolate(output_frames_middle.reshape(-1, 3, H, W), (256, 256), mode='bilinear', align_corners=False)
94
95 metrics['psnr_novel'] = compute_psnr(input_frames_middle_256, output_frames_middle_256).mean()
96 metrics['ssim_novel'] = compute_ssim(input_frames_middle_256, output_frames_middle_256).mean()
97 metrics['lpips_novel'] = compute_lpips(input_frames_middle_256 * 2 - 1, output_frames_middle_256 * 2 - 1).mean()
98 else:
99 metrics['psnr_novel'] = torch.tensor(0.0, device=input_frames.device)
100 metrics['ssim_novel'] = torch.tensor(0.0, device=input_frames.device)
101 metrics['lpips_novel'] = torch.tensor(0.0, device=input_frames.device)
102
103 # First and last frames (given views)
104 if V >= 1:
105 indices = [0]
106 if V > 1:
107 indices.append(V - 1)
108
109 input_frames_ends = input_frames[:, indices]
110 output_frames_ends = output_frames[:, indices]
111
112 input_frames_ends_256 = F.interpolate(input_frames_ends.reshape(-1, 3, H, W), (256, 256), mode='bilinear', align_corners=False)
113 output_frames_ends_256 = F.interpolate(output_frames_ends.reshape(-1, 3, H, W), (256, 256), mode='bilinear', align_corners=False)
114
115 metrics['psnr_given'] = compute_psnr(input_frames_ends_256, output_frames_ends_256).mean()
116 metrics['ssim_given'] = compute_ssim(input_frames_ends_256, output_frames_ends_256).mean()
117 metrics['lpips_given'] = compute_lpips(input_frames_ends_256 * 2 - 1, output_frames_ends_256 * 2 - 1).mean()
118 else: # Should not happen if V >= 1
119 metrics['psnr_given'] = torch.tensor(0.0, device=input_frames.device)
120 metrics['ssim_given'] = torch.tensor(0.0, device=input_frames.device)
121 metrics['lpips_given'] = torch.tensor(0.0, device=input_frames.device)
122
123 psnr = metrics['psnr']
124
125 if self.opt.skip:
126 # skip the frame_0

Callers 1

forwardMethod · 0.95

Calls 3

compute_psnrFunction · 0.90
compute_ssimFunction · 0.90
compute_lpipsFunction · 0.90

Tested by

no test coverage detected