`np.cumsum` with equivalent implementation for torch. Args: a: input data to compute cumsum. axis: expected axis to compute cumsum. kwargs: if `a` is PyTorch Tensor, additional args for `torch.cumsum`, more details: https://pytorch.org/docs/stable/genera
(a: NdarrayOrTensor, axis=None, **kwargs)
| 317 | |
| 318 | |
| 319 | def cumsum(a: NdarrayOrTensor, axis=None, **kwargs) -> NdarrayOrTensor: |
| 320 | """ |
| 321 | `np.cumsum` with equivalent implementation for torch. |
| 322 | |
| 323 | Args: |
| 324 | a: input data to compute cumsum. |
| 325 | axis: expected axis to compute cumsum. |
| 326 | kwargs: if `a` is PyTorch Tensor, additional args for `torch.cumsum`, more details: |
| 327 | https://pytorch.org/docs/stable/generated/torch.cumsum.html. |
| 328 | |
| 329 | """ |
| 330 | |
| 331 | if isinstance(a, np.ndarray): |
| 332 | return np.cumsum(a, axis) # type: ignore |
| 333 | if axis is None: |
| 334 | return torch.cumsum(a[:], 0, **kwargs) |
| 335 | return torch.cumsum(a, dim=axis, **kwargs) |
| 336 | |
| 337 | |
| 338 | def isfinite(x: NdarrayOrTensor) -> NdarrayOrTensor: |
no outgoing calls
no test coverage detected
searching dependent graphs…