Similar to torch.max, but with mask
(input: torch.Tensor, mask: torch.BoolTensor, dim: int = None, keepdim: bool = False)
| 305 | |
| 306 | |
| 307 | def masked_max(input: torch.Tensor, mask: torch.BoolTensor, dim: int = None, keepdim: bool = False) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: |
| 308 | """Similar to torch.max, but with mask |
| 309 | """ |
| 310 | if dim is None: |
| 311 | return torch.where(mask, input, torch.tensor(-torch.inf, dtype=input.dtype, device=input.device)).max() |
| 312 | else: |
| 313 | return torch.where(mask, input, torch.tensor(-torch.inf, dtype=input.dtype, device=input.device)).max(dim=dim, keepdim=keepdim) |
| 314 | |
| 315 | |
| 316 | def bounding_rect(mask: torch.BoolTensor): |