(x)
| 3041 | c = mx.array(c).astype(mx.float32) |
| 3042 | |
| 3043 | def hadamard_transform(x): |
| 3044 | return h @ x / mx.sqrt(x.shape[-1]) |
| 3045 | |
| 3046 | out = mx.vjp(hadamard_transform, [x], [c]) |
| 3047 | out_t = mx.vjp(mx.hadamard_transform, [x], [c]) |
no outgoing calls
no test coverage detected