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,
)
| 167 | |
| 168 | |
| 169 | def 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 | |
| 194 | def compute_psnr( |
nothing calls this directly
no test coverage detected