MCPcopy Create free account
hub / github.com/PeizeSun/TransTrack / NestedTensor

Class NestedTensor

util/misc.py:329–354  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

327
328
329class NestedTensor(object):
330 def __init__(self, tensors, mask: Optional[Tensor]):
331 self.tensors = tensors
332 self.mask = mask
333
334 def to(self, device, non_blocking=False):
335 # type: (Device) -> NestedTensor # noqa
336 cast_tensor = self.tensors.to(device, non_blocking=non_blocking)
337 mask = self.mask
338 if mask is not None:
339 assert mask is not None
340 cast_mask = mask.to(device, non_blocking=non_blocking)
341 else:
342 cast_mask = None
343 return NestedTensor(cast_tensor, cast_mask)
344
345 def record_stream(self, *args, **kwargs):
346 self.tensors.record_stream(*args, **kwargs)
347 if self.mask is not None:
348 self.mask.record_stream(*args, **kwargs)
349
350 def decompose(self):
351 return self.tensors, self.mask
352
353 def __repr__(self):
354 return str(self.tensors)
355
356
357def setup_for_distributed(is_master):

Callers 9

forwardMethod · 0.90
forward_onceMethod · 0.90
forward_trainMethod · 0.90
forwardMethod · 0.90
forwardMethod · 0.90
forwardMethod · 0.90
forward_trainMethod · 0.90
toMethod · 0.85

Calls

no outgoing calls

Tested by 2

forwardMethod · 0.72
forwardMethod · 0.72