minfer Metal Backend Design

How minfer runs the compute graph on an Apple GPU: the MetalBackend graph executor, the src/metal/ device/kernel layer, the src/metal/kernels/ shaders, the command-buffer rhythm, and the memory and safety rules that hold it together.

Status. Landed. The backend arrived with compute-graph Phase 3, and the G1–G6 wiring passes plus the objc2 migration brought it to parity with the pre-graph imperative path. Every mechanism described here is implemented in the tree. Baseline: HEAD = 62a0a3e (2026-09-14).

Provenance. New document. It mirrors the layout of docs/CUDA-BACKEND-DESIGN.md and is the design of record for the Metal backend; the optimization campaign and its measurements stay in docs/METAL_OPTIMIZATIONS.md.

Related records. docs/METAL_OPTIMIZATIONS.md is the optimization history and current-state ledger (including §0.1, the graph-path integration status), docs/METAL_OBJC-ECOSYSTEM.md and docs/METAL-OBJC2-MIGRATION-PLAN.md cover the objc2 crate migration, docs/LLAMA_METAL_E2E.md is the llama.cpp Metal reference, docs/GPU_SAFETY.md holds the hard safety rules (Metal sections), and docs/inference_e2e_walkthrough/14-metal-backend.md narrates the backend for a first-time reader. The graph contract this backend implements is docs/COMPUTE-GRAPH-DESIGN.md §3.5/§7.


1. Design Goals and Outcome

1.1 Goal

Implement MetalBackend (src/graph/metal_backend.rs) as the macOS backend of the compute graph by wiring the existing per-op kernels on MpsState (src/metal/) — no new kernels required — so that on Apple Silicon the whole per-layer chain runs on the GPU through the standard build → assign → fuse → alloc → execute pipeline, with the same correctness contract as CPU: backend placement is decided at build time, kernel-invariant violations return Err, and there is never a silent mid-run fallback.

1.2 Outcome

GoalLanded outcomeEvidence
MetalBackend implements the Backend trait over MpsStateFull trait: shared-memory buffer pool, weight_buf offsets, per-op dispatch, direct host views, split-scoped command buffer§2, §4.1–§4.2
One command buffer per split, submitted at boundariescb() creates it lazily, synchronize() submits it; Drop flushes a pending one§2.4, §4.8
Per-node placement decided at build timesupports_op + the model-level weights_on_gpu all-or-nothing gate feed CParams.gpu; the scheduler splits the graph§4.3, §4.6
Zero-copy weightsGGUF parts are wrapped with newBufferWithBytesNoCopy; weights are (buffer, byte offset) pairs; a warm-up read moves the ~44 ms first-touch page cost to model load§2.3
Attention dispatch matches the old pathG1 wires nt==1 and nt>1 to the flash/split/parallel/classic kernels with the same gates and env vars§4.4, §4.6
Decode fusions on MetalG4 FusedQKV, G5 FusedFFN, G6 FusedQkvNorm (Qwen3) are built as single nodes§4.4, §4.6
Same correctness gates as CPUPer-op parity tests, cross-backend copy tests, model-level CPU-vs-Metal logits / greedy equality, kernel isolation tests§7
PerformanceGraph path at or above the old imperative path, re-measured 2026-10-06 at 6b95763 on macbook (macOS 27.0.1, Apple M4 Pro) (minfer bench -p <P> -n 128 -r 3): 0.5B Q4_0 decode 306.19 ± 1.00 tok/s / prefill pp440 6249.60 ± 10.52 tok/s; 7B Q4_K_M decode 48.52 ± 0.18 tok/s / prefill pp206 406.51 ± 0.83 tok/s; Qwen3-4B Q4_K_M decode 74.11 ± 0.19 tok/s (llama-Metal 79.7). Metal correctness: the external oracle graph_metal_matches_llama_reference reproduces the pinned greedy prefix, the model-level graph_metal_matches_cpu_logits compares a Layers(0) CPU engine against the full-plan Metal engine (restored by #324), and the per-op metal_*_matches_cpu gates are greenMETAL_OPTIMIZATIONS.md §0.1

1.3 Non-goals

  • A whole-layer layer_gpu fast path. The pre-graph imperative path was removed; src/metal/ is the per-op device/kernel layer only (the name survives solely in the legacy CUDA code).
  • f16/bf16 activations. Graph activations are f32; the f16 story is the KV cache (MINFER_CACHE_TYPE) and the flash-attention f16-KV kernel variants.
  • New kernels for the graph. The graph path is wiring plus the G2 rms_norm_256 selection and the decode *_off kernel variants that already existed.
  • Multi-GPU / device selection. One integrated GPU per Mac.
  • Training or fine-tuning.
TopicWhere
Optimization history, current state, graph-path status, env gatesdocs/METAL_OPTIMIZATIONS.md
objc 0.2 → objc2 crate migration (phases, gotchas, checklist)docs/METAL-OBJC2-MIGRATION-PLAN.md, docs/METAL_OBJC-ECOSYSTEM.md
llama.cpp Metal (MPS) end-to-end path (reference baseline)docs/LLAMA_METAL_E2E.md
GPU safety rules (barriers, bounded submit, capture windows)docs/GPU_SAFETY.md
Graph contract (IR, allocator, scheduler, backend trait)docs/COMPUTE-GRAPH-DESIGN.md
Beginner narrative of this backenddocs/inference_e2e_walkthrough/14-metal-backend.md
Kernel-level analysesdocs/metal-inference-analysis.md, docs/multi-token-kernel-analysis.md
Qwen3-4B vs llama.cpp (Metal)docs/PERF-QWEN3-4B-VS-LLAMACPP.md

2. Architecture at a Glance

2.1 Three layers

LayerFileRole
Graph executorsrc/graph/metal_backend.rsImplements Backend: shared-memory buffer pool, name→buffer-offset weight resolution, per-op dispatch, split-scoped MpsCommandBuffer, capture staging, error contract
Device/kernel layersrc/metal/MpsState singleton: device/queue/library init, zero-copy weight registry, MpsCommandBuffer (encoder, barriers, submit), and one Rust method per op/kernel
Shaderssrc/metal/kernels/The Metal Shading Language kernels (norms, matmul tiers, attention variants, elementwise, KV store, get_rows, fused epilogues)
Build chainbuild.rsRuntime shader compilation (default) or a precompiled metallib (MINFER_METALLIB_FILE/_PATH); the objc2 framework links

The split follows CUDA: src/metal/ is the only place that touches Objective-C/Metal APIs, and metal_backend.rs is the only place that knows about graph nodes.

2.2 The MetalBackend surface

#![allow(unused)]
fn main() {
pub struct MetalBackend {
    state: &'static MpsState,                 // process-wide MPS singleton
    pool: Vec<MetalBuffer>,                   // f32-element pool: id -> shared MTLBuffer
    free: Vec<usize>,                         // exact-byte-length free list
    staging: Vec<MetalBuffer>,                // trace/viz capture staging (blit targets)
    free_staging: Vec<usize>,
    cb_ptr: *mut MpsCommandBuffer<'static>,   // pending command buffer (null = none)
}
}

Pool buffers are StorageModeShared MTLBuffers — host and GPU see the same memory — so read_host/write_host are direct memory views and a cross-backend copy is a host round trip, not a staged transfer. The command buffer is stored as a leaked box pointer because MpsCommandBuffer is !Send/!Sync; all access happens sequentially through &self/&mut self, and the struct is unsafe impl Send/Sync on that basis.

2.3 Weight residency and registry

  • Zero-copy parts. MpsState::register_part wraps an mmap'd GGUF part with newBufferWithBytesNoCopy (StorageModeShared). The base must be page-aligned (16 KiB on Apple Silicon); mmap guarantees it, and the code debug_assert!s it. #39 audited the remaining debug_assert!s and kept this one deliberately: it is an OS contract on a stdlib call (mmap's guarantee), not a shape or kernel invariant, and register_part returns () so there is no Result for the value to travel in — the one asymmetry the release-build rule allows.
  • Per-weight offsets. register_weight(name, data) records (part buffer, byte offset); the executor resolves any weight through state.weight_buf(name) -> Option<(MetalBuffer, u64)> and passes the offset to the kernel. MINFER_WEIGHT_COPY=1 forces a copy per weight for A/B.
  • Load-time warm-up. The first GPU access to file-backed (mmap) pages costs a one-time page/TLB setup (~44 ms measured); a dummy full-buffer read at model load moves that cost out of the first prefill (METAL_OPTIMIZATIONS.md §0 Done #39).
  • KV element type (per engine, #44 part (b)). models::load_model_configured resolves MINFER_CACHE_TYPE once (kvformat::auto_device_format picks f16 when n_layers × n_kv_embd >= 8192 — the 7B class, measured ~−1 ms/token at 2K context — and f32 otherwise; f16 measured ~3% slower on the 0.5B, dispatch-latency-bound). The answer is stored on MetalBackend as kv_format, stamped through GraphAllocator::set_kv_format, and passed as an explicit f16 argument to every store/attention/fused decode. There is no process-wide Metal tag any more (ADR-0005, ADR-0006): two engines with different dims hold different layouts in one process. A third value, q8_0, is the packed cache the CPU, CUDA and Metal kernels read; Metal's packed path is enabled (#310): READS_PACKED_KV is true, two read mechanisms cover every shape (mechanism A native packed decode, mechanism B an f32 staging window for the fast prefill/window families), and the classic packed kernels stay the fallback — MINFER_CACHE_TYPE=q8_0 loads and runs. f16 remains the device default.

2.4 Command buffers and submission

One MpsCommandBuffer is kept for the current split:

  • cb() creates it on the first op of a split (a leaked box, so the returned reference is 'static and callers can still touch the pool).
  • Every dispatch helper ends with memoryBarrierWithScope(MTLBarrierScope::Buffers), so kernels in one command buffer see each other's writes.
  • synchronize() submits it; the scheduler calls that at split boundaries and once at the end.
  • Drop submits a pending buffer so an unterminated encoder can never be left behind.
  • MpsCommandBuffer::submit() commits with a dispatch-semaphore completion handler (avoiding the ~20 ms scheduler wakeup of waitUntilCompleted), waits with a bound and checks the command buffer status. On a failure it returns Err; the backend expects it, and device-configuration problems go through gpu_abort instead.

2.5 Legacy surface

The pre-graph whole-layer layer_gpu path was removed when the graph became the default: there is no layer_gpu function in src/metal/ (the name survives only in the legacy CUDA code), and src/graph/metal_backend.rs is the only live consumer of the device layer. What remains is a #[allow(dead_code)] block in src/metal/ holding the old-forward scaffolding and a few methods kept for tests (e.g. matmul_on_gpu_buf); the loaders, the graph backend and the kernel tests are the live callers. (The legacy KVCache type in src/cache.rs was likewise unused by the graph path; KV lives in the allocator's persistent regions, and the type — with the ModelDef::forward argument that kept it alive — was deleted in #252.)


3. llama.cpp Metal Reference Map

docs/LLAMA_METAL_E2E.md documents llama.cpp's Metal path end to end. What minfer borrowed, and what it deliberately does not do:

llama.cpp Metal conceptminfer analogStatus
Backend interface + scheduler splitsGraph Backend trait + split_graph; per-node execute_nodeBorrowed, reshaped
Multi-command-buffer scheme with status trackingOne command buffer per split; submit at boundariesSimplified
Per-op encoder + explicit memory barriersOne compute encoder per split, memoryBarrierWithScope(Buffers) after every dispatchBorrowed
MUL_MAT three-way kernel selection (matrix/tile/vector)quant_matmul_f32_on_gpu_buf tiers: simdgroup GEMM, _multi, single-tokenBorrowed in spirit, own kernels
FLASH_ATTN_EXT variant selection (DK/DV, f16 KV)gqa_attn_flash (nt==1) and attn_flash_prefill (nt>1) with hd ∈ {64,128} guardsBorrowed in spirit, own port
Fusion rules in the Metal backendminfer fuses at the graph level: FusionPass SwiGLU + build-time FusedQKV/FusedFFN/FusedQkvNormDiverged (graph-level fusion)
Unified-memory weight buffers (mmap or copies)newBufferWithBytesNoCopy over mmap'd GGUF parts, per-weight (buffer, offset)Borrowed
KV cache element-type policyModelDef::kv_format auto f16 for the 7B class, per engineBorrowed
MPS MPSGraph / higher-level MPS APIs—Not used (all kernels are hand-written MSL)
Multi-GPU / device selection, command-buffer concurrency tuning—Out of scope

4. Design

4.1 MetalBackend lifecycle and state

new() returns None when MPS is unavailable or MINFER_DISABLE_MPS is set (both handled inside MpsState::try_new); otherwise it stores the 'static singleton reference and starts with an empty pool and a null command-buffer pointer. MpsState::init() runs once at model load.

Pool rules:

  • Exact byte-length reuse only — alloc_buffer matches size * 4 against the free list, else allocates a StorageModeShared f32 MTLBuffer.
  • free_buffer never releases — the id goes back to the free list so persistent KV regions survive rebuilds; the MTLBuffer stays alive for the process.
  • alloc_fresh always allocates — split-boundary staging must not be recycled out from under in-flight node buffers.
  • No generation counter. Unlike CUDA there is no captured-exec cache, so no pointer invalidation machinery is needed.

Capture staging (staging / free_staging) exists only for trace/viz: staging_alloc reuses an exact-length staging buffer or allocates one, capture_split(src_ids) encodes blits at the end of the split's command buffer (after all kernels, so the staging holds this step's output) and returns ids valid only after the next submit, read_staging(id) reads them back, and release_staging_all() returns them to the free list after that split's readback.

In-place aliasing is the allocator's decision (sole consumer + same backend). The executor calls copy_in(dst, src) only when in_bufs[0] != out_buf — the non-aliased case — so an in-place kernel runs directly on out_buf, which for an aliased node is the producer's buffer.

MINFER_OP_PROFILE=1 accumulates host encode time per op label and per-submit GPU wait; the first submit prints a top-20 table, later submits print one line each; zero overhead when unset. Drop submits any pending command buffer (never leaving an unterminated encoder) and prints the profile.

4.2 Backend trait mapping

Trait methodMetal implementation
name()"metal"
supports_op(op, dtype)§4.3
supports_fused(fused)matches!(fused, FusedOp::SwiGLU) — the only variant in the enum
alloc_buffer / free_buffer / alloc_fresh§4.1
execute_node(node, in_bufs, out_buf, kv_pair)opens the split's command buffer once, then dispatches; guards return Err
read_host(id) / write_host(id, data)direct &[f32] views over shared memory (no staging, no copy)
synchronize()submit_pending()
graph_replay(..)not overridden (the trait default returns false) — Metal has no capture/replay path

Alongside the trait, MetalBackend exposes the capture helpers (capture_split, read_staging, release_staging_all) that the scheduler's trace/viz path calls, and the module exposes metal_available().

4.3 Eligibility

supports_op is:

OpMetal
Inputyes (any dtype)
Add, Mul, Silu, RmsNorm, QkNorm, SwiGLUF32
MatMulF32 activation (the weight type rides in MatMulMeta.weight_ttype)
GetRows, RoPE, Attn, KvcacheStore, KvcacheLoadF32
FusedQKV, FusedQkvNorm, FusedFFNF32
View, Reshape, Permuteyes (identity copy)
Scale, Softmax, BatchMatMulno (vocabulary only)
QkvBiasRopeStoreno — the mixed-quant decode epilogue is CUDA-only; on macOS the builder never emits it (the #52 recorded decision, with the dispatch arithmetic, is in docs/SUPPORT-MATRIX.md)

Two layers of checks are deliberately elsewhere:

  1. Weight-type eligibility is a model-level all-or-nothing gate (weights_on_gpu, §4.6): either every graph-referenced weight is registered on the GPU, or the model runs entirely on CPU.
  2. Shape invariants are enforced in execute_node and return Err: attention requires nkt == n_head_kv * hd (the kernel strides KV by nk*hd) and hd == hd_kv (it uses the query head dim); the KV pair must exist; a missing weight is an error; the fast attention paths are limited to hd ∈ {64, 128} and fall back to the classic kernel otherwise.

One Metal/CUDA divergence worth stating: Op::RoPE carries the style into the kernel (rope_style 0 = non-interleaved/Qwen2, 1 = interleaved/LLaMA), so Metal supports both styles; CUDA's supports_op gates RoPE to NonInterleaved only. All supported models are non-interleaved.

4.4 Execution dispatch

execute_node opens the split's command buffer once, then dispatches:

OpMetal path
Inputno-op (host-filled)
Silucopy_in if not aliased, then silu_f32 in place
Add / Muladd_f32 / mul_f32
RmsNormweight from NormMeta; rms_norm_256 when rms_norm_256_enabled(), else rms_norm. A missing gain (no NormMeta, no weight_name, or a name the device never registered) is a loud Err — never the weightless kernel (#40)
QkNormsame kernels with d = hd, n = len/hd over the flat [nt*nh, hd] rows
MatMulquant_matmul_f32_on_gpu_buf (tier below) + optional add_bias_f32
GetRows + Embed metaembed_tokens_gpu (per-weight-type row gather + dequant)
GetRows + no metaget_rows_f32 (the G3 tail-row selection)
RoPEcopy_in if not aliased, then rope_f32 with the node's rope_style
SwiGLUswiglu_f32
KvcacheStorekv_pair required; two store_kv calls (K then V); f32 or f16 by the engine's per-instance kv_format
KvcacheLoadno-op — the output buffer is the persistent K region
Attn§4.4.1
View / Reshape / Permutecopy_in when the output differs, else no-op
FusedQKVconcat matmul (blk.{i}.attn_qkv) + attn_bias_rope_store; refuses nt != 1 with Err
FusedFFNconcat matmul (blk.{i}.ffn_gu, od = 2*nf) + in-place swiglu_f32_off; refuses nt != 1 with Err
FusedQkvNormconcat matmul + two in-place per-head rms_norm[_256] (q at offset 0, k at byte offset nqt*4) + attn_rope_store; refuses nt != 1 with Err
Scale / Softmax / BatchMatMulErr("op ... unsupported on Metal (Phase 3)")
QkvBiasRopeStoreErr("op ... unsupported on Metal (CUDA-only)") — reaching it is a scheduling invariant violation

MatMul tiers (quant_matmul_f32_on_gpu_buf): a simdgroup GEMM kernel (64×32 tile, 128 threads, 8 KiB threadgroup scratch) when nt >= 2 && (od >= 2048 || nt >= 9) && gemm_enabled(); otherwise the _multi kernel for nt > 1 (one threadgroup per two output rows); otherwise the single-token kernel. Every supported quant type has the tiers that matter. Guards: K-quant id % 256 != 0 aborts via gpu_abort; the GEMM checks the threadgroup-memory request against the device limit queried at init. MINFER_GEMM=0 disables the GEMM tier for A/B. An unregistered weight dtype is refused with Err (#329) — the catch-all _ arm is a guard, not a fallback (see the f32 paragraph below).

f16 weights run on the device (#164, landed on a Mac 2026-10-06). The tiers above exist for the quantized types; f16 and f32 have their own single-token arms. The loader registers TensorType::F16 raw (2 B/element — no registration-time f32 copy) via the Metal branch's matches!(ttype, F32 | F16 | BF16), and quant_matmul_f32_on_gpu_buf's TensorType::F16 arm dispatches kernel_f16_f32_matmul (src/metal/kernels/f16.metal), a f32-activation matmul that promotes each half weight in-register over NR0*NSG = 8 output rows per 64-thread threadgroup (grid (ceil(od/8), 1), the token loop inside so a prefill re-streams a weight row once per threadgroup). The embedding gather is the sibling kernel_get_rows_f16, selected by embed_tokens_gpu's F16 arm (one element per thread, nb = ne). Both are listed in build.rs's SHADER_SOURCES and their pipelines (pl_f16_f32, pl_get_rows_f16) are built in try_new, so Qwen2Graph::weights_on_gpu passes and an f16 GGUF is a Metal model. Like CUDA, an f16 prefill runs this f32-activation kernel, not a simdgroup GEMM; 1-D norms/biases stay f32 (the file contract), so an f16 norm can never reach a kernel. Before #164 the type was refused here — registering a weight type a kernel cannot consume would make the device claim true while the op silently ran the wrong (or no) kernel, exactly what the registration gate exists to prevent. A second 2 B/element dtype (bf16, #208) is the subject of the paragraph after the f32 one. Per docs/SUPPORT-MATRIX.md, f16 is on both device columns.

f32 weights run on the device too (#317, landed on a Mac 2026-10-06). The loader's matches!(ttype, F32 | F16 | BF16) branch registers an f32 2-D weight raw (4 B/element), and quant_matmul_f32_on_gpu_buf's TensorType::F32 arm dispatches kernel_f32_f32_matmul (src/metal/kernels/f32.metal, the pl_f32_f32 pipeline) — the f32 twin of the f16 kernel and the peer of CUDA's launch_f32_f32_matmul. Before #317 an f32 weight had no arm and hit the catch-all _ arm (the Q4_0 kernel): kernel_q4_0_f32_matmul reads the f32 bytes as Q4_0 blocks (the first two bytes of 1.0f32 are 0x0000, an f16 scale of 0) and writes zeros. The gap was invisible because the reporting test, graph::op_matrix::matrix_cases_match_their_reference, only ran its Metal column when some earlier test in the process had already initialized MpsState; its Metal arm now calls MpsState::init() explicitly, exactly as its CUDA arm calls CudaState::init(), so the column no longer depends on test order. That silent-wrong-kernel fallback is the registration-gate failure the f16 paragraph above describes, and docs/SUPPORT-MATRIX.md's footnote 2 is updated to match. Two scope notes: no model in the gate set carries a 2-D f32 weight, so this path is covered by the synthetic metal_matmul_f32_matches_cpu gate and the op-matrix case only; and the catch-all _ arm that silently ran the Q4_0 kernel is gone — since #329 it returns Err naming the node, the observed dtype and the kernel that would have run (pl_q4_0_f32 / _multi), so the next unregistered dtype aborts instead of repeating #317's silent zero. That refusal is driven through the production dispatch by graph::metal_backend::tests::metal_matmul_refuses_an_unkerneled_weight_dtype, whose control arm builds the same graph with an F32 weight and asserts it still computes, so the dtype is the only difference (rule 2).

bf16 weights run on the device too (#208, landed on a Mac 2026-10-06). The second 2 B/element dtype — the Metal half of the ticket whose CUDA half is PR #321. kernel_bf16_f32_matmul + kernel_get_rows_bf16 (src/metal/kernels/bf16.metal, the pl_bf16_f32 / pl_get_rows_bf16 pipelines built in try_new and listed in build.rs's SHADER_SOURCES) are dispatched by the TensorType::BF16 arms of quant_matmul_f32_on_gpu_buf / embed_tokens_gpu — the f16 pair's geometry with an in-register as_type<float>(bits << 16) promotion (the device twin of crate::block::bf16_to_f32, exact for every value including NaNs). Its own kernel, not a dtype flag on the f16 one — the same decision #208's CUDA half made: bf16 and f16 are different 2 B/element layouts, so a shared kernel would branch per element in the hottest device kernel. Both loaders' Metal arm (matches!(ttype, F32 | F16 | BF16)) registers it raw (2 B/element, no f32 copy), so weights_on_gpu passes and both architectures are Metal models; 1-D norms/biases stay f32. The kernel-exactness gates bf16_matmul_matches_the_exact_shift_reference / bf16_embed_gather_matches_the_reference assert bitwise against crate::block::bf16_to_f32 (a wrong kernel — the f16 or f32 one — is red), and the ignored real-model gate f208_bf16_weights_run_on_the_metal_device measures 169 bf16 matmul + 1 embed nodes all on Backend::METAL, 942.4 MiB of device weights, max |Δlogit| 1.889e-3 absolute / 1.025e-4 relative (bar 0.05 / 5e-3) with an identical greedy continuation. Per docs/SUPPORT-MATRIX.md, bf16 is now on both device columns.

Aliasing. Only Silu, RoPE and the view ops call copy_in(dst, src), and only when the allocator did not alias them; an aliased node runs its in-place kernel directly on out_buf.

4.4.1 Attention dispatch

Pre-dispatch guards return Err: nkt == n_head_kv * hd (the classic kernel strides KV by nk*hd), hd == hd_kv (it uses the query head dim), and the layer's KV pair must exist.

  • Decode (nt == 1): flash_attn_enabled(hd) → gqa_attn_flash (chunked, partials merged by the shared combine kernel); else hd ∈ {64,128} and MINFER_NO_SPLIT_ATTN != "1" → gqa_attn_split_f32 (two-pass KV-parallel); else the classic gqa_attn_f32.
  • Prefill (nt > 1) with hd ∈ {64,128}: prefill_flash_enabled(hd) → attn_flash_prefill (the llama flash_attn_ext_blk port, with the tail-pad kernel for a partial last KV block); else matmul_attn_enabled() → attn_parallel_prefill (3-pass scores → masked softmax → output); else the classic gqa_attn_f32.
  • Other head dims always take the classic kernel.

Window modes (E1, issue #44 part (a), landed on a Mac 2026-10-06). The dispatch above is the causal path: every kernel derives token t's window from positions (nkv = positions[t] + 1, or a host max_pos + 1 for prefill). When the node is Op::Attn { explicit_span: true } — more than one sequence in a batch, or a run that does not start at cell 0 — the arm selects the mode from the size of the window input (topology, fixed at build time), mirroring CUDA's arm:

  • in_bufs[3].len == 2 * nt → the attn_span layout (one [lo, hi) pair per query). A prefill (nt > 1) at hd ∈ {64,128} with MINFER_NO_WINDOW_FLASH unset takes the fast windowed family kernel_flash_attn_window_blk_{f32,f16} / _hd128_{f32,f16} (src/metal/kernels/fa_window.metal, issue #359): a copy of the causal kernel_flash_attn_blk_* tile structure (Q=8 × C=64 simdgroup GEMM, inline online softmax, the kernel_kv_tail_pad tail) whose mask reads each query's explicit [window[t], window[nt+t]) instead of the causal [0, positions[t] + 1). The threadgroup processes the launch's global [lo_min, hi_max) union; the host advances K/V by lo_min for the tail pad and passes lo_min for the mask, so a windowed prefill does the causal tile work with extra blocks masked out. Every other shape — nt == 1 decode, any hd outside {64,128}, the opt-out — keeps the correctness kernel gqa_attn_window_f32 / _f16 (src/metal/kernels/attn_window.metal), which reads K/V at the run's cells [lo, hi) with the classic kernel's Bc = 32 tiling (the CPU gather cpu_gqa_attn_runs is the structural reference). f32/f16 is selected from the engine's KV format for both families; the fast family is the one measured below, the correctness family stays its reference.
  • in_bufs[3].len == nt * KV_MAP_MAX_SPANS * 2 → the kv_map layout (a list of (cell, len) runs per query, C8b S2/S4). A prefill (nt > 1) at hd ∈ {64,128} with MINFER_NO_WINDOW_FLASH unset takes the fast map family kernel_flash_attn_window_map_{f32,f16} / _hd128_{f32,f16} (also in src/metal/kernels/fa_window.metal, issue #369): the same Q=8 × C=64 tile over the launch's global [lo_min, hi_max) union, with the mask a run-membership walk (fwin_map_has, ≤ KV_MAP_MAX_SPANS runs) instead of [lo, hi) — the arithmetic CUDA does in attn_map_nkv / kv_cell. Every other shape — nt == 1, any hd outside {64,128}, the opt-out — keeps the #362 correctness kernels gqa_attn_map_f32 / _f16 (also in src/metal/kernels/attn_window.metal), which resolve each flat window row to a cell by walking the ≤ 4 runs instead of the window's lo + ki. The correctness one-range and map kernels are a separate family on purpose: their instruction streams are the #315 measured contract, so their code path is byte-untouched. A sharing sequence's window (a shared prefix plus a private run) is exactly this shape; the input's size selects the layout.
  • anything else → a loud Err naming the accepted sizes.

supports_attn_span() is now true (SUPPORTS_ATTN_SPAN) and Device::gathers_attn_map is now true for Metal, so both explicit layouts are read on the device. The causal paths (flash / split / parallel-prefill / classic) are byte-untouched: the windowed families are used only for an explicit window, so a single-sequence causal forward keeps its previous numbers.

Packed q8_0 KV is enabled (#310). Metal reads a packed region: READS_PACKED_KV = true (the registry's reads_packed_kv, which GraphAllocator::ensure_kv and KvFormat::Q8_0::supports read), so MINFER_CACHE_TYPE=q8_0 loads and runs on Metal and a C5 session file carries FLAG_PACKED. The store (kernel_store_kv_q8_0, byte-identical to the CPU quantizer) and two read mechanisms cover every attention shape, selected by the pure crate::metal::packed_attn_route:

  • Mechanism A — native packed reads in the decode flash family. kernel_flash_attn_ext_q8_0 and kernel_flash_attn_ext_hd128_q8_0 (fa_decode.metal) are the f16 kernels' twins — the same threadgroup layout, chunk loop, barriers, MINF_MAXHALF masking and partial buffer, so the shared kernel_gqa_attn_combine_f32 merges them unchanged. Only the per-lane K/V read changes: one 34-byte block's four dequantized elements via dequant_q8_0_kv4 (dequantize.h, four scalar byte loads — a block begins at row*row_bytes + 34*b, and 34*b is 2 mod 4 for odd b, so no vector load is alignment-safe). row_bytes is the one added argument (buffer 11), so buffers 0..10 keep their indices. nt == 1, hd ∈ {64,128}, causal.
  • Mechanism B — an f32 staging window for every other packed path. kernel_dequant_kv_q8_0_to_f32 (kv.metal) dequantizes the needed window of packed cells into a transient, arena-addressed f32 scratch buffer (row i = cell i), then the unchanged f32 prefill (kernel_flash_attn_blk_*) or windowed-flash (kernel_flash_attn_window_{blk,map}_*) family runs against it with f16 = false. Because the stage is arena-addressed, the windowed kernels keep their absolute lo_min addressing and the tail-pad kernel its relative offset — no kernel in those families is edited. nt > 1 causal prefill and every nt > 1 explicit window.
  • Anything neither covers — a small/odd hd, an nt == 1 explicit window, any MINFER_NO_* opt-out — keeps the classic packed kernels kernel_gqa_attn_q8_0 / kernel_gqa_attn_window_q8_0 / kernel_gqa_attn_map_q8_0, still reachable and still the correctness fallback.

MINFER_PACKED_KV_STAGE=1 routes a shape mechanism A would take through mechanism B, so the decode A/B is measurable; the differential gate (metal_packed_decode_stage_matches_the_native_read) drives that switch through a #[cfg(test)] thread-local rather than the environment.

Why f32 staging, not f16. Staging to f16 would round every already-Q8_0-dequantized cell element a second time. The extra rounding is small per element, but on Qwen3-0.6B it is amplified through 8 decode steps: the interleaved f16-stage run measured max |Δlogit| 16.9 (at the argmax 9.46) against the f32 engine, over the same prompt where the classic f32-dequant packed path measures 1.29 / 0.43 and the f16 cache measures 0.011. The prefill-step delta is tiny (0.19) but compounds; the mechanism is the decode trajectory's sensitivity to the specific error realization, not a kernel fault. f32 staging removes the second rounding and restores the classic / CUDA class (docs/SUPPORT-MATRIX.md): the ignored real-model gate measures 1.28 / 0.36 on Qwen3-0.6B and 2.37 / 0.46 on Qwen2.5-0.5B, both inside the inherited ≤ 4.0 tail / ≤ 1.0 argmax bars — at parity speed (below).

Measured (issue #310, macbook (macOS 27.0.1, Apple M4 Pro), hostname macbookpro-ysw, 2026-10-08). Size and speed with the capability enabled and the two mechanisms in place, MINFER_CACHE_TYPE=f16 vs q8_0, 3 interleaved runs (minfer bench -p 1024 -n 64 -r 3 --n-ctx 2048, medians):

modelKV regions f32 vs q8_0pp1024 f16 → q8_0tg64 f16 → q8_0
Qwen3-0.6B Q8_0 (hd 128)58 720 256 B → 15 597 568 B (3.76×)4 844 → 4 683 tok/s (0.967×)193.5 → 176.6 tok/s (0.913×)
Qwen2.5-0.5B Q4_0 (hd 64)6 291 456 B → 1 671 168 B (3.76×)6 175 → 6 154 tok/s (0.997×)295.7 → 245.5 tok/s (0.830×)

The memory win is 3.76×. Prefill (mechanism B) lands within noise of f16 (0.997× / 0.967×); decode (mechanism A) is 1.20× / 1.10× slower than f16 — the native packed kernel issues five 2-byte loads per lane against the f16 kernel's one vectorized half4, so at these short, dispatch-bound contexts it is LSU-bound rather than bandwidth-bound and the 34 B/32-element read does not pay off. That is inside the CUDA precedent after #144 (1.01–1.24×). MINFER_PACKED_KV_STAGE=1 on the 0.5B tg64 shape (mechanism B decode: dequantize, then the f32 flash) measures 238.8 tok/s, ~2.7% below mechanism A's 245.5 — the stage cost is small and the mechanism-A decode is the slower one. The gate metal_packed_decode_stage_matches_the_native_read pins A vs B at max |Δ| 1.7e-8 (f32 reduction order). This is the enabled behaviour, not a refusal.

Measured (issue [#315], macbook (macOS 27.0.1, Apple M4 Pro), hostname macbookpro-ysw, 2026-10-07). Bar named before the run (docs/GATE-CONTRACT.md rules 3 and 5): the windowed arm's tokens/s must reach >= 0.8x the causal prefill's at the same total token count, on the median of interleaved matched rounds (rule 4). It was not met on either model. Three runs each of

cargo test --release --bin minfer -- --ignored --nocapture a_windowed_prefill_is_not_materially_slower
MINFER_315_MODEL=~/.cache/minfer/models/hf/Qwen/Qwen3-0.6B-GGUF/Qwen3-0.6B-Q8_0.gguf \
  cargo test --release --bin minfer -- --ignored --nocapture a_windowed_prefill_is_not_materially_slower

(a 5-round interleaved harness, #[cfg(target_os = "macos")] #[ignore]d, in src/models/qwen2/graph/tests/batching.rs; it prints the model, n_ctx, the node/attention-node counts and the kernel each arm took). Both models use n_ctx = 1024, n_total = 512; the second exercises the f16 windowed kernel and hd = 128, the first the f32 kernel at hd = 64.

Model (KV, hd)causal single seq (attn_flash_prefill)windowed two-seq batchwindowed one-seq @ cell 128
Qwen2.5-0.5B Q4_0 (f32, 64)6387 ± 10 tok/s3439 ± 52410 ± 3
Qwen3-0.6B Q8_0 (f16, 128)5258 ± 17 tok/s838 ± 1505 ± 1
Pair (median of the 3 runs)0.5BQwen3-0.6B
primary — one 512-token causal sequence vs a two-sequence batch of the same 512 tokens0.539x (0.537 / 0.541 / 0.539)0.159x (0.159 / 0.160 / 0.159)
shape-matched — one causal sequence vs the same sequence behind a 128-cell holder, same nt/n_out/tokens, only the kernel differs0.376x (0.375 / 0.377 / 0.376)0.096x (0.096 / 0.097 / 0.095)

So the windowed kernel runs at ~0.54x / ~0.16x the causal prefill's tokens/s for the two-sequence batch (~1.9x / ~6.3x slower) and ~0.38x / ~0.10x at the same shape (~2.7x / ~10.4x slower); the f16/hd = 128 instantiation is the worse of the two.

The causal kernel and code path are untouched (the round adds only the harness); the gap is the simple windowed kernel's serial per-(query, KV-head) walk over the run, which the tuned attn_flash_prefill tiles. A windowed fast path — tiling the run list the way fa_prefill.metal tiles a contiguous window, without disturbing the causal kernels (the #137 lesson) — was therefore warranted and is filed as #359; the correctness kernel stays the reference.

Measured (issue [#359], macbook (macOS 27.0.1, Apple M4 Pro), hostname macbookpro-ysw, 2026-10-07). Same harness, bar and n_ctx/n_total as the #315 run above (median of 5 interleaved rounds, three runs each). The fast family is now selected for every windowed prefill arm; the causal arm and its kernel are unchanged, and both windowed arms print kernel_gqa_attn_window_* because the harness's kernel label is hard-coded to the correctness family (the harness is the yardstick and was re-run unchanged).

Pair (median of the 3 runs)0.5B (before → after)Qwen3-0.6B (before → after)
primary — one 512-token causal sequence vs a two-sequence batch of the same 512 tokens0.540x → 0.968x0.158x → 0.918x
shape-matched — one causal sequence vs the same sequence behind a 128-cell holder, same nt/n_out/tokens, only the kernel differs0.379x → 0.996x0.093x → 0.995x

Both arms clear the >= 0.8x bar on both models (0.5B primary 0.968 / 0.968 / 0.971, shape 0.995 / 0.996 / 0.997; Qwen3-0.6B primary 0.918 / 0.917 / 0.920, shape 0.995 / 0.997 / 0.995). The residual gap in each primary pair is the two-sequence batch's own per-forward overhead, not the attention kernel: the windowed-global launch runs the same tile count the causal prefill does. With MINFER_NO_WINDOW_FLASH=1 both pairs fall back to the #315 numbers above, which is the A/B control for the fast family.

Measured (issue [#369], macbook (macOS 27.0.1, Apple M4 Pro), hostname macbookpro-ysw, 2026-10-07). The same harness gained a fourth arm for the set-valued kv_map layout: a donor owns the first N_TOTAL/2 positions and a subject reads them in place while prefilling its own N_TOTAL tokens after them, so its window is a shared run plus a private run. First-fit puts the subject's run at cell N_TOTAL/2, adjacent to the donor, so the map's global union is one contiguous range and the fast map kernel pays the same tile cost the one-range fast kernel does (the run walk is the only extra work). Bar and protocol are the #315/#359 ones (median of 5 interleaved rounds, three runs each). Before: the #362 correctness map kernel; after: the #369 fast map kernel.

Map arm (median of the 3 runs)0.5B (before → after)Qwen3-0.6B (before → after)
map — one 512-token causal sequence vs one shared-prefix kv_map subject, same N_TOTAL tokens0.229x → 0.943x0.046x → 0.914x

Both reach the >= 0.8x bar (0.5B 0.942 / 0.947 / 0.943; Qwen3-0.6B 0.912 / 0.915 / 0.914). MINFER_NO_WINDOW_FLASH=1 restores the before numbers on both layouts (map 0.229x, one-range primary 0.539x / shape 0.378x), the A/B control the win is attributed through. The map arm's residual gap to the causal prefill is the adjacent shared run (the subject's attention window is N_TOTAL/2 rows wider than a causal query at the same position), not the run walk: the union is contiguous, so fwin_map_has runs over the same tiles the one-range kernel would.

Measured: the map's remaining shape limit. A window whose runs are far apart pays the gap between them — the fast map kernel tiles the global [lo_min, hi_max) union, while the correctness kernel walks just the runs; a map whose shared prefix and private run are non-adjacent would therefore do the union's tile work, not the runs'. The harness's first-fit layout (the production server's) is adjacent and does not pay it; a kv_defrag-spread layout could. This is the recorded tradeoff, not a silent one.

Gate: metal::tests::window_map_flash_matches_the_cpu_reference drives the fast map kernel directly at hd 64 / hd 128 (f32 and f16), with a shared run plus a growing private run, both adjacent (gap = 0) and gapped (gap = 8, so the run walk is load-bearing), nt = 197 and a non-zero lo_min. Bars named before measuring: a one-cell map returns the named cell's V row bitwise (0 on every case), the two-run map matches the CPU reference to <= 0.01 (f32) / <= 0.05 (f16) — measured 1.0e-5 to 2.0e-5 — and a wrong shared base changes the output. The #362 map gates (metal_map_single_cell_matches_the_v_row, metal_map_matches_the_span_and_a_wrong_base_differs) and the #44/#359 gates are unchanged and green.

Gate: metal::tests::window_flash_matches_the_cpu_reference drives the fast kernel directly at hd 64 (f32 and f16) and hd 128 (f32 and f16), with lo_min ∈ {0, 64, 96}, nt = 197 (a non-multiple of 8, so the Q-tile tail padding runs) and nkv = 197 (a partial 64-row tail, so kernel_kv_tail_pad runs). Its bars were named before measuring: a one-cell window returns the named cell's V row bitwise (max|Δ| = 0 on every case), the growing window matches the CPU reference to <= 0.01 (f32) / <= 0.05 (f16) — measured 1.5e-4 / 2.7e-4 — and a window shifted one cell changes the output (the rule-2 control). The #44 correctness gates (metal_attn_span_matches_cpu, metal_attn_span_multi_key_matches_cpu, metal_attn_span_nonzero_start, metal_map_single_cell_matches_the_v_row, metal_map_matches_the_span_and_a_wrong_base_differs) are unchanged and green.

4.4.2 KV write/move side (issue #44 part (b), landed on a Mac 2026-10-06)

  • C3 row move — Backend::copy_cells. CUDA-rejecting on both directions (dst_row above or below src_row, overlapping): the arm opens one MTLBlitCommandEncoder in the current command buffer and copies row by row in the overlap-safe order — ascending when the run slides down, descending when it slides up (dst_row <= src_row), the mirror of CUDA's kv_move_rows. A bulk blit is not an option: Apple documents an overlapping same-buffer copy as undefined, and a separate submission would overlap the producer split (#137). f16 halves elems_per_cell (half[cell * nkt] is nkt / 2 f32 words apart); a packed Q8_0 cell is already whole f32 words including padding and is passed through unchanged (the caller passes region.elems / n_ctx, which ensure_kv sized to KvFormat::Q8_0.row_elems), so the move is a plain whole-word copy that keeps every packed block and its padding verbatim (#310). Gates: metal_copy_cells_moves_overlapping_rows_in_both_directions, metal_f16_kv_cell_move_strides_by_row_bytes.
  • C2 shift / C5 sessions — GraphAllocator::copy_kv_to_cpu. The CPU || CUDA hardcode is gone; the read goes through the registry host_read hook for every backend, so a Metal session shifts (kv_rm/kv_shift) and saves (kv_save*) instead of re-rendering. The read is ordered after the split's submission (copy_kv_to_cpu takes &mut self and calls sync_backend first — MetalBackend::read_host takes &self and does not submit, so a pending buffer would be read stale, the #301 shape). Gate metal_copy_kv_to_cpu_reads_after_the_pending_split leaves a real store dispatch un-submitted and reads through copy_kv_to_cpu: without the flush the read is the region's zeros. The C2 re-rope stays f32-only: an f16 region has no host map (the raw halves would be rotated as f32), so kv_rm — and its start == 0 spelling kv_shift — refuses it loudly, naming #306; CUDA is exposed too and is not fixed here. That refusal is the recorded end state, pinned by graph::alloc::tests::kv_shift::an_f16_region_refuses_the_physical_shift_and_the_other_formats_take_it (a genuine set_kv_format(F16) region, with f32 and packed q8_0 controls that do shift on the same fixture), and the caller re-renders the retained window instead — measured on dgxspark (aarch64, GB10 sm_121), 2026-10-07, 7B Q4_K_M --cnv overflow at --n-ctx 512: 499–505 tokens / 0.30–0.35 s re-prefilled per overflow, against the shift's 22-token / 0.26 s delta.
  • Per-engine kv_format. MetalBackend now carries its own KvFormat, stamped from GraphAllocator::set_kv_format (and read from the allocator stamp on enable_metal), and every attention/store dispatch takes it as an explicit f16 argument instead of reading a process-wide tag. The old metal::KV_F16 OnceLock, kv_cache_is_f16, set_kv_cache_type and the #[cfg(test)] override are deleted: two engines in one process now hold their own layouts and a C5 session header is described under the width its region really uses (the registry kv_format hook reads the allocator's stamp so it answers before the pool is enabled — the CUDA hook's shape). Gate metal_kv_format_is_per_engine.
  • Packed q8_0 is enabled since #310 (READS_PACKED_KV = true); §4.4 records the two mechanisms and the measurements, and docs/SUPPORT-MATRIX.md carries the format row.

The decode chunk count is MINFER_ATTN_CHUNKS or ((max_pos + 1 + 31) / 32).clamp(1, 16) — one chunk per 32 KV rows, capped at 16. nkv for the prefill kernels and the chunk count come from a host read of the positions buffer; that is safe because positions are host-written input data, never GPU-computed. The read is bounded by the node's logical length (positions_max(positions, nt), BufRef::len), not the class-rounded pool buffer: a recycled positions buffer keeps a prior graph's tail, and scanning it derived a stale max_pos that made a causal prefill attend past its own window — the reused_cache_across_prompts_matches_a_fresh_cache failure fixed in #44 part (b). (CUDA instead derives the bound on device so nothing host-side enters a captured graph; Metal has no replay to protect, so the host read is free.)

One measured order-dependence. a_compaction_between_steps_keeps_the_continuation compares two paths (a prefill whose KV is compacted mid-session against one that is not); the two take different kernels, and under the full-suite Metal state the drift is ~0.0058 while it is 0 when the gate runs alone — the same class as matrix_cases_match_their_reference. The gate therefore keeps its behavioural assertion (the greedy token) rather than a tight numeric one, and the number is recorded here so a future full-suite run does not read it as new.

4.4.3 Prefill flash tail-pad overlap (issue #314, landed on a Mac 2026-10-06)

The flash prefill kernel (kernel_flash_attn_blk_f32/_f16, incl. the hd=128 variants) tiles KV into C = 64 blocks. For a partial last block (nkv % C != 0) it read the last C rows (pos0 = nkv - C) from the [2][64][nkt] tail pad. That window overlaps the previous full block whenever nkv % C != 0 && nkv > C (e.g. nkv = 72: block 0 covers rows 0–63, the "partial" block covers 8–71), so the online softmax counted the overlapped rows twice — a real attention error (measured 0.05–0.16 on a synthetic reference), which surfaced as a 1.445 (Qwen2.5-0.5B) / 0.382 (Qwen3-0.6B) chunked-vs-unchunked logit drift because every chunk > 64 tokens hit it, and as a wrong answer for any >64-token prefill. The fix: the partial block reads its own [ic, ic + C) window (pos0 = ic) and kernel_kv_tail_pad pads from ic = (nkv / 64) * 64, so the rows past nkv are zero+masked (kpos0 <= qpos) instead of the leading rows being double-counted. CPU is bitwise across chunk shapes; after the fix Metal's residual drift is the ordinary cross-shape accumulation class (named 0.1, measured 0.0087 on the 0.5B and 0.0078 on Qwen3-0.6B), which the a_chunked_prefill_answers_like_an_unchunked_one gate now carries alongside CUDA. Why no gate caught it for so long: the only cross-shape Metal gate was graph_metal_matches_cpu_logits, whose prompt is 33 tokens — below C = 64, so it never reached a partial block, and the drift was invisible by construction. The isolating experiment named the kernel rather than an accumulation-order effect: MINFER_NO_PREFILL_FLASH=1 (the 3-pass parallel attention) drops the drift from 1.445 to 0.007, and a chunk sweep (1.71 / 1.61 / 1.68 / 1.45 for different chunk counts) is not proportional to the number of chunks, which rules out accumulation order and leaves the flash-prefill kernel as the entry. Gates: flash_prefill_matches_the_cpu_reference_at_every_kv_tail (nkv 72/96/98 red before at 0.164, all ≤ 2e-4 after) and the real-model chunked gate.

4.5 Allocator and scheduler integration

  • Assignment priority is Metal → CUDA → CPU; enable_metal() mirrors enable_cuda().
  • KV regions are created by ensure_kv on the layer's assigned backend, so a Metal-assigned layer's K/V live in the Metal pool and kv_pair resolves to pool ids the executor passes to the kernels.
  • Cross-backend values are staged by the allocator (alloc_fresh on the consumer's backend) and copied through host memory — with shared MTLBuffers that is a plain copy_nonoverlapping.
  • The scheduler calls sync_backend (→ MetalBackend::synchronize) at every backend change and after the last split, which is exactly when the split's single command buffer is submitted.
  • Because Metal supports the whole Qwen2/Qwen3 op set (including the embedding and tail gathers), a normal forward is a single Metal split; CPU splits appear only for an op Metal refuses (Scale/Softmax) or in synthetic/mixed graphs — and the tests deliberately exercise the multi-split alternation.

4.6 Model wiring

metal_on = metal_available() && weights_on_gpu(model)   // #[cfg(target_os = "macos")]
CParams.gpu = metal_on || cuda_on
  • weights_on_gpu is registration-only: it builds the exact list of weight names the graph reads (tok_embd, output_norm, output, output_b, and per layer attn_norm, wq, [bq Qwen2], wk, [bk], wv, [bv], wo, ffn_norm, ffn_gate, ffn_up, ffn_down, plus Qwen3's q_norm/k_norm) and requires mps.has_weight(name) for every one. It is all-or-nothing: one missing name and the model runs entirely on CPU. Unlike CUDA's weights_on_cuda there is no type whitelist here and no diagnostic print — a Metal gate failure is silent.
  • Type support is enforced at loader registration: load_ti registers a tensor only when its type is Q4_0/Q4_1/Q4_K/Q5_0/Q5_1/Q5_K/Q6_K/Q8_0, or F32/F16/BF16 (the last two since #164 and #208). An unsupported type is simply never registered, so the gate fails and the model falls back to CPU. Names are namespaced (mps.register_part per mmap'd part; {ns}{tensor} for the registry) so a second model cannot collide with the first.
  • Concat weights are built once at load with metal::concat_rows and registered: Qwen2 and Qwen3 both register blk.{i}.attn_qkv (all three of wq/wk/wv present, same type/dim, block-aligned) and blk.{i}.ffn_gu, the latter only when nf <= 16384 && MINFER_NO_FUSE_FFN != "1" (the 7B concat would otherwise hold ~2 GiB of weights no node reads).
  • Fused nodes per model: Qwen2 builds FusedQKV for decode QKV; Qwen3 builds FusedQkvNorm (per-head Q/K norm, which the Qwen2 bias+rope+store kernel cannot express); both build FusedFFN when nf <= 16384. The mixed-quant QkvBiasRopeStore class is CUDA-only, so those layers keep the unfused chain on macOS. Gates: CParams.fuse_qkv / fuse_ffn (Qwen3's fuse_qkv is metal_on-only), so no CUDA presence enables a Metal-incompatible fusion.
  • FusionPass gets [cpu, metal, maybe cuda] and a backend_of that maps CPU→0, Metal→1 and CUDA→the index found by name; the SwiGLU rewrite is gated by supports_fused(SwiGLU) and applies to the unfused FFN path (when FusedFFN is built, silu+mul are inside the fused kernel).

4.7 Memory, residency and staging

  • Unified memory changes the copy math. Activations are StorageModeShared MTLBuffers, so host reads/writes are direct views and copy_across is a host round trip; there is no pinned staging ring and no pool_gen.
  • Weights are zero-copy over the mmap'd GGUF parts, with the ~44 ms first-touch page cost paid at load by a dummy warm-up read (METAL_OPTIMIZATIONS.md §0 Done #39). MINFER_WEIGHT_COPY=1 forces per-weight copies.
  • E4 charges the registered weights (#299). The registry (MpsStateInner::weights) stores each weight's own byte extent, and MpsState::weights_bytes sums it (recovering a poisoned lock rather than reporting 0, the CUDA twin); MetalBackend::weights_bytes forwards it, so weights + pooled + request > budget is the one comparison on Metal too. Until #299 the trait-default 0 let the E4 gate ignore the resident weights the E5 auto fit already charged from the GGUF index — it could admit weights against a budget that never counted them. The recorded Mac measurement is in the plan's #299 entry. The gate is graph::alloc::tests::budget::metal_registered_weights_are_charged_in_the_budget_gate; with weights_bytes forced back to 0 its budget refusal does not happen and the gate fails at the unwrap_err. The per-weight extent is tensor.data().len(): 2 B/element for f16/bf16, the block size for a quant — the same bytes CUDA's registry sums.
  • KV element type. The persistent regions stay f32-shaped in the IR, but the engine's per-instance kv_format picks f16 for the 7B class (KV-bandwidth-bound; measured ~−1 ms/token at 2K) and f32 for small models (f16 measured ~3% slower there); MINFER_CACHE_TYPE overrides. The arm is GraphAllocator::set_kv_format (per engine, #44 part (b)).
  • Capture staging exists only while trace/live capture is armed; per-split blits write node outputs into host-readable staging at the end of the command buffer, read back after submit, then released.
  • Device limits are queried once at init (maxThreadgroupMemoryLength) and referenced by dispatch-time guards; the remaining hardcoded numbers are kernel-declared array sizes, documented in GPU_SAFETY.md §3.
  • Profiling. MINFER_OP_PROFILE=1 reports host-encode time per op and per-submit GPU wait, which is how the decode dispatch-cost story in METAL_OPTIMIZATIONS.md was measured.

4.8 Command buffers, submission and trace capture

One MpsCommandBuffer per split, submitted at boundaries, is the whole execution model:

  • Every dispatch helper ends with memoryBarrierWithScope(MTLBarrierScope::Buffers). Metal does not guarantee write visibility between dispatches in one compute encoder; the missing barrier caused intermittent last-2-token corruption on 1.5B/7B prefill before the 2026-08-19 fix.
  • submit() commits with a dispatch-semaphore completion handler (avoiding the ~20 ms scheduler wakeup of waitUntilCompleted), waits with a 10 s bound, and requires MTLCommandBufferStatus::Completed; otherwise it returns Err carrying the recent dispatch trace. The backend expects the submit result; device-configuration problems go through gpu_abort (print the actual limits and exit) rather than degrading silently.
  • Trace/viz: the scheduler pushes each non-KV Metal node's output buffer, and at the split end capture_split encodes the blits; after sync_backend submits, flush_metal_captures reads every staging buffer and records its stats. KV regions are skipped (a full region per layer would dominate the trace).
  • No graph replay. Metal has no CUDA-Graph equivalent here; graph_replay is the trait default. MINFER_METAL_CAPTURE=1 instead starts an MTLCaptureManager GPU trace at init for Xcode, and MINFER_TRACE=1 records a 16-deep per-dispatch label ring used to diagnose a GPU fault (note: MINFER_TRACE also names the graph trace path in the CLI, which is a separate mechanism).

4.9 GPU safety (Metal edition)

docs/GPU_SAFETY.md holds the rules and the incident history; the backend implements them as:

  1. Bounded submit with a status check — never DISPATCH_TIME_FOREVER, never a silent non-Completed status.
  2. No early return past a threadgroup_barrier — the attention kernels were rewritten to run a dummy head instead of returning early, then skip the output write via a valid_head flag.
  3. Device limits queried at runtime — threadgroup memory and thread limits come from the device at init; guards compare against the queried values.
  4. Barriers between dispatches — a buffer written by one dispatch and read by the next needs the explicit scope barrier; a reused threadgroup-memory buffer needs a threadgroup_barrier between the last read and the first write.
  5. Err from execute_node, never a CPU fallback — missing weight, bad shapes, missing KV regions, unsupported op.
    • KV-store bound (#38, gap-table G1). Every Metal arm that writes the persistent K/V region — KvcacheStore (row = cells), FusedQKV and FusedQkvNorm (row = positions) — reads the small StorageModeShared index buffer back and returns Err naming the offending cell and the region's n_ctx before any dispatch (MetalBackend::check_kv_store_rows). A kernel-side range check would be a silent no-write, which this rule forbids. The allocator bounds the same input on the fill_input_i32 path (GraphAllocator::check_positions_bound), so the arm guard closes the fill paths that bypass it. Gate: metal_kvcache_store_refuses_a_cell_past_the_arena. Scope: the check compares against the layer region's n_ctx, not the allocator's arena max when a per-layer n_ctx is smaller; and the FusedQKV/FusedQkvNorm arms share the helper but have no dispatch gate of their own. Record: the plan's #38 entry.
    • Decode-fusion shape guards (#39, gap-table G2). The three decode-only fused arms — FusedFFN, FusedQKV, FusedQkvNorm — refuse a non-decode shape (nt != 1) with an Err naming the node and the observed nt, and the check runs before the weight lookup (shape validation is weight-independent and cheaper, so a bad shape on a weightless node reports the shape, not the missing weight; this matches CUDA's FusedQKV/QkvBiasRopeStore order). Before #39 the arms asserted this with debug_assert!, which a release build compiles out — the asymmetry CUDA never had. Gates: metal_fused_ffn_refuses_nt_other_than_one, metal_fused_qkv_refuses_nt_other_than_one, metal_fused_qkv_norm_refuses_nt_other_than_one (src/graph/metal_backend/tests/fusion_shape.rs). Record: the plan's #39 entry.
    • Norm-weight guards (#40, gap-table G3). Op::RmsNorm / Op::QkNorm no longer fall through to the weightless rms_norm kernel when the gain cannot be resolved. Both None meanings — NormMeta::weight_name absent, or a set name the device never registered — are a missing gain, and the old path produced plausible-looking output from a wrong computation, the failure mode this section forbids. Both arms call MetalBackend::norm_weight, which returns Err naming the node and the missing tensor (the CUDA twin is CudaBackend::norm_weight); no supported producer (models/qwen2, models/qwen3) builds a weightless norm, so there is no legitimate path to preserve. Gates: metal_rms_norm_refuses_a_weight_not_on_gpu, metal_rms_norm_refuses_a_weightless_node, metal_qk_norm_refuses_a_weight_not_on_gpu (src/graph/metal_backend/tests/norm_weight.rs). Scope: the CPU arm keeps its None => rms_norm_f32 fall-through (src/graph/cpu_backend.rs) — untouched, since #40 is a Metal ticket — so CPU remains exposed to the same weightless computation. Record: the plan's #40 entry.
    • KV-store row count (#305). Op::KvcacheStore derives nt from the K input's logical length (BufRef::len), not self.pool[id].length(): the pool allocates at the E4 S2 size class, so the physical length over-counts nt whenever nkt * nt is not itself a class size, and the store would read cells past the filled prefix into the class's uninitialised tail and write those garbage rows into an arena other nodes read. The other self.pool[..].length() uses were audited and are harmless: Silu/Add/Mul/QkNorm over-process their own output padding, and FusedQKV/FusedQkvNorm slice to the rounded length but read only p[0] for the concat. Record: the plan's #305 entry.
  6. gpu_abort for configurations the GPU path cannot run — dimension misalignment, device-limit overruns, kernel-array overflow: print the actual values and exit.
  7. Recurrence playbook — reproduce with one app and a bounded -n; bisect with MINFER_GEMM=0 and MINFER_CACHE_TYPE=f32; on a freeze, spindump over SSH and check the diagnostic reports.

Accepted and documented (not fixed): the audit's L1/L2 landmines, and the fact that a kernel-level fault is not always attributable to a single dispatch (hence the MINFER_TRACE ring).

5. Implementation Phases

5.1 Graph-backend phases

Metal arrived as Phase 3 of the compute-graph rewrite and was then wired to parity with the pre-graph path through the G-series:

PhaseContentStatus
3MetalBackend per-op adapter (metal_backend.rs) + cross-backend scheduling✅
4–6scheduler assign/split/execute, FusionPass, DOT/cache; Qwen2 graph build; imperative forward.rs deleted✅
G1Attn dispatch mirrors the old path (flash / split / parallel / classic) with the same gates✅ 4e105ce
G2RmsNorm selects rms_norm_256 when enabled (~2× faster per dispatch)✅ 4e105ce
G3n_out tail-row GetRows + two allocator liveness fixes✅ d81af71 (docs 8d7cb38)
G4Decode Op::FusedQKV (concat matmul + attn_bias_rope_store)✅ bd28047 (docs 96404fb)
G5Decode Op::FusedFFN (concat gate+up + in-place swiglu), gated nf <= 16384✅ 1dee1b5 (docs ec922f1)
G6Qwen3 Op::QkNorm + Op::FusedQkvNorm (per-head norm + no-bias rope/store)✅ 283c7d6, 94d57ac, d5b8023
objc2 migrationmetal 0.28 / block + vendor/block patch → objc2-metal / block2 / objc2-foundation, Phases 0–6✅ 9e238bd, 6a382a3, be3df55, ee9b65b

The graph-era details live in docs/COMPUTE-GRAPH-DESIGN.md §17 (Phases 1–11, deviations 18–26) and METAL_OPTIMIZATIONS.md §0.1 / §4.3.

5.2 Optimization and cold-start record

Hash caveat. The commit hashes printed in docs/METAL_OPTIMIZATIONS.md predate a repository history rewrite and no longer resolve. The hashes below are subject-matched equivalents from the current history; each was verified with git cat-file -t. Cite these, and the METAL_OPTIMIZATIONS.md section, together.

WorkstreamWhat it changedDoc referenceRepresentative commits
Initial Metal portDevice/queue/encoder skeleton, the full kernel set, RMSNorm/GQA simdgroup parallelism, SwiGLU fusion§3.1, §5.1ac2cb3d, e7df395, 6811da3, 2a76d0c, b0819e7, 2f484f8
Correctness foundationRoPE freq_scale, output_b, softmax max, dynamic hd, Q5_K formula, GQA simd_max partial-tile divergence, first isolation suites§3.1 (#1–#5)df2da9f, c34b7d8, 87fec18
GPU-safety hardeningBounded submit() + status check, dispatch-label trace ring, no early return past a barrier, runtime guards, autorelease retain fix§3.1 (#6–#7), GPU_SAFETY.md §1–§4aef6cde7, 5f5a42d
Decode fusions + split attention (old path)Fused QKV/FFN decode matmuls, 2-pass KV-parallel split attention, float4 acc, adaptive chunks, KV geometric growth§3.3e5db3aa, 39c2ba9, eb0f812, e031282
GEMM (simdgroup) workllama kernel_mul_mm port (64×32 tile, 4 simdgroups), per-quant GEMMs, hot-loop unroll, ik-loop simdgroup_barrier, the partial-tile race + memoryBarrier fix§3.4 (#11/#12/#28/#29/#30)c34b7d8, d83bd25, 5202548, e7be3d3, 69a47d4, f52628c, e997b99
Decode matmul layout portsq6_K stride-2/float4 (72→209 GB/s), q4_K stride-4/sc16 (7B decode ~51→~19.3 ms/token)§3.3 (#27)36c9b01, b59c8c8
Flash-attention portsDecode flash_attn_ext_vec hd=64/hd=128, prefill flash_attn_blk hd=64/hd=128, tail-pad kernel§3.3/§3.4 (#22/#24–#26), §5.5c1177f7, b56b5dd, ac5e1ea, a78620c
Parallel prefill attention + RMSNorm-2563-pass barrier-free prefill attention; 256-thread RMSNorm; chunk cap 16; drop the KV→CPU sync§3.4/§3.5 (#16/#17)35a9659, 89aa1fa
KV f16store_kv_f16 + _f16 attention kernels; auto-select f16 for the 7B class§3.5 (#13/#37)0aad968, ff60ed7
n_out tail-row reductionFinal norm + lm_head on output rows; last-layer FFN + both residuals on the tail rows§3.7 (#32/#34)59a97f9, 660f59e
Cold-start axisEmbedded precompiled metallib, GGUF mmap + zero-copy weights, load-time warm-up read§4.2b7cb5bc, 1032ca9, a5709e7
Embedding coverage + GEMM thresholdGPU get_rows for all 8 quant types; adaptive GEMM dispatch `nt≥2 && (od≥2048nt≥9)`
Prefill-GEMM investigationGrid-shape probe, exact-shape replay, 7B decomposition, structural-equivalence audit, gap acceptance§3.6 (§4.3.1–§4.3.10)6205f7e, 9a201b5, c900127, 89e2468
Compute-graph IR + allocatorPhases 1–6 of the rewrite (IR/builder/allocator, CPU backend, Metal backend, scheduler/fusion/cache, Qwen2 graph build)COMPUTE-GRAPH-DESIGN.md §17.1a163a07, 308cc74, be35b1b, 8091b61, 941d34f, e54070c

Numbers like #27 are local to a METAL_OPTIMIZATIONS.md §0 table — always qualify them, because the Done / To-do / Decided tables reuse the same numbers.


6. Risks and Open Questions

#Risk / questionStatus
1Cross-dispatch write visibility (2026-08-19 incident)Fixed: scope barrier after every dispatch; rule recorded in GPU_SAFETY.md
2Early return past a threadgroup_barrier in attentionFixed: no early returns; invalid heads run the loop and skip the store
3Device-limit guessesFixed: limits queried at init, never hardcoded
4Concurrent MPS access under parallel tests flips fused/unfused greedy tokensMitigated for tests: metal_test_lock() + MpsState::init(); production runs one scheduler thread
5Q8_0 multi-token matmul race (missing trailing threadgroup_barrier)Fixed; pinned by metal_prefill_determinism
6Metal gate failure is silent (no diagnostic like CUDA's CUDA GATE:)Open (diagnostics): a missing/unregistered weight drops the model to CPU with only the init line to explain it
7Whole-layer layer_gpu reference path still referenced in comments/testsOpen (cleanup): the function is gone from src/metal/; the #[allow(dead_code)] block and comments remain
87B prefill ~10% behind the old pathAccepted: GEMM-bound; the GEMM transfers fully, and the prefill-GEMM investigation closed as "decided not to change" (METAL_OPTIMIZATIONS.md §3.6)
9Warm-up / cold-start costsTracked: mmap first-touch solved by the load-time warm-up; remaining cold-start to-dos in METAL_OPTIMIZATIONS.md §4.2

7. Verification

7.1 Device test suite

src/graph/metal_backend.rs carries 14 graph-backend tests, all serialized by metal_test_lock() and skipped with MPS unavailable; skipping when Metal is absent:

  • Elementwise / norm: metal_elementwise_matches_cpu (silu+add, bit-for-bit), metal_rmsnorm_matches_cpu, metal_rmsnorm_real_scale (d=896, nt=8).
  • Cross-backend: metal_cross_backend_copy, metal_cross_backend_copy_large, metal_embed_then_rmsnorm_cross_backend, metal_multi_split_alternation.
  • Matmul: metal_matmul_q8_matches_cpu (real wk shape), metal_matmul_q4_matches_reference (layer-0 wq, [896, 896], vs a manual Q4_0×f32 reference).
  • KV + attention: metal_attn_kv_matches_cpu (bit-exact), metal_attn_kv_real_scale (nh=14, nk=2, hd=64, nkt=128, nt=30), metal_attn_decode_step (nt=1 with 30 stored rows), metal_store_after_gpu_op (KV store whose K comes from a GPU op), metal_store_real_dims (n_ctx=32768).

Model-level Metal tests live with the models and skip on non-macOS: graph_metal_matches_cpu_logits, graph_metal_layer0_isolation, graph_metal_real_wk_matmul (Qwen2), fused_qkv_matches_unfused_decode (Qwen2, with the metal_test_lock rationale recorded), fused_qkv_norm_matches_unfused_decode (Qwen3), graph_metal_matches_llama_reference (Qwen3; pins the first 9 tokens of a 60-token byte-identical llama-Metal run) and metal_prefill_determinism (Qwen3; the Q8_0 multi-token matmul race regression).

7.2 Kernel isolation suites

tests/ carries four macOS-only integration binaries (#![cfg(target_os = "macos")], so they do not exist on Linux). Each builds its own deterministic inputs and embeds a scalar CPU reference — none needs an external dump:

FileTestsCoverage
tests/flash_attn_isolation.rsflash_attn_ext_isolation (hd 64 + hd 128), flash_attn_matches_splitdecode flash vs a scalar online-softmax reference and vs the split path; partial/empty KV chunks, nt 1–2, nkv up to 4097
tests/flash_attn_blk_isolation.rsflash_attn_blk_isolationprefill blk port (hd 64/128, NSG=4), partial last KV block via the tail-pad kernel, nt up to 200, GQA heads, f32 + f16 KV, vs classic
tests/gqa_attn_isolation.rsgqa_attn_isolation, gqa_attn_split_isolation, gqa_attn_split_timingclassic + split attention vs a scalar reference, including the nkv % 32 != 0 divergent-simd_max case
tests/gemm_isolation.rsgemm_isolation, qkv_row_concat_layout, non_q4_0_gemm_isolation, get_rows_q4_k_isolation, get_rows_multi_type_isolationGEMM determinism + correctness vs scalar, concat layout, per-type row-gather

7.3 Acceptance gates

  • Bit-exactness where the math is identical: elementwise and KV/attention round trips are checked bit-for-bit; matmul and norm are checked within float tolerance.
  • CPU-vs-Metal logits: #324 restored the model-level graph_metal_matches_cpu_logits. The pre-#324 body compared ModelDef::forward with ModelDef::forward_graph, but in the graph era both route through Qwen2Graph::forward under the same device decision, so its max |Δ| = 0 compared the Metal graph with itself — a vacuous gate. It now loads the same cached 0.5B Q4_0 twice — a Layers(0) CPU engine and an explicit full-plan Metal engine — drives the same greedy continuation through forward_graph_cached, asserts the two built graphs genuinely differ in backend assignment (metal arm 248 METAL / 0 CPU nodes; cpu arm 0 METAL / 440 CPU nodes), pins the greedy continuation to a literal, and compares the final-step logits. Bar named before measuring: 1.0 absolute (the q8_0 weight-quantisation class of docs/GGUF-TOOLING.md, 3× the observed) and 5e-2 relative. Measured on macbook (macOS 27.0.1, Apple M4 Pro), 2026-10-07, cargo test --release --bin minfer -- --nocapture graph_metal_matches_cpu_logits: max |Δlogit| 0.347 absolute / 1.57e-2 relative against max |logit| 22.04, greedy [12095, 11, 323, 432] identical. The class is quantization, not accumulation order: the CPU quantizes activations to Q8_0 while Metal reads f32 (rule 9), which is why the bar is looser than the f16 #164 gate's 0.05. The row is still backed more widely by the per-op metal_*_matches_cpu gates and the external oracle graph_metal_matches_llama_reference.
  • Fused-vs-unfused decode must be bit-identical, with the unfused side running the FusionPass.
  • llama-reference oracle: Qwen3's first 9 greedy tokens are pinned against llama-Metal.
  • Determinism: metal_prefill_determinism and the isolation suites' repeat-run checks.

Suite baselines: METAL_OPTIMIZATIONS.md / KNOWN-CPU-ISSUES-2026-08-29.md record the single-threaded and parallel runs (e.g. 152 passed / 3 ignored single-threaded main-bin, plus the integration binaries green on device). Run cargo test --release on macOS for the device suites; on Linux the Metal-only tests self-skip and the isolation binaries are empty.

7.4 The macOS suite baseline (#298, recorded by #54 at the round's end)

The current machine-checked baseline is not on this page. It is docs/TEST-BASELINES.md ("macOS unit", box macbook (macOS 27.0.1, Apple M4 Pro), 2026-10-08) backed by the macos-unit row of scripts/test-baselines.toml, together with that page's macOS real-model row. Edit that ledger, never a counter. A later Mac gate run diffs against it: anything new is a regression.

The round-era record is in the plan. The round's final master was 6b95763; the round's start was 97823e4, where the same suite had 21 failures, and the enumeration that followed — the three root causes they decomposed into, the two order-dependent flakes, and the eleven-day build-macos blind spot that hid the macOS test target from cdf41b2 to 4add59f — is the plan's record: G7's row (#54) for the baseline, the #298/#231 suite enumeration for the failures, and the #303 record for the blind spot.

The purpose is unchanged: a Mac gate run compares against a dated baseline, so a green suite cannot hide a new failure. Only the baseline's home moved.


8. Out of Scope / Future

  • Not planned: MPSGraph / higher-level MPS APIs, multi-GPU, f16 activations, training. (f16 weights landed on Metal in #164 and bf16 in #208 — see §4.4 and docs/SUPPORT-MATRIX.md; it is the activation dtype that stays f32.)
  • Closed by measurement (METAL_OPTIMIZATIONS.md §3.6/§4): the prefill GEMM gap (params match llama's; decided not to change), and the flash/split/parallel attention lineup.
  • Remaining research (METAL_OPTIMIZATIONS.md §4.1/§4.2): cold-start items and the residual 7B-prefill gap; see that document's roadmap section rather than duplicating it here.
  • Cleanup candidates: the #[allow(dead_code)] legacy block/comments in src/metal/, and adding a Metal gate diagnostic to match CUDA's CUDA GATE: line (§6 #6/#7).

Decisions governing this document

This page is the current contract; the decisions behind it are frozen in the ADR corpus:

  • ADR-0001 — Inference runs through one declarative compute graph
  • ADR-0002 — Topology is a function of GraphParams alone, so positions cannot be structure
  • ADR-0005 — Metal becomes a first-class backend
  • ADR-0006 — The KV storage format is a per-engine gate, not a process-wide global
  • ADR-0008 — GPU safety: bounded waits, no early return past a barrier, runtime device limits
  • ADR-0009 — A failure is an error, never a silent fallback
  • ADR-0010 — The identity gate: bitwise by default, a named tolerance class otherwise
  • ADR-0013 — The CPU quantizes activations to Q8_0; a device reads f32
  • ADR-0014 — A KV session is a versioned, checksummed file — never a memory dump
  • ADR-0021 — bf16 is a round-to-nearest-even cast, and 1-D tensors stay f32
  • ADR-0012 — Device is the first axis, the layer the second — and no premature common
  • ADR-0015 — The offload auto fit takes a prefix, not a knapsack
  • ADR-0016 — A failed device-memory query is not a zero budget
  • ADR-0032 — The packed-KV staging window is f32, not f16
  • ADR-0036 — bf16 weights get their own device kernels, not a dtype flag on the f16 ones
  • ADR-0037 — A Metal weight dtype a kernel cannot consume is refused, never run as a wrong kernel