diff --git a/meson.build b/meson.build index 693584c..70373f5 100644 --- a/meson.build +++ b/meson.build @@ -203,11 +203,18 @@ executable('ssm_timing_test', sources : [tests/'ssm_timing_test.cpp'] + engine_sources_wyut, kwargs : test_engine_args) -# M1.7a: WY/UT parallel scan vs decode oracle (SYCL). +# M1.7a: WY/UT parallel scan vs decode oracle (SYCL). Uses the wyut source +# set (calls gpu_deltanet_wyut_b by name). executable('deltanet_wyut_test', sources : [tests/'deltanet_wyut_test.cpp'] + engine_sources_wyut, kwargs : test_engine_args) +# M1.7b.12: naive register-recurrence vs decode oracle (SYCL). Uses the +# naive source set (gpu_deltanet_batched dispatches to the naive kernel). +executable('deltanet_naive_test', + sources : [tests/'deltanet_naive_test.cpp'] + engine_sources_naive, + kwargs : test_engine_args) + # gemm_tf32 batched GEMM vs CPU fp32 ref (SYCL). executable('gemm_tf32_test', sources : [tests/'gemm_tf32_test.cpp', src/'qxmx_gemm_tf32.cpp'], diff --git a/src/qxmx_deltanet_naive.cpp b/src/qxmx_deltanet_naive.cpp index 6f49599..b52d0bf 100644 --- a/src/qxmx_deltanet_naive.cpp +++ b/src/qxmx_deltanet_naive.cpp @@ -1,26 +1,36 @@ -/* qxmx_deltanet_naive: M1.7b.12 naive per-token DeltaNet recurrence (the - * llama.cpp algorithm — one warp per state column, state shard in registers, - * warp_reduce across rows; no matmuls, no SLM, no UT solve). See - * docs/memory qxm_m1_7b_11_llama_deltanet_cmp.md for the algorithm + the - * measured win over WY/UT at S_v=256. +/* qxmx_deltanet_naive: M1.7b.12 naive per-token DeltaNet recurrence. * - * STATUS: STUB. The meson `naive` engine source set compiles this so the - * target links and the build wiring is exercised, but gpu_deltanet_batched - * here is NOT yet a correct recurrence — it aborts. Fill in the naive kernel - * (port the warp-per-column s_shard[rows_per_lane] form from - * llama.cpp/ggml/src/ggml-sycl/gated_delta_net.cpp) to make this live. */ + * Port of llama.cpp's gated_delta_net algorithm (ggml/src/ggml-sycl/ + * gated_delta_net.cpp): the state column lives in REGISTERS (one private + * array per thread), and the per-token recurrence is a scalar loop — no + * matmuls, no SLM, no WY/UT decomposition. At S_v=HD=128 this beats the + * WY/UT chunked form (which pays matmul-launch + a T-sequential UT solve) + * and scales with ubatch. + * + * Layout (matches our M1.5 gpu_deltanet_rmsnorm_b math exactly, but with the + * HD*HD state column moved from SLM to a per-thread private array): + * - One WG per (v-head hv) = NVH=48 groups; LWS = HD = 128 threads. + * - Each thread owns one column `col` of S, holding all HD rows in a + * private float Sd_col[HD] (128 floats/thread). + * - q,k are L2-normalized per KEY head (hk = hv % NKH) via a group reduce. + * - Per token: decay S, accumulate kv_mem (per-thread, no reduce — each + * thread owns its full column), delta = (v-kv_mem)*beta_s, rank-1 update, + * accumulate output, then folded group_rmsnorm over the HD output values. + * - State loaded from global once at start, written once at end. + * + * Bit-exact vs gpu_deltanet_rmsnorm (decode oracle) per (token, head, col): + * same fp ops in the same order. See deltanet_wyut_test for the gate. + * + * Strides: pass 0 to use natural contiguous strides; ssm_forward_b passes + * (CONVD, CONVD, VALD) since proj_buf is fused [n_tok][CONVD]. */ #include "qxmx_deltanet.h" #include "qxmx.h" -#include +#include #include -#include #include namespace qx { -/* TODO M1.7b.12: naive register-recurrence. Until implemented, abort so no - * caller silently gets wrong output. The wyut engine source set provides the - * real recurrence for production. */ void gpu_deltanet_batched(sycl::queue& q, const float* d_qk, const float* d_v, const float* d_alpha, const float* d_beta, @@ -29,12 +39,84 @@ void gpu_deltanet_batched(sycl::queue& q, float* d_state, float* d_out, int n_tok, size_t qk_stride, size_t v_stride, size_t out_stride, wyut_workspace* ws) { - (void)q; (void)d_qk; (void)d_v; (void)d_alpha; (void)d_beta; - (void)d_a; (void)d_dtb; (void)d_norm_w; (void)d_state; (void)d_out; - (void)n_tok; (void)qk_stride; (void)v_stride; (void)out_stride; (void)ws; - std::fprintf(stderr, "gpu_deltanet_batched (naive): not yet implemented -- " - "build with the wyut engine source set\n"); - assert(0 && "naive DeltaNet not implemented"); + (void)ws; /* naive uses no workspace */ + constexpr int HD = QX_SSM_STATE, NKH = QX_SSM_GROUPS, NVH = QX_SSM_DT_RANK; + constexpr int LWS = HD; + const float eps = QX_RMS_EPS; + const float scale_qk = 1.0f / sqrtf((float)HD); + const size_t qk_s = qk_stride ? qk_stride : (size_t)2 * NKH * HD; + const size_t v_s = v_stride ? v_stride : (size_t)NVH * HD; + const size_t out_s = out_stride? out_stride: v_s; + + q.submit([&](sycl::handler& h) { + /* ql/kl: L2-normalized q,k for this v-head's key head, shared across + * all HD columns via SLM (each thread reads its col). pr holds the + * (decay, beta_s) pair computed by lane 0. lss is the rmsnorm sumsq + * reduction buffer. */ + sycl::local_accessor ql(sycl::range<1>(HD),h), kl(sycl::range<1>(HD),h); + sycl::local_accessor pr(sycl::range<1>(2),h); + sycl::local_accessor lss(sycl::range<1>(LWS),h); + h.parallel_for(sycl::nd_range<1>(NVH*LWS, LWS), [=](sycl::nd_item<1> it){ + int hv=it.get_group(0), hk=hv%NKH, col=it.get_local_id(0); + float* S = d_state + (size_t)hv*HD*HD; + /* State column in registers/private memory — each thread owns all + * HD rows of its column. Loaded from global once, written once. */ + float Sd_col[HD]; + #pragma unroll + for(int row=0;row()); + double ks=sycl::reduce_over_group(it.get_group(),kk,sycl::plus()); + ql[col]*=1.0f/sycl::sqrt((float)(qs+eps)); kl[col]*=1.0f/sycl::sqrt((float)(ks+eps)); + it.barrier(sycl::access::fence_space::local_space); + if(col==0){float apb=d_alpha[(size_t)t*NVH+hv]+d_dtb[hv]; + float sp=apb>20?apb:(apb<-20?sycl::exp(apb):sycl::log1p(sycl::exp(apb))); + pr[0]=sycl::exp(d_a[hv]*sp); pr[1]=1.0f/(1.0f+sycl::exp(-d_beta[(size_t)t*NVH+hv]));} + it.barrier(sycl::access::fence_space::local_space); + float decay=pr[0], beta_s=pr[1], v_reg=v_t[hv*HD+col]; + /* Pass 1: decay S in registers, accumulate kv_mem (per-thread; + * no reduce — each thread owns its full column). */ + float kv_mem=0; + #pragma unroll + for(int row=0;row 0; s >>= 1) { + if (col < s) lss[col] += lss[col + s]; + it.barrier(sycl::access::fence_space::local_space); + } + float inv = 1.0f / sycl::sqrt((float)(lss[0] / HD) + QX_RMS_EPS); + d_out[(size_t)t * out_s + hv*HD+col] = out_val * inv * d_norm_w[col]; + it.barrier(sycl::access::fence_space::local_space); + } + /* Write final S to global once. */ + #pragma unroll + for(int row=0;row +#include +#include +#include +#include +#include + +#include + +#include "qxmx.h" +#include "qxmx_deltanet.h" + +using namespace qx; + +static std::mt19937_64 rng(0xD7A5ULL); + +static constexpr int HD = QX_SSM_STATE; /* 128 */ +static constexpr int NKH = QX_SSM_GROUPS; /* 16 */ +static constexpr int NVH = QX_SSM_DT_RANK; /* 48 */ +static constexpr int QK_STRIDE = 2 * NKH * HD; /* 4096 per token */ +static constexpr int V_STRIDE = NVH * HD; /* 6144 per token */ +static constexpr int S_SIZE = NVH * HD * HD; + +static void fill_gauss(float* p, size_t n, float scale, std::mt19937_64& r) { + std::normal_distribution g(0.0f, 1.0f); + for (size_t i = 0; i < n; i++) p[i] = g(r) * scale; +} + +static int cmp_f32(const char* what, const float* a, const float* b, size_t n, + double& max_abs, double& max_rel) { + int bad = 0; + for (size_t i = 0; i < n && bad < 8; i++) { + double df = fabs((double)a[i] - (double)b[i]); + if (df > max_abs) max_abs = df; + double rel = df / (fabs((double)b[i]) > 1e-3 ? fabs((double)b[i]) : 1e-3); + if (rel > max_rel) max_rel = rel; + if (df > 1e-3 && bad < 8) { + printf(" %s MISMATCH @%zu: %a vs %a (diff %g)\n", what, i, a[i], b[i], a[i] - b[i]); + bad++; + } + } + return bad; +} + +static int run_config(sycl::queue& q, const char* label, int n_tok, + bool zero_s0, int decay_mode) { + std::vector qk((size_t)n_tok * QK_STRIDE); + std::vector v((size_t)n_tok * V_STRIDE); + std::vector alpha((size_t)n_tok * NVH); + std::vector beta((size_t)n_tok * NVH); + std::vector a_log(NVH), dtb(NVH), norm_w(HD); + std::vector s0(S_SIZE); + + fill_gauss(qk.data(), qk.size(), 1.0f, rng); + fill_gauss(v.data(), v.size(), 1.0f, rng); + fill_gauss(alpha.data(), alpha.size(), 2.0f, rng); + fill_gauss(beta.data(), beta.size(), 2.0f, rng); + fill_gauss(dtb.data(), dtb.size(), 1.0f, rng); + for (int i = 0; i < HD; i++) norm_w[i] = 1.0f + 0.1f * (float)(rng() % 100) / 100.0f; + + for (int h = 0; h < NVH; h++) { + if (decay_mode == 0) a_log[h] = -0.01f * (float)(h % 3); + else if (decay_mode == 1) a_log[h] = -1.0f - 0.1f * (float)(h % 5); + else a_log[h] = (h % 2 == 0) ? -0.01f : -1.5f; + } + + if (zero_s0) std::memset(s0.data(), 0, S_SIZE * sizeof(float)); + else fill_gauss(s0.data(), S_SIZE, 0.5f, rng); + + auto H2D = [&](const std::vector& vv) -> float* { + float* d = sycl::malloc_shared(vv.size(), q); + memcpy(d, vv.data(), vv.size() * sizeof(float)); + return d; + }; + float* d_qk = H2D(qk); + float* d_v = H2D(v); + float* d_alpha = H2D(alpha); + float* d_beta = H2D(beta); + float* d_a = H2D(a_log); + float* d_dtb = H2D(dtb); + float* d_nw = H2D(norm_w); + + float* d_S_oracle = H2D(s0); + float* d_S_wyut = H2D(s0); + float* d_out_oracle = sycl::malloc_shared((size_t)n_tok * V_STRIDE, q); + float* d_out_wyut = sycl::malloc_shared((size_t)n_tok * V_STRIDE, q); + memset(d_out_oracle, 0xEE, (size_t)n_tok * V_STRIDE * 4); + memset(d_out_wyut, 0xEE, (size_t)n_tok * V_STRIDE * 4); + + /* oracle: per-token decode calls */ + for (int t = 0; t < n_tok; t++) { + gpu_deltanet_rmsnorm(q, + d_qk + (size_t)t * QK_STRIDE, + d_v + (size_t)t * V_STRIDE, + d_alpha + (size_t)t * NVH, + d_beta + (size_t)t * NVH, + d_a, d_dtb, d_nw, + d_S_oracle, d_out_oracle + (size_t)t * V_STRIDE); + } + q.wait(); + + /* Batched (meson-selected: naive or wyut) */ + gpu_deltanet_batched(q, d_qk, d_v, d_alpha, d_beta, d_a, d_dtb, d_nw, + d_S_wyut, d_out_wyut, n_tok); + q.wait(); + + double max_abs = 0, max_rel = 0; + int bad = 0; + bad += cmp_f32("out", d_out_wyut, d_out_oracle, (size_t)n_tok * V_STRIDE, max_abs, max_rel); + bad += cmp_f32("state", d_S_wyut, d_S_oracle, S_SIZE, max_abs, max_rel); + + printf(" %-16s n_tok=%-4d s0=%-6s decay=%-5s : abs %.3g (gate 0.05) rel %.3g (gate 0.5) %s\n", + label, n_tok, zero_s0 ? "zero" : "rand", + decay_mode == 0 ? "slow" : decay_mode == 1 ? "fast" : "mixed", + max_abs, max_rel, + (max_abs <= 0.05 && max_rel <= 0.5) ? "PASSED" : "FAILED"); + + sycl::free(d_qk, q); sycl::free(d_v, q); + sycl::free(d_alpha, q); sycl::free(d_beta, q); + sycl::free(d_a, q); sycl::free(d_dtb, q); sycl::free(d_nw, q); + sycl::free(d_S_oracle, q); sycl::free(d_S_wyut, q); + sycl::free(d_out_oracle, q); sycl::free(d_out_wyut, q); + return (max_abs <= 0.05 && max_rel <= 0.5) ? 0 : 1; +} + +int main() { + sycl::queue q(sycl::gpu_selector_v, sycl::property::queue::in_order{}); + printf("deltanet_naive_test (M1.7b.12): gpu_deltanet_batched vs decode oracle\n"); + int total = 0; + /* The naive recurrence has no GEMM tile constraint, so any n_tok is + * valid. Cover the chunk sizes the engine uses (CHUNK=256, mid-chunks + * that are multiples of 16) plus a few ragged sizes the WY/UT test skips. */ + total += run_config(q, "chunk-1", 1, true, 0); + total += run_config(q, "chunk-13", 13, false, 1); + total += run_config(q, "chunk-32", 32, true, 0); + total += run_config(q, "chunk-64", 64, false, 1); + total += run_config(q, "chunk-100", 100, true, 2); + total += run_config(q, "chunk-128", 128, false, 2); + total += run_config(q, "chunk-256", 256, true, 0); + total += run_config(q, "chunk-512", 512, false, 1); + printf(total ? "\nFAILED: %d config(s)\n" + : "\nALL PASSED — bit-exact vs decode oracle\n", total); + return total ? 1 : 0; +} \ No newline at end of file