deepgrove-team commited on
Commit
e8be321
·
verified ·
1 Parent(s): ad3e2a0

Speed up Maple decode with and without the native extension

Browse files
Files changed (1) hide show
  1. maple.py +190 -59
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 = w + head * HEAD_DIM;
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 theta = pos * inv_freq[p];
 
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", "w", "inv_freq", "pos_eps"],
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
- result[r]+=maple_qdot(w+row*(K/4)+k/4+lid*4,xt,(float)sc[row],(float)bi[row],sum);
 
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
- """Compact repeated affine metadata, without changing quantized weights."""
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
- sources = (p.weight, p.scales, p.biases)
 
471
  state = p.get("_maple_row_state")
472
- if state is not None and state["sources"].matches(sources):
 
 
 
 
 
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=[("K", k), ("N", n), ("GATHER", gather), ("SPLIT", split)],
 
 
 
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 so the pos/eps pair is always uniform.
637
- pos_eps = mx.array([float(offset), self._eps], dtype=mx.float32)
 
 
638
  return _qk_norm_rope_kernel(
639
- inputs=[qkv, self._qk_w, self._inv_freq, pos_eps],
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
- logits = _router_gemv_kernel(
1105
- inputs=[hn, router_weight],
1106
- grid=(2048, 1, 1),
1107
- threadgroup=(128, 1, 1),
1108
- output_shapes=[(1, 1, 256)],
1109
- output_dtypes=[mx.float32],
1110
- )[0]
 
 
 
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
- not hasattr(_maple_native, "ArraySnapshot")
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 not state["sources"].matches(sources)
 
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
- top = mx.argpartition(self.centroids(hv), kth=-self.n_probes, axis=-1)[
 
1425
  ..., -self.n_probes :
1426
- ] # [1, n_probes]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- sources = (lm_head.weight, lm_head.scales, lm_head.biases, self._force_ids)
1442
- if self._force_sources is None or not self._force_sources.matches(sources):
1443
- self._force_rows = (
1444
- lm_head.weight[self._force_ids],
1445
- lm_head.scales[self._force_ids],
1446
- lm_head.biases[self._force_ids],
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
- )[0]
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(