MCPcopy Create free account
hub / github.com/apple/ml-pointersect / compute_rmse

Function compute_rmse

pointersect/inference/inference_utils.py:169–191  ·  view source on GitHub ↗

Compute the root mean squared error between arr and ref. The average is taken over the d_shape. Args: arr: (*b_shape, *d_shape) ref: (*b_shape, *d_shape) ndim_b: number of dimension of b_shape. If None, = 0. valid_mask: (*b_shape, *d_shape)

(
        arr: torch.Tensor,
        ref: torch.Tensor,
        ndim_b: int = None,
        valid_mask: torch.Tensor = None,
)

Source from the content-addressed store, hash-verified

167
168
169def compute_rmse(
170 arr: torch.Tensor,
171 ref: torch.Tensor,
172 ndim_b: int = None,
173 valid_mask: torch.Tensor = None,
174):
175 """
176 Compute the root mean squared error between arr and ref.
177 The average is taken over the d_shape.
178
179 Args:
180 arr: (*b_shape, *d_shape)
181 ref: (*b_shape, *d_shape)
182 ndim_b:
183 number of dimension of b_shape. If None, = 0.
184 valid_mask: (*b_shape, *d_shape)
185
186 Returns:
187 mse: (*b_shape,)
188 """
189 mse = compute_mse(arr=arr, ref=ref, ndim_b=ndim_b, valid_mask=valid_mask) # (*b,)
190 rmse = mse ** 0.5 # (*b,)
191 return rmse
192
193
194def compute_psnr(

Callers

nothing calls this directly

Calls 1

compute_mseFunction · 0.85

Tested by

no test coverage detected