Compute average l1 distance between arr and ref. 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) Returns: err: (*b_shape,)
(
arr: torch.Tensor,
ref: torch.Tensor,
ndim_b: int = None,
valid_mask: torch.Tensor = None,
)
| 308 | |
| 309 | |
| 310 | def compute_l1( |
| 311 | arr: torch.Tensor, |
| 312 | ref: torch.Tensor, |
| 313 | ndim_b: int = None, |
| 314 | valid_mask: torch.Tensor = None, |
| 315 | ): |
| 316 | """ |
| 317 | Compute average l1 distance between arr and ref. |
| 318 | |
| 319 | Args: |
| 320 | arr: (*b_shape, *d_shape) |
| 321 | ref: (*b_shape, *d_shape) |
| 322 | ndim_b: |
| 323 | number of dimension of b_shape. If None, = 0. |
| 324 | valid_mask: (*b_shape, *d_shape) |
| 325 | |
| 326 | Returns: |
| 327 | err: (*b_shape,) |
| 328 | """ |
| 329 | if ndim_b is None: |
| 330 | ndim_b = 0 |
| 331 | |
| 332 | err = (arr - ref).abs() # (*b, *d) |
| 333 | err = err.reshape(*(arr.shape[:ndim_b]), -1) # (*b, numel_d) |
| 334 | if valid_mask is None: |
| 335 | err = err.mean(dim=-1) # (*b,) |
| 336 | else: |
| 337 | valid_mask = valid_mask.view( |
| 338 | *(valid_mask.shape), *([1] * (arr.ndim - valid_mask.ndim))).expand_as(arr) |
| 339 | valid_mask = valid_mask.reshape(*(arr.shape[:ndim_b]), -1) # (*b, numel_d) |
| 340 | err = (err * valid_mask).sum(dim=-1) / valid_mask.sum(-1) |
| 341 | |
| 342 | return err |
| 343 | |
| 344 | |
| 345 | def compute_area( |