A single-model (Bonsai-27B) inference engine for Intel Arc GPUs
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174/* fa_kernel_test: M1.7b.1 — bit-exact equivalence of the fused FlashAttention * kernel (qxmx_fa.cpp, gpu_flash_attn) vs the M1.4 single-token decode oracle * (gpu_attn_softmax) and the M1.4 naive prefill (gpu_attn_prefill). * * Mirrors attn_prefill_test's check 2: fills history+chunk caches via the * decode quant oracle, then compares the fused FA output to the per-token * decode oracle across the same 18 configs (3 K formats x 6 {pos_start, n_tok}). * Gate: bitwise-identical per (token, head) is too strict (FA reassociates); * use the M1.x reassociation gate: mean|d| <= 0.10 AND greedy-token match. * * Build: meson compile -C build fa_kernel_test && ./build/fa_kernel_test */#include <cmath>#include <cstdint>#include <cstdio>#include <cstring>#include <random>#include <type_traits>#include <vector>
#include <sycl/sycl.hpp>
#include "qxmx.h"#include "qxmx_gpu.h"#include "qxmx_fa.h"
using namespace qx;
static std::mt19937_64 rng(0xFA17ULL);
static void fill_tile(float* p, int n_tok, int dim, int seed_off) { std::normal_distribution<float> g(0.0f, 1.0f); for (int t = 0; t < n_tok; t++) { float sc = (((t + seed_off) % 17) == 3) ? 0.0f : powf(10.0f, -3.0f + 4.0f * (((t + seed_off) % 13) / 12.0f)); for (int i = 0; i < dim; i++) p[(size_t)t * dim + i] = g(rng) * sc; if (((t + seed_off) % 31) == 5) p[(size_t)t * dim + 7] = 123.0f; }}
template <typename Kct>static void quant_row(sycl::queue& q, const float* src, uint8_t* dst) { if constexpr (std::is_same_v<Kct, KFp16>) gpu_fp32_to_fp16(q, src, reinterpret_cast<sycl::half*>(dst), QX_KV_DIM); else if constexpr (std::is_same_v<Kct, KQ8_0>) gpu_quant_k_row_q8(q, src, dst); else gpu_quant_k_row_fp8(q, src, dst);}
static float mean_abs_diff(const float* a, const float* b, size_t n) { double s = 0.0; for (size_t i = 0; i < n; i++) s += fabs((double)a[i] - (double)b[i]); return (float)(s / (double)n);}
template <typename Kct>static int run_format(sycl::queue& q, const char* kname, int64_t pos_start, int n_tok) { constexpr int NQ = QX_N_HEAD, HD = QX_HEAD_DIM, KV = QX_KV_DIM; constexpr size_t QOUT = (size_t)NQ * HD; const size_t krow = std::is_same_v<Kct, KFp16> ? KROW_BYTES_F16 : KROW_BYTES_Q8; const int max_ctx = (int)(pos_start + n_tok); const size_t kvk_bytes = (size_t)max_ctx * krow; const size_t kvv_u32 = (size_t)max_ctx * VROW_U32;
std::vector<float> hist_k((size_t)pos_start * KV), hist_v((size_t)pos_start * KV); std::vector<float> chk_k((size_t)n_tok * KV), chk_v((size_t)n_tok * KV); std::vector<float> Qt((size_t)n_tok * QOUT); if (pos_start > 0) { fill_tile(hist_k.data(), pos_start, KV, 1); fill_tile(hist_v.data(), pos_start, KV, 2); } fill_tile(chk_k.data(), n_tok, KV, 3); fill_tile(chk_v.data(), n_tok, KV, 4); fill_tile(Qt.data(), n_tok, QOUT, 5);
auto H2D = [&](const std::vector<float>& v) { float* d = sycl::malloc_shared<float>(v.size(), q); memcpy(d, v.data(), v.size() * sizeof(float)); return d; }; float* d_hist_k = pos_start ? H2D(hist_k) : nullptr; float* d_hist_v = pos_start ? H2D(hist_v) : nullptr; float* d_chk_k = H2D(chk_k); float* d_chk_v = H2D(chk_v); float* d_Q = H2D(Qt);
uint8_t* kvk = sycl::malloc_shared<uint8_t>(kvk_bytes, q); uint32_t* kvv = sycl::malloc_shared<uint32_t>(kvv_u32, q); memset(kvk, 0xEE, kvk_bytes); memset(kvv, 0xEE, kvv_u32 * 4);
/* fill cache via decode oracle (ground truth). */ for (int t = 0; t < pos_start; t++) { quant_row<Kct>(q, d_hist_k + (size_t)t * KV, kvk + (size_t)t * krow); gpu_quant_v_row(q, d_hist_v + (size_t)t * KV, kvv + (size_t)t * VROW_U32); } for (int t = 0; t < n_tok; t++) { quant_row<Kct>(q, d_chk_k + (size_t)t * KV, kvk + (size_t)(pos_start + t) * krow); gpu_quant_v_row(q, d_chk_v + (size_t)t * KV, kvv + (size_t)(pos_start + t) * VROW_U32); } q.wait();
/* ref: per-token decode oracle. fa: fused kernel. */ float* ref_out = sycl::malloc_shared<float>((size_t)n_tok * QOUT, q); float* fa_out = sycl::malloc_shared<float>((size_t)n_tok * QOUT, q); memset(fa_out, 0xEE, (size_t)n_tok * QOUT * 4); for (int t = 0; t < n_tok; t++) gpu_attn_softmax<Kct>(q, d_Q + (size_t)t * QOUT, kvk, kvv, ref_out + (size_t)t * QOUT, pos_start + t); q.wait();
/* FA requires n_tok % BLK_Q (128) == 0; pad n_tok up to a multiple of 128 * by extending the chunk (extra tokens past pos_start+n_tok are masked by * the causal cap). Simplest: only test configs where n_tok % 128 == 0 in * the FA path; for non-multiples, compare the first floor(n_tok/128)*128 * rows. Here we just call on the full n_tok when it's a multiple of 128; * otherwise skip the FA call and report it as a "tail" config (M1.7b.3 * will handle tails via gpu_attn_prefill fallback). */ int fa_n = (n_tok / 128) * 128; if (fa_n == 0) { printf(" K=%-5s pos_start=%-4lld n_tok=%-4d : SKIP (n_tok < 128, tail fallback not wired)\n", kname, (long long)pos_start, n_tok); } else { gpu_flash_attn<Kct>(q, d_Q, kvk, kvv, fa_out, pos_start, fa_n); q.wait(); float md = mean_abs_diff(fa_out, ref_out, (size_t)fa_n * QOUT); /* Gate: mean|d| <= 0.10 (M1.x reassociation gate). The attention output * is per-head, not logits -- no greedy-token check here (that gate * applies to engine prefill-then-decode, not this isolated kernel). */ const char* verdict = (md <= 0.10f) ? "PASS" : "FAIL"; printf(" K=%-5s pos_start=%-4lld n_tok=%-4d fa_n=%-3d : mean|d|=%.4e %s\n", kname, (long long)pos_start, n_tok, fa_n, md, verdict); if (md > 0.10f) { for (size_t i = 0; i < (size_t)fa_n * QOUT; i++) { if (fabsf(fa_out[i] - ref_out[i]) > 1e-3f) { printf(" first big diff @%zu: fa=%.6f ref=%.6f\n", i, fa_out[i], ref_out[i]); break; } } } }
int bad = 0; (void)bad; 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;}
int main() { sycl::queue q(sycl::gpu_selector_v, sycl::property::queue::in_order{}); printf("fa_kernel_test (M1.7b.1): fused FlashAttention vs decode oracle\n"); printf(" gate: mean|d| <= 0.10 (reassociation; attention output, not logits)\n"); int total = 0; struct Cfg { int64_t pos; int n_tok; const char* note; }; Cfg cfgs[] = { {0, 128, "pure-causal chunk (1 Q-tile)"}, {0, 256, "pure-causal chunk (2 Q-tiles)"}, {37, 128, "chunk appended to history"}, {37, 256, "larger appended"}, {5, 7, "ragged tail (skip -- n_tok<128)"}, {40, 96, "ragged tail (skip -- n_tok<128)"}, }; for (const auto& c : cfgs) { printf(" [%s]\n", c.note); total += run_format<KFp16>(q, "fp16", c.pos, c.n_tok); total += run_format<KQ8_0>(q, "q8_0", c.pos, c.n_tok); total += run_format<KFp8 >(q, "fp8", c.pos, c.n_tok); } printf(total ? "\nFAILED\n" : "\nALL PASSED\n"); return total ? 1 : 0;}