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

Function compute_l1

pointersect/inference/inference_utils.py:310–342  ·  view source on GitHub ↗

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,
)

Source from the content-addressed store, hash-verified

308
309
310def 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
345def compute_area(

Callers

nothing calls this directly

Calls 1

reshapeMethod · 0.45

Tested by

no test coverage detected