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

Function compute_mse

pointersect/inference/inference_utils.py:134–166  ·  view source on GitHub ↗

Compute the 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) R

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

Source from the content-addressed store, hash-verified

132
133
134def compute_mse(
135 arr: torch.Tensor,
136 ref: torch.Tensor,
137 ndim_b: int = None,
138 valid_mask: torch.Tensor = None,
139):
140 """
141 Compute the mean squared error between arr and ref.
142 The average is taken over the d_shape.
143
144 Args:
145 arr: (*b_shape, *d_shape)
146 ref: (*b_shape, *d_shape)
147 ndim_b:
148 number of dimension of b_shape. If None, = 0.
149 valid_mask: (*b_shape, *d_shape)
150
151 Returns:
152 mse: (*b_shape,)
153 """
154 if ndim_b is None:
155 ndim_b = 0
156
157 squared_error = (arr - ref) ** 2 # (*b, *d)
158 squared_error = squared_error.reshape(*(arr.shape[:ndim_b]), -1) # (*b, numel_d)
159 if valid_mask is None:
160 mse = squared_error.mean(dim=-1) # (*b,)
161 else:
162 valid_mask = valid_mask.view(
163 *(valid_mask.shape), *([1] * (arr.ndim - valid_mask.ndim))).expand_as(arr)
164 valid_mask = valid_mask.reshape(*(arr.shape[:ndim_b]), -1) # (*b, numel_d)
165 mse = (squared_error * valid_mask).sum(dim=-1) / valid_mask.sum(-1)
166 return mse
167
168
169def compute_rmse(

Callers 2

compute_rmseFunction · 0.85
compute_psnrFunction · 0.85

Calls 1

reshapeMethod · 0.45

Tested by

no test coverage detected