MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / NestedTensor

Class NestedTensor

util/misc.py:369–442  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

367
368
369class NestedTensor(object):
370 def __init__(self, tensors, mask: Optional[Tensor]):
371 self.tensors = tensors
372 self.mask = mask
373 if mask == 'auto':
374 self.mask = torch.zeros_like(tensors).to(tensors.device)
375 if self.mask.dim() == 3:
376 self.mask = self.mask.sum(0).to(bool)
377 elif self.mask.dim() == 4:
378 self.mask = self.mask.sum(1).to(bool)
379 else:
380 raise ValueError(
381 'tensors dim must be 3 or 4 but {}({})'.format(
382 self.tensors.dim(), self.tensors.shape))
383
384 def imgsize(self):
385 res = []
386 for i in range(self.tensors.shape[0]):
387 mask = self.mask[i]
388 maxH = (~mask).sum(0).max()
389 maxW = (~mask).sum(1).max()
390 res.append(torch.Tensor([maxH, maxW]))
391 return res
392
393 def to(self, device):
394 # type: (Device) -> NestedTensor # noqa
395 cast_tensor = self.tensors.to(device)
396 mask = self.mask
397 if mask is not None:
398 assert mask is not None
399 cast_mask = mask.to(device)
400 else:
401 cast_mask = None
402 return NestedTensor(cast_tensor, cast_mask)
403
404 def to_img_list_single(self, tensor, mask):
405 assert tensor.dim() == 3, 'dim of tensor should be 3 but {}'.format(
406 tensor.dim())
407 maxH = (~mask).sum(0).max()
408 maxW = (~mask).sum(1).max()
409 img = tensor[:, :maxH, :maxW]
410 return img
411
412 def to_img_list(self):
413 """remove the padding and convert to img list
414 Returns:
415 [type]: [description]
416 """
417 if self.tensors.dim() == 3:
418 return self.to_img_list_single(self.tensors, self.mask)
419 else:
420 res = []
421 for i in range(self.tensors.shape[0]):
422 tensor_i = self.tensors[i]
423 mask_i = self.mask[i]
424 res.append(self.to_img_list_single(tensor_i, mask_i))
425 return res
426

Callers 9

forwardMethod · 0.90
prepare_targetsMethod · 0.90
forwardMethod · 0.90
prepare_targetsMethod · 0.90
forwardMethod · 0.90
forwardMethod · 0.90
toMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected