Framework-agnostic version of `numpy.transpose` that will work on torch/TensorFlow/Jax tensors as well as NumPy arrays.
(array, axes=None)
| 608 | |
| 609 | |
| 610 | def transpose(array, axes=None): |
| 611 | """ |
| 612 | Framework-agnostic version of `numpy.transpose` that will work on torch/TensorFlow/Jax tensors as well as NumPy |
| 613 | arrays. |
| 614 | """ |
| 615 | if is_numpy_array(array): |
| 616 | return np.transpose(array, axes=axes) |
| 617 | elif is_torch_tensor(array): |
| 618 | return array.T if axes is None else array.permute(*axes) |
| 619 | elif is_tf_tensor(array): |
| 620 | import tensorflow as tf |
| 621 | |
| 622 | return tf.transpose(array, perm=axes) |
| 623 | elif is_jax_tensor(array): |
| 624 | import jax.numpy as jnp |
| 625 | |
| 626 | return jnp.transpose(array, axes=axes) |
| 627 | else: |
| 628 | raise ValueError(f"Type not supported for transpose: {type(array)}.") |
| 629 | |
| 630 | |
| 631 | def reshape(array, newshape): |