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

Function ssimse_loss

model/midas_loss.py:99–152  ·  view source on GitHub ↗

Scale and translation invariant MSE loss Args: pred_depth: predicted depth [B, 1, H, W] target_depth: ground truth depth [B, 1, H, W] mask: validity mask [B, 1, H, W] ignore_large_loss: threshold for filtering large losses eps: small constant for nume

(pred_depth, target_depth, mask=None, ignore_large_loss=0.0, eps=1e-8)

Source from the content-addressed store, hash-verified

97 return loss_val + utilization_loss
98
99def ssimse_loss(pred_depth, target_depth, mask=None, ignore_large_loss=0.0, eps=1e-8):
100 """
101 Scale and translation invariant MSE loss
102 Args:
103 pred_depth: predicted depth [B, 1, H, W]
104 target_depth: ground truth depth [B, 1, H, W]
105 mask: validity mask [B, 1, H, W]
106 ignore_large_loss: threshold for filtering large losses
107 eps: small constant for numerical stability
108 """
109 if mask is None:
110 mask = torch.ones_like(target_depth)
111
112 # Apply mask and flatten
113 mask = mask.float()
114 pred_depth = pred_depth * mask
115 target_depth = target_depth * mask
116
117 pred_depth = pred_depth.flatten(1) # [B, H*W]
118 target_depth = target_depth.flatten(1)
119 mask_flat = mask.flatten(1)
120
121 # Compute means on valid pixels only
122 valid_pixels = mask_flat.sum(1) + eps
123 gt_mean = (target_depth * mask_flat).sum(1) / valid_pixels
124 pred_mean = (pred_depth * mask_flat).sum(1) / valid_pixels
125
126 # Center the depths
127 pred_centered = pred_depth - pred_mean[:, None]
128 target_centered = target_depth - gt_mean[:, None]
129
130 # Compute scale factor (least squares)
131 numerator = (pred_centered * target_centered * mask_flat).sum(1)
132 denominator = (target_centered**2 * mask_flat).sum(1) + eps
133
134 # Clamp scale factor for stability
135 s = (numerator / denominator).clamp(-10, 10)
136 t = pred_mean - s * gt_mean
137
138 # Apply scale and translation
139 pred_aligned = s[:, None] * pred_depth + t[:, None]
140
141 # Compute loss on valid pixels only
142 delta = (pred_aligned - target_depth) * mask_flat
143
144 if ignore_large_loss > 0:
145 valid_mask = ((delta ** 2) < ignore_large_loss) & (mask_flat > 0)
146 if valid_mask.any():
147 delta = delta[valid_mask]
148 else:
149 return torch.tensor(0.0, device=pred_depth.device)
150
151 loss_val = (delta ** 2).mean()
152 return loss_val

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected