A single-model (Bonsai-27B) inference engine for Intel Arc GPUs
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194/* elementwise_batch_test: M1.6 S2 — bit-exact equivalence of the batched * elementwise kernels (_b) vs M calls to the decode single-token versions. * * Covers: gpu_group_rmsnorm_b, gpu_apply_rope_b, gpu_attn_split_qgate_b, * gpu_silu_inplace_b. All are pure elementwise / per-row reductions, so the * batched form (add a token grid dimension, identical per-row math) must be * bit-exact per row vs the decode oracle — no fp reassociation is involved. * * Build: meson compile -C build elementwise_batch_test && ./build/elementwise_batch_test */#include <cmath>#include <cstdint>#include <cstdio>#include <cstring>#include <random>#include <vector>
#include <sycl/sycl.hpp>
#include "qxmx.h"#include "qxmx_gpu.h"
using namespace qx;
static std::mt19937_64 rng(0x525ULL);
static int g_fail = 0;
/* Fill [n_tok*dim] with varied row scales (matches gemm_engine_test dist). */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; }}
static int cmp_f32(const char* what, const float* a, const float* b, size_t n) { int bad = 0; for (size_t i = 0; i < n && bad < 8; i++) if (a[i] != b[i]) { printf(" %s MISMATCH @%zu: %a vs %a (diff %g)\n", what, i, a[i], b[i], a[i] - b[i]); bad++; } return bad;}
/* ---------------- group_rmsnorm_b ---------------- */static int test_group_rmsnorm(sycl::queue& q, int n_heads, int head_dim, int n_tok) { size_t n = (size_t)n_tok * n_heads * head_dim; std::vector<float> w(n_heads * head_dim); { std::normal_distribution<float> g(0.0f, 1.0f); for (size_t i = 0; i < w.size(); i++) w[i] = 1.0f + 0.1f * g(rng); } float* d_x = sycl::malloc_shared<float>(n, q); float* d_w = sycl::malloc_shared<float>(w.size(), q); float* d_ob = sycl::malloc_shared<float>(n, q); float* d_od = sycl::malloc_shared<float>(n, q); fill_tile(d_x, n_tok, n_heads * head_dim, 1); memcpy(d_w, w.data(), w.size() * sizeof(float));
/* oracle: per-token decode calls */ for (int t = 0; t < n_tok; t++) gpu_group_rmsnorm(q, d_x + (size_t)t * n_heads * head_dim, d_w, d_od + (size_t)t * n_heads * head_dim, n_heads, head_dim); q.wait(); /* batched */ gpu_group_rmsnorm_b(q, d_x, d_w, d_ob, n_heads, head_dim, n_tok); q.wait(); int bad = cmp_f32("out", d_ob, d_od, n); printf(" group_rmsnorm_b n_heads=%d head_dim=%d n_tok=%-4d : %s\n", n_heads, head_dim, n_tok, bad ? "FAILED" : "PASSED (bit-exact)"); sycl::free(d_x, q); sycl::free(d_w, q); sycl::free(d_ob, q); sycl::free(d_od, q); return bad;}
/* ---------------- silu_inplace_b ---------------- */static int test_silu_inplace(sycl::queue& q, int n, int n_tok) { size_t sz = (size_t)n_tok * n; float* d_xb = sycl::malloc_shared<float>(sz, q); float* d_xd = sycl::malloc_shared<float>(sz, q); fill_tile(d_xb, n_tok, n, 2); memcpy(d_xd, d_xb, sz * sizeof(float));
for (int t = 0; t < n_tok; t++) gpu_silu_inplace(q, d_xd + (size_t)t * n, n); gpu_silu_inplace_b(q, d_xb, n, n_tok); q.wait(); int bad = cmp_f32("x", d_xb, d_xd, sz); printf(" silu_inplace_b n=%d n_tok=%-4d : %s\n", n, n_tok, bad ? "FAILED" : "PASSED (bit-exact)"); sycl::free(d_xb, q); sycl::free(d_xd, q); return bad;}
/* ---------------- attn_split_qgate_b ---------------- */static int test_split_qgate(sycl::queue& q, int n_heads, int head_dim, int n_tok) { int n = n_heads * head_dim; size_t proj_sz = (size_t)n_tok * 2 * n; size_t out_sz = (size_t)n_tok * n; float* d_proj = sycl::malloc_shared<float>(proj_sz, q); float* d_qb = sycl::malloc_shared<float>(out_sz, q); float* d_gb = sycl::malloc_shared<float>(out_sz, q); float* d_qd = sycl::malloc_shared<float>(out_sz, q); float* d_gd = sycl::malloc_shared<float>(out_sz, q); fill_tile(d_proj, n_tok, 2 * n, 3);
for (int t = 0; t < n_tok; t++) gpu_attn_split_qgate(q, d_proj + (size_t)t * 2 * n, d_qd + (size_t)t * n, d_gd + (size_t)t * n, n_heads, head_dim); gpu_attn_split_qgate_b(q, d_proj, d_qb, d_gb, n_heads, head_dim, n_tok); q.wait(); int bad = cmp_f32("q", d_qb, d_qd, out_sz); bad += cmp_f32("gate", d_gb, d_gd, out_sz); printf(" split_qgate_b n_heads=%d head_dim=%d n_tok=%-4d : %s\n", n_heads, head_dim, n_tok, bad ? "FAILED" : "PASSED (bit-exact)"); sycl::free(d_proj, q); sycl::free(d_qb, q); sycl::free(d_gb, q); sycl::free(d_qd, q); sycl::free(d_gd, q); return bad;}
/* ---------------- apply_rope_b ---------------- */static int test_apply_rope(sycl::queue& q, int n_q, int n_kv, int head_dim, int rot_dim, int64_t pos_start, int n_tok) { int half = rot_dim / 2; size_t qsz = (size_t)n_tok * n_q * head_dim; size_t ksz = (size_t)n_tok * n_kv * head_dim; size_t tabsz = (size_t)n_tok * half; float* d_qb = sycl::malloc_shared<float>(qsz, q); float* d_kb = sycl::malloc_shared<float>(ksz, q); float* d_qd = sycl::malloc_shared<float>(qsz, q); float* d_kd = sycl::malloc_shared<float>(ksz, q); float* rc = sycl::malloc_shared<float>(tabsz, q); float* rs = sycl::malloc_shared<float>(tabsz, q); fill_tile(d_qb, n_tok, n_q * head_dim, 4); fill_tile(d_kb, n_tok, n_kv * head_dim, 5); memcpy(d_qd, d_qb, qsz * sizeof(float)); memcpy(d_kd, d_kb, ksz * sizeof(float)); engine_compute_rope_table(pos_start, n_tok, rc, rs);
/* oracle: per-token decode with the token's single-position table */ for (int t = 0; t < n_tok; t++) gpu_apply_rope(q, d_qd + (size_t)t * n_q * head_dim, d_kd + (size_t)t * n_kv * head_dim, n_q, n_kv, head_dim, rot_dim, rc + (size_t)t * half, rs + (size_t)t * half); /* batched */ gpu_apply_rope_b(q, d_qb, d_kb, n_q, n_kv, head_dim, rot_dim, rc, rs, n_tok); q.wait(); int bad = cmp_f32("q", d_qb, d_qd, qsz); bad += cmp_f32("k", d_kb, d_kd, ksz); printf(" apply_rope_b pos=%-4lld n_tok=%-4d : %s\n", (long long)pos_start, n_tok, bad ? "FAILED" : "PASSED (bit-exact)"); sycl::free(d_qb, q); sycl::free(d_kb, q); sycl::free(d_qd, q); sycl::free(d_kd, q); sycl::free(rc, q); sycl::free(rs, q); return bad;}
int main() { sycl::queue q(sycl::gpu_selector_v, sycl::property::queue::in_order{}); printf("elementwise_batch_test (M1.6 S2): _b kernels vs decode oracle\n");
/* real model dims (qxmx.h): N_HEAD=24, N_HEAD_KV=4, HEAD_DIM=256, ROT_DIM=64 */ g_fail += test_group_rmsnorm(q, QX_N_HEAD, QX_HEAD_DIM, 1); g_fail += test_group_rmsnorm(q, QX_N_HEAD, QX_HEAD_DIM, 16); g_fail += test_group_rmsnorm(q, QX_N_HEAD, QX_HEAD_DIM, 256); g_fail += test_group_rmsnorm(q, QX_N_HEAD_KV, QX_HEAD_DIM, 256); g_fail += test_group_rmsnorm(q, QX_N_HEAD, QX_HEAD_DIM, 13); /* ragged */
g_fail += test_silu_inplace(q, QX_SSM_CONVD, 1); g_fail += test_silu_inplace(q, QX_SSM_CONVD, 16); g_fail += test_silu_inplace(q, QX_SSM_CONVD, 256); g_fail += test_silu_inplace(q, QX_SSM_CONVD, 13);
g_fail += test_split_qgate(q, QX_N_HEAD, QX_HEAD_DIM, 1); g_fail += test_split_qgate(q, QX_N_HEAD, QX_HEAD_DIM, 16); g_fail += test_split_qgate(q, QX_N_HEAD, QX_HEAD_DIM, 256); g_fail += test_split_qgate(q, QX_N_HEAD, QX_HEAD_DIM, 13);
g_fail += test_apply_rope(q, QX_N_HEAD, QX_N_HEAD_KV, QX_HEAD_DIM, QX_ROT_DIM, 0, 1); g_fail += test_apply_rope(q, QX_N_HEAD, QX_N_HEAD_KV, QX_HEAD_DIM, QX_ROT_DIM, 0, 16); g_fail += test_apply_rope(q, QX_N_HEAD, QX_N_HEAD_KV, QX_HEAD_DIM, QX_ROT_DIM, 0, 256); g_fail += test_apply_rope(q, QX_N_HEAD, QX_N_HEAD_KV, QX_HEAD_DIM, QX_ROT_DIM, 37, 64); g_fail += test_apply_rope(q, QX_N_HEAD, QX_N_HEAD_KV, QX_HEAD_DIM, QX_ROT_DIM, 5, 13);
printf(g_fail ? "\nFAILED: %d group(s)\n" : "\nALL PASSED — bit-exact per (token, row) across configs\n", g_fail); return g_fail ? 1 : 0;}