| 110 | } |
| 111 | |
| 112 | void time_strided_ops() { |
| 113 | int M = 50, N = 50, O = 50, P = 50; |
| 114 | auto a = mx::random::uniform({M, N, O, P}); |
| 115 | auto b = mx::random::uniform({M, N, O, P}); |
| 116 | auto device = mx::default_device(); |
| 117 | mx::eval(a, b); |
| 118 | TIMEM("non-strided", mx::add, a, b, device); |
| 119 | a = mx::transpose(a, {1, 0, 2, 3}); |
| 120 | b = mx::transpose(b, {3, 2, 0, 1}); |
| 121 | mx::eval(a, b); |
| 122 | TIMEM("strided", mx::add, a, b, device); |
| 123 | } |
| 124 | |
| 125 | void time_comparisons() { |
| 126 | int M = 1000, N = 100, K = 10; |