Shuffle `arr` along `axis`. Args: arr: (*,) axis: Returns: (*,)
(arr: np.ndarray, axis: int, rng=None)
| 90 | |
| 91 | |
| 92 | def shuffle_along_axis(arr: np.ndarray, axis: int, rng=None) -> np.ndarray: |
| 93 | """ |
| 94 | Shuffle `arr` along `axis`. |
| 95 | Args: |
| 96 | arr: |
| 97 | (*,) |
| 98 | axis: |
| 99 | |
| 100 | Returns: |
| 101 | (*,) |
| 102 | """ |
| 103 | if rng is None: |
| 104 | rng = np.random |
| 105 | |
| 106 | idx = rng.rand(*arr.shape).argsort(axis=axis) |
| 107 | return np.take_along_axis(arr, idx, axis=axis) |