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,
)
| 132 | |
| 133 | |
| 134 | def 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 | |
| 169 | def compute_rmse( |
no test coverage detected