diff --git a/docs/backends.md b/docs/backends.md index 79c8805..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 @@ -144,23 +153,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 @@ -192,6 +205,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 @@ -250,14 +268,8 @@ 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. -- `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/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/cuda.h b/include/cuda.h index 59123a7..9860359 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 } } @@ -289,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 @@ -517,6 +500,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; @@ -1143,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) { @@ -1156,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, @@ -2075,8 +2054,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..d169dca 100644 --- a/include/gpu_abi.h +++ b/include/gpu_abi.h @@ -19,20 +19,22 @@ #include #include #include +#include +#include +#include #include "profile.h" 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 +46,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 +59,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_", @@ -176,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/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/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}; 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 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; }