Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 22 additions & 17 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -6,33 +6,38 @@ adhere to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).

## [Unreleased]

## [0.3.8]

Batched and shared-prefix decode attention: cascade, paged sparse, q8 KV
operands, routed-expert shed.

### Added
- `route_shed(indices, scores, slot_table)`: GPU-side routed-expert slot
remap plus residency shed for streamed MoE decode; non-resident experts
are shed and reported (miss ids and scores) without a host sync.
- `sdpa_decode_gqa` optional `starts` (int32 [B]): per-batch-row key start
offsets for left-padded batched KV caches; padded-out key chunks are
skipped, not staged.
- `sdpa_decode_gqa` q8 KV operands (affine wire, bits 8, group 64): batched
decode attends over quantized KV directly, dequantizing on the staged
tile; up to 1.9x/call at depth vs dequantize-then-attend.
- Env-gated small-M qmm experiment kernels (`KQ_QMM_SPLITK`,
`KQ_QMM_SPLITK_NAX`, `KQ_MV_EXT_SB`, `KQ_MV_EXT_NX`, `KQ_MV_EXT_HD`):
the NAX split-K path lifts the collapsed M9-16 band 65-76%; the rest
measured flat to negative on M5 and stay off by default.
- `sdpa_decode_gqa_cascade`: fused shared-prefix batched decode; one KV
walk serves the shared prefix for every batch row, private suffixes read
per row. 1.6-4.2x vs per-row calls at P 14k-32k, hd128/hd256.
- `sdpa_fa_verify` head_dim 64/128 tiles; `return_lse` on
`sdpa_decode_gqa` and `sdpa_fa_verify`.
- `sdpa_decode_gqa_paged`: page-gather decode over per-kv-head page lists
for top-k sparse attention, with `starts` for left-padded batch rows.
- Verify width (qL 1-8) on the cascade op: end-aligned causal over each
row's private slab with full shared-prefix visibility; `lse` gains the
qL axis.
- q8 KV operands (bits 8, group 64) on the cascade op, dequantized on the
staged tiles in both passes; bit-exact vs the fp16 cascade on
dequantized arrays (head_dim 512 declines).
- `sdpa_decode_gqa` optional `starts` (int32 [B]): per-batch-row key start
offsets for left-padded batched KV caches; padded-out key chunks are
skipped, not staged.
- `sdpa_decode_gqa` q8 KV operands (affine wire, bits 8, group 64): batched
decode attends over quantized KV directly, dequantizing on the staged
tile; up to 1.9x/call at depth vs dequantize-then-attend.
- `sdpa_decode_gqa_paged`: page-gather decode over per-kv-head page lists
for top-k sparse attention, with `starts` for left-padded batch rows.
- `sdpa_fa_verify` head_dim 64/128 tiles; `return_lse` on
`sdpa_decode_gqa` and `sdpa_fa_verify`.
- `route_shed(indices, scores, slot_table)`: GPU-side routed-expert slot
remap plus residency shed for streamed MoE decode; non-resident experts
are shed and reported (miss ids and scores) without a host sync.
- Env-gated small-M qmm experiment kernels (`KQ_QMM_SPLITK`,
`KQ_QMM_SPLITK_NAX`, `KQ_MV_EXT_SB`, `KQ_MV_EXT_NX`, `KQ_MV_EXT_HD`,
`KQ_MV_EXT_TS`): the NAX split-K path lifts the collapsed M9-16 band
65-76%; the rest measured flat to negative on M5 and stay off by default.

## [0.3.7]

Expand Down
26 changes: 25 additions & 1 deletion docs/kernels.md
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,15 @@ Tuning levers (defaults are right for normal use):
- `KQ_MV_EXT_NR` - `2` selects the two-rows-per-thread `mv_ext` variant (q6_k, M 5-12), which
halves activation cache traffic but measured no faster than the shipped kernels. Kept as a probe
for future silicon. Default `1` (shipped behavior).
- `KQ_QMM_SPLITK_NAX` - split-K on the NAX BM=32 tile; the value is the target slice count (`1` =
auto 32, `0` off). Lifts the collapsed M 9-16 band 65-76% on M5 Max. q6_k/q8_0, M <= 32; read
live per call. Default off.
- `KQ_QMM_SPLITK` - split-K for the plain small-M qmm (target slice count, `0` off). Measured flat
to negative on M5 Max; kept as a probe. K-quants plus q8_0, M <= 32. Default off.
- `KQ_MV_EXT_SB` / `KQ_MV_EXT_NX` / `KQ_MV_EXT_HD` / `KQ_MV_EXT_TS` - `mv_ext` activation-traffic
experiments: shuffle-broadcast (`1`), wide nxpsg (`16`/`32`), half-precision chunk dots (`1`),
threadgroup-staged activations (`1`). q6_k M 4-12 only. `HD` measured +4-5% at M 8; the rest flat
to negative on M5 Max. Kept as probes. Default off.

## MoE GLU

Expand Down Expand Up @@ -95,8 +104,18 @@ mechanism below.
fused vector path does not cover.
- **`sdpa_decode_gqa`** - decode/verify GQA tuned for long KV caches: the key axis splits into coarse
chunks streamed through threadgroup-staged K/V tiles shared by the GQA group, so device memory reads
the KV once per chunk.
the KV once per chunk. Optional `starts` (int32 `[B]`) restricts row b to keys `[starts[b], kL)` for
left-padded batches, skipping fully padded-out chunks. Optional affine q8 K/V operands (scales and
biases, bits 8, group 64) dequantize on the tile stage. `return_lse=True` adds per-row log-sum-exp.
- **`sdpa_decode_gqa_cascade`** - shared-prefix batched decode: every row attends one common prefix
plus its own private suffix. The prefix is walked once for all rows on the matrix-unit tile, private
suffixes run per row, one merge pass folds both; 1.6-4.2x over per-row calls at 14k-32k prefixes.
qL 1-8 (verify width, end-aligned causal); takes `starts` and the q8 operands on either region.
- **`sdpa_decode_gqa_paged`** - sparse page-gather decode: attends only the K/V pages listed per
(batch, kv-head), so cost tracks the selected keys rather than the cache length. The page unit is
the staged tile height (32 rows at head dim 64/128, 16 at 256, 8 at 512); takes `starts`.
- **`sdpa_fa_verify`** - speculative-verify attention on the matrix units for a GQA-folded query tile.
Head dims 64 through 512; `return_lse` as above.

## DeepSeek/GLM sparse attention (DSA)

Expand Down Expand Up @@ -148,3 +167,8 @@ The zero-copy arena buffers and shared-event stream primitives (`arena_alloc`, `
`event_wait`, `shared_event_*`, `zero_copy_view_count`, `verify_zero_copy_views`, `load_gguf`) support
a producer/consumer decode loop and are a separate subsystem; see
[docs/feeder/DESIGN.md](feeder/DESIGN.md).

- **`route_shed`** - routed-expert slot remap plus residency shed for streamed MoE decode: expert ids
map to arena slots through a resident-slot table, non-resident experts are shed with their gate
mass renormalized onto the kept ones, and the misses come back (ids and scores) for between-token
prestaging. No host sync.
2 changes: 1 addition & 1 deletion mlx_kquant/_version.py
Original file line number Diff line number Diff line change
@@ -1 +1 @@
__version__ = "0.3.7"
__version__ = "0.3.8"