From 3989eb9e69152572fd16032f5d19464c20f2b3c1 Mon Sep 17 00:00:00 2001 From: yhirose Date: Tue, 22 Sep 2026 00:48:46 -0400 Subject: [PATCH 1/3] gpu: xent_bwd and adam_step take the tier-0 route where a kernel declines Every other fused op already had a gpu::generic composition to fall back to; these two eager training ops still dropped straight to the host -- adam_step's own comment even said so. Both now compose from the same tier-0 primitives the rest of gpu::generic uses (broadcast subtract and exp plus a one-hot scatter for xent_bwd; a handful of affines, adds and a divide reused through three scratch buffers for adam_step), so a backend with no fused kernel keeps the update on the device instead of paying a round trip. No launch in either composition reads and writes the same buffer, the WebGPU constraint the rest of gpu::generic already works around. adam_step's chain folds an affine into each of its producing mul and sqrt launches (14 launches down to 12); the algebra is exact under the shared op's *scale+offset epilogue, not an approximation. Verified against the fused kernel on the host reference backend and native Metal, under WebGPU via Deno (the one backend that actually lacks both kernels, so this is the composition genuinely running), and by a clean CUDA host-side trace diff (new test launches only, no existing kernel call changed). Mutations of both compositions were each caught by the existing parity test. --- docs/backends.md | 28 ++++++++------ include/array.h | 11 ++++-- include/gpu_ops.h | 93 +++++++++++++++++++++++++++++++++------------ test/test_array.cpp | 34 +++++++++++++++++ 4 files changed, 126 insertions(+), 40 deletions(-) diff --git a/docs/backends.md b/docs/backends.md index 79c8805..f97c2cc 100644 --- a/docs/backends.md +++ b/docs/backends.md @@ -144,23 +144,27 @@ differ, not to avoid reordering a params struct. The ops fall in two tiers. Tier 0 — elementwise, broadcast, reduction, GEMM, copy and index — closes the array surface: a backend with those kernels runs -every graph. Tier 1 is the fused ops of the model path — `rmsnorm`, `swiglu`, +every graph. Tier 1 is the fused ops — the model path's `rmsnorm`, `swiglu`, the decode GEMV, the cache writes, `rope`, the attention, `split_heads` / -`merge_heads`, `argmax` — each a kernel a backend *may* have. Under each of -them `gpu_ops.h` holds one generic composition out of tier 0 (`gpu::generic`), -and the op takes it when the shared launch declines or `own` has no member: -several launches and a scratch buffer or two where the kernel is one launch, -so a backend that cares for the decode loop writes the kernel, and a backend -that has only tier 0 still runs the whole model path. The compositions are -f32: an operand the tier has no reader for — a bf16 cache or weight, int4 -weights — still declines, and a model keeps to the array ops there +`merge_heads`, `argmax`, plus two training ops off that path, `xent_bwd` and +`adam_step` — each a kernel a backend *may* have. Under each of them +`gpu_ops.h` holds one generic composition out of tier 0 (`gpu::generic`), and +the op +takes it when the shared launch declines or `own` has no member: several +launches and a scratch buffer or two where the kernel is one launch, so a +backend that cares writes the kernel, and a backend that has only tier 0 +still runs the whole model path — and, for `adam_step`, still updates its +state on the device rather than a host round trip. The compositions are f32: +an operand the tier has no reader for — a bf16 cache or weight, int4 weights +— still declines, and a model keeps to the array ops there (`caps::row_gemv`, `caps::bf16_gemm`). A composition names no backend and no array: spans, tier-0 ops, and the device core's `alloc` / `release` / `cpu_barrier` / `sync_to_host`. -Two tests hold the two routes together: the model-path test checks the op, -however it ran, against the array oracle, and a second runs each composition -beside the backend's kernel and requires the same numbers. +Two tests hold the two routes together: the model-path test (and, for the two +training ops off that path, their own eager tests) checks the op, however it +ran, against the array oracle, and a second runs each composition beside the +backend's kernel and requires the same numbers. ## Launch policy diff --git a/include/array.h b/include/array.h index b56b13b..06204c9 100644 --- a/include/array.h +++ b/include/array.h @@ -3534,10 +3534,13 @@ struct graph { inv_bc2)) { return true; } - // No kernel for it on this backend: the host loop below. The caller - // updates in place and cannot compose its way out, so this has to be - // total -- and Metal's unified memory makes the round trip a flush and a - // memcpy, the same bargain its other CPU fallbacks make. + // gpu::adam_step already falls to its tier-0 composition when the + // backend has no fused kernel, so this is the device itself declining + // -- a scratch it could not allocate, or a tier-0 launch it refused: + // the host loop below. The caller updates in place and cannot compose + // its way out, so this has to be total -- and Metal's unified memory + // makes the round trip a flush and a memcpy, the same bargain its + // other CPU fallbacks make. } // data() brings any device copy home first, the same as every other host // path. A CUDA build hands every buffer a mirror key, so "has a native diff --git a/include/gpu_ops.h b/include/gpu_ops.h index 98dd0c9..41449e7 100644 --- a/include/gpu_ops.h +++ b/include/gpu_ops.h @@ -202,30 +202,6 @@ inline bool gather_from_axis(span src, span idx, span o, int64_t n, policy::flat(n)); } -// Softmax cross-entropy's pullback from the forward's row logsumexp: -// out[i,j] = g[i] * (exp(x[i,j] - lse[i]) - [j == tgt[i]]). x and out are -// [rows, cols]; lse, tgt and g hold one value a row. -inline bool xent_bwd(span x, span lse, span tgt, span g, span o, int64_t rows, - int64_t cols) { - const int64_t n = rows * cols; - if (n <= 0) return false; - return launch(kop::xent_bwd_, {in(x), in(lse), in(tgt), in(g), out(o)}, - xent_bwd_params{static_cast(cols), - static_cast(n)}, - policy::flat(n)); -} - -// Adam's update in place over n contiguous elements: m and v advance, p moves -// by the bias-corrected ratio the host folded into lr_over_bc1 and inv_bc2. -inline bool adam_step(span p, span m, span v, span g, int64_t n, float beta1, - float beta2, float eps, float lr_over_bc1, float inv_bc2) { - if (n <= 0) return false; - return launch(kop::adam_step_, {inout(p), inout(m), inout(v), in(g)}, - adam_params{beta1, beta2, eps, lr_over_bc1, inv_bc2, - static_cast(n)}, - policy::flat(n)); -} - // tanh / sin / cos: unary's shape under their own vocabulary (gpu_abi.h). inline bool unary_ext(unary_ext_op op, span a, span o, int64_t n, float scale, float offset) { @@ -671,6 +647,45 @@ inline bool argmax(span a, int64_t n, int64_t* out_idx) { return true; } +// Softmax cross-entropy's pullback: exp(x - lse) row-broadcast, a one-hot +// scatter of the targets subtracted, g row-broadcast into the result. +inline bool xent_bwd(span x, span lse, span tgt, span g, span o, int64_t rows, + int64_t cols) { + const int64_t n = rows * cols; + scratch diff(n), p(n), onehot(n), ones(rows, true); + if (!diff || !p || !onehot || !ones) return false; + for (int64_t i = 0; i < rows; i++) ones.contents[i] = 1.0f; + return binary_bcast(kop::bsub, x, cols, 1, lse, 1, 0, diff, rows, cols, 1.0f, + 0.0f) && + unary(kop::exp_, diff, p, n, 1.0f, 0.0f) && + scatter_to_axis(tgt, ones, onehot, rows, cols) && + binary(kop::sub, p, onehot, diff, n, 1.0f, 0.0f) && + binary_bcast(kop::bmul, diff, cols, 1, g, 1, 0, o, rows, cols, 1.0f, + 0.0f); +} + +// Adam's update: m and v each a sum of two affines of their old value and g, +// p moved by their ratio through a scratch of its own (p is read only before +// it is ever the write side, so the update never aliases a launch's input). +inline bool adam_step(span p, span m, span v, span g, int64_t n, float beta1, + float beta2, float eps, float lr_over_bc1, + float inv_bc2) { + scratch a(n), b(n), c(n); + if (!a || !b || !c) return false; + return unary(kop::affine, m, a, n, beta1, 0.0f) && + unary(kop::affine, g, b, n, 1.0f - beta1, 0.0f) && + binary(kop::add, a, b, m, n, 1.0f, 0.0f) && + binary(kop::mul, g, g, c, n, 1.0f - beta2, 0.0f) && + unary(kop::affine, v, b, n, beta2, 0.0f) && + binary(kop::add, b, c, v, n, 1.0f, 0.0f) && + unary(kop::affine, v, a, n, inv_bc2, 0.0f) && + unary(kop::sqrt_, a, c, n, 1.0f, eps) && + unary(kop::affine, m, a, n, lr_over_bc1, 0.0f) && + binary(kop::div, a, c, b, n, 1.0f, 0.0f) && + binary(kop::sub, p, b, a, n, 1.0f, 0.0f) && + unary(kop::affine, a, p, n, 1.0f, 0.0f); +} + } // namespace generic // ---- the model path: what a decoder runs on raw device buffers between its @@ -908,6 +923,36 @@ inline bool attn_prefill_dkv(span q, span K, span V, span dO, span stats, } } +// ---- the training ops: eager, off the decode path, each the backend's fused +// kernel where it has one and the tier-0 composition where it does not. + +// Softmax cross-entropy's pullback from the forward's row logsumexp: +// out[i,j] = g[i] * (exp(x[i,j] - lse[i]) - [j == tgt[i]]). x and out are +// [rows, cols]; lse, tgt and g hold one value a row. +inline bool xent_bwd(span x, span lse, span tgt, span g, span o, int64_t rows, + int64_t cols) { + const int64_t n = rows * cols; + if (n <= 0) return false; + return launch(kop::xent_bwd_, {in(x), in(lse), in(tgt), in(g), out(o)}, + xent_bwd_params{static_cast(cols), + static_cast(n)}, + policy::flat(n)) || + generic::xent_bwd(x, lse, tgt, g, o, rows, cols); +} + +// Adam's update in place over n contiguous elements: m and v advance, p moves +// by the bias-corrected ratio the host folded into lr_over_bc1 and inv_bc2. +inline bool adam_step(span p, span m, span v, span g, int64_t n, float beta1, + float beta2, float eps, float lr_over_bc1, float inv_bc2) { + if (n <= 0) return false; + return launch(kop::adam_step_, {inout(p), inout(m), inout(v), in(g)}, + adam_params{beta1, beta2, eps, lr_over_bc1, inv_bc2, + static_cast(n)}, + policy::flat(n)) || + generic::adam_step(p, m, v, g, n, beta1, beta2, eps, lr_over_bc1, + inv_bc2); +} + // ---- the graph-capture forms (caps::graph_capture): a decode step whose // position is a device scalar, so one captured step replays as the cache // advances. `partials` is attn_dpos_partials_bytes() of scratch. diff --git a/test/test_array.cpp b/test/test_array.cpp index 512acd7..7c5e1ef 100644 --- a/test/test_array.cpp +++ b/test/test_array.cpp @@ -3085,6 +3085,40 @@ TEST_CASE("generic compositions agree with the fused kernels") { CHECK(ib == 4321); } + SUBCASE("xent_bwd and adam_step") { + const int64_t rows = 5, cols = 37; + array x = dev(random_array({rows, cols}, 980)); + array lse = dev(random_array({rows}, 981)); + array tgt = dev(array::from({3, 0, 36, 17, 9}, {rows})); + array g = dev(random_array({rows}, 982)); + array oa = array::empty({rows, cols}), ob = array::empty({rows, cols}); + REQUIRE(gpu::xent_bwd(x.device_span(), lse.device_span(), tgt.device_span(), + g.device_span(), oa.device_span(), rows, cols)); + REQUIRE(gen::xent_bwd(x.device_span(), lse.device_span(), tgt.device_span(), + g.device_span(), ob.device_span(), rows, cols)); + tl::gpu::flush(); + CHECK(same(oa, ob, 1e-5f)); + + const int64_t n = 200; + auto p0 = random_array({n}, 983), m0 = random_array({n}, 984); + auto v0 = tl::pow(random_array({n}, 985), 2.0f).eval(); // v is a square + array grad = dev(random_array({n}, 986)); + array pa = dev(p0), ma = dev(m0), va = dev(v0); + array pb = dev(p0), mb = dev(m0), vb = dev(v0); + const float beta1 = 0.9f, beta2 = 0.999f, eps = 1e-8f, lr_over_bc1 = 0.01f, + inv_bc2 = 1.05f; + REQUIRE(gpu::adam_step(pa.device_span(), ma.device_span(), va.device_span(), + grad.device_span(), n, beta1, beta2, eps, + lr_over_bc1, inv_bc2)); + REQUIRE(gen::adam_step(pb.device_span(), mb.device_span(), vb.device_span(), + grad.device_span(), n, beta1, beta2, eps, + lr_over_bc1, inv_bc2)); + tl::gpu::flush(); + CHECK(same(pa, pb, 1e-4f)); + CHECK(same(ma, mb, 1e-5f)); + CHECK(same(va, vb, 1e-4f)); + } + tl::device_ = prev; } From c38d0a965efd9b014e37c8345f8c18e4f9dcb3b5 Mon Sep 17 00:00:00 2001 From: yhirose Date: Tue, 22 Sep 2026 09:25:09 -0400 Subject: [PATCH 2/3] gpu: Metal's GEMM tiles and attention variants key their own pipelines kop carried two families no other backend implemented -- Metal's ten GEMM tile ids and its eighteen attention-kernel variants (by D and whether the KV cache is bf16) -- kept there only because bind_/pso_ took a kop directly. Metal now caches any (id, name) pair through one pso_id_, so its own gtile and attn_op enums use it the same way kop does, offset into disjoint ranges (gtile negative, attn_op further negative) so the shared cache never confuses one family's pipeline for another's -- a static_assert now checks the gap directly rather than leaving it to a comment. CUDA's one f32 GEMM fallback used to borrow kop::sgemm32 as an arbitrary cache key -- its cached_(slot, name) helper already didn't need a kop at all, so it gets its own named slot instead. docs/backends.md folds the CUDA wave plan's design reasoning into the existing paragraph that already says a backend's own kernel choice stays with the backend, rather than repeating it as an open gap. Verified on native Metal (117 cases, 11640 assertions with real device access, plus check_qwen's greedy tokens against the numpy oracle) and gpu_host's no-GPU reference build, under WebGPU via Deno, and by a clean CUDA host-side trace diff (no existing kernel call changed). Two collision mutations (gtile's key losing its offset, then attn_op's) each hung the real device outright rather than silently misbehaving -- about as decisive a catch as a test gets. --- docs/backends.md | 12 ++-- include/cuda.h | 15 ++--- include/gpu_abi.h | 32 +++------- include/metal.h | 157 +++++++++++++++++++++++++++++----------------- 4 files changed, 120 insertions(+), 96 deletions(-) diff --git a/docs/backends.md b/docs/backends.md index f97c2cc..c0a72e1 100644 --- a/docs/backends.md +++ b/docs/backends.md @@ -196,6 +196,11 @@ why CUDA's fill is a constant rather than a device query. Which of its kernels a backend launches — CUDA's f32 tile and its wave plan, Metal's STEEL bands — stays with the backend: a tile is that kernel family's ABI. What such a choice measures its grid against is the same `fill_groups`. +The wave plan itself (`sgemm_wave_chunk_`) has no Metal counterpart to share +it with (Metal's GEMM dispatch only ever picks a tile, never splits K), and +its own two-stage layering — a floor'd fill share, then spare slots folded +into a shorter tail — doesn't fit `gpu::policy::split_rule`'s single ceil +formula either, so it stays a CUDA-only function. ## Profiling @@ -257,11 +262,8 @@ asks; a new backend edits no test. - The mirror table (handle to host copy, device copy, size, `residency`) and the buffer pool exist twice, in `cuda.h` and `webgpu.h`. The state machine itself is shared. -- `kop` still lists kernel ids only one backend has (Metal's GEMM tiles and - attention variants). -- CUDA's f32 GEMM wave plan (`sgemm_wave_chunk_`: layers, spare slots and - rounds over a two-blocks-per-SM wave) is a split policy of its own, written - against `traits::fill_groups` but not yet a shared function. +- WebGPU's kernels still predate the canonical ABI and go through `marshal_` + (see above). ## Verifying a change diff --git a/include/cuda.h b/include/cuda.h index 59123a7..f3e4bf4 100644 --- a/include/cuda.h +++ b/include/cuda.h @@ -226,13 +226,6 @@ inline const char* kernel_name_(kop op) { case kop::merge_heads_: return "tl_merge_heads"; case kop::split_heads_: return "tl_split_heads"; case kop::argmax_: return "tl_argmax"; - // Every f32 GEMM id is the one general kernel here (the tiled fast path - // has its own names: sgemm_tiles below). - case kop::sgemm32: case kop::sgemm32x64: case kop::sgemm64x32: - case kop::sgemm64: case kop::steel: case kop::steel32x64: - case kop::steel_ta: case kop::steel_tb: case kop::steel32x64_ta: - case kop::steel32x64_tb: - return "tl_sgemm"; default: return nullptr; // no kernel here: dispatch declines } } @@ -517,6 +510,11 @@ struct context { CUfunction pad_() { return cached_(pad_fn, "tl_pad"); } CUfunction fold_() { return cached_(fold_fn, "tl_fold"); } + // The f32 gemm's general fallback (the tiled fast path above, sgemm_(), has + // its own names). + CUfunction sgemm_fallback_fn = nullptr; + CUfunction sgemm_fallback_() { return cached_(sgemm_fallback_fn, "tl_sgemm"); } + // Embedding-table lookup (index_select/index_add) and pooling-style // one-hot scatter (scatter_to_axis), cached the same way. CUfunction index_add_fn = nullptr, scatter_axis_fn = nullptr; @@ -2075,8 +2073,7 @@ inline bool own::gemm_batched(gpu::span a, int64_t lda, bool ta, int64_t sa, unsigned gx = (un + bx - 1) / bx, gy = (um + by - 1) / by; if (gx == 0) gx = 1; if (gy == 0) gy = 1; - // kop::sgemm32 is routed to tl_sgemm by kernel_name_. - return c.launch_(c.fn_(kop::sgemm32), {gx, gy}, {bx, by}, 0, pa, pb, po, + return c.launch_(c.sgemm_fallback_(), {gx, gy}, {bx, by}, 0, pa, pb, po, pbias, um, un, uk, ula, ulb, uta, utb, scale, offset); } diff --git a/include/gpu_abi.h b/include/gpu_abi.h index ea83003..e39f287 100644 --- a/include/gpu_abi.h +++ b/include/gpu_abi.h @@ -25,14 +25,13 @@ namespace tl { namespace gpu { -// Kernel ids. The shared ops name the ones every backend may implement; the -// rest are one backend's own (Metal's GEMM tiles, its attention variants), -// kept here only until that backend keys them privately. +// Kernel ids. The shared ops name the ones every backend may implement. +// Metal's own kernels -- the GEMM tiles and the attention variants, neither +// touched by another backend's dispatch -- are keyed privately in metal.h's +// own `gtile` and `attn_op` instead, not here. enum class kop { add, sub, mul, div, pow_, exp_, log_, sqrt_, sigmoid, relu, affine, badd, bsub, bmul, bdiv, bpow, // rank-2 broadcast binary (strided operands) - sgemm32, sgemm32x64, sgemm64x32, sgemm64, - steel, steel32x64, steel_ta, steel_tb, steel32x64_ta, steel32x64_tb, softmax, row_sum, row_max, pad, fold, index_select, index_add, scatter_axis, badd_nd, bsub_nd, bmul_nd, bdiv_nd, bpow_nd, // N-D broadcast binary @@ -44,14 +43,6 @@ enum class kop { pow_s_, gt_s_, lt_s_, ge_s_, le_s_, eq_s_, ne_s_, // scalar_op maps onto these layer_norm_, // the fused layer norm layer_norm_bwd_dx_, layer_norm_bwd_gb_, layer_norm_bwd_gb_fold_, // its pullback - attn_prefill_64_, attn_prefill_128_, // causal prefill attention, per D - attn_prefill_bf16_64_, attn_prefill_bf16_128_, // over a bf16 KV cache - attn_bwd_dq_64_, attn_bwd_dq_128_, attn_bwd_dkv_64_, attn_bwd_dkv_128_, - attn_decode_64_, attn_decode_128_, // fused decode attention, per D - attn_decode_split_64_, attn_decode_split_128_, // its split-KV pass - attn_combine_64_, attn_combine_128_, // and their partials - attn_decode_bf16_64_, attn_decode_bf16_128_, // the same over a bf16 cache - attn_decode_split_bf16_64_, attn_decode_split_bf16_128_, kv_append_, kv_append_bf16_, kv_fill_, kv_fill_bf16_, // the KV cache's writes argmax_, rmsnorm_, add_rmsnorm_, swiglu_, split_heads_, merge_heads_, // decode's rest gemv_f32_, gemv_bf16_, gemv_q4_, // decode GEMVs, per weight dtype @@ -65,22 +56,15 @@ inline constexpr size_t kKopCount = static_cast(kop::adam_step_) + 1; // The ids' names, for what reports by kernel (the census, tl::profile). inline constexpr const char* kKopNames[] = { "add", "sub", "mul", "div", "pow_", "exp_", "log_", "sqrt_", "sigmoid", - "relu", "affine", "badd", "bsub", "bmul", "bdiv", "bpow", "sgemm32", - "sgemm32x64", "sgemm64x32", "sgemm64", "steel", "steel32x64", "steel_ta", - "steel_tb", "steel32x64_ta", "steel32x64_tb", "softmax", "row_sum", - "row_max", "pad", "fold", "index_select", "index_add", "scatter_axis", + "relu", "affine", "badd", "bsub", "bmul", "bdiv", "bpow", "softmax", + "row_sum", "row_max", "pad", "fold", "index_select", "index_add", + "scatter_axis", "badd_nd", "bsub_nd", "bmul_nd", "bdiv_nd", "bpow_nd", "where_nd", "copy_nd", "gt_", "lt_", "ge_", "le_", "eq_", "ne_", "tanh_", "sin_", "cos_", "clamp_", "sum_to_", "sum_to_blocked_", "concat_part_", "rope_", "pow_s_", "gt_s_", "lt_s_", "ge_s_", "le_s_", "eq_s_", "ne_s_", "layer_norm_", "layer_norm_bwd_dx_", "layer_norm_bwd_gb_", - "layer_norm_bwd_gb_fold_", "attn_prefill_64_", "attn_prefill_128_", - "attn_prefill_bf16_64_", "attn_prefill_bf16_128_", "attn_bwd_dq_64_", - "attn_bwd_dq_128_", "attn_bwd_dkv_64_", "attn_bwd_dkv_128_", - "attn_decode_64_", "attn_decode_128_", "attn_decode_split_64_", - "attn_decode_split_128_", "attn_combine_64_", "attn_combine_128_", - "attn_decode_bf16_64_", "attn_decode_bf16_128_", - "attn_decode_split_bf16_64_", "attn_decode_split_bf16_128_", "kv_append_", + "layer_norm_bwd_gb_fold_", "kv_append_", "kv_append_bf16_", "kv_fill_", "kv_fill_bf16_", "argmax_", "rmsnorm_", "add_rmsnorm_", "swiglu_", "split_heads_", "merge_heads_", "gemv_f32_", "gemv_bf16_", "gemv_q4_", "gemv_combine_", "gemv_bf16_row_", diff --git a/include/metal.h b/include/metal.h index b2b80fa..7760934 100644 --- a/include/metal.h +++ b/include/metal.h @@ -48,6 +48,68 @@ using cmp_op = gpu::cmp_op; using unary_ext_op = gpu::unary_ext_op; using scalar_op = gpu::scalar_op; +// The f32 GEMM's kernel choice: no other backend has these names, so they +// stay out of the shared kop vocabulary. context::pso_id_ caches these +// alongside a kop's pipelines, keyed negative (gtile_key_) so the two id +// spaces never collide. +enum class gtile { + sgemm32, sgemm32x64, sgemm64x32, sgemm64, + steel, steel32x64, steel_ta, steel_tb, steel32x64_ta, steel32x64_tb, +}; +constexpr int gtile_key_(gtile t) { return -1 - static_cast(t); } +inline const char* gtile_name_(gtile t) { + switch (t) { + case gtile::sgemm32: return "sgemm_32_"; + case gtile::sgemm32x64: return "sgemm_32x64_"; + case gtile::sgemm64x32: return "sgemm_64x32_"; + case gtile::sgemm64: return "sgemm_64_"; + case gtile::steel: return "sgemm_steel_"; + case gtile::steel32x64: return "sgemm_steel_32x64_"; + case gtile::steel_ta: return "sgemm_steel_ta_"; + case gtile::steel_tb: return "sgemm_steel_tb_"; + case gtile::steel32x64_ta: return "sgemm_steel_32x64_ta_"; + case gtile::steel32x64_tb: return "sgemm_steel_32x64_tb_"; + } + return ""; +} + +// The attention kernels' variant, by D and whether the KV cache is bf16: also +// no other backend's, so also kept out of kop. Offset far enough past +// gtile_key_'s range that the two never collide either. +enum class attn_op { + prefill_64, prefill_128, prefill_bf16_64, prefill_bf16_128, + bwd_dq_64, bwd_dq_128, bwd_dkv_64, bwd_dkv_128, + decode_64, decode_128, decode_split_64, decode_split_128, + combine_64, combine_128, decode_bf16_64, decode_bf16_128, + decode_split_bf16_64, decode_split_bf16_128, +}; +constexpr int attn_op_key_(attn_op t) { return -101 - static_cast(t); } +inline const char* attn_op_name_(attn_op t) { + switch (t) { + case attn_op::prefill_64: return "attn_prefill_64_"; + case attn_op::prefill_128: return "attn_prefill_128_"; + case attn_op::prefill_bf16_64: return "attn_prefill_bf16_64_"; + case attn_op::prefill_bf16_128: return "attn_prefill_bf16_128_"; + case attn_op::bwd_dq_64: return "attn_bwd_dq_64_"; + case attn_op::bwd_dq_128: return "attn_bwd_dq_128_"; + case attn_op::bwd_dkv_64: return "attn_bwd_dkv_64_"; + case attn_op::bwd_dkv_128: return "attn_bwd_dkv_128_"; + case attn_op::decode_64: return "attn_decode_64_"; + case attn_op::decode_128: return "attn_decode_128_"; + case attn_op::decode_split_64: return "attn_decode_split_64_"; + case attn_op::decode_split_128: return "attn_decode_split_128_"; + case attn_op::combine_64: return "attn_combine_64_"; + case attn_op::combine_128: return "attn_combine_128_"; + case attn_op::decode_bf16_64: return "attn_decode_bf16_64_"; + case attn_op::decode_bf16_128: return "attn_decode_bf16_128_"; + case attn_op::decode_split_bf16_64: return "attn_decode_split_bf16_64_"; + case attn_op::decode_split_bf16_128: return "attn_decode_split_bf16_128_"; + } + return ""; +} +static_assert(gtile_key_(gtile::steel32x64_tb) > attn_op_key_(attn_op::prefill_64), + "gtile and attn_op pipeline keys must stay in disjoint ranges"); + struct mtl_size { unsigned long w, h, d; }; @@ -128,16 +190,6 @@ struct context { case kop::sigmoid: return "sigmoid_"; case kop::relu: return "relu_"; case kop::affine: return "affine_"; - case kop::sgemm32: return "sgemm_32_"; - case kop::sgemm32x64: return "sgemm_32x64_"; - case kop::sgemm64x32: return "sgemm_64x32_"; - case kop::sgemm64: return "sgemm_64_"; - case kop::steel: return "sgemm_steel_"; - case kop::steel32x64: return "sgemm_steel_32x64_"; - case kop::steel_ta: return "sgemm_steel_ta_"; - case kop::steel_tb: return "sgemm_steel_tb_"; - case kop::steel32x64_ta: return "sgemm_steel_32x64_ta_"; - case kop::steel32x64_tb: return "sgemm_steel_32x64_tb_"; case kop::softmax: return "softmax_"; case kop::row_sum: return "row_sum_"; case kop::row_max: return "row_max_"; @@ -178,24 +230,6 @@ struct context { case kop::layer_norm_bwd_dx_: return "layer_norm_bwd_dx_"; case kop::layer_norm_bwd_gb_: return "layer_norm_bwd_gb_"; case kop::layer_norm_bwd_gb_fold_: return "layer_norm_bwd_gb_fold_"; - case kop::attn_prefill_64_: return "attn_prefill_64_"; - case kop::attn_prefill_128_: return "attn_prefill_128_"; - case kop::attn_prefill_bf16_64_: return "attn_prefill_bf16_64_"; - case kop::attn_prefill_bf16_128_: return "attn_prefill_bf16_128_"; - case kop::attn_bwd_dq_64_: return "attn_bwd_dq_64_"; - case kop::attn_bwd_dq_128_: return "attn_bwd_dq_128_"; - case kop::attn_bwd_dkv_64_: return "attn_bwd_dkv_64_"; - case kop::attn_bwd_dkv_128_: return "attn_bwd_dkv_128_"; - case kop::attn_decode_64_: return "attn_decode_64_"; - case kop::attn_decode_128_: return "attn_decode_128_"; - case kop::attn_decode_split_64_: return "attn_decode_split_64_"; - case kop::attn_decode_split_128_: return "attn_decode_split_128_"; - case kop::attn_combine_64_: return "attn_combine_64_"; - case kop::attn_combine_128_: return "attn_combine_128_"; - case kop::attn_decode_bf16_64_: return "attn_decode_bf16_64_"; - case kop::attn_decode_bf16_128_: return "attn_decode_bf16_128_"; - case kop::attn_decode_split_bf16_64_: return "attn_decode_split_bf16_64_"; - case kop::attn_decode_split_bf16_128_: return "attn_decode_split_bf16_128_"; case kop::kv_append_: return "kv_append_"; case kop::kv_append_bf16_: return "kv_append_bf16_"; case kop::kv_fill_: return "kv_fill_"; @@ -223,16 +257,24 @@ struct context { // Every op binds its pipeline through here: the pipeline for `op` on the // pending encoder, opened if there is none. - void bind_(kop op) { - objc::id pso = pso_(op); + void bind_(kop op) { bind_id_(static_cast(op), kernel_name_(op)); } + void bind_(gtile t) { bind_id_(gtile_key_(t), gtile_name_(t)); } + void bind_(attn_op t) { bind_id_(attn_op_key_(t), attn_op_name_(t)); } + + void bind_id_(int id, const char* name) { + objc::id pso = pso_id_(id, name); if (cb && profile::active()) commit_(nullptr); // untimed dispatches ensure_encoder_(); objc::send(enc, "setComputePipelineState:", pso); - bound = kernel_name_(op); + bound = name; } - objc::id pso_(kop op) { - auto it = psos.find(static_cast(op)); + // psos is one cache for all three: a kop's id is never negative, gtile's + // (gtile_key_) counts down from -1, attn_op's (attn_op_key_) from -101 -- + // a gap wide enough that neither range reaches the other (the static_assert + // above checks it). + objc::id pso_id_(int id, const char* name) { + auto it = psos.find(id); if (it != psos.end()) return it->second; if (!library) { objc::id err = nullptr; @@ -245,18 +287,17 @@ struct context { objc::error_str(err)); } } - auto name = objc::send(objc::cls("NSString"), "stringWithUTF8String:", - kernel_name_(op)); - auto fn = objc::send(library, "newFunctionWithName:", name); + auto oname = objc::send(objc::cls("NSString"), "stringWithUTF8String:", + name); + auto fn = objc::send(library, "newFunctionWithName:", oname); objc::id err = nullptr; auto pso = objc::send(device, "newComputePipelineStateWithFunction:error:", fn, &err); if (!pso) { throw std::runtime_error("tl::metal: PSO creation failed for " + - std::string(kernel_name_(op)) + ": " + - objc::error_str(err)); + std::string(name) + ": " + objc::error_str(err)); } - psos[static_cast(op)] = pso; + psos[id] = pso; return pso; } @@ -532,15 +573,15 @@ inline bool own::gemm(gpu::span a, int64_t lda, bool ta, gpu::span b, // shapes take the simple-tile family, which reads transposed views in // place. Gates are provisional pending a full census vs PyTorch-MPS. bool steel = !(ta && tb) && m >= 16 && n >= 48 && k >= 16; - kop kk_; + gtile kk_; unsigned long bm, bn; uint32_t fast_a, fast_b; // STEEL reuses the a_fast slot for swizzle_log unsigned long gx, gy; if (steel) { bool band32 = m < 97; - kk_ = band32 ? (ta ? kop::steel32x64_ta - : tb ? kop::steel32x64_tb : kop::steel32x64) - : (ta ? kop::steel_ta : tb ? kop::steel_tb : kop::steel); + kk_ = band32 ? (ta ? gtile::steel32x64_ta + : tb ? gtile::steel32x64_tb : gtile::steel32x64) + : (ta ? gtile::steel_ta : tb ? gtile::steel_tb : gtile::steel); bm = band32 ? 32 : 64; bn = 64; unsigned long tiles_n = (static_cast(n) + bn - 1) / bn; @@ -553,7 +594,7 @@ inline bool own::gemm(gpu::span a, int64_t lda, bool ta, gpu::span b, gx = tiles_n << swizzle_log; gy = (tiles_m + ((1ul << swizzle_log) - 1)) >> swizzle_log; } else { - kk_ = m >= 64 ? kop::sgemm64x32 : kop::sgemm32; + kk_ = m >= 64 ? gtile::sgemm64x32 : gtile::sgemm32; bm = m >= 64 ? 64 : 32; bn = 32; // float4 loader eligibility: row-major operand only (Apple GPUs handle @@ -807,7 +848,7 @@ struct copy_nd_params { // shared with binary_bcast() above -- array.h's gpu_binary_bcast_nd_ passes // the same `bk` either kernel would take); map it to its own PSO/kernel name // here rather than caching the N-D kernel under the rank-2 op's slot in -// context::psos, which pso_() keys by this same enum value. +// context::psos, which pso_id_ keys by this same enum value. inline kop to_nd_(kop op) { switch (op) { case kop::badd: return kop::badd_nd; @@ -1104,9 +1145,9 @@ inline bool own::attn_prefill(gpu::span q, gpu::span K, gpu::span V, auto& c = context::get(); if (!c.device || (D != 64 && D != 128)) return false; if (n_kv_heads <= 0 || n_q_heads % n_kv_heads != 0 || T <= 0) return false; - c.bind_(kv_bf16 ? (D == 64 ? kop::attn_prefill_bf16_64_ - : kop::attn_prefill_bf16_128_) - : (D == 64 ? kop::attn_prefill_64_ : kop::attn_prefill_128_)); + c.bind_(kv_bf16 ? (D == 64 ? attn_op::prefill_bf16_64 + : attn_op::prefill_bf16_128) + : (D == 64 ? attn_op::prefill_64 : attn_op::prefill_128)); detail_::set_buf_(c.enc, q, 0ul); detail_::set_buf_(c.enc, K, 1ul); detail_::set_buf_(c.enc, V, 2ul); @@ -1210,14 +1251,14 @@ inline bool own::attn_decode(gpu::span q, gpu::span K, gpu::span V, } const bool split = chunk != 0; if (split) { - c.bind_(kv_bf16 ? (d64 ? kop::attn_decode_split_bf16_64_ - : kop::attn_decode_split_bf16_128_) - : (d64 ? kop::attn_decode_split_64_ - : kop::attn_decode_split_128_)); + c.bind_(kv_bf16 ? (d64 ? attn_op::decode_split_bf16_64 + : attn_op::decode_split_bf16_128) + : (d64 ? attn_op::decode_split_64 + : attn_op::decode_split_128)); } else { - c.bind_(kv_bf16 ? (d64 ? kop::attn_decode_bf16_64_ - : kop::attn_decode_bf16_128_) - : (d64 ? kop::attn_decode_64_ : kop::attn_decode_128_)); + c.bind_(kv_bf16 ? (d64 ? attn_op::decode_bf16_64 + : attn_op::decode_bf16_128) + : (d64 ? attn_op::decode_64 : attn_op::decode_128)); } detail_::set_buf_(c.enc, q, 0ul); detail_::set_buf_(c.enc, K, 1ul); @@ -1233,7 +1274,7 @@ inline bool own::attn_decode(gpu::span q, gpu::span K, gpu::span V, split ? splits : 1ul, 1}, {static_cast(D), 1, 1}); if (!split) return true; - c.bind_(d64 ? kop::attn_combine_64_ : kop::attn_combine_128_); + c.bind_(d64 ? attn_op::combine_64 : attn_op::combine_128); detail_::set_buf_(c.enc, dst, 0ul); detail_::set_buf_(c.enc, out, 1ul); detail_::attn_combine_params cp{static_cast(splits)}; @@ -1254,7 +1295,7 @@ inline bool own::attn_prefill_dq(gpu::span q, gpu::span K, gpu::span V, int64_t D, float scale) { auto& c = context::get(); if (!c.device || (D != 64 && D != 128) || H <= 0 || T <= 0) return false; - c.bind_(D == 64 ? kop::attn_bwd_dq_64_ : kop::attn_bwd_dq_128_); + c.bind_(D == 64 ? attn_op::bwd_dq_64 : attn_op::bwd_dq_128); const gpu::span views[] = {q, K, V, dO, O, dq, stats}; for (unsigned long i = 0; i < 7; i++) detail_::set_buf_(c.enc, views[i], i); detail_::attn_params p{static_cast(T), 0, 0, 0, scale}; @@ -1269,7 +1310,7 @@ inline bool own::attn_prefill_dkv(gpu::span q, gpu::span K, gpu::span V, float scale) { auto& c = context::get(); if (!c.device || (D != 64 && D != 128) || H <= 0 || T <= 0) return false; - c.bind_(D == 64 ? kop::attn_bwd_dkv_64_ : kop::attn_bwd_dkv_128_); + c.bind_(D == 64 ? attn_op::bwd_dkv_64 : attn_op::bwd_dkv_128); const gpu::span views[] = {q, K, V, dO, stats, dK, dV}; for (unsigned long i = 0; i < 7; i++) detail_::set_buf_(c.enc, views[i], i); detail_::attn_params p{static_cast(T), 0, 0, 0, scale}; From e85d5e68e47a904f44b24031769196b2278011fd Mon Sep 17 00:00:00 2001 From: yhirose Date: Tue, 22 Sep 2026 12:32:44 -0400 Subject: [PATCH 3/3] gpu: the mirror table and its buffer pool live once, in gpu_abi.h cuda.h and webgpu.h each kept their own copy of the same shape: a struct pairing a host buffer with a device handle and a residency state, a map from the native handle to it, and a size-keyed free list alloc()/release() recycle through. Only the device handle's type differed (CUdeviceptr, wgpu::Buffer). gpu::mirror_table (gpu_abi.h) holds that shape once; each backend instantiates it for its own handle type and keeps doing its own copying (before_kernel_, sync_to_host), which is where the two drivers actually diverge. mirror_table::insert moves its Handle parameter into the entry rather than copying it, so a ref-counted handle (wgpu::Buffer) picks up only one AddRef per allocation instead of two. Verified on native Metal (117 cases, 11640 assertions with real device access, plus check_qwen's greedy tokens against the numpy oracle) and gpu_host's no-GPU reference build, under WebGPU via Deno, and by a clean CUDA host-side trace diff against the pre-refactor commit (no kernel call changed). A mutation that let take() hand out an already-live pooled buffer twice was caught hard by the WebGPU suite (32 of 117 cases failed). --- docs/backends.md | 12 +++++++--- include/cuda.h | 57 ++++++++++++++++------------------------------- include/gpu_abi.h | 51 ++++++++++++++++++++++++++++++++++++++++++ include/webgpu.h | 44 ++++++++++++------------------------ 4 files changed, 93 insertions(+), 71 deletions(-) diff --git a/docs/backends.md b/docs/backends.md index c0a72e1..82f13bb 100644 --- a/docs/backends.md +++ b/docs/backends.md @@ -60,6 +60,15 @@ buffer, so an output into a buffer whose live bytes are the host's has to bring them up or lose the rest, while an output into a fresh buffer, which is nearly every output, has nothing to bring. +The table that holds those copies — host pointer, device handle, size, +`residency` — and the size-keyed free list `alloc`/`release` recycle buffers +through are `gpu::mirror_table` (`gpu_abi.h`), one instantiation per +mirrored backend (`mirror_table` in `cuda.h`, +`mirror_table` in `webgpu.h`). Only `Handle` — the backend's own +device buffer type — varies; the map, the pool and the copying policy do not. +A backend still does the actual copy (`before_kernel_`, `sync_to_host`) since +that call differs by driver. + ## The kernel ABI An op hands the backend a kernel id, an ordered list of views, a params @@ -259,9 +268,6 @@ asks; a new backend edits no test. ## What is not shared yet -- The mirror table (handle to host copy, device copy, size, `residency`) and - the buffer pool exist twice, in `cuda.h` and `webgpu.h`. The state machine - itself is shared. - WebGPU's kernels still predate the canonical ABI and go through `marshal_` (see above). diff --git a/include/cuda.h b/include/cuda.h index f3e4bf4..9860359 100644 --- a/include/cuda.h +++ b/include/cuda.h @@ -282,29 +282,19 @@ struct context { std::vector timed; std::vector spare_events; - // Host/device mirror per allocation, keyed by the device pointer (== the - // `native` handle stored in storage). Views sharing a storage share the key, - // so one state serves every view. When to copy is gpu::residency's decision - // (gpu_abi.h); the copies are made here. - struct mirror { - float* host = nullptr; // CPU-side buffer (storage.contents/ptr) - CUdeviceptr dev = 0; // device buffer (storage.native) - size_t bytes = 0; - gpu::residency live; - }; - std::unordered_map mirrors; - - // Size-keyed free list (like Metal's MTLBuffer pool). Released buffers are - // recycled, not cuMemFree'd — repeated large alloc/free otherwise fragments - // the driver allocator (decode benches, training that churns activations). - // Buffers persist until the (leaked) context tears down. Keyed by exact byte - // size; the workloads that churn reuse identical shapes. - std::unordered_map>> pool; - - mirror* mirror_(void* native) { - auto it = mirrors.find(reinterpret_cast(native)); - return it == mirrors.end() ? nullptr : &it->second; - } + // Host/device mirror per allocation and its size-keyed free list (shared + // shape with webgpu.h; `gpu::mirror_table`, gpu_abi.h). Keyed by the device + // pointer, reinterpreted as the `native` handle stored in storage. Views + // sharing a storage share the key, so one state serves every view. When to + // copy is gpu::residency's decision; the copies are made here. Released + // buffers are recycled, not cuMemFree'd — repeated large alloc/free + // otherwise fragments the driver allocator (decode benches, training that + // churns activations); they persist until the (leaked) context tears down. + using mirror = gpu::mirror_table::entry; + gpu::mirror_table mt; + + mirror* mirror_(void* native) { return mt.find(native); } + // A kernel is about to touch this buffer as `a`: bring the host copy up if // residency says so. Async on the stream like the meta uploads (a blocking // copy would wait out every kernel already queued and stall the pipeline @@ -1141,12 +1131,7 @@ inline void* alloc(int64_t bytes, float** contents, bool host_fill = false) { size_t nb = bytes > 0 ? (size_t)bytes : 4; CUdeviceptr dev = 0; float* host = nullptr; - auto it = c.pool.find(nb); // reuse a recycled buffer of this exact size - if (it != c.pool.end() && !it->second.empty()) { - dev = it->second.back().first; - host = it->second.back().second; - it->second.pop_back(); - } else { + if (!c.mt.take(nb, dev, host)) { // no recycled buffer of this exact size if (c.d.MemAlloc(&dev, nb) != 0) return nullptr; host = static_cast(std::malloc(nb)); if (!host) { @@ -1154,22 +1139,18 @@ inline void* alloc(int64_t bytes, float** contents, bool host_fill = false) { return nullptr; } } - c.mirrors[dev] = context::mirror{host, dev, nb, gpu::residency(host_fill)}; + void* native = reinterpret_cast(dev); + c.mt.insert(native, dev, host, nb, host_fill); if (contents) *contents = host; - return reinterpret_cast(dev); + return native; } inline void release(void* buf, int64_t, float*) { auto& c = context::get(); if (!c.ready || !buf) return; - CUdeviceptr dev = reinterpret_cast(buf); - auto it = c.mirrors.find(dev); - if (it == c.mirrors.end()) { - c.d.MemFree(dev); // untracked (shouldn't happen); free outright - return; + if (!c.mt.release(buf)) { + c.d.MemFree(reinterpret_cast(buf)); // untracked (shouldn't happen); free outright } - c.pool[it->second.bytes].push_back({dev, it->second.host}); // recycle - c.mirrors.erase(it); } // Reconcile a buffer for a CPU access: when the device holds the live copy, diff --git a/include/gpu_abi.h b/include/gpu_abi.h index e39f287..d169dca 100644 --- a/include/gpu_abi.h +++ b/include/gpu_abi.h @@ -19,6 +19,9 @@ #include #include #include +#include +#include +#include #include "profile.h" @@ -160,6 +163,54 @@ struct residency { void uploaded() { where = both; } }; +// The host/device mirror and size-keyed buffer pool a mirrored backend (CUDA, +// WebGPU) keeps per allocation. `Handle` is the backend's own device buffer +// type (`CUdeviceptr`, `wgpu::Buffer`); the map, the pool and the entry shape +// are otherwise identical between them, so only `Handle` varies. A backend +// still does its own copying (`before_kernel_`, `sync_to_host`) and still +// keys lookups by whatever `void*` it hands out as `native`; this only owns +// where the state lives. +template +struct mirror_table { + struct entry { + float* host = nullptr; + Handle dev{}; + size_t bytes = 0; + residency live; + }; + std::unordered_map mirrors; + std::unordered_map>> pool; + + entry* find(void* native) { + auto it = mirrors.find(native); + return it == mirrors.end() ? nullptr : &it->second; + } + + // A released buffer of this exact size, if the pool has one. + bool take(size_t bytes, Handle& dev, float*& host) { + auto it = pool.find(bytes); + if (it == pool.end() || it->second.empty()) return false; + dev = it->second.back().first; + host = it->second.back().second; + it->second.pop_back(); + return true; + } + + void insert(void* key, Handle dev, float* host, size_t bytes, bool host_fill) { + mirrors[key] = entry{host, std::move(dev), bytes, residency(host_fill)}; + } + + // Moves `native`'s entry into the free list, keyed by its size. False if + // `native` is not tracked (nothing to release). + bool release(void* native) { + auto it = mirrors.find(native); + if (it == mirrors.end()) return false; + pool[it->second.bytes].push_back({it->second.dev, it->second.host}); + mirrors.erase(it); + return true; + } +}; + // A launch's extent: how many groups, how many threads in each, and the bytes // of per-group scratch a kernel's reduction needs where the backend sizes it at // launch (CUDA's shared memory; Metal and WGSL size theirs in the kernel). diff --git a/include/webgpu.h b/include/webgpu.h index 6264257..70b478f 100644 --- a/include/webgpu.h +++ b/include/webgpu.h @@ -196,21 +196,16 @@ struct context { void* meta_ring_(float** host_out); uint32_t meta_reserve_slot_(); - // Host/device mirror per allocation, keyed by the opaque handle alloc() - // returns as `native`. Views sharing a storage share the key, so one dirty - // state serves every view. `where` tracks which copy is live. - struct mirror { - float* host = nullptr; // CPU-side buffer (storage.ptr) - wgpu::Buffer dev; // device buffer - size_t bytes = 0; - gpu::residency live; // when to copy (gpu_abi.h); the copies are made here - }; - std::unordered_map mirrors; - - // Size-keyed free lists (like Metal's MTLBuffer pool and CUDA's). Repeated - // alloc/free of identical shapes is the common case, and per-dispatch - // allocation would compound the fixed dispatch floor. - std::unordered_map>> pool; + // Host/device mirror per allocation and its size-keyed free list (shared + // shape with cuda.h; `gpu::mirror_table`, gpu_abi.h), keyed by the opaque + // handle alloc() returns as `native`. Views sharing a storage share the + // key, so one dirty state serves every view. + using mirror = gpu::mirror_table::entry; + gpu::mirror_table mt; + + // Mapped-for-readback staging buffers (sync_to_host's D2H), a distinct pool + // from the mirror table's device buffers: these are never bound as a + // kernel operand, only mapped. std::unordered_map> staging_pool; // Compute pipelines, keyed by WGSL entry point. Every one is built in the @@ -347,10 +342,7 @@ struct context { return instance.WaitAny(f, UINT64_MAX) == wgpu::WaitStatus::Success; } - mirror* mirror_(void* native) { - auto it = mirrors.find(native); - return it == mirrors.end() ? nullptr : &it->second; - } + mirror* mirror_(void* native) { return mt.find(native); } // A kernel is about to touch this buffer as `a`: bring the host copy up if // residency says so. @@ -512,12 +504,7 @@ inline void* alloc(int64_t bytes, float** contents, bool host_fill = false) { wgpu::Buffer dev; float* host = nullptr; - auto it = c.pool.find(nb); // reuse a recycled buffer of this exact size - if (it != c.pool.end() && !it->second.empty()) { - dev = it->second.back().first; - host = it->second.back().second; - it->second.pop_back(); - } else { + if (!c.mt.take(nb, dev, host)) { // no recycled buffer of this exact size wgpu::BufferDescriptor d = {}; d.size = nb; d.usage = wgpu::BufferUsage::Storage | wgpu::BufferUsage::CopyDst | @@ -532,7 +519,7 @@ inline void* alloc(int64_t bytes, float** contents, bool host_fill = false) { // host pointer is both, and malloc will not hand out the same address twice // while it is live. void* token = host; - c.mirrors[token] = context::mirror{host, dev, nb, gpu::residency(host_fill)}; + c.mt.insert(token, dev, host, nb, host_fill); if (contents) *contents = host; return token; } @@ -540,10 +527,7 @@ inline void* alloc(int64_t bytes, float** contents, bool host_fill = false) { inline void release(void* buf, int64_t, float*) { auto& c = context::get(); if (!c.ready || !buf) return; - auto it = c.mirrors.find(buf); - if (it == c.mirrors.end()) return; - c.pool[it->second.bytes].push_back({it->second.dev, it->second.host}); - c.mirrors.erase(it); + c.mt.release(buf); } // Reconcile a buffer for a CPU access: flush pending kernels, then D2H if the