|
| 1 | +"""Experimental Metal GELU/gate kernel, using MLX's same erf implementation. |
| 2 | +
|
| 3 | +The vendored erf/expm1 helpers are from MLX v0.32.2, copyright Apple and |
| 4 | +Norbert Juffa. Their original notices are retained in vendor/ and MLX_LICENSE. |
| 5 | +Only this experiment uses them; the production package is unaffected. |
| 6 | +""" |
| 7 | + |
| 8 | +from pathlib import Path |
| 9 | + |
| 10 | +import mlx.core as mx |
| 11 | +import mlx.nn as nn |
| 12 | + |
| 13 | +VENDOR = Path(__file__).with_name("vendor") |
| 14 | +helpers = "\n".join( |
| 15 | + line |
| 16 | + for filename in ("mlx_expm1f.h", "mlx_erf.h") |
| 17 | + for line in (VENDOR / filename).read_text().splitlines() |
| 18 | + if not line.startswith(("#include", "#pragma")) |
| 19 | +) |
| 20 | +GELU_GATE = mx.fast.metal_kernel( |
| 21 | + name="laya_experimental_exact_gelu_gate", |
| 22 | + input_names=["inp"], |
| 23 | + output_names=["out"], |
| 24 | + header="namespace laya_erf { using namespace metal;\n" + helpers + "\n}", |
| 25 | + source=""" |
| 26 | + uint elem = thread_position_in_grid.x; |
| 27 | + uint row = elem / I; |
| 28 | + uint col = elem % I; |
| 29 | + T value = inp[row * (2 * I) + col]; |
| 30 | + T gate = inp[row * (2 * I) + I + col]; |
| 31 | + T scaled = value / T(1.4142135623730951); |
| 32 | + T erf_value = T(laya_erf::erf(float(scaled))); |
| 33 | + T gelu = T(T(value * T(T(1) + erf_value)) / T(2)); |
| 34 | + out[elem] = gelu * gate; |
| 35 | + """, |
| 36 | + compile_options={"math_mode": "safe"}, |
| 37 | +) |
| 38 | + |
| 39 | + |
| 40 | +def metal_gelu_gate(x): |
| 41 | + if x.dtype != mx.float16: |
| 42 | + raise ValueError("This research kernel has only been implemented and tested for FP16") |
| 43 | + if x.shape[-1] % 2: |
| 44 | + raise ValueError("GELU/gate input needs two equal-width branches") |
| 45 | + intermediate = x.shape[-1] // 2 |
| 46 | + shape = (*x.shape[:-1], intermediate) |
| 47 | + return GELU_GATE( |
| 48 | + inputs=[x], |
| 49 | + template=[("T", x.dtype), ("I", intermediate)], |
| 50 | + grid=(x.size // 2, 1, 1), |
| 51 | + threadgroup=(256, 1, 1), |
| 52 | + output_shapes=[shape], |
| 53 | + output_dtypes=[x.dtype], |
| 54 | + )[0] |
| 55 | + |
| 56 | + |
| 57 | +class MetalMLP(nn.Module): |
| 58 | + def __init__(self, layer): |
| 59 | + super().__init__() |
| 60 | + self.Wi, self.Wo = layer.Wi, layer.Wo |
| 61 | + |
| 62 | + def __call__(self, x): |
| 63 | + return self.Wo(metal_gelu_gate(self.Wi(x))) |
0 commit comments