(x, w, b, eps)
| 8 | |
| 9 | |
| 10 | def layer_norm(x, w, b, eps): |
| 11 | ot = x.dtype |
| 12 | x = x.astype(mx.float32) |
| 13 | mu = mx.mean(x, -1, keepdims=True) |
| 14 | v = mx.var(x, -1, keepdims=True) |
| 15 | y = (x - mu) * mx.rsqrt(v + eps) |
| 16 | if w is not None: |
| 17 | y = y * w |
| 18 | if b is not None: |
| 19 | y = y + b |
| 20 | return y |
| 21 | |
| 22 | |
| 23 | def time_layer_norm(N, dt): |