Speed up Maple decode with and without the native extension
Browse files
maple.py
CHANGED
|
@@ -182,6 +182,7 @@ def _add_rms_norm_ok(dim, dtype, w, eps):
|
|
| 182 |
)
|
| 183 |
|
| 184 |
|
|
|
|
| 185 |
def _aggregate_add_rms_norm(h, r, scores, w, eps):
|
| 186 |
return _make_add_rms_norm_kernel(eps, aggregate=True)(
|
| 187 |
inputs=[h.reshape(-1), r.reshape(-1), scores.reshape(-1), w],
|
|
@@ -314,7 +315,7 @@ def _make_qk_norm_rope_kernel():
|
|
| 314 |
}
|
| 315 |
return;
|
| 316 |
}
|
| 317 |
-
const device T_* wh =
|
| 318 |
|
| 319 |
float ss = 0.0f;
|
| 320 |
for (int i = 0; i < per_lane; ++i) {
|
|
@@ -332,7 +333,8 @@ def _make_qk_norm_rope_kernel():
|
|
| 332 |
if (ROPE_DIM > 0 && j < ROPE_DIM) {
|
| 333 |
constexpr int rhalf = ROPE_DIM > 0 ? ROPE_DIM / 2 : 1;
|
| 334 |
int p = j < rhalf ? j : j - rhalf;
|
| 335 |
-
float
|
|
|
|
| 336 |
float c = metal::fast::cos(theta);
|
| 337 |
float s = metal::fast::sin(theta);
|
| 338 |
int j2 = j < rhalf ? j + rhalf : j - rhalf;
|
|
@@ -344,7 +346,7 @@ def _make_qk_norm_rope_kernel():
|
|
| 344 |
"""
|
| 345 |
return mx.fast.metal_kernel(
|
| 346 |
name="maple_qk_norm_rope",
|
| 347 |
-
input_names=["x", "
|
| 348 |
output_names=["out"],
|
| 349 |
source=source,
|
| 350 |
)
|
|
@@ -378,8 +380,8 @@ _row_qmv_kernel = mx.fast.metal_kernel(
|
|
| 378 |
uint expert=GATHER?ids[slot]:0;
|
| 379 |
uint row0=threadgroup_position_in_grid.x*8+sg*4;
|
| 380 |
const device uchar* w=(const device uchar*)weight+expert*N*(K/4);
|
| 381 |
-
const device bfloat16_t* sc=scales+expert*N;
|
| 382 |
-
const device bfloat16_t* bi=biases+expert*N;
|
| 383 |
const device bfloat16_t* xv=x+(SPLIT?slot*K:0);
|
| 384 |
float result[4]={0};
|
| 385 |
for(int k=0;k<K;k+=512) {
|
|
@@ -392,7 +394,8 @@ _row_qmv_kernel = mx.fast.metal_kernel(
|
|
| 392 |
}
|
| 393 |
for(int r=0;r<4;++r) {
|
| 394 |
uint row=row0+r;
|
| 395 |
-
|
|
|
|
| 396 |
}
|
| 397 |
}
|
| 398 |
for(int r=0;r<4;++r) {
|
|
@@ -415,8 +418,8 @@ _row_up_kernel = mx.fast.metal_kernel(
|
|
| 415 |
uint expert=ids[slot];
|
| 416 |
uint first=threadgroup_position_in_grid.x*(SGS*R)+sg*R;
|
| 417 |
const device uchar* w=(const device uchar*)weight+expert*1024*512;
|
| 418 |
-
const device bfloat16_t* sc=scales+expert*1024;
|
| 419 |
-
const device bfloat16_t* bi=biases+expert*1024;
|
| 420 |
float up[R]={0},gate[R]={0};
|
| 421 |
for(int k=0;k<2048;k+=512) {
|
| 422 |
const device bfloat16_t* xx=x+k+lid*16;
|
|
@@ -430,9 +433,9 @@ _row_up_kernel = mx.fast.metal_kernel(
|
|
| 430 |
for(int r=0;r<R;++r) {
|
| 431 |
uint row=first+r;
|
| 432 |
uint wi=row*512+k/4+lid*4;
|
| 433 |
-
uint si=row;
|
| 434 |
up[r]+=maple_qdot(w+wi,xt,(float)sc[si],(float)bi[si],sum);
|
| 435 |
-
gate[r]+=maple_qdot(w+wi+512*512,xt,(float)sc[si+512],(float)bi[si+512],sum);
|
| 436 |
}
|
| 437 |
}
|
| 438 |
for(int r=0;r<R;++r) {
|
|
@@ -453,10 +456,9 @@ _row_up_kernel = mx.fast.metal_kernel(
|
|
| 453 |
|
| 454 |
|
| 455 |
def _row_quantized_metadata(p):
|
| 456 |
-
"""
|
| 457 |
if (
|
| 458 |
mx.__version__ != "0.32.0"
|
| 459 |
-
or not hasattr(_maple_native, "ArraySnapshot")
|
| 460 |
or type(p) not in (nn.QuantizedLinear, QuantizedSwitchLinear)
|
| 461 |
or p.bits != 2
|
| 462 |
or p.group_size != 128
|
|
@@ -467,16 +469,16 @@ def _row_quantized_metadata(p):
|
|
| 467 |
or p.biases.dtype != mx.bfloat16
|
| 468 |
):
|
| 469 |
return None
|
| 470 |
-
|
|
|
|
| 471 |
state = p.get("_maple_row_state")
|
| 472 |
-
if
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 473 |
return state if state["supported"] else None
|
| 474 |
-
state = {
|
| 475 |
-
"sources": _maple_native.ArraySnapshot(sources),
|
| 476 |
-
"supported": False,
|
| 477 |
-
"qmv_ok": None,
|
| 478 |
-
}
|
| 479 |
-
p["_maple_row_state"] = state
|
| 480 |
k, n = p.scales.shape[-1] * 128, p.weight.shape[-2]
|
| 481 |
if (
|
| 482 |
k not in (512, 2048)
|
|
@@ -486,6 +488,23 @@ def _row_quantized_metadata(p):
|
|
| 486 |
or p.biases.shape != p.scales.shape
|
| 487 |
):
|
| 488 |
return None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 489 |
sc, bi = p.scales[..., :1], p.biases[..., :1]
|
| 490 |
same = mx.all(p.scales.view(mx.uint16) == sc.view(mx.uint16)) & mx.all(
|
| 491 |
p.biases.view(mx.uint16) == bi.view(mx.uint16)
|
|
@@ -507,7 +526,10 @@ def _row_qmv_arrays(x, weight, scales, biases, ids, gather, split):
|
|
| 507 |
shape = (8, 1, n) if gather else (*x.shape[:-1], n)
|
| 508 |
return _row_qmv_kernel(
|
| 509 |
inputs=[x, weight, scales, biases, ids],
|
| 510 |
-
template=[
|
|
|
|
|
|
|
|
|
|
| 511 |
grid=(n // 4 * 32, 8 if gather else 1, 1),
|
| 512 |
threadgroup=(64, 1, 1),
|
| 513 |
output_shapes=[shape],
|
|
@@ -520,6 +542,7 @@ def _row_experts_arrays(x, indices, uw, us, ub, dw, ds, db):
|
|
| 520 |
ids = indices.reshape(-1).astype(mx.uint32)
|
| 521 |
y = _row_up_kernel(
|
| 522 |
inputs=[x, uw, us, ub, ids],
|
|
|
|
| 523 |
grid=(4096, 8, 1),
|
| 524 |
threadgroup=(64, 1, 1),
|
| 525 |
output_shapes=[(8, 1, 512)],
|
|
@@ -630,17 +653,18 @@ class MapleAttention(nn.Module):
|
|
| 630 |
|
| 631 |
def _qk_fused(self, qkv, offset):
|
| 632 |
"""Normalize/rotate Q and K; copy any trailing V heads unchanged."""
|
| 633 |
-
self._ensure_qk_state()
|
| 634 |
-
|
| 635 |
# cache.offset is a Python int for a plain cache but an mx.array for
|
| 636 |
-
# the batched caches; coerce
|
| 637 |
-
pos_eps = mx.array(
|
|
|
|
|
|
|
| 638 |
return _qk_norm_rope_kernel(
|
| 639 |
-
inputs=[qkv, self.
|
| 640 |
template=[
|
| 641 |
("T_", qkv.dtype),
|
| 642 |
("HEAD_DIM", self.head_dim),
|
| 643 |
("ROPE_DIM", self.rope.dims if self.use_rope else 0),
|
|
|
|
| 644 |
("NQK", self.num_attention_heads + self.num_key_value_heads),
|
| 645 |
],
|
| 646 |
grid=(32, qkv.shape[0], 1),
|
|
@@ -1099,15 +1123,18 @@ _router_select_kernel = mx.fast.metal_kernel(
|
|
| 1099 |
|
| 1100 |
|
| 1101 |
@mx.compile
|
| 1102 |
-
def _norm_router_arrays(h, r, norm_weight, router_weight, eps):
|
| 1103 |
h, hn = _add_rms_norm(h, r, norm_weight, eps)
|
| 1104 |
-
|
| 1105 |
-
|
| 1106 |
-
|
| 1107 |
-
|
| 1108 |
-
|
| 1109 |
-
|
| 1110 |
-
|
|
|
|
|
|
|
|
|
|
| 1111 |
indices, scores = _router_select_kernel(
|
| 1112 |
inputs=[logits],
|
| 1113 |
grid=(128, 1, 1),
|
|
@@ -1125,7 +1152,7 @@ def _decode_norm_router(h, r, norm, gate):
|
|
| 1125 |
return hh, hn, indices, scores
|
| 1126 |
|
| 1127 |
if (
|
| 1128 |
-
|
| 1129 |
or type(norm) is not MapleRMSNorm
|
| 1130 |
or type(gate) is not MapleGate
|
| 1131 |
or h.shape != (1, 1, 2048)
|
|
@@ -1141,15 +1168,17 @@ def _decode_norm_router(h, r, norm, gate):
|
|
| 1141 |
or norm.eps <= 0
|
| 1142 |
):
|
| 1143 |
return fallback()
|
|
|
|
| 1144 |
sources = (norm.weight, gate.weight)
|
| 1145 |
state = gate.get("_maple_norm_router_state")
|
| 1146 |
if (
|
| 1147 |
state is None
|
| 1148 |
or state["eps"] != norm.eps
|
| 1149 |
-
or
|
|
|
|
| 1150 |
):
|
| 1151 |
state = {
|
| 1152 |
-
"sources": _maple_native.ArraySnapshot(sources),
|
| 1153 |
"eps": norm.eps,
|
| 1154 |
"ok": None,
|
| 1155 |
}
|
|
@@ -1158,7 +1187,7 @@ def _decode_norm_router(h, r, norm, gate):
|
|
| 1158 |
return fallback()
|
| 1159 |
|
| 1160 |
def fast():
|
| 1161 |
-
return _norm_router_arrays(h, r, norm.weight, gate.weight, norm.eps)
|
| 1162 |
|
| 1163 |
if state["ok"] is None:
|
| 1164 |
state["ok"] = False
|
|
@@ -1349,6 +1378,102 @@ def _flash_scatter(top, token_map, logits, force_ids, force_logits, vocab_size):
|
|
| 1349 |
)[0]
|
| 1350 |
|
| 1351 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1352 |
class FlashHead(nn.Module):
|
| 1353 |
"""Two-phase approximate lm_head for single-stream decode.
|
| 1354 |
|
|
@@ -1406,9 +1531,8 @@ class FlashHead(nn.Module):
|
|
| 1406 |
),
|
| 1407 |
}
|
| 1408 |
self._force_ids = mx.array(meta.get("force_tokens", []), dtype=mx.int32)
|
| 1409 |
-
self._force_rows = None
|
| 1410 |
-
self._force_sources = None
|
| 1411 |
self._scatter = None
|
|
|
|
| 1412 |
|
| 1413 |
def _scatter_reference(self, top, logits, force_logits, vocab_size):
|
| 1414 |
oids = self.token_map[top[0]].reshape(-1)
|
|
@@ -1421,9 +1545,26 @@ class FlashHead(nn.Module):
|
|
| 1421 |
|
| 1422 |
def __call__(self, h: mx.array, lm_head: nn.Module) -> mx.array:
|
| 1423 |
hv = h[:, -1, :]
|
| 1424 |
-
|
|
|
|
| 1425 |
..., -self.n_probes :
|
| 1426 |
-
]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1427 |
|
| 1428 |
logits = mx.gather_qmm(
|
| 1429 |
hv.reshape(1, 1, 1, 1, -1),
|
|
@@ -1438,27 +1579,17 @@ class FlashHead(nn.Module):
|
|
| 1438 |
|
| 1439 |
force_logits = logits[:0]
|
| 1440 |
if self._force_ids.size:
|
| 1441 |
-
|
| 1442 |
-
|
| 1443 |
-
|
| 1444 |
-
|
| 1445 |
-
|
| 1446 |
-
|
| 1447 |
-
)
|
| 1448 |
-
if hasattr(_maple_native, "ArraySnapshot"):
|
| 1449 |
-
mx.eval(*self._force_rows)
|
| 1450 |
-
self._force_sources = _maple_native.ArraySnapshot(sources)
|
| 1451 |
-
fw, fs, fb = self._force_rows
|
| 1452 |
-
force_logits = mx.quantized_matmul(
|
| 1453 |
-
hv,
|
| 1454 |
-
fw,
|
| 1455 |
-
scales=fs,
|
| 1456 |
-
biases=fb,
|
| 1457 |
transpose=True,
|
| 1458 |
group_size=lm_head.group_size,
|
| 1459 |
bits=lm_head.bits,
|
| 1460 |
mode=getattr(lm_head, "mode", "affine"),
|
| 1461 |
-
)
|
| 1462 |
vocab_size = lm_head.weight.shape[0]
|
| 1463 |
if mx.__version__ == "0.32.0" and self._force_ids.size <= 8:
|
| 1464 |
fast = lambda: _flash_scatter(
|
|
|
|
| 182 |
)
|
| 183 |
|
| 184 |
|
| 185 |
+
@mx.compile
|
| 186 |
def _aggregate_add_rms_norm(h, r, scores, w, eps):
|
| 187 |
return _make_add_rms_norm_kernel(eps, aggregate=True)(
|
| 188 |
inputs=[h.reshape(-1), r.reshape(-1), scores.reshape(-1), w],
|
|
|
|
| 315 |
}
|
| 316 |
return;
|
| 317 |
}
|
| 318 |
+
const device T_* wh = head < (uint)NQ ? qw : kw;
|
| 319 |
|
| 320 |
float ss = 0.0f;
|
| 321 |
for (int i = 0; i < per_lane; ++i) {
|
|
|
|
| 333 |
if (ROPE_DIM > 0 && j < ROPE_DIM) {
|
| 334 |
constexpr int rhalf = ROPE_DIM > 0 ? ROPE_DIM / 2 : 1;
|
| 335 |
int p = j < rhalf ? j : j - rhalf;
|
| 336 |
+
float inv = metal::exp2(-(float(p) / float(rhalf)) * pos_eps[2]);
|
| 337 |
+
float theta = pos * inv;
|
| 338 |
float c = metal::fast::cos(theta);
|
| 339 |
float s = metal::fast::sin(theta);
|
| 340 |
int j2 = j < rhalf ? j + rhalf : j - rhalf;
|
|
|
|
| 346 |
"""
|
| 347 |
return mx.fast.metal_kernel(
|
| 348 |
name="maple_qk_norm_rope",
|
| 349 |
+
input_names=["x", "qw", "kw", "pos_eps"],
|
| 350 |
output_names=["out"],
|
| 351 |
source=source,
|
| 352 |
)
|
|
|
|
| 380 |
uint expert=GATHER?ids[slot]:0;
|
| 381 |
uint row0=threadgroup_position_in_grid.x*8+sg*4;
|
| 382 |
const device uchar* w=(const device uchar*)weight+expert*N*(K/4);
|
| 383 |
+
const device bfloat16_t* sc=scales+expert*N*(GROUPED?K/128:1);
|
| 384 |
+
const device bfloat16_t* bi=biases+expert*N*(GROUPED?K/128:1);
|
| 385 |
const device bfloat16_t* xv=x+(SPLIT?slot*K:0);
|
| 386 |
float result[4]={0};
|
| 387 |
for(int k=0;k<K;k+=512) {
|
|
|
|
| 394 |
}
|
| 395 |
for(int r=0;r<4;++r) {
|
| 396 |
uint row=row0+r;
|
| 397 |
+
uint si=GROUPED?row*(K/128)+(k+lid*16)/128:row;
|
| 398 |
+
result[r]+=maple_qdot(w+row*(K/4)+k/4+lid*4,xt,(float)sc[si],(float)bi[si],sum);
|
| 399 |
}
|
| 400 |
}
|
| 401 |
for(int r=0;r<4;++r) {
|
|
|
|
| 418 |
uint expert=ids[slot];
|
| 419 |
uint first=threadgroup_position_in_grid.x*(SGS*R)+sg*R;
|
| 420 |
const device uchar* w=(const device uchar*)weight+expert*1024*512;
|
| 421 |
+
const device bfloat16_t* sc=scales+expert*1024*(GROUPED?16:1);
|
| 422 |
+
const device bfloat16_t* bi=biases+expert*1024*(GROUPED?16:1);
|
| 423 |
float up[R]={0},gate[R]={0};
|
| 424 |
for(int k=0;k<2048;k+=512) {
|
| 425 |
const device bfloat16_t* xx=x+k+lid*16;
|
|
|
|
| 433 |
for(int r=0;r<R;++r) {
|
| 434 |
uint row=first+r;
|
| 435 |
uint wi=row*512+k/4+lid*4;
|
| 436 |
+
uint si=GROUPED?row*16+(k+lid*16)/128:row;
|
| 437 |
up[r]+=maple_qdot(w+wi,xt,(float)sc[si],(float)bi[si],sum);
|
| 438 |
+
gate[r]+=maple_qdot(w+wi+512*512,xt,(float)sc[si+512*(GROUPED?16:1)],(float)bi[si+512*(GROUPED?16:1)],sum);
|
| 439 |
}
|
| 440 |
}
|
| 441 |
for(int r=0;r<R;++r) {
|
|
|
|
| 456 |
|
| 457 |
|
| 458 |
def _row_quantized_metadata(p):
|
| 459 |
+
"""Use compact row metadata with snapshots, live group metadata without."""
|
| 460 |
if (
|
| 461 |
mx.__version__ != "0.32.0"
|
|
|
|
| 462 |
or type(p) not in (nn.QuantizedLinear, QuantizedSwitchLinear)
|
| 463 |
or p.bits != 2
|
| 464 |
or p.group_size != 128
|
|
|
|
| 469 |
or p.biases.dtype != mx.bfloat16
|
| 470 |
):
|
| 471 |
return None
|
| 472 |
+
native = hasattr(_maple_native, "ArraySnapshot")
|
| 473 |
+
sources = (p.weight, p.scales, p.biases) if native else None
|
| 474 |
state = p.get("_maple_row_state")
|
| 475 |
+
if (
|
| 476 |
+
native
|
| 477 |
+
and state is not None
|
| 478 |
+
and state.get("sources") is not None
|
| 479 |
+
and state["sources"].matches(sources)
|
| 480 |
+
):
|
| 481 |
return state if state["supported"] else None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 482 |
k, n = p.scales.shape[-1] * 128, p.weight.shape[-2]
|
| 483 |
if (
|
| 484 |
k not in (512, 2048)
|
|
|
|
| 488 |
or p.biases.shape != p.scales.shape
|
| 489 |
):
|
| 490 |
return None
|
| 491 |
+
if not native:
|
| 492 |
+
if (
|
| 493 |
+
state is None
|
| 494 |
+
or state.get("sources") is not None
|
| 495 |
+
or (state.get("k"), state.get("n")) != (k, n)
|
| 496 |
+
):
|
| 497 |
+
state = dict(k=k, n=n, qmv_ok=None, zero=mx.array([0], mx.uint32), sources=None)
|
| 498 |
+
p["_maple_row_state"] = state
|
| 499 |
+
# Live arrays need no snapshot: no derived weight values are cached.
|
| 500 |
+
state.update(scales=p.scales, biases=p.biases)
|
| 501 |
+
return state
|
| 502 |
+
state = {
|
| 503 |
+
"sources": _maple_native.ArraySnapshot(sources),
|
| 504 |
+
"supported": False,
|
| 505 |
+
"qmv_ok": None,
|
| 506 |
+
}
|
| 507 |
+
p["_maple_row_state"] = state
|
| 508 |
sc, bi = p.scales[..., :1], p.biases[..., :1]
|
| 509 |
same = mx.all(p.scales.view(mx.uint16) == sc.view(mx.uint16)) & mx.all(
|
| 510 |
p.biases.view(mx.uint16) == bi.view(mx.uint16)
|
|
|
|
| 526 |
shape = (8, 1, n) if gather else (*x.shape[:-1], n)
|
| 527 |
return _row_qmv_kernel(
|
| 528 |
inputs=[x, weight, scales, biases, ids],
|
| 529 |
+
template=[
|
| 530 |
+
("K", k), ("N", n), ("GATHER", gather), ("SPLIT", split),
|
| 531 |
+
("GROUPED", scales.ndim == weight.ndim),
|
| 532 |
+
],
|
| 533 |
grid=(n // 4 * 32, 8 if gather else 1, 1),
|
| 534 |
threadgroup=(64, 1, 1),
|
| 535 |
output_shapes=[shape],
|
|
|
|
| 542 |
ids = indices.reshape(-1).astype(mx.uint32)
|
| 543 |
y = _row_up_kernel(
|
| 544 |
inputs=[x, uw, us, ub, ids],
|
| 545 |
+
template=[("GROUPED", us.ndim == uw.ndim)],
|
| 546 |
grid=(4096, 8, 1),
|
| 547 |
threadgroup=(64, 1, 1),
|
| 548 |
output_shapes=[(8, 1, 512)],
|
|
|
|
| 653 |
|
| 654 |
def _qk_fused(self, qkv, offset):
|
| 655 |
"""Normalize/rotate Q and K; copy any trailing V heads unchanged."""
|
|
|
|
|
|
|
| 656 |
# cache.offset is a Python int for a plain cache but an mx.array for
|
| 657 |
+
# the batched caches; coerce before constructing the kernel input.
|
| 658 |
+
pos_eps = mx.array(
|
| 659 |
+
[float(offset), self._eps, math.log2(self._rope_base)], dtype=mx.float32
|
| 660 |
+
)
|
| 661 |
return _qk_norm_rope_kernel(
|
| 662 |
+
inputs=[qkv, self.q_norm.weight, self.k_norm.weight, pos_eps],
|
| 663 |
template=[
|
| 664 |
("T_", qkv.dtype),
|
| 665 |
("HEAD_DIM", self.head_dim),
|
| 666 |
("ROPE_DIM", self.rope.dims if self.use_rope else 0),
|
| 667 |
+
("NQ", self.num_attention_heads),
|
| 668 |
("NQK", self.num_attention_heads + self.num_key_value_heads),
|
| 669 |
],
|
| 670 |
grid=(32, qkv.shape[0], 1),
|
|
|
|
| 1123 |
|
| 1124 |
|
| 1125 |
@mx.compile
|
| 1126 |
+
def _norm_router_arrays(h, r, norm_weight, router_weight, eps, portable=False):
|
| 1127 |
h, hn = _add_rms_norm(h, r, norm_weight, eps)
|
| 1128 |
+
if portable:
|
| 1129 |
+
logits = hn.astype(mx.float32) @ router_weight.astype(mx.float32).T
|
| 1130 |
+
else:
|
| 1131 |
+
logits = _router_gemv_kernel(
|
| 1132 |
+
inputs=[hn, router_weight],
|
| 1133 |
+
grid=(2048, 1, 1),
|
| 1134 |
+
threadgroup=(128, 1, 1),
|
| 1135 |
+
output_shapes=[(1, 1, 256)],
|
| 1136 |
+
output_dtypes=[mx.float32],
|
| 1137 |
+
)[0]
|
| 1138 |
indices, scores = _router_select_kernel(
|
| 1139 |
inputs=[logits],
|
| 1140 |
grid=(128, 1, 1),
|
|
|
|
| 1152 |
return hh, hn, indices, scores
|
| 1153 |
|
| 1154 |
if (
|
| 1155 |
+
mx.__version__ != "0.32.0"
|
| 1156 |
or type(norm) is not MapleRMSNorm
|
| 1157 |
or type(gate) is not MapleGate
|
| 1158 |
or h.shape != (1, 1, 2048)
|
|
|
|
| 1168 |
or norm.eps <= 0
|
| 1169 |
):
|
| 1170 |
return fallback()
|
| 1171 |
+
native = hasattr(_maple_native, "ArraySnapshot")
|
| 1172 |
sources = (norm.weight, gate.weight)
|
| 1173 |
state = gate.get("_maple_norm_router_state")
|
| 1174 |
if (
|
| 1175 |
state is None
|
| 1176 |
or state["eps"] != norm.eps
|
| 1177 |
+
or native != (state["sources"] is not None)
|
| 1178 |
+
or (native and not state["sources"].matches(sources))
|
| 1179 |
):
|
| 1180 |
state = {
|
| 1181 |
+
"sources": _maple_native.ArraySnapshot(sources) if native else None,
|
| 1182 |
"eps": norm.eps,
|
| 1183 |
"ok": None,
|
| 1184 |
}
|
|
|
|
| 1187 |
return fallback()
|
| 1188 |
|
| 1189 |
def fast():
|
| 1190 |
+
return _norm_router_arrays(h, r, norm.weight, gate.weight, norm.eps, not native)
|
| 1191 |
|
| 1192 |
if state["ok"] is None:
|
| 1193 |
state["ok"] = False
|
|
|
|
| 1378 |
)[0]
|
| 1379 |
|
| 1380 |
|
| 1381 |
+
# Exact bf16 top-k set: two radix histograms, with MLX's stable tie order.
|
| 1382 |
+
# One threadgroup owns selection and compaction; no cross-group synchronization.
|
| 1383 |
+
_head_select_kernel = mx.fast.metal_kernel(
|
| 1384 |
+
name="maple_head_select",
|
| 1385 |
+
input_names=["x"],
|
| 1386 |
+
output_names=["out"],
|
| 1387 |
+
header=r"""
|
| 1388 |
+
inline uint head_key(bfloat16_t x) {
|
| 1389 |
+
uint raw=as_type<ushort>(x), magnitude=raw&0x7fffu;
|
| 1390 |
+
if(magnitude>0x7f80u) return 0xffffu;
|
| 1391 |
+
if(magnitude<0x80u) return 0x8000u;
|
| 1392 |
+
return (raw&0x8000u)?((~raw)&0xffffu):(raw^0x8000u);
|
| 1393 |
+
}
|
| 1394 |
+
inline uint head_prefix(uint value,uint tid,threadgroup uint* groups) {
|
| 1395 |
+
uint lane=tid%32u,sg=tid/32u;
|
| 1396 |
+
uint prefix=simd_prefix_exclusive_sum(value);
|
| 1397 |
+
if(lane==31u) groups[sg]=prefix+value;
|
| 1398 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 1399 |
+
for(uint w=0;w<sg;++w) prefix+=groups[w];
|
| 1400 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 1401 |
+
return prefix;
|
| 1402 |
+
}
|
| 1403 |
+
""",
|
| 1404 |
+
source=r"""
|
| 1405 |
+
uint tid=thread_position_in_threadgroup.x;
|
| 1406 |
+
constexpr uint VALUES=(N+1023u)/1024u;
|
| 1407 |
+
uint keys[VALUES];
|
| 1408 |
+
threadgroup atomic_uint hist[256];
|
| 1409 |
+
threadgroup uint groups[32],cut[5];
|
| 1410 |
+
if(tid<256u) atomic_store_explicit(hist+tid,0u,memory_order_relaxed);
|
| 1411 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 1412 |
+
for(uint j=0;j<VALUES;++j) {
|
| 1413 |
+
uint i=tid*VALUES+j;
|
| 1414 |
+
keys[j]=i<N?head_key(x[i]):0u;
|
| 1415 |
+
uint digit=keys[j]>>8;
|
| 1416 |
+
uint matches=uint((simd_vote::vote_t)simd_ballot(i<N));
|
| 1417 |
+
for(uint b=0;b<8;++b) {
|
| 1418 |
+
bool bit=(digit&(1u<<b))!=0u;
|
| 1419 |
+
uint votes=uint((simd_vote::vote_t)simd_ballot(bit));
|
| 1420 |
+
matches &= bit?votes:~votes;
|
| 1421 |
+
}
|
| 1422 |
+
if(i<N && tid%32u==ctz(matches))
|
| 1423 |
+
atomic_fetch_add_explicit(hist+digit,popcount(matches),memory_order_relaxed);
|
| 1424 |
+
}
|
| 1425 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 1426 |
+
uint count=tid<256u?atomic_load_explicit(hist+tid,memory_order_relaxed):0u;
|
| 1427 |
+
uint prefix=head_prefix(count,tid,groups);
|
| 1428 |
+
if(tid<256u && prefix<=N-K && prefix+count>N-K) {cut[0]=tid;cut[1]=prefix;}
|
| 1429 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 1430 |
+
if(tid<256u) atomic_store_explicit(hist+tid,0u,memory_order_relaxed);
|
| 1431 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 1432 |
+
for(uint j=0;j<VALUES;++j) {
|
| 1433 |
+
if(tid*VALUES+j<N && (keys[j]>>8)==cut[0])
|
| 1434 |
+
atomic_fetch_add_explicit(hist+(keys[j]&255u),1u,memory_order_relaxed);
|
| 1435 |
+
}
|
| 1436 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 1437 |
+
count=tid<256u?atomic_load_explicit(hist+tid,memory_order_relaxed):0u;
|
| 1438 |
+
prefix=head_prefix(count,tid,groups);
|
| 1439 |
+
uint rank=N-K-cut[1];
|
| 1440 |
+
if(tid<256u && prefix<=rank && prefix+count>rank) {
|
| 1441 |
+
cut[2]=(cut[0]<<8)|tid;
|
| 1442 |
+
cut[3]=N-(cut[1]+prefix+count);
|
| 1443 |
+
cut[4]=count;
|
| 1444 |
+
}
|
| 1445 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 1446 |
+
uint greater=0u,equal=0u;
|
| 1447 |
+
for(uint j=0;j<VALUES;++j) if(tid*VALUES+j<N) {
|
| 1448 |
+
greater+=keys[j]>cut[2]; equal+=keys[j]==cut[2];
|
| 1449 |
+
}
|
| 1450 |
+
uint greater_prefix=head_prefix(greater,tid,groups);
|
| 1451 |
+
uint equal_prefix=head_prefix(equal,tid,groups);
|
| 1452 |
+
uint skip=cut[4]-(K-cut[3]);
|
| 1453 |
+
for(uint j=0;j<VALUES;++j) {
|
| 1454 |
+
uint i=tid*VALUES+j;
|
| 1455 |
+
if(i>=N) break;
|
| 1456 |
+
if(keys[j]>cut[2]) out[greater_prefix++]=i;
|
| 1457 |
+
else if(keys[j]==cut[2]) {
|
| 1458 |
+
if(equal_prefix>=skip) out[cut[3]+equal_prefix-skip]=i;
|
| 1459 |
+
equal_prefix++;
|
| 1460 |
+
}
|
| 1461 |
+
}
|
| 1462 |
+
""",
|
| 1463 |
+
)
|
| 1464 |
+
|
| 1465 |
+
|
| 1466 |
+
def _head_select(x, k):
|
| 1467 |
+
return _head_select_kernel(
|
| 1468 |
+
inputs=[x],
|
| 1469 |
+
template=[("N", x.size), ("K", k)],
|
| 1470 |
+
grid=(1024, 1, 1),
|
| 1471 |
+
threadgroup=(1024, 1, 1),
|
| 1472 |
+
output_shapes=[(1, k)],
|
| 1473 |
+
output_dtypes=[mx.uint32],
|
| 1474 |
+
)[0]
|
| 1475 |
+
|
| 1476 |
+
|
| 1477 |
class FlashHead(nn.Module):
|
| 1478 |
"""Two-phase approximate lm_head for single-stream decode.
|
| 1479 |
|
|
|
|
| 1531 |
),
|
| 1532 |
}
|
| 1533 |
self._force_ids = mx.array(meta.get("force_tokens", []), dtype=mx.int32)
|
|
|
|
|
|
|
| 1534 |
self._scatter = None
|
| 1535 |
+
self._select = None
|
| 1536 |
|
| 1537 |
def _scatter_reference(self, top, logits, force_logits, vocab_size):
|
| 1538 |
oids = self.token_map[top[0]].reshape(-1)
|
|
|
|
| 1545 |
|
| 1546 |
def __call__(self, h: mx.array, lm_head: nn.Module) -> mx.array:
|
| 1547 |
hv = h[:, -1, :]
|
| 1548 |
+
scores = self.centroids(hv)
|
| 1549 |
+
reference = lambda: mx.argpartition(scores, kth=-self.n_probes, axis=-1)[
|
| 1550 |
..., -self.n_probes :
|
| 1551 |
+
]
|
| 1552 |
+
if (
|
| 1553 |
+
mx.__version__ == "0.32.0"
|
| 1554 |
+
and mx.default_device() == mx.gpu
|
| 1555 |
+
and scores.dtype == mx.bfloat16
|
| 1556 |
+
and scores.ndim == 2
|
| 1557 |
+
and scores.shape[0] == 1
|
| 1558 |
+
and 0 < self.n_probes <= scores.size <= 8192
|
| 1559 |
+
):
|
| 1560 |
+
if self._select is None:
|
| 1561 |
+
self._select = _matches(
|
| 1562 |
+
lambda: (mx.sort(_head_select(scores, self.n_probes)),),
|
| 1563 |
+
lambda: (mx.sort(reference()),),
|
| 1564 |
+
)
|
| 1565 |
+
top = _head_select(scores, self.n_probes) if self._select else reference()
|
| 1566 |
+
else:
|
| 1567 |
+
top = reference()
|
| 1568 |
|
| 1569 |
logits = mx.gather_qmm(
|
| 1570 |
hv.reshape(1, 1, 1, 1, -1),
|
|
|
|
| 1579 |
|
| 1580 |
force_logits = logits[:0]
|
| 1581 |
if self._force_ids.size:
|
| 1582 |
+
force_logits = mx.gather_qmm(
|
| 1583 |
+
hv.reshape(1, 1, -1),
|
| 1584 |
+
lm_head.weight.reshape(lm_head.weight.shape[0], 1, -1),
|
| 1585 |
+
lm_head.scales.reshape(lm_head.weight.shape[0], 1, -1),
|
| 1586 |
+
lm_head.biases.reshape(lm_head.weight.shape[0], 1, -1),
|
| 1587 |
+
rhs_indices=self._force_ids,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1588 |
transpose=True,
|
| 1589 |
group_size=lm_head.group_size,
|
| 1590 |
bits=lm_head.bits,
|
| 1591 |
mode=getattr(lm_head, "mode", "affine"),
|
| 1592 |
+
).reshape(-1)
|
| 1593 |
vocab_size = lm_head.weight.shape[0]
|
| 1594 |
if mx.__version__ == "0.32.0" and self._force_ids.size <= 8:
|
| 1595 |
fast = lambda: _flash_scatter(
|