`np.any` with equivalent implementation for torch. For pytorch, convert to boolean for compatibility with older versions. Args: x: input array/tensor. axis: axis to perform `any` over. Returns: Return a contiguous flattened array/tensor.
(x: NdarrayOrTensor, axis: int | Sequence[int])
| 269 | |
| 270 | |
| 271 | def any_np_pt(x: NdarrayOrTensor, axis: int | Sequence[int]) -> NdarrayOrTensor: |
| 272 | """`np.any` with equivalent implementation for torch. |
| 273 | |
| 274 | For pytorch, convert to boolean for compatibility with older versions. |
| 275 | |
| 276 | Args: |
| 277 | x: input array/tensor. |
| 278 | axis: axis to perform `any` over. |
| 279 | |
| 280 | Returns: |
| 281 | Return a contiguous flattened array/tensor. |
| 282 | """ |
| 283 | if isinstance(x, np.ndarray): |
| 284 | return np.any(x, axis) # type: ignore |
| 285 | |
| 286 | # pytorch can't handle multiple dimensions to `any` so loop across them |
| 287 | axis = [axis] if not isinstance(axis, Sequence) else axis |
| 288 | for ax in axis: |
| 289 | try: |
| 290 | x = torch.any(x, ax) |
| 291 | except RuntimeError: |
| 292 | # older versions of pytorch require the input to be cast to boolean |
| 293 | x = torch.any(x.bool(), ax) |
| 294 | return x |
| 295 | |
| 296 | |
| 297 | def maximum(a: NdarrayOrTensor, b: NdarrayOrTensor) -> NdarrayOrTensor: |
no outgoing calls
no test coverage detected
searching dependent graphs…