| 299 | |
| 300 | |
| 301 | class NestedTensor(object): |
| 302 | def __init__(self, tensors, mask: Optional[Tensor]): |
| 303 | self.tensors = tensors |
| 304 | self.mask = mask |
| 305 | if mask == 'auto': |
| 306 | self.mask = torch.zeros_like(tensors).to(tensors.device) |
| 307 | if self.mask.dim() == 3: |
| 308 | self.mask = self.mask.sum(0).to(bool) |
| 309 | elif self.mask.dim() == 4: |
| 310 | self.mask = self.mask.sum(1).to(bool) |
| 311 | else: |
| 312 | raise ValueError("tensors dim must be 3 or 4 but {}({})".format(self.tensors.dim(), self.tensors.shape)) |
| 313 | |
| 314 | def imgsize(self): |
| 315 | res = [] |
| 316 | for i in range(self.tensors.shape[0]): |
| 317 | mask = self.mask[i] |
| 318 | maxH = (~mask).sum(0).max() |
| 319 | maxW = (~mask).sum(1).max() |
| 320 | res.append(torch.Tensor([maxH, maxW])) |
| 321 | return res |
| 322 | |
| 323 | def to(self, device): |
| 324 | # type: (Device) -> NestedTensor # noqa |
| 325 | cast_tensor = self.tensors.to(device) |
| 326 | mask = self.mask |
| 327 | if mask is not None: |
| 328 | assert mask is not None |
| 329 | cast_mask = mask.to(device) |
| 330 | else: |
| 331 | cast_mask = None |
| 332 | return NestedTensor(cast_tensor, cast_mask) |
| 333 | |
| 334 | def to_img_list_single(self, tensor, mask): |
| 335 | assert tensor.dim() == 3, "dim of tensor should be 3 but {}".format(tensor.dim()) |
| 336 | maxH = (~mask).sum(0).max() |
| 337 | maxW = (~mask).sum(1).max() |
| 338 | img = tensor[:, :maxH, :maxW] |
| 339 | return img |
| 340 | |
| 341 | def to_img_list(self): |
| 342 | """remove the padding and convert to img list |
| 343 | |
| 344 | Returns: |
| 345 | [type]: [description] |
| 346 | """ |
| 347 | if self.tensors.dim() == 3: |
| 348 | return self.to_img_list_single(self.tensors, self.mask) |
| 349 | else: |
| 350 | res = [] |
| 351 | for i in range(self.tensors.shape[0]): |
| 352 | tensor_i = self.tensors[i] |
| 353 | mask_i = self.mask[i] |
| 354 | res.append(self.to_img_list_single(tensor_i, mask_i)) |
| 355 | return res |
| 356 | |
| 357 | @property |
| 358 | def device(self): |
no outgoing calls
no test coverage detected