diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index dba0d2c..dee65be 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -6,20 +6,6 @@ on: pull_request: jobs: - # Static, no build needed: is a GPU op real on the same backends everywhere, - # or did one backend land ahead of the others and never get caught up (that - # is exactly what happened to concat_part/rope before this job existed — - # see tools/check_backend_parity.py's own docstring). Deliberate gaps (the - # M9 LLM decode fast path, still CUDA-only) are recorded in - # tools/backend_parity_allowlist.txt and don't fail this job; an - # unallowlisted asymmetry does. - backend-parity: - name: backend parity (CUDA / Metal / WebGPU) - runs-on: ubuntu-latest - steps: - - uses: actions/checkout@v4 - - run: python3 tools/check_backend_parity.py - # All three modes — cpu / gpu / auto — on every OS. # What each hosted runner actually offers: # - macos-15 (arm64): Metal works through a paravirtual GPU -> a real GPU test @@ -40,6 +26,27 @@ jobs: - name: Test (cpu / gpu / auto) run: ctest --test-dir build -C Release --output-on-failure + # The host reference backend (gpu_host.h): a "device" that is the CPU, so the + # shared GPU layer — gpu_ops.h, the launch policy, residency, the census — and + # the conformance tests run on runners with no GPU, where the plain `test` job + # above only ever takes the CPU fallback and never enters them. On macOS it + # also checks that asking for it by name wins over Metal. + host-backend: + name: host reference backend / ${{ matrix.os }} + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + os: [ubuntu-latest, macos-15] + steps: + - uses: actions/checkout@v4 + - name: Configure + run: cmake -B build -DCMAKE_BUILD_TYPE=Release -DTENSORLIB_HOST_GPU=ON + - name: Build + run: cmake --build build --config Release --target tensorlib_test + - name: Test (cpu / gpu / auto) + run: ctest --test-dir build -C Release --output-on-failure -R '^(cpu|gpu|auto)$' + # Build verification for the CUDA toolchain. Hosted runners have no NVIDIA # driver, so the kernels are only compiled and linked. Running the tests here # doubles as coverage of the dlopen fallback on a driver-less machine, which diff --git a/CMakeLists.txt b/CMakeLists.txt index 903886c..db29c7f 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -10,6 +10,7 @@ set(CMAKE_CXX_STANDARD 23) set(CMAKE_CXX_STANDARD_REQUIRED ON) option(TENSORLIB_CUDA "Build the CUDA backend (needs CUDA Toolkit at build time only)" OFF) +option(TENSORLIB_HOST_GPU "Select the host reference backend (gpu_host.h): runs the GPU layer and its conformance tests with no GPU" OFF) add_library(tensorlib INTERFACE) target_include_directories(tensorlib INTERFACE ${CMAKE_CURRENT_SOURCE_DIR}/include) @@ -19,6 +20,10 @@ if(APPLE) "-framework Accelerate" "-framework Metal" "-framework Foundation") endif() +if(TENSORLIB_HOST_GPU) + target_compile_definitions(tensorlib INTERFACE TENSORLIB_HOST_GPU) +endif() + if(TENSORLIB_CUDA) # The CUDA backend does NOT link CUDA — the driver is loaded at runtime # (dlopen on Unix, LoadLibrary(nvcuda.dll) on Windows; see cuda.h) and kernels diff --git a/README.md b/README.md index a20380c..0595506 100644 --- a/README.md +++ b/README.md @@ -161,11 +161,20 @@ that is derived from the pool's measured wake-up on first use (`TL_CPU_MIN_WORK` pins it; `misc/census_pool_latency.cpp` shows the arithmetic); below that it runs on the calling thread. -New ops sometimes land on one backend (usually CUDA) before the others catch -up. `tools/check_backend_parity.py` reports, per op, which of CUDA/Metal/ -WebGPU actually implement it rather than falling back to the CPU oracle, and -fails (in CI too) if an asymmetry isn't recorded in -`tools/backend_parity_allowlist.txt` as deliberate. +The GPU ops are written once, in `gpu_ops.h`, over views (`gpu::span`: a device +handle and a byte offset) and a kernel ABI every backend shares (`gpu_abi.h`). A +backend is a device core — memory, a kernel table, one `dispatch` — plus kernel +source, and the ops it runs its own way as members of its `own` struct; an op a +backend has no kernel for answers false and the evaluator falls back to the CPU. +`gpu::census(kernel)` counts launches, so a test can tell the two apart. No +machine here runs CUDA kernels, so `tools/cuda_trace` records what the CUDA +backend asks of the driver (kernel, grid, every argument) against a stand-in +`libcuda` in a Linux container, and `tools/cuda_trace/compare.sh ` diffs +that across a change to the backend's host side. +`-DTENSORLIB_HOST_GPU=ON` selects a reference backend whose device is the CPU +(`gpu_host.h`), which runs that whole layer and its conformance tests with no +GPU. [docs/backends.md](docs/backends.md) covers the layers, the kernel ABI, and +how to add an op or a backend. ### Profiling diff --git a/bench/cuda/check/check_attn64.cpp b/bench/cuda/check/check_attn64.cpp index 9a8d074..8b98b45 100644 --- a/bench/cuda/check/check_attn64.cpp +++ b/bench/cuda/check/check_attn64.cpp @@ -15,7 +15,7 @@ #ifndef TENSORLIB_CUDA #define TENSORLIB_CUDA #endif -#include "cuda.h" +#include "gpu.h" // cuda.h plus the shared ops (tl::gpu resolves to cuda here) #include "kv_cache.h" #include @@ -80,13 +80,13 @@ int main() { } sync_to_host(kn, true); sync_to_host(vn, true); - if (!cache.append(kn, vn)) { std::printf(" append failed pos %lld\n", (long long)pos); return 1; } + if (!cache.append({kn, 0}, {vn, 0})) { std::printf(" append failed pos %lld\n", (long long)pos); return 1; } if (ci < sizeof(checks) / sizeof(checks[0]) && step == checks[ci]) { ci++; for (int64_t i = 0; i < HQ * D; i++) hq[i] = rnd(); sync_to_host(q, true); - cache.attn(q, o, HQ, scale); + cache.attn({q, 0}, {o, 0}, HQ, scale); flush(); sync_to_host(o, false); // Lockstep guard: attn_dpos (split-KV on the capacity-static grid, @@ -94,7 +94,7 @@ int main() { // device attn_dpos_chunk twins the host attn_split_count/attn_split_chunk // heuristic, and the capacity grid's empty splits combine as exact zeros. upload_u32(dp, (unsigned)pos); - cache.attn_dpos(q, o2, HQ, dp, scale); + cache.attn_dpos({q, 0}, {o2, 0}, HQ, {dp, 0}, scale); flush(); sync_to_host(o2, false); if (std::memcmp(ho, ho2, (size_t)HQ * D * 4) != 0) { @@ -150,7 +150,7 @@ int main() { tl::kv_cache cache; if (!cache.init(HKV, MAXC, D)) { std::printf(" cache init failed\n"); return 1; } - cache.prefill(qp, ks, vs, op, T, HQ, scale); + cache.prefill({qp, 0}, {ks, 0}, {vs, 0}, {op, 0}, T, HQ, scale); flush(); sync_to_host(op, false); @@ -192,8 +192,8 @@ int main() { for (int64_t i = 0; i < HKV * D; i++) { hkn[i] = rnd(); hvn[i] = rnd(); } for (int64_t i = 0; i < HQ * D; i++) hq1[i] = rnd(); sync_to_host(kn, true); sync_to_host(vn, true); sync_to_host(q1, true); - cache.append(kn, vn); - cache.attn(q1, o1, HQ, scale); + cache.append({kn, 0}, {vn, 0}); + cache.attn({q1, 0}, {o1, 0}, HQ, scale); flush(); sync_to_host(o1, false); diff --git a/bench/cuda/check/check_cuda.cpp b/bench/cuda/check/check_cuda.cpp index 10ee91c..3de2f95 100644 --- a/bench/cuda/check/check_cuda.cpp +++ b/bench/cuda/check/check_cuda.cpp @@ -6,7 +6,7 @@ #ifndef TENSORLIB_CUDA #define TENSORLIB_CUDA // standalone build; the CMake build passes it as a flag #endif -#include "cuda.h" +#include "gpu.h" // the shared ops (tl::gpu) resolve to cuda here #include #include @@ -14,7 +14,7 @@ #include using namespace tl::cuda; -using kop = tl::metal::kop; +using kop = tl::gpu::kop; static int failures = 0; static void check(bool ok, const char* what) { @@ -49,7 +49,8 @@ int main() { cb[i] = d(g); ref[i] = (ca[i] + cb[i]) * 2.0f + 1.0f; } - bool launched = tl::cuda::binary(kop::add, a, 0, b, 0, o, 0, n, 2.0f, 1.0f); + bool launched = + tl::gpu::binary(kop::add, {a, 0}, {b, 0}, {o, 0}, n, 2.0f, 1.0f); sync_to_host(o, false); // D2H the device-written output before host read bool match = launched; for (int64_t i = 0; i < n; i++) @@ -72,7 +73,7 @@ int main() { ca[i] = d(g) * 4.0f; ref[i] = 1.0f / (1.0f + std::exp(-ca[i])); } - bool launched = tl::cuda::unary(kop::sigmoid, a, 0, o, 0, n, 1.0f, 0.0f); + bool launched = tl::gpu::unary(kop::sigmoid, {a, 0}, {o, 0}, n, 1.0f, 0.0f); sync_to_host(o, false); // D2H the device-written output before host read bool match = launched; for (int64_t i = 0; i < n; i++) @@ -100,7 +101,7 @@ int main() { for (int64_t p = 0; p < k; p++) s += ca[i * k + p] * cb[p * n + j]; ref[i * n + j] = s * 0.5f; } - bool launched = gemm(a, 0, k, false, b, 0, n, false, o, 0, m, n, k, 0.5f, 0); + bool launched = tl::gpu::gemm({a, 0}, k, false, {b, 0}, n, false, {o, 0}, m, n, k, 0.5f, 0); sync_to_host(o, false); // D2H the device-written output before host read bool match = launched; for (int64_t i = 0; i < m * n; i++) @@ -131,7 +132,7 @@ int main() { ref[i * n + j] = s; } // trans_b=true, ldb=k (the col stride of the logical k×n = row stride of bt) - bool launched = gemm(a, 0, k, false, b, 0, k, true, o, 0, m, n, k, 1.0f, 0); + bool launched = tl::gpu::gemm({a, 0}, k, false, {b, 0}, k, true, {o, 0}, m, n, k, 1.0f, 0); sync_to_host(o, false); // D2H the device-written output before host read bool match = launched; for (int64_t i = 0; i < m * n; i++) @@ -159,7 +160,7 @@ int main() { } ref[r] = s * 0.25f + 3.0f; } - bool launched = tl::cuda::row_op(kop::row_sum, in, 0, o, 0, rows, cols, 0.25f, 3.0f); + bool launched = tl::gpu::row_op(kop::row_sum, {in, 0}, {o, 0}, rows, cols, 0.25f, 3.0f); sync_to_host(o, false); // D2H the device-written output before host read bool match = launched; for (int64_t r = 0; r < rows; r++) @@ -187,7 +188,7 @@ int main() { for (int64_t c = 0; c < cols; c++) ref[r * cols + c] = std::exp(ci[r * cols + c] - mx) / sum; } - bool launched = tl::cuda::row_op(kop::softmax, in, 0, o, 0, rows, cols, 1.0f, 0.0f); + bool launched = tl::gpu::row_op(kop::softmax, {in, 0}, {o, 0}, rows, cols, 1.0f, 0.0f); sync_to_host(o, false); // D2H the device-written output before host read bool match = launched; for (int64_t i = 0; i < rows * cols; i++) diff --git a/bench/cuda/check/check_llm_decode.cpp b/bench/cuda/check/check_llm_decode.cpp index fa22176..fef2ce5 100644 --- a/bench/cuda/check/check_llm_decode.cpp +++ b/bench/cuda/check/check_llm_decode.cpp @@ -205,9 +205,9 @@ vec gpu_step(GpuModel& m, int64_t id, int64_t pos) { array k = array::rope(h.dot(L.Wk).reshape({HKV, hd}), pos); array v = h.dot(L.Wv); // [1, HKV*hd] == [HKV, hd] bytes q.eval(); k.eval(); v.eval(); - L.cache.append(k.native(), v.native()); + L.cache.append(k.device_span(), v.device_span()); array a_out = array::empty({HQ, hd}); - L.cache.attn(q.native(), a_out.native(), HQ, SCALE); + L.cache.attn(q.device_span(), a_out.device_span(), HQ, SCALE); array a = a_out.reshape({1, Dm}); array x1 = x + a.dot(L.Wo); array h2 = array::rmsnorm(x1, L.n2, EPS); diff --git a/bench/cuda/check/trace_sweep.cpp b/bench/cuda/check/trace_sweep.cpp new file mode 100644 index 0000000..b94f847 --- /dev/null +++ b/bench/cuda/check/trace_sweep.cpp @@ -0,0 +1,128 @@ +// What the launch trace (tools/cuda_trace) cannot reach through the test suite +// and the checkers: the kernels no test selects, and the calls no evaluator +// makes — above all a non-zero OUTPUT offset, which array.h never passes but +// the contract allows. It checks nothing. It exists to be traced, so that a +// change to cuda.h's host side shows up in a diff on a machine with no GPU. +#ifndef TENSORLIB_CUDA +#define TENSORLIB_CUDA +#endif +#include "gpu.h" // the shared ops (tl::gpu) resolve to cuda here + +#include +#include + +namespace cu = tl::cuda; +using kop = tl::gpu::kop; + +namespace { + +struct buf { + void* native = nullptr; + float* host = nullptr; + int64_t bytes = 0; + explicit buf(int64_t b) : bytes(b) { native = cu::alloc(b, &host); } + ~buf() { cu::release(native, bytes, host); } + buf(const buf&) = delete; + operator void*() const { return native; } +}; + +constexpr int64_t kOff = 64; // bytes; 16B-aligned so the tiled GEMM still takes it + +void elementwise() { + const int64_t n = 1000; + buf a(n * 4 + kOff), b(n * 4 + kOff), o(n * 4 + kOff); + for (kop op : {kop::add, kop::sub, kop::mul, kop::div, kop::pow_}) { + tl::gpu::binary(op, {a, kOff}, {b, 0}, {o, kOff}, n, 2.0f, 1.0f); + } + for (kop op : {kop::badd, kop::bsub, kop::bmul, kop::bdiv, kop::bpow}) { + tl::gpu::binary_bcast(op, {a, kOff}, 40, 1, {b, 0}, 0, 1, {o, kOff}, 25, 40, + 1.0f, 0.0f); + } + const int64_t shape[3] = {5, 8, 25}, as[3] = {200, 25, 1}, bs[3] = {0, 25, 1}; + for (kop op : {kop::badd, kop::bsub, kop::bmul, kop::bdiv, kop::bpow}) { + tl::gpu::binary_bcast_nd(op, {a, kOff}, as, {b, 0}, bs, {o, kOff}, shape, 3, n, 1.0f, 0.0f); + } + using cu::cmp_op; + for (cmp_op op : {cmp_op::gt, cmp_op::lt, cmp_op::ge, cmp_op::le, cmp_op::eq, + cmp_op::ne}) { + tl::gpu::compare(op, {a, kOff}, {b, 0}, {o, kOff}, n, 1); + } +} + +void gemm_bias() { + struct shape { int64_t m, n, k; }; + for (shape s : {shape{8, 8, 8}, {64, 64, 64}, {256, 512, 512}, + {256, 2048, 512}, {512, 896, 896}}) { + buf a(s.m * s.k * 4 + kOff), b(s.k * s.n * 4 + kOff), bias(s.n * 4), + o(s.m * s.n * 4 + kOff); + for (int layout = 0; layout < 4; layout++) { + const bool ta = layout & 1, tb = layout & 2; + tl::gpu::gemm_bias({a, kOff}, ta ? s.m : s.k, ta, {b, kOff}, tb ? s.k : s.n, tb, {bias, 0}, {o, kOff}, s.m, s.n, s.k, 1.0f, 0.0f); + } + } +} + +// The four ops that zero their output before scattering into it. +void zero_then_scatter() { + const int64_t a_shape[2] = {4, 6}, out_shape[2] = {4, 10}; + buf a(24 * 4 + kOff), o(40 * 4 + kOff); + tl::gpu::pad({a, kOff}, {o, kOff}, a_shape, out_shape, 2, 1, 2, 24, 40); + + const int64_t w_shape[3] = {4, 4, 3}, f_shape[2] = {4, 6}; + buf w(48 * 4 + kOff), f(24 * 4 + kOff); + tl::gpu::fold({w, kOff}, {f, kOff}, w_shape, f_shape, 3, 1, 1, 48, 24); + + buf idx(8 * 4 + kOff), vals(8 * 16 * 4 + kOff), table(32 * 16 * 4 + kOff); + tl::gpu::index_add({idx, kOff}, {vals, kOff}, {table, kOff}, 16, 8, 32 * 16); + tl::gpu::scatter_to_axis({idx, kOff}, {vals, kOff}, {table, kOff}, 8, 16); +} + +void llm() { + for (int64_t m : {8, 96, 512}) { + for (int64_t n : {896, 4864}) { + const int64_t k = 896; + buf a(m * k * 4), B(n * k * 2), o(m * n * 4); + tl::gpu::gemm_bf16_nt({a, 0}, {B, 0}, {o, 0}, m, n, k); + } + } + for (int64_t D : {64, 128}) { + const int64_t hq = 14, hkv = 2, kv_max = 4096; + buf q(hq * D * 4), K(hkv * kv_max * D * 2), V(hkv * kv_max * D * 2), + o(hq * D * 4); + for (int64_t ctx : {17, 700, 4000}) { + tl::gpu::attn_decode({q, 0}, {K, 0}, {V, 0}, {o, 0}, hq, hkv, ctx, kv_max, D, 0.125f, true); + } + } +} + +// The graph-capture group: device-position variants, recorded then replayed. +void capture(int64_t D) { + const int64_t hq = 14, hkv = 2, kv_max = 2048; + buf pos(4), x(hq * D * 4), xo(hq * D * 4), knew(hkv * D * 4), vnew(hkv * D * 4), + K(hkv * kv_max * D * 4), V(hkv * kv_max * D * 4), o(hq * D * 4), + partials(cu::attn_dpos_partials_bytes(hq, kv_max, D)); + cu::upload_u32(pos, 5); + if (!cu::capture_begin()) return; + tl::gpu::rope_dpos({x, 0}, {xo, 0}, hq, 1, D, {pos, 0}, 10000.0f); + tl::gpu::kv_append_dpos({K, 0}, {V, 0}, {knew, 0}, {vnew, 0}, {pos, 0}, kv_max, hkv, D); + tl::gpu::attn_decode_dpos({xo, 0}, {K, 0}, {V, 0}, {o, 0}, hq, hkv, {pos, 0}, kv_max, D, 0.125f, {partials, 0}); + tl::gpu::incr_u32({pos, 0}); + auto graph = cu::capture_end(); + cu::graph_launch(graph); + cu::flush(); + cu::graph_destroy(graph); +} + +} // namespace + +int main() { + if (!cu::available()) return std::printf("no CUDA driver: nothing to trace\n"), 0; + elementwise(); + gemm_bias(); + zero_then_scatter(); + llm(); + capture(64); + capture(128); + cu::flush(); + return 0; +} diff --git a/bench/cuda/speed/bench_attn_bwd.cpp b/bench/cuda/speed/bench_attn_bwd.cpp index 89d25d1..9d35d69 100644 --- a/bench/cuda/speed/bench_attn_bwd.cpp +++ b/bench/cuda/speed/bench_attn_bwd.cpp @@ -12,7 +12,7 @@ #ifndef TENSORLIB_CUDA #define TENSORLIB_CUDA #endif -#include "cuda.h" +#include "gpu.h" // cuda.h plus the shared ops (tl::gpu resolves to cuda here) #include #include @@ -72,13 +72,13 @@ int main() { fill_random(hG, n, 4); auto run_fwd = [&] { - attn_prefill(q, K, V, out, s.H, s.H, s.T, s.T, s.D, scale); + tl::gpu::attn_prefill({q, 0}, {K, 0}, {V, 0}, {out, 0}, s.H, s.H, s.T, s.T, s.D, scale); }; auto run_dq = [&] { - attn_prefill_dq(q, K, V, dO, out, dq, stats, s.H, s.T, s.D, scale); + tl::gpu::attn_prefill_dq({q, 0}, {K, 0}, {V, 0}, {dO, 0}, {out, 0}, {dq, 0}, {stats, 0}, s.H, s.T, s.D, scale); }; auto run_dkv = [&] { - attn_prefill_dkv(q, K, V, dO, stats, dK, dV, s.H, s.T, s.D, scale); + tl::gpu::attn_prefill_dkv({q, 0}, {K, 0}, {V, 0}, {dO, 0}, {stats, 0}, {dK, 0}, {dV, 0}, s.H, s.T, s.D, scale); }; auto time_it = [&](auto&& fn) { fn(); diff --git a/bench/cuda/speed/bench_attn_decode.cpp b/bench/cuda/speed/bench_attn_decode.cpp index 71c6d42..0082b81 100644 --- a/bench/cuda/speed/bench_attn_decode.cpp +++ b/bench/cuda/speed/bench_attn_decode.cpp @@ -14,7 +14,7 @@ #ifndef TENSORLIB_CUDA #define TENSORLIB_CUDA #endif -#include "cuda.h" +#include "gpu.h" // cuda.h plus the shared ops (tl::gpu resolves to cuda here) #include "kv_cache.h" #include @@ -74,12 +74,11 @@ int main() { } // Narrow the f32 K,V into the bf16 cache buffers (kv_fill: [H,ctx,D] -> cache // rows [0,ctx), kv_max=ctx so no padding). H_kv = H here (no GQA). - kv_fill(Kb, Vb, K, V, ctx, ctx, H, D, /*kv_bf16=*/true); + tl::gpu::kv_fill({Kb, 0}, {Vb, 0}, {K, 0}, {V, 0}, ctx, ctx, H, D, /*kv_bf16=*/true); flush(); auto run = [&](bool bf16) { - attn_decode(q, bf16 ? Kb : K, bf16 ? Vb : V, o, H, H, ctx, ctx, D, scale, - bf16); + tl::gpu::attn_decode({q, 0}, {bf16 ? Kb : K, 0}, {bf16 ? Vb : V, 0}, {o, 0}, H, H, ctx, ctx, D, scale, bf16); }; run(true); flush(); @@ -170,11 +169,10 @@ int main() { return (int32_t)st * (1.0f / 2147483648.0f); }; for (int64_t i = 0; i < HQ * D; i++) hq[i] = rnd(); for (int64_t i = 0; i < HKV * ctx * D; i++) { hK[i] = rnd(); hV[i] = rnd(); } - kv_fill(Kb, Vb, K, V, ctx, ctx, HKV, D, true); + tl::gpu::kv_fill({Kb, 0}, {Vb, 0}, {K, 0}, {V, 0}, ctx, ctx, HKV, D, true); flush(); auto run = [&](bool bf16) { - attn_decode(q, bf16 ? Kb : K, bf16 ? Vb : V, o, HQ, HKV, ctx, ctx, D, sc, - bf16); + tl::gpu::attn_decode({q, 0}, {bf16 ? Kb : K, 0}, {bf16 ? Vb : V, 0}, {o, 0}, HQ, HKV, ctx, ctx, D, sc, bf16); }; auto time = [&](bool bf16) { run(bf16); flush(); @@ -252,7 +250,7 @@ int main() { } sync_to_host(kn, true); // host just wrote → re-upload on next device read sync_to_host(vn, true); - if (!cache.append(kn, vn)) { + if (!cache.append({kn, 0}, {vn, 0})) { std::printf(" append failed at pos %lld\n", (long long)pos); return 1; } @@ -261,7 +259,7 @@ int main() { ci++; for (int64_t i = 0; i < HQ * D; i++) hq[i] = rnd(); sync_to_host(q, true); - cache.attn(q, o, HQ, scale); + cache.attn({q, 0}, {o, 0}, HQ, scale); flush(); sync_to_host(o, false); @@ -302,7 +300,7 @@ int main() { std::vector ms; for (int r = 0; r < ROUNDS; r++) { auto t0 = clk::now(); - for (int i = 0; i < R; i++) cache.attn(q, o, HQ, scale); + for (int i = 0; i < R; i++) cache.attn({q, 0}, {o, 0}, HQ, scale); flush(); ms.push_back( std::chrono::duration(clk::now() - t0).count() / R); @@ -358,7 +356,7 @@ int main() { std::printf(" cache init failed\n"); return 1; } - cache.prefill(qp, ks, vs, op, T, HQ, scale); + cache.prefill({qp, 0}, {ks, 0}, {vs, 0}, {op, 0}, T, HQ, scale); flush(); sync_to_host(op, false); @@ -412,8 +410,8 @@ int main() { sync_to_host(kn, true); sync_to_host(vn, true); sync_to_host(q1, true); - cache.append(kn, vn); - cache.attn(q1, o1, HQ, scale); + cache.append({kn, 0}, {vn, 0}); + cache.attn({q1, 0}, {o1, 0}, HQ, scale); flush(); sync_to_host(o1, false); @@ -457,8 +455,7 @@ int main() { for (int r = 0; r < ROUNDS; r++) { auto t0 = clk::now(); for (int i = 0; i < R; i++) - attn_prefill(qp, cache.K.native, cache.V.native, op, HQ, HKV, T, MAXC, D, - scale); + tl::gpu::attn_prefill({qp, 0}, {cache.K.native, 0}, {cache.V.native, 0}, {op, 0}, HQ, HKV, T, MAXC, D, scale); flush(); ms.push_back( std::chrono::duration(clk::now() - t0).count() / R); diff --git a/bench/cuda/speed/bench_bf16_gemv.cpp b/bench/cuda/speed/bench_bf16_gemv.cpp index d6a8ca9..5f21d35 100644 --- a/bench/cuda/speed/bench_bf16_gemv.cpp +++ b/bench/cuda/speed/bench_bf16_gemv.cpp @@ -16,7 +16,7 @@ #ifndef TENSORLIB_CUDA #define TENSORLIB_CUDA #endif -#include "cuda.h" +#include "gpu.h" // cuda.h plus the shared ops (tl::gpu resolves to cuda here) #include #include @@ -88,8 +88,8 @@ int main() { } // correctness: bf16-weight result vs f32-weight result - gemv_f32(a, Bf, yf, N, K); - gemv_bf16(a, Bb, yb, N, K); + tl::gpu::gemv_f32({a, 0}, {Bf, 0}, {yf, 0}, N, K); + tl::gpu::gemv_bf16({a, 0}, {Bb, 0}, {yb, 0}, N, K); flush(); sync_to_host(yf, false); sync_to_host(yb, false); @@ -113,8 +113,8 @@ int main() { } return median(ms); }; - double f32_ms = time_ms([&] { gemv_f32(a, Bf, yf, N, K); }); - double bf16_ms = time_ms([&] { gemv_bf16(a, Bb, yb, N, K); }); + double f32_ms = time_ms([&] { tl::gpu::gemv_f32({a, 0}, {Bf, 0}, {yf, 0}, N, K); }); + double bf16_ms = time_ms([&] { tl::gpu::gemv_bf16({a, 0}, {Bb, 0}, {yb, 0}, N, K); }); double f32_gbs = static_cast(K) * N * 4 / (f32_ms * 1e6); double bf16_gbs = static_cast(K) * N * 2 / (bf16_ms * 1e6); sum_f32 += f32_ms; diff --git a/bench/cuda/speed/bench_cuda_gemm.cpp b/bench/cuda/speed/bench_cuda_gemm.cpp index 64c3f27..5a6bbe0 100644 --- a/bench/cuda/speed/bench_cuda_gemm.cpp +++ b/bench/cuda/speed/bench_cuda_gemm.cpp @@ -19,7 +19,7 @@ #ifndef TENSORLIB_CUDA #define TENSORLIB_CUDA #endif -#include "cuda.h" +#include "gpu.h" // cuda.h plus the shared ops (tl::gpu resolves to cuda here) #include #include @@ -120,14 +120,13 @@ int main(int argc, char** argv) { dB, (int)n, dA, (int)k, &beta, (float*)out, (int)n); }; auto run_own = [&](void* out) { - gemm(At, 0, sh.ta ? m : k, sh.ta, Bt, 0, sh.tb ? k : n, sh.tb, out, 0, m, n, k, - 1.0f, 0.0f); + tl::gpu::gemm({At, 0}, sh.ta ? m : k, sh.ta, {Bt, 0}, sh.tb ? k : n, sh.tb, {out, 0}, m, n, k, 1.0f, 0.0f); }; // ---- correctness: own vs cuBLAS (cuBLAS is the trusted oracle here) ---- run_own(C); // also uploads A,B to device (unless ta/tb) if (sh.ta || sh.tb) // own read At/Bt: upload A,B through a plain NN gemm - gemm(A, 0, k, false, B, 0, n, false, Cref, 0, m, n, k, 1.0f, 0.0f); + tl::gpu::gemm({A, 0}, k, false, {B, 0}, n, false, {Cref, 0}, m, n, k, 1.0f, 0.0f); run_cublas(Cref); // reads the same device A,B cudaDeviceSynchronize(); std::vector ownv(m * n), refv(m * n); diff --git a/bench/cuda/speed/bench_q4_gemv.cpp b/bench/cuda/speed/bench_q4_gemv.cpp index 8c07a54..c3ae143 100644 --- a/bench/cuda/speed/bench_q4_gemv.cpp +++ b/bench/cuda/speed/bench_q4_gemv.cpp @@ -9,7 +9,7 @@ #ifndef TENSORLIB_CUDA #define TENSORLIB_CUDA #endif -#include "cuda.h" +#include "gpu.h" // cuda.h plus the shared ops (tl::gpu resolves to cuda here) #include #include @@ -90,7 +90,7 @@ int main() { } } - gemv_q4(a, qw, sc, y, N, K, G); + tl::gpu::gemv_q4({a, 0}, {qw, 0}, {sc, 0}, {y, 0}, N, K, G); flush(); sync_to_host(y, false); @@ -117,7 +117,7 @@ int main() { } return median(ms); }; - double ms = time_ms([&] { gemv_q4(a, qw, sc, y, N, K, G); }); + double ms = time_ms([&] { tl::gpu::gemv_q4({a, 0}, {qw, 0}, {sc, 0}, {y, 0}, N, K, G); }); double bytes = (double)N * K * 0.5 + (double)N * groups * 4; double gbs = bytes / (ms * 1e6); double bpw = bytes / ((double)N * K); diff --git a/bench/cuda/speed/bench_qwen_ctx.cpp b/bench/cuda/speed/bench_qwen_ctx.cpp index 52817b9..616fed1 100644 --- a/bench/cuda/speed/bench_qwen_ctx.cpp +++ b/bench/cuda/speed/bench_qwen_ctx.cpp @@ -103,7 +103,7 @@ int main(int argc, char** argv) { // (a MemFree sync + realloc) as S steps up across ctx 256..max, and those // reallocs would otherwise land inside timed tokens of the imperative curve. qm::set_cache_pos(M, max_ctx); - M.layers[0].cache.attn(M.scratch.qb.native, M.scratch.ab.native, qm::NH, qm::SCALE); + M.layers[0].cache.attn(M.scratch.qb.device_span(), M.scratch.ab.device_span(), qm::NH, qm::SCALE); cu::flush(); // ---- 1. per-position decode curve. @@ -136,9 +136,9 @@ int main(int argc, char** argv) { // attn() reads ctx = pos; attn_dpos() reads ctx = *d_pos + 1. Same ctx. cu::upload_u32(d_pos, (unsigned)(ctx - 1)); auto& L0 = M.layers[0]; - L0.cache.attn(M.scratch.qb.native, refb, qm::NH, qm::SCALE); + L0.cache.attn(M.scratch.qb.device_span(), {refb, 0}, qm::NH, qm::SCALE); cu::sync_to_host(refb, false); - L0.cache.attn_dpos(M.scratch.qb.native, gotb, qm::NH, d_pos, qm::SCALE); + L0.cache.attn_dpos(M.scratch.qb.device_span(), {gotb, 0}, qm::NH, {d_pos, 0}, qm::SCALE); cu::sync_to_host(gotb, false); bool bit_eq = std::memcmp(ref_host, got_host, qm::NH * qm::HD * 4) == 0; auto time_24x = [&](auto&& attn1) { // REPS x 24-layer attn -> ms per rep @@ -149,10 +149,10 @@ int main(int argc, char** argv) { return (now_ms() - t) / REPS; }; auto host1 = [&](qm::Layer& L) { - L.cache.attn(M.scratch.qb.native, M.scratch.ab.native, qm::NH, qm::SCALE); + L.cache.attn(M.scratch.qb.device_span(), M.scratch.ab.device_span(), qm::NH, qm::SCALE); }; auto dpos1 = [&](qm::Layer& L) { - L.cache.attn_dpos(M.scratch.qb.native, M.scratch.ab.native, qm::NH, d_pos, qm::SCALE); + L.cache.attn_dpos(M.scratch.qb.device_span(), M.scratch.ab.device_span(), qm::NH, {d_pos, 0}, qm::SCALE); }; double host_ms = 1e30, dpos_ms = 1e30; for (int r = 0; r < 3; r++) { // min of 3, host/dpos interleaved (WSL2 noise) diff --git a/bench/cuda/speed/bench_qwen_decode.cpp b/bench/cuda/speed/bench_qwen_decode.cpp index 3746d76..f3ec0ce 100644 --- a/bench/cuda/speed/bench_qwen_decode.cpp +++ b/bench/cuda/speed/bench_qwen_decode.cpp @@ -138,12 +138,12 @@ int main(int argc, char** argv) { // row P identically) — the dpos KERNELS must match the host-pos kernels. qm::stage_embed(M, next); qm::set_cache_pos(M, P); - qm::run_layers_(M, M.scratch.embed.native, P); + qm::run_layers_(M, qm::sp(M.scratch.embed), P); cu::sync_to_host(M.scratch.logits.native, false); std::vector ref(M.scratch.logits.ptr, M.scratch.logits.ptr + qm::VOCAB); qm::set_cache_pos(M, P); cu::upload_u32(d_pos, (unsigned)P); - qm::run_layers_(M, M.scratch.embed.native, P, d_pos); + qm::run_layers_(M, qm::sp(M.scratch.embed), P, qm::sp(cap.d_pos)); cu::sync_to_host(M.scratch.logits.native, false); const float* got = M.scratch.logits.ptr; int mism = 0; @@ -192,7 +192,7 @@ int main(int argc, char** argv) { double raw_min = bench([&](int64_t) { cu::graph_launch(cap.exec); }); double ra_min = bench([&](int64_t) { cu::graph_launch(cap.exec); - cu::argmax(M.scratch.logits.native, qm::VOCAB, &di); + tl::gpu::argmax({M.scratch.logits.native, 0}, qm::VOCAB, &di); }); double cap_min = bench([&](int64_t) { idx = cap.step(M, idx); }); std::printf("=== correct captured decode (device-pos) ===\n"); diff --git a/bench/cuda/speed/bench_qwen_gemv.cpp b/bench/cuda/speed/bench_qwen_gemv.cpp index 14e7cdd..0fe1da4 100644 --- a/bench/cuda/speed/bench_qwen_gemv.cpp +++ b/bench/cuda/speed/bench_qwen_gemv.cpp @@ -20,7 +20,7 @@ #ifndef TENSORLIB_CUDA #define TENSORLIB_CUDA #endif -#include "cuda.h" +#include "gpu.h" // cuda.h plus the shared ops (tl::gpu resolves to cuda here) #include #include @@ -95,13 +95,13 @@ struct Op { release(a, 0, nullptr); release(B, 0, nullptr); release(y, 0, nullptr); release(qw, 0, nullptr); release(sc, 0, nullptr); } - void run() const { gemv_bf16(a, B, y, N, K); } + void run() const { tl::gpu::gemv_bf16({a, 0}, {B, 0}, {y, 0}, N, K); } // Warp-per-row [N,K] variant (lever A). Speed is layout-agnostic (same K*N*2 // random bytes, same access footprint), so it reuses the same B buffer — only // the in-kernel interpretation differs. Compared head-to-head with split-K. - void run_row() const { gemv_bf16_row(a, B, y, N, K); } + void run_row() const { tl::gpu::gemv_bf16_row({a, 0}, {B, 0}, {y, 0}, N, K); } // q4 warp-per-row [N,K] (bandwidth lever): ~0.625 B/wt vs bf16's 2. - void run_q4() const { gemv_q4(a, qw, sc, y, N, K, G); } + void run_q4() const { tl::gpu::gemv_q4({a, 0}, {qw, 0}, {sc, 0}, {y, 0}, N, K, G); } }; static const int R = 50, ROUNDS = 7; diff --git a/bench/cuda/speed/bench_qwen_prefill.cpp b/bench/cuda/speed/bench_qwen_prefill.cpp index 93d00a3..48d127f 100644 --- a/bench/cuda/speed/bench_qwen_prefill.cpp +++ b/bench/cuda/speed/bench_qwen_prefill.cpp @@ -247,8 +247,8 @@ int main(int argc, char** argv) { bool gemm_ok = true; for (const Shape& s : shapes) { // Warm both paths (module load, first-touch upload) outside the timing. - cu::gemv_bf16_row(a, s.w->native(), y1, s.N, s.K); - cu::gemm_bf16_nt(a, s.w->native(), y, MB, s.N, s.K); + tl::gpu::gemv_bf16_row({a, 0}, {s.w->native(), 0}, {y1, 0}, s.N, s.K); + tl::gpu::gemm_bf16_nt({a, 0}, {s.w->native(), 0}, {y, 0}, MB, s.N, s.K); cu::flush(); // Correctness: GEMM row 0 must reproduce the GEMV of A's row 0. cu::sync_to_host(y1, false); @@ -261,10 +261,10 @@ int main(int argc, char** argv) { gemm_ok &= ok; double gemv_us = 1000.0 * min_ms(3, 20, [&] { - cu::gemv_bf16_row(a, s.w->native(), y1, s.N, s.K); + tl::gpu::gemv_bf16_row({a, 0}, {s.w->native(), 0}, {y1, 0}, s.N, s.K); }); double gemm_us = 1000.0 * min_ms(3, 5, [&] { - cu::gemm_bf16_nt(a, s.w->native(), y, MB, s.N, s.K); + tl::gpu::gemm_bf16_nt({a, 0}, {s.w->native(), 0}, {y, 0}, MB, s.N, s.K); }) / MB; // per prompt token sum_gemv += gemv_us; sum_gemm += gemm_us; @@ -299,14 +299,14 @@ int main(int argc, char** argv) { double per_tok = min_ms(3, 1, [&] { c.pos = 0; for (int64_t i = 0; i < T; i++) { - c.append(ks, vs); // same row values each step; only the cost matters - c.attn(qs, os, qm::NH, qm::SCALE); + c.append({ks, 0}, {vs, 0}); // same row values each step; only the cost matters + c.attn({qs, 0}, {os, 0}, qm::NH, qm::SCALE); } }); // Batched: bulk kv_fill + one causal tiled attn_prefill. double batched = min_ms(3, 1, [&] { c.pos = 0; - c.prefill(qs, ks, vs, os, T, qm::NH, qm::SCALE); + c.prefill({qs, 0}, {ks, 0}, {vs, 0}, {os, 0}, T, qm::NH, qm::SCALE); }); std::printf(" %6lld %14.3f %14.3f %7.1fx\n", (long long)T, per_tok, batched, per_tok / batched); diff --git a/bench/cuda/speed/bench_xent.cpp b/bench/cuda/speed/bench_xent.cpp index c61c96e..cc82633 100644 --- a/bench/cuda/speed/bench_xent.cpp +++ b/bench/cuda/speed/bench_xent.cpp @@ -12,7 +12,7 @@ #ifndef TENSORLIB_CUDA #define TENSORLIB_CUDA #endif -#include "cuda.h" +#include "gpu.h" // cuda.h plus the shared ops (tl::gpu resolves to cuda here) #include "array.h" // the second table: the same work through the graph @@ -24,6 +24,7 @@ #include using namespace tl::cuda; +namespace gpu = tl::gpu; // the shared ops; tl::gpu resolves to cuda here using clk = std::chrono::steady_clock; static double median(std::vector v) { @@ -94,17 +95,16 @@ int main() { }; const double lse_ms = time_it( - [&] { row_logsumexp(x, 0, lse, 0, s.rows, s.cols, 1.0f, 0.0f); }); - // Qualified because kop lives in tl::metal: an unqualified call would pull - // tl::metal::row_op into the overload set by ADL and be ambiguous. + [&] { gpu::row_logsumexp({x, 0}, {lse, 0}, s.rows, s.cols, 1.0f, 0.0f); }); const double sm_ms = time_it([&] { - tl::cuda::row_op(tl::metal::kop::softmax, x, 0, probs, 0, s.rows, s.cols, - 1.0f, 0.0f); + gpu::row_op(gpu::kop::softmax, {x, 0}, {probs, 0}, s.rows, s.cols, 1.0f, + 0.0f); + }); + const double gat_ms = time_it([&] { + gpu::gather_from_axis({x, 0}, {tgt, 0}, {picked, 0}, s.rows, s.cols); }); - const double gat_ms = time_it( - [&] { gather_from_axis(x, 0, tgt, 0, picked, 0, s.rows, s.cols); }); const double bwd_ms = time_it([&] { - xent_bwd(x, 0, lse, 0, tgt, 0, g, 0, dx, 0, s.rows, s.cols); + gpu::xent_bwd({x, 0}, {lse, 0}, {tgt, 0}, {g, 0}, {dx, 0}, s.rows, s.cols); }); char name[32]; diff --git a/bench/metal/speed/bench_attn_bwd.cpp b/bench/metal/speed/bench_attn_bwd.cpp index c0f8e01..2c3231b 100644 --- a/bench/metal/speed/bench_attn_bwd.cpp +++ b/bench/metal/speed/bench_attn_bwd.cpp @@ -5,7 +5,7 @@ // Direct metal:: API (timing via metal::flush + cpu_barrier and steady_clock), // so a change to one kernel can be measured without a consumer's build. -#include "metal.h" +#include "gpu.h" // metal.h plus the shared ops (tl::gpu resolves to metal here) #include #include @@ -67,13 +67,13 @@ int main() { fill_random(hG, n, 4); auto run_fwd = [&] { - attn_prefill(q, K, V, out, s.H, s.H, s.T, s.T, s.D, scale); + tl::gpu::attn_prefill({q, 0}, {K, 0}, {V, 0}, {out, 0}, s.H, s.H, s.T, s.T, s.D, scale); }; auto run_dq = [&] { - attn_prefill_dq(q, K, V, dO, out, dq, stats, s.H, s.T, s.D, scale); + tl::gpu::attn_prefill_dq({q, 0}, {K, 0}, {V, 0}, {dO, 0}, {out, 0}, {dq, 0}, {stats, 0}, s.H, s.T, s.D, scale); }; auto run_dkv = [&] { - attn_prefill_dkv(q, K, V, dO, stats, dK, dV, s.H, s.T, s.D, scale); + tl::gpu::attn_prefill_dkv({q, 0}, {K, 0}, {V, 0}, {dO, 0}, {stats, 0}, {dK, 0}, {dV, 0}, s.H, s.T, s.D, scale); }; auto time_it = [&](auto&& fn) { fn(); diff --git a/bench/models/qwen2.h b/bench/models/qwen2.h index 2ebe0b9..05ff69b 100644 --- a/bench/models/qwen2.h +++ b/bench/models/qwen2.h @@ -154,22 +154,13 @@ inline array load_w_T_row(const gg::model& m, const std::string& n, int64_t in, // (unified on Metal, the mirror on CUDA). inline storage scratch_f32(int64_t n) { return storage::make(n); } -// Slice a device f32 buffer by element offset. The result is a mid-buffer -// pointer: reads through it are ordered after whatever wrote the base and -// need no host sync, since the base is already device-live. -inline void* off_f32(void* p, int64_t nfloats) { - return static_cast(p) + nfloats * 4; -} - -// The same slice where a pointer cannot name one: `n` floats from element -// `from` of `src` into the start of `dst`. An affine unary carries a byte -// offset on every backend, a pointer only where gpu::caps::flat_addressing -// says a device pointer is an address. The model fuses its projections -// either way — what the cap decides is whether reading the fused output back -// costs a copy. -inline void copy_out(void* src, int64_t from, void* dst, int64_t n) { - gpu::unary(gpu::kop::affine, src, from * 4, dst, 0, n, 1.0f, 0.0f); -} +// Device views. A storage or an array names its whole buffer; at_f32 slices a +// view by element. The offset travels beside the handle (gpu::span), so a slice +// is a view on every backend — no copy, and reads through it are ordered after +// whatever wrote the base. +inline gpu::span sp(const storage& s) { return s.device_span(); } +inline gpu::span sp(const array& a) { return a.device_span(); } +inline gpu::span at_f32(gpu::span s, int64_t nfloats) { return s.at(nfloats * 4); } // Make a device buffer's bytes readable on the host: drain the queue (on // unified memory that is all it takes) and pull the mirror back where there is @@ -235,8 +226,6 @@ struct Scratch { storage h2b; // post-attn norm out [NE] storage qb; // [NH*HD] query fixture (bench_qwen_ctx's isolated- // attention timing; decode reads q as a slice of qkvb) - storage kb, vb; // [NKV*HD] k and v copied out of qkvb, where a - // pointer cannot name that slice (see copy_out) storage qkvb; // fused QKV out [(NH+2*NKV)*HD] = [1152] storage ab; // attn out [NH*HD] storage mb; // swiglu out [FF] @@ -256,8 +245,6 @@ struct Scratch { hb = scratch_f32(NE); h2b = scratch_f32(NE); qb = scratch_f32(NH * HD); - kb = scratch_f32(NKV * HD); - vb = scratch_f32(NKV * HD); qkvb = scratch_f32((NH + 2 * NKV) * HD); ab = scratch_f32(NH * HD); mb = scratch_f32(FF); @@ -280,11 +267,9 @@ struct PrefillScratch { storage h; // input-norm out [cap, NE] storage h2; // post-attn norm out [cap, NE] storage qkv; // fused QKV out [cap, (NH+2*NKV)*HD] - // q|k|v head-major. Fused: ONE buffer [NH+2*NKV, cap, HD], one split_heads - // pass over every head, k/v mid-buffer pointers into it. Unfused: one - // buffer and one pass each, with that head block's own bias (bq/bk/bv - // rather than the concatenated bqkv). - storage qkvh, qh, kh, vh; + // q|k|v head-major in ONE buffer [NH+2*NKV, cap, HD]: one split_heads pass + // over every head, and k / v are views into it. + storage qkvh; storage ah; // attn out head-major [NH, cap, HD] storage at; // attn out token-major [cap, NH*HD] storage gu; // fused gate|up [cap, 2*FF] @@ -302,13 +287,7 @@ struct PrefillScratch { h = scratch_f32(cap * NE); h2 = scratch_f32(cap * NE); qkv = scratch_f32(cap * (NH + 2 * NKV) * HD); - if (gpu::caps::flat_addressing) { - qkvh = scratch_f32((NH + 2 * NKV) * cap * HD); - } else { - qh = scratch_f32(NH * cap * HD); - kh = scratch_f32(NKV * cap * HD); - vh = scratch_f32(NKV * cap * HD); - } + qkvh = scratch_f32((NH + 2 * NKV) * cap * HD); ah = scratch_f32(NH * cap * HD); at = scratch_f32(cap * NH * HD); gu = scratch_f32(cap * 2 * FF); @@ -333,16 +312,17 @@ struct Model { }; // Decode GEMV picking the weight-dtype kernel: y(n) = a(1,k) @ W[k,n]. -inline bool gemv_w(const array& W, void* a, void* y, int64_t n, int64_t k) { - return W.dt() == tl::dtype::bf16 ? gpu::gemv_bf16(a, W.native(), y, n, k) - : gpu::gemv_f32(a, W.native(), y, n, k); +inline bool gemv_w(const array& W, gpu::span a, gpu::span y, int64_t n, + int64_t k) { + return W.dt() == tl::dtype::bf16 ? gpu::gemv_bf16(a, sp(W), y, n, k) + : gpu::gemv_f32(a, sp(W), y, n, k); } // Greedy token from the logits the most recent forward left in scratch (the one // terminal sync of a step lives inside gpu::argmax). inline int64_t argmax_logits(Model& M) { int64_t idx = 0; - gpu::argmax(M.scratch.logits.native, VOCAB, &idx); + gpu::argmax(sp(M.scratch.logits), VOCAB, &idx); return idx; } @@ -385,45 +365,33 @@ inline void prefill_chunk_(Model& M, int64_t T) { // same source rather than taking it as a parameter that could disagree. const int64_t pos0 = M.layers[0].cache.pos; constexpr int64_t QKVN = (NH + 2 * NKV) * HD; - void* x = P.emb.native; - gpu::rmsnorm(x, M.layers[0].an.native(), P.h.native, NE, EPS, T); + gpu::span x = sp(P.emb); + gpu::rmsnorm(x, sp(M.layers[0].an), sp(P.h), NE, EPS, T); for (int64_t l = 0; l < NL; l++) { Layer& L = M.layers[l]; - void* ro = P.res[l & 1].native; - gpu::gemm_bf16_nt(P.h.native, L.wqkv.native(), P.qkv.native, T, QKVN, NE); + const gpu::span ro = sp(P.res[l & 1]); + gpu::gemm_bf16_nt(sp(P.h), sp(L.wqkv), sp(P.qkv), T, QKVN, NE); // split_heads turns the [T, q|k|v] output into head-major [H, T, D] and // adds the bias — rope's own fused-bias form only indexes correctly at - // T == 1, so the bias rides along here instead. Its `off` names the - // column block, so the unfused form is the same pass three times, each - // with that block's own bias. - void *qh, *kh, *vh; - if (gpu::caps::flat_addressing) { - gpu::split_heads(P.qkv.native, L.bqkv.native(), P.qkvh.native, T, QKVN, 0, - NH + 2 * NKV, HD); - qh = P.qkvh.native; - kh = off_f32(qh, NH * T * HD); - vh = off_f32(qh, (NH + NKV) * T * HD); - } else { - qh = P.qh.native; - kh = P.kh.native; - vh = P.vh.native; - gpu::split_heads(P.qkv.native, L.bq.native(), qh, T, QKVN, 0, NH, HD); - gpu::split_heads(P.qkv.native, L.bk.native(), kh, T, QKVN, NH * HD, NKV, HD); - gpu::split_heads(P.qkv.native, L.bv.native(), vh, T, QKVN, - (NH + NKV) * HD, NKV, HD); - } + // T == 1, so the bias rides along here instead. One pass over every head; + // k and v are views into its output. + gpu::split_heads(sp(P.qkv), sp(L.bqkv), sp(P.qkvh), T, QKVN, 0, NH + 2 * NKV, + HD); + const gpu::span qh = sp(P.qkvh); + const gpu::span kh = at_f32(qh, NH * T * HD); + const gpu::span vh = at_f32(qh, (NH + NKV) * T * HD); // [H,T,D] flattened: row r = h*T + t, so rope's `pos + r % T` is pos0 + t. gpu::rope(qh, qh, NH * T, T, HD, pos0, ROPE_BASE); gpu::rope(kh, kh, NKV * T, T, HD, pos0, ROPE_BASE); - L.cache.prefill(qh, kh, vh, P.ah.native, T, NH, SCALE); - gpu::merge_heads(P.ah.native, P.at.native, T, NH, HD); - gpu::gemm_bf16_nt(P.at.native, L.wo_row.native(), ro, T, NE, NH * HD); - gpu::rmsnorm_res(ro, x, L.fn.native(), ro, P.h2.native, NE, EPS, T); - gpu::gemm_bf16_nt(P.h2.native, L.wgu.native(), P.gu.native, T, 2 * FF, NE); - gpu::swiglu(P.gu.native, P.mb.native, FF, T); - gpu::gemm_bf16_nt(P.mb.native, L.wd_row.native(), P.md.native, T, NE, FF); - void* nextw = (l + 1 < NL) ? M.layers[l + 1].an.native() : M.onorm.native(); - gpu::rmsnorm_res(ro, P.md.native, nextw, ro, P.h.native, NE, EPS, T); + L.cache.prefill(qh, kh, vh, sp(P.ah), T, NH, SCALE); + gpu::merge_heads(sp(P.ah), sp(P.at), T, NH, HD); + gpu::gemm_bf16_nt(sp(P.at), sp(L.wo_row), ro, T, NE, NH * HD); + gpu::rmsnorm_res(ro, x, sp(L.fn), ro, sp(P.h2), NE, EPS, T); + gpu::gemm_bf16_nt(sp(P.h2), sp(L.wgu), sp(P.gu), T, 2 * FF, NE); + gpu::swiglu(sp(P.gu), sp(P.mb), FF, T); + gpu::gemm_bf16_nt(sp(P.mb), sp(L.wd_row), sp(P.md), T, NE, FF); + const gpu::span nextw = sp((l + 1 < NL) ? M.layers[l + 1].an : M.onorm); + gpu::rmsnorm_res(ro, sp(P.md), nextw, ro, sp(P.h), NE, EPS, T); x = ro; } // Leaves P.h holding the final RMSNorm for every row of the chunk. Logits are @@ -474,10 +442,9 @@ inline int64_t prefill_batched(Model& M, const std::vector& ids, prefill_chunk_(M, last); } // Only the final row of the final chunk needs logits — one GEMV for the whole - // prompt rather than one per chunk. The row is copied into the decode's own - // norm buffer rather than pointed at in place — see copy_out. - copy_out(M.pscratch.h.native, (last - 1) * NE, M.scratch.hb.native, NE); - gemv_w(M.outwT, M.scratch.hb.native, M.scratch.logits.native, VOCAB, NE); + // prompt rather than one per chunk, read in place as a view of that row. + gemv_w(M.outwT, at_f32(sp(M.pscratch.h), (last - 1) * NE), sp(M.scratch.logits), + VOCAB, NE); return argmax_logits(M); } @@ -626,9 +593,9 @@ inline array forward(Model& M, int64_t id, int64_t pos, // see the writes. Removes 3 CtxSynchronize/layer. q.realize(); k.realize(); v.realize(); if (prof) { prof->qkv_eval += StepProf::now_ms() - t; t = StepProf::now_ms(); } - L.cache.append(k.native(), v.native()); + L.cache.append(sp(k), sp(v)); array a_out = array::empty({NH, HD}); - L.cache.attn(q.native(), a_out.native(), NH, SCALE); + L.cache.attn(sp(q), sp(a_out), NH, SCALE); if (prof) { prof->cache += StepProf::now_ms() - t; t = StepProf::now_ms(); } array x1 = x + a_out.reshape({1, NE}).dot(L.wo); array h2 = array::rmsnorm(x1, L.fn, EPS); @@ -679,7 +646,7 @@ inline int64_t step_greedy(Model& M, int64_t id, int64_t pos, logits.realize(); if (prof) { prof->logits_eval += StepProf::now_ms() - t; t = StepProf::now_ms(); } int64_t idx = 0; - if (!gpu::argmax(logits.native(), VOCAB, &idx)) { + if (!gpu::argmax(sp(logits), VOCAB, &idx)) { const float* p = logits.raw(); idx = 0; for (int64_t i = 1; i < VOCAB; i++) @@ -697,12 +664,12 @@ inline int64_t argmax(const std::vector& v) { } // q4 decode GEMV: y(N) = a(1,K) @ dequant(Wq). Wq is a q4 array (logical [K,N], -// storage [packed [N][K/2] | scales [N][K/32]]); the scales pointer is mid-buffer -// (rides along the base upload — same split as the array-path gpu_gemv_q4). -inline bool gemv_q4_w(const array& Wq, void* a, void* y) { +// storage [packed [N][K/2] | scales [N][K/32]]); the scales are a view into the +// same buffer (they ride along the base upload — the array path's split too). +inline bool gemv_q4_w(const array& Wq, gpu::span a, gpu::span y) { const int64_t K = Wq.shape()[0], N = Wq.shape()[1]; - void* scales = static_cast(Wq.native()) + N * K / 2; - return gpu::gemv_q4(a, Wq.native(), scales, y, N, K, tl::kQ4Group); + const gpu::span qw = sp(Wq); + return gpu::gemv_q4(a, qw, qw.at(N * K / 2), y, N, K, tl::kQ4Group); } // The 24 decoder layers + final RMSNorm + lm_head gemv as direct gpu:: calls @@ -720,10 +687,10 @@ inline bool gemv_q4_w(const array& Wq, void* a, void* y) { // device instead of the host `pos`, and a tl_incr_u32 at the tail advances it — // so the whole forward is CUDA-graph-capturable and one instantiated graph // replays correctly as pos grows (A-min). All caches share the one counter (every -// layer is at the same sequence position). d_pos==nullptr = the normal host path. -inline void run_layers_(Model& M, void* x0, int64_t pos, void* d_pos = nullptr, - void* logits_out = nullptr) { - const bool cap = d_pos != nullptr; +// layer is at the same sequence position). A null d_pos = the normal host path. +inline void run_layers_(Model& M, gpu::span x0, int64_t pos, + gpu::span d_pos = {}, gpu::span logits_out = {}) { + const bool cap = static_cast(d_pos); Scratch& S = M.scratch; // Fused seams (kills 4 elementwise launches/layer): q/k bias folds into rope // (gpu::rope's bias arg); the two residual adds fold into the following RMSNorms @@ -736,64 +703,56 @@ inline void run_layers_(Model& M, void* x0, int64_t pos, void* d_pos = nullptr, // layout is identical either way (contiguous [n]), so the fused-output slices // below are unchanged. const bool row = M.row; - auto gv = [&](const array& W, void* a, void* y, int64_t n, int64_t k) { - if (row) gpu::gemv_bf16_row(a, W.native(), y, n, k); + auto gv = [&](const array& W, gpu::span a, gpu::span y, int64_t n, int64_t k) { + if (row) gpu::gemv_bf16_row(a, sp(W), y, n, k); else gemv_w(W, a, y, n, k); }; - void* x = x0; - gpu::rmsnorm(x, M.layers[0].an.native(), S.hb.native, NE, EPS); + gpu::span x = x0; + gpu::rmsnorm(x, sp(M.layers[0].an), sp(S.hb), NE, EPS); for (int64_t l = 0; l < NL; l++) { Layer& L = M.layers[l]; - void* ro = S.res[l & 1].native; // res_out (x1 then x2), ping-pong + const gpu::span ro = sp(S.res[l & 1]); // res_out (x1 then x2), ping-pong // Fused QKV: one GEMV -> [q(NH*HD) | k(NKV*HD) | v(NKV*HD)] in S.qkvb, then - // slice: rope q & k in place, bias-add v. Same per-column split-K as the - // separate wq/wk/wv GEMVs (bx=1, chunk=32), so bit-identical per column. - // q is the slice at offset 0, which every backend can name; k and v are - // copied out where a pointer cannot (two 128-float copies a layer). - gv(L.wqkv, S.hb.native, S.qkvb.native, (NH + 2 * NKV) * HD, NE); - void* qp = S.qkvb.native; - void* kp = S.kb.native; - void* vp = S.vb.native; - if (gpu::caps::flat_addressing) { - kp = off_f32(qp, NH * HD); - vp = off_f32(qp, (NH + NKV) * HD); - } else { - copy_out(qp, NH * HD, kp, NKV * HD); - copy_out(qp, (NH + NKV) * HD, vp, NKV * HD); - } + // q, k and v are views of it: rope q & k in place, bias-add v. Same + // per-column split-K as the separate wq/wk/wv GEMVs (bx=1, chunk=32), so + // bit-identical per column. + gv(L.wqkv, sp(S.hb), sp(S.qkvb), (NH + 2 * NKV) * HD, NE); + const gpu::span qp = sp(S.qkvb); + const gpu::span kp = at_f32(qp, NH * HD); + const gpu::span vp = at_f32(qp, (NH + NKV) * HD); if (cap) { - gpu::rope_dpos(qp, qp, NH, 1, HD, d_pos, ROPE_BASE, L.bq.native()); - gpu::rope_dpos(kp, kp, NKV, 1, HD, d_pos, ROPE_BASE, L.bk.native()); + gpu::rope_dpos(qp, qp, NH, 1, HD, d_pos, ROPE_BASE, sp(L.bq)); + gpu::rope_dpos(kp, kp, NKV, 1, HD, d_pos, ROPE_BASE, sp(L.bk)); } else { - gpu::rope(qp, qp, NH, 1, HD, pos, ROPE_BASE, L.bq.native()); - gpu::rope(kp, kp, NKV, 1, HD, pos, ROPE_BASE, L.bk.native()); + gpu::rope(qp, qp, NH, 1, HD, pos, ROPE_BASE, sp(L.bq)); + gpu::rope(kp, kp, NKV, 1, HD, pos, ROPE_BASE, sp(L.bk)); } - gpu::binary(gpu::kop::add, vp, 0, L.bv.native(), 0, vp, 0, NKV * HD, 1, 0); + gpu::binary(gpu::kop::add, vp, sp(L.bv), vp, NKV * HD, 1, 0); if (cap) { L.cache.append_dpos(kp, vp, d_pos); - L.cache.attn_dpos(qp, S.ab.native, NH, d_pos, SCALE); + L.cache.attn_dpos(qp, sp(S.ab), NH, d_pos, SCALE); } else { L.cache.append(kp, vp); - L.cache.attn(qp, S.ab.native, NH, SCALE); + L.cache.attn(qp, sp(S.ab), NH, SCALE); } - gv(row ? L.wo_row : L.wo, S.ab.native, ro, NE, NH * HD); // ro = attn @ wo + gv(row ? L.wo_row : L.wo, sp(S.ab), ro, NE, NH * HD); // ro = attn @ wo // x1 = x + (attn@wo); h2 = rmsnorm(x1, fn) — fused. - gpu::rmsnorm_res(ro, x, L.fn.native(), ro, S.h2b.native, NE, EPS); + gpu::rmsnorm_res(ro, x, sp(L.fn), ro, sp(S.h2b), NE, EPS); // Fused gate|up: one GEMV -> [gate(FF) | up(FF)] in S.gub; swiglu reads both. - if (M.q4_mlp) gemv_q4_w(L.wgu_q4, S.h2b.native, S.gub.native); // [NE, 2*FF] - else gv(L.wgu, S.h2b.native, S.gub.native, 2 * FF, NE); - gpu::swiglu(S.gub.native, S.mb.native, FF); - if (M.q4_mlp) gemv_q4_w(L.wd_q4, S.mb.native, S.mdb.native); // [FF, NE] - else gv(row ? L.wd_row : L.wd, S.mb.native, S.mdb.native, NE, FF); + if (M.q4_mlp) gemv_q4_w(L.wgu_q4, sp(S.h2b), sp(S.gub)); // [NE, 2*FF] + else gv(L.wgu, sp(S.h2b), sp(S.gub), 2 * FF, NE); + gpu::swiglu(sp(S.gub), sp(S.mb), FF); + if (M.q4_mlp) gemv_q4_w(L.wd_q4, sp(S.mb), sp(S.mdb)); // [FF, NE] + else gv(row ? L.wd_row : L.wd, sp(S.mb), sp(S.mdb), NE, FF); // x2 = x1 + mlp; next input norm = rmsnorm(x2, next an | final onorm) — fused. - void* nextw = (l + 1 < NL) ? M.layers[l + 1].an.native() : M.onorm.native(); - gpu::rmsnorm_res(ro, S.mdb.native, nextw, ro, S.hb.native, NE, EPS); + const gpu::span nextw = sp((l + 1 < NL) ? M.layers[l + 1].an : M.onorm); + gpu::rmsnorm_res(ro, sp(S.mdb), nextw, ro, sp(S.hb), NE, EPS); x = ro; } // hb now holds the final RMSNorm output (folded into the last layer's seam). - void* logits = logits_out ? logits_out : S.logits.native; - if (M.q4_lmhead) gemv_q4_w(M.outwT_q4, S.hb.native, logits); // [NE, VOCAB] - else gemv_w(M.outwT, S.hb.native, logits, VOCAB, NE); + const gpu::span logits = logits_out ? logits_out : sp(S.logits); + if (M.q4_lmhead) gemv_q4_w(M.outwT_q4, sp(S.hb), logits); // [NE, VOCAB] + else gemv_w(M.outwT, sp(S.hb), logits, VOCAB, NE); // Tail of the captured region: advance the shared device pos so the next // graph replay reads pos+1 (lm_head is pos-independent, so order vs it is free). if (cap) gpu::incr_u32(d_pos); @@ -807,7 +766,7 @@ inline void run_layers_(Model& M, void* x0, int64_t pos, void* d_pos = nullptr, inline int64_t step_imperative(Model& M, int64_t id, int64_t pos) { array e = embed_row(M, id); e.realize(); // embed on device (native valid; uploaded on first read) - run_layers_(M, e.native(), pos); + run_layers_(M, sp(e), pos); return argmax_logits(M); } @@ -817,7 +776,7 @@ inline int64_t step_imperative(Model& M, int64_t id, int64_t pos) { inline const float* imperative_logits(Model& M, int64_t id, int64_t pos) { array e = embed_row(M, id); e.realize(); - run_layers_(M, e.native(), pos); + run_layers_(M, sp(e), pos); return host_read(M.scratch.logits); } @@ -872,16 +831,16 @@ struct captured_decoder { if (!gpu::graph_available()) return; d_pos = storage::make(1); if (!d_pos.native) return; - void* dp = d_pos.native; + const gpu::span dp = sp(d_pos); stage_embed(M, first_id); - gpu::upload_u32(dp, (unsigned)pos); + gpu::upload_u32(d_pos.native, (unsigned)pos); // Warm run: real launches, but its logits go to the throwaway sink so a // prompt's logits (already in scratch when begin() batched) survive. - run_layers_(M, M.scratch.embed.native, pos, dp, M.scratch.logits_sink.native); + run_layers_(M, sp(M.scratch.embed), pos, dp, sp(M.scratch.logits_sink)); gpu::flush(); - gpu::upload_u32(dp, (unsigned)pos); // reset after warm + gpu::upload_u32(d_pos.native, (unsigned)pos); // reset after warm if (!gpu::capture_begin()) return; - run_layers_(M, M.scratch.embed.native, pos, dp); // recorded, not executed + run_layers_(M, sp(M.scratch.embed), pos, dp); // recorded, not executed exec = gpu::capture_end(); } @@ -937,7 +896,7 @@ struct captured_decoder { if (cur_pos >= max_ctx) return false; stage_embed(M, id); // gather id's row -> S.embed (host + blocking H2D) if (exec) gpu::graph_launch(exec); // replay: append@d_pos, attn, logits, incr - else run_layers_(M, M.scratch.embed.native, cur_pos); // host-pos imperative + else run_layers_(M, sp(M.scratch.embed), cur_pos); // host-pos imperative cur_pos++; return true; } diff --git a/docs/backends.md b/docs/backends.md new file mode 100644 index 0000000..269ce4c --- /dev/null +++ b/docs/backends.md @@ -0,0 +1,247 @@ +# GPU backends + +How the GPU layer is put together, how to add an op to it, and how to add a +backend. The code is `include/gpu.h`, `gpu_abi.h`, `gpu_ops.h`, `gpu_null.h` +and one header per backend (`metal.h`, `cuda.h`, `webgpu.h`, `gpu_host.h`). + +## Layers + +``` +array.h / storage.h / kv_cache.h / models name no backend; call tl::gpu +────────────────────────────────────────────── +gpu_ops.h every op, written once shared +gpu_abi.h span, access, grid, kernel ABI, + params structs, launch policy +────────────────────────────────────────────── +device core lifecycle, memory, dispatch, one per backend: + own, traits, caps metal / cuda / webgpu / host / null +kernels .metal / .cu / .wgsl / C++ loops +``` + +`tl::gpu` is a namespace holding the shared layer plus a using-directive for +the selected backend. A name the shared layer declares (an op) is found +first; anything else (`alloc`, `flush`, `caps`) falls through to the backend. +`gpu.h` selects exactly one backend header by its gate and includes only that +one; a build that fits none gets `gpu_null.h`. + +`array.h` and `storage.h` never name a backend, and neither does an embedder: +everything goes through `tl::gpu` or the array API. + +## Views: `gpu::span` + +```cpp +struct span { void* buf; int64_t off; }; // off in bytes +``` + +Every op takes its buffers as spans. `buf` is whatever the backend's `alloc` +returned and means nothing outside that backend: an `MTLBuffer` handle on +Metal, a device address on CUDA, a key into a mirror table on WebGPU. The +offset travels beside the handle rather than inside it, so a view is +expressible on every backend, and arithmetic on a handle is not expressible at +all. `s.at(bytes)` slices a view; `array::device_span()` and +`storage::device_span()` produce one. + +Each view a kernel takes is tagged with how the kernel touches it: + +| access | meaning | what a mirrored backend does | +|---------|---------|------------------------------| +| `in` | read | upload first if the host holds the live copy | +| `out` | written | the device copy becomes the live one; upload first only if the host had filled the buffer | +| `inout` | read, then written | upload, then the device copy is the live one | + +A unified-memory backend ignores the tag. No op states residency by hand. + +When to copy is decided once, by `gpu::residency` (`gpu_abi.h`): a mirrored +backend keeps one per allocation beside its two copies, asks it +`before_kernel(access)` and `before_host(for_write)`, and does the copying. An +allocation starts `none` (nobody has filled it) unless `alloc` was told the host +fills it. That is what makes `out` safe on a view: a view may cover part of its +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 kernel ABI + +An op hands the backend a kernel id, an ordered list of views, a params +struct, and a grid: + +```cpp +inline bool rmsnorm_res(span x, span delta, span w, span xout, span hout, + int64_t n, float eps, int64_t rows = 1) { + if (n <= 0 || rows <= 0) return false; + return launch(kop::add_rmsnorm_, + {in(x), in(delta), in(w), out(xout), out(hout)}, + rmsnorm_params{static_cast(n), eps}, + policy::one_group_per_row(rows)); +} +``` + +For a kernel id, two things are the same on every backend: + +- **the views**, in the order the kernel declares its buffers; +- **the params struct**: a run of 4-byte fields (`uint32_t`, `int32_t`, + `float`) in the order the kernel takes its scalars. The structs live in + `gpu_abi.h`; the MSL and WGSL sources declare the same layouts on their side. + +That is what lets a backend realize a launch with no per-kernel host code: + +- **Metal** binds view *i* at buffer index *i* (`setBuffer:offset:atIndex:`) + and the params at index *n* (`setBytes`). +- **CUDA** builds `cuLaunchKernel`'s `argv` as the views' device addresses + followed by the params' fields, four bytes apiece. Its kernels take their + pointers first and 4-byte scalars after (116 of the 117 do; + `tools/cuda_trace/gen_kernel_sigs.py` reads this off the `.cu`). +- **WebGPU**'s kernels predate the ABI: every entry point reads one 96-byte + uniform layout, a family picks its operation by number, and the bind group is + fixed (A and B read, C written, D and E read). Its core carries a `marshal_` + from the canonical params into that layout, per kernel id. This stays inside + `webgpu.h`; a backend whose kernels follow the ABI needs none. + +Where two backends' kernels disagreed, the CUDA kernel's order is canonical, +because the `.cu` is the one source no development machine here can run, and +the MSL side can be reordered and tested locally. + +## Two kinds of op + +**A single-kernel op** is one kernel with the same buffers everywhere. It is a +function in `gpu_ops.h` over `launch`, as above, and exists on every backend +that has a kernel for its id. A backend without one returns false from +`dispatch`, and the evaluator falls back to the CPU. + +**A backend-own op** is one whose algorithm differs by backend: a scatter into a +zeroed buffer on CUDA against a gather on Metal and WebGPU; N-D shape metadata +uploaded to a device buffer against a params block; one kernel against a split +pass and a combine. Forcing these into one kernel ABI would mean rewriting +kernels whose differences are deliberate. Instead a backend declares the op as +a static member of its `own` struct, with the signature the op has in +`gpu_ops.h`, and writes it over its own `dispatch`: + +```cpp +// gpu_ops.h — the signature, once +TL_GPU_DETECT_OWN(rope) +template +inline bool rope(span x, span o, int64_t rows, int64_t T, int64_t D, + int64_t pos, float base, span bias = {}) { + if constexpr (detail::owns_rope::value) { + return Own::rope(x, o, rows, T, D, pos, base, bias); + } else { + return false; + } +} + +// metal.h — declared in `struct own`, defined among its helpers +inline bool own::rope(gpu::span x, gpu::span out, ...) { ... } +``` + +`gpu_ops.h` detects the member and forwards to it, or answers false. A backend +declares what it has and nothing else: there are no stubs. A member whose +signature drifts from the shared one is a compile error, not a silent fallback. + +Prefer the single-kernel form. Reach for `own` when the kernels genuinely +differ, not to avoid reordering a params struct. + +## Launch policy + +The shapes ops launch in live in `gpu::policy` (`gpu_abi.h`), shared host code: +`flat`, `flat_rows`, `one_group_per_row`, `row_reduce`, `per_head`, `cells_2d`. +A `grid` is groups x threads-per-group plus the bytes of per-group scratch a +reduction needs where the backend sizes it at launch (CUDA's shared memory; +Metal and WGSL size theirs in the kernel). + +What differs between backends' kernels comes in through the backend's `traits`. +Today that is one fact: whether a rank-2 elementwise kernel reads its cell from +a 2-D thread position or from a flat index (`traits::cells_2d`). + +## Adding an op + +1. If every backend can run it as one kernel with the same buffers: add its + params struct to `gpu_abi.h`, the function to `gpu_ops.h`, the kernel to each + backend's source with that layout, and the id to each backend's kernel table + (`kernel_name_` in `metal.h` and `cuda.h`, `marshal_` in `webgpu.h`). +2. Otherwise add the detecting wrapper to `gpu_ops.h` and the member to the + `own` struct of each backend that implements it. +3. A backend that gets neither simply declines the op. Nothing else changes. + +## Adding a backend + +1. Copy `gpu_null.h`. It is everything `gpu.h` asks of a backend, with nothing + filled in: lifecycle (`available`, `pending`, `flush`, `cpu_barrier`), memory + (`alloc`, `release`, `sync_to_host`, `upload`), `dispatch`, `own`, `traits`, + `caps` and the graph-capture plumbing. +2. Gate the header whole on the platform it builds for, and add one branch to + the selection in `gpu.h`. +3. Write `dispatch` against the kernel ABI, and kernels that follow it. Start + with the elementwise, broadcast, reduction, GEMM, copy and index families: + they close the array surface. Every op the backend has no kernel for falls + back to the CPU, so the suite passes from the first kernel on. +4. Declare in `own` whatever the backend runs its own way. +5. Run the suite in `--gpu` and `--auto` mode and check `gpu::census`: the + suite's oracle comparisons pass whether or not the GPU engaged. + +No existing backend's file is touched. + +`gpu_host.h` is a backend written this way, and the proof that the steps above +are the whole job: a "device" that is the CPU, with kernels that are plain +loops. Its core (lifecycle, memory, `dispatch`'s switch, `traits`, `caps`) is +about 190 lines; its kernels, about 270, cover the single-kernel ops; its `own` +struct, about 70, holds the four backend-own ops the conformance test expects +of every backend. `-DTENSORLIB_HOST_GPU=ON` selects it ahead of any real backend, and +the whole suite passes on it in `--gpu` and `--auto` mode, so the shared layer +and the conformance tests run on a machine with no GPU. Each kernel in it is +also the plainest statement of what its id computes. + +What a test may ask of a backend is asked in code, not by platform macro: +`gpu::has_` (whether the selected backend runs a backend-own op), +`gpu::caps`, `gpu::traits`. A test that needs to know whether a kernel exists +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. +- `tl::profile` hooks sit in each real backend's launch path, because each + stamps its launches with a device time its own way. A backend that records + none (`traits::profiles_launches = false`) gets a row per launch from the + shared layer, by kernel name, so it is profiled from its first kernel. +- `kop` still lists kernel ids only one backend has (Metal's GEMM tiles and + attention variants). +- Launch policy inside the own ops (CUDA's split-K and tile choices, Metal's + GEMM ladder) is still the backend's. +- There is no generic composition under the fused ops, so a backend without the + model-path kernels reports `caps::model_path = false` rather than running them + slowly. WebGPU is in that position. + +## Verifying a change + +All of these run on a development Mac. + +| what | command | +|------|---------| +| Metal, and the CPU | `cmake --build build && ctest --test-dir build` | +| a real model on Metal | `build/tensorlib_check_qwen` (greedy tokens against a numpy oracle) | +| WebGPU | `test/wasm/build.sh && deno run --allow-all test/wasm/deno_run.js` | +| CUDA's host side | `tools/cuda_trace/compare.sh ` | +| no backend | a Linux build without `TENSORLIB_CUDA`: `gpu.h` selects `gpu_null.h` | +| the shared layer, with no GPU | `cmake -B build-host -DTENSORLIB_HOST_GPU=ON`, then `ctest` | + +No machine here runs CUDA kernels, and CI compiles them without running them. +`tools/cuda_trace` puts a stand-in `libcuda.so.1` in front of the backend (it +`dlopen`s the driver, so nothing in `cuda.h` knows), builds the suite and the +CUDA checkers on Linux in a container, and records every launch: kernel, grid, +block, shared-memory bytes and each argument, with pointers printed as +(allocation, byte offset) so two runs diff. What a change to the backend's host +side must preserve is that the same kernels get the same arguments, and that is +what the trace holds. `bench/cuda/check/trace_sweep.cpp` reaches the kernels the +suite does not, so all 117 appear. `TL_CUDA_TRACE_CHECK=1 tools/cuda_trace/run.sh +` only compiles: every test, checker, bench and model driver against the +CUDA branch of the headers. Whether the kernels compute the right thing is +still a `ctest` on NVIDIA hardware. + +`gpu::census(kernel)` counts a shared op's launches and `gpu::ops_run()` every +op that ran on the device, shared or backend-own, since `gpu::census_reset()`. +An op that declines falls back to the CPU and the result is still right, so a +test that wants to know the device was reached has to ask. The suite has two +tests built on this: every op on views at non-zero offsets, inside buffers with +sentinels on both sides, against a plain host loop; and one graph per op family +through the evaluator in GPU mode, where the census has to move. diff --git a/include/array.h b/include/array.h index 557a1f4..b56b13b 100644 --- a/include/array.h +++ b/include/array.h @@ -391,6 +391,9 @@ class array { // buffer to an imperative cuda:: kernel (e.g. the kv_cache decode loop) — // eval() first, then pass native() as the q/k/v pointer. Contiguous, offset 0. void* native() const { return storage_.native; } + // This view as the GPU layer names it: the storage's device handle and the + // view's byte offset (gpu_abi.h). Null `buf` when the storage is heap. + gpu::span device_span() const { return {storage_.native, offset_ * 4}; } // Views (zero-copy on the materialized result) and copies. View // construction realizes the source without a sync — pending GPU kernels @@ -2832,9 +2835,8 @@ struct graph { auto out = array::empty(n.shape); if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; - if (!gpu::binary(*k, a.storage_.native, a.offset_ * 4, - b.storage_.native, b.offset_ * 4, out.storage_.native, - out.offset_ * 4, out.size(), n.scale, n.offset)) { + if (!gpu::binary(*k, a.device_span(), b.device_span(), out.device_span(), + out.size(), n.scale, n.offset)) { return std::nullopt; } return out; @@ -2855,9 +2857,8 @@ struct graph { auto out = array::empty(n.shape); if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; - if (!gpu::binary_bcast(*bk, a.storage_.native, a.offset_ * 4, ra[0], ra[1], - b.storage_.native, b.offset_ * 4, rb[0], rb[1], - out.storage_.native, out.offset_ * 4, n.shape[0], + if (!gpu::binary_bcast(*bk, a.device_span(), ra[0], ra[1], b.device_span(), + rb[0], rb[1], out.device_span(), n.shape[0], n.shape[1], n.scale, n.offset)) { return std::nullopt; } @@ -2879,11 +2880,9 @@ struct graph { if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; std::vector out_shape_v(n.shape.begin(), n.shape.end()); - if (!gpu::binary_bcast_nd(bk, a.storage_.native, a.offset_ * 4, ra.data(), - b.storage_.native, b.offset_ * 4, rb.data(), - out.storage_.native, out.offset_ * 4, - out_shape_v.data(), rank, out.size(), n.scale, - n.offset)) { + if (!gpu::binary_bcast_nd(bk, a.device_span(), ra.data(), b.device_span(), + rb.data(), out.device_span(), out_shape_v.data(), + rank, out.size(), n.scale, n.offset)) { return std::nullopt; } return out; @@ -2914,10 +2913,8 @@ struct graph { if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; std::vector out_shape_v(out_shape.begin(), out_shape.end()); - if (!gpu::where_nd(cond.storage_.native, cond.offset_ * 4, rc.data(), - a.storage_.native, a.offset_ * 4, ra.data(), - b.storage_.native, b.offset_ * 4, rb.data(), - out.storage_.native, out.offset_ * 4, + if (!gpu::where_nd(cond.device_span(), rc.data(), a.device_span(), + ra.data(), b.device_span(), rb.data(), out.device_span(), out_shape_v.data(), rank, out.size())) { return std::nullopt; } @@ -2937,9 +2934,8 @@ struct graph { if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; std::vector shape_v(a.shape().begin(), a.shape().end()); - if (!gpu::copy_nd(a.storage_.native, a.offset_ * 4, a.strides_.data(), - out.storage_.native, out.offset_ * 4, shape_v.data(), - rank, out.size())) { + if (!gpu::copy_nd(a.device_span(), a.strides_.data(), out.device_span(), + shape_v.data(), rank, out.size())) { return std::nullopt; } return out; @@ -3015,19 +3011,15 @@ struct graph { const array bias = wrap(*n.inputs[2]); if (bias.storage_.dt == tl::dtype::f32 && bias.contiguous() && bias.storage_.native && - gpu::gemm_bias(a.storage_.native, a.offset_ * 4, la->ld, la->trans, - b.storage_.native, b.offset_ * 4, lb->ld, lb->trans, - bias.storage_.native, bias.offset_ * 4, - out.storage_.native, out.offset_ * 4, m, nn, k, n.scale, - n.offset)) { + gpu::gemm_bias(a.device_span(), la->ld, la->trans, b.device_span(), + lb->ld, lb->trans, bias.device_span(), + out.device_span(), m, nn, k, n.scale, n.offset)) { bias_owed = false; return out.reshape(n.shape); } } - if (!gpu::gemm(a.storage_.native, a.offset_ * 4, la->ld, la->trans, - b.storage_.native, b.offset_ * 4, lb->ld, lb->trans, - out.storage_.native, out.offset_ * 4, m, nn, k, n.scale, - n.offset)) { + if (!gpu::gemm(a.device_span(), la->ld, la->trans, b.device_span(), lb->ld, + lb->trans, out.device_span(), m, nn, k, n.scale, n.offset)) { return std::nullopt; } return out.reshape(n.shape); @@ -3110,19 +3102,18 @@ struct graph { if (m == 0 || nn == 0 || batch == 0) return out; auto sa = batch_stride_(a), sb = batch_stride_(b); if (bdot_one_launch_ && sa && sb && - gpu::gemm_batched(a.storage_.native, a.offset_ * 4, la->ld, la->trans, - *sa, b.storage_.native, b.offset_ * 4, lb->ld, - lb->trans, *sb, out.storage_.native, out.offset_ * 4, - m, nn, k, batch, n.scale, n.offset)) { + gpu::gemm_batched(a.device_span(), la->ld, la->trans, *sa, + b.device_span(), lb->ld, lb->trans, *sb, + out.device_span(), m, nn, k, batch, n.scale, + n.offset)) { return out; } batch_walk_ w(r - 2); for (int64_t bi = 0; bi < batch; bi++, w.step(a)) { - if (!gpu::gemm(a.storage_.native, (a.offset_ + w.offset(a)) * 4, la->ld, - la->trans, b.storage_.native, - (b.offset_ + w.offset(b)) * 4, lb->ld, lb->trans, - out.storage_.native, (out.offset_ + bi * m * nn) * 4, m, - nn, k, n.scale, n.offset)) { + if (!gpu::gemm(a.device_span().at(w.offset(a) * 4), la->ld, la->trans, + b.device_span().at(w.offset(b) * 4), lb->ld, lb->trans, + out.device_span().at(bi * m * nn * 4), m, nn, k, n.scale, + n.offset)) { return std::nullopt; } } @@ -3194,17 +3185,17 @@ struct graph { array out = array::empty({int64_t{1}, nn}); if (!out.storage_.native) return std::nullopt; if (nn == 0) return out.reshape(n.shape); - bool ok = bf16 ? gpu::gemv_bf16(a->storage_.native, b.storage_.native, - out.storage_.native, nn, k) - : gpu::gemv_f32(a->storage_.native, b.storage_.native, - out.storage_.native, nn, k); + bool ok = bf16 ? gpu::gemv_bf16(a->device_span(), b.device_span(), + out.device_span(), nn, k) + : gpu::gemv_f32(a->device_span(), b.device_span(), + out.device_span(), nn, k); if (!ok) return std::nullopt; return out.reshape(n.shape); } // M8 int4-weight decode GEMV: a(1,K)f32 @ Wq(K,N)q4 -> (1,N)f32. Wq's logical // shape is [K,N]; its storage is packed [N,K] int4 + appended scales, so the - // scales pointer is native + N·K/2 bytes (one buffer). Gated to the decode + // scales are the view N·K/2 bytes into the same buffer. Gated to the decode // shape; non-decode / non-GPU dequantizes to F32 via the input funnel. static std::optional gpu_gemv_q4(const node& n, const array& a_in, const array& Wq) { @@ -3216,10 +3207,9 @@ struct graph { array out = array::empty({int64_t{1}, N}); if (!out.storage_.native) return std::nullopt; if (N == 0) return out.reshape(n.shape); - void* scales = reinterpret_cast( - reinterpret_cast(Wq.storage_.native) + N * K / 2); - if (!gpu::gemv_q4(a->storage_.native, Wq.storage_.native, scales, - out.storage_.native, N, K, tl::kQ4Group)) { + const gpu::span qw = Wq.device_span(); + if (!gpu::gemv_q4(a->device_span(), qw, qw.at(N * K / 2), out.device_span(), + N, K, tl::kQ4Group)) { return std::nullopt; } return out.reshape(n.shape); @@ -3272,9 +3262,8 @@ struct graph { if (!out.storage_.native) return std::nullopt; // Array path has no persistent cache: K/V are [H,ctx,D], so n_kv_heads==H // (no GQA) and kv_max==ctx (kv_stride==ctx*D degenerates to whole-buffer). - if (!gpu::attn_decode(q.storage_.native, K.storage_.native, - V.storage_.native, out.storage_.native, H, H, ctx, ctx, - D, n.arg0)) { + if (!gpu::attn_decode(q.device_span(), K.device_span(), V.device_span(), + out.device_span(), H, H, ctx, ctx, D, n.arg0)) { return std::nullopt; } return out; @@ -3309,9 +3298,8 @@ struct graph { if (!out.storage_.native) return std::nullopt; // No persistent cache on the array path: K/V are [H,T,D], so n_kv_heads==H // (no GQA) and kv_max==T (the whole buffer is the cache, filled from 0). - if (!gpu::attn_prefill(q.storage_.native, K.storage_.native, - V.storage_.native, out.storage_.native, H, H, T, T, - D, n.arg0)) { + if (!gpu::attn_prefill(q.device_span(), K.device_span(), V.device_span(), + out.device_span(), H, H, T, T, D, n.arg0)) { return std::nullopt; } return out; @@ -3362,10 +3350,8 @@ struct graph { profile::scope ps("xent_bwd"); auto out = array::empty(s); if (!out.storage_.native) return std::nullopt; - if (!gpu::xent_bwd(x.storage_.native, x.offset_ * 4, lse.storage_.native, - lse.offset_ * 4, tgt.storage_.native, tgt.offset_ * 4, - g.storage_.native, g.offset_ * 4, out.storage_.native, - out.offset_ * 4, rows, cols)) { + if (!gpu::xent_bwd(x.device_span(), lse.device_span(), tgt.device_span(), + g.device_span(), out.device_span(), rows, cols)) { return std::nullopt; } return out; @@ -3433,13 +3419,11 @@ struct graph { !stats.storage_.native || !partials.storage_.native) { return std::nullopt; } - if (!gpu::layer_norm_bwd(x.storage_.native, x.offset_ * 4, - gamma.storage_.native, gamma.offset_ * 4, - dout.storage_.native, dout.offset_ * 4, - dx.storage_.native, dg.storage_.native, - db.storage_.native, stats.storage_.native, - partials.storage_.native, rows, d, per_chunk, - chunks, eps)) { + if (!gpu::layer_norm_bwd(x.device_span(), gamma.device_span(), + dout.device_span(), dx.device_span(), + dg.device_span(), db.device_span(), + stats.device_span(), partials.device_span(), rows, + d, per_chunk, chunks, eps)) { return std::nullopt; } return std::array{dx, dg, db}; @@ -3545,10 +3529,9 @@ struct graph { const float lr_over_bc1 = lr / bc1, inv_bc2 = 1.0f / bc2; if (gpu_mode_(n, kernel_class::elementwise) && p.storage_.native && m.storage_.native && v.storage_.native && g.storage_.native) { - if (gpu::adam_step(p.storage_.native, p.offset_ * 4, m.storage_.native, - m.offset_ * 4, v.storage_.native, v.offset_ * 4, - g.storage_.native, g.offset_ * 4, n, beta1, beta2, - eps, lr_over_bc1, inv_bc2)) { + if (gpu::adam_step(p.device_span(), m.device_span(), v.device_span(), + g.device_span(), n, beta1, beta2, eps, lr_over_bc1, + inv_bc2)) { return true; } // No kernel for it on this backend: the host loop below. The caller @@ -3612,10 +3595,10 @@ struct graph { profile::scope ps("attn_prefill_bwd_dq"); array dq = array::empty(s), stats = array::empty({2, H, T}); if (!dq.storage_.native || !stats.storage_.native) return std::nullopt; - if (!gpu::attn_prefill_dq(q.storage_.native, K.storage_.native, - V.storage_.native, dout.storage_.native, - out.storage_.native, dq.storage_.native, - stats.storage_.native, H, T, D, scale)) { + if (!gpu::attn_prefill_dq(q.device_span(), K.device_span(), V.device_span(), + dout.device_span(), out.device_span(), + dq.device_span(), stats.device_span(), H, T, D, + scale)) { return std::nullopt; } return std::make_pair(dq, stats); @@ -3658,10 +3641,10 @@ struct graph { profile::scope ps("attn_prefill_bwd_dkv"); array dK = array::empty(s), dV = array::empty(s); if (!dK.storage_.native || !dV.storage_.native) return std::nullopt; - if (!gpu::attn_prefill_dkv(q.storage_.native, K.storage_.native, - V.storage_.native, dout.storage_.native, - stats.storage_.native, dK.storage_.native, - dV.storage_.native, H, T, D, scale)) { + if (!gpu::attn_prefill_dkv(q.device_span(), K.device_span(), + V.device_span(), dout.device_span(), + stats.device_span(), dK.device_span(), + dV.device_span(), H, T, D, scale)) { return std::nullopt; } return std::make_pair(dK, dV); @@ -3869,11 +3852,10 @@ struct graph { int64_t rows = x.size() / D; int64_t T = x.rank() == 3 ? x.shape()[1] : 1; if (!gpu_mode_(x.size(), kernel_class::elementwise)) return std::nullopt; - if (!x.contiguous() || x.offset_ != 0 || !x.storage_.native) - return std::nullopt; + if (!x.contiguous() || !x.storage_.native) return std::nullopt; array out = array::empty(x.shape()); if (!out.storage_.native) return std::nullopt; - if (!gpu::rope(x.storage_.native, out.storage_.native, rows, T, D, n.axis, + if (!gpu::rope(x.device_span(), out.device_span(), rows, T, D, n.axis, n.arg0)) return std::nullopt; return out; @@ -3892,10 +3874,9 @@ struct graph { if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; int64_t d = x.shape().back(); - if (!gpu::layer_norm(x.storage_.native, x.offset_ * 4, g.storage_.native, - g.offset_ * 4, b.storage_.native, b.offset_ * 4, - out.storage_.native, out.offset_ * 4, x.size() / d, d, - n.arg0, n.scale, n.offset)) + if (!gpu::layer_norm(x.device_span(), g.device_span(), b.device_span(), + out.device_span(), x.size() / d, d, n.arg0, n.scale, + n.offset)) return std::nullopt; return out; } @@ -3945,15 +3926,15 @@ struct graph { auto out = array::empty(out_shape); if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; - if (!gpu::row_op(k, a.storage_.native, a.offset_ * 4, out.storage_.native, - out.offset_ * 4, rows, cols, scale, offset)) { + if (!gpu::row_op(k, a.device_span(), out.device_span(), rows, cols, scale, + offset)) { return std::nullopt; } return out; } // The one-input elementwise GPU dispatches share this gate and output: - // `call(a_native, a_off, out_native, out_off, n)` is the backend call. + // `call(a, out, n)` is the backend call, on the two views. template static std::optional gpu_one_input_(const array& a, Call&& call) { if (!gpu_mode_(a.size(), kernel_class::elementwise) || !a.contiguous()) { @@ -3963,8 +3944,7 @@ struct graph { auto out = array::empty(a.shape()); if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; - if (!call(a.storage_.native, a.offset_ * 4, out.storage_.native, - out.offset_ * 4, out.size())) { + if (!call(a.device_span(), out.device_span(), out.size())) { return std::nullopt; } return out; @@ -3974,8 +3954,8 @@ struct graph { const array& a, float scale, float offset) { if (!k) return std::nullopt; // no kernel for this op on this backend - return gpu_one_input_(a, [&](void* an, int64_t ao, void* on, int64_t oo, int64_t n) { - return gpu::unary(*k, an, ao, on, oo, n, scale, offset); + return gpu_one_input_(a, [&](gpu::span in, gpu::span out, int64_t n) { + return gpu::unary(*k, in, out, n, scale, offset); }); } @@ -4011,14 +3991,13 @@ struct graph { int64_t before, const shape_t& out_shape) { return gpu_pad_fold_(a, out_shape, - [&](array& out, std::vector& a_shape, - std::vector& out_shape_v, int rank) { - return gpu::pad(a.storage_.native, a.offset_ * 4, - out.storage_.native, out.offset_ * 4, - a_shape.data(), out_shape_v.data(), - rank, static_cast(axis), before, - a.size(), out.size()); - }); + [&](array& out, std::vector& a_shape, + std::vector& out_shape_v, int rank) { + return gpu::pad(a.device_span(), out.device_span(), + a_shape.data(), out_shape_v.data(), + rank, static_cast(axis), before, + a.size(), out.size()); + }); } // unfold's inverse: scatter-add `a` back into a zero-initialized @@ -4028,14 +4007,13 @@ struct graph { int64_t step, const shape_t& out_shape) { return gpu_pad_fold_(a, out_shape, - [&](array& out, std::vector& a_shape, - std::vector& out_shape_v, int rank) { - return gpu::fold(a.storage_.native, a.offset_ * 4, - out.storage_.native, out.offset_ * 4, - a_shape.data(), out_shape_v.data(), - rank, static_cast(axis), step, - a.size(), out.size()); - }); + [&](array& out, std::vector& a_shape, + std::vector& out_shape_v, int rank) { + return gpu::fold(a.device_span(), out.device_span(), + a_shape.data(), out_shape_v.data(), + rank, static_cast(axis), step, + a.size(), out.size()); + }); } // The GPU-dispatch twin of ref::concat: one gpu::concat_part() launch per @@ -4063,10 +4041,9 @@ struct graph { int64_t offset = 0; for (auto& p : parts) { std::vector p_shape(p.shape().begin(), p.shape().end()); - if (!gpu::concat_part(p.storage_.native, p.offset_ * 4, - out.storage_.native, out.offset_ * 4, - p_shape.data(), out_shape_v.data(), rank, - static_cast(axis), offset, p.size())) { + if (!gpu::concat_part(p.device_span(), out.device_span(), p_shape.data(), + out_shape_v.data(), rank, static_cast(axis), + offset, p.size())) { return std::nullopt; } offset += p_shape[static_cast(axis)]; @@ -4088,10 +4065,8 @@ struct graph { if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; int64_t row_size = a.size() / a.shape()[0]; - if (!gpu::index_select(a.storage_.native, a.offset_ * 4, - indices.storage_.native, indices.offset_ * 4, - out.storage_.native, out.offset_ * 4, row_size, - out_shape[0])) { + if (!gpu::index_select(a.device_span(), indices.device_span(), + out.device_span(), row_size, out_shape[0])) { return std::nullopt; } return out; @@ -4145,10 +4120,9 @@ struct graph { if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; int64_t row_size = out.size() / target_shape[0]; - if (!gpu::index_add(indices.storage_.native, indices.offset_ * 4, - values.storage_.native, values.offset_ * 4, - out.storage_.native, out.offset_ * 4, row_size, - indices.shape()[0], out.size())) { + if (!gpu::index_add(indices.device_span(), values.device_span(), + out.device_span(), row_size, indices.shape()[0], + out.size())) { return std::nullopt; } return out; @@ -4170,10 +4144,9 @@ struct graph { auto out = array::empty(out_shape); if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; - if (!gpu::scatter_to_axis(indices.storage_.native, indices.offset_ * 4, - values.storage_.native, values.offset_ * 4, - out.storage_.native, out.offset_ * 4, - values.size(), out_shape.back())) { + if (!gpu::scatter_to_axis(indices.device_span(), values.device_span(), + out.device_span(), values.size(), + out_shape.back())) { return std::nullopt; } return out; @@ -4192,9 +4165,8 @@ struct graph { auto out = array::empty(out_shape); if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; - if (!gpu::gather_from_axis(src.storage_.native, src.offset_ * 4, - indices.storage_.native, indices.offset_ * 4, - out.storage_.native, out.offset_ * 4, out.size(), + if (!gpu::gather_from_axis(src.device_span(), indices.device_span(), + out.device_span(), out.size(), src.shape().back())) { return std::nullopt; } @@ -4217,8 +4189,7 @@ struct graph { auto out = array::empty(out_shape); if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; - if (!gpu::row_logsumexp(a.storage_.native, a.offset_ * 4, - out.storage_.native, out.offset_ * 4, rows, cols, + if (!gpu::row_logsumexp(a.device_span(), out.device_span(), rows, cols, scale, offset)) { return std::nullopt; } @@ -4245,9 +4216,9 @@ struct graph { for (int d = 0; d < rank; d++) { if (acc[d] == 0) reduced_n *= a_shape_v[d]; } - if (!gpu::sum_to(a.storage_.native, a.offset_ * 4, a_shape_v.data(), - a_strides_v.data(), acc.data(), rank, out.size(), - reduced_n, out.storage_.native, out.offset_ * 4)) { + if (!gpu::sum_to(a.device_span(), a_shape_v.data(), a_strides_v.data(), + acc.data(), rank, out.size(), reduced_n, + out.device_span())) { return std::nullopt; } return out; @@ -4308,8 +4279,7 @@ struct graph { auto out = array::empty(a.shape()); if (out.size() == 0) return out; if (!out.storage_.native) return std::nullopt; - if (!gpu::compare(c, a.storage_.native, a.offset_ * 4, b.storage_.native, - b.offset_ * 4, out.storage_.native, out.offset_ * 4, + if (!gpu::compare(c, a.device_span(), b.device_span(), out.device_span(), out.size(), bstride)) { return std::nullopt; } @@ -4329,8 +4299,8 @@ struct graph { case op_t::cos_: u = gpu::unary_ext_op::cos_; break; default: return std::nullopt; } - return gpu_one_input_(a, [&](void* an, int64_t ao, void* on, int64_t oo, int64_t n) { - return gpu::unary_ext(u, an, ao, on, oo, n, scale, offset); + return gpu_one_input_(a, [&](gpu::span in, gpu::span out, int64_t n) { + return gpu::unary_ext(u, in, out, n, scale, offset); }); } @@ -4339,8 +4309,8 @@ struct graph { // role scale/offset play elsewhere, and nothing composes a further affine // onto it today. static std::optional gpu_clamp_(const array& a, float lo, float hi) { - return gpu_one_input_(a, [&](void* an, int64_t ao, void* on, int64_t oo, int64_t n) { - return gpu::clamp(an, ao, on, oo, n, lo, hi); + return gpu_one_input_(a, [&](gpu::span in, gpu::span out, int64_t n) { + return gpu::clamp(in, out, n, lo, hi); }); } @@ -4358,8 +4328,8 @@ struct graph { case op_t::ne_s: k = gpu::scalar_op::ne; break; default: return std::nullopt; } - return gpu_one_input_(a, [&](void* an, int64_t ao, void* on, int64_t oo, int64_t len) { - return gpu::scalar_binary(k, an, ao, on, oo, len, n.arg0, n.scale, n.offset); + return gpu_one_input_(a, [&](gpu::span in, gpu::span out, int64_t len) { + return gpu::scalar_binary(k, in, out, len, n.arg0, n.scale, n.offset); }); } diff --git a/include/cuda.h b/include/cuda.h index 97bf05a..daf3600 100644 --- a/include/cuda.h +++ b/include/cuda.h @@ -20,31 +20,20 @@ // than on device memory — the roadmap's pre-authorized device-buffer pivot. // View offsets are folded host-side into the pointer passed to each kernel. // -// Real implementation is gated on TENSORLIB_CUDA && !__APPLE__ (Apple uses -// Metal; a plain build gets the stubs below). The API matches metal.h exactly -// — available/pending/flush/alloc/release/binary/unary/gemm/row_op — so the -// eval_one dispatch seam is backend-agnostic and carries no platform #ifdefs. +// The whole header is gated on TENSORLIB_CUDA && !__APPLE__ (Apple uses Metal): +// elsewhere it declares nothing, and gpu.h selects another backend (or +// gpu_null.h). What it provides is the device core gpu.h describes, so the +// eval seam is backend-agnostic and carries no platform #ifdefs. #include -#include "metal.h" // reuse tl::metal::kop (platform-independent op enum) +#include "gpu_abi.h" // the op vocabulary and the launch contract #include "profile.h" // tl::profile (per-launch attribution and timing) #include "shape.h" // tl::contiguous_strides_into (pad/fold meta upload) #include "types.h" // tl::dtype (KV cache storage width) -namespace tl { -namespace cuda { - -using kop = tl::metal::kop; -using cmp_op = tl::metal::cmp_op; -using unary_ext_op = tl::metal::unary_ext_op; -using scalar_op = tl::metal::scalar_op; - #if defined(TENSORLIB_CUDA) && !defined(__APPLE__) -} // namespace cuda -} // namespace tl - #ifdef _WIN32 #ifndef WIN32_LEAN_AND_MEAN #define WIN32_LEAN_AND_MEAN @@ -68,6 +57,11 @@ using scalar_op = tl::metal::scalar_op; namespace tl { namespace cuda { +using kop = gpu::kop; +using cmp_op = gpu::cmp_op; +using unary_ext_op = gpu::unary_ext_op; +using scalar_op = gpu::scalar_op; + // Dynamic-loader shim: dlopen/dlsym on Unix, LoadLibrary/GetProcAddress on // Windows (where the driver ships as nvcuda.dll). Symbols are cast to the // hand-declared function-pointer types by the caller, same as before. @@ -194,10 +188,52 @@ inline const char* kernel_name_(kop op) { case kop::sigmoid: return "tl_sigmoid"; case kop::relu: return "tl_relu"; case kop::affine: return "tl_affine"; + case kop::tanh_: return "tl_tanh"; + case kop::sin_: return "tl_sin"; + case kop::cos_: return "tl_cos"; case kop::softmax: return "tl_softmax"; case kop::row_sum: return "tl_row_sum"; case kop::row_max: return "tl_row_max"; - default: return "tl_sgemm"; // sgemm* / steel* all route to tl_sgemm + case kop::gt_: return "tl_gt"; + case kop::lt_: return "tl_lt"; + case kop::ge_: return "tl_ge"; + case kop::le_: return "tl_le"; + case kop::eq_: return "tl_eq"; + case kop::ne_: return "tl_ne"; + case kop::clamp_: return "tl_clamp"; + case kop::pow_s_: return "tl_pow_s"; + case kop::gt_s_: return "tl_gt_s"; + case kop::lt_s_: return "tl_lt_s"; + case kop::ge_s_: return "tl_ge_s"; + case kop::le_s_: return "tl_le_s"; + case kop::eq_s_: return "tl_eq_s"; + case kop::ne_s_: return "tl_ne_s"; + case kop::layer_norm_: return "tl_layer_norm"; + case kop::index_select: return "tl_index_select"; + case kop::gather_axis_: return "tl_gather_axis"; + case kop::row_logsumexp_: return "tl_row_logsumexp"; + case kop::xent_bwd_: return "tl_xent_bwd"; + case kop::adam_step_: return "tl_adam_step"; + case kop::rmsnorm_: return "tl_rmsnorm"; + case kop::add_rmsnorm_: return "tl_add_rmsnorm"; + case kop::swiglu_: return "tl_swiglu"; + case kop::gemv_bf16_row_: return "tl_gemv_bf16_row"; + case kop::gemv_q4_: return "tl_gemv_q4"; + case kop::kv_append_: return "tl_kv_append"; + case kop::kv_append_bf16_: return "tl_kv_append_bf16"; + case kop::kv_fill_: return "tl_kv_fill"; + case kop::kv_fill_bf16_: return "tl_kv_fill_bf16"; + 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 } } @@ -255,13 +291,13 @@ struct context { // 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 dirty state serves every view. loc tracks where the live copy is. - enum loc { HOST, DEVICE, BOTH }; + // 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; - loc where = HOST; + gpu::residency live; }; std::unordered_map mirrors; @@ -276,20 +312,20 @@ struct context { auto it = mirrors.find(reinterpret_cast(native)); return it == mirrors.end() ? nullptr : &it->second; } - // A kernel is about to READ this buffer: ensure the device copy is current. - // Async on the stream like the meta uploads (a blocking copy would wait out - // every kernel already queued and stall the pipeline mid-graph); the driver - // stages a pageable source during the call. - void device_read_(void* 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 + // mid-graph); the driver stages a pageable source during the call. + void before_kernel_(void* native, gpu::access a) { mirror* m = mirror_(native); - if (m && m->where == HOST) { - upload_(*m); - m->where = BOTH; - } - } - // The H2D behind device_read_ / device_rmw_: async on the stream when the - // driver has it. Profiled as a transfer either way (the blocking form with - // its wait). + if (m && m->live.before_kernel(a)) upload_(*m); + } + // The three accesses by name, for the ops that state residency themselves. + void device_read_(void* native) { before_kernel_(native, gpu::access::in); } + void device_write_(void* native) { before_kernel_(native, gpu::access::out); } + void device_rmw_(void* native) { before_kernel_(native, gpu::access::inout); } + // The H2D: async on the stream when the driver has it. Profiled as a + // transfer either way (the blocking form with its wait). void upload_(const mirror& m) { if (d.MemcpyHtoDAsync) { d.MemcpyHtoDAsync(m.dev, m.host, m.bytes, stream); @@ -299,22 +335,6 @@ struct context { profile::detail::blocked timing{"h2d", m.bytes}; d.MemcpyHtoD(m.dev, m.host, m.bytes); } - // A kernel is about to WRITE every element of this buffer: it becomes the - // live copy, and whatever the host held is dead. A kernel that reads it - // first, or writes only part of it, is device_rmw_ below. - void device_write_(void* native) { - if (mirror* m = mirror_(native)) m->where = DEVICE; - } - // A kernel is about to READ and then WRITE this buffer (an in-place update): - // a host-born copy comes up first, then the device copy is the live one. One - // probe for both -- the mirror map is every CUDA buffer, and an optimizer - // step does this per parameter. - void device_rmw_(void* native) { - mirror* m = mirror_(native); - if (!m) return; - if (m->where == HOST) upload_(*m); - m->where = DEVICE; - } static context& get() { static auto* c = new context(); // leaked: outlives all storage deleters @@ -427,7 +447,8 @@ struct context { int key = static_cast(op); auto it = fns.find(key); if (it != fns.end()) return it->second; - CUfunction f = load_(kernel_name_(op)); + const char* name = kernel_name_(op); + CUfunction f = name ? load_(name) : nullptr; fns[key] = f; return f; } @@ -490,19 +511,6 @@ struct context { CUfunction gemv_f32_() { return cached_(gemv_f32_fn, "tl_gemv_f32"); } CUfunction gemv_bf16_() { return cached_(gemv_bf16_fn, "tl_gemv_bf16"); } CUfunction gemv_bf16v8_() { return cached_(gemv_bf16v8_fn, "tl_gemv_bf16v8"); } - CUfunction gemv_bf16_row_fn = nullptr; - CUfunction gemv_bf16_row_() { - return cached_(gemv_bf16_row_fn, "tl_gemv_bf16_row"); - } - - // M8 int4-weight decode GEMV. - CUfunction gemv_q4_fn = nullptr; - CUfunction gemv_q4_() { return cached_(gemv_q4_fn, "tl_gemv_q4"); } - - // M9 batched-prefill layout kernels (token-major <-> head-major). - CUfunction split_heads_fn = nullptr, merge_heads_fn = nullptr; - CUfunction split_heads_() { return cached_(split_heads_fn, "tl_split_heads"); } - CUfunction merge_heads_() { return cached_(merge_heads_fn, "tl_merge_heads"); } // im2col's pad/fold, cached like split_heads/merge_heads. CUfunction pad_fn = nullptr, fold_fn = nullptr; @@ -511,33 +519,12 @@ struct context { // Embedding-table lookup (index_select/index_add) and pooling-style // one-hot scatter (scatter_to_axis), cached the same way. - CUfunction index_select_fn = nullptr, index_add_fn = nullptr, - scatter_axis_fn = nullptr; - CUfunction index_select_() { - return cached_(index_select_fn, "tl_index_select"); - } + CUfunction index_add_fn = nullptr, scatter_axis_fn = nullptr; CUfunction index_add_() { return cached_(index_add_fn, "tl_index_add"); } CUfunction scatter_axis_() { return cached_(scatter_axis_fn, "tl_scatter_axis"); } - // Cross-entropy's three: the trailing-axis gather (scatter_axis_'s dual), - // the one-pass row logsumexp its forward reduces with, and the pullback - // that rebuilds the softmax from that logsumexp. - CUfunction gather_axis_fn = nullptr, row_logsumexp_fn = nullptr, - xent_bwd_fn = nullptr; - CUfunction gather_axis_() { - return cached_(gather_axis_fn, "tl_gather_axis"); - } - CUfunction row_logsumexp_() { - return cached_(row_logsumexp_fn, "tl_row_logsumexp"); - } - CUfunction xent_bwd_() { return cached_(xent_bwd_fn, "tl_xent_bwd"); } - - // Adam's fused per-parameter update (the optimizer's whole step, one launch). - CUfunction adam_step_fn = nullptr; - CUfunction adam_step_() { return cached_(adam_step_fn, "tl_adam_step"); } - // N-D broadcast binary (any rank) and N-D broadcast ternary select // (Tensor.where's GPU dispatch) -- new capabilities, one kernel per op // like the rank-2 kop/fn_() vocabulary above, but not part of that @@ -572,55 +559,6 @@ struct context { return cached_(sum_to_blocked_fn, "tl_sum_to_blocked"); } - // Comparisons (gt/lt/ge/le/eq/ne): ReLU/LeakyReLU/Clip's backward gate - // and Tensor.gt/lt/... generally. Own vocabulary, not the kop table - // (see metal.h's cmp_op comment for why). - CUfunction gt_fn = nullptr, lt_fn = nullptr, ge_fn = nullptr, - le_fn = nullptr, eq_fn = nullptr, ne_fn = nullptr; - CUfunction compare_(cmp_op op) { - switch (op) { - case cmp_op::gt: return cached_(gt_fn, "tl_gt"); - case cmp_op::lt: return cached_(lt_fn, "tl_lt"); - case cmp_op::ge: return cached_(ge_fn, "tl_ge"); - case cmp_op::le: return cached_(le_fn, "tl_le"); - case cmp_op::eq: return cached_(eq_fn, "tl_eq"); - case cmp_op::ne: return cached_(ne_fn, "tl_ne"); - default: return nullptr; - } - } - - // tanh_/sin_/cos_ (RoPE's trig, RNN/LSTM's tanh) and clamp (Clip's - // forward). Own vocabulary, not the kop table (same reason as compare_ - // above). - CUfunction tanh_fn = nullptr, sin_fn = nullptr, cos_fn = nullptr, - clamp_fn = nullptr; - CUfunction unary_ext_(unary_ext_op op) { - switch (op) { - case unary_ext_op::tanh_: return cached_(tanh_fn, "tl_tanh"); - case unary_ext_op::sin_: return cached_(sin_fn, "tl_sin"); - case unary_ext_op::cos_: return cached_(cos_fn, "tl_cos"); - default: return nullptr; - } - } - CUfunction clamp_() { return cached_(clamp_fn, "tl_clamp"); } - - // Tensor-scalar ops (pow(x, s), x > s, ...): the scalar is a kernel argument. - CUfunction pow_s_fn = nullptr, gt_s_fn = nullptr, lt_s_fn = nullptr, - ge_s_fn = nullptr, le_s_fn = nullptr, eq_s_fn = nullptr, - ne_s_fn = nullptr; - CUfunction scalar_binary_(scalar_op op) { - switch (op) { - case scalar_op::pow: return cached_(pow_s_fn, "tl_pow_s"); - case scalar_op::gt: return cached_(gt_s_fn, "tl_gt_s"); - case scalar_op::lt: return cached_(lt_s_fn, "tl_lt_s"); - case scalar_op::ge: return cached_(ge_s_fn, "tl_ge_s"); - case scalar_op::le: return cached_(le_s_fn, "tl_le_s"); - case scalar_op::eq: return cached_(eq_s_fn, "tl_eq_s"); - case scalar_op::ne: return cached_(ne_s_fn, "tl_ne_s"); - default: return nullptr; - } - } - // M9 batched-prefill GEMM (bf16 [N,K] weights, the decode GEMV's own layout). CUfunction gemm_bf16_nt_fn = nullptr, gemm_bf16_nt_s_fn = nullptr, gemm_bf16_nt_sk_fn = nullptr; @@ -659,13 +597,6 @@ struct context { return cached_(attn_combine_fn, "tl_attn_combine"); } - // M9 KV cache append (scatter one token's k,v into the persistent cache). - CUfunction kv_append_fn = nullptr, kv_append_bf16_fn = nullptr; - CUfunction kv_append_(bool bf16 = false) { - return bf16 ? cached_(kv_append_bf16_fn, "tl_kv_append_bf16") - : cached_(kv_append_fn, "tl_kv_append"); - } - // RoPE (rotary position embedding) for q/k. CUfunction rope_fn = nullptr; CUfunction rope_() { return cached_(rope_fn, "tl_rope"); } @@ -685,26 +616,15 @@ struct context { : cached_(attn_split_dpos_fn, "tl_attn_decode_split_dpos"); } - // GPU argmax (greedy last-mile): kernel + a persistent 4-byte device result - // buffer so the per-token result is a 4-byte D2H, not the 608KB logits copy. - CUfunction argmax_fn = nullptr; - CUfunction argmax_() { return cached_(argmax_fn, "tl_argmax"); } + // GPU argmax (greedy last-mile): a persistent 4-byte device result buffer, so + // the per-token result is a 4-byte D2H, not the 608KB logits copy. CUdeviceptr argmax_res = 0; CUdeviceptr argmax_res_() { if (!argmax_res && d.MemAlloc(&argmax_res, 16) != 0) argmax_res = 0; return argmax_res; } - // Fused decode-step ops (imperative path): RMSNorm + SwiGLU. - CUfunction rmsnorm_fn = nullptr, swiglu_fn = nullptr; - CUfunction rmsnorm_() { return cached_(rmsnorm_fn, "tl_rmsnorm"); } - CUfunction swiglu_() { return cached_(swiglu_fn, "tl_swiglu"); } - CUfunction add_rmsnorm_fn = nullptr; - CUfunction add_rmsnorm_() { return cached_(add_rmsnorm_fn, "tl_add_rmsnorm"); } - - // The graph's fused layer norm, and its pullback's three kernels. - CUfunction layer_norm_fn = nullptr; - CUfunction layer_norm_() { return cached_(layer_norm_fn, "tl_layer_norm"); } + // The fused layer norm's pullback: three kernels. CUfunction layer_norm_bwd_dx_fn = nullptr; CUfunction layer_norm_bwd_dx_() { return cached_(layer_norm_bwd_dx_fn, "tl_layer_norm_bwd_dx"); @@ -722,12 +642,7 @@ struct context { CUfunction fill_rows_fn = nullptr; CUfunction fill_rows_() { return cached_(fill_rows_fn, "tl_fill_rows"); } - // M9 prefill: bulk cache fill + causal prefill attention. - CUfunction kv_fill_fn = nullptr, kv_fill_bf16_fn = nullptr; - CUfunction kv_fill_(bool bf16 = false) { - return bf16 ? cached_(kv_fill_bf16_fn, "tl_kv_fill_bf16") - : cached_(kv_fill_fn, "tl_kv_fill"); - } + // M9 prefill: causal prefill attention. CUfunction attn_prefill_tiled_fn = nullptr, attn_prefill_tiled_64_fn = nullptr, attn_prefill_tiled_bf16_fn = nullptr, attn_prefill_tiled_bf16_64_fn = nullptr; @@ -819,8 +734,14 @@ struct context { "kernel arg must be a pointer or a 4-byte scalar: an 8-byte " "one (int64_t/size_t/double/nullptr) shifts every arg after " "it. Cast to unsigned/float at the call site."); - if (!f) return false; void* argv[] = {&args...}; + return launch_argv_(f, grid, block, smem, argv); + } + + // The launch itself, once argv is built (by launch_ above, or by dispatch_). + bool launch_argv_(CUfunction f, dims grid, dims block, unsigned smem, + void** argv) { + if (!f) return false; pending = true; // Profiling: the launch under the open scope, and — outside a graph // capture, where an event record would become a graph node — an event on @@ -843,6 +764,32 @@ struct context { return ok; } + // The shared layer's launch (gpu_abi.h): residency from each view's access, + // then argv as the views' device addresses followed by the params' 4-byte + // fields. It is launch_'s contract read the other way round — every kernel + // takes its pointers first and 4-byte scalars after — so a params struct + // laid out in the kernel's argument order expands with no per-kernel code. + bool dispatch_(CUfunction f, const gpu::arg* args, size_t n, + const void* params, size_t params_bytes, const gpu::grid& g) { + constexpr size_t kMaxArgs = 32; + const size_t words = params_bytes / 4; + if (!f || params_bytes % 4 || n + words == 0 || n + words > kMaxArgs) { + return false; + } + void* ptrs[kMaxArgs]; + uint32_t scalars[kMaxArgs]; + void* argv[kMaxArgs]; + for (size_t i = 0; i < n; i++) { + before_kernel_(args[i].s.buf, args[i].a); + ptrs[i] = off_(args[i].s.buf, args[i].s.off); + argv[i] = &ptrs[i]; + } + std::memcpy(scalars, params, params_bytes); + for (size_t j = 0; j < words; j++) argv[n + j] = &scalars[j]; + return launch_argv_(f, {g.gx, g.gy, g.gz}, {g.tx, g.ty, g.tz}, + g.scratch_bytes, argv); + } + // The 1-D elementwise shape: 256-thread blocks covering n elements. template bool launch1d_(CUfunction f, unsigned n, Ts... args) { @@ -854,6 +801,101 @@ struct context { inline bool available() { return context::get().ready; } +inline bool dispatch(kop k, const gpu::arg* args, size_t n, const void* params, + size_t params_bytes, const gpu::grid& g); + +// The ops this backend runs its own way: a different algorithm, several +// kernels, or a kernel whose ABI is its own. gpu_ops.h forwards to whichever of +// these exist (TL_GPU_DETECT_OWN) and answers false for the rest, so a backend +// declares what it has and nothing else. Defined below, among their helpers. +struct own { + static bool split_heads(gpu::span src, gpu::span bias, gpu::span dst, + int64_t T, int64_t ld, int64_t off, int64_t H, + int64_t D); + static bool argmax(gpu::span a, int64_t n, int64_t* out_idx); + // The graph-capture forms: the position is a device scalar, so a captured + // step replays against an advancing cache row. + static bool rope_dpos(gpu::span x, gpu::span out, int64_t rows, int64_t T, + int64_t D, gpu::span d_pos, float base, + gpu::span bias = {}); + static bool kv_append_dpos(gpu::span Kc, gpu::span Vc, gpu::span k_new, + gpu::span v_new, gpu::span d_pos, int64_t kv_max, + int64_t n_kv_heads, int64_t D); + static bool attn_decode_dpos(gpu::span q, gpu::span K, gpu::span V, + gpu::span out, int64_t n_q_heads, + int64_t n_kv_heads, gpu::span d_pos, + int64_t kv_max, int64_t D, float scale, + gpu::span partials); + static bool incr_u32(gpu::span d_pos); + static bool binary_bcast_nd(kop op, gpu::span a, const int64_t* a_strides, + gpu::span b, const int64_t* b_strides, + gpu::span out, const int64_t* out_shape, int rank, + int64_t n, float scale, float offset); + static bool where_nd(gpu::span cond, const int64_t* c_strides, gpu::span a, + const int64_t* a_strides, gpu::span b, + const int64_t* b_strides, gpu::span out, + const int64_t* out_shape, int rank, int64_t n); + static bool copy_nd(gpu::span a, const int64_t* a_strides, gpu::span out, + const int64_t* out_shape, int rank, int64_t n); + static bool sum_to(gpu::span a, const int64_t* a_shape, + const int64_t* a_strides, const int64_t* acc, int rank, + int64_t out_n, int64_t reduced_n, gpu::span out); + static bool pad(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, int64_t before, + int64_t n, int64_t out_n); + static bool fold(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, int64_t step, + int64_t n, int64_t out_n); + static bool concat_part(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t before, int64_t n); + static bool index_add(gpu::span idx, gpu::span values, gpu::span out, + int64_t row_size, int64_t k, int64_t out_n); + static bool scatter_to_axis(gpu::span idx, gpu::span values, gpu::span out, + int64_t n, int64_t size); + static bool gemm(gpu::span a, int64_t lda, bool ta, gpu::span b, int64_t ldb, + bool tb, gpu::span out, int64_t m, int64_t n, int64_t k, + float scale, float offset); + static bool gemm_batched(gpu::span a, int64_t lda, bool ta, int64_t sa, + gpu::span b, int64_t ldb, bool tb, int64_t sb, + gpu::span out, int64_t m, int64_t n, int64_t k, + int64_t batch, float scale, float offset, + gpu::span bias = {}); + static bool gemm_bias(gpu::span a, int64_t lda, bool ta, gpu::span b, + int64_t ldb, bool tb, gpu::span bias, gpu::span out, + int64_t m, int64_t n, int64_t k, float scale, + float offset); + static bool gemv_f32(gpu::span a, gpu::span B, gpu::span y, int64_t n, + int64_t k); + static bool gemv_bf16(gpu::span a, gpu::span B, gpu::span y, int64_t n, + int64_t k); + static bool attn_decode(gpu::span q, gpu::span K, gpu::span V, gpu::span out, + int64_t n_q_heads, int64_t n_kv_heads, int64_t ctx, + int64_t kv_max, int64_t D, float scale, + bool kv_bf16 = false); + static bool attn_prefill(gpu::span q, gpu::span K, gpu::span V, gpu::span out, + int64_t n_q_heads, int64_t n_kv_heads, int64_t T, + int64_t kv_max, int64_t D, float scale, + bool kv_bf16 = false, int64_t pos0 = 0); + static bool attn_prefill_dq(gpu::span q, gpu::span K, gpu::span V, + gpu::span dO, gpu::span O, gpu::span dq, + gpu::span stats, int64_t H, int64_t T, int64_t D, + float scale); + static bool attn_prefill_dkv(gpu::span q, gpu::span K, gpu::span V, + gpu::span dO, gpu::span stats, gpu::span dK, + gpu::span dV, int64_t H, int64_t T, int64_t D, + float scale); + static bool gemm_bf16_nt(gpu::span a, gpu::span B, gpu::span out, int64_t m, + int64_t n, int64_t k); + static bool rope(gpu::span x, gpu::span out, int64_t rows, int64_t T, + int64_t D, int64_t pos, float base, gpu::span bias = {}); + static bool layer_norm_bwd(gpu::span x, gpu::span g, gpu::span dy, + gpu::span dx, gpu::span dg, gpu::span db, + gpu::span stats, gpu::span partials, int64_t rows, + int64_t cols, int64_t per_chunk, int64_t chunks, + float eps); +}; + // Blocks that keep the GPU busy: ~2 per SM on the 82-SM RTX 3090. The // threshold every tile and split-K choice below measures its grid against. constexpr long kFillBlocks = 164; @@ -961,6 +1003,54 @@ inline void flush() { // imperative decode step (no host sync / blocking copy mid-stream) is // capturable; embed staging + argmax happen outside the captured region. // What a model may ask of this backend beyond the kernel contract (gpu.h). + +// "No bias" is a null pointer here, which the kernel tests. +inline bool own::split_heads(gpu::span src, gpu::span bias, gpu::span dst, + int64_t T, int64_t ld, int64_t off, int64_t H, + int64_t D) { + struct { + uint32_t T, ld, off, D; + } p{static_cast(T), static_cast(ld), + static_cast(off), static_cast(D)}; + const gpu::arg args[] = {gpu::in(src), gpu::in(bias), gpu::out(dst)}; + return dispatch(kop::split_heads_, args, 3, &p, sizeof(p), + gpu::policy::per_head(H, T, D)); +} + +// One int comes back through a 4-byte device buffer the context keeps; the +// reduction carries a (value, index) pair per thread. +inline bool own::argmax(gpu::span a, int64_t n, int64_t* out_idx) { + auto& c = context::get(); + if (!c.ready) return false; + CUdeviceptr res = c.argmax_res_(); + if (!res) return false; + const uint32_t p = static_cast(n); + const gpu::arg args[] = {gpu::in(a), + gpu::out({reinterpret_cast(res), 0})}; + const uint32_t block = 256; + if (!dispatch(kop::argmax_, args, 2, &p, sizeof(p), + {1, 1, 1, block, 1, 1, + block * uint32_t(sizeof(float) + sizeof(int))})) { + return false; + } + flush(); // the result index must be ready before the 4-byte D2H + int h = 0; + if (c.d.MemcpyDtoH(&h, res, sizeof(int)) != 0) return false; + *out_idx = h; + return true; +} + +// What the shared launch policy (gpu_ops.h) may assume of this backend's +// kernels. +struct traits { + // A [rows, cols] elementwise kernel reads its cell from a flat index. + static constexpr bool cells_2d = false; + // Launches are recorded under tl::profile by this backend itself, with + // (times_launches) a device time on each. + static constexpr bool profiles_launches = true; + static constexpr bool times_launches = true; +}; + struct caps { // Whether the model-path row is real here, or answers false: a decoder // runs on raw buffers only where it is true, and keeps to the array ops @@ -969,9 +1059,6 @@ struct caps { static constexpr bool graph_capture = true; static constexpr bool row_gemv = true; // gemv_bf16_row: weights as [N,K] static constexpr bool bf16_gemm = true; // gemm_bf16_nt: the batched prefill - // A device pointer is an address, so base + n names a mid-buffer location - // and a fused kernel's output can be read back in slices. - static constexpr bool flat_addressing = true; }; using graph_exec = CUgraphExec; @@ -1030,7 +1117,7 @@ inline void upload(void* native, const float* src, int64_t n) { size_t bytes = std::min((size_t)n * sizeof(float), m->bytes); if (src != m->host) std::memcpy(m->host, src, bytes); c.d.MemcpyHtoD(m->dev, m->host, bytes); - m->where = context::BOTH; + m->live.uploaded(); } // Set a device u32 scalar (e.g. the capture pos counter) via its mirror. Raw @@ -1042,7 +1129,7 @@ inline void upload_u32(void* native, unsigned val) { if (!m || m->bytes < 4) return; std::memcpy(m->host, &val, 4); c.d.MemcpyHtoD(m->dev, m->host, 4); - m->where = context::BOTH; + m->live.uploaded(); } // Mirror allocation: a device buffer (returned as `native`) paired with a host @@ -1051,7 +1138,7 @@ inline void upload_u32(void* native, unsigned val) { // keeps native != contents, like Metal (MTLBuffer handle vs .contents pointer). // `host_fill` (the host writes it first) needs nothing here: the host copy is // its own allocation, which no queued work writes. -inline void* alloc(int64_t bytes, float** contents, bool /*host_fill*/ = false) { +inline void* alloc(int64_t bytes, float** contents, bool host_fill = false) { auto& c = context::get(); if (!c.ready) return nullptr; size_t nb = bytes > 0 ? (size_t)bytes : 4; @@ -1070,7 +1157,7 @@ inline void* alloc(int64_t bytes, float** contents, bool /*host_fill*/ = false) return nullptr; } } - c.mirrors[dev] = context::mirror{host, dev, nb, context::HOST}; + c.mirrors[dev] = context::mirror{host, dev, nb, gpu::residency(host_fill)}; if (contents) *contents = host; return reinterpret_cast(dev); } @@ -1099,65 +1186,19 @@ inline void sync_to_host(void* native, bool for_write) { if (!c.ready || !native) return; context::mirror* m = c.mirror_(native); if (!m) return; - if (m->where == context::DEVICE) { + if (m->live.before_host(for_write)) { if (c.pending) flush(); - { - profile::detail::blocked timing{"d2h", m->bytes}; - c.d.MemcpyDtoH(m->host, m->dev, m->bytes); - } - m->where = context::BOTH; + profile::detail::blocked timing{"d2h", m->bytes}; + c.d.MemcpyDtoH(m->host, m->dev, m->bytes); } - if (for_write) m->where = context::HOST; } -// out = (a OP b) * scale + offset, contiguous; offsets in bytes. -inline bool binary(kop op, void* a, int64_t ao, void* b, int64_t bo, void* out, - int64_t oo, int64_t n, float scale, float offset) { +// The device core's one way to run a kernel for the shared ops (gpu_ops.h). +inline bool dispatch(kop k, const gpu::arg* args, size_t n, const void* params, + size_t params_bytes, const gpu::grid& g) { auto& c = context::get(); if (!c.ready) return false; - c.device_read_(a); - c.device_read_(b); - c.device_write_(out); - float* pa = context::off_(a, ao); - float* pb = context::off_(b, bo); - float* po = context::off_(out, oo); - unsigned un = static_cast(n); - return c.launch1d_(c.fn_(op), un, pa, pb, po, un, scale, offset); -} - -// Rank-2 broadcast binary (bias / row-vector / column-vector / scalar) -- -// mirrors metal.h's own kernel; ars/acs/brs/bcs are element strides (0 on a -// broadcast axis), computed host-side by array.h's gpu_binary via the same -// broadcast_strides() the CPU oracle uses. -inline bool binary_bcast(kop op, void* a, int64_t ao, int64_t ars, int64_t acs, - void* b, int64_t bo, int64_t brs, int64_t bcs, - void* out, int64_t oo, int64_t m, int64_t n, - float scale, float offset) { - auto& c = context::get(); - if (!c.ready) return false; - c.device_read_(a); - c.device_read_(b); - c.device_write_(out); - float* pa = context::off_(a, ao); - float* pb = context::off_(b, bo); - float* po = context::off_(out, oo); - unsigned um = static_cast(m), un = static_cast(n); - return c.launch1d_(c.fn_(op), um * un, pa, pb, po, um, un, - static_cast(ars), static_cast(acs), - static_cast(brs), static_cast(bcs), - scale, offset); -} - -inline bool unary(kop op, void* a, int64_t ao, void* out, int64_t oo, int64_t n, - float scale, float offset) { - auto& c = context::get(); - if (!c.ready) return false; - c.device_read_(a); - c.device_write_(out); - float* pa = context::off_(a, ao); - float* po = context::off_(out, oo); - unsigned un = static_cast(n); - return c.launch1d_(c.fn_(op), un, pa, po, un, scale, offset); + return c.dispatch_(c.fn_(k), args, n, params, params_bytes, g); } // Rank cap shared with the kernel side (tensorlib_cuda.cu's @@ -1231,26 +1272,25 @@ inline const long long* upload_bcast_meta_( // BatchNorm-shaped [N,D] input). `a_strides`/`b_strides` are the broadcast // strides (0 on a broadcast axis) array.h's gpu_binary_nd_ computes via the // same broadcast_strides() the CPU oracle uses. -inline bool binary_bcast_nd(kop op, void* a_native, int64_t ao, - const int64_t* a_strides, void* b_native, - int64_t bo, const int64_t* b_strides, - void* out_native, int64_t oo, - const int64_t* out_shape, int rank, int64_t n, - float scale, float offset) { +inline bool own::binary_bcast_nd(kop op, gpu::span a, const int64_t* a_strides, + gpu::span b, const int64_t* b_strides, + gpu::span out, const int64_t* out_shape, + int rank, int64_t n, float scale, + float offset) { if (rank <= 0 || rank > kPadFoldMaxRank) return false; auto& c = context::get(); if (!c.ready) return false; CUfunction f = c.bcast_nd_(op); if (!f) return false; - c.device_read_(a_native); - c.device_read_(b_native); - c.device_write_(out_native); + c.device_read_(a.buf); + c.device_read_(b.buf); + c.device_write_(out.buf); const long long* pmeta = upload_bcast_meta_(c, out_shape, rank, {a_strides, b_strides}); if (!pmeta) return false; - float* pa = context::off_(a_native, ao); - float* pb = context::off_(b_native, bo); - float* po = context::off_(out_native, oo); + float* pa = context::off_(a.buf, a.off); + float* pb = context::off_(b.buf, b.off); + float* po = context::off_(out.buf, out.off); unsigned un = static_cast(n); return c.launch1d_(f, un, pa, pb, po, pmeta, rank, un, scale, offset); } @@ -1259,25 +1299,24 @@ inline bool binary_bcast_nd(kop op, void* a_native, int64_t ao, // on no backend before this (eval_one's where_ case always ran the CPU // map_ternary). Same flat-index decode as binary_bcast_nd above, one more // operand -- masking (attention/padding masks) is the concrete caller. -inline bool where_nd(void* cond_native, int64_t co, const int64_t* c_strides, - void* a_native, int64_t ao, const int64_t* a_strides, - void* b_native, int64_t bo, const int64_t* b_strides, - void* out_native, int64_t oo, const int64_t* out_shape, - int rank, int64_t n) { +inline bool own::where_nd(gpu::span cond, const int64_t* c_strides, gpu::span a, + const int64_t* a_strides, gpu::span b, + const int64_t* b_strides, gpu::span out, + const int64_t* out_shape, int rank, int64_t n) { if (rank <= 0 || rank > kPadFoldMaxRank) return false; auto& c = context::get(); if (!c.ready) return false; - c.device_read_(cond_native); - c.device_read_(a_native); - c.device_read_(b_native); - c.device_write_(out_native); + c.device_read_(cond.buf); + c.device_read_(a.buf); + c.device_read_(b.buf); + c.device_write_(out.buf); const long long* pmeta = upload_bcast_meta_( c, out_shape, rank, {c_strides, a_strides, b_strides}); if (!pmeta) return false; - float* pc = context::off_(cond_native, co); - float* pa = context::off_(a_native, ao); - float* pb = context::off_(b_native, bo); - float* po = context::off_(out_native, oo); + float* pc = context::off_(cond.buf, cond.off); + float* pa = context::off_(a.buf, a.off); + float* pb = context::off_(b.buf, b.off); + float* po = context::off_(out.buf, out.off); unsigned un = static_cast(n); return c.launch1d_(c.where_nd_(), un, pc, pa, pb, po, pmeta, rank, un); } @@ -1285,18 +1324,17 @@ inline bool where_nd(void* cond_native, int64_t co, const int64_t* c_strides, // N-D strided copy: clone()'s device arm for a view the flat one-input // kernels cannot read (a permute, a transpose). Same flat-index decode and // meta upload as where_nd above, one operand. -inline bool copy_nd(void* a_native, int64_t ao, const int64_t* a_strides, - void* out_native, int64_t oo, const int64_t* out_shape, - int rank, int64_t n) { +inline bool own::copy_nd(gpu::span a, const int64_t* a_strides, gpu::span out, + const int64_t* out_shape, int rank, int64_t n) { if (rank <= 0 || rank > kPadFoldMaxRank) return false; auto& c = context::get(); if (!c.ready) return false; - c.device_read_(a_native); - c.device_write_(out_native); + c.device_read_(a.buf); + c.device_write_(out.buf); const long long* pmeta = upload_bcast_meta_(c, out_shape, rank, {a_strides}); if (!pmeta) return false; - float* pa = context::off_(a_native, ao); - float* po = context::off_(out_native, oo); + float* pa = context::off_(a.buf, a.off); + float* po = context::off_(out.buf, out.off); unsigned un = static_cast(n); return c.launch1d_(c.copy_nd_(), un, pa, po, pmeta, rank, un); } @@ -1308,22 +1346,21 @@ inline bool copy_nd(void* a_native, int64_t ao, const int64_t* a_strides, // broadcast_strides(target, out_strides, a.shape()) -- 0 on every axis // being summed over. `reduced_n` is the product of a_shape over exactly // those zero-acc axes (1 if there are none). -inline bool sum_to(void* a_native, int64_t ao, const int64_t* a_shape, - const int64_t* a_strides, const int64_t* acc, int rank, - int64_t out_n, int64_t reduced_n, void* out_native, - int64_t oo) { +inline bool own::sum_to(gpu::span a, const int64_t* a_shape, + const int64_t* a_strides, const int64_t* acc, int rank, + int64_t out_n, int64_t reduced_n, gpu::span out) { if (rank <= 0 || rank > kPadFoldMaxRank) return false; auto& c = context::get(); if (!c.ready) return false; CUfunction f = c.sum_to_(); if (!f) return false; - c.device_read_(a_native); - c.device_write_(out_native); + c.device_read_(a.buf); + c.device_write_(out.buf); const long long* pmeta = upload_bcast_meta_(c, a_shape, rank, {a_strides, acc}); if (!pmeta) return false; - float* pa = context::off_(a_native, ao); - float* po = context::off_(out_native, oo); + float* pa = context::off_(a.buf, a.off); + float* po = context::off_(out.buf, out.off); unsigned un = static_cast(out_n); unsigned ured = static_cast(reduced_n); // A deep reduction (a bias gradient sums its column over every row) earns a @@ -1339,98 +1376,27 @@ inline bool sum_to(void* a_native, int64_t ao, const int64_t* a_shape, return c.launch1d_(f, un, pa, po, pmeta, rank, un, ured); } -// Elementwise comparison, same shape only (array.h's gpu_compare_ gates on -// that; ReLU/LeakyReLU/Clip's backward gate and the concrete Tensor.gt/... -// callers never need a broadcast form). Output is a F32 mask (1.0f/0.0f), -// matching the CPU oracle's own comparison ops. -inline bool compare(cmp_op op, void* a_native, int64_t ao, void* b_native, - int64_t bo, void* out_native, int64_t oo, int64_t n, - int64_t bstride) { - auto& c = context::get(); - if (!c.ready) return false; - CUfunction f = c.compare_(op); - if (!f) return false; - c.device_read_(a_native); - c.device_read_(b_native); - c.device_write_(out_native); - float* pa = context::off_(a_native, ao); - float* pb = context::off_(b_native, bo); - float* po = context::off_(out_native, oo); - unsigned un = static_cast(n); - unsigned ubs = static_cast(bstride); - return c.launch1d_(f, un, pa, pb, po, un, ubs); -} - -// tanh_/sin_/cos_: plain elementwise, same shape as tl_exp/tl_sqrt (scale/ -// offset epilogue included, same reason those have it). -inline bool unary_ext(unary_ext_op op, void* a_native, int64_t ao, - void* out_native, int64_t oo, int64_t n, float scale, - float offset) { - auto& c = context::get(); - if (!c.ready) return false; - CUfunction f = c.unary_ext_(op); - if (!f) return false; - c.device_read_(a_native); - c.device_write_(out_native); - float* pa = context::off_(a_native, ao); - float* po = context::off_(out_native, oo); - unsigned un = static_cast(n); - return c.launch1d_(f, un, pa, po, un, scale, offset); -} - -// clamp(x, lo, hi): Clip's forward. No epilogue -- lo/hi occupy the role -// scale/offset play elsewhere, and nothing composes a further affine onto -// it today (array.h's gpu_clamp_ doesn't thread one through). -inline bool clamp(void* a_native, int64_t ao, void* out_native, int64_t oo, - int64_t n, float lo, float hi) { - auto& c = context::get(); - if (!c.ready) return false; - CUfunction f = c.clamp_(); - if (!f) return false; - c.device_read_(a_native); - c.device_write_(out_native); - float* pa = context::off_(a_native, ao); - float* po = context::off_(out_native, oo); - unsigned un = static_cast(n); - return c.launch1d_(f, un, pa, po, un, lo, hi); -} - -// Tensor-scalar ops: out = f(a, s) * scale + offset (see metal.h's scalar_op). -inline bool scalar_binary(scalar_op op, void* a_native, int64_t ao, - void* out_native, int64_t oo, int64_t n, float s, - float scale, float offset) { - auto& c = context::get(); - if (!c.ready) return false; - CUfunction f = c.scalar_binary_(op); - if (!f) return false; - c.device_read_(a_native); - c.device_write_(out_native); - float* pa = context::off_(a_native, ao); - float* po = context::off_(out_native, oo); - unsigned un = static_cast(n); - return c.launch1d_(f, un, pa, po, un, s, scale, offset); -} - // Places `a` (contiguous) into a zero buffer of out_shape (array.h's // gpu_pad_ allocates `out` uninitialized via array::empty — this zeros the // device copy directly, no host round trip), shifted by `before` along // `axis`. No scale/offset — eval_one's shared epilogue applies those (see // array.h's op_t::pad_ case). -inline bool pad(void* a_native, int64_t ao, void* out_native, int64_t oo, - const int64_t* a_shape, const int64_t* out_shape, int rank, - int axis, int64_t before, int64_t n, int64_t out_n) { +inline bool own::pad(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t before, int64_t n, int64_t out_n) { if (rank <= 0 || rank > kPadFoldMaxRank) return false; auto& c = context::get(); if (!c.ready) return false; - c.device_read_(a_native); - c.device_write_(out_native); - zero_device_(reinterpret_cast(out_native), out_n); + c.device_read_(a.buf); + c.device_write_(out.buf); + zero_device_(reinterpret_cast(context::off_(out.buf, out.off)), + out_n); int64_t out_strides[kPadFoldMaxRank]; const long long* pmeta = upload_pad_fold_meta_(c, a_shape, rank, out_shape, rank, out_strides); if (!pmeta) return false; - float* pa = context::off_(a_native, ao); - float* po = context::off_(out_native, oo); + float* pa = context::off_(a.buf, a.off); + float* po = context::off_(out.buf, out.off); unsigned un = static_cast(n); unsigned ushift = static_cast(before * out_strides[axis]); return c.launch1d_(c.pad_(), un, pa, po, pmeta, rank, ushift, un); @@ -1439,21 +1405,22 @@ inline bool pad(void* a_native, int64_t ao, void* out_native, int64_t oo, // unfold's inverse: scatter-add `a` (contiguous; its last dim is the sliding // window) into a zero buffer of out_shape (zeroed the same way as pad() // above) — every overlap accumulates via atomicAdd, so it must start at 0. -inline bool fold(void* a_native, int64_t ao, void* out_native, int64_t oo, - const int64_t* a_shape, const int64_t* out_shape, int rank, - int axis, int64_t step, int64_t n, int64_t out_n) { +inline bool own::fold(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t step, int64_t n, int64_t out_n) { if (rank <= 0 || rank > kPadFoldMaxRank) return false; auto& c = context::get(); if (!c.ready) return false; - c.device_read_(a_native); - c.device_write_(out_native); - zero_device_(reinterpret_cast(out_native), out_n); + c.device_read_(a.buf); + c.device_write_(out.buf); + zero_device_(reinterpret_cast(context::off_(out.buf, out.off)), + out_n); int64_t out_strides[kPadFoldMaxRank]; const long long* pmeta = upload_pad_fold_meta_(c, a_shape, rank, out_shape, rank - 1, out_strides); if (!pmeta) return false; - float* pa = context::off_(a_native, ao); - float* po = context::off_(out_native, oo); + float* pa = context::off_(a.buf, a.off); + float* po = context::off_(out.buf, out.off); unsigned un = static_cast(n); unsigned ustep = static_cast(step); return c.launch1d_(c.fold_(), un, pa, po, pmeta, rank, axis, ustep, un); @@ -1465,61 +1432,40 @@ inline bool fold(void* a_native, int64_t ao, void* out_native, int64_t oo, // padding border to zero, unlike pad() above). Reuses pad's own kernel: // writing a same-shape source at an axis-shifted offset is exactly what // tl_pad already does per source element. -inline bool concat_part(void* a_native, int64_t ao, void* out_native, - int64_t oo, const int64_t* a_shape, - const int64_t* out_shape, int rank, int axis, - int64_t before, int64_t n) { +inline bool own::concat_part(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t before, int64_t n) { if (rank <= 0 || rank > kPadFoldMaxRank) return false; auto& c = context::get(); if (!c.ready) return false; - c.device_read_(a_native); - c.device_write_(out_native); + c.device_read_(a.buf); + c.device_write_(out.buf); int64_t out_strides[kPadFoldMaxRank]; const long long* pmeta = upload_pad_fold_meta_(c, a_shape, rank, out_shape, rank, out_strides); if (!pmeta) return false; - float* pa = context::off_(a_native, ao); - float* po = context::off_(out_native, oo); + float* pa = context::off_(a.buf, a.off); + float* po = context::off_(out.buf, out.off); unsigned un = static_cast(n); unsigned ushift = static_cast(before * out_strides[axis]); return c.launch1d_(c.pad_(), un, pa, po, pmeta, rank, ushift, un); } -// Row gather along axis 0: out[i] = a[indices[i]] (a, indices contiguous; -// indices float-valued, rounded on-device to match argmax's own -// convention). One thread per output element, no write conflicts — no -// zeroing needed (every element is written exactly once). -inline bool index_select(void* a_native, int64_t ao, void* idx_native, - int64_t idxo, void* out_native, int64_t oo, - int64_t row_size, int64_t k) { - auto& c = context::get(); - if (!c.ready) return false; - c.device_read_(a_native); - c.device_read_(idx_native); - c.device_write_(out_native); - float* pa = context::off_(a_native, ao); - float* pidx = context::off_(idx_native, idxo); - float* po = context::off_(out_native, oo); - unsigned un = static_cast(k * row_size); - unsigned urow = static_cast(row_size); - return c.launch1d_(c.index_select_(), un, pa, pidx, po, urow, un); -} - // index_select's dual: scatter-add `values` into `out` by row index. // Repeated indices really do collide (real write conflicts — the kernel // uses atomicAdd), so `out` must start zeroed, same as pad/fold above. -inline bool index_add(void* idx_native, int64_t idxo, void* values_native, - int64_t vo, void* out_native, int64_t oo, - int64_t row_size, int64_t k, int64_t out_n) { +inline bool own::index_add(gpu::span idx, gpu::span values, gpu::span out, + int64_t row_size, int64_t k, int64_t out_n) { auto& c = context::get(); if (!c.ready) return false; - c.device_read_(idx_native); - c.device_read_(values_native); - c.device_write_(out_native); - zero_device_(reinterpret_cast(out_native), out_n); - float* pidx = context::off_(idx_native, idxo); - float* pv = context::off_(values_native, vo); - float* po = context::off_(out_native, oo); + c.device_read_(idx.buf); + c.device_read_(values.buf); + c.device_write_(out.buf); + zero_device_(reinterpret_cast(context::off_(out.buf, out.off)), + out_n); + float* pidx = context::off_(idx.buf, idx.off); + float* pv = context::off_(values.buf, values.off); + float* po = context::off_(out.buf, out.off); unsigned un = static_cast(k * row_size); unsigned urow = static_cast(row_size); return c.launch1d_(c.index_add_(), un, pidx, pv, po, urow, un); @@ -1530,104 +1476,23 @@ inline bool index_add(void* idx_native, int64_t idxo, void* values_native, // to a distinct output slot (the axis is brand new), so — unlike // index_add above — there is no accumulation and no atomics; `out` still // starts zeroed since untouched slots must read back as 0. -inline bool scatter_to_axis(void* idx_native, int64_t idxo, - void* values_native, int64_t vo, void* out_native, - int64_t oo, int64_t n, int64_t size) { +inline bool own::scatter_to_axis(gpu::span idx, gpu::span values, gpu::span out, + int64_t n, int64_t size) { auto& c = context::get(); if (!c.ready) return false; - c.device_read_(idx_native); - c.device_read_(values_native); - c.device_write_(out_native); - zero_device_(reinterpret_cast(out_native), n * size); - float* pidx = context::off_(idx_native, idxo); - float* pv = context::off_(values_native, vo); - float* po = context::off_(out_native, oo); + c.device_read_(idx.buf); + c.device_read_(values.buf); + c.device_write_(out.buf); + zero_device_(reinterpret_cast(context::off_(out.buf, out.off)), + n * size); + float* pidx = context::off_(idx.buf, idx.off); + float* pv = context::off_(values.buf, values.off); + float* po = context::off_(out.buf, out.off); unsigned un = static_cast(n); unsigned usize = static_cast(size); return c.launch1d_(c.scatter_axis_(), un, pidx, pv, po, usize, un); } -// scatter_to_axis's dual: out[i] = src[i * size + indices[i]], taking the one -// element each position labels out of the trailing axis. One thread per -// output, so — like index_select — no conflicts and nothing to pre-zero. -inline bool gather_from_axis(void* src_native, int64_t so, void* idx_native, - int64_t idxo, void* out_native, int64_t oo, - int64_t n, int64_t size) { - auto& c = context::get(); - if (!c.ready) return false; - c.device_read_(src_native); - c.device_read_(idx_native); - c.device_write_(out_native); - float* ps = context::off_(src_native, so); - float* pidx = context::off_(idx_native, idxo); - float* po = context::off_(out_native, oo); - unsigned un = static_cast(n); - unsigned usize = static_cast(size); - return c.launch1d_(c.gather_axis_(), un, ps, pidx, po, usize, un); -} - -// Row logsumexp over the last axis: log(sum exp) per row, affine epilogue. -// One block per row like row_op, but each thread carries a running (max, sum) -// pair through the tree, so it takes two shared floats per thread instead of -// one — and reads the row once where a max pass plus a sum pass reads it twice. -inline bool row_logsumexp(void* in, int64_t io, void* out, int64_t oo, - int64_t rows, int64_t cols, float scale, - float offset) { - auto& c = context::get(); - if (!c.ready) return false; - c.device_read_(in); - c.device_write_(out); - float* pin = context::off_(in, io); - float* po = context::off_(out, oo); - unsigned ur = (unsigned)rows, uc = (unsigned)cols; - unsigned block = 256; - return c.launch_(c.row_logsumexp_(), {ur ? ur : 1}, {block}, - 2 * block * sizeof(float), pin, po, ur, uc, scale, offset); -} - -// Softmax cross-entropy's pullback, from the forward's row logsumexp: -// out[i,j] = g[i] · (exp(x[i,j] - lse[i]) - [j == targets[i]]). x and out are -// [rows, cols]; lse, targets and g are one value per row. -inline bool xent_bwd(void* x, int64_t xo, void* lse, int64_t lo, void* tgt, - int64_t to, void* g, int64_t go, void* out, int64_t oo, - int64_t rows, int64_t cols) { - auto& c = context::get(); - if (!c.ready) return false; - c.device_read_(x); - c.device_read_(lse); - c.device_read_(tgt); - c.device_read_(g); - c.device_write_(out); - float* px = context::off_(x, xo); - float* pl = context::off_(lse, lo); - float* pt = context::off_(tgt, to); - float* pg = context::off_(g, go); - float* po = context::off_(out, oo); - unsigned un = static_cast(rows * cols); - unsigned uc = static_cast(cols); - return c.launch1d_(c.xent_bwd_(), un, px, pl, pt, pg, po, uc, un); -} - -// Adam's per-parameter update in place: m and v advance, p moves by the -// bias-corrected ratio. p, m, v and g are one shape, contiguous; the host -// folds the bias correction into lr_over_bc1 = lr/bc1 and inv_bc2 = 1/bc2. -inline bool adam_step(void* p, int64_t po, void* m, int64_t mo, void* v, - int64_t vo, void* g, int64_t go, int64_t n, float beta1, - float beta2, float eps, float lr_over_bc1, - float inv_bc2) { - auto& c = context::get(); - if (!c.ready || n <= 0) return false; - c.device_read_(g); - c.device_rmw_(p); // read and written: an optimizer's state is host-born - c.device_rmw_(m); - c.device_rmw_(v); - unsigned un = static_cast(n); - return c.launch1d_(c.adam_step_(), un, context::off_(p, po), - context::off_(m, mo), context::off_(v, vo), - context::off_(g, go), beta1, beta2, eps, lr_over_bc1, - inv_bc2, un); -} - // M7 decode GEMV: y(n) = a(1,k) @ B(k,n), F32 accumulate. B is either f32 or // bf16 weights (bf16 halves the dominant K×N weight traffic — the decode // bandwidth lever). Buffers are opaque device pointers; the kernel interprets @@ -1638,8 +1503,7 @@ inline bool adam_step(void* p, int64_t po, void* m, int64_t mo, void* v, // partition K over gridDim.y, atomicAdd into a pre-zeroed y, so the kernel stays // bandwidth-bound rather than occupancy-bound. gridDim.y==1 stores directly. inline bool gemv_run_(CUfunction f, float* pa, float* pB, float* py, - void* y_native, unsigned un, unsigned uk, - unsigned vcols = 1) { + unsigned un, unsigned uk, unsigned vcols = 1) { auto& c = context::get(); unsigned per = 256u * vcols; // output columns covered by one block unsigned bx = (un + per - 1) / per; @@ -1661,35 +1525,37 @@ inline bool gemv_run_(CUfunction f, float* pa, float* pB, float* py, // Zero y for the split-K atomicAdd. Async on the stream (ordered before the // gemv on the same stream) so this stays capturable — a blocking MemsetD8 // is illegal mid CUDA-graph capture. - CUdeviceptr yd = reinterpret_cast(y_native); + CUdeviceptr yd = reinterpret_cast(py); if (c.d.MemsetD8Async) c.d.MemsetD8Async(yd, 0, (size_t)un * 4, c.stream); else c.d.MemsetD8(yd, 0, (size_t)un * 4); } return c.launch_(f, {bx, gy}, {256}, 0, pa, pB, py, un, uk, ksplit); } -inline bool gemv_f32(void* a, void* B, void* y, int64_t n, int64_t k) { +inline bool own::gemv_f32(gpu::span a, gpu::span B, gpu::span y, int64_t n, + int64_t k) { auto& c = context::get(); if (!c.ready) return false; - c.device_read_(a); - c.device_read_(B); - c.device_write_(y); - return gemv_run_(c.gemv_f32_(), context::off_(a, 0), context::off_(B, 0), - context::off_(y, 0), y, static_cast(n), - static_cast(k)); -} -inline bool gemv_bf16(void* a, void* B, void* y, int64_t n, int64_t k) { + c.device_read_(a.buf); + c.device_read_(B.buf); + c.device_write_(y.buf); + return gemv_run_(c.gemv_f32_(), context::off_(a.buf, a.off), + context::off_(B.buf, B.off), context::off_(y.buf, y.off), + static_cast(n), static_cast(k)); +} +inline bool own::gemv_bf16(gpu::span a, gpu::span B, gpu::span y, int64_t n, + int64_t k) { auto& c = context::get(); if (!c.ready) return false; - c.device_read_(a); - c.device_read_(B); // B reinterpreted as __nv_bfloat16* in-kernel - c.device_write_(y); + c.device_read_(a.buf); + c.device_read_(B.buf); // B reinterpreted as __nv_bfloat16* in-kernel + c.device_write_(y.buf); // Vectorized 8-cols/thread path when n%8==0 (all transformer dims) — 16-byte // bf16 loads close the bandwidth gap to f32; scalar fallback otherwise. bool v8 = (n % 8) == 0; - return gemv_run_(v8 ? c.gemv_bf16v8_() : c.gemv_bf16_(), context::off_(a, 0), - context::off_(B, 0), context::off_(y, 0), y, - static_cast(n), static_cast(k), - v8 ? 8u : 1u); + return gemv_run_(v8 ? c.gemv_bf16v8_() : c.gemv_bf16_(), + context::off_(a.buf, a.off), context::off_(B.buf, B.off), + context::off_(y.buf, y.off), static_cast(n), + static_cast(k), v8 ? 8u : 1u); } // Block size (32..256 threads) for the one-block-per-row GEMVs (tl_gemv_bf16_row, @@ -1713,39 +1579,19 @@ inline unsigned gemv_row_smem(unsigned block) { return block > 32 ? (block >> 5) * (unsigned)sizeof(float) : 0u; } -// Warp-per-row bf16 decode GEMV (lever A): y(N) = a(1,K) @ W[N,K], W row-major -// (K contiguous per output row). ONE BLOCK per output row (grid.x == N), no -// split-K — no memset, no atomic combine. The small-N floor-bound lever; see -// tl_gemv_bf16_row. Requires K % 8 == 0 (host-gated; caller falls back to the -// split-K [K,N] path otherwise). -inline bool gemv_bf16_row(void* a, void* B, void* y, int64_t n, int64_t k) { - auto& c = context::get(); - if (!c.ready || (k % 8) != 0) return false; - c.device_read_(a); - c.device_read_(B); // B reinterpreted as __nv_bfloat16* [N][K] in-kernel - c.device_write_(y); - float* pa = context::off_(a, 0); - float* pB = context::off_(B, 0); - float* py = context::off_(y, 0); - unsigned uN = static_cast(n), uK = static_cast(k); - unsigned block = gemv_row_block_size(k); - return c.launch_(c.gemv_bf16_row_(), {uN}, {block}, gemv_row_smem(block), pa, - pB, py, uN, uK); -} - // M9 batched-prefill GEMM: C(M,N) = A(M,K) @ W[N,K]^T, W the same row-major // bf16 weight gemv_bf16_row consumes — so a batched prompt reuses the decode // weights as-is. Requires K % 8 == 0. See tl_gemm_bf16_nt. -inline bool gemm_bf16_nt(void* a, void* B, void* out, int64_t m, int64_t n, - int64_t k) { +inline bool own::gemm_bf16_nt(gpu::span a, gpu::span B, gpu::span out, + int64_t m, int64_t n, int64_t k) { auto& c = context::get(); if (!c.ready || (k % 8) != 0 || m <= 0 || n <= 0) return false; - c.device_read_(a); - c.device_read_(B); // B reinterpreted as __nv_bfloat16* [N][K] in-kernel - c.device_write_(out); - float* pa = context::off_(a, 0); - float* pB = context::off_(B, 0); - float* po = context::off_(out, 0); + c.device_read_(a.buf); + c.device_read_(B.buf); // B reinterpreted as __nv_bfloat16* [N][K] in-kernel + c.device_write_(out.buf); + float* pa = context::off_(a.buf, a.off); + float* pB = context::off_(B.buf, B.off); + float* po = context::off_(out.buf, out.off); unsigned uM = (unsigned)m, uN = (unsigned)n, uK = (unsigned)k; // Tile choice by fill (big_tile_): a prefill chunk is only a few hundred // tokens, so a narrow projection (N=896, M=512) is 28 big blocks against 82 @@ -1786,60 +1632,6 @@ inline bool gemm_bf16_nt(void* a, void* B, void* out, int64_t m, int64_t n, // ---- M9 batched-prefill layout moves between token-major projections and // head-major attention. -// [T, ld] token-major -> [H, T, D] head-major, adding an optional [H*D] bias. -// `off` picks a column block of a fused projection output (q|k|v from one GEMM). -inline bool split_heads(void* src, void* bias, void* dst, int64_t T, int64_t ld, - int64_t off, int64_t H, int64_t D) { - auto& c = context::get(); - if (!c.ready || T <= 0 || H <= 0 || D <= 0) return false; - c.device_read_(src); - if (bias) c.device_read_(bias); - c.device_write_(dst); - float* ps = context::off_(src, 0); - float* pb = bias ? context::off_(bias, 0) : nullptr; - float* pd = context::off_(dst, 0); - unsigned uT = (unsigned)T, uld = (unsigned)ld, uoff = (unsigned)off, - uD = (unsigned)D; - return c.launch_(c.split_heads_(), {(unsigned)H, uT}, {uD}, 0, ps, pb, pd, uT, - uld, uoff, uD); -} - -// [H, T, D] head-major -> [T, H*D] token-major (inverse of split_heads). -inline bool merge_heads(void* src, void* dst, int64_t T, int64_t H, int64_t D) { - auto& c = context::get(); - if (!c.ready || T <= 0 || H <= 0 || D <= 0) return false; - c.device_read_(src); - c.device_write_(dst); - float* ps = context::off_(src, 0); - float* pd = context::off_(dst, 0); - unsigned uT = (unsigned)T, uH = (unsigned)H, uD = (unsigned)D; - return c.launch_(c.merge_heads_(), {uH, uT}, {uD}, 0, ps, pd, uT, uH, uD); -} - -// M8 int4-weight decode GEMV: y(N) = a(1,K) @ dequant(Wq[N,K]), F32 accumulate. -// qw = packed int4 [N][K/8] words, scales = f32 [N][K/group]. ONE BLOCK per -// output row (grid.x == N), K-adaptive block size — see gemv_bf16_row. -// K % group == 0, group % 8 == 0 (host-gated); the kernel's per-thread tail -// guard lifts the old K % 256 requirement (Qwen K=896 = 3×256+128 works). -inline bool gemv_q4(void* a, void* qw, void* scales, void* y, int64_t N, - int64_t K, int64_t group) { - auto& c = context::get(); - if (!c.ready || group <= 0 || (K % group) != 0 || (group % 8) != 0) - return false; - c.device_read_(a); - c.device_read_(qw); - c.device_read_(scales); - c.device_write_(y); - float* pa = context::off_(a, 0); - float* pq = context::off_(qw, 0); - float* ps = context::off_(scales, 0); - float* py = context::off_(y, 0); - unsigned uN = (unsigned)N, uK = (unsigned)K, uG = (unsigned)group; - unsigned block = gemv_row_block_size(K); - return c.launch_(c.gemv_q4_(), {uN}, {block}, gemv_row_smem(block), pa, pq, - ps, py, uN, uK, uG); -} - // Split-KV split-count heuristic: how many ctx-splits make grid = heads×S fill // the SMs (heads alone is ~32 blocks « 82 SMs; target ~4 blocks/SM, and each // split needs >=128 keys to amortize its fixed cost). Shared by attn_decode @@ -1906,20 +1698,21 @@ struct attn_partials { // one pass. q [n_q_heads,D], out [n_q_heads,D]; K/V are a [n_kv_heads,kv_max,D] // cache read over its valid prefix [0,ctx) (kv_max==ctx is the no-cache case). // GQA: q head h reads kv head h/(n_q_heads/n_kv_heads). Contiguous, D∈{64,128}. -inline bool attn_decode(void* q, void* K, void* V, void* out, int64_t n_q_heads, - int64_t n_kv_heads, int64_t ctx, int64_t kv_max, - int64_t D, float scale, bool kv_bf16 = false) { +inline bool own::attn_decode(gpu::span q, gpu::span K, gpu::span V, + gpu::span out, int64_t n_q_heads, + int64_t n_kv_heads, int64_t ctx, int64_t kv_max, + int64_t D, float scale, bool kv_bf16) { auto& c = context::get(); if (!c.ready || (D != 128 && D != 64)) return false; if (n_kv_heads <= 0 || n_q_heads % n_kv_heads != 0) return false; - c.device_read_(q); - c.device_read_(K); - c.device_read_(V); - c.device_write_(out); - float* pq = context::off_(q, 0); - float* pk = context::off_(K, 0); - float* pv = context::off_(V, 0); - float* po = context::off_(out, 0); + c.device_read_(q.buf); + c.device_read_(K.buf); + c.device_read_(V.buf); + c.device_write_(out.buf); + float* pq = context::off_(q.buf, q.off); + float* pk = context::off_(K.buf, K.off); + float* pv = context::off_(V.buf, V.off); + float* po = context::off_(out.buf, out.off); unsigned uh = static_cast(n_q_heads), uctx = static_cast(ctx); unsigned kv_stride = static_cast(kv_max * D); unsigned group = static_cast(n_q_heads / n_kv_heads); @@ -1953,48 +1746,49 @@ inline bool attn_decode(void* q, void* K, void* V, void* out, int64_t n_q_heads, // from the cache CAPACITY (max_ctx) and lets the kernel bound the work by pos. // RoPE reading pos from *d_pos (else identical to rope()). -inline bool rope_dpos(void* x, void* out, int64_t rows, int64_t T, int64_t D, - void* d_pos, float base, void* bias = nullptr) { +inline bool own::rope_dpos(gpu::span x, gpu::span out, int64_t rows, int64_t T, + int64_t D, gpu::span d_pos, float base, + gpu::span bias) { auto& c = context::get(); if (!c.ready || D <= 0 || (D & 1)) return false; - c.device_read_(x); - if (bias) c.device_read_(bias); - c.device_write_(out); - c.device_read_(d_pos); - float* px = context::off_(x, 0); - float* pbias = bias ? context::off_(bias, 0) : nullptr; - float* po = context::off_(out, 0); - float* pp = context::off_(d_pos, 0); + c.device_read_(x.buf); + if (bias.buf) c.device_read_(bias.buf); + c.device_write_(out.buf); + c.device_read_(d_pos.buf); + float* px = context::off_(x.buf, x.off); + float* pbias = bias.buf ? context::off_(bias.buf, bias.off) : nullptr; + float* po = context::off_(out.buf, out.off); + float* pp = context::off_(d_pos.buf, d_pos.off); unsigned uT = (unsigned)T, uD = (unsigned)D; return c.launch_(c.rope_dpos_(), {(unsigned)rows}, {(unsigned)(D / 2)}, 0, px, pbias, po, uT, uD, pp, base); } // One-thread *d_pos += 1 (tail of a captured forward; advances the counter). -inline bool incr_u32(void* d_pos) { +inline bool own::incr_u32(gpu::span d_pos) { auto& c = context::get(); if (!c.ready) return false; - c.device_write_(d_pos); - float* pp = context::off_(d_pos, 0); + c.device_write_(d_pos.buf); + float* pp = context::off_(d_pos.buf, d_pos.off); return c.launch_(c.incr_u32_(), {1}, {1}, 0, pp); } // KV append with write-row = *d_pos (else identical to kv_append(); f32 KV). -inline bool kv_append_dpos(void* Kc, void* Vc, void* k_new, void* v_new, - void* d_pos, int64_t kv_max, int64_t n_kv_heads, - int64_t D) { +inline bool own::kv_append_dpos(gpu::span Kc, gpu::span Vc, gpu::span k_new, + gpu::span v_new, gpu::span d_pos, + int64_t kv_max, int64_t n_kv_heads, int64_t D) { auto& c = context::get(); if (!c.ready || (D != 128 && D != 64)) return false; - c.device_read_(k_new); - c.device_read_(v_new); - c.device_read_(d_pos); - c.device_write_(Kc); - c.device_write_(Vc); - float* pKc = context::off_(Kc, 0); - float* pVc = context::off_(Vc, 0); - float* pk = context::off_(k_new, 0); - float* pv = context::off_(v_new, 0); - float* pp = context::off_(d_pos, 0); + c.device_read_(k_new.buf); + c.device_read_(v_new.buf); + c.device_read_(d_pos.buf); + c.device_write_(Kc.buf); + c.device_write_(Vc.buf); + float* pKc = context::off_(Kc.buf, Kc.off); + float* pVc = context::off_(Vc.buf, Vc.off); + float* pk = context::off_(k_new.buf, k_new.off); + float* pv = context::off_(v_new.buf, v_new.off); + float* pp = context::off_(d_pos.buf, d_pos.off); unsigned kv_stride = (unsigned)(kv_max * D); return c.launch_(c.kv_append_dpos_(), {(unsigned)n_kv_heads}, {(unsigned)D}, 0, pKc, pVc, pk, pv, pp, kv_stride); @@ -2012,29 +1806,30 @@ inline bool kv_append_dpos(void* Kc, void* Vc, void* k_new, void* v_new, // — like d_pos, it is graph-lifetime state (a captured graph bakes its address // in), so it must not be a shared growable scratch; kv_cache owns one per // cache. f32 KV. -inline bool attn_decode_dpos(void* q, void* K, void* V, void* out, - int64_t n_q_heads, int64_t n_kv_heads, void* d_pos, - int64_t kv_max, int64_t D, float scale, - void* partials) { +inline bool own::attn_decode_dpos(gpu::span q, gpu::span K, gpu::span V, + gpu::span out, int64_t n_q_heads, + int64_t n_kv_heads, gpu::span d_pos, + int64_t kv_max, int64_t D, float scale, + gpu::span partials) { auto& c = context::get(); - if (!c.ready || (D != 128 && D != 64) || !partials) return false; + if (!c.ready || (D != 128 && D != 64) || !partials.buf) return false; if (n_kv_heads <= 0 || n_q_heads % n_kv_heads != 0) return false; - c.device_read_(q); - c.device_read_(K); - c.device_read_(V); - c.device_read_(d_pos); - c.device_write_(out); - c.device_write_(partials); - float* pq = context::off_(q, 0); - float* pk = context::off_(K, 0); - float* pv = context::off_(V, 0); - float* po = context::off_(out, 0); - float* pp = context::off_(d_pos, 0); + c.device_read_(q.buf); + c.device_read_(K.buf); + c.device_read_(V.buf); + c.device_read_(d_pos.buf); + c.device_write_(out.buf); + c.device_write_(partials.buf); + float* pq = context::off_(q.buf, q.off); + float* pk = context::off_(K.buf, K.off); + float* pv = context::off_(V.buf, V.off); + float* po = context::off_(out.buf, out.off); + float* pp = context::off_(d_pos.buf, d_pos.off); unsigned uh = (unsigned)n_q_heads, uD = (unsigned)D; unsigned kv_stride = (unsigned)(kv_max * D); unsigned group = (unsigned)(n_q_heads / n_kv_heads); unsigned S = attn_split_count(uh, kv_max); - attn_partials p(context::off_(partials, 0), (size_t)uh * S); + attn_partials p(context::off_(partials.buf, partials.off), (size_t)uh * S); if (!c.launch_(c.attn_split_dpos_(D), {uh, S}, {uD}, 0, pq, pk, pv, p.pm, p.pl, p.pacc, pp, kv_stride, group, scale)) { return false; @@ -2042,72 +1837,28 @@ inline bool attn_decode_dpos(void* q, void* K, void* V, void* out, return p.combine(c, po, uh, uD, S); } -// M9 KV cache append: scatter one decode step's k,v (each [n_kv_heads,D] device -// buffers) into the cache (K,V each [n_kv_heads,kv_max,D]) at row `pos`. -inline bool kv_append(void* Kc, void* Vc, void* k_new, void* v_new, int64_t pos, - int64_t kv_max, int64_t n_kv_heads, int64_t D, - bool kv_bf16 = false) { - auto& c = context::get(); - if (!c.ready || (D != 128 && D != 64)) return false; - c.device_read_(k_new); - c.device_read_(v_new); - c.device_write_(Kc); - c.device_write_(Vc); - float* pKc = context::off_(Kc, 0); - float* pVc = context::off_(Vc, 0); - float* pk = context::off_(k_new, 0); - float* pv = context::off_(v_new, 0); - unsigned upos = static_cast(pos); - unsigned kv_stride = static_cast(kv_max * D); - return c.launch_(c.kv_append_(kv_bf16), {static_cast(n_kv_heads)}, - {static_cast(D)}, 0, pKc, pVc, pk, pv, upos, - kv_stride); -} - -// M9 prefill: bulk-copy a block of k,v (each [n_kv_heads,T,D] device buffers) -// into the cache (K,V each [n_kv_heads,kv_max,D]) rows [pos0, pos0+T), so a long -// prompt can be filled in chunks. grid=(n_kv_heads,T). -inline bool kv_fill(void* Kc, void* Vc, void* K, void* V, int64_t T, - int64_t kv_max, int64_t n_kv_heads, int64_t D, - bool kv_bf16 = false, int64_t pos0 = 0) { - auto& c = context::get(); - if (!c.ready || (D != 128 && D != 64)) return false; - c.device_read_(K); - c.device_read_(V); - c.device_write_(Kc); - c.device_write_(Vc); - float* pKc = context::off_(Kc, 0); - float* pVc = context::off_(Vc, 0); - float* pk = context::off_(K, 0); - float* pv = context::off_(V, 0); - unsigned uT = static_cast(T); - unsigned kv_stride = static_cast(kv_max * D); - unsigned up0 = static_cast(pos0); - return c.launch_(c.kv_fill_(kv_bf16), {static_cast(n_kv_heads), uT}, - {static_cast(D)}, 0, pKc, pVc, pk, pv, uT, - kv_stride, up0); -} - // M9 causal prefill attention: q,out [n_q_heads,T,D]; K/V a [n_kv_heads,kv_max,D] // cache read over [0,pos0+T). Query p is at absolute position pos0+p and attends // keys 0..pos0+p, so a long prompt can be run in chunks (and a later turn // appended to a live cache). GQA via group. D∈{64,128}. // One block per (head, query tile); grid = (n_q_heads, ceil(T/tile)). -inline bool attn_prefill(void* q, void* K, void* V, void* out, int64_t n_q_heads, - int64_t n_kv_heads, int64_t T, int64_t kv_max, int64_t D, - float scale, bool kv_bf16 = false, int64_t pos0 = 0) { +inline bool own::attn_prefill(gpu::span q, gpu::span K, gpu::span V, + gpu::span out, int64_t n_q_heads, + int64_t n_kv_heads, int64_t T, int64_t kv_max, + int64_t D, float scale, bool kv_bf16, + int64_t pos0) { auto& c = context::get(); if (!c.ready || (D != 128 && D != 64)) return false; if (n_kv_heads <= 0 || n_q_heads % n_kv_heads != 0) return false; if (T <= 0 || T > 65535) return false; // one call is one prompt chunk - c.device_read_(q); - c.device_read_(K); - c.device_read_(V); - c.device_write_(out); - float* pq = context::off_(q, 0); - float* pk = context::off_(K, 0); - float* pv = context::off_(V, 0); - float* po = context::off_(out, 0); + c.device_read_(q.buf); + c.device_read_(K.buf); + c.device_read_(V.buf); + c.device_write_(out.buf); + float* pq = context::off_(q.buf, q.off); + float* pk = context::off_(K.buf, K.off); + float* pv = context::off_(V.buf, V.off); + float* po = context::off_(out.buf, out.off); unsigned uT = static_cast(T); unsigned kv_stride = static_cast(kv_max * D); unsigned group = static_cast(n_q_heads / n_kv_heads); @@ -2128,68 +1879,70 @@ inline bool attn_prefill(void* q, void* K, void* V, void* out, int64_t n_q_heads // all [H,T,D] contiguous — a training shape, so no KV cache, no GQA and no // chunked positions — and `stats` [2,H,T] the row logsumexp and dO·O the dK/dV // half reads. D∈{64,128}. One block per (head, query tile). -inline bool attn_prefill_dq(void* q, void* K, void* V, void* dO, void* O, - void* dq, void* stats, int64_t H, int64_t T, - int64_t D, float scale) { +inline bool own::attn_prefill_dq(gpu::span q, gpu::span K, gpu::span V, + gpu::span dO, gpu::span O, gpu::span dq, + gpu::span stats, int64_t H, int64_t T, + int64_t D, float scale) { auto& c = context::get(); if (!c.ready || (D != 128 && D != 64)) return false; if (H <= 0 || T <= 0 || T > 65535) return false; - c.device_read_(q); - c.device_read_(K); - c.device_read_(V); - c.device_read_(dO); - c.device_read_(O); - c.device_write_(dq); - c.device_write_(stats); + c.device_read_(q.buf); + c.device_read_(K.buf); + c.device_read_(V.buf); + c.device_read_(dO.buf); + c.device_read_(O.buf); + c.device_write_(dq.buf); + c.device_write_(stats.buf); unsigned uT = static_cast(T); const unsigned tile = attn_bwd_tile(D); return c.launch_(c.attn_bwd_dq_(D), {static_cast(H), (uT + tile - 1) / tile}, - {attn_tile_threads}, 0, context::off_(q, 0), - context::off_(K, 0), context::off_(V, 0), - context::off_(dO, 0), context::off_(O, 0), - context::off_(dq, 0), context::off_(stats, 0), uT, scale); + {attn_tile_threads}, 0, context::off_(q.buf, q.off), + context::off_(K.buf, K.off), context::off_(V.buf, V.off), + context::off_(dO.buf, dO.off), context::off_(O.buf, O.off), + context::off_(dq.buf, dq.off), context::off_(stats.buf, stats.off), uT, scale); } // The key/value half, reading the stats the call above wrote: q, K, V, dO, dK // and dV all [H,T,D] contiguous, `stats` [2,H,T]. One block per (head, key // tile), and the head count reaches the kernel as gridDim.x — it is what the // stats' plane stride is made of. -inline bool attn_prefill_dkv(void* q, void* K, void* V, void* dO, void* stats, - void* dK, void* dV, int64_t H, int64_t T, - int64_t D, float scale) { +inline bool own::attn_prefill_dkv(gpu::span q, gpu::span K, gpu::span V, + gpu::span dO, gpu::span stats, gpu::span dK, + gpu::span dV, int64_t H, int64_t T, int64_t D, + float scale) { auto& c = context::get(); if (!c.ready || (D != 128 && D != 64)) return false; if (H <= 0 || T <= 0 || T > 65535) return false; - c.device_read_(q); - c.device_read_(K); - c.device_read_(V); - c.device_read_(dO); - c.device_read_(stats); - c.device_write_(dK); - c.device_write_(dV); + c.device_read_(q.buf); + c.device_read_(K.buf); + c.device_read_(V.buf); + c.device_read_(dO.buf); + c.device_read_(stats.buf); + c.device_write_(dK.buf); + c.device_write_(dV.buf); unsigned uT = static_cast(T); const unsigned tile = attn_bwd_tile(D); return c.launch_(c.attn_bwd_dkv_(D), {static_cast(H), (uT + tile - 1) / tile}, - {attn_tile_threads}, 0, context::off_(q, 0), - context::off_(K, 0), context::off_(V, 0), - context::off_(dO, 0), context::off_(stats, 0), - context::off_(dK, 0), context::off_(dV, 0), uT, scale); + {attn_tile_threads}, 0, context::off_(q.buf, q.off), + context::off_(K.buf, K.off), context::off_(V.buf, V.off), + context::off_(dO.buf, dO.off), context::off_(stats.buf, stats.off), + context::off_(dK.buf, dK.off), context::off_(dV.buf, dV.off), uT, scale); } // RoPE: rotate a contiguous [rows, D] buffer (rows = H*T). Row r's position is // pos + (r % T); half-split (GPT-NeoX / HF-llama) convention. D must be even. -inline bool rope(void* x, void* out, int64_t rows, int64_t T, int64_t D, - int64_t pos, float base, void* bias = nullptr) { +inline bool own::rope(gpu::span x, gpu::span out, int64_t rows, int64_t T, + int64_t D, int64_t pos, float base, gpu::span bias) { auto& c = context::get(); if (!c.ready || D <= 0 || (D & 1)) return false; - c.device_read_(x); - if (bias) c.device_read_(bias); - c.device_write_(out); - float* px = context::off_(x, 0); - float* pbias = bias ? context::off_(bias, 0) : nullptr; - float* po = context::off_(out, 0); + c.device_read_(x.buf); + if (bias.buf) c.device_read_(bias.buf); + c.device_write_(out.buf); + float* px = context::off_(x.buf, x.off); + float* pbias = bias.buf ? context::off_(bias.buf, bias.off) : nullptr; + float* po = context::off_(out.buf, out.off); unsigned uT = static_cast(T), uD = static_cast(D), upos = static_cast(pos); return c.launch_(c.rope_(), {static_cast(rows)}, @@ -2202,89 +1955,6 @@ inline bool rope(void* x, void* out, int64_t rows, int64_t T, int64_t D, // what keeps the two paths from drifting numerically. Buffers are [rows, n] // contiguous; the weight is [n], shared by every row. -// xout = x + delta; hout = rmsnorm(xout) * w, per row. xout may alias x. Folds a -// layer's residual add into the following norm (the o-proj->norm and -// mlp->next-input-norm seams), writing both the residual sum (the next residual -// base) and its normalized form. -inline bool rmsnorm_res(void* x, void* delta, void* w, void* xout, void* hout, - int64_t n, float eps, int64_t rows = 1) { - auto& c = context::get(); - if (!c.ready || n <= 0 || rows <= 0) return false; - c.device_read_(x); - c.device_read_(delta); - c.device_read_(w); - c.device_write_(xout); - c.device_write_(hout); - float* pa = context::off_(x, 0); - float* pb = context::off_(delta, 0); - float* pw = context::off_(w, 0); - float* px = context::off_(xout, 0); - float* ph = context::off_(hout, 0); - unsigned un = (unsigned)n, block = 256; - return c.launch_(c.add_rmsnorm_(), {(unsigned)rows}, {block}, - block * sizeof(float), pa, pb, pw, px, ph, un, eps); -} - -// GPU argmax over a length-n device vector (contiguous, offset 0). Reduces on -// device and D2H's only the 4-byte index — replaces the per-token 608KB logits -// copy + host scan that greedy decoding otherwise pays. `in` is a native -// device-buffer handle (e.g. an evaluated logits array's native()); stream -// ordering means the prior gemv that filled it need not be host-synced first. -// Returns the argmax index, tie-broken to the smallest index (matches the -// host `v[i] > v[bi]` loop) so greedy output stays bit-identical. -inline bool argmax(void* in, int64_t n, int64_t* out_idx) { - auto& c = context::get(); - if (!c.ready || !in || n <= 0 || !out_idx) return false; - c.device_read_(in); - float* pin = context::off_(in, 0); - CUdeviceptr res = c.argmax_res_(); - if (!res) return false; - int* pres = reinterpret_cast(res); - unsigned un = static_cast(n); - unsigned block = 256; - if (!c.launch_(c.argmax_(), {1}, {block}, - block * (sizeof(float) + sizeof(int)), pin, pres, un)) { - return false; - } - flush(); // the result index must be ready before the 4-byte D2H - int h = 0; - if (c.d.MemcpyDtoH(&h, res, sizeof(int)) != 0) return false; - *out_idx = h; - return true; -} - -// Fused RMSNorm over one length-n row: out = x * 1/sqrt(mean(x^2)+eps) * w. -// x/w/out are native device handles (offset 0). One block; matches the array -// composition numerically (see tl_rmsnorm). In-place safe (out may alias x). -inline bool rmsnorm(void* x, void* w, void* out, int64_t n, float eps, - int64_t rows = 1) { - auto& c = context::get(); - if (!c.ready || n <= 0 || rows <= 0) return false; - c.device_read_(x); - c.device_read_(w); - c.device_write_(out); - float* px = context::off_(x, 0); - float* pw = context::off_(w, 0); - float* po = context::off_(out, 0); - unsigned un = (unsigned)n, block = 256; - return c.launch_(c.rmsnorm_(), {(unsigned)rows}, {block}, - block * sizeof(float), px, pw, po, un, eps); -} - -// out[rows, ff] = silu(gate) * up, read out of the FUSED gate|up buffer -// gu[rows, 2*ff] that both paths already produce (up is gate + ff in each row). -inline bool swiglu(void* gu, void* out, int64_t ff, int64_t rows = 1) { - auto& c = context::get(); - if (!c.ready || ff <= 0 || rows <= 0) return false; - c.device_read_(gu); - c.device_write_(out); - float* pg = context::off_(gu, 0); - float* po = context::off_(out, 0); - unsigned uff = (unsigned)ff, block = 256; - unsigned gx = (uff + block - 1) / block; - return c.launch_(c.swiglu_(), {gx, (unsigned)rows}, {block}, 0, pg, po, uff); -} - // The KV cache itself is tl::kv_cache (kv_cache.h), written once over the // gpu:: facade; its graph-capture forms call kv_append_dpos / attn_decode_dpos // above and size their partials here. @@ -2342,21 +2012,21 @@ inline sgemm_splitk sgemm_splitk_(long base_blocks, unsigned k, const sgemm_tile // plain GEMM and keeps the one-output-per-thread fallback (tl_sgemm) for the // layouts the fast path declines; batch > 1 has no fallback here and returns // false, so the caller loops per slice. -inline bool gemm_batched(void* a, int64_t ao, int64_t lda, bool ta, int64_t sa, - void* b, int64_t bo, int64_t ldb, bool tb, int64_t sb, - void* out, int64_t oo, int64_t m, int64_t n, int64_t k, - int64_t batch, float scale, float offset, - void* bias = nullptr, int64_t biaso = 0) { +inline bool own::gemm_batched(gpu::span a, int64_t lda, bool ta, int64_t sa, + gpu::span b, int64_t ldb, bool tb, int64_t sb, + gpu::span out, int64_t m, int64_t n, int64_t k, + int64_t batch, float scale, float offset, + gpu::span bias) { auto& c = context::get(); if (!c.ready || batch < 1) return false; - c.device_read_(a); - c.device_read_(b); - if (bias) c.device_read_(bias); - c.device_write_(out); - float* pa = context::off_(a, ao); - float* pb = context::off_(b, bo); - float* pbias = bias ? context::off_(bias, biaso) : nullptr; - float* po = context::off_(out, oo); + c.device_read_(a.buf); + c.device_read_(b.buf); + if (bias.buf) c.device_read_(bias.buf); + c.device_write_(out.buf); + float* pa = context::off_(a.buf, a.off); + float* pb = context::off_(b.buf, b.off); + float* pbias = bias.buf ? context::off_(bias.buf, bias.off) : nullptr; + float* po = context::off_(out.buf, out.off); unsigned um = (unsigned)m, un = (unsigned)n, uk = (unsigned)k; // Tiled fast path (tl_sgemm_cp*, one per tile per operand layout): each @@ -2368,7 +2038,7 @@ inline bool gemm_batched(void* a, int64_t ao, int64_t lda, bool ta, int64_t sa, // 16B-aligned: multiples of 4 floats (C is stored per element, so only its // m·n stride has to fit). M and N block edges are predicated in-kernel. // Strided views, odd K and unaligned offsets fall to tl_sgemm. - bool aligned = (ao % 16 == 0) && (bo % 16 == 0) && (oo % 16 == 0); + bool aligned = (a.off % 16 == 0) && (b.off % 16 == 0) && (out.off % 16 == 0); bool a_ok = ta ? (lda == m && m % 4 == 0) : (lda == k); bool b_ok = tb ? (ldb == k) : (ldb == n && n % 4 == 0); bool batch_ok = sa % 4 == 0 && sb % 4 == 0 && sa <= (int64_t)UINT32_MAX && @@ -2418,89 +2088,50 @@ inline bool gemm_batched(void* a, int64_t ao, int64_t lda, bool ta, int64_t sa, } // C(m,n) = (A @ B) * scale + offset: the batch == 1 case of gemm_batched. -inline bool gemm(void* a, int64_t ao, int64_t lda, bool ta, void* b, int64_t bo, - int64_t ldb, bool tb, void* out, int64_t oo, int64_t m, - int64_t n, int64_t k, float scale, float offset) { - return gemm_batched(a, ao, lda, ta, 0, b, bo, ldb, tb, 0, out, oo, m, n, k, 1, - scale, offset); +inline bool own::gemm(gpu::span a, int64_t lda, bool ta, gpu::span b, + int64_t ldb, bool tb, gpu::span out, int64_t m, int64_t n, + int64_t k, float scale, float offset) { + return gemm_batched(a, lda, ta, 0, b, ldb, tb, 0, out, m, n, k, 1, scale, + offset); } // C(m,n) = (A @ B) * scale + offset + bias[j]: addmm's shape, the row bias // added in the gemm's own store rather than by a second pass over C. -inline bool gemm_bias(void* a, int64_t ao, int64_t lda, bool ta, void* b, - int64_t bo, int64_t ldb, bool tb, void* bias, - int64_t biaso, void* out, int64_t oo, int64_t m, - int64_t n, int64_t k, float scale, float offset) { - return gemm_batched(a, ao, lda, ta, 0, b, bo, ldb, tb, 0, out, oo, m, n, k, 1, - scale, offset, bias, biaso); -} - -// Row op over the last axis: softmax writes rows×cols; row_sum/row_max write -// one value per row, affine epilogue. One block per row, 256 threads. -inline bool row_op(kop op, void* in, int64_t io, void* out, int64_t oo, - int64_t rows, int64_t cols, float scale, float offset) { - auto& c = context::get(); - if (!c.ready) return false; - c.device_read_(in); - c.device_write_(out); - float* pin = context::off_(in, io); - float* po = context::off_(out, oo); - unsigned ur = (unsigned)rows, uc = (unsigned)cols; - unsigned block = 256; - return c.launch_(c.fn_(op), {ur ? ur : 1}, {block}, block * sizeof(float), - pin, po, ur, uc, scale, offset); -} - -// Layer norm over the last axis: out = (x - mu) · 1/sqrt(var + eps) · g + b per -// row, affine epilogue; g and b are contiguous d-vectors. One block per row, -// 256 threads, like row_op. -inline bool layer_norm(void* x, int64_t xo, void* g, int64_t go, void* b, - int64_t bo, void* out, int64_t oo, int64_t rows, - int64_t cols, float eps, float scale, float offset) { - auto& c = context::get(); - if (!c.ready) return false; - c.device_read_(x); - c.device_read_(g); - c.device_read_(b); - c.device_write_(out); - float* px = context::off_(x, xo); - float* pg = context::off_(g, go); - float* pb = context::off_(b, bo); - float* po = context::off_(out, oo); - unsigned ur = (unsigned)rows, uc = (unsigned)cols; - unsigned block = 256; - return c.launch_(c.layer_norm_(), {ur ? ur : 1}, {block}, - block * sizeof(float), px, pg, pb, po, ur, uc, eps, scale, - offset); +inline bool own::gemm_bias(gpu::span a, int64_t lda, bool ta, gpu::span b, + int64_t ldb, bool tb, gpu::span bias, gpu::span out, + int64_t m, int64_t n, int64_t k, float scale, + float offset) { + return gemm_batched(a, lda, ta, 0, b, ldb, tb, 0, out, m, n, k, 1, scale, + offset, bias); } // Layer norm's pullback: dx [rows, cols], dg and db [cols] from x and dy // [rows, cols] and the d-vector g, all contiguous; dx/dg/db, `stats` [2, rows] // and `partials` [2, chunks, cols] are fresh buffers of the caller's, the rows // taken `per_chunk` at a time. A row kernel, a column-strip kernel, a fold. -inline bool layer_norm_bwd(void* x, int64_t xo, void* g, int64_t go, void* dy, - int64_t dyo, void* dx, void* dg, void* db, - void* stats, void* partials, int64_t rows, - int64_t cols, int64_t per_chunk, int64_t chunks, - float eps) { +inline bool own::layer_norm_bwd(gpu::span x, gpu::span g, gpu::span dy, + gpu::span dx, gpu::span dg, gpu::span db, + gpu::span stats, gpu::span partials, + int64_t rows, int64_t cols, int64_t per_chunk, + int64_t chunks, float eps) { auto& c = context::get(); if (!c.ready || rows <= 0 || cols <= 0 || chunks <= 0) return false; - c.device_read_(x); - c.device_read_(g); - c.device_read_(dy); - c.device_write_(dx); - c.device_write_(dg); - c.device_write_(db); - c.device_write_(stats); - c.device_write_(partials); - float* px = context::off_(x, xo); - float* pg = context::off_(g, go); - float* pdy = context::off_(dy, dyo); - float* pdx = context::off_(dx, 0); - float* pdg = context::off_(dg, 0); - float* pdb = context::off_(db, 0); - float* ps = context::off_(stats, 0); - float* pp = context::off_(partials, 0); + c.device_read_(x.buf); + c.device_read_(g.buf); + c.device_read_(dy.buf); + c.device_write_(dx.buf); + c.device_write_(dg.buf); + c.device_write_(db.buf); + c.device_write_(stats.buf); + c.device_write_(partials.buf); + float* px = context::off_(x.buf, x.off); + float* pg = context::off_(g.buf, g.off); + float* pdy = context::off_(dy.buf, dy.off); + float* pdx = context::off_(dx.buf, dx.off); + float* pdg = context::off_(dg.buf, dg.off); + float* pdb = context::off_(db.buf, db.off); + float* ps = context::off_(stats.buf, stats.off); + float* pp = context::off_(partials.buf, partials.off); unsigned ur = (unsigned)rows, uc = (unsigned)cols, uk = (unsigned)chunks; unsigned up = (unsigned)per_chunk; unsigned block = 256; @@ -2515,133 +2146,6 @@ inline bool layer_norm_bwd(void* x, int64_t xo, void* g, int64_t go, void* dy, return c.launch1d_(c.layer_norm_bwd_gb_fold_(), uc, pp, pdg, pdb, uk, uc); } -#else // stubs (Apple, or a build without TENSORLIB_CUDA) - -// Only the gpu:: facade surface is stubbed — what array.h/storage.h dispatch -// through, so a consumer can name tl::cuda:: unconditionally and get "no device -// here". The LLM-path entry points (gemv/attention/kv_cache/graph capture/the -// fused decode ops) deliberately have NO stubs: they are reachable only from -// code that is itself CUDA-gated (bench/cuda/*, which needs kv_cache and the -// capture types anyway), so a stub could never be linked — it would just be an -// unreachable `return false` claiming an API that isn't really there. - -inline bool available() { return false; } -inline bool pending() { return false; } -inline void flush() {} -inline void* alloc(int64_t, float**, bool = false) { return nullptr; } -inline void release(void*, int64_t, float*) {} -inline bool binary(kop, void*, int64_t, void*, int64_t, void*, int64_t, int64_t, - float, float) { - return false; -} -inline bool binary_bcast(kop, void*, int64_t, int64_t, int64_t, void*, int64_t, - int64_t, int64_t, void*, int64_t, int64_t, int64_t, - float, float) { - return false; -} -inline bool unary(kop, void*, int64_t, void*, int64_t, int64_t, float, float) { - return false; -} -inline bool gemm(void*, int64_t, int64_t, bool, void*, int64_t, int64_t, bool, - void*, int64_t, int64_t, int64_t, int64_t, float, float) { - return false; -} -inline bool gemm_batched(void*, int64_t, int64_t, bool, int64_t, void*, - int64_t, int64_t, bool, int64_t, void*, int64_t, - int64_t, int64_t, int64_t, int64_t, float, float) { - return false; -} -inline bool gemm_bias(void*, int64_t, int64_t, bool, void*, int64_t, int64_t, - bool, void*, int64_t, void*, int64_t, int64_t, int64_t, - int64_t, float, float) { - return false; -} -inline bool row_op(kop, void*, int64_t, void*, int64_t, int64_t, int64_t, float, - float) { - return false; -} -inline bool layer_norm(void*, int64_t, void*, int64_t, void*, int64_t, void*, - int64_t, int64_t, int64_t, float, float, float) { - return false; -} -inline bool layer_norm_bwd(void*, int64_t, void*, int64_t, void*, int64_t, - void*, void*, void*, void*, void*, int64_t, int64_t, - int64_t, int64_t, float) { - return false; -} -inline bool pad(void*, int64_t, void*, int64_t, const int64_t*, - const int64_t*, int, int, int64_t, int64_t, int64_t) { - return false; -} -inline bool fold(void*, int64_t, void*, int64_t, const int64_t*, - const int64_t*, int, int, int64_t, int64_t, int64_t) { - return false; -} -inline bool index_select(void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool index_add(void*, int64_t, void*, int64_t, void*, int64_t, int64_t, - int64_t, int64_t) { - return false; -} -inline bool scatter_to_axis(void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool gather_from_axis(void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool row_logsumexp(void*, int64_t, void*, int64_t, int64_t, int64_t, - float, float) { - return false; -} -inline bool xent_bwd(void*, int64_t, void*, int64_t, void*, int64_t, void*, - int64_t, void*, int64_t, int64_t, int64_t) { - return false; -} -inline bool adam_step(void*, int64_t, void*, int64_t, void*, int64_t, void*, - int64_t, int64_t, float, float, float, float, float) { - return false; -} -inline bool binary_bcast_nd(kop, void*, int64_t, const int64_t*, void*, - int64_t, const int64_t*, void*, int64_t, - const int64_t*, int, int64_t, float, float) { - return false; -} -inline bool where_nd(void*, int64_t, const int64_t*, void*, int64_t, - const int64_t*, void*, int64_t, const int64_t*, void*, - int64_t, const int64_t*, int, int64_t) { - return false; -} -inline bool copy_nd(void*, int64_t, const int64_t*, void*, int64_t, - const int64_t*, int, int64_t) { - return false; -} -inline bool sum_to(void*, int64_t, const int64_t*, const int64_t*, - const int64_t*, int, int64_t, int64_t, void*, int64_t) { - return false; -} -inline bool compare(cmp_op, void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool unary_ext(unary_ext_op, void*, int64_t, void*, int64_t, int64_t, - float, float) { - return false; -} -inline bool clamp(void*, int64_t, void*, int64_t, int64_t, float, float) { - return false; -} -inline bool scalar_binary(scalar_op, void*, int64_t, void*, int64_t, int64_t, - float, float, float) { - return false; -} -inline void sync_to_host(void*, bool) {} - -#endif - // CPU-read barrier: sync the GPU before any host read of a managed buffer. // Nothing to wait for before a CPU access in general: kernels write only // device copies, so sync_to_host waits per buffer, when that buffer's live copy @@ -2652,3 +2156,5 @@ inline void cpu_barrier() {} } // namespace cuda } // namespace tl + +#endif // TENSORLIB_CUDA && !__APPLE__ diff --git a/include/gpu.h b/include/gpu.h index 7cbd5f4..c011e2a 100644 --- a/include/gpu.h +++ b/include/gpu.h @@ -1,66 +1,80 @@ #pragma once -// The GPU-backend facade: array.h and storage.h dispatch through tl::gpu, so -// the eval seam carries no platform #ifdefs. Every backend header exposes the -// identical API and shares tl::metal::kop, which makes the alias below a -// drop-in. The contract, i.e. everything array.h/storage.h may call: -// lifecycle available / pending / flush / cpu_barrier -// memory alloc / release / sync_to_host / upload -// kernels binary / binary_bcast / binary_bcast_nd / where_nd / copy_nd / unary / -// gemm / gemm_batched / gemm_bias / row_op / pad / fold / -// index_select / index_add / scatter_to_axis / gather_from_axis / -// sum_to / compare / unary_ext / clamp / scalar_binary / -// concat_part / rope / layer_norm / layer_norm_bwd / -// row_logsumexp / xent_bwd / adam_step -// LLM path gemv_f32 / gemv_bf16 / gemv_q4 / attn_decode / attn_prefill / -// attn_prefill_dq / attn_prefill_dkv -// model path kv_append / kv_fill / argmax / rmsnorm / rmsnorm_res / swiglu / -// split_heads / merge_heads / gemv_bf16_row / gemm_bf16_nt -// A backend with no kernel for one of these returns false and the evaluator -// falls back to the CPU — so the LLM row is real on CUDA and stubs elsewhere -// (tools/check_backend_parity.py checks every name on both lines here -// against cuda.h/metal.h/webgpu.h and fails CI if one drifts unannounced). -// The model path is what a decoder (kv_cache.h, bench/models) runs on raw -// device buffers between its GEMVs and attention; it has no CPU fallback, -// so a model checks the return and keeps to the array ops where it is false. +// The GPU facade: array.h and storage.h dispatch through tl::gpu, so the eval +// seam carries no platform #ifdefs. tl::gpu is two layers: // -// Beyond the kernels, each backend states what a model may assume of it in -// `caps` (model_path, graph_capture, row_gemv, bf16_gemm, flat_addressing), -// and carries the graph-capture group — graph_available / capture_begin / -// capture_end / -// graph_launch / graph_destroy / upload_u32 / incr_u32 / rope_dpos / -// kv_append_dpos / attn_decode_dpos / attn_dpos_partials_bytes — as no-ops -// where the capability is false, so a decoder is written once and branches on -// caps. `upload` (staging host bytes into a device buffer) is memory, not -// capture: every decoder needs it for the embedding row it feeds each step. +// shared gpu_abi.h — the op vocabulary, `span` (a device handle and a byte +// offset: the currency every op takes), the kernel ABI and the +// launch policy — and gpu_ops.h, where every op is written once. +// An op's signature exists there and nowhere else. +// backend the selected backend header, reached through the using-directive +// below. What a backend provides, its device core: +// lifecycle available / pending / flush / cpu_barrier +// memory alloc / release / sync_to_host / upload +// launch dispatch(kernel, views, params, grid): the one way +// a shared op runs a kernel; false for a kernel id the +// backend has none for +// own the ops it runs its own way (a different algorithm, +// several kernels, a host round trip), as static +// members of `struct own` with the shared signature. +// gpu_ops.h forwards to the ones that exist; an op a +// backend does not declare has no stub to keep in step +// traits what the launch policy may assume of its kernels +// caps what a model may assume (model_path, graph_capture, +// row_gemv, bf16_gemm), plus the graph-capture plumbing +// — graph_available / capture_begin / capture_end / +// graph_launch / graph_destroy / upload_u32 / +// attn_dpos_partials_bytes — as no-ops where absent // -// Each backend compiles to stubs unless its own gate holds, so including all of -// them is free: metal.h is real only on __APPLE__, cuda.h only on -// TENSORLIB_CUDA && !__APPLE__, webgpu.h only on TENSORLIB_WEBGPU && -// __EMSCRIPTEN__. The alias picks the one that can do real work. +// An op answers false when the backend has no kernel for it, and the evaluator +// falls back to the CPU. The model path (what a decoder runs on raw device +// buffers between its GEMVs and attention: kv_cache.h, bench/models) has no CPU +// fallback, so a model checks the return and keeps to the array ops where it +// is false. gpu::census(kernel) counts launches, which is how a test tells a +// kernel that ran from an op that quietly fell back. // -// This lived at the bottom of cuda.h until M10 — the only place both namespaces -// happened to be visible. That stopped scaling at the third backend, since -// adding a browser GPU meant editing the CUDA header. - -#include "cuda.h" -#include "metal.h" -#include "webgpu.h" - -namespace tl { +// One backend is selected below, by the gate its header is written under: +// webgpu.h under TENSORLIB_WEBGPU && __EMSCRIPTEN__, cuda.h under TENSORLIB_CUDA +// && !__APPLE__, metal.h under __APPLE__, and gpu_null.h — no device, every op +// declines — for a build none of them fits. TENSORLIB_HOST_GPU asks for +// gpu_host.h instead of any of them: the reference backend, whose device is the +// CPU and whose kernels are plain loops. gpu_null.h is also the template: +// it is everything this file asks of a backend, with nothing in it. Adding a +// backend is a header that fills that in, its kernels, and one branch here. // WebGPU is checked first: a wasm build defines neither __APPLE__ nor // TENSORLIB_CUDA, but a host build could define TENSORLIB_WEBGPU by accident -// and should not silently take a backend that cannot work there — webgpu:: -// is stubs unless __EMSCRIPTEN__ too, so the order is safe either way. -#if defined(TENSORLIB_WEBGPU) && defined(__EMSCRIPTEN__) -namespace gpu = webgpu; +// and should not take a backend that cannot work there. +#if defined(TENSORLIB_HOST_GPU) // asked for by name: the reference backend +#include "gpu_host.h" +#define TL_GPU_BACKEND host_gpu +#elif defined(TENSORLIB_WEBGPU) && defined(__EMSCRIPTEN__) +#include "webgpu.h" +#define TL_GPU_BACKEND webgpu #elif defined(TENSORLIB_CUDA) && !defined(__APPLE__) -namespace gpu = cuda; +#include "cuda.h" +#define TL_GPU_BACKEND cuda +#elif defined(__APPLE__) +#include "metal.h" +#define TL_GPU_BACKEND metal #else -namespace gpu = metal; +#include "gpu_null.h" +#define TL_GPU_BACKEND null_gpu #endif +namespace tl { + +// A using-directive rather than an alias, so tl::gpu can hold the shared layer +// too: a name declared in tl::gpu itself (a shared op) is found first, and one +// that is not falls through to the backend. +namespace gpu { +using namespace TL_GPU_BACKEND; +} // namespace gpu + inline bool gpu_available() { return gpu::available(); } } // namespace tl + +#include "gpu_ops.h" + +#undef TL_GPU_BACKEND diff --git a/include/gpu_abi.h b/include/gpu_abi.h new file mode 100644 index 0000000..41036c9 --- /dev/null +++ b/include/gpu_abi.h @@ -0,0 +1,316 @@ +#pragma once + +// What the shared GPU layer and a backend's device core agree on. No backend is +// named here, and none needs to be: a backend is a device core (memory, a +// kernel table, one `dispatch`) plus kernel source, and everything above it — +// the ops in gpu_ops.h, the launch policy below — is written once against +// these types. +// +// The kernel ABI. An op hands `dispatch` a kernel id, an ordered list of +// buffer views with how each is accessed, a params struct, and a grid. For a +// kernel id, the order of the views and the layout of the params are the same +// on every backend: views in the order the kernel declares its buffers, and +// the params struct a run of 4-byte fields (uint32_t / int32_t / float) in the +// order the kernel takes its scalars. That is what lets each backend realize a +// launch generically — Metal binds view i at buffer index i and the params +// after them; CUDA builds cuLaunchKernel's argv as the view addresses followed +// by the params' fields, four bytes apiece. + +#include +#include + +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. +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 + where_nd, copy_nd, // N-D select / clone()'s strided gather + gt_, lt_, ge_, le_, eq_, ne_, // comparisons -- cmp_op maps onto these + tanh_, sin_, cos_, // unary_ext_op maps onto these + clamp_, sum_to_, sum_to_blocked_, // dedicated ops, mirroring cuda.h's own + concat_part_, rope_, // ditto -- Tensor.concat / RoPE's own dispatch + 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 + gemv_combine_, // their split-K partials + gemv_bf16_row_, // ... and the [N,K] weight layout's own + gemm_bf16_nt_, gemm_bf16_nt32_, // the prefill's bf16 GEMM, per M tile + gather_axis_, row_logsumexp_, xent_bwd_, adam_step_ // cross-entropy, Adam +}; +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", + "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_", + "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_", + "gemm_bf16_nt_", "gemm_bf16_nt32_", "gather_axis_", "row_logsumexp_", + "xent_bwd_", "adam_step_", +}; +static_assert(sizeof(kKopNames) / sizeof(kKopNames[0]) == kKopCount, + "kKopNames lists every kop, in order"); +inline const char* kop_name(kop k) { return kKopNames[static_cast(k)]; } + +// Comparisons (gt/lt/ge/le/eq/ne), the extra unaries and the tensor-scalar +// ops are their own small vocabularies rather than kop values: each arrived +// on one backend ahead of the others, and an op keyed by its own enum can be +// declined (return false, fall back to the CPU) where a kop could not. +enum class cmp_op { gt, lt, ge, le, eq, ne }; +enum class unary_ext_op { tanh_, sin_, cos_ }; +// pow(x, s) and the comparisons against a scalar, with s a kernel argument +// instead of a rank-0 operand buffer (an allocation and an upload per call). +enum class scalar_op { pow, gt, lt, ge, le, eq, ne }; + +// A view into a device buffer: the currency every op takes. `buf` is whatever +// the backend's alloc() returned — an MTLBuffer handle, a device address, a +// key into a mirror table — and means nothing outside that backend. `off` is +// bytes. Because the offset travels beside the handle instead of inside it, a +// view is expressible on every backend, and arithmetic on a handle is not +// expressible at all. +struct span { + void* buf = nullptr; + int64_t off = 0; + + span at(int64_t bytes) const { return {buf, off + bytes}; } + explicit operator bool() const { return buf != nullptr; } +}; + +// How a kernel touches a view. It drives the residency of a mirrored backend +// (`residency` below), so no op says any of that by hand: an `in` is uploaded +// if the host holds the live copy; an `out` makes the device copy the live one, +// uploading first only if the host had filled the buffer (a view may be part +// of it); an `inout` is uploaded and then becomes live. +enum class access : uint8_t { in, out, inout }; + +struct arg { + span s; + access a; +}; +inline arg in(span s) { return {s, access::in}; } +inline arg out(span s) { return {s, access::out}; } +inline arg inout(span s) { return {s, access::inout}; } + +// Where an allocation's live bytes are, for a backend whose device memory is +// not the host's (a mirrored backend keeps one of these per allocation, next to +// the two copies; a unified one has nothing to track). The backend does the +// copying; when to copy is decided here, once, from how each kernel and each +// host access touches the buffer. +// +// `none` is a fresh allocation nobody has filled. It matters for `out`: a view +// may cover only part of its buffer, so a kernel's output into a buffer whose +// live bytes are the host's has to bring them up first or lose the rest of +// them — but an output into a fresh buffer, which is nearly every output, has +// nothing to bring. +struct residency { + enum state : uint8_t { none, host, device, both }; + state where = none; + + explicit residency(bool host_filled = false) : where(host_filled ? host : none) {} + + // A kernel is about to touch the buffer as `a`. True: upload the host copy + // first. (A read of a `none` buffer uploads too: the host may have filled it + // without saying so, and an unfilled one costs a transfer of garbage that no + // correct program pays.) + bool before_kernel(access a) { + const bool upload = where == host || (where == none && a != access::out); + if (a == access::in) { + if (upload) where = both; + } else { + where = device; + } + return upload; + } + // The host is about to read the buffer or, with for_write, overwrite it. + // True: download the device copy first. (A backend whose download can fail + // asks needs_download(), and reports downloaded() / host_wrote() itself.) + bool before_host(bool for_write) { + const bool download = needs_download(); + if (download) downloaded(); + if (for_write) host_wrote(); + return download; + } + bool needs_download() const { return where == device; } + void downloaded() { where = both; } + void host_wrote() { where = host; } + // The backend copied host bytes up outside a kernel (an explicit upload). + void uploaded() { where = both; } +}; + +// 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). +struct grid { + uint32_t gx = 1, gy = 1, gz = 1; + uint32_t tx = 1, ty = 1, tz = 1; + uint32_t scratch_bytes = 0; +}; + +// Params structs, one per kernel family: the second half of the kernel ABI. +// 4-byte fields only, in the order the kernel takes its scalars (so a backend +// may pass them as one block of bytes or field by field, and they mean the +// same). The MSL and WGSL sources declare the same layouts on their side. +struct ew_params { // elementwise: out[i] = f(...) * scale + offset, i < n + uint32_t n; + float scale, offset; +}; +struct bcast_params { // rank-2 broadcast binary into a contiguous [m, n] + uint32_t m, n; + uint32_t ars, acs, brs, bcs; // each operand's row / column stride, elements + float scale, offset; +}; +struct cmp_params { // out[i] = a[i] CMP b[i * bstride] (bstride 0: a scalar) + uint32_t n, bstride; +}; +struct clamp_params { + uint32_t n; + float lo, hi; +}; +struct scalar_params { // out[i] = (a[i] OP s) * scale + offset + uint32_t n; + float s, scale, offset; +}; +struct reduce_params { // a row op over the last axis of [rows, cols] + uint32_t rows, cols; + float scale, offset; +}; +struct layer_norm_params { + uint32_t rows, cols; + float eps, scale, offset; +}; +struct gather_params { // index_select: rows of `row_size`, n output elements + uint32_t row_size, n; +}; +struct gather_axis_params { // out[i] = src[i * size + idx[i]], i < n + uint32_t size, n; +}; +struct xent_bwd_params { + uint32_t cols, n; // n = rows * cols +}; +struct adam_params { + float b1, b2, eps, lr_over_bc1, inv_bc2; + uint32_t n; +}; +struct rmsnorm_params { // per row of n: x * rsqrt(mean(x^2) + eps) * w + uint32_t n; + float eps; +}; +struct swiglu_params { + uint32_t ff; +}; +struct gemv_row_params { // y[n] = a[k] . W[n, k], one group an output row + uint32_t n, k; +}; +struct gemv_q4_params { + uint32_t n, k, group; +}; +struct kv_append_params { + uint32_t pos, kv_stride; // kv_stride = kv_max * D, elements a head +}; +struct kv_fill_params { + uint32_t T, kv_stride, pos0; +}; +struct merge_heads_params { + uint32_t T, H, D; +}; + +// Launch policy: the shapes ops launch in, in one place. Host code shared by +// every backend; what differs between devices will come in as traits. +namespace policy { + +// One thread an element, in 1-D groups of `threads`. +inline grid flat(int64_t n, uint32_t threads = 256) { + uint32_t groups = static_cast((n + threads - 1) / threads); + return {groups ? groups : 1, 1, 1, threads, 1, 1, 0}; +} + +// One group a row, its threads reducing the row between them; each thread +// keeps `floats_per_thread` of scratch for the reduction tree. +inline grid one_group_per_row(int64_t rows, uint32_t floats_per_thread = 1, + uint32_t threads = 256) { + uint32_t groups = static_cast(rows); + return {groups ? groups : 1, 1, 1, threads, 1, 1, + threads * floats_per_thread * static_cast(sizeof(float))}; +} + +// One thread an element of a row of n, `rows` of them stacked on y. +inline grid flat_rows(int64_t n, int64_t rows, uint32_t threads = 256) { + grid g = flat(n, threads); + g.gy = static_cast(rows); + return g; +} + +// One group an output row, its threads striding that row's k inputs in steps +// of `step` elements and reducing between them: the smallest group (a +// multiple of 32, at most 256) that leaves each thread the fewest steps, so a +// narrow row is not spread over threads with nothing to do. Scratch is one +// float per 32-thread lane set, needed once there is more than one. +inline grid row_reduce(int64_t rows, int64_t k, uint32_t step = 8) { + uint32_t threads = 32; + int64_t fewest = INT64_MAX; + for (uint32_t t = 32; t <= 256; t += 32) { + const int64_t steps = (k + int64_t(step) * t - 1) / (int64_t(step) * t); + if (steps < fewest) { + fewest = steps; + threads = t; + } + } + const uint32_t groups = static_cast(rows); + return {groups ? groups : 1, 1, 1, threads, 1, 1, + threads > 32 ? (threads >> 5) * static_cast(sizeof(float)) : 0}; +} + +// One group per (head, row), a thread per head-dim element: the attention +// and KV-cache kernels' shape. +inline grid per_head(int64_t heads, int64_t rows, int64_t D) { + return {static_cast(heads), static_cast(rows), 1, + static_cast(D), 1, 1, 0}; +} + +// A thread per cell of [rows, cols], for a kernel that reads its cell from a +// 2-D thread position (x the column). 32x8 groups. +inline grid cells_2d(int64_t rows, int64_t cols) { + return {static_cast((cols + 31) / 32), + static_cast((rows + 7) / 8), 1, 32, 8, 1, 0}; +} + +} // namespace policy + +} // namespace gpu +} // namespace tl diff --git a/include/gpu_host.h b/include/gpu_host.h new file mode 100644 index 0000000..6c9a5e9 --- /dev/null +++ b/include/gpu_host.h @@ -0,0 +1,532 @@ +#pragma once + +// The host reference backend: a "device" that is the CPU, whose kernels are +// plain loops run at once. Build with TENSORLIB_HOST_GPU and gpu.h selects it +// ahead of any real one. +// +// It is here for two reasons. It is the fourth backend, written the way +// docs/backends.md says one is — a device core and kernels that follow the +// kernel ABI, no existing backend's file touched — so the claim that a backend +// is that small is something the build checks rather than something the docs +// say. And it lets a machine with no GPU run the shared layer (gpu_ops.h, the +// launch policy, the census) and the conformance tests, which a CPU fallback +// never enters. +// +// A kernel here is also the plainest statement of what its id computes: the +// views in order, the params struct's fields, one loop. Speed is not a goal; +// the threaded CPU path (cpu.h) is what a build without a GPU should use. + +#if defined(TENSORLIB_HOST_GPU) + +#include +#include +#include +#include +#include +#include + +#include "gpu_abi.h" +#include "profile.h" +#include "types.h" // bf16 <-> f32 + +namespace tl { +namespace host_gpu { + +using kop = gpu::kop; + +// ---- lifecycle. A kernel has run by the time dispatch returns, but the +// backend keeps a device's form: a launch leaves work pending until a flush +// "waits" for it, so the evaluator's flush and barrier paths, and what +// tl::profile records of them, are exercised here as on a real device. +namespace detail_ { +inline bool pending_ = false; +} +inline bool available() { return true; } +inline bool pending() { return detail_::pending_; } +inline void flush() { + if (!detail_::pending_) return; + profile::detail::blocked waiting; + detail_::pending_ = false; +} +inline void cpu_barrier() { + if (pending()) flush(); +} + +// ---- memory: unified, so `native` is the address itself and a view's bytes +// are at buf + off. +inline void* alloc(int64_t bytes, float** contents, bool /*host_fill*/ = false) { + const size_t n = (static_cast(bytes > 0 ? bytes : 4) + 63) & ~size_t(63); +#ifdef _WIN32 + void* p = _aligned_malloc(n, 64); +#else + void* p = std::aligned_alloc(64, n); +#endif + if (contents) *contents = static_cast(p); + return p; +} +inline void release(void* buf, int64_t, float*) { +#ifdef _WIN32 + _aligned_free(buf); +#else + std::free(buf); +#endif +} +inline void sync_to_host(void*, bool) {} +inline void upload(void* native, const float* src, int64_t n) { + if (native && src != native && n > 0) { + std::memcpy(native, src, static_cast(n) * sizeof(float)); + } +} + +namespace detail_ { + +template +inline T* at(const gpu::arg& a) { + return reinterpret_cast(static_cast(a.s.buf) + a.s.off); +} +template +inline const P& as(const void* params) { + return *static_cast(params); +} + +// out[i] = f(a[i], b[i]) * scale + offset +template +inline bool binary(const gpu::arg* v, const void* params, F f) { + const auto& p = as(params); + const float *a = at(v[0]), *b = at(v[1]); + float* o = at(v[2]); + for (uint32_t i = 0; i < p.n; i++) o[i] = f(a[i], b[i]) * p.scale + p.offset; + return true; +} +template +inline bool unary(const gpu::arg* v, const void* params, F f) { + const auto& p = as(params); + const float* a = at(v[0]); + float* o = at(v[1]); + for (uint32_t i = 0; i < p.n; i++) o[i] = f(a[i]) * p.scale + p.offset; + return true; +} +template +inline bool bcast(const gpu::arg* v, const void* params, F f) { + const auto& p = as(params); + const float *a = at(v[0]), *b = at(v[1]); + float* o = at(v[2]); + for (uint32_t r = 0; r < p.m; r++) { + for (uint32_t c = 0; c < p.n; c++) { + o[size_t(r) * p.n + c] = + f(a[size_t(r) * p.ars + size_t(c) * p.acs], + b[size_t(r) * p.brs + size_t(c) * p.bcs]) * p.scale + p.offset; + } + } + return true; +} +template +inline bool compare(const gpu::arg* v, const void* params, F f) { + const auto& p = as(params); + const float *a = at(v[0]), *b = at(v[1]); + float* o = at(v[2]); + for (uint32_t i = 0; i < p.n; i++) { + o[i] = f(a[i], b[size_t(i) * p.bstride]) ? 1.0f : 0.0f; + } + return true; +} +template +inline bool scalar(const gpu::arg* v, const void* params, F f) { + const auto& p = as(params); + const float* a = at(v[0]); + float* o = at(v[1]); + for (uint32_t i = 0; i < p.n; i++) o[i] = f(a[i], p.s) * p.scale + p.offset; + return true; +} +// One value a row of [rows, cols]: reduce(row) * scale + offset. +template +inline bool row_reduce(const gpu::arg* v, const void* params, F reduce) { + const auto& p = as(params); + const float* a = at(v[0]); + float* o = at(v[1]); + for (uint32_t r = 0; r < p.rows; r++) { + o[r] = reduce(a + size_t(r) * p.cols, p.cols) * p.scale + p.offset; + } + return true; +} +// sum(exp(x - max)) and the max: what softmax and logsumexp both reduce to. +// Subtracting the max first keeps huge logits from overflowing. +inline double sum_exp(const float* x, uint32_t n, float* max_out) { + float mx = x[0]; + for (uint32_t c = 1; c < n; c++) mx = std::max(mx, x[c]); + *max_out = mx; + if (std::isinf(mx)) return mx > 0 ? 1.0 : 0.0; // a row with no finite max + double s = 0; + for (uint32_t c = 0; c < n; c++) s += std::exp(double(x[c]) - mx); + return s; +} +inline float row_logsumexp(const float* x, uint32_t n) { + float mx; + const double s = sum_exp(x, n, &mx); + return std::isinf(mx) ? mx : mx + static_cast(std::log(s)); +} +inline bool softmax(const gpu::arg* v, const void* params) { + const auto& p = as(params); + const float* a = at(v[0]); + float* o = at(v[1]); + for (uint32_t r = 0; r < p.rows; r++) { + const float* x = a + size_t(r) * p.cols; + float mx; + const double s = sum_exp(x, p.cols, &mx); + for (uint32_t c = 0; c < p.cols; c++) { + const float y = static_cast(std::exp(double(x[c]) - mx) / s); + o[size_t(r) * p.cols + c] = y * p.scale + p.offset; + } + } + return true; +} +inline bool layer_norm(const gpu::arg* v, const void* params) { + const auto& p = as(params); + const float *x = at(v[0]), *g = at(v[1]), + *b = at(v[2]); + float* o = at(v[3]); + for (uint32_t r = 0; r < p.rows; r++) { + const float* row = x + size_t(r) * p.cols; + double mean = 0, var = 0; + for (uint32_t c = 0; c < p.cols; c++) mean += row[c]; + mean /= p.cols; + for (uint32_t c = 0; c < p.cols; c++) var += (row[c] - mean) * (row[c] - mean); + const double inv = 1.0 / std::sqrt(var / p.cols + p.eps); + for (uint32_t c = 0; c < p.cols; c++) { + const float y = static_cast((row[c] - mean) * inv) * g[c] + b[c]; + o[size_t(r) * p.cols + c] = y * p.scale + p.offset; + } + } + return true; +} +// hout = v * rsqrt(mean(v^2) + eps) * w per row, v = x (+ delta, also stored). +inline bool rmsnorm(const gpu::arg* v, const void* params, const gpu::grid& g, + bool add) { + const auto& p = as(params); + const float* x = at(v[0]); + const float* delta = add ? at(v[1]) : nullptr; + const float* w = at(v[add ? 2 : 1]); + float* xout = add ? at(v[3]) : nullptr; + float* hout = at(v[add ? 4 : 2]); + for (uint32_t r = 0; r < g.gx; r++) { // one group a row + const size_t base = size_t(r) * p.n; + double ss = 0; + for (uint32_t i = 0; i < p.n; i++) { + const float val = add ? x[base + i] + delta[base + i] : x[base + i]; + if (add) xout[base + i] = val; + ss += double(val) * val; + } + const float inv = static_cast(1.0 / std::sqrt(ss / p.n + p.eps)); + const float* src = add ? xout : x; + for (uint32_t i = 0; i < p.n; i++) hout[base + i] = src[base + i] * inv * w[i]; + } + return true; +} +inline bool swiglu(const gpu::arg* v, const void* params, const gpu::grid& g) { + const uint32_t ff = as(params).ff; + const float* gu = at(v[0]); + float* o = at(v[1]); + for (uint32_t r = 0; r < g.gy; r++) { // rows ride on the grid's y + for (uint32_t f = 0; f < ff; f++) { + const float gate = gu[size_t(r) * 2 * ff + f], up = gu[size_t(r) * 2 * ff + ff + f]; + o[size_t(r) * ff + f] = gate / (1.0f + std::exp(-gate)) * up; + } + } + return true; +} +inline bool index_select(const gpu::arg* v, const void* params) { + const auto& p = as(params); + const float *a = at(v[0]), *idx = at(v[1]); + float* o = at(v[2]); + for (uint32_t i = 0; i < p.n; i++) { + const size_t row = i / p.row_size, col = i % p.row_size; + o[i] = a[static_cast(idx[row] + 0.5f) * p.row_size + col]; + } + return true; +} +inline bool gather_axis(const gpu::arg* v, const void* params) { + const auto& p = as(params); + const float *src = at(v[0]), *idx = at(v[1]); + float* o = at(v[2]); + for (uint32_t i = 0; i < p.n; i++) { + o[i] = src[size_t(i) * p.size + static_cast(idx[i] + 0.5f)]; + } + return true; +} +inline bool xent_bwd(const gpu::arg* v, const void* params) { + const auto& p = as(params); + const float *x = at(v[0]), *lse = at(v[1]), + *tgt = at(v[2]), *g = at(v[3]); + float* o = at(v[4]); + for (uint32_t i = 0; i < p.n; i++) { + const size_t row = i / p.cols, col = i % p.cols; + float prob = std::exp(x[i] - lse[row]); + if (col == static_cast(tgt[row] + 0.5f)) prob -= 1.0f; + o[i] = prob * g[row]; + } + return true; +} +inline bool adam_step(const gpu::arg* v, const void* params) { + const auto& q = as(params); + float *p = at(v[0]), *m = at(v[1]), *vv = at(v[2]); + const float* g = at(v[3]); + for (uint32_t i = 0; i < q.n; i++) { + m[i] = q.b1 * m[i] + (1.0f - q.b1) * g[i]; + vv[i] = q.b2 * vv[i] + (1.0f - q.b2) * g[i] * g[i]; + p[i] -= (m[i] * q.lr_over_bc1) / (std::sqrt(vv[i] * q.inv_bc2) + q.eps); + } + return true; +} +// y[n] = a[k] . W[n, k], the weight row-major: bf16, or int4 with a scale a +// group of `group` weights (8 nibbles a word, low nibble first, offset 8). +inline bool gemv_bf16_row(const gpu::arg* v, const void* params) { + const auto& p = as(params); + const float* a = at(v[0]); + const uint16_t* W = at(v[1]); + float* y = at(v[2]); + for (uint32_t r = 0; r < p.n; r++) { + double acc = 0; + for (uint32_t c = 0; c < p.k; c++) acc += double(a[c]) * bf16_to_f32(W[size_t(r) * p.k + c]); + y[r] = static_cast(acc); + } + return true; +} +inline bool gemv_q4(const gpu::arg* v, const void* params) { + const auto& p = as(params); + const float* a = at(v[0]); + const uint32_t* qw = at(v[1]); + const float* scales = at(v[2]); + float* y = at(v[3]); + for (uint32_t r = 0; r < p.n; r++) { + double acc = 0; + for (uint32_t c = 0; c < p.k; c++) { + const uint32_t word = qw[size_t(r) * (p.k / 8) + c / 8]; + const int q = static_cast((word >> ((c % 8) * 4)) & 0xFu) - 8; + acc += double(a[c]) * (scales[size_t(r) * (p.k / p.group) + c / p.group] * q); + } + y[r] = static_cast(acc); + } + return true; +} +// The KV cache's writes, f32 or bf16: heads ride on the grid's x, a fill's T +// rows on its y, and the head dim is the group's thread count. +template +inline bool kv_write(const gpu::arg* v, const gpu::grid& g, uint32_t kv_stride, + uint32_t pos0, uint32_t T, Narrow narrow) { + KT *Kc = at(v[0]), *Vc = at(v[1]); + const float *k = at(v[2]), *val = at(v[3]); + const uint32_t D = g.tx; + for (uint32_t h = 0; h < g.gx; h++) { + for (uint32_t t = 0; t < T; t++) { + for (uint32_t d = 0; d < D; d++) { + const size_t dst = size_t(h) * kv_stride + size_t(pos0 + t) * D + d; + const size_t src = (size_t(h) * T + t) * D + d; + Kc[dst] = narrow(k[src]); + Vc[dst] = narrow(val[src]); + } + } + } + return true; +} +inline bool kv_append(const gpu::arg* v, const void* params, const gpu::grid& g, + bool bf16) { + const auto& p = as(params); + return bf16 ? kv_write(v, g, p.kv_stride, p.pos, 1, f32_to_bf16) + : kv_write(v, g, p.kv_stride, p.pos, 1, [](float f) { return f; }); +} +inline bool kv_fill(const gpu::arg* v, const void* params, const gpu::grid& g, + bool bf16) { + const auto& p = as(params); + return bf16 ? kv_write(v, g, p.kv_stride, p.pos0, p.T, f32_to_bf16) + : kv_write(v, g, p.kv_stride, p.pos0, p.T, [](float f) { return f; }); +} +inline bool merge_heads(const gpu::arg* v, const void* params) { + const auto& p = as(params); + const float* src = at(v[0]); + float* dst = at(v[1]); + for (uint32_t h = 0; h < p.H; h++) + for (uint32_t t = 0; t < p.T; t++) + for (uint32_t d = 0; d < p.D; d++) + dst[(size_t(t) * p.H + h) * p.D + d] = src[(size_t(h) * p.T + t) * p.D + d]; + return true; +} + +} // namespace detail_ + +// ---- launch: the kernel table and the kernels in one switch. An id with no +// case declines, and the op above falls back to the CPU. +inline bool dispatch(kop k, const gpu::arg* v, size_t /*n*/, const void* params, + size_t /*params_bytes*/, const gpu::grid& g) { + namespace d = detail_; + d::pending_ = true; // declining leaves it set too: a flush then waits on nothing + switch (k) { + case kop::add: return d::binary(v, params, [](float a, float b) { return a + b; }); + case kop::sub: return d::binary(v, params, [](float a, float b) { return a - b; }); + case kop::mul: return d::binary(v, params, [](float a, float b) { return a * b; }); + case kop::div: return d::binary(v, params, [](float a, float b) { return a / b; }); + case kop::pow_: return d::binary(v, params, [](float a, float b) { return std::pow(a, b); }); + case kop::exp_: return d::unary(v, params, [](float a) { return std::exp(a); }); + case kop::log_: return d::unary(v, params, [](float a) { return std::log(a); }); + case kop::sqrt_: return d::unary(v, params, [](float a) { return std::sqrt(a); }); + case kop::sigmoid: return d::unary(v, params, [](float a) { return 1.0f / (1.0f + std::exp(-a)); }); + case kop::relu: return d::unary(v, params, [](float a) { return a > 0 ? a : 0.0f; }); + case kop::affine: return d::unary(v, params, [](float a) { return a; }); + case kop::tanh_: return d::unary(v, params, [](float a) { return std::tanh(a); }); + case kop::sin_: return d::unary(v, params, [](float a) { return std::sin(a); }); + case kop::cos_: return d::unary(v, params, [](float a) { return std::cos(a); }); + case kop::badd: return d::bcast(v, params, [](float a, float b) { return a + b; }); + case kop::bsub: return d::bcast(v, params, [](float a, float b) { return a - b; }); + case kop::bmul: return d::bcast(v, params, [](float a, float b) { return a * b; }); + case kop::bdiv: return d::bcast(v, params, [](float a, float b) { return a / b; }); + case kop::bpow: return d::bcast(v, params, [](float a, float b) { return std::pow(a, b); }); + case kop::gt_: return d::compare(v, params, [](float a, float b) { return a > b; }); + case kop::lt_: return d::compare(v, params, [](float a, float b) { return a < b; }); + case kop::ge_: return d::compare(v, params, [](float a, float b) { return a >= b; }); + case kop::le_: return d::compare(v, params, [](float a, float b) { return a <= b; }); + case kop::eq_: return d::compare(v, params, [](float a, float b) { return a == b; }); + case kop::ne_: return d::compare(v, params, [](float a, float b) { return a != b; }); + case kop::clamp_: { + const auto& p = d::as(params); + const float* a = d::at(v[0]); + float* o = d::at(v[1]); + for (uint32_t i = 0; i < p.n; i++) o[i] = std::min(std::max(a[i], p.lo), p.hi); + return true; + } + case kop::pow_s_: return d::scalar(v, params, [](float a, float s) { return std::pow(a, s); }); + case kop::gt_s_: return d::scalar(v, params, [](float a, float s) { return a > s ? 1.0f : 0.0f; }); + case kop::lt_s_: return d::scalar(v, params, [](float a, float s) { return a < s ? 1.0f : 0.0f; }); + case kop::ge_s_: return d::scalar(v, params, [](float a, float s) { return a >= s ? 1.0f : 0.0f; }); + case kop::le_s_: return d::scalar(v, params, [](float a, float s) { return a <= s ? 1.0f : 0.0f; }); + case kop::eq_s_: return d::scalar(v, params, [](float a, float s) { return a == s ? 1.0f : 0.0f; }); + case kop::ne_s_: return d::scalar(v, params, [](float a, float s) { return a != s ? 1.0f : 0.0f; }); + case kop::softmax: return d::softmax(v, params); + case kop::row_sum: + return d::row_reduce(v, params, [](const float* x, uint32_t n) { + double s = 0; + for (uint32_t c = 0; c < n; c++) s += x[c]; + return static_cast(s); + }); + case kop::row_max: + return d::row_reduce(v, params, [](const float* x, uint32_t n) { + return *std::max_element(x, x + n); + }); + case kop::row_logsumexp_: return d::row_reduce(v, params, d::row_logsumexp); + case kop::layer_norm_: return d::layer_norm(v, params); + case kop::index_select: return d::index_select(v, params); + case kop::gather_axis_: return d::gather_axis(v, params); + case kop::xent_bwd_: return d::xent_bwd(v, params); + case kop::adam_step_: return d::adam_step(v, params); + case kop::rmsnorm_: return d::rmsnorm(v, params, g, false); + case kop::add_rmsnorm_: return d::rmsnorm(v, params, g, true); + case kop::swiglu_: return d::swiglu(v, params, g); + case kop::gemv_bf16_row_: return d::gemv_bf16_row(v, params); + case kop::gemv_q4_: return d::gemv_q4(v, params); + case kop::kv_append_: return d::kv_append(v, params, g, false); + case kop::kv_append_bf16_: return d::kv_append(v, params, g, true); + case kop::kv_fill_: return d::kv_fill(v, params, g, false); + case kop::kv_fill_bf16_: return d::kv_fill(v, params, g, true); + case kop::merge_heads_: return d::merge_heads(v, params); + default: return false; + } +} + +// ---- the ops this backend runs its own way: the ones with no single-kernel +// form in gpu_ops.h that the conformance test expects of every backend. Each +// is the definition of its op, as a loop. +struct own { + // C(m,n) = (A @ B) * scale + offset; a transposed operand is read in place. + static bool gemm(gpu::span a, int64_t lda, bool ta, gpu::span b, int64_t ldb, + bool tb, gpu::span out, int64_t m, int64_t n, int64_t k, + float scale, float offset) { + const float *A = view(a), *B = view(b); + float* C = view(out); + detail_::pending_ = true; + for (int64_t i = 0; i < m; i++) { + for (int64_t j = 0; j < n; j++) { + double acc = 0; + for (int64_t q = 0; q < k; q++) { + acc += double(ta ? A[q * lda + i] : A[i * lda + q]) * + (tb ? B[j * ldb + q] : B[q * ldb + j]); + } + C[i * n + j] = static_cast(acc) * scale + offset; + } + } + return true; + } + + // `a` into a zero buffer of out_shape, shifted by `before` along `axis`. + static bool pad(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, int64_t before, + int64_t n, int64_t out_n) { + const float* A = view(a); + float* O = view(out); + std::fill(O, O + out_n, 0.0f); + int64_t inner = 1; + for (int d = axis + 1; d < rank; d++) inner *= a_shape[d]; + const int64_t a_axis = a_shape[axis], o_axis = out_shape[axis]; + for (int64_t i = 0; i < n; i++) { + const int64_t outer = i / (a_axis * inner), rest = i % (a_axis * inner); + O[outer * o_axis * inner + before * inner + rest] = A[i]; + } + return true; + } + + // out row idx[i] += values row i, for k rows, into a zeroed out. + static bool index_add(gpu::span idx, gpu::span values, gpu::span out, + int64_t row_size, int64_t k, int64_t out_n) { + const float *I = view(idx), *V = view(values); + float* O = view(out); + std::fill(O, O + out_n, 0.0f); + for (int64_t i = 0; i < k; i++) { + const int64_t row = static_cast(I[i] + 0.5f); + for (int64_t c = 0; c < row_size; c++) O[row * row_size + c] += V[i * row_size + c]; + } + return true; + } + + // out[i, idx[i]] = values[i], zero elsewhere, over a trailing axis of `size`. + static bool scatter_to_axis(gpu::span idx, gpu::span values, gpu::span out, + int64_t n, int64_t size) { + const float *I = view(idx), *V = view(values); + float* O = view(out); + std::fill(O, O + n * size, 0.0f); + for (int64_t i = 0; i < n; i++) O[i * size + static_cast(I[i] + 0.5f)] = V[i]; + return true; + } + + private: + template + static T* view(gpu::span s) { + return reinterpret_cast(static_cast(s.buf) + s.off); + } +}; + +// ---- what the shared layer may assume. +struct traits { + static constexpr bool cells_2d = false; // a cell is read from a flat index + // Whether this backend records its own launches under tl::profile (else the + // shared layer does), and whether each carries a device time. + static constexpr bool profiles_launches = false; + static constexpr bool times_launches = false; +}; +struct caps { + // The decoder's single-kernel ops are here, but not attention, rope or the + // GEMVs a model also needs, so a model keeps to the array ops. + static constexpr bool model_path = false; + static constexpr bool graph_capture = false; + static constexpr bool row_gemv = true; + static constexpr bool bf16_gemm = false; +}; +using graph_exec = void*; +inline bool graph_available() { return false; } +inline bool capture_begin() { return false; } +inline graph_exec capture_end() { return nullptr; } +inline bool graph_launch(graph_exec) { return false; } +inline void graph_destroy(graph_exec) {} +inline void upload_u32(void*, unsigned) {} +inline int64_t attn_dpos_partials_bytes(int64_t, int64_t, int64_t) { return 0; } + +} // namespace host_gpu +} // namespace tl + +#endif // TENSORLIB_HOST_GPU diff --git a/include/gpu_null.h b/include/gpu_null.h new file mode 100644 index 0000000..5ab255a --- /dev/null +++ b/include/gpu_null.h @@ -0,0 +1,74 @@ +#pragma once + +// The backend a build gets when no GPU backend's gate holds: no device, so +// every op declines and the evaluator runs on the CPU. It is also the whole of +// what gpu.h asks of a backend, with nothing in it — the list a new backend +// fills in: a device core (lifecycle, memory, one `dispatch`), what it runs its +// own way (`own`), and what the shared layer may assume of it (`traits`, +// `caps`, the graph-capture plumbing). + +#include +#include + +#include "gpu_abi.h" + +namespace tl { +namespace null_gpu { + +// ---- lifecycle +inline bool available() { return false; } +inline bool pending() { return false; } // work encoded but not yet waited on +inline void flush() {} // submit it and block until it is done +inline void cpu_barrier() {} // before any host read of a buffer + +// ---- memory. `native` is this backend's handle for a buffer — whatever +// gpu::span::buf should carry — and `contents` the host-readable bytes behind +// it (the same memory where it is unified, a mirror where it is not). +inline void* alloc(int64_t, float** /*contents*/, bool /*host_fill*/ = false) { + return nullptr; // null: storage falls back to the heap +} +inline void release(void*, int64_t, float*) {} +inline void sync_to_host(void*, bool /*for_write*/) {} // before a host access +inline void upload(void*, const float*, int64_t) {} // stage host floats in + +// ---- launch: the one way a shared op (gpu_ops.h) runs a kernel. View i is the +// kernel's i-th buffer at its byte offset, `params` a block of 4-byte fields +// in the kernel's argument order (gpu_abi.h). False for a kernel id this +// backend has no kernel for. +inline bool dispatch(gpu::kop, const gpu::arg*, size_t, const void* /*params*/, + size_t /*params_bytes*/, const gpu::grid&) { + return false; +} + +// ---- the ops this backend runs its own way, as static members with the +// signature the op has in gpu_ops.h. None. +struct own {}; + +// ---- what the shared launch policy may assume of this backend's kernels. +struct traits { + static constexpr bool cells_2d = false; + // Whether this backend records its own launches under tl::profile (else the + // shared layer does), and whether each carries a device time. + static constexpr bool profiles_launches = false; + static constexpr bool times_launches = false; +}; + +// ---- what a model may assume, and the graph-capture plumbing behind +// caps::graph_capture. +struct caps { + static constexpr bool model_path = false; + static constexpr bool graph_capture = false; + static constexpr bool row_gemv = false; + static constexpr bool bf16_gemm = false; +}; +using graph_exec = void*; +inline bool graph_available() { return false; } +inline bool capture_begin() { return false; } +inline graph_exec capture_end() { return nullptr; } +inline bool graph_launch(graph_exec) { return false; } +inline void graph_destroy(graph_exec) {} +inline void upload_u32(void*, unsigned) {} +inline int64_t attn_dpos_partials_bytes(int64_t, int64_t, int64_t) { return 0; } + +} // namespace null_gpu +} // namespace tl diff --git a/include/gpu_ops.h b/include/gpu_ops.h new file mode 100644 index 0000000..b94af6d --- /dev/null +++ b/include/gpu_ops.h @@ -0,0 +1,705 @@ +#pragma once + +// The GPU ops, written once. An op is a kernel id, the views it touches and +// how, a params struct, and a launch shape (gpu_abi.h); `launch` hands those to +// the selected backend's `dispatch`. Nothing here names a backend, so an op +// added here exists on all of them, and a backend without the kernel declines +// (false) and the caller falls back, exactly as before. +// +// Included by gpu.h, after the backend is selected. + +#include +#include +#include +#include + +#include "gpu_abi.h" +#include "profile.h" + +namespace tl { +namespace gpu { + +// The census: how a test tells work that ran on the device from an op that +// quietly fell back to the CPU, which no comparison against an oracle can (the +// fallback is right too). census(k) counts the shared ops' launches per kernel +// id, ops_run() every op that ran here, shared or backend-own, since the last +// census_reset(). +namespace detail { +inline std::array census_counts{}; +inline uint64_t census_ops_run = 0; +// A launch the backend did not record under tl::profile itself is recorded +// here, by name, so a backend is profiled from its first kernel; one that +// stamps its launches with a device time says so (traits::profiles_launches) +// and records its own. +inline void profile_launch(const char* name) { + if (!traits::profiles_launches && profile::active()) profile::detail::launch(name); +} +inline bool ran(const char* op, bool ok) { + if (ok) { + census_ops_run++; + profile_launch(op); + } + return ok; +} +} // namespace detail + +inline uint64_t census(kop k) { + return detail::census_counts[static_cast(k)]; +} +inline uint64_t ops_run() { return detail::census_ops_run; } +inline void census_reset() { + detail::census_counts.fill(0); + detail::census_ops_run = 0; +} + +// A backend may run an op its own way — a different algorithm, several +// kernels, a host round trip — by declaring it as a static member of its `own` +// struct, with the signature the op has here. An op that allows this asks +// TL_GPU_OWNS whether the selected backend did; one that did not declare it +// simply has no such member, so there are no stubs to keep in step. +#define TL_GPU_DETECT_OWN(name) \ + namespace detail { \ + template \ + struct owns_##name : std::false_type {}; \ + template \ + struct owns_##name> : std::true_type {}; \ + } \ + inline constexpr bool has_##name = detail::owns_##name::value; + +// Every shared op launches through here. +template +inline bool launch(kop k, std::initializer_list args, const P& params, + const grid& g) { + static_assert(std::is_trivially_copyable

::value && sizeof(P) % 4 == 0 && + alignof(P) == 4, + "kernel params are a run of 4-byte fields (gpu_abi.h)"); + if (!dispatch(k, args.begin(), args.size(), ¶ms, sizeof(P), g)) { + return false; + } + detail::census_counts[static_cast(k)]++; + detail::census_ops_run++; + detail::profile_launch(kop_name(k)); + return true; +} + +// out[i] = (a[i] OP b[i]) * scale + offset over n contiguous elements. +inline bool binary(kop op, span a, span b, span o, int64_t n, float scale, + float offset) { + return launch(op, {in(a), in(b), out(o)}, + ew_params{static_cast(n), scale, offset}, + policy::flat(n)); +} + +// out[i] = OP(a[i]) * scale + offset over n contiguous elements. +inline bool unary(kop op, span a, span o, int64_t n, float scale, + float offset) { + return launch(op, {in(a), out(o)}, + ew_params{static_cast(n), scale, offset}, + policy::flat(n)); +} + +// Rank-2 broadcast binary: out[r,c] = f(a[r*ars + c*acs], b[r*brs + c*bcs]) +// into a contiguous [m, n] output, affine epilogue. One stride-parameterized +// kernel covers every rank-2 broadcast (row vector, column vector, per-row +// scalar), which keeps bias/gamma/beta chains on the device. Strides are +// elements and non-negative. A backend's kernel reads its cell either from a +// flat index or from a 2-D thread position; traits::cells_2d says which. +inline bool binary_bcast(kop op, span a, int64_t ars, int64_t acs, span b, + int64_t brs, int64_t bcs, span o, int64_t m, int64_t n, + float scale, float offset) { + if (m <= 0 || n <= 0 || ars < 0 || acs < 0 || brs < 0 || bcs < 0) return false; + auto u = [](int64_t v) { return static_cast(v); }; + return launch(op, {in(a), in(b), out(o)}, + bcast_params{u(m), u(n), u(ars), u(acs), u(brs), u(bcs), scale, + offset}, + traits::cells_2d ? policy::cells_2d(m, n) : policy::flat(m * n)); +} + +// out[i] = a[i] CMP b[i * bstride] as 1.0 / 0.0 (bstride 0 broadcasts a +// scalar b). No epilogue. +inline bool compare(cmp_op op, span a, span b, span o, int64_t n, + int64_t bstride) { + kop k = kop::gt_; + switch (op) { + case cmp_op::gt: k = kop::gt_; break; + case cmp_op::lt: k = kop::lt_; break; + case cmp_op::ge: k = kop::ge_; break; + case cmp_op::le: k = kop::le_; break; + case cmp_op::eq: k = kop::eq_; break; + case cmp_op::ne: k = kop::ne_; break; + } + return launch(k, {in(a), in(b), out(o)}, + cmp_params{static_cast(n), static_cast(bstride)}, + policy::flat(n)); +} + +// out[i] = min(max(a[i], lo), hi): Clip's forward. No epilogue. +inline bool clamp(span a, span o, int64_t n, float lo, float hi) { + return launch(kop::clamp_, {in(a), out(o)}, + clamp_params{static_cast(n), lo, hi}, policy::flat(n)); +} + +// Tensor-scalar ops: out[i] = (a[i] OP s) * scale + offset, the scalar a +// kernel argument rather than a rank-0 operand buffer. +inline bool scalar_binary(scalar_op op, span a, span o, int64_t n, float s, + float scale, float offset) { + kop k = kop::pow_s_; + switch (op) { + case scalar_op::pow: k = kop::pow_s_; break; + case scalar_op::gt: k = kop::gt_s_; break; + case scalar_op::lt: k = kop::lt_s_; break; + case scalar_op::ge: k = kop::ge_s_; break; + case scalar_op::le: k = kop::le_s_; break; + case scalar_op::eq: k = kop::eq_s_; break; + case scalar_op::ne: k = kop::ne_s_; break; + } + return launch(k, {in(a), out(o)}, + scalar_params{static_cast(n), s, scale, offset}, + policy::flat(n)); +} + +// Row op over the last axis of [rows, cols]: softmax writes rows x cols; +// row_sum / row_max write one value a row. Affine epilogue. +inline bool row_op(kop op, span a, span o, int64_t rows, int64_t cols, + float scale, float offset) { + if (rows <= 0 || cols <= 0) return false; + return launch(op, {in(a), out(o)}, + reduce_params{static_cast(rows), + static_cast(cols), scale, offset}, + policy::one_group_per_row(rows)); +} + +// log(sum exp) per row, affine epilogue: row_op's shape, one pass over the row +// (each thread carries a running max and sum, hence two floats of scratch). +inline bool row_logsumexp(span a, span o, int64_t rows, int64_t cols, + float scale, float offset) { + if (rows <= 0 || cols <= 0) return false; + return launch(kop::row_logsumexp_, {in(a), out(o)}, + reduce_params{static_cast(rows), + static_cast(cols), scale, offset}, + policy::one_group_per_row(rows, 2)); +} + +// out = (x - mean) / sqrt(var + eps) * g + b per row of [rows, cols], affine +// epilogue; g and b are contiguous cols-vectors. +inline bool layer_norm(span x, span g, span b, span o, int64_t rows, + int64_t cols, float eps, float scale, float offset) { + if (rows <= 0 || cols <= 0) return false; + return launch(kop::layer_norm_, {in(x), in(g), in(b), out(o)}, + layer_norm_params{static_cast(rows), + static_cast(cols), eps, scale, offset}, + policy::one_group_per_row(rows)); +} + +// Row gather along axis 0: out row i = a row idx[i], for k rows of row_size. +inline bool index_select(span a, span idx, span o, int64_t row_size, + int64_t k) { + const int64_t n = k * row_size; + if (n <= 0) return false; + return launch(kop::index_select, {in(a), in(idx), out(o)}, + gather_params{static_cast(row_size), + static_cast(n)}, + policy::flat(n)); +} + +// out[i] = src[i * size + idx[i]]: the element each of n positions labels +// along a trailing axis of `size`. +inline bool gather_from_axis(span src, span idx, span o, int64_t n, + int64_t size) { + if (n <= 0) return false; + return launch(kop::gather_axis_, {in(src), in(idx), out(o)}, + gather_axis_params{static_cast(size), + static_cast(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)); +} + +// ---- the model path: what a decoder runs on raw device buffers between its +// GEMVs and attention. No CPU fallback sits under these, so a model checks the +// return and keeps to the array ops where it is false. + +// out = x * rsqrt(mean(x^2) + eps) * w per row of [rows, n]. out may alias x. +inline bool rmsnorm(span x, span w, span o, int64_t n, float eps, + int64_t rows = 1) { + if (n <= 0 || rows <= 0) return false; + return launch(kop::rmsnorm_, {in(x), in(w), out(o)}, + rmsnorm_params{static_cast(n), eps}, + policy::one_group_per_row(rows)); +} + +// xout = x + delta and hout = rmsnorm(xout) * w: a residual add folded into +// the norm that follows it. xout may alias x. +inline bool rmsnorm_res(span x, span delta, span w, span xout, span hout, + int64_t n, float eps, int64_t rows = 1) { + if (n <= 0 || rows <= 0) return false; + return launch(kop::add_rmsnorm_, + {in(x), in(delta), in(w), out(xout), out(hout)}, + rmsnorm_params{static_cast(n), eps}, + policy::one_group_per_row(rows)); +} + +// out[rows, ff] = silu(gate) * up out of the fused gate|up buffer [rows, 2ff]. +inline bool swiglu(span gu, span o, int64_t ff, int64_t rows = 1) { + if (ff <= 0 || rows <= 0) return false; + return launch(kop::swiglu_, {in(gu), out(o)}, + swiglu_params{static_cast(ff)}, + policy::flat_rows(ff, rows)); +} + +// y[1,N] = a[1,K] . W[N,K]^T with the bf16 weight row-major (GGML-native): one +// group an output row. Requires k % 8 == 0. +inline bool gemv_bf16_row(span a, span W, span y, int64_t n, int64_t k) { + if (n <= 0 || k <= 0 || k % 8 != 0) return false; + return launch(kop::gemv_bf16_row_, {in(a), in(W), out(y)}, + gemv_row_params{static_cast(n), static_cast(k)}, + policy::row_reduce(n, k)); +} + +// The same over int4 weights: qw is [N][K/8] packed words and scales +// [N][K/group] floats, two views that may share a buffer. K % group == 0 and +// group % 8 == 0. +inline bool gemv_q4(span a, span qw, span scales, span y, int64_t N, int64_t K, + int64_t group) { + if (N <= 0 || K <= 0 || group <= 0 || K % group != 0 || group % 8 != 0) { + return false; + } + return launch(kop::gemv_q4_, {in(a), in(qw), in(scales), out(y)}, + gemv_q4_params{static_cast(N), static_cast(K), + static_cast(group)}, + policy::row_reduce(N, K)); +} + +// One decode step's k, v (each [n_kv_heads, D]) into row `pos` of a +// [n_kv_heads, kv_max, D] cache, f32 or bf16. +inline bool kv_append(span Kc, span Vc, span k_new, span v_new, int64_t pos, + int64_t kv_max, int64_t n_kv_heads, int64_t D, + bool kv_bf16 = false) { + if ((D != 64 && D != 128) || n_kv_heads <= 0) return false; + return launch(kv_bf16 ? kop::kv_append_bf16_ : kop::kv_append_, + {out(Kc), out(Vc), in(k_new), in(v_new)}, + kv_append_params{static_cast(pos), + static_cast(kv_max * D)}, + policy::per_head(n_kv_heads, 1, D)); +} + +// A prefill's k, v (each [n_kv_heads, T, D]) into cache rows [pos0, pos0 + T). +inline bool kv_fill(span Kc, span Vc, span K, span V, int64_t T, int64_t kv_max, + int64_t n_kv_heads, int64_t D, bool kv_bf16 = false, + int64_t pos0 = 0) { + if ((D != 64 && D != 128) || n_kv_heads <= 0 || T <= 0) return false; + return launch(kv_bf16 ? kop::kv_fill_bf16_ : kop::kv_fill_, + {out(Kc), out(Vc), in(K), in(V)}, + kv_fill_params{static_cast(T), + static_cast(kv_max * D), + static_cast(pos0)}, + policy::per_head(n_kv_heads, T, D)); +} + +// Head-major [H, T, D] -> token-major [T, H*D]: split_heads' inverse. +inline bool merge_heads(span src, span dst, int64_t T, int64_t H, int64_t D) { + if (T <= 0 || H <= 0 || D <= 0) return false; + return launch(kop::merge_heads_, {in(src), out(dst)}, + merge_heads_params{static_cast(T), + static_cast(H), + static_cast(D)}, + policy::per_head(H, T, D)); +} + +// Token-major [T, ld] -> head-major [H, T, D] from column block `off`, adding +// the optional per-head bias [H, D] (a null view: none). Backend-own: the +// kernels disagree on how "no bias" is said. +TL_GPU_DETECT_OWN(split_heads) +template +inline bool split_heads(span src, span bias, span dst, int64_t T, int64_t ld, + int64_t off, int64_t H, int64_t D) { + if (T <= 0 || H <= 0 || D <= 0) return false; + if constexpr (detail::owns_split_heads::value) { + return detail::ran("split_heads", Own::split_heads(src, bias, dst, T, ld, off, H, D)); + } else { + return false; + } +} + +// The argmax of a length-n vector, the smallest index on ties: greedy decoding +// reads one int back rather than the logits. Drains the queue. Backend-own: +// the result's staging buffer and its read-back are the backend's. +TL_GPU_DETECT_OWN(argmax) +template +inline bool argmax(span a, int64_t n, int64_t* out_idx) { + if (!a || n <= 0 || !out_idx) return false; + if constexpr (detail::owns_argmax::value) { + return detail::ran("argmax", Own::argmax(a, n, out_idx)); + } else { + return false; + } +} + +// 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) { + kop k = kop::tanh_; + switch (op) { + case unary_ext_op::tanh_: k = kop::tanh_; break; + case unary_ext_op::sin_: k = kop::sin_; break; + case unary_ext_op::cos_: k = kop::cos_; break; + } + return unary(k, a, o, n, scale, offset); +} + +// ---- ops a backend runs its own way. Same contract as the rest — views in, +// false when it declines — but the algorithm under each differs by backend +// (a scatter into a zeroed buffer against a gather, shape metadata in a device +// buffer against a params block, one kernel against a split and a combine), so +// each backend writes its own over its `dispatch`, as a member of `own`. + +// N-D broadcast binary into a contiguous out of out_shape: strides are +// elements per axis (0 on a broadcast axis), rank <= 8. +TL_GPU_DETECT_OWN(binary_bcast_nd) +template +inline bool binary_bcast_nd(kop op, span a, const int64_t* a_strides, span b, + const int64_t* b_strides, span o, + const int64_t* out_shape, int rank, int64_t n, + float scale, float offset) { + if constexpr (detail::owns_binary_bcast_nd::value) { + return detail::ran("binary_bcast_nd", Own::binary_bcast_nd(op, a, a_strides, b, b_strides, o, out_shape, rank, n, scale, offset)); + } else { + return false; + } +} + +// out = cond ? a : b over out_shape, each operand through its own strides. +TL_GPU_DETECT_OWN(where_nd) +template +inline bool where_nd(span cond, const int64_t* c_strides, span a, + const int64_t* a_strides, span b, const int64_t* b_strides, + span o, const int64_t* out_shape, int rank, int64_t n) { + if constexpr (detail::owns_where_nd::value) { + return detail::ran("where_nd", Own::where_nd(cond, c_strides, a, a_strides, b, b_strides, o, out_shape, rank, n)); + } else { + return false; + } +} + +// clone()'s strided gather: a through a_strides into a contiguous out. +TL_GPU_DETECT_OWN(copy_nd) +template +inline bool copy_nd(span a, const int64_t* a_strides, span o, + const int64_t* out_shape, int rank, int64_t n) { + if constexpr (detail::owns_copy_nd::value) { + return detail::ran("copy_nd", Own::copy_nd(a, a_strides, o, out_shape, rank, n)); + } else { + return false; + } +} + +// Un-broadcast a gradient: every out element sums the `a` elements that +// broadcast onto it (`acc` the per-axis accumulation strides). +TL_GPU_DETECT_OWN(sum_to) +template +inline bool sum_to(span a, const int64_t* a_shape, const int64_t* a_strides, + const int64_t* acc, int rank, int64_t out_n, + int64_t reduced_n, span o) { + if constexpr (detail::owns_sum_to::value) { + return detail::ran("sum_to", Own::sum_to(a, a_shape, a_strides, acc, rank, out_n, reduced_n, o)); + } else { + return false; + } +} + +// `a` (contiguous) into a zero buffer of out_shape, shifted by `before` along +// `axis`. No epilogue. +TL_GPU_DETECT_OWN(pad) +template +inline bool pad(span a, span o, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, int64_t before, + int64_t n, int64_t out_n) { + if constexpr (detail::owns_pad::value) { + return detail::ran("pad", Own::pad(a, o, a_shape, out_shape, rank, axis, before, n, out_n)); + } else { + return false; + } +} + +// unfold's inverse: scatter-add a's trailing windows into out_shape. +TL_GPU_DETECT_OWN(fold) +template +inline bool fold(span a, span o, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, int64_t step, + int64_t n, int64_t out_n) { + if constexpr (detail::owns_fold::value) { + return detail::ran("fold", Own::fold(a, o, a_shape, out_shape, rank, axis, step, n, out_n)); + } else { + return false; + } +} + +// One part of a concat: `a` into out at `before` along `axis`. +TL_GPU_DETECT_OWN(concat_part) +template +inline bool concat_part(span a, span o, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t before, int64_t n) { + if constexpr (detail::owns_concat_part::value) { + return detail::ran("concat_part", Own::concat_part(a, o, a_shape, out_shape, rank, axis, before, n)); + } else { + return false; + } +} + +// index_select's dual: out row idx[i] += values row i, for k rows. +TL_GPU_DETECT_OWN(index_add) +template +inline bool index_add(span idx, span values, span o, int64_t row_size, + int64_t k, int64_t out_n) { + if constexpr (detail::owns_index_add::value) { + return detail::ran("index_add", Own::index_add(idx, values, o, row_size, k, out_n)); + } else { + return false; + } +} + +// One-hot scatter into a new trailing axis: out[i, idx[i]] = values[i]. +TL_GPU_DETECT_OWN(scatter_to_axis) +template +inline bool scatter_to_axis(span idx, span values, span o, int64_t n, + int64_t size) { + if constexpr (detail::owns_scatter_to_axis::value) { + return detail::ran("scatter_to_axis", Own::scatter_to_axis(idx, values, o, n, size)); + } else { + return false; + } +} + +// C(m,n) = (A @ B) * scale + offset. lda/ldb are row strides and the trans +// flags read a transposed view in place. +TL_GPU_DETECT_OWN(gemm) +template +inline bool gemm(span a, int64_t lda, bool ta, span b, int64_t ldb, bool tb, + span o, int64_t m, int64_t n, int64_t k, float scale, + float offset) { + if constexpr (detail::owns_gemm::value) { + return detail::ran("gemm", Own::gemm(a, lda, ta, b, ldb, tb, o, m, n, k, scale, offset)); + } else { + return false; + } +} + +// `batch` GEMMs in one launch, slices sa / sb elements apart, with an optional +// row bias added in the store. +TL_GPU_DETECT_OWN(gemm_batched) +template +inline bool gemm_batched(span a, int64_t lda, bool ta, int64_t sa, span b, + int64_t ldb, bool tb, int64_t sb, span o, int64_t m, + int64_t n, int64_t k, int64_t batch, float scale, + float offset, span bias = {}) { + if constexpr (detail::owns_gemm_batched::value) { + return detail::ran("gemm_batched", Own::gemm_batched(a, lda, ta, sa, b, ldb, tb, sb, o, m, n, k, batch, scale, offset, bias)); + } else { + return false; + } +} + +// addmm's shape: a GEMM with its row bias added in the store. +TL_GPU_DETECT_OWN(gemm_bias) +template +inline bool gemm_bias(span a, int64_t lda, bool ta, span b, int64_t ldb, + bool tb, span bias, span o, int64_t m, int64_t n, + int64_t k, float scale, float offset) { + if constexpr (detail::owns_gemm_bias::value) { + return detail::ran("gemm_bias", Own::gemm_bias(a, lda, ta, b, ldb, tb, bias, o, m, n, k, scale, offset)); + } else { + return false; + } +} + +// Rotary embedding over [rows, T, D] at position `pos`, adding the optional +// per-row bias first. +TL_GPU_DETECT_OWN(rope) +template +inline bool rope(span x, span o, int64_t rows, int64_t T, int64_t D, + int64_t pos, float base, span bias = {}) { + if constexpr (detail::owns_rope::value) { + return detail::ran("rope", Own::rope(x, o, rows, T, D, pos, base, bias)); + } else { + return false; + } +} + +// Layer norm's pullback: dx, and dg / db reduced over rows in chunks. +TL_GPU_DETECT_OWN(layer_norm_bwd) +template +inline bool layer_norm_bwd(span x, span g, span dy, span dx, span dg, span db, + span stats, span partials, int64_t rows, + int64_t cols, int64_t per_chunk, int64_t chunks, + float eps) { + if constexpr (detail::owns_layer_norm_bwd::value) { + return detail::ran("layer_norm_bwd", Own::layer_norm_bwd(x, g, dy, dx, dg, db, stats, partials, rows, cols, per_chunk, chunks, eps)); + } else { + return false; + } +} + +// ---- the LLM path. + +// y[1,n] = a[1,k] . B[k,n], the weight column-major, f32 or bf16. +TL_GPU_DETECT_OWN(gemv_f32) +template +inline bool gemv_f32(span a, span B, span y, int64_t n, int64_t k) { + if constexpr (detail::owns_gemv_f32::value) { + return detail::ran("gemv_f32", Own::gemv_f32(a, B, y, n, k)); + } else { + return false; + } +} + +TL_GPU_DETECT_OWN(gemv_bf16) +template +inline bool gemv_bf16(span a, span B, span y, int64_t n, int64_t k) { + if constexpr (detail::owns_gemv_bf16::value) { + return detail::ran("gemv_bf16", Own::gemv_bf16(a, B, y, n, k)); + } else { + return false; + } +} + +// C[M,N] = A[M,K] . B[N,K]^T with B bf16: a prefill chunk's projection. +TL_GPU_DETECT_OWN(gemm_bf16_nt) +template +inline bool gemm_bf16_nt(span A, span B, span C, int64_t M, int64_t N, + int64_t K) { + if constexpr (detail::owns_gemm_bf16_nt::value) { + return detail::ran("gemm_bf16_nt", Own::gemm_bf16_nt(A, B, C, M, N, K)); + } else { + return false; + } +} + +// One decode step: q [n_q_heads, D] against a [n_kv_heads, kv_max, D] cache +// read over [0, ctx). GQA by the head ratio; the cache is f32 or bf16. +TL_GPU_DETECT_OWN(attn_decode) +template +inline bool attn_decode(span q, span K, span V, span o, int64_t n_q_heads, + int64_t n_kv_heads, int64_t ctx, int64_t kv_max, + int64_t D, float scale, bool kv_bf16 = false) { + if constexpr (detail::owns_attn_decode::value) { + return detail::ran("attn_decode", Own::attn_decode(q, K, V, o, n_q_heads, n_kv_heads, ctx, kv_max, D, scale, kv_bf16)); + } else { + return false; + } +} + +// Causal prefill: q, out [n_q_heads, T, D] against the cache rows +// [0, pos0 + T); query t sits at absolute position pos0 + t. +TL_GPU_DETECT_OWN(attn_prefill) +template +inline bool attn_prefill(span q, span K, span V, span o, int64_t n_q_heads, + int64_t n_kv_heads, int64_t T, int64_t kv_max, + int64_t D, float scale, bool kv_bf16 = false, + int64_t pos0 = 0) { + if constexpr (detail::owns_attn_prefill::value) { + return detail::ran("attn_prefill", Own::attn_prefill(q, K, V, o, n_q_heads, n_kv_heads, T, kv_max, D, scale, kv_bf16, pos0)); + } else { + return false; + } +} + +// The causal prefill's pullback, query half then key/value half. +TL_GPU_DETECT_OWN(attn_prefill_dq) +template +inline bool attn_prefill_dq(span q, span K, span V, span dO, span O, span dq, + span stats, int64_t H, int64_t T, int64_t D, + float scale) { + if constexpr (detail::owns_attn_prefill_dq::value) { + return detail::ran("attn_prefill_dq", Own::attn_prefill_dq(q, K, V, dO, O, dq, stats, H, T, D, scale)); + } else { + return false; + } +} + +TL_GPU_DETECT_OWN(attn_prefill_dkv) +template +inline bool attn_prefill_dkv(span q, span K, span V, span dO, span stats, + span dK, span dV, int64_t H, int64_t T, int64_t D, + float scale) { + if constexpr (detail::owns_attn_prefill_dkv::value) { + return detail::ran("attn_prefill_dkv", Own::attn_prefill_dkv(q, K, V, dO, stats, dK, dV, H, T, D, scale)); + } else { + return false; + } +} + +// ---- 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. +TL_GPU_DETECT_OWN(rope_dpos) +template +inline bool rope_dpos(span x, span o, int64_t rows, int64_t T, int64_t D, + span d_pos, float base, span bias = {}) { + if constexpr (detail::owns_rope_dpos::value) { + return detail::ran("rope_dpos", Own::rope_dpos(x, o, rows, T, D, d_pos, base, bias)); + } else { + return false; + } +} + +TL_GPU_DETECT_OWN(kv_append_dpos) +template +inline bool kv_append_dpos(span Kc, span Vc, span k_new, span v_new, span d_pos, + int64_t kv_max, int64_t n_kv_heads, int64_t D) { + if constexpr (detail::owns_kv_append_dpos::value) { + return detail::ran("kv_append_dpos", Own::kv_append_dpos(Kc, Vc, k_new, v_new, d_pos, kv_max, n_kv_heads, D)); + } else { + return false; + } +} + +TL_GPU_DETECT_OWN(attn_decode_dpos) +template +inline bool attn_decode_dpos(span q, span K, span V, span o, int64_t n_q_heads, + int64_t n_kv_heads, span d_pos, int64_t kv_max, + int64_t D, float scale, span partials) { + if constexpr (detail::owns_attn_decode_dpos::value) { + return detail::ran("attn_decode_dpos", Own::attn_decode_dpos(q, K, V, o, n_q_heads, n_kv_heads, d_pos, kv_max, D, scale, partials)); + } else { + return false; + } +} + +// *d_pos += 1 on the device. +TL_GPU_DETECT_OWN(incr_u32) +template +inline bool incr_u32(span d_pos) { + if constexpr (detail::owns_incr_u32::value) { + return detail::ran("incr_u32", Own::incr_u32(d_pos)); + } else { + return false; + } +} + +#undef TL_GPU_DETECT_OWN + +} // namespace gpu +} // namespace tl diff --git a/include/kv_cache.h b/include/kv_cache.h index 4c86513..5417e50 100644 --- a/include/kv_cache.h +++ b/include/kv_cache.h @@ -44,50 +44,51 @@ struct kv_cache { return K.native && V.native; // a heap fallback has no device buffer } - // k_new/v_new: [n_kv_heads, D] device buffers (this step's projected k, v). - bool append(void* k_new, void* v_new) { + // k_new/v_new: [n_kv_heads, D] device views (this step's projected k, v). + bool append(gpu::span k_new, gpu::span v_new) { if (pos >= max_ctx) return false; - if (!gpu::kv_append(K.native, V.native, k_new, v_new, pos, max_ctx, - n_kv_heads, D, kv_bf16)) { + if (!gpu::kv_append(K.device_span(), V.device_span(), k_new, v_new, pos, + max_ctx, n_kv_heads, D, kv_bf16)) { return false; } pos++; return true; } - // q/out: [n_q_heads, D] device buffers. Attends over the cached prefix. - bool attn(void* q, void* out, int64_t n_q_heads, float scale) { - return gpu::attn_decode(q, K.native, V.native, out, n_q_heads, n_kv_heads, - pos, max_ctx, D, scale, kv_bf16); + // q/out: [n_q_heads, D] device views. Attends over the cached prefix. + bool attn(gpu::span q, gpu::span out, int64_t n_q_heads, float scale) { + return gpu::attn_decode(q, K.device_span(), V.device_span(), out, n_q_heads, + n_kv_heads, pos, max_ctx, D, scale, kv_bf16); } // T tokens at once: k_src/v_src [n_kv_heads, T, D] appended at `pos`, and // the causal attention of q/out [n_q_heads, T, D] over everything cached // before them. Leaves pos advanced by T. - bool prefill(void* q, void* k_src, void* v_src, void* out, int64_t T, - int64_t n_q_heads, float scale) { + bool prefill(gpu::span q, gpu::span k_src, gpu::span v_src, gpu::span out, + int64_t T, int64_t n_q_heads, float scale) { if (T <= 0 || pos + T > max_ctx) return false; - if (!gpu::kv_fill(K.native, V.native, k_src, v_src, T, max_ctx, n_kv_heads, - D, kv_bf16, pos)) { + if (!gpu::kv_fill(K.device_span(), V.device_span(), k_src, v_src, T, max_ctx, + n_kv_heads, D, kv_bf16, pos)) { return false; } const int64_t p0 = pos; pos += T; - return gpu::attn_prefill(q, K.native, V.native, out, n_q_heads, n_kv_heads, - T, max_ctx, D, scale, kv_bf16, p0); + return gpu::attn_prefill(q, K.device_span(), V.device_span(), out, + n_q_heads, n_kv_heads, T, max_ctx, D, scale, kv_bf16, + p0); } // Graph-capture forms (f32 cache only): the position is *d_pos. - bool append_dpos(void* k_new, void* v_new, void* d_pos) { - return gpu::kv_append_dpos(K.native, V.native, k_new, v_new, d_pos, max_ctx, - n_kv_heads, D); + bool append_dpos(gpu::span k_new, gpu::span v_new, gpu::span d_pos) { + return gpu::kv_append_dpos(K.device_span(), V.device_span(), k_new, v_new, + d_pos, max_ctx, n_kv_heads, D); } // The split-KV partials a captured graph bakes in are this cache's own (a // shared scratch could be freed under a live graph), sized once from the // capacity on the first call. storage dpos_partials; - bool attn_dpos(void* q, void* out, int64_t n_q_heads, void* d_pos, - float scale) { + bool attn_dpos(gpu::span q, gpu::span out, int64_t n_q_heads, + gpu::span d_pos, float scale) { if (!dpos_partials.native) { const int64_t bytes = gpu::attn_dpos_partials_bytes(n_q_heads, max_ctx, D); @@ -95,9 +96,9 @@ struct kv_cache { dpos_partials = storage::make(bytes / 4); if (!dpos_partials.native) return false; } - return gpu::attn_decode_dpos(q, K.native, V.native, out, n_q_heads, - n_kv_heads, d_pos, max_ctx, D, scale, - dpos_partials.native); + return gpu::attn_decode_dpos(q, K.device_span(), V.device_span(), out, + n_q_heads, n_kv_heads, d_pos, max_ctx, D, scale, + dpos_partials.device_span()); } }; diff --git a/include/metal.h b/include/metal.h index db3151b..2ce1835 100644 --- a/include/metal.h +++ b/include/metal.h @@ -16,11 +16,13 @@ // - Kernels JIT-compile once from the #embed'd MSL source on first GPU // dispatch. Editing metal_kernels.metal requires rebuilding the host. // -// On non-Apple builds everything is an inline stub returning false/null, so -// callers carry no platform conditionals. +// The whole header is gated on __APPLE__: elsewhere it declares nothing, and +// gpu.h selects another backend (or gpu_null.h). #include +#include "gpu_abi.h" + #ifdef __APPLE__ #include @@ -36,65 +38,15 @@ extern "C" void* MTLCreateSystemDefaultDevice(void); extern "C" void* objc_autoreleasePoolPush(void); extern "C" void objc_autoreleasePoolPop(void*); -#endif - namespace tl { namespace metal { -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 - where_nd, copy_nd, // N-D select / clone()'s strided gather - gt_, lt_, ge_, le_, eq_, ne_, // comparisons -- cmp_op maps onto these - tanh_, sin_, cos_, // unary_ext_op maps onto these - clamp_, sum_to_, sum_to_blocked_, // dedicated ops, mirroring cuda.h's own - concat_part_, rope_, // ditto -- Tensor.concat / RoPE's own dispatch - 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_, swiglu_, split_heads_, merge_heads_, // the decode step's rest - gemv_f32_, gemv_bf16_, gemv_q4_, // decode GEMVs, per weight dtype - gemv_combine_, // their split-K partials - gemv_bf16_row_, // ... and the [N,K] weight layout's own - gemm_bf16_nt_, gemm_bf16_nt32_, // the prefill's bf16 GEMM, per M tile - gather_axis_, row_logsumexp_, xent_bwd_, adam_step_ // cross-entropy, Adam -}; - -// Comparisons (gt/lt/ge/le/eq/ne) are deliberately NOT kop values: kop is -// called unconditionally through this file's own pso_()-based binary(), -// which throws rather than declining an enum value it has no MSL kernel -// for -- fine for ops every backend already implements, not for a -// CUDA-only addition landing ahead of its Metal/WebGPU kernels. compare() -// below is its own small vocabulary so an unimplemented backend can just -// return false, same as index_select/index_add/scatter_axis/sum_to. -enum class cmp_op { gt, lt, ge, le, eq, ne }; - -// tanh_/sin_/cos_: same reasoning as cmp_op above -- a CUDA-only addition, -// so its own vocabulary rather than a new kop. clamp (2 node-specific -// scalars, no epilogue) gets its own dedicated function below instead of -// an enum value, same as index_select/sum_to's own dedicated functions. -enum class unary_ext_op { tanh_, sin_, cos_ }; - -// Tensor-scalar ops: pow(x, s) and the comparisons against a scalar, with s a -// kernel argument instead of a rank-0 operand buffer (an allocation and an -// upload per call). Own vocabulary, like cmp_op. -enum class scalar_op { pow, gt, lt, ge, le, eq, ne }; - -#ifdef __APPLE__ +// The op vocabulary is the shared layer's (gpu_abi.h); the names stay reachable +// as metal::kop for the code that spelled them that way. +using kop = gpu::kop; +using cmp_op = gpu::cmp_op; +using unary_ext_op = gpu::unary_ext_op; +using scalar_op = gpu::scalar_op; struct mtl_size { unsigned long w, h, d; @@ -250,6 +202,7 @@ struct context { case kop::kv_fill_bf16_: return "kv_fill_bf16_"; case kop::argmax_: return "argmax_"; case kop::rmsnorm_: return "rmsnorm_"; + case kop::add_rmsnorm_: return "add_rmsnorm_"; case kop::swiglu_: return "swiglu_"; case kop::split_heads_: return "split_heads_"; case kop::merge_heads_: return "merge_heads_"; @@ -327,6 +280,75 @@ struct context { inline bool available() { return context::get().device != nullptr; } +// The ops this backend runs its own way: a different algorithm, several +// kernels, or a kernel whose ABI is its own. gpu_ops.h forwards to whichever of +// these exist (TL_GPU_DETECT_OWN) and answers false for the rest, so a backend +// declares what it has and nothing else. Defined below, among their helpers. +struct own { + static bool split_heads(gpu::span src, gpu::span bias, gpu::span dst, + int64_t T, int64_t ld, int64_t off, int64_t H, + int64_t D); + static bool argmax(gpu::span a, int64_t n, int64_t* out_idx); + static bool binary_bcast_nd(kop op, gpu::span a, const int64_t* a_strides, + gpu::span b, const int64_t* b_strides, + gpu::span out, const int64_t* out_shape, int rank, + int64_t n, float scale, float offset); + static bool where_nd(gpu::span cond, const int64_t* c_strides, gpu::span a, + const int64_t* a_strides, gpu::span b, + const int64_t* b_strides, gpu::span out, + const int64_t* out_shape, int rank, int64_t n); + static bool copy_nd(gpu::span a, const int64_t* a_strides, gpu::span out, + const int64_t* out_shape, int rank, int64_t n); + static bool sum_to(gpu::span a, const int64_t* a_shape, + const int64_t* a_strides, const int64_t* acc, int rank, + int64_t out_n, int64_t reduced_n, gpu::span out); + static bool pad(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, int64_t before, + int64_t n, int64_t out_n); + static bool fold(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, int64_t step, + int64_t n, int64_t out_n); + static bool concat_part(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t before, int64_t n); + static bool index_add(gpu::span idx, gpu::span values, gpu::span out, + int64_t row_size, int64_t k, int64_t out_n); + static bool scatter_to_axis(gpu::span idx, gpu::span values, gpu::span out, + int64_t n, int64_t size); + static bool gemm(gpu::span a, int64_t lda, bool ta, gpu::span b, int64_t ldb, + bool tb, gpu::span out, int64_t m, int64_t n, int64_t k, + float scale, float offset); + static bool gemv_f32(gpu::span a, gpu::span B, gpu::span y, int64_t n, + int64_t k); + static bool gemv_bf16(gpu::span a, gpu::span B, gpu::span y, int64_t n, + int64_t k); + static bool attn_decode(gpu::span q, gpu::span K, gpu::span V, gpu::span out, + int64_t n_q_heads, int64_t n_kv_heads, int64_t ctx, + int64_t kv_max, int64_t D, float scale, + bool kv_bf16 = false); + static bool attn_prefill(gpu::span q, gpu::span K, gpu::span V, gpu::span out, + int64_t n_q_heads, int64_t n_kv_heads, int64_t T, + int64_t kv_max, int64_t D, float scale, + bool kv_bf16 = false, int64_t pos0 = 0); + static bool attn_prefill_dq(gpu::span q, gpu::span K, gpu::span V, + gpu::span dO, gpu::span O, gpu::span dq, + gpu::span stats, int64_t H, int64_t T, int64_t D, + float scale); + static bool attn_prefill_dkv(gpu::span q, gpu::span K, gpu::span V, + gpu::span dO, gpu::span stats, gpu::span dK, + gpu::span dV, int64_t H, int64_t T, int64_t D, + float scale); + static bool gemm_bf16_nt(gpu::span A, gpu::span B, gpu::span C, int64_t M, + int64_t N, int64_t K); + static bool rope(gpu::span x, gpu::span out, int64_t rows, int64_t T, + int64_t D, int64_t pos, float base, gpu::span bias = {}); + static bool layer_norm_bwd(gpu::span x, gpu::span g, gpu::span dy, + gpu::span dx, gpu::span dg, gpu::span db, + gpu::span stats, gpu::span partials, int64_t rows, + int64_t cols, int64_t per_chunk, int64_t chunks, + float eps); +}; + inline bool pending() { return context::get().pending; } // End the batch: commit and block until the GPU finishes (MLX-style eval). @@ -431,97 +453,32 @@ inline void dispatch_grid_(objc::id enc, mtl_size grid, mtl_size tg) { } } -struct ew_params { - float scale; - float offset; - uint32_t n; -}; - -inline void dispatch_(objc::id enc, const ew_params& p, - unsigned long params_index) { - objc::send(enc, "setBytes:length:atIndex:", static_cast(&p), - static_cast(sizeof(p)), params_index); - unsigned long groups = (p.n + 255ul) / 256ul; - dispatch_grid_(enc, {groups, 1, 1}, {256, 1, 1}); -} - } // namespace detail_ -// Contiguous elementwise dispatches; offsets in bytes. Epilogue (scale, -// offset) applies inside the kernel. Encodes without committing. -inline bool binary(kop op, void* a, int64_t ao, void* b, int64_t bo, void* out, - int64_t oo, int64_t n, float scale, float offset) { - auto& c = context::get(); - if (!c.device) return false; - c.bind_(op); - objc::send(c.enc, "setBuffer:offset:atIndex:", a, - static_cast(ao), 0ul); - objc::send(c.enc, "setBuffer:offset:atIndex:", b, - static_cast(bo), 1ul); - objc::send(c.enc, "setBuffer:offset:atIndex:", out, - static_cast(oo), 2ul); - detail_::dispatch_(c.enc, {scale, offset, static_cast(n)}, 3ul); - return true; -} - -inline bool unary(kop op, void* a, int64_t ao, void* out, int64_t oo, - int64_t n, float scale, float offset) { +// The device core's one way to run a kernel (gpu_abi.h has the contract): +// view i goes to buffer index i at its byte offset, the params follow at index +// n, and the grid is threadgroups x threads. Encodes without committing. +// Unified memory, so an arg's access says nothing this backend has to act on. +inline bool dispatch(kop k, const gpu::arg* args, size_t n, const void* params, + size_t params_bytes, const gpu::grid& g) { auto& c = context::get(); - if (!c.device) return false; - c.bind_(op); - objc::send(c.enc, "setBuffer:offset:atIndex:", a, - static_cast(ao), 0ul); - objc::send(c.enc, "setBuffer:offset:atIndex:", out, - static_cast(oo), 1ul); - detail_::dispatch_(c.enc, {scale, offset, static_cast(n)}, 2ul); + if (!c.device || !*context::kernel_name_(k)) return false; + c.bind_(k); + for (size_t i = 0; i < n; i++) { + objc::send(c.enc, "setBuffer:offset:atIndex:", args[i].s.buf, + static_cast(args[i].s.off), + static_cast(i)); + } + objc::send(c.enc, "setBytes:length:atIndex:", params, + static_cast(params_bytes), + static_cast(n)); + detail_::dispatch_grid_(c.enc, {g.gx, g.gy, g.gz}, {g.tx, g.ty, g.tz}); return true; } namespace detail_ { -struct ew_bcast_params { - float scale; - float offset; - uint32_t M, N; - uint32_t ars, acs, brs, bcs; // per-operand row/col strides (elements) -}; } // namespace detail_ -// Rank-2 broadcast binary: out[r,c] = f(a[r*ars+c*acs], b[r*brs+c*bcs]) into a -// contiguous [M,N] output. One kernel covers every rank-2 broadcast (row -// vector, column vector, per-row scalar), keeping bias/gamma/beta chains on -// the GPU instead of a CPU fallback that drains the pipeline. Encodes without -// committing, like binary(). -inline bool binary_bcast(kop op, void* a, int64_t ao, int64_t ars, int64_t acs, - void* b, int64_t bo, int64_t brs, int64_t bcs, - void* out, int64_t oo, int64_t m, int64_t n, - float scale, float offset) { - auto& c = context::get(); - if (!c.device) return false; - c.bind_(op); - objc::send(c.enc, "setBuffer:offset:atIndex:", a, - static_cast(ao), 0ul); - objc::send(c.enc, "setBuffer:offset:atIndex:", b, - static_cast(bo), 1ul); - objc::send(c.enc, "setBuffer:offset:atIndex:", out, - static_cast(oo), 2ul); - detail_::ew_bcast_params p{ - scale, - offset, - static_cast(m), - static_cast(n), - static_cast(ars), - static_cast(acs), - static_cast(brs), - static_cast(bcs)}; - objc::send(c.enc, "setBytes:length:atIndex:", static_cast(&p), - static_cast(sizeof(p)), 3ul); - detail_::dispatch_grid_(c.enc, - {(static_cast(n) + 31ul) / 32ul, - (static_cast(m) + 7ul) / 8ul, 1}, - {32, 8, 1}); - return true; -} - namespace detail_ { struct gemm_params { @@ -530,20 +487,13 @@ struct gemm_params { float scale, offset; }; -struct reduce_params { - uint32_t rows, cols; - float scale, offset; -}; - -struct layer_norm_params { - uint32_t rows, cols; - float eps, scale, offset; -}; - inline void set_buf_(objc::id enc, void* buf, int64_t off, unsigned long idx) { objc::send(enc, "setBuffer:offset:atIndex:", buf, static_cast(off), idx); } +inline void set_buf_(objc::id enc, gpu::span s, unsigned long idx) { + set_buf_(enc, s.buf, s.off, idx); +} template inline void set_bytes_(objc::id enc, const P& p, unsigned long idx) { @@ -556,9 +506,9 @@ inline void set_bytes_(objc::id enc, const P& p, unsigned long idx) { // C(m,n) = (A @ B) * scale + offset. lda/ldb are row strides; trans flags // let a transposed view be read in place. Buffers are raw MTLBuffers; byte // offsets fold the view offset in. Encodes without committing. -inline bool gemm(void* a, int64_t ao, int64_t lda, bool ta, void* b, - int64_t bo, int64_t ldb, bool tb, void* out, int64_t oo, - int64_t m, int64_t n, int64_t k, float scale, float offset) { +inline bool own::gemm(gpu::span a, int64_t lda, bool ta, gpu::span b, + int64_t ldb, bool tb, gpu::span out, int64_t m, int64_t n, + int64_t k, float scale, float offset) { auto& c = context::get(); if (!c.device) return false; // Dispatch ladder. STEEL (BN=64 bands) covers NN and single-transposed @@ -599,9 +549,9 @@ inline bool gemm(void* a, int64_t ao, int64_t lda, bool ta, void* b, gy = (static_cast(m) + bm - 1) / bm; } c.bind_(kk_); - detail_::set_buf_(c.enc, a, ao, 0ul); - detail_::set_buf_(c.enc, b, bo, 1ul); - detail_::set_buf_(c.enc, out, oo, 2ul); + detail_::set_buf_(c.enc, a, 0ul); + detail_::set_buf_(c.enc, b, 1ul); + detail_::set_buf_(c.enc, out, 2ul); detail_::gemm_params p{static_cast(m), static_cast(n), static_cast(k), static_cast(lda), static_cast(ldb), ta ? 1u : 0u, @@ -614,55 +564,6 @@ inline bool gemm(void* a, int64_t ao, int64_t lda, bool ta, void* b, return true; } -// Batched GEMM in one launch: not on this backend yet — array.h's batched dot -// loops gemm per slice when this declines (CUDA folds the batch into its grid). -inline bool gemm_batched(void*, int64_t, int64_t, bool, int64_t, void*, - int64_t, int64_t, bool, int64_t, void*, int64_t, - int64_t, int64_t, int64_t, int64_t, float, float, - void* = nullptr, int64_t = 0) { - return false; -} - -// Row-wise op over the last axis (cols): softmax writes rows×cols; row_sum/ -// row_max write one value per row (rows), with the affine epilogue. -inline bool row_op(kop op, void* in, int64_t io, void* out, int64_t oo, - int64_t rows, int64_t cols, float scale, float offset) { - auto& c = context::get(); - if (!c.device) return false; - c.bind_(op); - detail_::set_buf_(c.enc, in, io, 0ul); - detail_::set_buf_(c.enc, out, oo, 1ul); - detail_::reduce_params p{static_cast(rows), - static_cast(cols), scale, offset}; - objc::send(c.enc, "setBytes:length:atIndex:", static_cast(&p), - static_cast(sizeof(p)), 2ul); - detail_::dispatch_grid_(c.enc, {static_cast(rows), 1, 1}, - {256, 1, 1}); - return true; -} - -// Layer norm over the last axis: out = (x - mu) · 1/sqrt(var + eps) · g + b per -// row, affine epilogue; g and b are contiguous d-vectors. One threadgroup per -// row, like row_op. -inline bool layer_norm(void* x, int64_t xo, void* g, int64_t go, void* b, - int64_t bo, void* out, int64_t oo, int64_t rows, - int64_t cols, float eps, float scale, float offset) { - auto& c = context::get(); - if (!c.device) return false; - c.bind_(kop::layer_norm_); - detail_::set_buf_(c.enc, x, xo, 0ul); - detail_::set_buf_(c.enc, g, go, 1ul); - detail_::set_buf_(c.enc, b, bo, 2ul); - detail_::set_buf_(c.enc, out, oo, 3ul); - detail_::layer_norm_params p{static_cast(rows), - static_cast(cols), eps, scale, offset}; - objc::send(c.enc, "setBytes:length:atIndex:", static_cast(&p), - static_cast(sizeof(p)), 4ul); - detail_::dispatch_grid_(c.enc, {static_cast(rows), 1, 1}, - {256, 1, 1}); - return true; -} - // Rank cap for pad_/fold_'s GPU dispatch — matches cuda.h's kPadFoldMaxRank // and metal_kernels.metal's own copy (an MSL kernel can't see a host-side // C++ constant), and bounds pad_fold_params' fixed-size arrays. @@ -719,11 +620,11 @@ inline bool dispatch_pad_fold_(kop op, void* a_native, int64_t ao, // gpu_pad_/gpu_fold_, so the metal_kernels.metal side derives a_strides from // a_shape rather than have them uploaded. Encodes without committing, like // every other dispatch above. -inline bool pad(void* a_native, int64_t ao, void* out_native, int64_t oo, - const int64_t* a_shape, const int64_t* out_shape, int rank, - int axis, int64_t before, int64_t n, int64_t out_n) { +inline bool own::pad(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t before, int64_t n, int64_t out_n) { (void)n; - return detail_::dispatch_pad_fold_(kop::pad, a_native, ao, out_native, oo, + return detail_::dispatch_pad_fold_(kop::pad, a.buf, a.off, out.buf, out.off, a_shape, out_shape, rank, rank, axis, before, out_n); } @@ -733,20 +634,16 @@ inline bool pad(void* a_native, int64_t ao, void* out_native, int64_t oo, // (fold's own out has one fewer axis than `a`); `a`'s last dim is the // sliding window (size a_shape[rank-1]), and `a`'s `axis` dim is the window // count. -inline bool fold(void* a_native, int64_t ao, void* out_native, int64_t oo, - const int64_t* a_shape, const int64_t* out_shape, int rank, - int axis, int64_t step, int64_t n, int64_t out_n) { +inline bool own::fold(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t step, int64_t n, int64_t out_n) { (void)n; - return detail_::dispatch_pad_fold_(kop::fold, a_native, ao, out_native, oo, + return detail_::dispatch_pad_fold_(kop::fold, a.buf, a.off, out.buf, out.off, a_shape, out_shape, rank, rank - 1, axis, step, out_n); } namespace detail_ { -struct gather_params { - uint32_t row_size; - uint32_t n; -}; struct index_add_params { uint32_t row_size; uint32_t k; @@ -778,31 +675,18 @@ inline bool dispatch_gather3_(kop op, void* buf0, int64_t off0, void* buf1, } } // namespace detail_ -// Row gather along axis 0: out[i] = a[indices[row(i)]] (a, indices -// contiguous). One thread per output element, no write conflicts. -inline bool index_select(void* a_native, int64_t ao, void* idx_native, - int64_t idxo, void* out_native, int64_t oo, - int64_t row_size, int64_t k) { - detail_::gather_params p{static_cast(row_size), - static_cast(k * row_size)}; - return detail_::dispatch_gather3_(kop::index_select, a_native, ao, - idx_native, idxo, out_native, oo, &p, - sizeof(p), p.n); -} - // index_select's dual as a gather (float atomics would need MSL 3): each // output row sums the source rows whose index matches it, in source order -- // no zeroing needed. -inline bool index_add(void* idx_native, int64_t idxo, void* values_native, - int64_t vo, void* out_native, int64_t oo, - int64_t row_size, int64_t k, int64_t out_n) { +inline bool own::index_add(gpu::span idx, gpu::span values, gpu::span out, + int64_t row_size, int64_t k, int64_t out_n) { auto& c = context::get(); if (!c.device || row_size <= 0) return false; // A threadgroup per output row (see index_add_ in the MSL). c.bind_(kop::index_add); - detail_::set_buf_(c.enc, idx_native, idxo, 0ul); - detail_::set_buf_(c.enc, values_native, vo, 1ul); - detail_::set_buf_(c.enc, out_native, oo, 2ul); + detail_::set_buf_(c.enc, idx, 0ul); + detail_::set_buf_(c.enc, values, 1ul); + detail_::set_buf_(c.enc, out, 2ul); detail_::set_bytes_(c.enc, detail_::index_add_params{static_cast(row_size), static_cast(k)}, @@ -815,60 +699,18 @@ inline bool index_add(void* idx_native, int64_t idxo, void* values_native, // One-hot scatter into a new trailing axis, as a gather: out[pos,k] = // values[pos] where indices[pos] == k, else 0. Every output element reads, // never writes twice, so — like index_select above — no zeroing needed. -inline bool scatter_to_axis(void* idx_native, int64_t idxo, - void* values_native, int64_t vo, void* out_native, - int64_t oo, int64_t n, int64_t size) { +inline bool own::scatter_to_axis(gpu::span idx, gpu::span values, gpu::span out, + int64_t n, int64_t size) { detail_::scatter_axis_params p{static_cast(size), static_cast(n * size)}; - return detail_::dispatch_gather3_(kop::scatter_axis, idx_native, idxo, - values_native, vo, out_native, oo, &p, + return detail_::dispatch_gather3_(kop::scatter_axis, idx.buf, idx.off, + values.buf, values.off, out.buf, out.off, &p, sizeof(p), p.n); } // Cross-entropy's three: the trailing-axis gather, the one-pass row logsumexp // and the pullback that reads it. -// out[i] = src[i*size + idx[i]]: the element each position labels along a -// trailing axis of `size`. One thread per output. -inline bool gather_from_axis(void* src_native, int64_t so, void* idx_native, - int64_t idxo, void* out_native, int64_t oo, - int64_t n, int64_t size) { - if (n <= 0) return false; - const uint32_t p[2] = {static_cast(size), static_cast(n)}; - return detail_::dispatch_gather3_(kop::gather_axis_, src_native, so, - idx_native, idxo, out_native, oo, p, - sizeof(p), p[1]); -} - -// Row logsumexp over the last axis, affine epilogue: row_op's shape, one pass -// over the row (a running max and sum per thread). -inline bool row_logsumexp(void* in, int64_t io, void* out, int64_t oo, - int64_t rows, int64_t cols, float scale, - float offset) { - if (rows <= 0) return false; - return row_op(kop::row_logsumexp_, in, io, out, oo, rows, cols, scale, - offset); -} - -// Softmax cross-entropy's pullback from the forward's row logsumexp: x and -// out [rows, cols]; lse, targets and g one value per row. -inline bool xent_bwd(void* x, int64_t xo, void* lse, int64_t lo, void* tgt, - int64_t to, void* g, int64_t go, void* out, int64_t oo, - int64_t rows, int64_t cols) { - auto& c = context::get(); - if (!c.device || rows <= 0 || cols <= 0) return false; - c.bind_(kop::xent_bwd_); - detail_::set_buf_(c.enc, x, xo, 0ul); - detail_::set_buf_(c.enc, lse, lo, 1ul); - detail_::set_buf_(c.enc, tgt, to, 2ul); - detail_::set_buf_(c.enc, g, go, 3ul); - detail_::set_buf_(c.enc, out, oo, 4ul); - const uint32_t n = static_cast(rows * cols); - const uint32_t p[2] = {static_cast(cols), n}; - detail_::set_bytes_(c.enc, p, 5ul); - detail_::dispatch_grid_(c.enc, {(n + 255ul) / 256ul, 1, 1}, {256, 1, 1}); - return true; -} namespace detail_ { struct layer_norm_bwd_params { uint32_t rows, cols, rows_per_chunk, chunks; @@ -878,11 +720,11 @@ struct layer_norm_bwd_params { // Layer norm's pullback: cuda.h's contract and its three launches -- a row // kernel (dx and the row stats), a column-strip kernel, a fold. -inline bool layer_norm_bwd(void* x, int64_t xo, void* g, int64_t go, void* dy, - int64_t dyo, void* dx, void* dg, void* db, - void* stats, void* partials, int64_t rows, - int64_t cols, int64_t per_chunk, int64_t chunks, - float eps) { +inline bool own::layer_norm_bwd(gpu::span x, gpu::span g, gpu::span dy, + gpu::span dx, gpu::span dg, gpu::span db, + gpu::span stats, gpu::span partials, + int64_t rows, int64_t cols, int64_t per_chunk, + int64_t chunks, float eps) { auto& c = context::get(); if (!c.device || rows <= 0 || cols <= 0 || chunks <= 0) return false; detail_::layer_norm_bwd_params p{ @@ -893,60 +735,34 @@ inline bool layer_norm_bwd(void* x, int64_t xo, void* g, int64_t go, void* dy, const auto uk = static_cast(chunks); c.bind_(kop::layer_norm_bwd_dx_); - detail_::set_buf_(c.enc, x, xo, 0ul); - detail_::set_buf_(c.enc, g, go, 1ul); - detail_::set_buf_(c.enc, dy, dyo, 2ul); - detail_::set_buf_(c.enc, dx, 0, 3ul); - detail_::set_buf_(c.enc, stats, 0, 4ul); + detail_::set_buf_(c.enc, x, 0ul); + detail_::set_buf_(c.enc, g, 1ul); + detail_::set_buf_(c.enc, dy, 2ul); + detail_::set_buf_(c.enc, dx, 3ul); + detail_::set_buf_(c.enc, stats, 4ul); detail_::set_bytes_(c.enc, p, 5ul); detail_::dispatch_grid_(c.enc, {ur, 1, 1}, {256, 1, 1}); c.bind_(kop::layer_norm_bwd_gb_); - detail_::set_buf_(c.enc, x, xo, 0ul); - detail_::set_buf_(c.enc, dy, dyo, 1ul); - detail_::set_buf_(c.enc, stats, 0, 2ul); - detail_::set_buf_(c.enc, partials, 0, 3ul); + detail_::set_buf_(c.enc, x, 0ul); + detail_::set_buf_(c.enc, dy, 1ul); + detail_::set_buf_(c.enc, stats, 2ul); + detail_::set_buf_(c.enc, partials, 3ul); detail_::set_bytes_(c.enc, p, 4ul); detail_::dispatch_grid_(c.enc, {(uc + 31) / 32, uk, 1}, {32, 8, 1}); c.bind_(kop::layer_norm_bwd_gb_fold_); - detail_::set_buf_(c.enc, partials, 0, 0ul); - detail_::set_buf_(c.enc, dg, 0, 1ul); - detail_::set_buf_(c.enc, db, 0, 2ul); + detail_::set_buf_(c.enc, partials, 0ul); + detail_::set_buf_(c.enc, dg, 1ul); + detail_::set_buf_(c.enc, db, 2ul); detail_::set_bytes_(c.enc, p, 3ul); detail_::dispatch_grid_(c.enc, {(uc + 255) / 256, 1, 1}, {256, 1, 1}); return true; } namespace detail_ { -struct adam_params { - float b1, b2, eps, lr_over_bc1, inv_bc2; - uint32_t n; -}; } // namespace detail_ -// Adam's update in place: p, m and v read and written, g read, all one -// contiguous shape; the bias correction folded into lr_over_bc1 and inv_bc2. -inline bool adam_step(void* p, int64_t po, void* m, int64_t mo, void* v, - int64_t vo, void* g, int64_t go, int64_t n, float beta1, - float beta2, float eps, float lr_over_bc1, - float inv_bc2) { - auto& c = context::get(); - if (!c.device || n <= 0) return false; - c.bind_(kop::adam_step_); - detail_::set_buf_(c.enc, p, po, 0ul); - detail_::set_buf_(c.enc, m, mo, 1ul); - detail_::set_buf_(c.enc, v, vo, 2ul); - detail_::set_buf_(c.enc, g, go, 3ul); - detail_::adam_params ap{beta1, beta2, eps, lr_over_bc1, inv_bc2, - static_cast(n)}; - detail_::set_bytes_(c.enc, ap, 4ul); - detail_::dispatch_grid_( - c.enc, {(static_cast(n) + 255ul) / 256ul, 1, 1}, - {256, 1, 1}); - return true; -} - namespace detail_ { struct bcast_nd_params { uint32_t out_shape[kPadFoldMaxRank]; @@ -994,18 +810,17 @@ inline kop to_nd_(kop op) { // a_strides/b_strides are the broadcast strides (0 on a broadcast axis) // array.h computes host-side via the same broadcast_strides() the CPU // oracle uses -- mirrors cuda.h's own binary_bcast_nd exactly. -inline bool binary_bcast_nd(kop op, void* a_native, int64_t ao, - const int64_t* a_strides, void* b_native, - int64_t bo, const int64_t* b_strides, - void* out_native, int64_t oo, - const int64_t* out_shape, int rank, int64_t n, - float scale, float offset) { +inline bool own::binary_bcast_nd(kop op, gpu::span a, const int64_t* a_strides, + gpu::span b, const int64_t* b_strides, + gpu::span out, const int64_t* out_shape, + int rank, int64_t n, float scale, + float offset) { auto& c = context::get(); if (!c.device || rank <= 0 || rank > kPadFoldMaxRank) return false; c.bind_(detail_::to_nd_(op)); - detail_::set_buf_(c.enc, a_native, ao, 0ul); - detail_::set_buf_(c.enc, b_native, bo, 1ul); - detail_::set_buf_(c.enc, out_native, oo, 2ul); + detail_::set_buf_(c.enc, a, 0ul); + detail_::set_buf_(c.enc, b, 1ul); + detail_::set_buf_(c.enc, out, 2ul); detail_::bcast_nd_params p{}; for (int d = 0; d < rank; d++) { p.out_shape[d] = static_cast(out_shape[d]); @@ -1026,18 +841,17 @@ inline bool binary_bcast_nd(kop op, void* a_native, int64_t ao, // N-D broadcast ternary select: Tensor.where's GPU dispatch. Same flat-index // decode as binary_bcast_nd above, one more operand -- mirrors cuda.h's own // where_nd exactly. -inline bool where_nd(void* cond_native, int64_t co, const int64_t* c_strides, - void* a_native, int64_t ao, const int64_t* a_strides, - void* b_native, int64_t bo, const int64_t* b_strides, - void* out_native, int64_t oo, const int64_t* out_shape, - int rank, int64_t n) { +inline bool own::where_nd(gpu::span cond, const int64_t* c_strides, gpu::span a, + const int64_t* a_strides, gpu::span b, + const int64_t* b_strides, gpu::span out, + const int64_t* out_shape, int rank, int64_t n) { auto& c = context::get(); if (!c.device || rank <= 0 || rank > kPadFoldMaxRank) return false; c.bind_(kop::where_nd); - detail_::set_buf_(c.enc, cond_native, co, 0ul); - detail_::set_buf_(c.enc, a_native, ao, 1ul); - detail_::set_buf_(c.enc, b_native, bo, 2ul); - detail_::set_buf_(c.enc, out_native, oo, 3ul); + detail_::set_buf_(c.enc, cond, 0ul); + detail_::set_buf_(c.enc, a, 1ul); + detail_::set_buf_(c.enc, b, 2ul); + detail_::set_buf_(c.enc, out, 3ul); detail_::where_nd_params p{}; for (int d = 0; d < rank; d++) { p.out_shape[d] = static_cast(out_shape[d]); @@ -1057,14 +871,13 @@ inline bool where_nd(void* cond_native, int64_t co, const int64_t* c_strides, // clone()'s device arm for a strided view: a gather into a contiguous output, // where_nd's decode with one operand -- mirrors cuda.h's own copy_nd. The // .metal banner says why a clone must not go through the host here. -inline bool copy_nd(void* a_native, int64_t ao, const int64_t* a_strides, - void* out_native, int64_t oo, const int64_t* out_shape, - int rank, int64_t n) { +inline bool own::copy_nd(gpu::span a, const int64_t* a_strides, gpu::span out, + const int64_t* out_shape, int rank, int64_t n) { auto& c = context::get(); if (!c.device || rank <= 0 || rank > kPadFoldMaxRank) return false; c.bind_(kop::copy_nd); - detail_::set_buf_(c.enc, a_native, ao, 0ul); - detail_::set_buf_(c.enc, out_native, oo, 1ul); + detail_::set_buf_(c.enc, a, 0ul); + detail_::set_buf_(c.enc, out, 1ul); detail_::copy_nd_params p{}; for (int d = 0; d < rank; d++) { p.out_shape[d] = static_cast(out_shape[d]); @@ -1079,49 +892,6 @@ inline bool copy_nd(void* a_native, int64_t ao, const int64_t* a_strides, return true; } namespace detail_ { -inline kop to_cmp_(cmp_op op) { - switch (op) { - case cmp_op::gt: return kop::gt_; - case cmp_op::lt: return kop::lt_; - case cmp_op::ge: return kop::ge_; - case cmp_op::le: return kop::le_; - case cmp_op::eq: return kop::eq_; - case cmp_op::ne: return kop::ne_; - } - return kop::gt_; -} -inline kop to_unary_ext_(unary_ext_op op) { - switch (op) { - case unary_ext_op::tanh_: return kop::tanh_; - case unary_ext_op::sin_: return kop::sin_; - case unary_ext_op::cos_: return kop::cos_; - } - return kop::tanh_; -} -struct cmp_params { - uint32_t n; - uint32_t bstride; -}; -struct clamp_params { - float lo, hi; - uint32_t n; -}; -struct scalar_params { - float s, scale, offset; - uint32_t n; -}; -inline kop to_scalar_(scalar_op op) { - switch (op) { - case scalar_op::pow: return kop::pow_s_; - case scalar_op::gt: return kop::gt_s_; - case scalar_op::lt: return kop::lt_s_; - case scalar_op::ge: return kop::ge_s_; - case scalar_op::le: return kop::le_s_; - case scalar_op::eq: return kop::eq_s_; - case scalar_op::ne: return kop::ne_s_; - } - return kop::pow_s_; -} struct sum_to_params { uint32_t a_shape[kPadFoldMaxRank]; uint32_t a_strides[kPadFoldMaxRank]; @@ -1132,81 +902,20 @@ struct sum_to_params { }; } // namespace detail_ -// gt/lt/ge/le/eq/ne (array.h's comparison ops, ReLU/LeakyReLU/Clip's -// backward gate): same-shape only (bstride=1) or a scalar b (bstride=0) -- -// the two shapes array.h's gpu_compare_ ever dispatches. cmp_op stays its -// own vocabulary at the array.h boundary (see this file's cmp_op comment); -// mapped onto a dedicated kop slot here so it shares pso_()'s caching, same -// idea as binary_bcast_nd's to_nd_ above. Mirrors cuda.h's own compare(). -inline bool compare(cmp_op op, void* a, int64_t ao, void* b, int64_t bo, - void* out, int64_t oo, int64_t n, int64_t bstride) { - auto& c = context::get(); - if (!c.device) return false; - c.bind_(detail_::to_cmp_(op)); - detail_::set_buf_(c.enc, a, ao, 0ul); - detail_::set_buf_(c.enc, b, bo, 1ul); - detail_::set_buf_(c.enc, out, oo, 2ul); - detail_::cmp_params p{static_cast(n), - static_cast(bstride)}; - objc::send(c.enc, "setBytes:length:atIndex:", static_cast(&p), - static_cast(sizeof(p)), 3ul); - unsigned long groups = (static_cast(n) + 255ul) / 256ul; - detail_::dispatch_grid_(c.enc, {groups, 1, 1}, {256, 1, 1}); - return true; -} -// tanh_/sin_/cos_ (RoPE's trig, RNN/LSTM's tanh): plain elementwise, same -// shape as exp_/sqrt_ above -- unary() already does exactly this dispatch, -// just keyed by a kop array.h doesn't see directly. -inline bool unary_ext(unary_ext_op op, void* a, int64_t ao, void* out, - int64_t oo, int64_t n, float scale, float offset) { - return unary(detail_::to_unary_ext_(op), a, ao, out, oo, n, scale, offset); -} -// clamp(x, lo, hi): Clip's forward. No epilogue -- lo/hi occupy the role -// scale/offset play elsewhere (mirrors cuda.h's own clamp). -inline bool clamp(void* a, int64_t ao, void* out, int64_t oo, int64_t n, - float lo, float hi) { - auto& c = context::get(); - if (!c.device) return false; - c.bind_(kop::clamp_); - detail_::set_buf_(c.enc, a, ao, 0ul); - detail_::set_buf_(c.enc, out, oo, 1ul); - detail_::clamp_params p{lo, hi, static_cast(n)}; - objc::send(c.enc, "setBytes:length:atIndex:", static_cast(&p), - static_cast(sizeof(p)), 2ul); - unsigned long groups = (static_cast(n) + 255ul) / 256ul; - detail_::dispatch_grid_(c.enc, {groups, 1, 1}, {256, 1, 1}); - return true; -} -// Tensor-scalar ops (scalar_op): mirrors cuda.h's own scalar_binary. -inline bool scalar_binary(scalar_op op, void* a, int64_t ao, void* out, - int64_t oo, int64_t n, float s, float scale, - float offset) { - auto& c = context::get(); - if (!c.device) return false; - c.bind_(detail_::to_scalar_(op)); - detail_::set_buf_(c.enc, a, ao, 0ul); - detail_::set_buf_(c.enc, out, oo, 1ul); - detail_::scalar_params p{s, scale, offset, static_cast(n)}; - objc::send(c.enc, "setBytes:length:atIndex:", static_cast(&p), - static_cast(sizeof(p)), 2ul); - unsigned long groups = (static_cast(n) + 255ul) / 256ul; - detail_::dispatch_grid_(c.enc, {groups, 1, 1}, {256, 1, 1}); - return true; -} // sum_to (un-broadcast a gradient): gather, mirrors cuda.h's tl_sum_to -- // one thread per OUTPUT element sums every `a` element that broadcasts // onto it, so no atomics (unlike index_add). -inline bool sum_to(void* a, int64_t ao, const int64_t* a_shape, - const int64_t* a_strides, const int64_t* acc, int rank, - int64_t out_n, int64_t reduced_n, void* out, int64_t oo) { +inline bool own::sum_to(gpu::span a, const int64_t* a_shape, + const int64_t* a_strides, const int64_t* acc, int rank, + int64_t out_n, int64_t reduced_n, gpu::span out) { auto& c = context::get(); if (!c.device || rank <= 0 || rank > kPadFoldMaxRank) return false; // A deep reduction (a bias gradient sums its column over every row) earns a // threadgroup per output, as on CUDA; a shallow one keeps a thread per output. const bool blocked = reduced_n >= 64; c.bind_(blocked ? kop::sum_to_blocked_ : kop::sum_to_); - detail_::set_buf_(c.enc, a, ao, 0ul); - detail_::set_buf_(c.enc, out, oo, 1ul); + detail_::set_buf_(c.enc, a, 0ul); + detail_::set_buf_(c.enc, out, 1ul); detail_::sum_to_params p{}; for (int d = 0; d < rank; d++) { p.a_shape[d] = static_cast(a_shape[d]); @@ -1243,9 +952,9 @@ struct concat_part_params { // over OUTPUT elements (needed for pad's zero border) rather than the much // smaller SOURCE (this part's own) element count concat wants, so this // gets its own small kernel instead of reusing pad_'s PSO. -inline bool concat_part(void* a, int64_t ao, void* out, int64_t oo, - const int64_t* a_shape, const int64_t* out_shape, - int rank, int axis, int64_t before, int64_t n) { +inline bool own::concat_part(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t before, int64_t n) { auto& c = context::get(); if (!c.device || rank <= 0 || rank > kPadFoldMaxRank) return false; int64_t out_strides[kPadFoldMaxRank]; @@ -1255,8 +964,8 @@ inline bool concat_part(void* a, int64_t ao, void* out, int64_t oo, acc *= out_shape[d]; } c.bind_(kop::concat_part_); - detail_::set_buf_(c.enc, a, ao, 0ul); - detail_::set_buf_(c.enc, out, oo, 1ul); + detail_::set_buf_(c.enc, a, 0ul); + detail_::set_buf_(c.enc, out, 1ul); detail_::concat_part_params p{}; for (int d = 0; d < rank; d++) { p.out_strides[d] = static_cast(out_strides[d]); @@ -1291,16 +1000,16 @@ struct rope_params { // in; array.h's gpu_rope_ never passes one). Dispatched flat over rows*(D/2) // (CUDA instead grids by row, blocks by D/2 -- this file's own kernels are // all flat-1D, same as sum_to/compare/unary_ext above). -inline bool rope(void* x, void* out, int64_t rows, int64_t T, int64_t D, - int64_t pos, float base, void* bias = nullptr) { +inline bool own::rope(gpu::span x, gpu::span out, int64_t rows, int64_t T, + int64_t D, int64_t pos, float base, gpu::span bias) { auto& c = context::get(); if (!c.device || D <= 0 || (D & 1)) return false; int64_t half = D / 2; int64_t n = rows * half; c.bind_(kop::rope_); - detail_::set_buf_(c.enc, x, 0, 0ul); - detail_::set_buf_(c.enc, out, 0, 1ul); - detail_::set_buf_(c.enc, bias ? bias : x, 0, 3ul); // a slot must be bound + detail_::set_buf_(c.enc, x, 0ul); + detail_::set_buf_(c.enc, out, 1ul); + detail_::set_buf_(c.enc, bias ? bias : x, 3ul); // a slot must be bound detail_::rope_params p{}; p.T = static_cast(T); p.D = static_cast(D); @@ -1308,7 +1017,7 @@ inline bool rope(void* x, void* out, int64_t rows, int64_t T, int64_t D, p.half_ = static_cast(half); p.base = base; p.n = static_cast(n); - p.has_bias = bias ? 1u : 0u; + p.has_bias = bias.buf ? 1u : 0u; objc::send(c.enc, "setBytes:length:atIndex:", static_cast(&p), static_cast(sizeof(p)), 2ul); unsigned long groups = (static_cast(n) + 255ul) / 256ul; @@ -1334,18 +1043,10 @@ struct gemv_combine_params { uint32_t n, parts; }; -struct gemv_row_params { - uint32_t n, k; -}; - struct gemm_nt_params { uint32_t M, N, K; }; -struct gemv_q4_params { - uint32_t n, k, group; -}; - struct attn_decode_params { uint32_t ctx, kv_stride, group, chunk; float scale; @@ -1355,27 +1056,6 @@ struct attn_combine_params { uint32_t splits; }; -struct kv_params { - uint32_t pos, kv_stride, T; -}; - -struct argmax_params { - uint32_t n; -}; - -struct rmsnorm_params { - uint32_t n, add; - float eps; -}; - -struct swiglu_params { - uint32_t ff; -}; - -struct heads_params { - uint32_t T, ld, off, has_bias; -}; - // Keys per split, or 0 for the single-pass kernel: one threadgroup a head // leaves most of a 16-core GPU idle when a model has few heads, so cut the // keys until there are enough threadgroups (cuda's attn_split_count). @@ -1403,20 +1083,21 @@ inline void attn_dispatch_(objc::id enc, const attn_params& p, // q,out [n_q_heads,T,D]; K/V a [n_kv_heads,kv_max,D] cache read over // [0,pos0+T), query p at absolute position pos0+p. GQA via the head ratio. // D∈{64,128}; the cache is f32, or bf16 with kv_bf16. -inline bool attn_prefill(void* q, void* K, void* V, void* out, - int64_t n_q_heads, int64_t n_kv_heads, int64_t T, - int64_t kv_max, int64_t D, float scale, - bool kv_bf16 = false, int64_t pos0 = 0) { +inline bool own::attn_prefill(gpu::span q, gpu::span K, gpu::span V, + gpu::span out, int64_t n_q_heads, + int64_t n_kv_heads, int64_t T, int64_t kv_max, + int64_t D, float scale, bool kv_bf16, + int64_t pos0) { 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_)); - detail_::set_buf_(c.enc, q, 0, 0ul); - detail_::set_buf_(c.enc, K, 0, 1ul); - detail_::set_buf_(c.enc, V, 0, 2ul); - detail_::set_buf_(c.enc, out, 0, 3ul); + detail_::set_buf_(c.enc, q, 0ul); + detail_::set_buf_(c.enc, K, 1ul); + detail_::set_buf_(c.enc, V, 2ul); + detail_::set_buf_(c.enc, out, 3ul); detail_::attn_params p{static_cast(T), static_cast(kv_max * D), static_cast(n_q_heads / n_kv_heads), @@ -1428,7 +1109,8 @@ inline bool attn_prefill(void* q, void* K, void* V, void* out, // y[1,N] = a[1,K] · B[K,N] with f32 or bf16 weights, all contiguous: the // decode projection. A narrow layer's N/256 threadgroups leave the GPU idle, // so K is split across the grid's y and a combine pass sums the slices. -inline bool gemv_(kop op, void* a, void* B, void* y, int64_t n, int64_t k) { +inline bool gemv_(kop op, gpu::span a, gpu::span B, gpu::span y, int64_t n, + int64_t k) { auto& c = context::get(); if (!c.device || n <= 0 || k <= 0) return false; constexpr unsigned long NT = 256, kGroupsWanted = 64; @@ -1437,23 +1119,23 @@ inline bool gemv_(kop op, void* a, void* B, void* y, int64_t n, int64_t k) { unsigned long chunk = (static_cast(k) + parts - 1) / parts; chunk = (chunk + NT - 1) / NT * NT; // whole `a` tiles parts = (static_cast(k) + chunk - 1) / chunk; - void* out = y; + gpu::span out = y; if (parts > 1) { - out = detail_::scratch_(static_cast(parts) * n * 4); + out = {detail_::scratch_(static_cast(parts) * n * 4), 0}; if (!out) parts = 1, chunk = static_cast(k), out = y; } c.bind_(op); - detail_::set_buf_(c.enc, a, 0, 0ul); - detail_::set_buf_(c.enc, B, 0, 1ul); - detail_::set_buf_(c.enc, out, 0, 2ul); + detail_::set_buf_(c.enc, a, 0ul); + detail_::set_buf_(c.enc, B, 1ul); + detail_::set_buf_(c.enc, out, 2ul); detail_::gemv_params p{static_cast(n), static_cast(k), static_cast(chunk)}; detail_::set_bytes_(c.enc, p, 3ul); detail_::dispatch_grid_(c.enc, {cols, parts, 1}, {NT, 1, 1}); if (parts == 1) return true; c.bind_(kop::gemv_combine_); - detail_::set_buf_(c.enc, out, 0, 0ul); - detail_::set_buf_(c.enc, y, 0, 1ul); + detail_::set_buf_(c.enc, out, 0ul); + detail_::set_buf_(c.enc, y, 1ul); detail_::gemv_combine_params cp{static_cast(n), static_cast(parts)}; detail_::set_bytes_(c.enc, cp, 2ul); @@ -1461,53 +1143,28 @@ inline bool gemv_(kop op, void* a, void* B, void* y, int64_t n, int64_t k) { return true; } -inline bool gemv_f32(void* a, void* B, void* y, int64_t n, int64_t k) { +inline bool own::gemv_f32(gpu::span a, gpu::span B, gpu::span y, int64_t n, + int64_t k) { return gemv_(kop::gemv_f32_, a, B, y, n, k); } -inline bool gemv_bf16(void* a, void* B, void* y, int64_t n, int64_t k) { +inline bool own::gemv_bf16(gpu::span a, gpu::span B, gpu::span y, int64_t n, + int64_t k) { return gemv_(kop::gemv_bf16_, a, B, y, n, k); } -// y[1,N] = a[1,K] · W[N,K]ᵀ with the weight row-major (GGML-native). One -// threadgroup per output row, so a narrow layer parallelizes over N without -// the [K,N] path's split-K; the threadgroup size follows cuda's own policy, -// the smallest that still gives every thread ~one 8-wide step. Requires -// k % 8 == 0 (host-gated; the caller keeps to the [K,N] GEMV otherwise). -inline bool gemv_bf16_row(void* a, void* B, void* y, int64_t n, int64_t k) { - auto& c = context::get(); - if (!c.device || n <= 0 || k <= 0 || (k % 8) != 0) return false; - unsigned long nt = 32; - for (unsigned long bs = 64; bs <= 256; bs += 32) { - if ((k + 8 * (int64_t)nt - 1) / (8 * (int64_t)nt) > - (k + 8 * (int64_t)bs - 1) / (8 * (int64_t)bs)) { - nt = bs; - } - } - c.bind_(kop::gemv_bf16_row_); - detail_::set_buf_(c.enc, a, 0, 0ul); - detail_::set_buf_(c.enc, B, 0, 1ul); - detail_::set_buf_(c.enc, y, 0, 2ul); - detail_::gemv_row_params p{static_cast(n), - static_cast(k)}; - detail_::set_bytes_(c.enc, p, 3ul); - detail_::dispatch_grid_(c.enc, {static_cast(n), 1, 1}, - {nt, 1, 1}); - return true; -} - // C[M,N] = A[M,K] · B[N,K]ᵀ, B bf16: the batched prefill's projection, where // one weight serves a whole chunk of prompt tokens. The 32-row tile for a // short chunk (more M blocks to fill the GPU), the 64-row one otherwise. -inline bool gemm_bf16_nt(void* A, void* B, void* C, int64_t M, int64_t N, - int64_t K) { +inline bool own::gemm_bf16_nt(gpu::span A, gpu::span B, gpu::span C, int64_t M, + int64_t N, int64_t K) { auto& c = context::get(); if (!c.device || M <= 0 || N <= 0 || K <= 0) return false; const unsigned long BM = M <= 64 ? 32 : 64; c.bind_(M <= 64 ? kop::gemm_bf16_nt32_ : kop::gemm_bf16_nt_); - detail_::set_buf_(c.enc, A, 0, 0ul); - detail_::set_buf_(c.enc, B, 0, 1ul); - detail_::set_buf_(c.enc, C, 0, 2ul); + detail_::set_buf_(c.enc, A, 0ul); + detail_::set_buf_(c.enc, B, 1ul); + detail_::set_buf_(c.enc, C, 2ul); detail_::gemm_nt_params p{static_cast(M), static_cast(N), static_cast(K)}; detail_::set_bytes_(c.enc, p, 3ul); @@ -1517,34 +1174,12 @@ inline bool gemm_bf16_nt(void* A, void* B, void* C, int64_t M, int64_t N, return true; } -// int4 weights: one threadgroup per output row. `scales` is the caller's -// pointer arithmetic on the q4 buffer (a device address on CUDA); here the two -// name one MTLBuffer, so the difference is the scales block's byte offset. -inline bool gemv_q4(void* a, void* qw, void* scales, void* y, int64_t N, - int64_t K, int64_t group) { - auto& c = context::get(); - if (!c.device || N <= 0 || K <= 0 || group <= 0) return false; - if (K % group != 0 || group % 8 != 0) return false; - const unsigned long soff = static_cast( - static_cast(scales) - static_cast(qw)); - c.bind_(kop::gemv_q4_); - detail_::set_buf_(c.enc, a, 0, 0ul); - detail_::set_buf_(c.enc, qw, 0, 1ul); - detail_::set_buf_(c.enc, qw, soff, 2ul); - detail_::set_buf_(c.enc, y, 0, 3ul); - detail_::gemv_q4_params p{static_cast(N), static_cast(K), - static_cast(group)}; - detail_::set_bytes_(c.enc, p, 4ul); - detail_::dispatch_grid_(c.enc, {static_cast(N), 1, 1}, - {256, 1, 1}); - return true; -} - // One decode step: q [n_q_heads,D] against a [n_kv_heads,kv_max,D] cache read // over [0,ctx). GQA via the head ratio; the cache is f32, or bf16 with kv_bf16. -inline bool attn_decode(void* q, void* K, void* V, void* out, int64_t n_q_heads, - int64_t n_kv_heads, int64_t ctx, int64_t kv_max, - int64_t D, float scale, bool kv_bf16 = false) { +inline bool own::attn_decode(gpu::span q, gpu::span K, gpu::span V, + gpu::span out, int64_t n_q_heads, + int64_t n_kv_heads, int64_t ctx, int64_t kv_max, + int64_t D, float scale, bool kv_bf16) { 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 || ctx <= 0) return false; @@ -1552,11 +1187,11 @@ inline bool attn_decode(void* q, void* K, void* V, void* out, int64_t n_q_heads, unsigned long chunk = detail_::attn_split_chunk_(n_q_heads, ctx); const unsigned long splits = chunk ? (static_cast(ctx) + chunk - 1) / chunk : 1; - void* dst = out; + gpu::span dst = out; if (splits > 1) { // pm[H*S] | pl[H*S] | pacc[H*S*D], the layout attn_combine_ reads. const int64_t hs = n_q_heads * static_cast(splits); - dst = detail_::scratch_((hs * 2 + hs * D) * 4); + dst = {detail_::scratch_((hs * 2 + hs * D) * 4), 0}; if (!dst) chunk = 0, dst = out; } const bool split = chunk != 0; @@ -1570,10 +1205,10 @@ inline bool attn_decode(void* q, void* K, void* V, void* out, int64_t n_q_heads, : kop::attn_decode_bf16_128_) : (d64 ? kop::attn_decode_64_ : kop::attn_decode_128_)); } - detail_::set_buf_(c.enc, q, 0, 0ul); - detail_::set_buf_(c.enc, K, 0, 1ul); - detail_::set_buf_(c.enc, V, 0, 2ul); - detail_::set_buf_(c.enc, dst, 0, 3ul); + detail_::set_buf_(c.enc, q, 0ul); + detail_::set_buf_(c.enc, K, 1ul); + detail_::set_buf_(c.enc, V, 2ul); + detail_::set_buf_(c.enc, dst, 3ul); detail_::attn_decode_params p{static_cast(ctx), static_cast(kv_max * D), static_cast(n_q_heads / n_kv_heads), @@ -1585,8 +1220,8 @@ inline bool attn_decode(void* q, void* K, void* V, void* out, int64_t n_q_heads, {static_cast(D), 1, 1}); if (!split) return true; c.bind_(d64 ? kop::attn_combine_64_ : kop::attn_combine_128_); - detail_::set_buf_(c.enc, dst, 0, 0ul); - detail_::set_buf_(c.enc, out, 0, 1ul); + detail_::set_buf_(c.enc, dst, 0ul); + detail_::set_buf_(c.enc, out, 1ul); detail_::attn_combine_params cp{static_cast(splits)}; detail_::set_bytes_(c.enc, cp, 2ul); detail_::dispatch_grid_(c.enc, {static_cast(n_q_heads), 1, 1}, @@ -1597,349 +1232,86 @@ inline bool attn_decode(void* q, void* K, void* V, void* out, int64_t n_q_heads, // ---- the KV cache's writes and the decode step's remaining kernels -------- // cuda.h's contracts: raw device buffers, so a model's step builds no graph. -// One step's k/v (each [n_kv_heads, D]) into row `pos` of the [n_kv_heads, -// kv_max, D] cache; narrowed to bf16 when the cache is. -inline bool kv_append(void* Kc, void* Vc, void* k_new, void* v_new, int64_t pos, - int64_t kv_max, int64_t n_kv_heads, int64_t D, - bool kv_bf16 = false) { - auto& c = context::get(); - if (!c.device || (D != 64 && D != 128) || n_kv_heads <= 0) return false; - c.bind_(kv_bf16 ? kop::kv_append_bf16_ : kop::kv_append_); - void* bufs[] = {Kc, Vc, k_new, v_new}; - for (unsigned long i = 0; i < 4; i++) detail_::set_buf_(c.enc, bufs[i], 0, i); - detail_::kv_params p{static_cast(pos), - static_cast(kv_max * D), 1}; - detail_::set_bytes_(c.enc, p, 4ul); - detail_::dispatch_grid_(c.enc, {static_cast(n_kv_heads), 1, 1}, - {static_cast(D), 1, 1}); - return true; -} - -// A prefill's k/v (each [n_kv_heads, T, D]) into cache rows [pos0, pos0 + T). -inline bool kv_fill(void* Kc, void* Vc, void* K, void* V, int64_t T, - int64_t kv_max, int64_t n_kv_heads, int64_t D, - bool kv_bf16 = false, int64_t pos0 = 0) { - auto& c = context::get(); - if (!c.device || (D != 64 && D != 128) || n_kv_heads <= 0 || T <= 0) { - return false; - } - c.bind_(kv_bf16 ? kop::kv_fill_bf16_ : kop::kv_fill_); - void* bufs[] = {Kc, Vc, K, V}; - for (unsigned long i = 0; i < 4; i++) detail_::set_buf_(c.enc, bufs[i], 0, i); - detail_::kv_params p{static_cast(pos0), - static_cast(kv_max * D), - static_cast(T)}; - detail_::set_bytes_(c.enc, p, 4ul); - detail_::dispatch_grid_(c.enc, - {static_cast(n_kv_heads), - static_cast(T), 1}, - {static_cast(D), 1, 1}); - return true; -} - -// The argmax of a length-n vector, the smallest index on ties: greedy -// decoding reads one int back rather than the logits. Drains the queue. -inline bool argmax(void* in, int64_t n, int64_t* out_idx) { - auto& c = context::get(); - if (!c.device || !in || n <= 0 || !out_idx) return false; - if (!c.argmax_res) { - c.argmax_res = alloc(4, &c.argmax_res_contents); - if (!c.argmax_res) return false; - } - c.bind_(kop::argmax_); - detail_::set_buf_(c.enc, in, 0, 0ul); - detail_::set_buf_(c.enc, c.argmax_res, 0, 1ul); - detail_::argmax_params p{static_cast(n)}; - detail_::set_bytes_(c.enc, p, 2ul); - detail_::dispatch_grid_(c.enc, {1, 1, 1}, {256, 1, 1}); - flush(); - *out_idx = *reinterpret_cast(c.argmax_res_contents); - return true; -} - -// hout = x · rsqrt(mean(x²) + eps) · w, per row of the [rows, n] buffer. -inline bool rmsnorm(void* x, void* w, void* out, int64_t n, float eps, - int64_t rows = 1) { - auto& c = context::get(); - if (!c.device || n <= 0 || rows <= 0) return false; - c.bind_(kop::rmsnorm_); - void* bufs[] = {x, x, w, out, out}; // no delta: slots 1 and 3 idle - for (unsigned long i = 0; i < 5; i++) detail_::set_buf_(c.enc, bufs[i], 0, i); - detail_::rmsnorm_params p{static_cast(n), 0, eps}; - detail_::set_bytes_(c.enc, p, 5ul); - detail_::dispatch_grid_(c.enc, {static_cast(rows), 1, 1}, - {256, 1, 1}); - return true; -} - -// xout = x + delta and hout = rmsnorm(xout) · w: a residual add folded into -// the norm that follows it. xout may alias x. -inline bool rmsnorm_res(void* x, void* delta, void* w, void* xout, void* hout, - int64_t n, float eps, int64_t rows = 1) { - auto& c = context::get(); - if (!c.device || n <= 0 || rows <= 0) return false; - c.bind_(kop::rmsnorm_); - void* bufs[] = {x, delta, w, xout, hout}; - for (unsigned long i = 0; i < 5; i++) detail_::set_buf_(c.enc, bufs[i], 0, i); - detail_::rmsnorm_params p{static_cast(n), 1, eps}; - detail_::set_bytes_(c.enc, p, 5ul); - detail_::dispatch_grid_(c.enc, {static_cast(rows), 1, 1}, - {256, 1, 1}); - return true; -} - -// out[rows, ff] = silu(gate) · up out of the fused gate|up buffer [rows, 2ff]. -inline bool swiglu(void* gu, void* out, int64_t ff, int64_t rows = 1) { - auto& c = context::get(); - if (!c.device || ff <= 0 || rows <= 0) return false; - c.bind_(kop::swiglu_); - detail_::set_buf_(c.enc, gu, 0, 0ul); - detail_::set_buf_(c.enc, out, 0, 1ul); - detail_::swiglu_params p{static_cast(ff)}; - detail_::set_bytes_(c.enc, p, 2ul); - detail_::dispatch_grid_(c.enc, - {(static_cast(ff) + 255) / 256, - static_cast(rows), 1}, - {256, 1, 1}); - return true; -} - -// Token-major [T, ld] -> head-major [H, T, D] from column block `off`, adding -// the optional per-head bias [H, D]. -inline bool split_heads(void* src, void* bias, void* dst, int64_t T, int64_t ld, - int64_t off, int64_t H, int64_t D) { - auto& c = context::get(); - if (!c.device || T <= 0 || H <= 0 || D <= 0) return false; - c.bind_(kop::split_heads_); - detail_::set_buf_(c.enc, src, 0, 0ul); - detail_::set_buf_(c.enc, bias ? bias : src, 0, 1ul); - detail_::set_buf_(c.enc, dst, 0, 2ul); - detail_::heads_params p{static_cast(T), static_cast(ld), - static_cast(off), bias ? 1u : 0u}; - detail_::set_bytes_(c.enc, p, 3ul); - detail_::dispatch_grid_( - c.enc, {static_cast(H), static_cast(T), 1}, - {static_cast(D), 1, 1}); - return true; -} - -// Head-major [H, T, D] -> token-major [T, H·D], split_heads' inverse. -inline bool merge_heads(void* src, void* dst, int64_t T, int64_t H, int64_t D) { - auto& c = context::get(); - if (!c.device || T <= 0 || H <= 0 || D <= 0) return false; - c.bind_(kop::merge_heads_); - detail_::set_buf_(c.enc, src, 0, 0ul); - detail_::set_buf_(c.enc, dst, 0, 1ul); - detail_::heads_params p{static_cast(T), 0, 0, 0}; - detail_::set_bytes_(c.enc, p, 2ul); - detail_::dispatch_grid_( - c.enc, {static_cast(H), static_cast(T), 1}, - {static_cast(D), 1, 1}); - return true; -} - // The query half of the pullback: q, K, V, dO, O and dq all [H,T,D] // contiguous, `stats` [2,H,T] the row logsumexp and dO·O. -inline bool attn_prefill_dq(void* q, void* K, void* V, void* dO, void* O, - void* dq, void* stats, int64_t H, int64_t T, - int64_t D, float scale) { +inline bool own::attn_prefill_dq(gpu::span q, gpu::span K, gpu::span V, + gpu::span dO, gpu::span O, gpu::span dq, + gpu::span stats, int64_t H, int64_t T, + 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_); - void* bufs[] = {q, K, V, dO, O, dq, stats}; - for (unsigned long i = 0; i < 7; i++) detail_::set_buf_(c.enc, bufs[i], 0, i); + 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}; detail_::attn_dispatch_(c.enc, p, 7ul, H, T); return true; } // The key/value half, reading the stats the call above wrote. -inline bool attn_prefill_dkv(void* q, void* K, void* V, void* dO, void* stats, - void* dK, void* dV, int64_t H, int64_t T, - int64_t D, float scale) { +inline bool own::attn_prefill_dkv(gpu::span q, gpu::span K, gpu::span V, + gpu::span dO, gpu::span stats, gpu::span dK, + gpu::span dV, int64_t H, int64_t T, 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_dkv_64_ : kop::attn_bwd_dkv_128_); - void* bufs[] = {q, K, V, dO, stats, dK, dV}; - for (unsigned long i = 0; i < 7; i++) detail_::set_buf_(c.enc, bufs[i], 0, i); + 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}; detail_::attn_dispatch_(c.enc, p, 7ul, H, T); return true; } -#else // !__APPLE__ — stubs so callers carry no platform conditionals - -inline bool available() { return false; } -inline bool pending() { return false; } -inline void flush() {} -inline void* alloc(int64_t, float**, bool = false) { return nullptr; } -inline void release(void*, int64_t, float*) {} -inline void upload(void*, const float*, int64_t) {} -inline bool binary(kop, void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, float, float) { - return false; -} -inline bool binary_bcast(kop, void*, int64_t, int64_t, int64_t, void*, int64_t, - int64_t, int64_t, void*, int64_t, int64_t, int64_t, - float, float) { - return false; -} -inline bool unary(kop, void*, int64_t, void*, int64_t, int64_t, float, float) { - return false; -} -inline bool gemm(void*, int64_t, int64_t, bool, void*, int64_t, int64_t, bool, - void*, int64_t, int64_t, int64_t, int64_t, float, float) { - return false; -} -inline bool gemm_batched(void*, int64_t, int64_t, bool, int64_t, void*, - int64_t, int64_t, bool, int64_t, void*, int64_t, - int64_t, int64_t, int64_t, int64_t, float, float, - void* = nullptr, int64_t = 0) { - return false; -} -inline bool row_op(kop, void*, int64_t, void*, int64_t, int64_t, int64_t, - float, float) { - return false; -} -inline bool layer_norm(void*, int64_t, void*, int64_t, void*, int64_t, void*, - int64_t, int64_t, int64_t, float, float, float) { - return false; -} -inline bool pad(void*, int64_t, void*, int64_t, const int64_t*, - const int64_t*, int, int, int64_t, int64_t, int64_t) { - return false; -} -inline bool fold(void*, int64_t, void*, int64_t, const int64_t*, - const int64_t*, int, int, int64_t, int64_t, int64_t) { - return false; -} -inline bool index_select(void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool index_add(void*, int64_t, void*, int64_t, void*, int64_t, int64_t, - int64_t, int64_t) { - return false; -} -inline bool scatter_to_axis(void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool gather_from_axis(void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool row_logsumexp(void*, int64_t, void*, int64_t, int64_t, int64_t, - float, float) { - return false; -} -inline bool xent_bwd(void*, int64_t, void*, int64_t, void*, int64_t, void*, - int64_t, void*, int64_t, int64_t, int64_t) { - return false; -} -inline bool layer_norm_bwd(void*, int64_t, void*, int64_t, void*, int64_t, - void*, void*, void*, void*, void*, int64_t, int64_t, - int64_t, int64_t, float) { - return false; -} -inline bool attn_prefill(void*, void*, void*, void*, int64_t, int64_t, int64_t, - int64_t, int64_t, float, bool = false, int64_t = 0) { - return false; -} -inline bool attn_prefill_dq(void*, void*, void*, void*, void*, void*, void*, - int64_t, int64_t, int64_t, float) { - return false; -} -inline bool attn_prefill_dkv(void*, void*, void*, void*, void*, void*, void*, - int64_t, int64_t, int64_t, float) { - return false; -} -inline bool adam_step(void*, int64_t, void*, int64_t, void*, int64_t, void*, - int64_t, int64_t, float, float, float, float, float) { - return false; -} -inline bool binary_bcast_nd(kop, void*, int64_t, const int64_t*, void*, - int64_t, const int64_t*, void*, int64_t, - const int64_t*, int, int64_t, float, float) { - return false; -} -inline bool where_nd(void*, int64_t, const int64_t*, void*, int64_t, - const int64_t*, void*, int64_t, const int64_t*, void*, - int64_t, const int64_t*, int, int64_t) { - return false; -} -inline bool copy_nd(void*, int64_t, const int64_t*, void*, int64_t, - const int64_t*, int, int64_t) { - return false; -} -inline bool sum_to(void*, int64_t, const int64_t*, const int64_t*, - const int64_t*, int, int64_t, int64_t, void*, int64_t) { - return false; -} -inline bool compare(cmp_op, void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool unary_ext(unary_ext_op, void*, int64_t, void*, int64_t, int64_t, - float, float) { - return false; -} -inline bool clamp(void*, int64_t, void*, int64_t, int64_t, float, float) { - return false; -} -inline bool scalar_binary(scalar_op, void*, int64_t, void*, int64_t, int64_t, - float, float, float) { - return false; -} -inline bool concat_part(void*, int64_t, void*, int64_t, const int64_t*, - const int64_t*, int, int, int64_t, int64_t) { - return false; -} -inline bool rope(void*, void*, int64_t, int64_t, int64_t, int64_t, float, - void* = nullptr) { - return false; -} -inline bool kv_append(void*, void*, void*, void*, int64_t, int64_t, int64_t, - int64_t, bool = false) { - return false; -} -inline bool kv_fill(void*, void*, void*, void*, int64_t, int64_t, int64_t, - int64_t, bool = false, int64_t = 0) { - return false; -} -inline bool argmax(void*, int64_t, int64_t*) { return false; } -inline bool rmsnorm(void*, void*, void*, int64_t, float, int64_t = 1) { - return false; -} -inline bool rmsnorm_res(void*, void*, void*, void*, void*, int64_t, float, - int64_t = 1) { - return false; -} -inline bool swiglu(void*, void*, int64_t, int64_t = 1) { return false; } -inline bool split_heads(void*, void*, void*, int64_t, int64_t, int64_t, int64_t, - int64_t) { - return false; -} -inline bool merge_heads(void*, void*, int64_t, int64_t, int64_t) { return false; } -inline bool gemv_f32(void*, void*, void*, int64_t, int64_t) { return false; } -inline bool gemv_bf16(void*, void*, void*, int64_t, int64_t) { return false; } -inline bool gemv_bf16_row(void*, void*, void*, int64_t, int64_t) { return false; } -inline bool gemm_bf16_nt(void*, void*, void*, int64_t, int64_t, int64_t) { - return false; -} -inline bool gemv_q4(void*, void*, void*, void*, int64_t, int64_t, int64_t) { - return false; -} -inline bool attn_decode(void*, void*, void*, void*, int64_t, int64_t, int64_t, - int64_t, int64_t, float, bool = false) { - return false; -} - -#endif - // What a model may ask of this backend beyond the kernel contract (gpu.h // lists every backend's). No graph capture: a Metal command buffer is cheap to // encode, so a decode step re-encodes each token. + +// "No bias" is a flag here: a kernel cannot test a buffer for null, so an +// absent bias binds src in its place and is never read. +inline bool own::split_heads(gpu::span src, gpu::span bias, gpu::span dst, + int64_t T, int64_t ld, int64_t off, int64_t H, + int64_t D) { + struct { + uint32_t T, ld, off, has_bias; + } p{static_cast(T), static_cast(ld), + static_cast(off), bias ? 1u : 0u}; + const gpu::arg args[] = {gpu::in(src), gpu::in(bias ? bias : src), + gpu::out(dst)}; + return dispatch(kop::split_heads_, args, 3, &p, sizeof(p), + gpu::policy::per_head(H, T, D)); +} + +// One int comes back through a 4-byte buffer the context keeps. +inline bool own::argmax(gpu::span a, int64_t n, int64_t* out_idx) { + auto& c = context::get(); + if (!c.device) return false; + if (!c.argmax_res) { + c.argmax_res = alloc(4, &c.argmax_res_contents); + if (!c.argmax_res) return false; + } + const uint32_t p = static_cast(n); + const gpu::arg args[] = {gpu::in(a), gpu::out({c.argmax_res, 0})}; + if (!dispatch(kop::argmax_, args, 2, &p, sizeof(p), {1, 1, 1, 256, 1, 1, 0})) { + return false; + } + flush(); + *out_idx = *reinterpret_cast(c.argmax_res_contents); + return true; +} + +// What the shared launch policy (gpu_ops.h) may assume of this backend's +// kernels. +struct traits { + // A [rows, cols] elementwise kernel reads its cell from a 2-D thread + // position rather than a flat index. + static constexpr bool cells_2d = true; + // Launches are recorded under tl::profile by this backend itself, with + // (times_launches) a device time on each. + static constexpr bool profiles_launches = true; + static constexpr bool times_launches = true; +}; + struct caps { // Whether the model-path row is real here, or answers false: a decoder // runs on raw buffers only where it is true, and keeps to the array ops @@ -1948,11 +1320,6 @@ struct caps { static constexpr bool graph_capture = false; static constexpr bool row_gemv = true; // gemv_bf16_row: weights as [N,K] static constexpr bool bf16_gemm = true; // gemm_bf16_nt: a bf16-weight GEMM - // `native` is an MTLBuffer handle, not an address: arithmetic on it names - // nothing, so a model writes each piece to its own buffer rather than - // slicing one kernel's output. (The generic kernels still take byte - // offsets — it is only a *pointer* that cannot carry one.) - static constexpr bool flat_addressing = false; }; // The capture group, cuda.h's contracts: absent here, so each answers false @@ -1965,30 +1332,8 @@ inline graph_exec capture_end() { return nullptr; } inline bool graph_launch(graph_exec) { return false; } inline void graph_destroy(graph_exec) {} inline void upload_u32(void*, unsigned) {} -inline bool incr_u32(void*) { return false; } -inline bool rope_dpos(void*, void*, int64_t, int64_t, int64_t, void*, float, - void* = nullptr) { - return false; -} -inline bool kv_append_dpos(void*, void*, void*, void*, void*, int64_t, int64_t, - int64_t) { - return false; -} -inline bool attn_decode_dpos(void*, void*, void*, void*, int64_t, int64_t, void*, - int64_t, int64_t, float, void*) { - return false; -} inline int64_t attn_dpos_partials_bytes(int64_t, int64_t, int64_t) { return 0; } -// A gemm with its row bias added in the store: CUDA-first; the evaluator adds -// the bias with the broadcast kernel after gemm here. Outside the #if/#else -// like the ops below. -inline bool gemm_bias(void*, int64_t, int64_t, bool, void*, int64_t, int64_t, - bool, void*, int64_t, void*, int64_t, int64_t, int64_t, - int64_t, float, float) { - return false; -} - // Every CPU-side buffer read funnels through array::raw()/data(), which call // this: one choke point makes mixed CPU/GPU graphs safe. inline void cpu_barrier() { @@ -2007,3 +1352,5 @@ inline void sync_to_host(void*, bool) {} // live in gpu.h — one place, so array.h's eval seam stays #ifdef-free. } // namespace tl + +#endif // __APPLE__ diff --git a/include/metal_kernels.metal b/include/metal_kernels.metal index de29b61..074b71c 100644 --- a/include/metal_kernels.metal +++ b/include/metal_kernels.metal @@ -11,10 +11,13 @@ #include using namespace metal; +// Params structs follow the shared kernel ABI (gpu_abi.h): 4-byte fields in +// the order gpu_ops.h declares them, which is the order the CUDA kernel of the +// same name takes its scalars. struct ew_params { + uint n; float scale; float offset; - uint n; }; // bf16 is the top 16 bits of the f32 pattern: widen by a shift, narrow with @@ -58,11 +61,11 @@ EW_BINARY(pow_, pow(a[i], b[i])) // instead of falling back to the CPU mid-graph (each fallback costs a full // pipeline flush). Output is contiguous row-major [M, N]. struct ew_bcast_params { - float scale; - float offset; uint M; uint N; uint ars, acs, brs, bcs; + float scale; + float offset; }; #define EW_BCAST(name, expr) \ @@ -133,8 +136,8 @@ EW_CMP(ne_, a[i] != bv) // clamp(x, lo, hi): Clip's forward. No epilogue -- lo/hi occupy the role // scale/offset play elsewhere. struct clamp_params { - float lo, hi; uint n; + float lo, hi; }; kernel void clamp_(device const float* a [[buffer(0)]], @@ -148,8 +151,8 @@ kernel void clamp_(device const float* a [[buffer(0)]], // Tensor-scalar: out = f(a, s) * scale + offset -- mirrors tensorlib_cuda.cu's // TL_EW_SCALAR. struct scalar_params { - float s, scale, offset; uint n; + float s, scale, offset; }; #define EW_SCALAR(name, expr) \ @@ -2503,8 +2506,11 @@ attn_combine_<128>(device const float*, device float*, // One step's k/v ([n_kv_heads, D]) into cache row `pos`; grid (n_kv_heads), // D threads. KT is the cache element type. -struct kv_params { - uint pos, kv_stride, T; +struct kv_append_params { + uint pos, kv_stride; +}; +struct kv_fill_params { + uint T, kv_stride, pos0; }; template @@ -2512,7 +2518,7 @@ kernel void kv_append_(device KT* Kc [[buffer(0)]], device KT* Vc [[buffer(1)]], device const float* k_new [[buffer(2)]], device const float* v_new [[buffer(3)]], - constant kv_params& p [[buffer(4)]], + constant kv_append_params& p [[buffer(4)]], uint h [[threadgroup_position_in_grid]], uint d [[thread_index_in_threadgroup]], uint D [[threads_per_threadgroup]]) { @@ -2523,10 +2529,10 @@ kernel void kv_append_(device KT* Kc [[buffer(0)]], } template [[host_name("kv_append_")]] kernel void kv_append_(device float*, device float*, device const float*, - device const float*, constant kv_params&, uint, uint, uint); + device const float*, constant kv_append_params&, uint, uint, uint); template [[host_name("kv_append_bf16_")]] kernel void kv_append_(device ushort*, device ushort*, device const float*, - device const float*, constant kv_params&, uint, uint, uint); + device const float*, constant kv_append_params&, uint, uint, uint); // A prefill's k/v ([n_kv_heads, T, D]) into cache rows [pos, pos + T); grid // (n_kv_heads, T), D threads. @@ -2535,22 +2541,22 @@ kernel void kv_fill_(device KT* Kc [[buffer(0)]], device KT* Vc [[buffer(1)]], device const float* K [[buffer(2)]], device const float* V [[buffer(3)]], - constant kv_params& p [[buffer(4)]], + constant kv_fill_params& p [[buffer(4)]], uint2 g [[threadgroup_position_in_grid]], uint d [[thread_index_in_threadgroup]], uint2 nt [[threads_per_threadgroup]]) { const uint h = g.x, t = g.y, D = nt.x; - const uint dst = h * p.kv_stride + (p.pos + t) * D + d; + const uint dst = h * p.kv_stride + (p.pos0 + t) * D + d; const uint src = (h * p.T + t) * D + d; narrow_(Kc + dst, K[src]); narrow_(Vc + dst, V[src]); } template [[host_name("kv_fill_")]] kernel void kv_fill_(device float*, device float*, device const float*, - device const float*, constant kv_params&, uint2, uint, uint2); + device const float*, constant kv_fill_params&, uint2, uint, uint2); template [[host_name("kv_fill_bf16_")]] kernel void kv_fill_(device ushort*, device ushort*, device const float*, - device const float*, constant kv_params&, uint2, uint, uint2); + device const float*, constant kv_fill_params&, uint2, uint, uint2); // argmax of one length-n vector, the smallest index on ties (the host scan's // `v[i] > best`), so greedy decoding reads one int back instead of the logits. @@ -2607,27 +2613,24 @@ kernel void argmax_(device const float* in [[buffer(0)]], // layer's residual add folded into the norm that follows it; xout may alias // x). One threadgroup a row, 256 threads; 1/sqrt, as the array composition. struct rmsnorm_params { - uint n, add; + uint n; float eps; }; -kernel void rmsnorm_(device const float* x [[buffer(0)]], - device const float* delta [[buffer(1)]], - device const float* w [[buffer(2)]], - device float* xout [[buffer(3)]], - device float* hout [[buffer(4)]], - constant rmsnorm_params& p [[buffer(5)]], - uint row [[threadgroup_position_in_grid]], - uint tid [[thread_index_in_threadgroup]], - uint nt [[threads_per_threadgroup]], - uint sgid [[simdgroup_index_in_threadgroup]], - uint lane [[thread_index_in_simdgroup]]) { - threadgroup float red[8]; +// hout = v * rsqrt(mean(v^2) + eps) * w per row, where v is x, or with ADD +// x + delta (also stored to xout): the residual add folded into the norm that +// follows it. Without ADD, delta and xout are never touched. +template +static inline void rmsnorm_core_(device const float* x, device const float* delta, + device const float* w, device float* xout, + device float* hout, constant rmsnorm_params& p, + threadgroup float* red, uint row, uint tid, + uint nt, uint sgid, uint lane) { const uint base = row * p.n; float acc = 0.0f; for (uint i = tid; i < p.n; i += nt) { float v = x[base + i]; - if (p.add) { + if (ADD) { v += delta[base + i]; xout[base + i] = v; } @@ -2639,10 +2642,38 @@ kernel void rmsnorm_(device const float* x [[buffer(0)]], float ss = 0.0f; for (uint s = 0; s < nt / 32; s++) ss += red[s]; const float inv = 1.0f / sqrt(ss / float(p.n) + p.eps); - device const float* src = p.add ? xout : x; + device const float* src = ADD ? xout : x; for (uint i = tid; i < p.n; i += nt) hout[base + i] = src[base + i] * inv * w[i]; } +kernel void rmsnorm_(device const float* x [[buffer(0)]], + device const float* w [[buffer(1)]], + device float* out [[buffer(2)]], + constant rmsnorm_params& p [[buffer(3)]], + uint row [[threadgroup_position_in_grid]], + uint tid [[thread_index_in_threadgroup]], + uint nt [[threads_per_threadgroup]], + uint sgid [[simdgroup_index_in_threadgroup]], + uint lane [[thread_index_in_simdgroup]]) { + threadgroup float red[8]; + rmsnorm_core_(x, x, w, out, out, p, red, row, tid, nt, sgid, lane); +} + +kernel void add_rmsnorm_(device const float* x [[buffer(0)]], + device const float* delta [[buffer(1)]], + device const float* w [[buffer(2)]], + device float* xout [[buffer(3)]], + device float* hout [[buffer(4)]], + constant rmsnorm_params& p [[buffer(5)]], + uint row [[threadgroup_position_in_grid]], + uint tid [[thread_index_in_threadgroup]], + uint nt [[threads_per_threadgroup]], + uint sgid [[simdgroup_index_in_threadgroup]], + uint lane [[thread_index_in_simdgroup]]) { + threadgroup float red[8]; + rmsnorm_core_(x, delta, w, xout, hout, p, red, row, tid, nt, sgid, lane); +} + // out[row, f] = silu(gate) · up out of the fused gate|up projection // gu[row, 2·ff]; grid (ceil(ff / 256), rows). struct swiglu_params { @@ -2680,9 +2711,13 @@ kernel void split_heads_(device const float* src [[buffer(0)]], dst[(h * p.T + t) * D + d] = v; } +struct merge_heads_params { + uint T, H, D; +}; + kernel void merge_heads_(device const float* src [[buffer(0)]], device float* dst [[buffer(1)]], - constant heads_params& p [[buffer(2)]], + constant merge_heads_params& p [[buffer(2)]], uint2 g [[threadgroup_position_in_grid]], uint2 ng [[threadgroups_per_grid]], uint d [[thread_index_in_threadgroup]], diff --git a/include/storage.h b/include/storage.h index 3a15e86..087384b 100644 --- a/include/storage.h +++ b/include/storage.h @@ -33,6 +33,10 @@ struct storage { float* data() const { return ptr; } + // The whole buffer as the GPU layer names it (gpu_abi.h): a view at offset 0, + // from which span::at takes a slice. Null `buf` on a heap storage. + gpu::span device_span() const { return {native, 0}; } + // `host_fill`: the host fills it before any kernel touches it (see // gpu::alloc). static storage make(int64_t n, dtype dt = dtype::f32, bool host_fill = false) { diff --git a/include/webgpu.h b/include/webgpu.h index 28a5fa8..11ad0bc 100644 --- a/include/webgpu.h +++ b/include/webgpu.h @@ -37,23 +37,12 @@ #include -#include "metal.h" // reuse tl::metal::kop (platform-independent op enum) +#include "gpu_abi.h" // the op vocabulary and the launch contract #include "profile.h" #include "types.h" -namespace tl { -namespace webgpu { - -using kop = tl::metal::kop; -using cmp_op = tl::metal::cmp_op; -using unary_ext_op = tl::metal::unary_ext_op; -using scalar_op = tl::metal::scalar_op; - #if defined(TENSORLIB_WEBGPU) && defined(__EMSCRIPTEN__) -} // namespace webgpu -} // namespace tl - #include // emdawnwebgpu declares emscripten_webgpu_get_device() in webgpu.h itself, @@ -70,6 +59,11 @@ using scalar_op = tl::metal::scalar_op; namespace tl { namespace webgpu { +using kop = gpu::kop; +using cmp_op = gpu::cmp_op; +using unary_ext_op = gpu::unary_ext_op; +using scalar_op = gpu::scalar_op; + inline const char* wgsl_source_() { static const char* src = #include "tensorlib_webgpu_wgsl.inc" @@ -205,12 +199,11 @@ struct context { // 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. - enum loc { HOST, DEVICE, BOTH }; struct mirror { float* host = nullptr; // CPU-side buffer (storage.ptr) wgpu::Buffer dev; // device buffer size_t bytes = 0; - loc where = HOST; + gpu::residency live; // when to copy (gpu_abi.h); the copies are made here }; std::unordered_map mirrors; @@ -359,25 +352,20 @@ struct context { return it == mirrors.end() ? nullptr : &it->second; } - // A kernel is about to READ this buffer: ensure the device copy is current. + // A kernel is about to touch this buffer as `a`: bring the host copy up if + // residency says so. // // WriteBuffer executes in queue order, i.e. ahead of anything still sitting - // in the unsubmitted encoder. That is safe precisely because a buffer in - // HOST state has no encoded command touching it: a pending kernel write - // would have set DEVICE, and a pending kernel read would have come through - // here and set BOTH. - void device_read_(void* native) { + // in the unsubmitted encoder. That is safe precisely because a buffer whose + // live bytes are the host's has no encoded command touching it: a pending + // kernel write would have made the device copy live, and a pending kernel + // read would have come through here and uploaded already. + void before_kernel_(void* native, gpu::access a) { mirror* m = mirror_(native); - if (m && m->where == HOST) { - queue.WriteBuffer(m->dev, 0, m->host, m->bytes); - m->where = BOTH; - } - } - - // A kernel is about to WRITE this buffer: it becomes the live copy. - void device_write_(void* native) { - if (mirror* m = mirror_(native)) m->where = DEVICE; + if (m && m->live.before_kernel(a)) queue.WriteBuffer(m->dev, 0, m->host, m->bytes); } + void device_read_(void* native) { before_kernel_(native, gpu::access::in); } + void device_write_(void* native) { before_kernel_(native, gpu::access::out); } // The one place a dispatch is encoded. Every op differs only in which // pipeline, which params and what grid — keeping the bind group, uniform @@ -450,6 +438,42 @@ struct context { }; inline bool available() { return context::get().ready; } + +// The ops this backend runs its own way: a different algorithm, several +// kernels, or a kernel whose ABI is its own. gpu_ops.h forwards to whichever of +// these exist (TL_GPU_DETECT_OWN) and answers false for the rest, so a backend +// declares what it has and nothing else. Defined below, among their helpers. +struct own { + static bool binary_bcast_nd(kop op, gpu::span a, const int64_t* a_strides, + gpu::span b, const int64_t* b_strides, + gpu::span out, const int64_t* out_shape, int rank, + int64_t n, float scale, float offset); + static bool where_nd(gpu::span cond, const int64_t* c_strides, gpu::span a, + const int64_t* a_strides, gpu::span b, + const int64_t* b_strides, gpu::span out, + const int64_t* out_shape, int rank, int64_t n); + static bool sum_to(gpu::span a, const int64_t* a_shape, + const int64_t* a_strides, const int64_t* acc, int rank, + int64_t out_n, int64_t reduced_n, gpu::span out); + static bool pad(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, int64_t before, + int64_t n, int64_t out_n); + static bool fold(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, int64_t step, + int64_t n, int64_t out_n); + static bool concat_part(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t before, int64_t n); + static bool index_add(gpu::span idx, gpu::span values, gpu::span out, + int64_t row_size, int64_t k, int64_t out_n); + static bool scatter_to_axis(gpu::span idx, gpu::span values, gpu::span out, + int64_t n, int64_t size); + static bool gemm(gpu::span a, int64_t lda, bool ta, gpu::span b, int64_t ldb, + bool tb, gpu::span out, int64_t m, int64_t n, int64_t k, + float scale, float offset); + static bool rope(gpu::span x, gpu::span out, int64_t rows, int64_t T, + int64_t D, int64_t pos, float base, gpu::span bias = {}); +}; inline bool pending() { return context::get().pending; } // End the batch: submit the accumulated encoder and block until the GPU @@ -480,7 +504,7 @@ inline void context::flush_() { flush(); } // anything dereferenceable; it is only ever a key back into `mirrors`. // `host_fill` (the host writes it first) needs nothing here: the host copy is // its own malloc, which kernels never write and an upload copies when queued. -inline void* alloc(int64_t bytes, float** contents, bool /*host_fill*/ = false) { +inline void* alloc(int64_t bytes, float** contents, bool host_fill = false) { auto& c = context::get(); if (!c.ready) return nullptr; size_t nb = bytes > 0 ? (size_t)bytes : 4; @@ -508,7 +532,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, context::HOST}; + c.mirrors[token] = context::mirror{host, dev, nb, gpu::residency(host_fill)}; if (contents) *contents = host; return token; } @@ -531,7 +555,7 @@ inline void sync_to_host(void* native, bool for_write) { context::mirror* m = c.mirror_(native); if (!m) return; if (c.pending) flush(); - if (m->where == context::DEVICE) { + if (m->live.needs_download()) { // No CPU-visible pointer to read from: copy device -> a MapRead staging // buffer, map it (the second and last suspend point), memcpy out. wgpu::Buffer stg = c.staging_(m->bytes); @@ -549,21 +573,21 @@ inline void sync_to_host(void* native, bool for_write) { ok) { if (const void* src = stg.GetConstMappedRange(0, m->bytes)) { std::memcpy(m->host, src, m->bytes); - m->where = context::BOTH; + m->live.downloaded(); } stg.Unmap(); } - // Only a completed memcpy makes the host copy current. Declaring BOTH on a + // Only a completed memcpy makes the host copy current. Declaring it so on a // failed readback would leave stale bytes permanently believed live, and no // later sync_to_host would retry — so say so loudly instead, as the WGSL // compile failure above does. - if (m->where == context::DEVICE) { + if (m->live.needs_download()) { std::fprintf(stderr, "tensorlib webgpu: readback of %zu bytes failed\n", m->bytes); } c.staging_pool[m->bytes].push_back(stg); } - if (for_write) m->where = context::HOST; + if (for_write) m->live.host_wrote(); } // Shared host-side prologue for every op: resolve the operand mirrors and @@ -593,18 +617,18 @@ inline bool elem_off_(int64_t byte_off, uint32_t* out) { } // out = op(a) @ op(b) * scale + offset. -inline bool gemm(void* a, int64_t ao, int64_t lda, bool ta, void* b, int64_t bo, - int64_t ldb, bool tb, void* out, int64_t oo, int64_t m, - int64_t n, int64_t k, float scale, float offset) { +inline bool own::gemm(gpu::span a, int64_t lda, bool ta, gpu::span b, + int64_t ldb, bool tb, gpu::span out, int64_t m, int64_t n, + int64_t k, float scale, float offset) { auto& c = context::get(); if (!c.ready || m <= 0 || n <= 0 || k <= 0) return false; params p = {}; - if (!elem_off_(ao, &p.a_off) || !elem_off_(bo, &p.b_off) || - !elem_off_(oo, &p.c_off)) { + if (!elem_off_(a.off, &p.a_off) || !elem_off_(b.off, &p.b_off) || + !elem_off_(out.off, &p.c_off)) { return false; } context::mirror *ma, *mb, *mo; - if (!operands_(c, a, b, out, &ma, &mb, &mo)) return false; + if (!operands_(c, a.buf, b.buf, out.buf, &ma, &mb, &mo)) return false; p.M = (uint32_t)m; p.N = (uint32_t)n; @@ -619,24 +643,143 @@ inline bool gemm(void* a, int64_t ao, int64_t lda, bool ta, void* b, int64_t bo, return c.encode_("sgemm", ma, mb, mo, p, (n + 63) / 64, (m + 63) / 64); } -// Contiguous elementwise binary over n elements. -inline bool binary(kop op, void* a, int64_t ao, void* b, int64_t bo, void* out, - int64_t oo, int64_t n, float scale, float offset) { +// This backend's kernels predate the shared kernel ABI (gpu_abi.h): every WGSL +// entry point reads the one `params` layout above, and a family picks its +// operation by number. So each kernel id the shared ops may dispatch is +// marshalled here, from its canonical params into that layout; the entry +// point comes back, or null for an id this backend has no kernel for. +// +// `in_off` holds the element offsets of the views the kernel reads. A and B's +// are placed by dispatch; a kernel that reads D or E at an offset moves it +// into its own field and clears it here, so an offset nobody consumed is +// caught rather than dropped. +inline const char* marshal_(kop k, const void* canonical, uint32_t* in_off, + params& p) { + switch (k) { + case kop::badd: case kop::bsub: case kop::bmul: case kop::bdiv: + case kop::bpow: { + const auto& q = *static_cast(canonical); + p.M = q.m; + p.N = q.n; + p.ars = q.ars; + p.acs = q.acs; + p.brs = q.brs; + p.bcs = q.bcs; + p.op = kernel_op_(k); + p.scale = q.scale; + p.offset = q.offset; + return "ew_bcast"; + } + case kop::gt_: case kop::lt_: case kop::ge_: case kop::le_: case kop::eq_: + case kop::ne_: { + const auto& q = *static_cast(canonical); + if (q.n == 0) return nullptr; + p.M = q.n; + p.ars = q.bstride; + p.op = static_cast(k) - static_cast(kop::gt_); + return "cmp"; + } + case kop::clamp_: { // no epilogue: scale/offset carry lo/hi + const auto& q = *static_cast(canonical); + if (q.n == 0) return nullptr; + p.M = q.n; + p.scale = q.lo; + p.offset = q.hi; + return "clamp_"; + } + case kop::pow_s_: case kop::gt_s_: case kop::lt_s_: case kop::ge_s_: + case kop::le_s_: case kop::eq_s_: case kop::ne_s_: { + const auto& q = *static_cast(canonical); + if (q.n == 0) return nullptr; + p.M = q.n; + p.op = static_cast(k) - static_cast(kop::pow_s_); + p.arg = q.s; + p.scale = q.scale; + p.offset = q.offset; + return "ew_scalar"; + } + case kop::softmax: case kop::row_sum: case kop::row_max: { + const auto& q = *static_cast(canonical); + p.M = q.rows; + p.N = q.cols; + p.op = kernel_op_(k); + p.scale = q.scale; + p.offset = q.offset; + return k == kop::softmax ? "softmax" : "row_reduce"; + } + case kop::layer_norm_: { // A = x, B = g, D = b at pad3, arg = eps + const auto& q = *static_cast(canonical); + p.M = q.rows; + p.N = q.cols; + p.arg = q.eps; + p.scale = q.scale; + p.offset = q.offset; + p.pad3 = in_off[2]; + in_off[2] = 0; + return "layer_norm"; + } + case kop::index_select: { + const auto& q = *static_cast(canonical); + p.M = q.n; + p.pad0 = q.row_size; + return "index_select"; + } + case kop::add: case kop::sub: case kop::mul: case kop::div: case kop::pow_: + case kop::exp_: case kop::log_: case kop::sqrt_: case kop::sigmoid: + case kop::relu: case kop::affine: case kop::tanh_: case kop::sin_: + case kop::cos_: { + const auto& q = *static_cast(canonical); + if (q.n == 0) return nullptr; + p.M = q.n; + p.op = kernel_op_(k); + p.scale = q.scale; + p.offset = q.offset; + return k <= kop::pow_ ? "ew_binary" : "ew_unary"; + } + default: return nullptr; + } +} + +// The device core's one way to run a kernel for the shared ops. The bind +// group is fixed — A and B read, C written, D and E read — so the view a +// kernel writes is C and the ones it reads fill A, B, D, E in order; a +// one-input kernel binds its input twice. Offsets ride in the uniform as +// element counts (A, B and C only: D and E are bound whole). +inline bool dispatch(kop k, const gpu::arg* args, size_t n, + const void* canonical, size_t /*params_bytes*/, + const gpu::grid& g) { auto& c = context::get(); - if (!c.ready || n <= 0) return false; + if (!c.ready) return false; params p = {}; - if (!elem_off_(ao, &p.a_off) || !elem_off_(bo, &p.b_off) || - !elem_off_(oo, &p.c_off)) { - return false; + context::mirror* in[4] = {}; + uint32_t in_off[4] = {}; + context::mirror* out = nullptr; + size_t ins = 0; + for (size_t i = 0; i < n; i++) { + context::mirror* m = c.mirror_(args[i].s.buf); + uint32_t off = 0; + if (!m || !elem_off_(args[i].s.off, &off)) return false; // CPU's + if (args[i].a == gpu::access::in) { + if (ins == 4) return false; + in_off[ins] = off; + in[ins++] = m; + } else { + if (out) return false; // one writable binding + out = m; + p.c_off = off; + } } - context::mirror *ma, *mb, *mo; - if (!operands_(c, a, b, out, &ma, &mb, &mo)) return false; - - p.M = (uint32_t)n; - p.op = kernel_op_(op); - p.scale = scale; - p.offset = offset; - return c.encode_("ew_binary", ma, mb, mo, p, (n + 255) / 256, 1); + if (!out || ins == 0) return false; + if (ins == 1) { + in[1] = in[0]; + in_off[1] = in_off[0]; + } + p.a_off = in_off[0]; + p.b_off = in_off[1]; + const char* entry = marshal_(k, canonical, in_off, p); + if (!entry || in_off[2] || in_off[3]) return false; + for (size_t i = 0; i < n; i++) c.before_kernel_(args[i].s.buf, args[i].a); + return c.encode_(entry, in[0], in[1], out, p, g.gx, g.gy, in[2], in[3]); } // A one-input elementwise dispatch over n elements: the kernel binds its one @@ -656,106 +799,6 @@ inline bool encode_one_input_(const char* entry, void* a, int64_t ao, void* out, return c.encode_(entry, ma, mb, mo, p, (n + 255) / 256, 1); } -inline bool unary(kop op, void* a, int64_t ao, void* out, int64_t oo, int64_t n, - float scale, float offset) { - return encode_one_input_("ew_unary", a, ao, out, oo, n, [&](params& p) { - p.op = kernel_op_(op); - p.scale = scale; - p.offset = offset; - }); -} - -// Rank-2 broadcast binary: out[r,c] = f(a[r*ars + c*acs], b[r*brs + c*bcs]) -// into a contiguous [m,n] output. One stride-parameterized kernel covers every -// rank-2 broadcast, which keeps bias/gamma/beta chains on the GPU — falling -// back mid-graph would cost a full submit-and-wait. -inline bool binary_bcast(kop op, void* a, int64_t ao, int64_t ars, int64_t acs, - void* b, int64_t bo, int64_t brs, int64_t bcs, - void* out, int64_t oo, int64_t m, int64_t n, - float scale, float offset) { - auto& c = context::get(); - if (!c.ready || m <= 0 || n <= 0) return false; - // Broadcast strides are non-negative here (broadcast_strides only ever - // zeroes an axis); a negative one would wrap as u32 in the kernel. - if (ars < 0 || acs < 0 || brs < 0 || bcs < 0) return false; - params p = {}; - if (!elem_off_(ao, &p.a_off) || !elem_off_(bo, &p.b_off) || - !elem_off_(oo, &p.c_off)) { - return false; - } - context::mirror *ma, *mb, *mo; - if (!operands_(c, a, b, out, &ma, &mb, &mo)) return false; - - p.M = (uint32_t)m; - p.N = (uint32_t)n; - p.ars = (uint32_t)ars; - p.acs = (uint32_t)acs; - p.brs = (uint32_t)brs; - p.bcs = (uint32_t)bcs; - p.op = kernel_op_(op); - p.scale = scale; - p.offset = offset; - return c.encode_("ew_bcast", ma, mb, mo, p, (n + 31) / 32, (m + 7) / 8); -} - -// Batched GEMM in one launch: not on this backend yet — array.h's batched dot -// loops gemm per slice when this declines (CUDA folds the batch into its grid). -inline bool gemm_batched(void*, int64_t, int64_t, bool, int64_t, void*, - int64_t, int64_t, bool, int64_t, void*, int64_t, - int64_t, int64_t, int64_t, int64_t, float, float, - void* = nullptr, int64_t = 0) { - return false; -} - -// Row-wise op over the last axis: softmax writes rows x cols; row_sum/row_max -// write one value per row, with the affine epilogue. One workgroup per row. -inline bool row_op(kop op, void* in, int64_t io, void* out, int64_t oo, - int64_t rows, int64_t cols, float scale, float offset) { - auto& c = context::get(); - if (!c.ready || rows <= 0 || cols <= 0) return false; - params p = {}; - if (!elem_off_(io, &p.a_off) || !elem_off_(oo, &p.c_off)) return false; - context::mirror *ma, *mb, *mo; - if (!operands_(c, in, nullptr, out, &ma, &mb, &mo)) return false; - - p.b_off = p.a_off; - p.M = (uint32_t)rows; - p.N = (uint32_t)cols; - p.op = kernel_op_(op); - p.scale = scale; - p.offset = offset; - const char* entry = op == kop::softmax ? "softmax" : "row_reduce"; - return c.encode_(entry, ma, mb, mo, p, rows, 1); -} - -// Layer norm over the last axis: out = (x - mu) · 1/sqrt(var + eps) · g + b per -// row, affine epilogue; g and b are contiguous d-vectors. One workgroup per -// row, like row_op. A = x, B = g, D = b (p.pad3 its element offset), p.arg = -// eps. b's mirror is resolved before operands_ stages anything, so a decline -// leaves no half-staged operands behind. -inline bool layer_norm(void* x, int64_t xo, void* g, int64_t go, void* b, - int64_t bo, void* out, int64_t oo, int64_t rows, - int64_t cols, float eps, float scale, float offset) { - auto& c = context::get(); - if (!c.ready || rows <= 0 || cols <= 0) return false; - params p = {}; - if (!elem_off_(xo, &p.a_off) || !elem_off_(go, &p.b_off) || - !elem_off_(bo, &p.pad3) || !elem_off_(oo, &p.c_off)) { - return false; - } - context::mirror* mb = c.mirror_(b); - if (!mb) return false; - context::mirror *mx, *mg, *mo; - if (!operands_(c, x, g, out, &mx, &mg, &mo)) return false; - c.device_read_(b); - p.M = static_cast(rows); - p.N = static_cast(cols); - p.arg = eps; - p.scale = scale; - p.offset = offset; - return c.encode_("layer_norm", mx, mg, mo, p, rows, 1, mb); -} - // A ring, not one reused buffer: queue.WriteBuffer runs ahead of whatever is // still sitting in the unsubmitted encoder (see device_read_'s comment // above), so two pad/fold calls batched into the same unflushed pass would @@ -833,7 +876,7 @@ inline context::mirror* commit_meta_(context& c, void* ring_tok, // mirror never legitimately settles into a single steady HOST/DEVICE state // (each slot is written once, read once, never again). c.queue.WriteBuffer(mm->dev, word_off * 4, words, word_count * 4); - mm->where = context::BOTH; + mm->live.uploaded(); return mm; } @@ -846,19 +889,19 @@ inline context::mirror* commit_meta_(context& c, void* ring_tok, // binding as bit-reinterpreted u32 — WGSL's fixed Params uniform (used by // every other kernel here) has no room for a variable-length array, and // WriteBuffer is a raw byte copy regardless of the binding's declared type. -inline bool pad(void* a_native, int64_t ao, void* out_native, int64_t oo, - const int64_t* a_shape, const int64_t* out_shape, int rank, - int axis, int64_t before, int64_t n, int64_t out_n) { +inline bool own::pad(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t before, int64_t n, int64_t out_n) { (void)n; auto& c = context::get(); if (!c.ready || rank <= 0 || rank > kPadFoldMaxRank) return false; params p = {}; - if (!elem_off_(ao, &p.a_off) || !elem_off_(oo, &p.c_off)) return false; - context::mirror* ma = c.mirror_(a_native); - context::mirror* mo = c.mirror_(out_native); + if (!elem_off_(a.off, &p.a_off) || !elem_off_(out.off, &p.c_off)) return false; + context::mirror* ma = c.mirror_(a.buf); + context::mirror* mo = c.mirror_(out.buf); if (!ma || !mo) return false; - c.device_read_(a_native); - c.device_write_(out_native); + c.device_read_(a.buf); + c.device_write_(out.buf); uint32_t word_off, *raw; void* ring_tok = reserve_meta_(c, &word_off, &raw); @@ -882,19 +925,19 @@ inline bool pad(void* a_native, int64_t ao, void* out_native, int64_t oo, // atomicAdd fold, needed because WGSL has no float atomicAdd. `a`'s own // strides aren't part of the metadata: `a` is contiguous (gpu_fold_'s // contract), so the kernel derives them from `a_shape` itself. -inline bool fold(void* a_native, int64_t ao, void* out_native, int64_t oo, - const int64_t* a_shape, const int64_t* out_shape, int rank, - int axis, int64_t step, int64_t n, int64_t out_n) { +inline bool own::fold(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t step, int64_t n, int64_t out_n) { (void)n; auto& c = context::get(); if (!c.ready || rank <= 0 || rank > kPadFoldMaxRank) return false; params p = {}; - if (!elem_off_(ao, &p.a_off) || !elem_off_(oo, &p.c_off)) return false; - context::mirror* ma = c.mirror_(a_native); - context::mirror* mo = c.mirror_(out_native); + if (!elem_off_(a.off, &p.a_off) || !elem_off_(out.off, &p.c_off)) return false; + context::mirror* ma = c.mirror_(a.buf); + context::mirror* mo = c.mirror_(out.buf); if (!ma || !mo) return false; - c.device_read_(a_native); - c.device_write_(out_native); + c.device_read_(a.buf); + c.device_write_(out.buf); int out_rank = rank - 1; uint32_t word_off, *raw; @@ -914,45 +957,22 @@ inline bool fold(void* a_native, int64_t ao, void* out_native, int64_t oo, return c.encode_("fold", ma, mm, mo, p, (out_n + 255) / 256, 1); } -// Row gather along axis 0: out[i] = a[indices[row(i)]] (a, indices -// contiguous). A = a, B = idx, C = out; p.pad0 = row_size. -inline bool index_select(void* a_native, int64_t ao, void* idx_native, - int64_t idxo, void* out_native, int64_t oo, - int64_t row_size, int64_t k) { - auto& c = context::get(); - int64_t n = k * row_size; - if (!c.ready || n <= 0) return false; - params p = {}; - if (!elem_off_(ao, &p.a_off) || !elem_off_(idxo, &p.b_off) || - !elem_off_(oo, &p.c_off)) { - return false; - } - context::mirror *ma, *mb, *mo; - if (!operands_(c, a_native, idx_native, out_native, &ma, &mb, &mo)) { - return false; - } - p.M = static_cast(n); - p.pad0 = static_cast(row_size); - return c.encode_("index_select", ma, mb, mo, p, (n + 255) / 256, 1); -} - // index_select's dual, rewritten as a gather: WGSL has no float atomicAdd, // the same gap pad/fold above work around, so this sums over every source // row matching each OUTPUT row instead of scattering into a pre-zeroed // buffer -- no zeroing needed. A = idx, B = values, C = out; p.pad0 = // row_size, p.pad1 = k (number of source rows to scan). -inline bool index_add(void* idx_native, int64_t idxo, void* values_native, - int64_t vo, void* out_native, int64_t oo, - int64_t row_size, int64_t k, int64_t out_n) { +inline bool own::index_add(gpu::span idx, gpu::span values, gpu::span out, + int64_t row_size, int64_t k, int64_t out_n) { auto& c = context::get(); if (!c.ready || out_n <= 0) return false; params p = {}; - if (!elem_off_(idxo, &p.a_off) || !elem_off_(vo, &p.b_off) || - !elem_off_(oo, &p.c_off)) { + if (!elem_off_(idx.off, &p.a_off) || !elem_off_(values.off, &p.b_off) || + !elem_off_(out.off, &p.c_off)) { return false; } context::mirror *ma, *mb, *mo; - if (!operands_(c, idx_native, values_native, out_native, &ma, &mb, &mo)) { + if (!operands_(c, idx.buf, values.buf, out.buf, &ma, &mb, &mo)) { return false; } p.M = static_cast(out_n); @@ -965,19 +985,18 @@ inline bool index_add(void* idx_native, int64_t idxo, void* values_native, // values[pos] where indices[pos] == k, else 0. Every output element reads, // never writes twice, so -- like index_select above -- no zeroing needed. // A = idx, B = values, C = out; p.pad0 = size. -inline bool scatter_to_axis(void* idx_native, int64_t idxo, - void* values_native, int64_t vo, void* out_native, - int64_t oo, int64_t n, int64_t size) { +inline bool own::scatter_to_axis(gpu::span idx, gpu::span values, gpu::span out, + int64_t n, int64_t size) { auto& c = context::get(); int64_t out_n = n * size; if (!c.ready || out_n <= 0) return false; params p = {}; - if (!elem_off_(idxo, &p.a_off) || !elem_off_(vo, &p.b_off) || - !elem_off_(oo, &p.c_off)) { + if (!elem_off_(idx.off, &p.a_off) || !elem_off_(values.off, &p.b_off) || + !elem_off_(out.off, &p.c_off)) { return false; } context::mirror *ma, *mb, *mo; - if (!operands_(c, idx_native, values_native, out_native, &ma, &mb, &mo)) { + if (!operands_(c, idx.buf, values.buf, out.buf, &ma, &mb, &mo)) { return false; } p.M = static_cast(out_n); @@ -985,36 +1004,6 @@ inline bool scatter_to_axis(void* idx_native, int64_t idxo, return c.encode_("scatter_axis", ma, mb, mo, p, (out_n + 255) / 256, 1); } -// Cross-entropy's three: the trailing-axis gather, the one-pass row logsumexp -// and the pullback that reads it. CUDA-first (allowlisted); array.h composes -// the same values here. -inline bool gather_from_axis(void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool row_logsumexp(void*, int64_t, void*, int64_t, int64_t, int64_t, - float, float) { - return false; -} -inline bool xent_bwd(void*, int64_t, void*, int64_t, void*, int64_t, void*, - int64_t, void*, int64_t, int64_t, int64_t) { - return false; -} -// Layer norm's pullback. CUDA-first (allowlisted); the caller composes the -// unfused form when this declines. -inline bool layer_norm_bwd(void*, int64_t, void*, int64_t, void*, int64_t, - void*, void*, void*, void*, void*, int64_t, int64_t, - int64_t, int64_t, float) { - return false; -} - -// Adam's fused per-parameter update. CUDA-first (allowlisted); array.h takes -// its host loop when this declines, a D2H and H2D round trip here. -inline bool adam_step(void*, int64_t, void*, int64_t, void*, int64_t, void*, - int64_t, int64_t, float, float, float, float, float) { - return false; -} - // N-D broadcast binary: generalizes binary_bcast() above to any rank (a // Transformer's [N,S,D] LayerNorm broadcasting a [N,S,1] mean, rank 3). // a_strides/b_strides are the broadcast strides (0 on a broadcast axis) @@ -1022,21 +1011,20 @@ inline bool adam_step(void*, int64_t, void*, int64_t, void*, int64_t, void*, // oracle uses. A = a, B = b, D = meta [out_shape(rank), a_strides(rank), // b_strides(rank)] (a and b already fill A/B, unlike pad/fold where B was // free for this); p.pad0 = rank, p.pad3 = meta's word offset into D. -inline bool binary_bcast_nd(kop op, void* a_native, int64_t ao, - const int64_t* a_strides, void* b_native, - int64_t bo, const int64_t* b_strides, - void* out_native, int64_t oo, - const int64_t* out_shape, int rank, int64_t n, - float scale, float offset) { +inline bool own::binary_bcast_nd(kop op, gpu::span a, const int64_t* a_strides, + gpu::span b, const int64_t* b_strides, + gpu::span out, const int64_t* out_shape, + int rank, int64_t n, float scale, + float offset) { auto& c = context::get(); if (!c.ready || rank <= 0 || rank > kPadFoldMaxRank || n <= 0) return false; params p = {}; - if (!elem_off_(ao, &p.a_off) || !elem_off_(bo, &p.b_off) || - !elem_off_(oo, &p.c_off)) { + if (!elem_off_(a.off, &p.a_off) || !elem_off_(b.off, &p.b_off) || + !elem_off_(out.off, &p.c_off)) { return false; } context::mirror *ma, *mb, *mo; - if (!operands_(c, a_native, b_native, out_native, &ma, &mb, &mo)) { + if (!operands_(c, a.buf, b.buf, out.buf, &ma, &mb, &mo)) { return false; } @@ -1068,28 +1056,27 @@ inline bool binary_bcast_nd(kop op, void* a_native, int64_t ao, // rank, p.pad3 = b's element offset into D, p.pad4 = meta's word offset // into E. Bypasses operands_() (built for two real operands) since this one // needs three, the same way pad()/fold() above do their own mirror lookups. -inline bool where_nd(void* cond_native, int64_t co, const int64_t* c_strides, - void* a_native, int64_t ao, const int64_t* a_strides, - void* b_native, int64_t bo, const int64_t* b_strides, - void* out_native, int64_t oo, const int64_t* out_shape, - int rank, int64_t n) { +inline bool own::where_nd(gpu::span cond, const int64_t* c_strides, gpu::span a, + const int64_t* a_strides, gpu::span b, + const int64_t* b_strides, gpu::span out, + const int64_t* out_shape, int rank, int64_t n) { auto& c = context::get(); if (!c.ready || rank <= 0 || rank > kPadFoldMaxRank || n <= 0) return false; params p = {}; uint32_t b_elem_off; - if (!elem_off_(co, &p.a_off) || !elem_off_(ao, &p.b_off) || - !elem_off_(bo, &b_elem_off) || !elem_off_(oo, &p.c_off)) { + if (!elem_off_(cond.off, &p.a_off) || !elem_off_(a.off, &p.b_off) || + !elem_off_(b.off, &b_elem_off) || !elem_off_(out.off, &p.c_off)) { return false; } - context::mirror* mcond = c.mirror_(cond_native); - context::mirror* ma = c.mirror_(a_native); - context::mirror* mb = c.mirror_(b_native); - context::mirror* mo = c.mirror_(out_native); + context::mirror* mcond = c.mirror_(cond.buf); + context::mirror* ma = c.mirror_(a.buf); + context::mirror* mb = c.mirror_(b.buf); + context::mirror* mo = c.mirror_(out.buf); if (!mcond || !ma || !mb || !mo) return false; - c.device_read_(cond_native); - c.device_read_(a_native); - c.device_read_(b_native); - c.device_write_(out_native); + c.device_read_(cond.buf); + c.device_read_(a.buf); + c.device_read_(b.buf); + c.device_write_(out.buf); uint32_t word_off, *raw; void* ring_tok = reserve_meta_(c, &word_off, &raw); @@ -1114,31 +1101,23 @@ inline bool where_nd(void* cond_native, int64_t co, const int64_t* c_strides, return c.encode_("where_nd", mcond, ma, mo, p, (n + 255) / 256, 1, mb, mm); } -// clone()'s device arm, CUDA-first: no WebGPU kernel yet, so a clone of a -// device buffer takes array.h's host copy. -inline bool copy_nd(void*, int64_t, const int64_t*, void*, int64_t, - const int64_t*, int, int64_t) { - return false; -} - // sum_to (un-broadcast a gradient): gather, mirrors cuda.h's tl_sum_to and // metal.h's own sum_to -- one invocation per OUTPUT element sums every `a` // element that broadcasts onto it, so no atomics (unlike index_add). Only // one real tensor operand (`a`), so -- like pad/fold above -- B is free for // the meta ring: [a_shape(rank), a_strides(rank), acc(rank)]. -inline bool sum_to(void* a_native, int64_t ao, const int64_t* a_shape, - const int64_t* a_strides, const int64_t* acc, int rank, - int64_t out_n, int64_t reduced_n, void* out_native, - int64_t oo) { +inline bool own::sum_to(gpu::span a, const int64_t* a_shape, + const int64_t* a_strides, const int64_t* acc, int rank, + int64_t out_n, int64_t reduced_n, gpu::span out) { auto& c = context::get(); if (!c.ready || rank <= 0 || rank > kPadFoldMaxRank) return false; params p = {}; - if (!elem_off_(ao, &p.a_off) || !elem_off_(oo, &p.c_off)) return false; - context::mirror* ma = c.mirror_(a_native); - context::mirror* mo = c.mirror_(out_native); + if (!elem_off_(a.off, &p.a_off) || !elem_off_(out.off, &p.c_off)) return false; + context::mirror* ma = c.mirror_(a.buf); + context::mirror* mo = c.mirror_(out.buf); if (!ma || !mo) return false; - c.device_read_(a_native); - c.device_write_(out_native); + c.device_read_(a.buf); + c.device_write_(out.buf); uint32_t word_off, *raw; void* ring_tok = reserve_meta_(c, &word_off, &raw); @@ -1162,76 +1141,6 @@ inline bool sum_to(void* a_native, int64_t ao, const int64_t* a_shape, return c.encode_("sum_to", ma, mm, mo, p, (out_n + 255) / 256, 1); } -// gt/lt/ge/le/eq/ne (array.h's comparison ops, ReLU/LeakyReLU/Clip's -// backward gate): same-shape only (bstride=1) or a scalar b (bstride=0) -- -// the two shapes array.h's gpu_compare_ ever dispatches. p.ars (unused by -// this family) carries bstride. Own entry point/op numbering, not folded -// into ew_binary's binary_op -- it returns a bool-as-float mask rather than -// composing with scale/offset. -inline bool compare(cmp_op op, void* a, int64_t ao, void* b, int64_t bo, - void* out, int64_t oo, int64_t n, int64_t bstride) { - auto& c = context::get(); - if (!c.ready || n <= 0) return false; - params p = {}; - if (!elem_off_(ao, &p.a_off) || !elem_off_(bo, &p.b_off) || - !elem_off_(oo, &p.c_off)) { - return false; - } - context::mirror *ma, *mb, *mo; - if (!operands_(c, a, b, out, &ma, &mb, &mo)) return false; - p.M = static_cast(n); - p.ars = static_cast(bstride); - switch (op) { - case cmp_op::gt: p.op = 0; break; - case cmp_op::lt: p.op = 1; break; - case cmp_op::ge: p.op = 2; break; - case cmp_op::le: p.op = 3; break; - case cmp_op::eq: p.op = 4; break; - case cmp_op::ne: p.op = 5; break; - } - return c.encode_("cmp", ma, mb, mo, p, (n + 255) / 256, 1); -} - -// tanh_/sin_/cos_ (RoPE's trig, RNN/LSTM's tanh): plain elementwise, same -// shape as exp_/sqrt_ -- unary() above already does exactly this dispatch -// through ew_unary, just keyed by a kop array.h doesn't see directly. -inline bool unary_ext(unary_ext_op op, void* a, int64_t ao, void* out, - int64_t oo, int64_t n, float scale, float offset) { - kop k; - switch (op) { - case unary_ext_op::tanh_: k = kop::tanh_; break; - case unary_ext_op::sin_: k = kop::sin_; break; - case unary_ext_op::cos_: k = kop::cos_; break; - } - // Qualified: unqualified unary(...) is ambiguous here -- kop is really - // tl::metal::kop, so ADL pulls in metal.h's non-Apple unary() stub - // alongside this namespace's own. - return tl::webgpu::unary(k, a, ao, out, oo, n, scale, offset); -} - -// clamp(x, lo, hi): Clip's forward. No epilogue -- p.scale/p.offset carry -// lo/hi instead (mirrors cuda.h's/metal.h's own clamp). -inline bool clamp(void* a, int64_t ao, void* out, int64_t oo, int64_t n, - float lo, float hi) { - return encode_one_input_("clamp_", a, ao, out, oo, n, [&](params& p) { - p.scale = lo; - p.offset = hi; - }); -} - -// Tensor-scalar ops (metal.h's scalar_op): s rides in p.arg; p.op is -// scalar_op's value (0 pow, then cmp_op's order). -inline bool scalar_binary(scalar_op op, void* a, int64_t ao, void* out, - int64_t oo, int64_t n, float s, float scale, - float offset) { - return encode_one_input_("ew_scalar", a, ao, out, oo, n, [&](params& p) { - p.op = static_cast(op); - p.arg = s; - p.scale = scale; - p.offset = offset; - }); -} - // concat_part (Tensor.concat along an arbitrary axis, KV-cache append): // scatters `a` (this part) into `out` at a flat element shift along one // axis -- see kernels/tensorlib_webgpu.wgsl's own concat_part for why this @@ -1241,21 +1150,20 @@ inline bool scalar_binary(scalar_op op, void* a, int64_t ao, void* out, // (`a`), so -- like pad/fold/sum_to above -- B is free for the meta ring: // [a_shape(rank), out_strides(rank)] (out_strides computed host-side, // mirrors cuda.h's own upload_pad_fold_meta_). -inline bool concat_part(void* a_native, int64_t ao, void* out_native, - int64_t oo, const int64_t* a_shape, - const int64_t* out_shape, int rank, int axis, - int64_t before, int64_t n) { +inline bool own::concat_part(gpu::span a, gpu::span out, const int64_t* a_shape, + const int64_t* out_shape, int rank, int axis, + int64_t before, int64_t n) { auto& c = context::get(); if (!c.ready || rank <= 0 || rank > kPadFoldMaxRank || n <= 0) { return false; } params p = {}; - if (!elem_off_(ao, &p.a_off) || !elem_off_(oo, &p.c_off)) return false; - context::mirror* ma = c.mirror_(a_native); - context::mirror* mo = c.mirror_(out_native); + if (!elem_off_(a.off, &p.a_off) || !elem_off_(out.off, &p.c_off)) return false; + context::mirror* ma = c.mirror_(a.buf); + context::mirror* mo = c.mirror_(out.buf); if (!ma || !mo) return false; - c.device_read_(a_native); - c.device_write_(out_native); + c.device_read_(a.buf); + c.device_write_(out.buf); int64_t out_strides[kPadFoldMaxRank]; int64_t acc = 1; @@ -1289,16 +1197,17 @@ inline bool concat_part(void* a_native, int64_t ao, void* out_native, // the WGSL). x/out always view at offset 0 (array.h's gpu_rope_ requires // x.offset_ == 0 and hands a fresh allocation for out), so there is no // ao/oo in this signature to convert. -inline bool rope(void* x, void* out, int64_t rows, int64_t T, int64_t D, - int64_t pos, float base, void* bias = nullptr) { +inline bool own::rope(gpu::span x, gpu::span out, int64_t rows, int64_t T, + int64_t D, int64_t pos, float base, gpu::span bias) { auto& c = context::get(); - if (!c.ready || D <= 0 || (D & 1) || bias) return false; // no fused bias + if (!c.ready || D <= 0 || (D & 1) || bias.buf) return false; // no fused bias int64_t half = D / 2; int64_t n = rows * half; if (n <= 0) return false; params p = {}; + if (!elem_off_(x.off, &p.a_off) || !elem_off_(out.off, &p.c_off)) return false; context::mirror *ma, *mb, *mo; - if (!operands_(c, x, nullptr, out, &ma, &mb, &mo)) return false; + if (!operands_(c, x.buf, nullptr, out.buf, &ma, &mb, &mo)) return false; p.b_off = p.a_off; p.M = static_cast(n); p.N = static_cast(T); @@ -1309,195 +1218,22 @@ inline bool rope(void* x, void* out, int64_t rows, int64_t T, int64_t D, return c.encode_("rope", ma, mb, mo, p, (n + 255) / 256, 1); } -#else // !(TENSORLIB_WEBGPU && __EMSCRIPTEN__) — stubs, as in metal.h - -inline bool available() { return false; } -inline bool pending() { return false; } -inline void flush() {} -inline void* alloc(int64_t, float**, bool = false) { return nullptr; } -inline void release(void*, int64_t, float*) {} -inline void sync_to_host(void*, bool) {} -inline bool binary(kop, void*, int64_t, void*, int64_t, void*, int64_t, int64_t, - float, float) { - return false; -} -inline bool binary_bcast(kop, void*, int64_t, int64_t, int64_t, void*, int64_t, - int64_t, int64_t, void*, int64_t, int64_t, int64_t, - float, float) { - return false; -} -inline bool unary(kop, void*, int64_t, void*, int64_t, int64_t, float, float) { - return false; -} -inline bool gemm(void*, int64_t, int64_t, bool, void*, int64_t, int64_t, bool, - void*, int64_t, int64_t, int64_t, int64_t, float, float) { - return false; -} -inline bool gemm_batched(void*, int64_t, int64_t, bool, int64_t, void*, - int64_t, int64_t, bool, int64_t, void*, int64_t, - int64_t, int64_t, int64_t, int64_t, float, float, - void* = nullptr, int64_t = 0) { - return false; -} -inline bool row_op(kop, void*, int64_t, void*, int64_t, int64_t, int64_t, float, - float) { - return false; -} -inline bool layer_norm(void*, int64_t, void*, int64_t, void*, int64_t, void*, - int64_t, int64_t, int64_t, float, float, float) { - return false; -} -inline bool pad(void*, int64_t, void*, int64_t, const int64_t*, - const int64_t*, int, int, int64_t, int64_t, int64_t) { - return false; -} -inline bool fold(void*, int64_t, void*, int64_t, const int64_t*, - const int64_t*, int, int, int64_t, int64_t, int64_t) { - return false; -} -inline bool index_select(void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool index_add(void*, int64_t, void*, int64_t, void*, int64_t, int64_t, - int64_t, int64_t) { - return false; -} -inline bool scatter_to_axis(void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool gather_from_axis(void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool row_logsumexp(void*, int64_t, void*, int64_t, int64_t, int64_t, - float, float) { - return false; -} -inline bool xent_bwd(void*, int64_t, void*, int64_t, void*, int64_t, void*, - int64_t, void*, int64_t, int64_t, int64_t) { - return false; -} -// Layer norm's pullback. CUDA-first (allowlisted); the caller composes the -// unfused form when this declines. -inline bool layer_norm_bwd(void*, int64_t, void*, int64_t, void*, int64_t, - void*, void*, void*, void*, void*, int64_t, int64_t, - int64_t, int64_t, float) { - return false; -} -inline bool adam_step(void*, int64_t, void*, int64_t, void*, int64_t, void*, - int64_t, int64_t, float, float, float, float, float) { - return false; -} -inline bool binary_bcast_nd(kop, void*, int64_t, const int64_t*, void*, - int64_t, const int64_t*, void*, int64_t, - const int64_t*, int, int64_t, float, float) { - return false; -} -inline bool where_nd(void*, int64_t, const int64_t*, void*, int64_t, - const int64_t*, void*, int64_t, const int64_t*, void*, - int64_t, const int64_t*, int, int64_t) { - return false; -} -inline bool copy_nd(void*, int64_t, const int64_t*, void*, int64_t, - const int64_t*, int, int64_t) { - return false; -} -inline bool sum_to(void*, int64_t, const int64_t*, const int64_t*, - const int64_t*, int, int64_t, int64_t, void*, int64_t) { - return false; -} -inline bool compare(cmp_op, void*, int64_t, void*, int64_t, void*, int64_t, - int64_t, int64_t) { - return false; -} -inline bool unary_ext(unary_ext_op, void*, int64_t, void*, int64_t, int64_t, - float, float) { - return false; -} -inline bool clamp(void*, int64_t, void*, int64_t, int64_t, float, float) { - return false; -} -inline bool scalar_binary(scalar_op, void*, int64_t, void*, int64_t, int64_t, - float, float, float) { - return false; -} -inline bool concat_part(void*, int64_t, void*, int64_t, const int64_t*, - const int64_t*, int, int, int64_t, int64_t) { - return false; -} -inline bool rope(void*, void*, int64_t, int64_t, int64_t, int64_t, float, - void* = nullptr) { - return false; -} - -#endif - -// A gemm with its row bias added in the store: CUDA-first; the evaluator adds -// the bias with the broadcast kernel after gemm here. Outside the #if/#else -// like the ops below. -inline bool gemm_bias(void*, int64_t, int64_t, bool, void*, int64_t, int64_t, - bool, void*, int64_t, void*, int64_t, int64_t, int64_t, - int64_t, float, float) { - return false; -} - -// ---- Ops with no WGSL kernel yet (the LLM decode path). Outside the #if/#else -// on purpose: both branches would define them identically, and returning false -// is the whole implementation either way — it routes the op to CPU, which is -// why each porting phase lands in a working state. -inline bool gemv_f32(void*, void*, void*, int64_t, int64_t) { return false; } -inline bool gemv_bf16(void*, void*, void*, int64_t, int64_t) { return false; } -inline bool attn_decode(void*, void*, void*, void*, int64_t, int64_t, int64_t, - int64_t, int64_t, float, bool = false) { - return false; -} -inline bool attn_prefill(void*, void*, void*, void*, int64_t, int64_t, int64_t, - int64_t, int64_t, float, bool = false, int64_t = 0) { - return false; -} -inline bool attn_prefill_dq(void*, void*, void*, void*, void*, void*, void*, - int64_t, int64_t, int64_t, float) { - return false; -} -inline bool attn_prefill_dkv(void*, void*, void*, void*, void*, void*, void*, - int64_t, int64_t, int64_t, float) { - return false; -} -inline bool gemv_q4(void*, void*, void*, void*, int64_t, int64_t, int64_t) { - return false; -} -inline bool kv_append(void*, void*, void*, void*, int64_t, int64_t, int64_t, - int64_t, bool = false) { - return false; -} -inline bool kv_fill(void*, void*, void*, void*, int64_t, int64_t, int64_t, - int64_t, bool = false, int64_t = 0) { - return false; -} -inline bool argmax(void*, int64_t, int64_t*) { return false; } -inline bool rmsnorm(void*, void*, void*, int64_t, float, int64_t = 1) { - return false; -} -inline bool rmsnorm_res(void*, void*, void*, void*, void*, int64_t, float, - int64_t = 1) { - return false; -} -inline bool swiglu(void*, void*, int64_t, int64_t = 1) { return false; } -inline bool split_heads(void*, void*, void*, int64_t, int64_t, int64_t, int64_t, - int64_t) { - return false; -} -inline bool merge_heads(void*, void*, int64_t, int64_t, int64_t) { return false; } -inline bool gemv_bf16_row(void*, void*, void*, int64_t, int64_t) { return false; } -inline bool gemm_bf16_nt(void*, void*, void*, int64_t, int64_t, int64_t) { - return false; -} // What a model may ask of this backend beyond the kernel contract (gpu.h), // and the graph-capture group it names: none of it here, so each answers // false or does nothing and a decoder takes its host-position path. +// What the shared launch policy (gpu_ops.h) may assume of this backend's +// kernels. +struct traits { + // A [rows, cols] elementwise kernel reads its cell from a 2-D thread + // position rather than a flat index. + static constexpr bool cells_2d = true; + // Launches are recorded under tl::profile by this backend itself, with + // (times_launches) a device time on each. + static constexpr bool profiles_launches = true; + static constexpr bool times_launches = false; +}; + struct caps { // Whether the model-path row is real here, or answers false: a decoder // runs on raw buffers only where it is true, and keeps to the array ops @@ -1506,7 +1242,6 @@ struct caps { static constexpr bool graph_capture = false; static constexpr bool row_gemv = false; static constexpr bool bf16_gemm = false; - static constexpr bool flat_addressing = false; // `native` is a WGPUBuffer }; using graph_exec = void*; inline bool graph_available() { return false; } @@ -1516,19 +1251,6 @@ inline bool graph_launch(graph_exec) { return false; } inline void graph_destroy(graph_exec) {} inline void upload(void*, const float*, int64_t) {} inline void upload_u32(void*, unsigned) {} -inline bool incr_u32(void*) { return false; } -inline bool rope_dpos(void*, void*, int64_t, int64_t, int64_t, void*, float, - void* = nullptr) { - return false; -} -inline bool kv_append_dpos(void*, void*, void*, void*, void*, int64_t, int64_t, - int64_t) { - return false; -} -inline bool attn_decode_dpos(void*, void*, void*, void*, int64_t, int64_t, void*, - int64_t, int64_t, float, void*) { - return false; -} inline int64_t attn_dpos_partials_bytes(int64_t, int64_t, int64_t) { return 0; } // Every CPU-side buffer read funnels through array::raw()/data(), which call @@ -1539,3 +1261,5 @@ inline void cpu_barrier() { } // namespace webgpu } // namespace tl + +#endif // TENSORLIB_WEBGPU && __EMSCRIPTEN__ diff --git a/test/test_array.cpp b/test/test_array.cpp index c01d36e..7171f2a 100644 --- a/test/test_array.cpp +++ b/test/test_array.cpp @@ -3,6 +3,8 @@ #include #include +#include +#include #include using tl::array; @@ -1426,12 +1428,11 @@ TEST_CASE("layer_norm_bwd matches the composed pullback on the GPU and the own C if (on_gpu) tl::use_gpu(); auto got = tl::array::layer_norm_bwd(x, g, dy); tl::use_cpu(); - // The own CPU takes every contiguous input, and so do the CUDA and Metal - // kernels; WebGPU declines and its caller composes the form above. + // The own CPU takes every contiguous input, and so does a backend that + // has the kernel; one without it declines and its caller composes the + // form above. if (!on_gpu) REQUIRE(got.has_value()); -#if !defined(TENSORLIB_WEBGPU) - REQUIRE(got.has_value()); -#endif + if (tl::gpu::has_layer_norm_bwd) REQUIRE(got.has_value()); if (!got) return; auto want = composed(x, g, dy); for (int i = 0; i < 3; i++) { @@ -1774,13 +1775,10 @@ static void check_attn_bwd_dq(bool on_gpu, int64_t D = 64) { auto out = tl::array::attn_prefill(q, K, V, scale); auto got = tl::array::attn_prefill_bwd_dq(q, K, V, dO, out, scale); - // The own CPU always takes it, and so do the CUDA and Metal kernels — while - // WebGPU declines and its caller composes the unfused form, which the - // gradient tests above cover. - if (!on_gpu) REQUIRE(got.has_value()); -#if !defined(TENSORLIB_WEBGPU) - REQUIRE(got.has_value()); -#endif + // The own CPU always takes it, and so does a backend that has the kernel — + // while one without it declines and its caller composes the unfused form, + // which the gradient tests above cover. + if (!on_gpu || tl::gpu::has_attn_prefill_dq) REQUIRE(got.has_value()); if (!got) { MESSAGE("no fused attention pullback on this backend — skipping"); return; @@ -1858,10 +1856,7 @@ static void check_attn_bwd_dkv(bool on_gpu, int64_t D = 64) { auto got = dqs ? tl::array::attn_prefill_bwd_dkv(q, K, V, dO, dqs->second, scale) : std::nullopt; - if (!on_gpu) REQUIRE(got.has_value()); -#if !defined(TENSORLIB_WEBGPU) - REQUIRE(got.has_value()); // as above -#endif + if (!on_gpu || tl::gpu::has_attn_prefill_dkv) REQUIRE(got.has_value()); // as above if (!got) { MESSAGE("no fused attention pullback on this backend — skipping"); return; @@ -2849,14 +2844,14 @@ TEST_CASE("profile: scopes nest into paths and launches land under them") { const row* wait = find("phase/mm", row::kind_t::wait); REQUIRE(wait); CHECK(wait->count >= 1); -#if !defined(TENSORLIB_WEBGPU) - for (const row* r : launches) { - CHECK(r->device_timed == r->count); - CHECK(r->device_us > 0); + if (tl::gpu::traits::times_launches) { + for (const row* r : launches) { + CHECK(r->device_timed == r->count); + CHECK(r->device_us > 0); + } } -#endif -#ifdef __APPLE__ - CHECK(tl::profile::summarize().batches >= 1); +#if defined(__APPLE__) && !defined(TENSORLIB_HOST_GPU) + CHECK(tl::profile::summarize().batches >= 1); // Metal times a batch too #endif } else { CHECK(launches.empty()); // the CPU path launches nothing @@ -2955,12 +2950,12 @@ TEST_CASE("the KV cache and the decode step's kernels match their array forms") for (int64_t t = 0; t < T; t++) { array k = dev(K.slice(1, t, 1).clone()); // [HKV,1,D] array v = dev(V.slice(1, t, 1).clone()); - REQUIRE(cache.append(k.native(), v.native())); + REQUIRE(cache.append(k.device_span(), v.device_span())); } CHECK(cache.pos == T); array q = dev(random_array({HQ, D}, 902)); array out = array::empty({HQ, D}); - REQUIRE(cache.attn(q.native(), out.native(), HQ, scale)); + REQUIRE(cache.attn(q.device_span(), out.device_span(), HQ, scale)); // Explicit: each q head attends its kv head's T rows. auto Kr = kv == tl::dtype::bf16 ? K.to_bf16().to_f32() : K; auto Vr = kv == tl::dtype::bf16 ? V.to_bf16().to_f32() : V; @@ -2992,18 +2987,18 @@ TEST_CASE("the KV cache and the decode step's kernels match their array forms") array qp = dev(random_array({HQ, T, D}, 903)); array Kd = dev(K), Vd = dev(V); array op = array::empty({HQ, T, D}); - REQUIRE(c2.prefill(qp.native(), Kd.native(), Vd.native(), op.native(), T, + REQUIRE(c2.prefill(qp.device_span(), Kd.device_span(), Vd.device_span(), op.device_span(), T, HQ, scale)); CHECK(c2.pos == T); array out2 = array::empty({HQ, D}); - REQUIRE(c2.attn(q.native(), out2.native(), HQ, scale)); + REQUIRE(c2.attn(q.device_span(), out2.device_span(), HQ, scale)); tl::gpu::flush(); CHECK(same(out2, out, 1e-5f)); // Row t of the prefill output is the decode of q row t over keys 0..t. const int64_t t = T - 1; array qt = dev(qp.slice(1, t, 1).reshape({HQ, D}).clone()); array ot = array::empty({HQ, D}); - REQUIRE(c2.attn(qt.native(), ot.native(), HQ, scale)); + REQUIRE(c2.attn(qt.device_span(), ot.device_span(), HQ, scale)); tl::gpu::flush(); CHECK(same(ot, op.slice(1, t, 1).reshape({HQ, D}), 1e-4f)); } @@ -3016,9 +3011,9 @@ TEST_CASE("the KV cache and the decode step's kernels match their array forms") array w = dev(random_array({n}, 912)); array h = array::empty({rows, n}), xo = array::empty({rows, n}), h2 = array::empty({rows, n}); - REQUIRE(gpu::rmsnorm(x.native(), w.native(), h.native(), n, 1e-6f, rows)); - REQUIRE(gpu::rmsnorm_res(x.native(), d.native(), w.native(), xo.native(), - h2.native(), n, 1e-6f, rows)); + REQUIRE(gpu::rmsnorm(x.device_span(), w.device_span(), h.device_span(), n, 1e-6f, rows)); + REQUIRE(gpu::rmsnorm_res(x.device_span(), d.device_span(), w.device_span(), xo.device_span(), + h2.device_span(), n, 1e-6f, rows)); tl::gpu::flush(); CHECK(same(h, array::rmsnorm(x, w, 1e-6f), 1e-5f)); CHECK(same(xo, x + d, 1e-6f)); @@ -3029,13 +3024,13 @@ TEST_CASE("the KV cache and the decode step's kernels match their array forms") // by a factor of ten rather than the last bits. array tiny = dev(x * 1e-4f); array ht = array::empty({rows, n}); - REQUIRE(gpu::rmsnorm(tiny.native(), w.native(), ht.native(), n, 1e-6f, rows)); + REQUIRE(gpu::rmsnorm(tiny.device_span(), w.device_span(), ht.device_span(), n, 1e-6f, rows)); tl::gpu::flush(); CHECK(same(ht, array::rmsnorm(tiny, w, 1e-6f), 1e-5f)); array gu = dev(random_array({rows, 2 * ff}, 913)); array o = array::empty({rows, ff}); - REQUIRE(gpu::swiglu(gu.native(), o.native(), ff, rows)); + REQUIRE(gpu::swiglu(gu.device_span(), o.device_span(), ff, rows)); tl::gpu::flush(); array gate = gu.slice(1, 0, ff), up = gu.slice(1, ff, ff); CHECK(same(o, array::swiglu(gate, up), 1e-5f)); @@ -3045,7 +3040,7 @@ TEST_CASE("the KV cache and the decode step's kernels match their array forms") const int64_t H = 14, D = 64; array x = dev(random_array({H, D}, 920)), b = dev(random_array({H, D}, 921)); array o = array::empty({H, D}); - REQUIRE(gpu::rope(x.native(), o.native(), H, 1, D, 37, 1e6f, b.native())); + REQUIRE(gpu::rope(x.device_span(), o.device_span(), H, 1, D, 37, 1e6f, b.device_span())); tl::gpu::flush(); CHECK(same(o, array::rope(x + b, 37, 1e6f), 1e-5f)); } @@ -3056,7 +3051,7 @@ TEST_CASE("the KV cache and the decode step's kernels match their array forms") v[77778] = 5.0f; // a tie: the smaller index wins, as the host scan does array a = dev(array::from(v, {(int64_t)v.size()})); int64_t idx = -1; - REQUIRE(gpu::argmax(a.native(), (int64_t)v.size(), &idx)); + REQUIRE(gpu::argmax(a.device_span(), (int64_t)v.size(), &idx)); CHECK(idx == 77777); } @@ -3076,14 +3071,14 @@ TEST_CASE("the KV cache and the decode step's kernels match their array forms") if (gpu::caps::row_gemv) { array a = dev(random_array({1, K}, 941)); array y = array::empty({1, N}); - REQUIRE(gpu::gemv_bf16_row(a.native(), Wb.native(), y.native(), N, K)); + REQUIRE(gpu::gemv_bf16_row(a.device_span(), Wb.device_span(), y.device_span(), N, K)); tl::gpu::flush(); CHECK(same(y, a.dot(Wt), 1e-4f)); } if (gpu::caps::bf16_gemm) { array A = dev(random_array({M, K}, 942)); array C = array::empty({M, N}); - REQUIRE(gpu::gemm_bf16_nt(A.native(), Wb.native(), C.native(), M, N, K)); + REQUIRE(gpu::gemm_bf16_nt(A.device_span(), Wb.device_span(), C.device_span(), M, N, K)); tl::gpu::flush(); CHECK(same(C, A.dot(Wt), 1e-4f)); // Row m of the GEMM is the GEMV of row m — the prefill and the decode @@ -3092,7 +3087,7 @@ TEST_CASE("the KV cache and the decode step's kernels match their array forms") if (gpu::caps::row_gemv) { array row = dev(A.slice(0, 7, 1).clone()); array y1 = array::empty({1, N}); - REQUIRE(gpu::gemv_bf16_row(row.native(), Wb.native(), y1.native(), N, K)); + REQUIRE(gpu::gemv_bf16_row(row.device_span(), Wb.device_span(), y1.device_span(), N, K)); tl::gpu::flush(); CHECK(same(y1, C.slice(0, 7, 1), 1e-4f)); } @@ -3103,9 +3098,9 @@ TEST_CASE("the KV cache and the decode step's kernels match their array forms") const int64_t T = 5, H = 18, D = 64, ld = H * D; array src = dev(random_array({T, ld}, 930)), bias = dev(random_array({H, D}, 931)); array heads = array::empty({H, T, D}), back = array::empty({T, ld}); - REQUIRE(gpu::split_heads(src.native(), bias.native(), heads.native(), T, ld, 0, + REQUIRE(gpu::split_heads(src.device_span(), bias.device_span(), heads.device_span(), T, ld, 0, H, D)); - REQUIRE(gpu::merge_heads(heads.native(), back.native(), T, H, D)); + REQUIRE(gpu::merge_heads(heads.device_span(), back.device_span(), T, H, D)); tl::gpu::flush(); // heads[h, t, :] = src[t, h*D:(h+1)*D] + bias[h] CHECK(same(heads, src.reshape({T, H, D}).transpose({1, 0, 2}) + @@ -3115,3 +3110,359 @@ TEST_CASE("the KV cache and the decode step's kernels match their array forms") tl::device_ = prev; } + +// The shared ops (gpu_ops.h) take views: a device handle and a byte offset. +// Each one here runs on views at distinct non-zero offsets into larger +// buffers and is checked against the same arithmetic on the host; the census +// says the kernel ran on the device rather than the evaluator falling back, +// which the suite's oracle comparisons cannot tell apart. +TEST_CASE("shared gpu ops: views at non-zero offsets, counted by the census") { + if (!tl::gpu_available()) return; + auto prev = tl::device_; + tl::use_gpu(); + namespace gpu = tl::gpu; + auto dev = [](const array& a) { + array c = a.clone(); + c.eval(); + return c; + }; + auto same = [](const array& got, const array& want, float tol) { + return tl::allclose(got, want, tol, tol); + }; + const int64_t n = 1000, pa = 3, pb = 7, po = 11; // elements of slack in front + array A = dev(random_array({n + pa}, 940)), B = dev(random_array({n + pb}, 941)); + array a = A.slice(0, pa, n), b = B.slice(0, pb, n); + + SUBCASE("binary") { + array O = dev(array::zeros({n + po})); + gpu::census_reset(); + REQUIRE(gpu::binary(gpu::kop::add, a.device_span(), b.device_span(), + O.slice(0, po, n).device_span(), n, 2.0f, 1.0f)); + tl::gpu::flush(); + CHECK(gpu::census(gpu::kop::add) == 1); + tl::use_cpu(); + CHECK(same(O.slice(0, po, n), (a + b) * 2.0f + 1.0f, 1e-6f)); + CHECK(same(O.slice(0, 0, po), array::zeros({po}), 0.0f)); // slack untouched + } + + SUBCASE("unary and unary_ext") { + array O = dev(array::zeros({n + po})), P = dev(array::zeros({n + po})); + gpu::census_reset(); + REQUIRE(gpu::unary(gpu::kop::sigmoid, a.device_span(), + O.slice(0, po, n).device_span(), n, 1.0f, 0.0f)); + REQUIRE(gpu::unary_ext(gpu::unary_ext_op::tanh_, b.device_span(), + P.slice(0, po, n).device_span(), n, 3.0f, -1.0f)); + tl::gpu::flush(); + CHECK(gpu::census(gpu::kop::sigmoid) == 1); + CHECK(gpu::census(gpu::kop::tanh_) == 1); + tl::use_cpu(); + CHECK(same(O.slice(0, po, n), a.sigmoid(), 1e-6f)); + CHECK(same(P.slice(0, po, n), b.tanh() * 3.0f - 1.0f, 1e-5f)); + CHECK(same(P.slice(0, 0, po), array::zeros({po}), 0.0f)); + } + + SUBCASE("the evaluator's elementwise graph reaches the device") { + gpu::census_reset(); + array y = (a.clone() + b.clone()).exp(); + y.eval(); + CHECK(gpu::census(gpu::kop::add) + gpu::census(gpu::kop::exp_) >= 1); + } + + tl::device_ = prev; +} + +// The backend conformance test: every op a backend may implement, called on +// views at non-zero byte offsets, against a plain host loop. A buffer is staged +// with a run of sentinels before and after its values, so a kernel that drops a +// view's offset, or writes past its end, shows in the values or in the +// sentinels. It asks nothing of any one backend: an op the backend declines is +// skipped, and one it accepts has to be right. `must` names the ops every +// backend here is expected to take, so a regression that makes one decline +// fails rather than skips. +TEST_CASE("gpu ops on views at non-zero offsets") { + if (!tl::gpu_available()) return; + auto prev = tl::device_; + tl::use_gpu(); + namespace gpu = tl::gpu; + constexpr float kSentinel = -777.0f; + constexpr int64_t kPad = 5; // elements of slack on each side + + struct staged { + array buf; + int64_t n = 0; + gpu::span view() const { return buf.device_span().at(kPad * 4); } + array values() const { return buf.slice(0, kPad, n); } + bool sentinels_intact() const { + auto s = array::full({kPad}, kSentinel); + return tl::allclose(buf.slice(0, 0, kPad), s, 0.0f, 0.0f) && + tl::allclose(buf.slice(0, kPad + n, kPad), s, 0.0f, 0.0f); + } + }; + auto stage = [&](const std::vector& v) { + std::vector padded(v.size() + 2 * kPad, kSentinel); + std::copy(v.begin(), v.end(), padded.begin() + kPad); + staged s; + s.n = static_cast(v.size()); + s.buf = array::from(std::move(padded)).clone(); + s.buf.eval(); + return s; + }; + auto out_of = [&](int64_t n) { return stage(std::vector((size_t)n, kSentinel)); }; + auto rnd = [](int64_t n, unsigned seed, float lo = -1.0f, float hi = 1.0f) { + std::mt19937 rng(seed); + std::uniform_real_distribution dist(lo, hi); + std::vector v((size_t)n); + for (auto& x : v) x = dist(rng); + return v; + }; + // Whether `ran`: a declined op leaves its output alone and is not counted. + auto check = [&](bool ran, bool must, staged& o, const std::vector& want, + float tol) { + if (!ran) { + CHECK_FALSE(must); + return; + } + tl::gpu::flush(); + auto keep = tl::device_; + tl::use_cpu(); + CHECK(tl::allclose(o.values(), array::from(want), tol, tol)); + CHECK(o.sentinels_intact()); + tl::device_ = keep; + }; + + SUBCASE("elementwise, broadcast, compare, clamp, scalar") { + const int64_t m = 7, n = 33; + auto va = rnd(m * n, 1), vb = rnd(m * n, 2), vrow = rnd(n, 3); + staged a = stage(va), b = stage(vb), row = stage(vrow); + { + staged o = out_of(m * n); + std::vector want(m * n); + for (int64_t i = 0; i < m * n; i++) want[i] = (va[i] - vb[i]) * 3.0f + 0.5f; + gpu::census_reset(); + check(gpu::binary(gpu::kop::sub, a.view(), b.view(), o.view(), m * n, 3.0f, 0.5f), + true, o, want, 1e-6f); + CHECK(gpu::census(gpu::kop::sub) == 1); + } + { + staged o = out_of(m * n); + std::vector want(m * n); + for (int64_t r = 0; r < m; r++) + for (int64_t c = 0; c < n; c++) want[r * n + c] = va[r * n + c] * vrow[c] + 1.0f; + check(gpu::binary_bcast(gpu::kop::bmul, a.view(), n, 1, row.view(), 0, 1, o.view(), + m, n, 1.0f, 1.0f), + true, o, want, 1e-6f); + } + { + staged o = out_of(m * n); + std::vector want(m * n); + for (int64_t i = 0; i < m * n; i++) want[i] = va[i] > vb[i] ? 1.0f : 0.0f; + check(gpu::compare(gpu::cmp_op::gt, a.view(), b.view(), o.view(), m * n, 1), true, + o, want, 0.0f); + } + { + staged o = out_of(m * n); + std::vector want(m * n); + for (int64_t i = 0; i < m * n; i++) want[i] = std::min(std::max(va[i], -0.25f), 0.5f); + check(gpu::clamp(a.view(), o.view(), m * n, -0.25f, 0.5f), true, o, want, 0.0f); + } + { + staged o = out_of(m * n); + std::vector want(m * n); + for (int64_t i = 0; i < m * n; i++) want[i] = (va[i] < 0.1f ? 1.0f : 0.0f) * 2.0f; + check(gpu::scalar_binary(gpu::scalar_op::lt, a.view(), o.view(), m * n, 0.1f, 2.0f, + 0.0f), + true, o, want, 0.0f); + } + { + staged o = out_of(m * n); + std::vector want(m * n); + for (int64_t i = 0; i < m * n; i++) want[i] = std::cos(va[i]); + check(gpu::unary_ext(gpu::unary_ext_op::cos_, a.view(), o.view(), m * n, 1.0f, 0.0f), + true, o, want, 1e-6f); + } + } + + SUBCASE("row reductions, layer norm, rmsnorm, swiglu") { + const int64_t rows = 4, cols = 300; + auto vx = rnd(rows * cols, 11), vg = rnd(cols, 12), vb = rnd(cols, 13), + vd = rnd(rows * cols, 14); + staged x = stage(vx), g = stage(vg), b = stage(vb), d = stage(vd); + auto row_stat = [&](int64_t r, auto f) { + double acc = 0; + for (int64_t c = 0; c < cols; c++) acc += f(vx[r * cols + c]); + return acc; + }; + { + staged o = out_of(rows); + std::vector want(rows); + for (int64_t r = 0; r < rows; r++) + want[r] = (float)(row_stat(r, [](float v) { return (double)v; }) * 0.5 + 1.0); + check(gpu::row_op(gpu::kop::row_sum, x.view(), o.view(), rows, cols, 0.5f, 1.0f), + true, o, want, 1e-4f); + } + { + staged o = out_of(rows * cols); + std::vector want(rows * cols); + for (int64_t r = 0; r < rows; r++) { + const double z = row_stat(r, [](float v) { return std::exp((double)v); }); + for (int64_t c = 0; c < cols; c++) + want[r * cols + c] = (float)(std::exp((double)vx[r * cols + c]) / z); + } + check(gpu::row_op(gpu::kop::softmax, x.view(), o.view(), rows, cols, 1.0f, 0.0f), + true, o, want, 1e-6f); + } + { + staged o = out_of(rows); + std::vector want(rows); + for (int64_t r = 0; r < rows; r++) + want[r] = (float)std::log(row_stat(r, [](float v) { return std::exp((double)v); })); + check(gpu::row_logsumexp(x.view(), o.view(), rows, cols, 1.0f, 0.0f), false, o, want, + 1e-5f); + } + { + staged o = out_of(rows * cols); + std::vector want(rows * cols); + for (int64_t r = 0; r < rows; r++) { + const double mean = row_stat(r, [](float v) { return (double)v; }) / cols; + const double var = + row_stat(r, [&](float v) { return ((double)v - mean) * ((double)v - mean); }) / cols; + const double inv = 1.0 / std::sqrt(var + 1e-5); + for (int64_t c = 0; c < cols; c++) + want[r * cols + c] = (float)(((double)vx[r * cols + c] - mean) * inv * vg[c] + vb[c]); + } + check(gpu::layer_norm(x.view(), g.view(), b.view(), o.view(), rows, cols, 1e-5f, 1.0f, + 0.0f), + true, o, want, 1e-5f); + } + const bool model = gpu::caps::model_path; + { + staged xo = out_of(rows * cols), ho = out_of(rows * cols); + std::vector wx(rows * cols), wh(rows * cols); + for (int64_t r = 0; r < rows; r++) { + double ss = 0; + for (int64_t c = 0; c < cols; c++) { + wx[r * cols + c] = vx[r * cols + c] + vd[r * cols + c]; + ss += (double)wx[r * cols + c] * wx[r * cols + c]; + } + const double inv = 1.0 / std::sqrt(ss / cols + 1e-6); + for (int64_t c = 0; c < cols; c++) + wh[r * cols + c] = (float)(wx[r * cols + c] * inv * vg[c]); + } + const bool ran = gpu::rmsnorm_res(x.view(), d.view(), g.view(), xo.view(), ho.view(), + cols, 1e-6f, rows); + check(ran, model, xo, wx, 1e-6f); + check(ran, model, ho, wh, 1e-5f); + } + { + const int64_t ff = cols / 2; // x as [rows, 2*ff]: gate | up + staged o = out_of(rows * ff); + std::vector want(rows * ff); + for (int64_t r = 0; r < rows; r++) + for (int64_t f = 0; f < ff; f++) { + const double gate = vx[r * cols + f], up = vx[r * cols + ff + f]; + want[r * ff + f] = (float)(gate / (1.0 + std::exp(-gate)) * up); + } + check(gpu::swiglu(x.view(), o.view(), ff, rows), model, o, want, 1e-6f); + } + } + + SUBCASE("index family, gemm, pad: the ops backends run their own way") { + const int64_t table_rows = 9, row_size = 16, k = 6; + auto vt = rnd(table_rows * row_size, 21); + std::vector vi = {3, 0, 8, 3, 5, 1}; + staged t = stage(vt), idx = stage(vi); + { + staged o = out_of(k * row_size); + std::vector want(k * row_size); + for (int64_t i = 0; i < k; i++) + for (int64_t c = 0; c < row_size; c++) + want[i * row_size + c] = vt[(int64_t)vi[i] * row_size + c]; + check(gpu::index_select(t.view(), idx.view(), o.view(), row_size, k), true, o, want, + 0.0f); + } + { + auto vv = rnd(k * row_size, 22); + staged v = stage(vv), o = out_of(table_rows * row_size); + std::vector want(table_rows * row_size, 0.0f); + for (int64_t i = 0; i < k; i++) + for (int64_t c = 0; c < row_size; c++) + want[(int64_t)vi[i] * row_size + c] += vv[i * row_size + c]; + check(gpu::index_add(idx.view(), v.view(), o.view(), row_size, k, + table_rows * row_size), + true, o, want, 1e-6f); + } + { + const int64_t size = 9; // one-hot width + auto vv = rnd(k, 23); + staged v = stage(vv), o = out_of(k * size); + std::vector want(k * size, 0.0f); + for (int64_t i = 0; i < k; i++) want[i * size + (int64_t)vi[i]] = vv[i]; + check(gpu::scatter_to_axis(idx.view(), v.view(), o.view(), k, size), true, o, want, + 0.0f); + } + { + const int64_t m = 5, kk = 8, n = 12; + auto va = rnd(m * kk, 24), vb = rnd(kk * n, 25); + staged a = stage(va), b = stage(vb), o = out_of(m * n); + std::vector want(m * n); + for (int64_t i = 0; i < m; i++) + for (int64_t j = 0; j < n; j++) { + double acc = 0; + for (int64_t q = 0; q < kk; q++) acc += (double)va[i * kk + q] * vb[q * n + j]; + want[i * n + j] = (float)(acc * 2.0 - 1.0); + } + gpu::census_reset(); + check(gpu::gemm(a.view(), kk, false, b.view(), n, false, o.view(), m, n, kk, 2.0f, + -1.0f), + true, o, want, 1e-5f); + CHECK(gpu::ops_run() == 1); // a backend-own op is counted too + } + { + const int64_t a_shape[2] = {4, 6}, out_shape[2] = {4, 10}; + auto va = rnd(24, 26); + staged a = stage(va), o = out_of(40); + std::vector want(40, 0.0f); + for (int64_t r = 0; r < 4; r++) + for (int64_t c = 0; c < 6; c++) want[r * 10 + c + 3] = va[r * 6 + c]; + check(gpu::pad(a.view(), o.view(), a_shape, out_shape, 2, 1, 3, 24, 40), true, o, + want, 0.0f); + } + } + + tl::device_ = prev; +} + +// What the oracle comparisons cannot see: whether the evaluator's GPU mode +// reached the device at all, or fell back op by op and was right anyway. One +// graph a family every backend implements, and the census has to move. +TEST_CASE("the evaluator reaches the device in gpu mode") { + if (!tl::gpu_available()) return; + auto prev = tl::device_; + tl::use_gpu(); + namespace gpu = tl::gpu; + array a = random_array({6, 40}, 31), b = random_array({6, 40}, 32), + row = random_array({40}, 33), w = random_array({40, 8}, 34); + struct family { + const char* name; + std::function graph; + }; + const family families[] = { + {"elementwise", [&] { return (a + b).exp(); }}, + {"broadcast", [&] { return a * row; }}, + {"matmul", [&] { return a.dot(w); }}, + {"softmax", [&] { return a.softmax(); }}, + {"row sum", [&] { return a.sum(1); }}, + {"compare", [&] { return a > b; }}, + {"clamp", [&] { return a.clamp(-0.5f, 0.5f); }}, + {"pow scalar", [&] { return tl::pow(a * a + 1.0f, 0.5f); }}, + {"pad", [&] { return a.pad(1, 2, 3); }}, + }; + for (const family& f : families) { + CAPTURE(f.name); + array g = f.graph(); + gpu::census_reset(); + g.eval(); + CHECK(gpu::ops_run() >= 1); + } + tl::device_ = prev; +} diff --git a/tools/backend_parity_allowlist.txt b/tools/backend_parity_allowlist.txt deleted file mode 100644 index 014befa..0000000 --- a/tools/backend_parity_allowlist.txt +++ /dev/null @@ -1,68 +0,0 @@ -# Ops allowed to stay CUDA-only for now (checked by check_backend_parity.py). -# One name per line; '#' starts a comment. When an op here becomes REAL on -# every backend, the script says so -- remove its line at that point rather -# than let this list grow stale. -# -# The M9 LLM decode fast path: fused single-token attention, and the -# bf16/int4-quantized GEMVs it composes with for local-LLM inference -# throughput. Deliberately landed CUDA-first (see project memory -# project_cpp_tensorlib_handoff.md); Metal has them all now, and WebGPU still -# takes the CPU fallback for every one. -gemv_f32 -gemv_bf16 -gemv_q4 -attn_decode -attn_prefill - -# The model path (gpu.h): what a decoder runs on raw buffers between its -# GEMVs and attention. CUDA and Metal have all of it, including the two [N,K] -# weight layouts (the row bf16 GEMV and the GEMM the batched prefill is built -# on); WebGPU has none of it yet. -kv_append -kv_fill -argmax -rmsnorm -rmsnorm_res -swiglu -split_heads -merge_heads -gemv_bf16_row -gemm_bf16_nt - -# Causal prefill attention's pullback, the training counterpart of the row -# above: same CUDA-first reasoning, and it is reachable only from a backward -# pass that falls back to the composed form everywhere else. -attn_prefill_dq -attn_prefill_dkv - -# One-launch batched GEMM (the batch on gridDim.z next to split-K). CUDA-first -# for attention's probs·v; Metal/WebGPU loop gemm per slice meanwhile. -gemm_batched - -# GEMM with a fused row bias added in its store (addmm's shape). CUDA-first -# for the transformer's projections and FFN; Metal/WebGPU add the bias with -# their broadcast kernel after gemm meanwhile. -gemm_bias - -# N-D strided copy: clone()'s device arm for a permuted/transposed view. -# CUDA and Metal gather on the device; WebGPU still takes the host copy, a -# D2H and an H2D round trip, meanwhile. -copy_nd - -# Cross-entropy without a [rows, classes] intermediate: the trailing-axis -# gather (scatter_to_axis's dual), the one-pass row logsumexp the forward -# reduces with, and the pullback that rebuilds the softmax from it. CUDA-first; -# elsewhere the same values compose out of softmax/scatter/log meanwhile. -gather_from_axis -row_logsumexp -xent_bwd - -# Layer norm's pullback in three launches (dx with the row stats, the column -# partials, their fold). CUDA-first; elsewhere the backward composes the -# same gradients out of the unfused ops meanwhile. -layer_norm_bwd - -# Adam's per-parameter update in one launch. CUDA-first; elsewhere array.h's -# own host loop runs after the data comes home -- a flush and a memcpy on -# Metal's unified memory, a D2H and H2D round trip on WebGPU -- meanwhile. -adam_step diff --git a/tools/check_backend_parity.py b/tools/check_backend_parity.py deleted file mode 100755 index be273bb..0000000 --- a/tools/check_backend_parity.py +++ /dev/null @@ -1,276 +0,0 @@ -#!/usr/bin/env python3 -"""Backend parity checker: which GPU ops does each backend (CUDA/Metal/ -WebGPU) actually implement, vs. which are still a `return false;` stub. - -The op surface it checks is the exact "kernels" / "LLM path" / "model -path" lists in include/gpu.h's own doc comment -- that comment is the contract array.h/ -storage.h dispatch through, and every CUDA-first landing so far has already -had to update it, so it stays in sync without this script needing its own -copy of the list. - -For each op name, every backend header (cuda.h / metal.h / webgpu.h) -follows the same convention: the *first* `inline bool (...)` in the -file is that backend's real-platform implementation (Apple for metal.h, -TENSORLIB_CUDA for cuda.h, TENSORLIB_WEBGPU&&__EMSCRIPTEN__ for webgpu.h); -a disabled-platform build gets a second definition later in the file with -unnamed parameters, always `return false;`. So "is this backend's first -definition's body anything other than exactly `return false;`" is a -reliable REAL/STUB signal without needing a real C++ parser. - -That convention is also checked, not just assumed: an op defined ONLY -between the header's platform `#else` and its `#endif` is missing from the -platform that actually compiles the backend, and only that platform's build -says so. See misplaced_ops() below. - -Usage: tools/check_backend_parity.py -It also compares each op's PARAMETER LIST across backends. A stub that -drifts from the real signature compiles fine everywhere except the one -platform that dispatches through it -- exactly how an attn_decode stub -missing its kv_bf16 parameter reached the wasm job and nothing else. - -Exit status: 0 if every op is REAL or STUB the same way on all three -backends (a uniform stub is fine -- that just means nobody has ported it -yet) and their signatures agree; 1 if any op's status or signature differs -across backends, or if any definition sits only in a disabled-platform -block, printing what is missing where. -""" -import re -import sys -from pathlib import Path - -ROOT = Path(__file__).resolve().parent.parent -BACKENDS = { - "cuda": ROOT / "include/cuda.h", - "metal": ROOT / "include/metal.h", - "webgpu": ROOT / "include/webgpu.h", -} - - -def canonical_ops(): - """Pull the op-name list straight out of gpu.h's own contract comment.""" - text = (ROOT / "include/gpu.h").read_text() - m = re.search( - r"^//\s*kernels\s+(.*?)^// A backend with no kernel", - text, - re.MULTILINE | re.DOTALL, - ) - if not m: - sys.exit("check_backend_parity: couldn't find gpu.h's kernels/LLM " - "path comment -- has its wording changed?") - body = m.group(1) - # Drop the leading "//" (and any inherited indentation) from every - # wrapped comment line, and a row's label ("LLM path", "model path": a - # row is a label of one or two words, then its ops), before splitting on - # "/", so a label never reads as an identifier. - lines = [re.sub(r"^\s*//\s*", "", ln) for ln in body.splitlines()] - lines = [re.sub(r"^[A-Za-z]+(?: [a-z]+)?\s{2,}", "", ln) for ln in lines] - lines = [re.sub(r"^(LLM path|model path)\s+", "", ln) for ln in lines] - ops = [] - for ln in lines: - ops += [tok.strip() for tok in ln.split("/") if tok.strip()] - # Keep discovery order but drop duplicates. - seen = set() - ordered = [] - for op in ops: - if op not in seen: - seen.add(op) - ordered.append(op) - return ordered - - -def function_body(text, name): - """Return the first `inline bool (...) { ... }` body's source, or - None if `name` isn't defined as an `inline bool` function in `text`.""" - m = re.search(rf"^inline bool {re.escape(name)}\s*\(", text, re.MULTILINE) - if not m: - return None - i = text.index("(", m.end() - 1) - depth = 1 - i += 1 - while depth: - if text[i] == "(": - depth += 1 - elif text[i] == ")": - depth -= 1 - i += 1 - while text[i] != "{": - i += 1 - body_start = i + 1 - depth = 1 - i = body_start - while depth: - if text[i] == "{": - depth += 1 - elif text[i] == "}": - depth -= 1 - i += 1 - return text[body_start:i - 1] - - -def signature(text, name): - """The op's parameter list, normalized to compare across backends: types - only, since a stub names none of its parameters and a real one names all - of them, and without defaults, which each definition may state or not.""" - m = re.search(rf"^inline bool {re.escape(name)}\s*\(", text, re.MULTILINE) - if not m: - return None - i = text.index("(", m.end() - 1) - depth, start = 1, i + 1 - i += 1 - while depth: - if text[i] == "(": - depth += 1 - elif text[i] == ")": - depth -= 1 - i += 1 - # A stub writes types alone ("int64_t"), a real definition writes a name - # after them ("int64_t pos"), so a trailing token that is not itself part - # of a type is the name. - TYPE_WORDS = {"void", "bool", "char", "short", "int", "long", "float", - "double", "unsigned", "signed", "const", "int64_t", - "uint32_t", "uint16_t", "size_t", "kop", "cmp_op", - "scalar_op", "unary_ext_op", "dtype"} - params = [] - for p in re.split(r",(?![^<]*>)", text[start:i - 1]): - p = p.split("=", 1)[0] # drop a default argument - p = re.sub(r"\s+", " ", p).replace(" *", "* ").strip() - toks = p.split() - if len(toks) > 1 and toks[-1] not in TYPE_WORDS and "*" not in toks[-1]: - toks.pop() - p = " ".join(toks).replace("* ", "*").strip() - if p: - params.append(p) - return ", ".join(params) - - -def status(text, name): - body = function_body(text, name) - if body is None: - return "MISSING" - return "stub" if re.sub(r"\s+", " ", body).strip() == "return false;" else "REAL" - - -def platform_block(text): - """(else, endif) offsets of the header's top-level platform #if block.""" - depth = 0 - else_pos = None - for m in re.finditer(r"^#(if\w*|else|elif|endif)", text, re.MULTILINE): - kind = m.group(1) - if kind.startswith("if"): - depth += 1 - elif kind == "endif": - if depth == 1 and else_pos is not None: - return else_pos, m.start() - depth -= 1 - elif kind == "else" and depth == 1 and else_pos is None: - else_pos = m.start() - return None, None - - -def misplaced_ops(text, ops): - """Ops defined only between the platform #else and #endif. - - The real platform then has no such function at all, so its build breaks - on the first caller -- and nothing else does, which is why this is worth - a check rather than a convention. Ops both branches would define the - same way live after the #endif (gemm_bias, the LLM decode ops); those - are outside the block and fine. - """ - else_pos, endif_pos = platform_block(text) - if else_pos is None: - return [] - out = [] - for op in ops: - spots = [m.start() for m in re.finditer( - rf"^inline bool {re.escape(op)}\s*\(", text, re.MULTILINE)] - if spots and all(else_pos < s < endif_pos for s in spots): - out.append(op) - return out - - -ALLOWLIST_PATH = ROOT / "tools/backend_parity_allowlist.txt" - - -def allowlist(): - if not ALLOWLIST_PATH.exists(): - return set() - lines = ALLOWLIST_PATH.read_text().splitlines() - return {ln.split("#", 1)[0].strip() for ln in lines if ln.split("#", 1)[0].strip()} - - -def main(): - ops = canonical_ops() - sources = {b: p.read_text() for b, p in BACKENDS.items()} - rows = [] - for op in ops: - rows.append((op, {b: status(src, op) for b, src in sources.items()})) - sigs = {op: {b: signature(src, op) for b, src in sources.items()} - for op in ops} - - names = list(BACKENDS) - width = max(len(op) for op, _ in rows) + 2 - print(f"{'op':<{width}}" + "".join(f"{b:<10}" for b in names)) - mismatches = [] - for op, st in rows: - print(f"{op:<{width}}" + "".join(f"{st[b]:<10}" for b in names)) - if len(set(st.values())) > 1: - mismatches.append((op, st)) - - allowed = allowlist() - unallowed = [(op, st) for op, st in mismatches if op not in allowed] - stale = sorted(allowed - {op for op, _ in mismatches}) - - if not mismatches: - print(f"\nbackend-parity OK ({len(ops)} ops, all backends agree)") - else: - print(f"\nbackend-parity: {len(mismatches)}/{len(ops)} ops differ " - f"across backends:") - for op, st in mismatches: - real = [b for b in names if st[b] == "REAL"] - behind = [b for b in names if st[b] != "REAL"] - tag = " (allowlisted)" if op in allowed else " (NOT allowlisted)" - print(f" {op}: real on {real or '(none)'}, " - f"missing/stub on {behind}{tag}") - - if stale: - print(f"\n{ALLOWLIST_PATH.name} lists ops that are no longer " - f"asymmetric -- remove these lines:") - for op in stale: - print(f" {op}") - - # Signatures: every backend that defines the op must take the same - # parameters, or the one platform dispatching through the odd one out is - # the only build that fails. - drifted = [] - for op in ops: - seen = {b: sg for b, sg in sigs[op].items() if sg is not None} - if len(set(seen.values())) > 1: - drifted.append((op, seen)) - if drifted: - print(f"\nFAIL: {len(drifted)} op(s) whose parameter lists disagree " - f"across backends:") - for op, seen in drifted: - print(f" {op}:") - for b, sg in seen.items(): - print(f" {b:<8} ({sg})") - - misplaced = [(b, op) for b, src in sources.items() - for op in misplaced_ops(src, ops)] - if misplaced: - print(f"\nFAIL: {len(misplaced)} definition(s) sit only inside a " - f"disabled-platform #else block, so the platform that compiles " - f"that backend has no such function:") - for b, op in misplaced: - print(f" {b}: {op} -- add it to the real block too " - f"(a `return false;` stub is fine)") - - if unallowed: - print(f"\nFAIL: {len(unallowed)} unallowlisted mismatch(es). Either " - f"port the missing backend(s), or add the op to " - f"{ALLOWLIST_PATH.name} with a reason if the gap is " - f"deliberate.") - return 1 if (unallowed or misplaced or drifted) else 0 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/tools/cuda_trace/Dockerfile b/tools/cuda_trace/Dockerfile new file mode 100644 index 0000000..348c4f4 --- /dev/null +++ b/tools/cuda_trace/Dockerfile @@ -0,0 +1,6 @@ +# The host toolchain the launch trace is built with. No CUDA toolkit: the +# backend's host side needs only a C++ compiler, and the driver is the fake. +FROM debian:trixie-slim +RUN apt-get update -qq \ + && apt-get install -y -qq --no-install-recommends g++ python3 \ + && rm -rf /var/lib/apt/lists/* diff --git a/tools/cuda_trace/compare.sh b/tools/cuda_trace/compare.sh new file mode 100755 index 0000000..31fd872 --- /dev/null +++ b/tools/cuda_trace/compare.sh @@ -0,0 +1,47 @@ +#!/bin/bash +# Did a change to the backend's host side alter what reaches the driver? +# +# tools/cuda_trace/compare.sh [base-ref] (default: origin/master) +# +# Traces `base-ref` and the working tree with the same tooling (this +# directory's, so a base that predates it still traces) and diffs them. No +# output and exit 0 means every kernel still gets the same arguments in the +# same order. A diff is not a failure by itself — a fix is supposed to change +# the trace, and a new test appends to it — but every line of it should be one +# you meant. Removed lines are the ones to read first: a refactor that only +# moves code removes none. +set -euo pipefail + +here=$(cd "$(dirname "$0")" && pwd) +repo=$(cd "$here/../.." && pwd) +base=${1:-origin/master} + +work=$(mktemp -d "${TMPDIR:-/tmp}/cuda_trace.XXXXXX") +mkdir -p "$work/base_tree" "$work/base" "$work/head" +git -C "$repo" archive "$base" | tar -x -C "$work/base_tree" + +"$here/run.sh" "$work/base" "$work/base_tree" +"$here/run.sh" "$work/head" "$repo" + +status=0 +for t in "$work"/head/*.trace; do + name=$(basename "$t") + if [ ! -f "$work/base/$name" ]; then + echo "== $name: new in the working tree, nothing to compare" + continue + fi + # The closing line counts allocations, so it moves whenever anything is added. + if ! diff <(grep -av '^end ' "$work/base/$name") <(grep -av '^end ' "$t") \ + > "$work/$name.diff"; then + status=1 + echo "== $name: $(grep -c '^<' "$work/$name.diff") removed, $(grep -c '^>' "$work/$name.diff") added" + grep -a '^<' "$work/$name.diff" | head -10 || true + fi +done +if [ "$status" -eq 0 ]; then + echo "cuda_trace: identical to $base" + rm -rf "$work" +else + echo "cuda_trace: diffs kept in $work" +fi +exit "$status" diff --git a/tools/cuda_trace/fake_libcuda.cpp b/tools/cuda_trace/fake_libcuda.cpp new file mode 100644 index 0000000..fd1da24 --- /dev/null +++ b/tools/cuda_trace/fake_libcuda.cpp @@ -0,0 +1,196 @@ +// A stand-in libcuda.so.1 that records what the CUDA backend asks of the +// driver instead of doing it. cuda.h dlopens the driver and looks kernels up +// by name, so putting this library first on LD_LIBRARY_PATH is the whole +// installation: nothing in the backend knows it is being watched. +// +// "Device" memory is host memory, copies are memcpy, and a launch runs no +// kernel — it writes one trace line. What a refactor of the backend's host +// side must preserve is that the same kernels get the same arguments in the +// same order, and that is exactly what the trace holds; whether the kernels +// compute the right thing is the GPU suite's question, not this one's. +// +// Addresses never appear in the trace: a pointer prints as #N+off, the N-th +// allocation and a byte offset into it, so two runs compare with diff. +// +// TL_CUDA_TRACE names the output file (default: stderr). +#include +#include +#include +#include +#include +#include + +namespace { + +struct kernel { + const char* name; + const char* sig; // P pointer, u 4-byte integer, f float +}; +const kernel kKernels[] = { +#include "kernel_sigs.inc" +}; + +struct block { + size_t bytes; + int id; +}; +std::map g_live; // by base address +int g_next_id = 0; +long g_unresolved = 0; + +FILE* out() { + static FILE* f = [] { + const char* path = std::getenv("TL_CUDA_TRACE"); + FILE* file = path ? std::fopen(path, "w") : nullptr; + return file ? file : stderr; + }(); + return f; +} + +// "#N+off" for an address inside (or one past) a live allocation. +std::string where(uintptr_t p) { + if (!p) return "null"; + auto it = g_live.upper_bound(p); + if (it != g_live.begin()) { + --it; + if (p <= it->first + it->second.bytes) { + return "#" + std::to_string(it->second.id) + "+" + + std::to_string(p - it->first); + } + } + g_unresolved++; + return "UNRESOLVED"; +} + +struct at_exit { + ~at_exit() { + std::fprintf(out(), "end allocations=%d unresolved=%ld\n", g_next_id, + g_unresolved); + std::fflush(out()); + } +} g_at_exit; + +} // namespace + +extern "C" { + +using CUresult = int; +using CUdeviceptr = unsigned long long; + +CUresult cuInit(unsigned) { return 0; } +CUresult cuDeviceGetCount(int* n) { return *n = 1, 0; } +CUresult cuDeviceGet(int* dev, int) { return *dev = 0, 0; } +CUresult cuDevicePrimaryCtxRetain(void** ctx, int) { + static int the_context; + return *ctx = &the_context, 0; +} +CUresult cuCtxSetCurrent(void*) { return 0; } +CUresult cuCtxSynchronize() { return std::fprintf(out(), "sync\n"), 0; } + +CUresult cuModuleLoadData(void** mod, const void*) { + static int the_module; + return *mod = &the_module, 0; +} +CUresult cuModuleGetFunction(void** fn, void*, const char* name) { + for (const kernel& k : kKernels) { + if (!std::strcmp(k.name, name)) { + return *fn = const_cast(static_cast(&k)), 0; + } + } + std::fprintf(out(), "nofunc %s\n", name); + return *fn = nullptr, 500; // CUDA_ERROR_NOT_FOUND +} + +CUresult cuLaunchKernel(void* fn, unsigned gx, unsigned gy, unsigned gz, + unsigned bx, unsigned by, unsigned bz, unsigned smem, + void* /*stream*/, void** argv, void**) { + const auto* k = static_cast(fn); + std::fprintf(out(), "launch %s grid=%u,%u,%u block=%u,%u,%u smem=%u", k->name, + gx, gy, gz, bx, by, bz, smem); + for (size_t i = 0; k->sig[i]; i++) { + if (k->sig[i] == 'P') { + uintptr_t p; + std::memcpy(&p, argv[i], sizeof p); + std::fprintf(out(), " %s", where(p).c_str()); + } else if (k->sig[i] == 'u') { + uint32_t v; + std::memcpy(&v, argv[i], 4); + std::fprintf(out(), " u:%u", v); + } else { + float v; + std::memcpy(&v, argv[i], 4); + std::fprintf(out(), " f:%.9g", static_cast(v)); + } + } + std::fputc('\n', out()); + return 0; +} + +CUresult cuMemAlloc_v2(CUdeviceptr* p, size_t bytes) { + void* mem = std::calloc(bytes ? bytes : 1, 1); + if (!mem) return 2; // CUDA_ERROR_OUT_OF_MEMORY + g_live[reinterpret_cast(mem)] = {bytes, g_next_id}; + std::fprintf(out(), "alloc #%d %zu\n", g_next_id++, bytes); + return *p = reinterpret_cast(mem), 0; +} +CUresult cuMemFree_v2(CUdeviceptr p) { + auto it = g_live.find(static_cast(p)); + if (it == g_live.end()) return 1; + std::fprintf(out(), "free #%d\n", it->second.id); + g_live.erase(it); + std::free(reinterpret_cast(static_cast(p))); + return 0; +} +CUresult cuMemcpyHtoD_v2(CUdeviceptr dst, const void* src, size_t bytes) { + std::fprintf(out(), "h2d %s %zu\n", where(dst).c_str(), bytes); + std::memcpy(reinterpret_cast(static_cast(dst)), src, bytes); + return 0; +} +CUresult cuMemcpyHtoDAsync_v2(CUdeviceptr dst, const void* src, size_t bytes, + void*) { + std::fprintf(out(), "h2d_async %s %zu\n", where(dst).c_str(), bytes); + std::memcpy(reinterpret_cast(static_cast(dst)), src, bytes); + return 0; +} +CUresult cuMemcpyDtoH_v2(void* dst, CUdeviceptr src, size_t bytes) { + std::fprintf(out(), "d2h %s %zu\n", where(src).c_str(), bytes); + std::memcpy(dst, reinterpret_cast(static_cast(src)), bytes); + return 0; +} +CUresult cuMemsetD8_v2(CUdeviceptr dst, unsigned char v, size_t bytes) { + std::fprintf(out(), "memset %s %zu %u\n", where(dst).c_str(), bytes, v); + std::memset(reinterpret_cast(static_cast(dst)), v, bytes); + return 0; +} +CUresult cuMemsetD8Async(CUdeviceptr dst, unsigned char v, size_t bytes, void*) { + std::fprintf(out(), "memset_async %s %zu %u\n", where(dst).c_str(), bytes, v); + std::memset(reinterpret_cast(static_cast(dst)), v, bytes); + return 0; +} + +// Streams and graph capture: enough for the capture group to take its real +// path. A captured region's launches are traced as they are recorded; a replay +// is one line. +CUresult cuStreamCreate(void** s, unsigned) { + return *s = std::malloc(1), std::fprintf(out(), "stream_create\n"), 0; +} +CUresult cuStreamDestroy_v2(void* s) { return std::free(s), 0; } +CUresult cuStreamSynchronize(void*) { + return std::fprintf(out(), "stream_sync\n"), 0; +} +CUresult cuStreamBeginCapture_v2(void*, int mode) { + return std::fprintf(out(), "capture_begin mode=%d\n", mode), 0; +} +CUresult cuStreamEndCapture(void*, void** graph) { + return *graph = std::malloc(1), std::fprintf(out(), "capture_end\n"), 0; +} +CUresult cuGraphInstantiateWithFlags(void** exec, void*, unsigned long long) { + return *exec = std::malloc(1), std::fprintf(out(), "graph_instantiate\n"), 0; +} +CUresult cuGraphLaunch(void*, void*) { + return std::fprintf(out(), "graph_launch\n"), 0; +} +CUresult cuGraphExecDestroy(void* e) { return std::free(e), 0; } +CUresult cuGraphDestroy(void* g) { return std::free(g), 0; } + +} // extern "C" diff --git a/tools/cuda_trace/gen_kernel_sigs.py b/tools/cuda_trace/gen_kernel_sigs.py new file mode 100755 index 0000000..4d1a446 --- /dev/null +++ b/tools/cuda_trace/gen_kernel_sigs.py @@ -0,0 +1,82 @@ +#!/usr/bin/env python3 +"""Every __global__ kernel's argument layout, read from the .cu itself. + +cuLaunchKernel receives `void** argv` and nothing else, so the fake driver +cannot tell a pointer slot from a 4-byte one without knowing the kernel. This +runs the C preprocessor over kernels/tensorlib_cuda.cu (so macro-generated +kernels appear under their real names) and prints one table row a kernel: + + {"tl_add", "PPPuff"}, + +P = pointer, u = 4-byte integer, f = float. Anything else is an error: the +launch contract is "a pointer or a 4-byte scalar", and a kernel that breaks it +should fail here, loudly, not shift its arguments at run time. +""" +import re +import shutil +import subprocess +import sys + +INT4 = {"unsigned", "int", "unsigned int", "uint32_t", "int32_t"} + + +def preprocess(path): + src = open(path).read() + # The only #include is a CUDA header this host does not have; nothing in a + # signature needs it, and g++ stops preprocessing at a missing include. + src = re.sub(r"^\s*#\s*include[^\n]*$", "", src, flags=re.M) + cc = next((c for c in ("c++", "g++", "clang++") if shutil.which(c)), None) + if not cc: + sys.exit("gen_kernel_sigs: no C++ compiler to preprocess with") + out = subprocess.run([cc, "-E", "-P", "-x", "c++", "-"], input=src, + capture_output=True, text=True) + if out.returncode != 0: + sys.exit("gen_kernel_sigs: preprocessing failed:\n" + out.stderr) + return out.stdout + + +def params_of(text, open_paren): + depth = 0 + for j in range(open_paren, len(text)): + if text[j] == "(": + depth += 1 + elif text[j] == ")": + depth -= 1 + if depth == 0: + return text[open_paren + 1:j] + sys.exit("gen_kernel_sigs: unbalanced parameter list") + + +def layout(name, params): + sig = "" + for p in params.split(","): + p = " ".join(p.split()) + if not p: + continue + if "*" in p: + sig += "P" + continue + ty = " ".join(w for w in p.split()[:-1] if w != "const") + if ty in INT4: + sig += "u" + elif ty == "float": + sig += "f" + else: + sys.exit(f"gen_kernel_sigs: {name}: parameter `{p}` is neither a " + "pointer nor a 4-byte scalar") + return sig + + +def main(): + text = preprocess(sys.argv[1]) + rows = {} + for m in re.finditer(r"__global__\s+void\s+(\w+)\s*\(", text): + rows[m.group(1)] = layout(m.group(1), params_of(text, m.end() - 1)) + if not rows: + sys.exit("gen_kernel_sigs: no kernels found") + for name in sorted(rows): + print(f'{{"{name}", "{rows[name]}"}},') + + +if __name__ == "__main__": + main() diff --git a/tools/cuda_trace/in_container.sh b/tools/cuda_trace/in_container.sh new file mode 100755 index 0000000..ca908e5 --- /dev/null +++ b/tools/cuda_trace/in_container.sh @@ -0,0 +1,92 @@ +#!/bin/bash +# Runs inside the container (see run.sh): /tools is this directory, /tree the +# tensorlib checkout being traced, /out where the traces go. +set -euo pipefail + +build=/tmp/build +mkdir -p "$build" + +# The kernel table comes from the traced tree's own .cu, so a kernel added or +# re-ordered there is known here without anyone editing a list. +python3 /tools/gen_kernel_sigs.py /tree/kernels/tensorlib_cuda.cu > "$build/kernel_sigs.inc" +g++ -std=c++17 -O1 -shared -fPIC -I"$build" /tools/fake_libcuda.cpp -o "$build/libcuda.so.1" + +# The PTX is nvcc's output and the fake driver never reads it: one zero byte +# stands in for the generated byte list. +echo 0x00 > "$build/tensorlib_cuda_ptx.inc" + +flags=(-std=c++23 -O0 -w -DTENSORLIB_CUDA -I/tree/include -I"$build") + +# TL_CUDA_TRACE_CHECK=1: only ask whether everything still compiles against the +# CUDA branch of the headers (a minute, against the full trace's two). +if [ "${TL_CUDA_TRACE_CHECK:-0}" = 1 ]; then + pids=() + for src in /tree/test/test_array.cpp /tree/bench/cuda/check/*.cpp \ + /tree/bench/cuda/speed/*.cpp /tree/bench/models/*.cpp; do + grep -q '#include "/out/$name.log" 2>&1 || status=$? + if [ "$status" -ge 128 ]; then + echo "cuda_trace: $name died on signal $((status - 128)); see $name.log" >&2 + exit 1 + fi + if grep -aq 'UNRESOLVED\|^nofunc' "/out/$name.trace"; then + echo "cuda_trace: $name passed the driver a pointer outside every allocation, or asked for an unknown kernel" >&2 + grep -an 'UNRESOLVED\|^nofunc' "/out/$name.trace" | head -5 >&2 + exit 1 + fi +} +trace tensorlib_test.gpu --gpu +trace tensorlib_test.auto --auto +trace check_cuda +trace check_llm_decode +trace check_attn64 +if [ -f "$sweep" ]; then trace trace_sweep; fi + +# Which kernels the traces reach, against every kernel the .cu defines. +grep -aho '^launch [a-z0-9_]*' /out/*.trace | sort -u | cut -d' ' -f2 > /out/kernels_reached.txt +grep -o '"tl_[a-z0-9_]*"' "$build/kernel_sigs.inc" | cut -d'"' -f2 | sort -u \ + | grep -vxFf /out/kernels_reached.txt > /out/kernels_unreached.txt || true +echo "cuda_trace: $(wc -l < /out/kernels_reached.txt) kernels reached, $(wc -l < /out/kernels_unreached.txt) not" diff --git a/tools/cuda_trace/run.sh b/tools/cuda_trace/run.sh new file mode 100755 index 0000000..a7b8de7 --- /dev/null +++ b/tools/cuda_trace/run.sh @@ -0,0 +1,23 @@ +#!/bin/bash +# Trace what a tensorlib tree's CUDA backend asks of the driver, with no GPU: +# +# tools/cuda_trace/run.sh [tree] +# +# `tree` defaults to the checkout this script lives in. The test suite and the +# CUDA checkers are built on Linux in a container against a recording driver +# (fake_libcuda.cpp), and each writes /.trace. +# TL_CUDA_TRACE_CHECK=1 stops after compiling: does the CUDA branch still build? +set -euo pipefail + +here=$(cd "$(dirname "$0")" && pwd) +out=${1:?usage: run.sh [tree]} +tree=$(cd "${2:-$here/../..}" && pwd) +mkdir -p "$out" +out=$(cd "$out" && pwd) + +image=tensorlib-cuda-trace +if ! docker image inspect "$image" > /dev/null 2>&1; then + docker build -q -t "$image" "$here" > /dev/null +fi +docker run --rm -e TL_CUDA_TRACE_CHECK="${TL_CUDA_TRACE_CHECK:-0}" -v "$here":/tools:ro -v "$tree":/tree:ro -v "$out":/out \ + "$image" bash /tools/in_container.sh