diff --git a/src/qxmx_gpu.cpp b/src/qxmx_gpu.cpp index 571e4bd..7300f2b 100644 --- a/src/qxmx_gpu.cpp +++ b/src/qxmx_gpu.cpp @@ -468,19 +468,49 @@ void gpu_attn_split_qgate_b(sycl::queue& q, const float* proj, }); } -/* Online-softmax attention (FlashAttention-style). One work-group per query - * head (HD=256 threads = head_dim). Each step does a cooperative HD-dim dot - * product (group reduction), then updates the running max/sum/accumulator. - * No scores buffer needed — everything in registers/SLM. Writes attn_out - * [NQ*HD]. Reads q [NQ*HD], kv_k [per K-cache format], kv_v [Phase 1 V]. +/* M4 FlashDecoding: max KV splits for split-KV decode attention. Partial + * acc [S][NQ][HD] + (m,l) [S][NQ][2] live in engine-hoisted USM. */ +constexpr int FD_MAX_SPLIT = 64; + +/* Cross-split merge: per (head, dim), + * out = sum_s exp(m_s-M)*acc_s / sum_s exp(m_s-M)*l_s, M = max_s m_s. */ +static void attn_fd_merge(sycl::queue& q, const float* fd_acc, + const float* fd_ml, float* out, int S) { + q.submit([&](sycl::handler& h) { + h.parallel_for( + sycl::nd_range<1>((size_t)QX_N_HEAD * QX_HEAD_DIM, QX_HEAD_DIM), + [=](sycl::nd_item<1> it) { + int hh = it.get_group(0); + int lid = it.get_local_id(0); + float M = -1e30f; + for (int s = 0; s < S; s++) + M = sycl::fmax(M, fd_ml[((size_t)s * QX_N_HEAD + hh) * 2]); + float num = 0.0f, den = 0.0f; + for (int s = 0; s < S; s++) { + float w = sycl::exp(fd_ml[((size_t)s * QX_N_HEAD + hh) * 2] - M); + num += w * fd_acc[((size_t)s * QX_N_HEAD + hh) * QX_HEAD_DIM + lid]; + den += w * fd_ml[((size_t)s * QX_N_HEAD + hh) * 2 + 1]; + } + out[(size_t)hh * QX_HEAD_DIM + lid] = num / den; + }); + }); +} + +/* Online-softmax attention (FlashAttention-style). One work-group per + * (query head, KV split): HD=256 threads = head_dim. Each step does a + * cooperative HD-dim dot product (group reduction), then updates the running + * max/sum/accumulator. No scores buffer needed — everything in registers/ + * SLM. Writes attn_out [NQ*HD] (S==1) or partials for attn_fd_merge (S>1). + * Reads q [NQ*HD], kv_k [per K-cache format], kv_v [Phase 1 V]. * GQA: hk=hh/GQA. Templated on K cache format (fp16 | q8_0 | fp8); the K * dequant is selected at compile time via if constexpr — no runtime branch * in the hot per-token loop. V dequant (TurboQuant WHT+4-bit) is identical * across K types. */ template void gpu_attn_softmax(sycl::queue& q, const float* q_vec, - const uint8_t* kv_k, const uint32_t* kv_v, - float* out, int64_t pos) { + const uint8_t* kv_k, const uint32_t* kv_v, + float* out, int64_t pos, + float* fd_acc, float* fd_ml, int fd_max_split) { constexpr int HD = QX_HEAD_DIM; constexpr int NQ = QX_N_HEAD; constexpr int GQA = NQ / QX_N_HEAD_KV; @@ -491,25 +521,44 @@ void gpu_attn_softmax(sycl::queue& q, const float* q_vec, * identical for both. */ constexpr int BLK_PER_HEAD = HD / VBLK; constexpr float NRM = 0.17677669529663687f; /* 1/sqrt(32) for inverse WHT */ + /* FlashDecoding split count: one WG per (head, split) of ~QXMX_FD_CHUNK + * tokens turns 24 serial WG scans into NQ*S parallel ones. S==1 keeps + * the exact single-pass oracle path (shallow ctx, no scratch, QXMX_FD=0). */ + static const int fd_chunk = [] { + const char* e = std::getenv("QXMX_FD_CHUNK"); + return e ? std::atoi(e) : 256; + }(); + static const bool fd_disabled = [] { + const char* e = std::getenv("QXMX_FD"); + return e && e[0] == '0'; + }(); + const int64_t n_kv = pos + 1; + if (fd_max_split < 1) fd_max_split = 1; + int S = 1; + if (fd_acc && fd_ml && !fd_disabled && fd_chunk > 0) { + int64_t s64 = (n_kv + fd_chunk - 1) / fd_chunk; + S = (int)(s64 < fd_max_split ? s64 : (int64_t)fd_max_split); + } + const int64_t fd_span = (n_kv + S - 1) / S; q.submit([&](sycl::handler& h) { sycl::local_accessor ql(sycl::range<1>(HD), h); - /* V dequant staging: this head's 8 blocks × 32 elems = 256 floats. - * Each block's 32 lanes occupy contiguous SLM (lid = blk_local*32+lane), - * so the inverse-WHT butterfly can run per-block with no cross-block - * interference. A second buffer is used for the double-buffered WHT. */ - sycl::local_accessor vs(sycl::range<1>(HD), h); - sycl::local_accessor vs2(sycl::range<1>(HD), h); - h.parallel_for(sycl::nd_range<1>(NQ * HD, HD), [=](sycl::nd_item<1> it) { - int hh = it.get_group(0); + h.parallel_for(sycl::nd_range<1>((size_t)NQ * S * HD, HD), [=](sycl::nd_item<1> it) + [[sycl::reqd_sub_group_size(VBLK)]] { + int g = it.get_group(0); + int hh = g % NQ; + int sp = g / NQ; int hk = hh / GQA; int lid = it.get_local_id(0); int blk_local = lid >> 5; /* 0..7 within the head */ int lane = lid & 31; /* 0..31 within the block */ int blk_row = hk * BLK_PER_HEAD + blk_local; + auto sg = it.get_sub_group(); ql[lid] = q_vec[(size_t)hh * HD + lid]; sycl::group_barrier(it.get_group()); float acc = 0.0f, max_s = -1e30f, sum_s = 0.0f; - for (int64_t t = 0; t <= pos; t++) { + const int64_t t0 = (int64_t)sp * fd_span; + const int64_t t1 = (t0 + fd_span < n_kv) ? t0 + fd_span : n_kv; + for (int64_t t = t0; t < t1; t++) { /* K dequant — compile-time selected per Kct. Thread lid reads * its own K element (hk*HD + lid) from the per-format layout. */ float k_val; @@ -549,32 +598,34 @@ void gpu_attn_softmax(sycl::queue& q, const float* q_vec, int off = (lane & 7) * 4; uint32_t word = (w == 0) ? c0 : (w == 1) ? c1 : (w == 2) ? c2 : c3; uint32_t code = (word >> off) & 0xFu; - vs[lid] = V_LM_LUT[code] * vscale; - sycl::group_barrier(it.get_group()); - /* Inverse WHT: 5 butterfly stages within each 32-lane block. - * Pairs are (lid, lid^bit); only lanes within the same block - * pair (blk_local matches because bit < 32). Double-buffer. */ - auto* cur = &vs; - auto* nxt = &vs2; + float v_val = V_LM_LUT[code] * vscale; + /* Inverse WHT via sub-group register shuffles — no SLM, no + * barriers. reqd_sub_group_size(VBLK=32) makes the SG span + * exactly one 32-elem V block; xor by bit in {1..16} pairs + * lanes within the block. High lane (lane&bit) emits partner + * - v, low emits v + partner (self-inverse butterfly). */ for (int bit = 1; bit < VBLK; bit <<= 1) { - float a = (*cur)[lid]; - float b = (*cur)[lid ^ bit]; - sycl::group_barrier(it.get_group()); - /* High lane (lane&bit) outputs b-a; low outputs a+b. - * Self-inverse butterfly: fwht(fwht(x)) = 32*x. */ - (*nxt)[lid] = (lane & bit) ? (b - a) : (a + b); - sycl::group_barrier(it.get_group()); - auto* tmp = cur; cur = nxt; nxt = tmp; + float partner = sycl::permute_group_by_xor(sg, v_val, bit); + v_val = (lane & bit) ? (partner - v_val) : (v_val + partner); } /* Apply 1/sqrt(32) norm + the fp16 norm-correction factor to * restore the block's L2 norm (quantization shrinks it). */ - float v_val = (*cur)[lid] * NRM * vnorm_corr; + v_val = v_val * NRM * vnorm_corr; acc = f * acc + e * v_val; max_s = m_new; } - out[(size_t)hh * HD + lid] = acc / sum_s; + if (S == 1) { + out[(size_t)hh * HD + lid] = acc / sum_s; + return; + } + fd_acc[((size_t)sp * NQ + hh) * HD + lid] = acc; + if (lid == 0) { + fd_ml[((size_t)sp * NQ + hh) * 2] = max_s; + fd_ml[((size_t)sp * NQ + hh) * 2 + 1] = sum_s; + } }); }); + if (S > 1) attn_fd_merge(q, fd_acc, fd_ml, out, S); } @@ -745,8 +796,9 @@ void attn_forward(sycl::queue& q, const dev_attn_blk* a, engine::KCT k_ctk, float* part, float* proj_buf, float* qvec, float* gate, float* kvec, float* vvec, - float* attn_out_buf, - const float* rope_c, const float* rope_s) { + float* attn_out_buf, + const float* rope_c, const float* rope_s, + float* fd_acc, float* fd_ml, int fd_max_split) { /* One rmsnorm_quant; q/k/v all share the packed activation. */ gpu_rmsnorm_quant(q, x, a->attn_norm, QX_D_MODEL, g_bvnni_d, g_suma_d, g_act_scale_d); @@ -773,9 +825,9 @@ void attn_forward(sycl::queue& q, const dev_attn_blk* a, /* Phase 1: V cache stored as packed 4-bit TurboQuant blocks. */ gpu_quant_v_row(q, vvec, kv_v + (size_t)pos*VROW_U32); /* Attn dequant — compile-time-selected template per K format. */ - if (k_ctk == engine::KCT::fp16) gpu_attn_softmax(q, qvec, kv_k, kv_v, attn_out_buf, pos); - else if (k_ctk == engine::KCT::q8_0) gpu_attn_softmax(q, qvec, kv_k, kv_v, attn_out_buf, pos); - else gpu_attn_softmax(q, qvec, kv_k, kv_v, attn_out_buf, pos); + if (k_ctk == engine::KCT::fp16) gpu_attn_softmax(q, qvec, kv_k, kv_v, attn_out_buf, pos, fd_acc, fd_ml, fd_max_split); + else if (k_ctk == engine::KCT::q8_0) gpu_attn_softmax(q, qvec, kv_k, kv_v, attn_out_buf, pos, fd_acc, fd_ml, fd_max_split); + else gpu_attn_softmax(q, qvec, kv_k, kv_v, attn_out_buf, pos, fd_acc, fd_ml, fd_max_split); /* sigmoid_mul producer packs; o-proj writes into x with accum=x (T6). */ gpu_sigmoid_mul_quant(q, attn_out_buf, gate, QX_Q_DIM, g_bvnni_d, g_suma_d, g_act_scale_d); @@ -1053,6 +1105,8 @@ bool engine::init(qx_model* m, int mc) { attn_k = dev_shared_alloc(q, (size_t)C * QX_KV_DIM); attn_v = dev_shared_alloc(q, (size_t)C * QX_KV_DIM); attn_out = dev_shared_alloc(q, (size_t)C * QX_Q_DIM); + fd_acc = sycl::malloc_device((size_t)FD_MAX_SPLIT * QX_N_HEAD * QX_HEAD_DIM, q); + fd_ml = sycl::malloc_device((size_t)FD_MAX_SPLIT * QX_N_HEAD * 2, q); bvnni_d = (int8_t*)sycl::malloc_device((size_t)QX_FF_INTER * C, q); suma_d = sycl::malloc_device((size_t)(QX_FF_INTER / QX_QK2) * C, q); act_scale_d = sycl::malloc_shared(C, q); @@ -1206,7 +1260,7 @@ void engine::forward(int slot, int token_id) { attn_forward(q, &attns[ai], x, st.kv_k[ai], st.kv_v[ai], pos, k_ctk, part_buf, proj_buf, attn_q, attn_gate, attn_k, attn_v, attn_out, - rope_c, rope_s); + rope_c, rope_s, fd_acc, fd_ml, FD_MAX_SPLIT); ai++; phase(t_attn); } else { @@ -1520,9 +1574,9 @@ int engine::prefill_cached(int slot, const int* token_ids, int n, /* Explicit instantiations of the templated single-query-token attention * oracle so the M1.4 test (attn_prefill_test) and any other TU that sees the * template decl in qxmx_gpu.h can link without re-defining the kernel. */ -template void gpu_attn_softmax(sycl::queue&, const float*, const uint8_t*, const uint32_t*, float*, int64_t); -template void gpu_attn_softmax(sycl::queue&, const float*, const uint8_t*, const uint32_t*, float*, int64_t); -template void gpu_attn_softmax(sycl::queue&, const float*, const uint8_t*, const uint32_t*, float*, int64_t); +template void gpu_attn_softmax(sycl::queue&, const float*, const uint8_t*, const uint32_t*, float*, int64_t, float*, float*, int); +template void gpu_attn_softmax(sycl::queue&, const float*, const uint8_t*, const uint32_t*, float*, int64_t, float*, float*, int); +template void gpu_attn_softmax(sycl::queue&, const float*, const uint8_t*, const uint32_t*, float*, int64_t, float*, float*, int); /* M1.3: batched per-row rmsnorm+quant (oracle: gpu_rmsnorm_quant). */ void gpu_rmsnorm_quant(sycl::queue& q, const float* x, const float* w, diff --git a/src/qxmx_gpu.h b/src/qxmx_gpu.h index a03334a..c906fa3 100644 --- a/src/qxmx_gpu.h +++ b/src/qxmx_gpu.h @@ -212,6 +212,8 @@ struct engine : engine_i { float* attn_k; /* [CHUNK][KV_DIM] */ float* attn_v; /* [CHUNK][KV_DIM] */ float* attn_out; /* [CHUNK][Q_DIM] */ + float* fd_acc; /* [FD_MAX_SPLIT][N_HEAD][HEAD_DIM] FlashDecoding partial acc */ + float* fd_ml; /* [FD_MAX_SPLIT][N_HEAD][2] partial (m, l) */ /* device-side activation staging (avoid per-call malloc/free). * bvnni_d layout = [K/4][n_tok] int32 = K*n_tok bytes; max K = FF_INTER. */ int8_t* bvnni_d; /* device [FF_INTER*CHUNK] bytes */ @@ -280,11 +282,15 @@ void gpu_quant_k_row_fp8(sycl::queue& q, const float* src, uint8_t* dst, int n_t void gpu_quant_v_row(sycl::queue& q, const float* src, uint32_t* dst_u32, int n_tok = 1); /* Templated causal attention. gpu_attn_softmax = single query token (decode - * oracle). Explicitly instantiated in qxmx_gpu.cpp for KFp16/KQ8_0/KFp8. */ + * oracle). With fd_acc/fd_ml scratch provided it runs FlashDecoding split-KV + * (partials + merge); with defaults it stays the exact single-pass oracle. + * Explicitly instantiated in qxmx_gpu.cpp for KFp16/KQ8_0/KFp8. */ template void gpu_attn_softmax(sycl::queue& q, const float* q_vec, const uint8_t* kv_k, const uint32_t* kv_v, - float* out, int64_t pos); + float* out, int64_t pos, + float* fd_acc = nullptr, float* fd_ml = nullptr, + int fd_max_split = 0); /* Set the file-static engine quant-stage globals (g_bvnni_d / g_suma_d / * g_act_scale_d) so the decode *_forward oracles run outside engine::init @@ -293,8 +299,7 @@ void gpu_attn_softmax(sycl::queue& q, const float* q_vec, void engine_set_quant_stage(int8_t* bvnni_d, int32_t* suma_d, float* act_scale_d); /* Decode single-token block forward paths (qxmx_gpu.cpp). Exported as the - * per-token oracle for the M1.6 batched _b variants. The xs/bvnni/suma args - * are vestigial (dead inside; cleanup parked) — pass nullptr from tests. */ + * per-token oracle for the M1.6 batched _b variants. */ void ffn_forward(sycl::queue& q, const dev_ffn_blk* f, float* x, float* part, float* g, float* u, float* t); @@ -308,7 +313,9 @@ void attn_forward(sycl::queue& q, const dev_attn_blk* a, float* proj_buf, float* qvec, float* gate, float* kvec, float* vvec, float* attn_out_buf, - const float* rope_c, const float* rope_s); + const float* rope_c, const float* rope_s, + float* fd_acc = nullptr, float* fd_ml = nullptr, + int fd_max_split = 0); void ssm_forward(sycl::queue& q, const dev_ssm_blk* s, float* ssm_state, float* ssm_conv, float* x, float* part, diff --git a/tests/fa_kernel_test.cpp b/tests/fa_kernel_test.cpp index 704bff3..b3a5938 100644 --- a/tests/fa_kernel_test.cpp +++ b/tests/fa_kernel_test.cpp @@ -108,6 +108,28 @@ static int run_format(sycl::queue& q, const char* kname, ref_out + (size_t)t * QOUT, pos_start + t); q.wait(); + /* M4: FlashDecoding split-KV decode path vs the same oracle. Default + * QXMX_FD_CHUNK=256 -> positions >255 exercise S>=2 splits + merge. + * Reassociation-only drift: gate mean|d| <= 1e-3, worst <= 5e-2. */ + constexpr int FD_MAX = 64; + float* fd_acc = sycl::malloc_device((size_t)FD_MAX * NQ * HD, q); + float* fd_ml = sycl::malloc_device((size_t)FD_MAX * NQ * 2, q); + float* fd_out = sycl::malloc_shared((size_t)n_tok * QOUT, q); + for (int t = 0; t < n_tok; t++) + gpu_attn_softmax(q, d_Q + (size_t)t * QOUT, kvk, kvv, + fd_out + (size_t)t * QOUT, pos_start + t, + fd_acc, fd_ml, FD_MAX); + q.wait(); + float fd_md = mean_abs_diff(fd_out, ref_out, (size_t)n_tok * QOUT); + float fd_worst = 0.0f; + for (size_t i = 0; i < (size_t)n_tok * QOUT; i++) { + float d = fabsf(fd_out[i] - ref_out[i]); + if (d > fd_worst) fd_worst = d; + } + bool fd_ok = fd_md <= 1e-3f && fd_worst <= 5e-2f; + printf(" K=%-5s pos_start=%-4lld n_tok=%-4d : FD split-KV mean|d|=%.4e worst=%.4e %s\n", + kname, (long long)pos_start, n_tok, fd_md, fd_worst, fd_ok ? "PASS" : "FAIL"); + /* M1.7b.7: 3-kernel split (gemm_tf32 QK -> softmax -> gemm_tf32 PV) needs * n_tok % 8 (gemm_tf32_batched M constraint). The engine prefill path * only calls FA with M%16 (CHUNK=256 or 16-multiple mid-chunks); ragged @@ -117,12 +139,13 @@ static int run_format(sycl::queue& q, const char* kname, if (fa_n % 8 != 0) { printf(" K=%-5s pos_start=%-4lld n_tok=%-4d : SKIP (n_tok %% 8 != 0; tail uses decode, not FA)\n", kname, (long long)pos_start, n_tok); + sycl::free(fd_acc, q); sycl::free(fd_ml, q); sycl::free(fd_out, q); sycl::free(ref_out, q); sycl::free(fa_out, q); sycl::free(kvk, q); sycl::free(kvv, q); if (d_hist_k) sycl::free(d_hist_k, q); if (d_hist_v) sycl::free(d_hist_v, q); sycl::free(d_chk_k, q); sycl::free(d_chk_v, q); sycl::free(d_Q, q); - return 0; + return fd_ok ? 0 : 1; } fa_split_workspace fa_ws = fa_split_alloc(q, fa_n); gpu_flash_attn(q, d_Q, kvk, kvv, fa_out, pos_start, fa_n, &fa_ws); @@ -145,13 +168,13 @@ static int run_format(sycl::queue& q, const char* kname, } } - int bad = 0; (void)bad; + sycl::free(fd_acc, q); sycl::free(fd_ml, q); sycl::free(fd_out, q); sycl::free(ref_out, q); sycl::free(fa_out, q); sycl::free(kvk, q); sycl::free(kvv, q); if (d_hist_k) sycl::free(d_hist_k, q); if (d_hist_v) sycl::free(d_hist_v, q); sycl::free(d_chk_k, q); sycl::free(d_chk_v, q); sycl::free(d_Q, q); - return bad; + return fd_ok ? 0 : 1; } int main() { @@ -167,6 +190,7 @@ int main() { {37, 256, "larger appended"}, {5, 7, "ragged tail (skip -- n_tok<128)"}, {40, 96, "ragged tail (skip -- n_tok<128)"}, + {2048, 256, "deep history (FD split count >= 8)"}, }; for (const auto& c : cfgs) { printf(" [%s]\n", c.note);