|
| 1 | +# Vendored verbatim from ml-explore/mlx-lm PR #1217 |
| 2 | +# (mlx_lm/models/gated_delta.py) |
| 3 | +# Upstream license: MIT, Copyright © 2023 Apple Inc. |
| 4 | +# |
| 5 | +# Public surface used by cppmega_v4 Path E: |
| 6 | +# - gated_delta_update(...) -> (y, state) |
| 7 | +# - gated_delta_kernel(...) — low-level Metal kernel entry |
| 8 | +# - gated_delta_ops(...) — pure-MLX reference for prefill |
| 9 | +# - compute_g(A_log, a, dt_bias) — gate transform |
| 10 | +# No edits — pure copy. cppmega_v4/nn/path_e_adapter.py wraps our (g, beta) |
| 11 | +# API to this module's (a, b, A_log, dt_bias) API. |
| 12 | + |
| 13 | +from functools import partial |
| 14 | +from typing import Optional, Tuple |
| 15 | + |
| 16 | +import mlx.core as mx |
| 17 | +import mlx.nn as nn |
| 18 | + |
| 19 | + |
| 20 | +@partial(mx.compile, shapeless=True) |
| 21 | +def compute_g(A_log, a, dt_bias): |
| 22 | + return mx.exp(-mx.exp(A_log.astype(mx.float32)) * nn.softplus(a + dt_bias)) |
| 23 | + |
| 24 | + |
| 25 | +def _make_gated_delta_kernel(has_mask=False, vectorized=False): |
| 26 | + if not mx.metal.is_available(): |
| 27 | + return None |
| 28 | + mask_source = "mask[b_idx * T + t]" if has_mask else "true" |
| 29 | + |
| 30 | + # Configure g indexing based on whether gating is vectorized |
| 31 | + if vectorized: |
| 32 | + g_comment = "// g: [B, T, Hv, Dk]" |
| 33 | + g_setup = "auto g_ = g + (b_idx * T * Hv + hv_idx) * Dk;" |
| 34 | + g_access = "g_[s_idx]" |
| 35 | + g_advance = "g_ += Hv * Dk;" |
| 36 | + else: |
| 37 | + g_comment = "// g: [B, T, Hv]" |
| 38 | + g_setup = "auto g_ = g + b_idx * T * Hv;" |
| 39 | + g_access = "g_[hv_idx]" |
| 40 | + g_advance = "g_ += Hv;" |
| 41 | + |
| 42 | + source = f""" |
| 43 | + auto n = thread_position_in_grid.z; |
| 44 | + auto b_idx = n / Hv; |
| 45 | + auto hv_idx = n % Hv; |
| 46 | + auto hk_idx = hv_idx / (Hv / Hk); |
| 47 | + constexpr int n_per_t = Dk / 32; |
| 48 | +
|
| 49 | + // q, k: [B, T, Hk, Dk] |
| 50 | + auto q_ = q + b_idx * T * Hk * Dk + hk_idx * Dk; |
| 51 | + auto k_ = k + b_idx * T * Hk * Dk + hk_idx * Dk; |
| 52 | +
|
| 53 | + // v, y: [B, T, Hv, Dv] |
| 54 | + auto v_ = v + b_idx * T * Hv * Dv + hv_idx * Dv; |
| 55 | + y += b_idx * T * Hv * Dv + hv_idx * Dv; |
| 56 | +
|
| 57 | + auto dk_idx = thread_position_in_threadgroup.x; |
| 58 | + auto dv_idx = thread_position_in_grid.y; |
| 59 | +
|
| 60 | + // state_in, state_out: [B, Hv, Dv, Dk] |
| 61 | + auto i_state = state_in + (n * Dv + dv_idx) * Dk; |
| 62 | + auto o_state = state_out + (n * Dv + dv_idx) * Dk; |
| 63 | +
|
| 64 | + float state[n_per_t]; |
| 65 | + for (int i = 0; i < n_per_t; ++i) {{ |
| 66 | + auto s_idx = n_per_t * dk_idx + i; |
| 67 | + state[i] = static_cast<float>(i_state[s_idx]); |
| 68 | + }} |
| 69 | +
|
| 70 | + {g_comment} |
| 71 | + {g_setup} |
| 72 | + auto beta_ = beta + b_idx * T * Hv; |
| 73 | +
|
| 74 | + for (int t = 0; t < T; ++t) {{ |
| 75 | + if ({mask_source}) {{ |
| 76 | + float kv_mem = 0.0f; |
| 77 | + for (int i = 0; i < n_per_t; ++i) {{ |
| 78 | + auto s_idx = n_per_t * dk_idx + i; |
| 79 | + state[i] = state[i] * {g_access}; |
| 80 | + kv_mem += state[i] * k_[s_idx]; |
| 81 | + }} |
| 82 | + kv_mem = simd_sum(kv_mem); |
| 83 | +
|
| 84 | + auto delta = (v_[dv_idx] - kv_mem) * beta_[hv_idx]; |
| 85 | +
|
| 86 | + float out = 0.0f; |
| 87 | + for (int i = 0; i < n_per_t; ++i) {{ |
| 88 | + auto s_idx = n_per_t * dk_idx + i; |
| 89 | + state[i] = state[i] + k_[s_idx] * delta; |
| 90 | + out += state[i] * q_[s_idx]; |
| 91 | + }} |
| 92 | + out = simd_sum(out); |
| 93 | + if (thread_index_in_simdgroup == 0) {{ |
| 94 | + y[dv_idx] = static_cast<InT>(out); |
| 95 | + }} |
| 96 | + }} else {{ |
| 97 | + y[dv_idx] = static_cast<InT>(0); |
| 98 | + }} |
| 99 | + // Increment data pointers to next time step |
| 100 | + q_ += Hk * Dk; |
| 101 | + k_ += Hk * Dk; |
| 102 | + v_ += Hv * Dv; |
| 103 | + y += Hv * Dv; |
| 104 | + {g_advance} |
| 105 | + beta_ += Hv; |
| 106 | + }} |
| 107 | + for (int i = 0; i < n_per_t; ++i) {{ |
| 108 | + auto s_idx = n_per_t * dk_idx + i; |
| 109 | + o_state[s_idx] = static_cast<StT>(state[i]); |
| 110 | + }} |
| 111 | + """ |
| 112 | + inputs = ["q", "k", "v", "g", "beta", "state_in", "T"] |
| 113 | + if has_mask: |
| 114 | + inputs.append("mask") |
| 115 | + |
| 116 | + suffix = "" |
| 117 | + if vectorized: |
| 118 | + suffix += "_vec" |
| 119 | + if has_mask: |
| 120 | + suffix += "_mask" |
| 121 | + |
| 122 | + return mx.fast.metal_kernel( |
| 123 | + name=f"gated_delta_step{suffix}", |
| 124 | + input_names=inputs, |
| 125 | + output_names=["y", "state_out"], |
| 126 | + source=source, |
| 127 | + ) |
| 128 | + |
| 129 | + |
| 130 | +_gated_delta_kernel = _make_gated_delta_kernel(has_mask=False, vectorized=False) |
| 131 | +_gated_delta_kernel_masked = _make_gated_delta_kernel(has_mask=True, vectorized=False) |
| 132 | +_gated_delta_kernel_vec = _make_gated_delta_kernel(has_mask=False, vectorized=True) |
| 133 | +_gated_delta_kernel_vec_masked = _make_gated_delta_kernel( |
| 134 | + has_mask=True, vectorized=True |
| 135 | +) |
| 136 | + |
| 137 | + |
| 138 | +@mx.compile |
| 139 | +def _gated_delta_step_ops( |
| 140 | + q: mx.array, |
| 141 | + k: mx.array, |
| 142 | + v: mx.array, |
| 143 | + g: mx.array, |
| 144 | + beta: mx.array, |
| 145 | + state: mx.array, |
| 146 | + mask: Optional[mx.array] = None, |
| 147 | +) -> Tuple[mx.array, mx.array]: |
| 148 | + """ |
| 149 | + Ops-based reference implementation for a single recurrent step. |
| 150 | +
|
| 151 | + Shapes: |
| 152 | + - q, k: [B, H, Dk] |
| 153 | + - v: [B, H, Dv] |
| 154 | + - g: [B, H] or [B, H, Dk] |
| 155 | + - beta: [B, H] |
| 156 | + - state: [B, H, Dv, Dk] |
| 157 | + Returns: |
| 158 | + - y: [B, H, Dv] |
| 159 | + - new_state: [B, H, Dv, Dk] |
| 160 | + """ |
| 161 | + |
| 162 | + # Decay |
| 163 | + old_state = state |
| 164 | + if g.ndim == 2: |
| 165 | + decay = g[..., None, None] |
| 166 | + elif g.ndim == 3: |
| 167 | + decay = g[..., None, :] |
| 168 | + else: |
| 169 | + raise ValueError(f"Unsupported gating shape {g.shape}") |
| 170 | + state = state * decay |
| 171 | + kv_mem = (state * k[..., None, :]).sum(axis=-1) # [B, H, Dv] |
| 172 | + delta = (v - kv_mem) * beta[..., None] # [B, H, Dv] |
| 173 | + state = state + k[..., None, :] * delta[..., None] |
| 174 | + # Output projection along key dim with q |
| 175 | + y = (state * q[..., None, :]).sum(axis=-1) # [B, H, Dv] |
| 176 | + |
| 177 | + if mask is not None: |
| 178 | + mask = mx.expand_dims(mask, axis=(1, 2, 3)) |
| 179 | + state = mx.where(mask, state, old_state) |
| 180 | + return y.astype(q.dtype), state |
| 181 | + |
| 182 | + |
| 183 | +def gated_delta_kernel( |
| 184 | + q: mx.array, |
| 185 | + k: mx.array, |
| 186 | + v: mx.array, |
| 187 | + g: mx.array, |
| 188 | + beta: mx.array, |
| 189 | + state: mx.array, |
| 190 | + mask: Optional[mx.array] = None, |
| 191 | +) -> Tuple[mx.array, mx.array]: |
| 192 | + B, T, Hk, Dk = k.shape |
| 193 | + Hv, Dv = v.shape[2:] |
| 194 | + input_type = q.dtype |
| 195 | + state_type = state.dtype |
| 196 | + if g.ndim == 4: |
| 197 | + kernel = _gated_delta_kernel_vec |
| 198 | + inputs = [q, k, v, g, beta, state, T] |
| 199 | + if mask is not None: |
| 200 | + kernel = _gated_delta_kernel_vec_masked |
| 201 | + inputs.append(mask) |
| 202 | + else: |
| 203 | + kernel = _gated_delta_kernel |
| 204 | + inputs = [q, k, v, g, beta, state, T] |
| 205 | + if mask is not None: |
| 206 | + kernel = _gated_delta_kernel_masked |
| 207 | + inputs.append(mask) |
| 208 | + |
| 209 | + return kernel( |
| 210 | + inputs=inputs, |
| 211 | + template=[ |
| 212 | + ("InT", input_type), |
| 213 | + ("StT", state_type), |
| 214 | + ("Dk", Dk), |
| 215 | + ("Dv", Dv), |
| 216 | + ("Hk", Hk), |
| 217 | + ("Hv", Hv), |
| 218 | + ], |
| 219 | + grid=(32, Dv, B * Hv), |
| 220 | + threadgroup=(32, 4, 1), |
| 221 | + output_shapes=[(B, T, Hv, Dv), state.shape], |
| 222 | + output_dtypes=[input_type, state_type], |
| 223 | + ) |
| 224 | + |
| 225 | + |
| 226 | +def gated_delta_ops( |
| 227 | + q: mx.array, |
| 228 | + k: mx.array, |
| 229 | + v: mx.array, |
| 230 | + g: mx.array, |
| 231 | + beta: mx.array, |
| 232 | + state: Optional[mx.array] = None, |
| 233 | + mask: Optional[mx.array] = None, |
| 234 | +) -> Tuple[mx.array, mx.array]: |
| 235 | + """ |
| 236 | + Ops-based reference implementation for prompt prefill (sequential loop). |
| 237 | + Supports both scalar and vectorized gating. |
| 238 | +
|
| 239 | + Shapes: |
| 240 | + - q, k: [B, T, Hk, Dk] |
| 241 | + - v: [B, T, Hv, Dv] |
| 242 | + - g: [B, T, Hv] (scalar) or [B, T, Hv, Dk] (vectorized) |
| 243 | + - beta: [B, T, Hv] |
| 244 | + - state: [B, Hv, Dv, Dk] |
| 245 | + Returns: |
| 246 | + - y: [B, T, Hv, Dv] |
| 247 | + - state: [B, Hv, Dv, Dk] |
| 248 | + """ |
| 249 | + B, T, Hk, Dk = q.shape |
| 250 | + Hv, Dv = v.shape[-2:] |
| 251 | + if state is None: |
| 252 | + state = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32) |
| 253 | + |
| 254 | + if (repeat_factor := Hv // Hk) > 1: |
| 255 | + q = mx.repeat(q, repeat_factor, -2) |
| 256 | + k = mx.repeat(k, repeat_factor, -2) |
| 257 | + |
| 258 | + ys = [] |
| 259 | + for t in range(T): |
| 260 | + y, state = _gated_delta_step_ops( |
| 261 | + q[:, t], |
| 262 | + k[:, t], |
| 263 | + v[:, t], |
| 264 | + g[:, t], |
| 265 | + beta[:, t], |
| 266 | + state, |
| 267 | + None if mask is None else mask[:, t], |
| 268 | + ) |
| 269 | + ys.append(y) |
| 270 | + y = mx.stack(ys, axis=1) |
| 271 | + return y, state |
| 272 | + |
| 273 | + |
| 274 | +def gated_delta_update( |
| 275 | + q: mx.array, |
| 276 | + k: mx.array, |
| 277 | + v: mx.array, |
| 278 | + a: mx.array, |
| 279 | + b: mx.array, |
| 280 | + A_log: mx.array, |
| 281 | + dt_bias: mx.array, |
| 282 | + state: Optional[mx.array] = None, |
| 283 | + mask: Optional[mx.array] = None, |
| 284 | + use_kernel: bool = True, |
| 285 | + training: bool = False, |
| 286 | +) -> Tuple[mx.array, mx.array]: |
| 287 | + if training: |
| 288 | + # Chunked VJP path with O(T/chunk) autodiff graph — fits T≥2048 |
| 289 | + # on 36 GB Apple Silicon where the Python-ops path OOMs. |
| 290 | + # Metal backward kernel is 8–11× faster than the Python reference, |
| 291 | + # but only handles the unmasked GPU path with Dk%32==0 and Dv%4==0. |
| 292 | + Dk = q.shape[-1] |
| 293 | + Dv = v.shape[-1] |
| 294 | + can_use_metal = ( |
| 295 | + mx.metal.is_available() |
| 296 | + and mx.default_device() == mx.gpu |
| 297 | + and mask is None |
| 298 | + and Dk % 32 == 0 |
| 299 | + and Dv % 4 == 0 |
| 300 | + ) |
| 301 | + if can_use_metal: |
| 302 | + try: |
| 303 | + from .gated_delta_vjp_metal import gated_delta_update_vjp_metal |
| 304 | + |
| 305 | + return gated_delta_update_vjp_metal( |
| 306 | + q, k, v, a, b, A_log, dt_bias, state, mask |
| 307 | + ) |
| 308 | + except ImportError: |
| 309 | + pass |
| 310 | + from .gated_delta_vjp import gated_delta_update_vjp |
| 311 | + |
| 312 | + return gated_delta_update_vjp(q, k, v, a, b, A_log, dt_bias, state, mask) |
| 313 | + |
| 314 | + beta = mx.sigmoid(b) |
| 315 | + g = compute_g(A_log, a, dt_bias) |
| 316 | + if state is None: |
| 317 | + B, _, Hk, Dk = q.shape |
| 318 | + Hv, Dv = v.shape[-2:] |
| 319 | + state = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32) |
| 320 | + |
| 321 | + if not use_kernel or mx.default_device() != mx.gpu or not mx.metal.is_available(): |
| 322 | + return gated_delta_ops(q, k, v, g, beta, state, mask) |
| 323 | + return gated_delta_kernel(q, k, v, g, beta, state, mask) |
0 commit comments