diff --git a/src/qxmx_deltanet_wyut.cpp b/src/qxmx_deltanet_wyut.cpp index 0322b65..577548a 100644 --- a/src/qxmx_deltanet_wyut.cpp +++ b/src/qxmx_deltanet_wyut.cpp @@ -5,6 +5,7 @@ * The meson `wyut` engine source set compiles this; `gpu_deltanet_batched` * dispatches to gpu_deltanet_wyut_b. */ #include "qxmx_deltanet.h" +#include "qxmx_device.h" #include "qxmx_gemm_tf32.h" #include "qxmx.h" #include @@ -149,23 +150,23 @@ void gpu_deltanet_wyut_b(sycl::queue& q, float *S0_scaled, *u_scaled; bool own = (ws == nullptr); if (own) { - Kb=sycl::malloc_shared((size_t)NVH*T*HD,q); - Qb=sycl::malloc_shared((size_t)NVH*T*HD,q); - KbT=sycl::malloc_shared((size_t)NVH*HD*T,q); - QbT=sycl::malloc_shared((size_t)NVH*HD*T,q); - KK=sycl::malloc_shared((size_t)NVH*T*T,q); - KQ=sycl::malloc_shared((size_t)NVH*T*T,q); - p=sycl::malloc_shared((size_t)NVH*T*HD,q); - u=sycl::malloc_shared((size_t)NVH*T*HD,q); - S0Q=sycl::malloc_shared((size_t)NVH*T*HD,q); - gl=sycl::malloc_shared((size_t)NVH*T,q); - Ga=sycl::malloc_shared((size_t)NVH*T,q); - bs=sycl::malloc_shared((size_t)NVH*T,q); - g_one=sycl::malloc_shared((size_t)NVH*T,q); - out_acc=sycl::malloc_shared((size_t)NVH*T*HD,q); - GKQ=sycl::malloc_shared((size_t)NVH*T*T,q); - S0_scaled=sycl::malloc_shared((size_t)NVH*HD*HD,q); - u_scaled=sycl::malloc_shared((size_t)NVH*T*HD,q); + Kb=chk_shared(q,(size_t)NVH*T*HD,"wyut.Kb"); + Qb=chk_shared(q,(size_t)NVH*T*HD,"wyut.Qb"); + KbT=chk_shared(q,(size_t)NVH*HD*T,"wyut.KbT"); + QbT=chk_shared(q,(size_t)NVH*HD*T,"wyut.QbT"); + KK=chk_shared(q,(size_t)NVH*T*T,"wyut.KK"); + KQ=chk_shared(q,(size_t)NVH*T*T,"wyut.KQ"); + p=chk_shared(q,(size_t)NVH*T*HD,"wyut.p"); + u=chk_shared(q,(size_t)NVH*T*HD,"wyut.u"); + S0Q=chk_shared(q,(size_t)NVH*T*HD,"wyut.S0Q"); + gl=chk_shared(q,(size_t)NVH*T,"wyut.gl"); + Ga=chk_shared(q,(size_t)NVH*T,"wyut.Ga"); + bs=chk_shared(q,(size_t)NVH*T,"wyut.bs"); + g_one=chk_shared(q,(size_t)NVH*T,"wyut.g_one"); + out_acc=chk_shared(q,(size_t)NVH*T*HD,"wyut.out_acc"); + GKQ=chk_shared(q,(size_t)NVH*T*T,"wyut.GKQ"); + S0_scaled=chk_shared(q,(size_t)NVH*HD*HD,"wyut.S0_scaled"); + u_scaled=chk_shared(q,(size_t)NVH*T*HD,"wyut.u_scaled"); } else { Kb=ws->Kb; Qb=ws->Qb; KbT=ws->KbT; QbT=ws->QbT; KK=ws->KK; KQ=ws->KQ; p=ws->p; u=ws->u; S0Q=ws->S0Q; gl=ws->gl; Ga=ws->Ga; bs=ws->bs; diff --git a/src/qxmx_device.cpp b/src/qxmx_device.cpp index 7a1f3e7..185802d 100644 --- a/src/qxmx_device.cpp +++ b/src/qxmx_device.cpp @@ -25,7 +25,7 @@ std::uint64_t device::checksum(const void* p, size_t nbytes) { constexpr size_t LWS = 64; constexpr size_t CHUNK = 4096; /* bytes per work-group */ size_t ngroups = (nbytes + CHUNK - 1) / CHUNK; - std::uint64_t* partials = sycl::malloc_shared(ngroups, q); + std::uint64_t* partials = chk_shared(q, ngroups, "device.checksum.partials"); q.parallel_for( sycl::nd_range<1>(ngroups * LWS, sycl::range<1>(LWS)), [=](sycl::nd_item<1> it) { diff --git a/src/qxmx_device.h b/src/qxmx_device.h index f8792a9..295a218 100644 --- a/src/qxmx_device.h +++ b/src/qxmx_device.h @@ -6,6 +6,8 @@ #include #include +#include +#include #include @@ -66,5 +68,33 @@ inline std::string device_name(const sycl::queue& q) { } } +/* USM allocators that abort with a diagnostic on null. sycl::malloc_shared / + * malloc_device return nullptr on failure (OOM, or shared USM unsupported by + * the active backend — e.g. WSL Level-Zero); a later q.fill/q.memset on that + * null pointer then throws the opaque "NULL pointer argument in memory fill + * operation", hiding the cause. These name the buffer, size, and device. */ +[[noreturn]] inline void abort_bad_alloc(sycl::queue& q, const char* what, + size_t bytes, const char* kind) { + std::fprintf(stderr, + "qxmx: %s returned null for '%s' (%zu bytes) on device [%s].\n" + " Usually out of device memory, or shared USM unsupported by\n" + " the active SYCL backend (e.g. WSL Level-Zero).\n", + kind, what, bytes, device_name(q).c_str()); + std::abort(); +} + +template +inline T* chk_shared(sycl::queue& q, size_t n, const char* what) { + T* p = sycl::malloc_shared(n, q); + if (!p) abort_bad_alloc(q, what, n * sizeof(T), "malloc_shared"); + return p; +} +template +inline T* chk_device(sycl::queue& q, size_t n, const char* what) { + T* p = sycl::malloc_device(n, q); + if (!p) abort_bad_alloc(q, what, n * sizeof(T), "malloc_device"); + return p; +} + } // namespace qx #endif diff --git a/src/qxmx_fa_fused.cpp b/src/qxmx_fa_fused.cpp index c01247c..190593d 100644 --- a/src/qxmx_fa_fused.cpp +++ b/src/qxmx_fa_fused.cpp @@ -42,6 +42,7 @@ * the engine's %16 mid-chunks work unchanged. */ #include "qxmx_fa.h" +#include "qxmx_device.h" #include #include @@ -127,8 +128,8 @@ static void fa_shadow_ensure(sycl::queue& q, fa_split_workspace* ws, q.wait(); if (ws->Kh) sycl::free(ws->Kh, q); if (ws->Vh) sycl::free(ws->Vh, q); - ws->Kh = sycl::malloc_shared((size_t)cap * KV_DIM, q); - ws->Vh = sycl::malloc_shared((size_t)cap * KV_DIM, q); + ws->Kh = chk_shared(q, (size_t)cap * KV_DIM, "fa_ws.Kh"); + ws->Vh = chk_shared(q, (size_t)cap * KV_DIM, "fa_ws.Vh"); ws->shadow_tok = cap; } diff --git a/src/qxmx_gpu.cpp b/src/qxmx_gpu.cpp index b7055a9..9cb0feb 100644 --- a/src/qxmx_gpu.cpp +++ b/src/qxmx_gpu.cpp @@ -3,6 +3,7 @@ #include "qxmx_gemm.h" #include "qxmx_fa.h" #include "qxmx_prefix_cache.h" +#include "qxmx_device.h" /* qx::device_name + chk_shared/chk_device */ #include #include @@ -1482,8 +1483,8 @@ static dev_q2 upload_q2(sycl::queue& q, const qx_weight* w) { else repack_q2_g128((const qx_block_q2_g128*)w->data, M, K, codes_h, sc_h); dev_q2 d; - d.codes = sycl::malloc_shared(codes_sz, q); - d.scales = sycl::malloc_shared(sc_sz, q); + d.codes = chk_shared(q, codes_sz, "q2.codes"); + d.scales = chk_shared(q, sc_sz, "q2.scales"); d.M = M; d.K = K; d.src = w; q.memcpy(d.codes, codes_h, codes_sz).wait(); q.memcpy(d.scales, sc_h, sc_sz * sizeof(float)).wait(); @@ -1492,17 +1493,19 @@ static dev_q2 upload_q2(sycl::queue& q, const qx_weight* w) { } static float* upload_f32(sycl::queue& q, const qx_weight* w) { - float* d = sycl::malloc_shared(w->n_elem, q); + float* d = chk_shared(q, w->n_elem, "upload_f32"); q.memcpy(d, w->data, w->nbytes).wait(); return d; } -/* Device-resident shared USM for hot activation buffers — avoids the H2D paging - * that malloc_host (host-pinned) triggers on every GPU read. Host can still - * access ( coherent), needed for x (host writes embedding) and logits (host - * reads argmax). */ -static float* dev_shared_alloc(sycl::queue& q, size_t n) { - return sycl::malloc_shared(n, q); +/* USM allocation failure -> nullptr. A subsequent q.fill/q.memset on that + * null pointer throws the cryptic "NULL pointer argument in memory fill + * operation" from the oneAPI runtime, hiding the real cause (OOM, or shared + * USM unsupported under a given backend — e.g. WSL Level-Zero). The checked + * allocators (qx::chk_shared/chk_device, header-only in qxmx_device.h) abort + * with the buffer name, size, and device instead. */ +static float* dev_shared_alloc(sycl::queue& q, size_t n, const char* what) { + return qx::chk_shared(q, n, what); } bool engine::init(qx_model* m, int mc) { @@ -1557,70 +1560,70 @@ bool engine::init(qx_model* m, int mc) { cv, chunk); } int C = chunk; - x = dev_shared_alloc(q, (size_t)C * QX_D_MODEL); - proj_buf = dev_shared_alloc(q, (size_t)C * 2 * QX_Q_DIM); /* > VOCAB at C=256 */ - part_buf = dev_shared_alloc(q, (size_t)8 * QX_N_VOCAB); - logits = dev_shared_alloc(q, QX_N_VOCAB); - alpha_buf = dev_shared_alloc(q, (size_t)C * QX_SSM_DT_RANK); - beta_buf = dev_shared_alloc(q, (size_t)C * QX_SSM_DT_RANK); - ssm_out_buf = dev_shared_alloc(q, (size_t)C * QX_SSM_VALD); - zg_buf = dev_shared_alloc(q, (size_t)C * QX_SSM_VALD); - ffn_g = dev_shared_alloc(q, (size_t)C * QX_FF_INTER); - ffn_u = dev_shared_alloc(q, (size_t)C * QX_FF_INTER); - ffn_t = dev_shared_alloc(q, (size_t)C * QX_FF_INTER); - attn_q = dev_shared_alloc(q, (size_t)C * QX_Q_DIM); - attn_gate = dev_shared_alloc(q, (size_t)C * QX_Q_DIM); - 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); - logits_batch_buf = dev_shared_alloc(q, (size_t)16 * QX_N_VOCAB); - 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); - rope_c = dev_shared_alloc(q, (size_t)C * (QX_ROT_DIM / 2)); - rope_s = dev_shared_alloc(q, (size_t)C * (QX_ROT_DIM / 2)); + x = dev_shared_alloc(q, (size_t)C * QX_D_MODEL, "x"); + proj_buf = dev_shared_alloc(q, (size_t)C * 2 * QX_Q_DIM, "proj_buf"); /* > VOCAB at C=256 */ + part_buf = dev_shared_alloc(q, (size_t)8 * QX_N_VOCAB, "part_buf"); + logits = dev_shared_alloc(q, QX_N_VOCAB, "logits"); + alpha_buf = dev_shared_alloc(q, (size_t)C * QX_SSM_DT_RANK, "alpha_buf"); + beta_buf = dev_shared_alloc(q, (size_t)C * QX_SSM_DT_RANK, "beta_buf"); + ssm_out_buf = dev_shared_alloc(q, (size_t)C * QX_SSM_VALD, "ssm_out_buf"); + zg_buf = dev_shared_alloc(q, (size_t)C * QX_SSM_VALD, "zg_buf"); + ffn_g = dev_shared_alloc(q, (size_t)C * QX_FF_INTER, "ffn_g"); + ffn_u = dev_shared_alloc(q, (size_t)C * QX_FF_INTER, "ffn_u"); + ffn_t = dev_shared_alloc(q, (size_t)C * QX_FF_INTER, "ffn_t"); + attn_q = dev_shared_alloc(q, (size_t)C * QX_Q_DIM, "attn_q"); + attn_gate = dev_shared_alloc(q, (size_t)C * QX_Q_DIM, "attn_gate"); + attn_k = dev_shared_alloc(q, (size_t)C * QX_KV_DIM, "attn_k"); + attn_v = dev_shared_alloc(q, (size_t)C * QX_KV_DIM, "attn_v"); + attn_out = dev_shared_alloc(q, (size_t)C * QX_Q_DIM, "attn_out"); + fd_acc = chk_device(q, (size_t)FD_MAX_SPLIT * QX_N_HEAD * QX_HEAD_DIM, "fd_acc"); + fd_ml = chk_device(q, (size_t)FD_MAX_SPLIT * QX_N_HEAD * 2, "fd_ml"); + logits_batch_buf = dev_shared_alloc(q, (size_t)16 * QX_N_VOCAB, "logits_batch_buf"); + bvnni_d = chk_device(q, (size_t)QX_FF_INTER * C, "bvnni_d"); + suma_d = chk_device(q, (size_t)(QX_FF_INTER / QX_QK2) * C, "suma_d"); + act_scale_d = chk_shared(q, C, "act_scale_d"); + rope_c = dev_shared_alloc(q, (size_t)C * (QX_ROT_DIM / 2), "rope_c"); + rope_s = dev_shared_alloc(q, (size_t)C * (QX_ROT_DIM / 2), "rope_s"); /* M1.7a WY/UT workspace (sized for C=CHUNK=256, NVH=48, HD=128). */ { const int T = chunk, Hd = QX_SSM_STATE, Nvh = QX_SSM_DT_RANK; - wy_ws.Kb = sycl::malloc_shared((size_t)Nvh*T*Hd, q); - wy_ws.Qb = sycl::malloc_shared((size_t)Nvh*T*Hd, q); - wy_ws.KbT = sycl::malloc_shared((size_t)Nvh*Hd*T, q); - wy_ws.QbT = sycl::malloc_shared((size_t)Nvh*Hd*T, q); - wy_ws.KK = sycl::malloc_shared((size_t)Nvh*T*T, q); - wy_ws.KQ = sycl::malloc_shared((size_t)Nvh*T*T, q); - wy_ws.p = sycl::malloc_shared((size_t)Nvh*T*Hd, q); - wy_ws.u = sycl::malloc_shared((size_t)Nvh*T*Hd, q); - wy_ws.S0Q = sycl::malloc_shared((size_t)Nvh*T*Hd, q); - wy_ws.gl = sycl::malloc_shared((size_t)Nvh*T, q); - wy_ws.Ga = sycl::malloc_shared((size_t)Nvh*T, q); - wy_ws.bs = sycl::malloc_shared((size_t)Nvh*T, q); - wy_ws.g_one = sycl::malloc_shared((size_t)Nvh*T, q); - wy_ws.out_acc = sycl::malloc_shared((size_t)Nvh*T*Hd, q); - wy_ws.GKQ = sycl::malloc_shared((size_t)Nvh*T*T, q); - wy_ws.S0_scaled = sycl::malloc_shared((size_t)Nvh*Hd*Hd, q); - wy_ws.u_scaled = sycl::malloc_shared((size_t)Nvh*T*Hd, q); + wy_ws.Kb = chk_shared(q, (size_t)Nvh*T*Hd, "wy_ws.Kb"); + wy_ws.Qb = chk_shared(q, (size_t)Nvh*T*Hd, "wy_ws.Qb"); + wy_ws.KbT = chk_shared(q, (size_t)Nvh*Hd*T, "wy_ws.KbT"); + wy_ws.QbT = chk_shared(q, (size_t)Nvh*Hd*T, "wy_ws.QbT"); + wy_ws.KK = chk_shared(q, (size_t)Nvh*T*T, "wy_ws.KK"); + wy_ws.KQ = chk_shared(q, (size_t)Nvh*T*T, "wy_ws.KQ"); + wy_ws.p = chk_shared(q, (size_t)Nvh*T*Hd, "wy_ws.p"); + wy_ws.u = chk_shared(q, (size_t)Nvh*T*Hd, "wy_ws.u"); + wy_ws.S0Q = chk_shared(q, (size_t)Nvh*T*Hd, "wy_ws.S0Q"); + wy_ws.gl = chk_shared(q, (size_t)Nvh*T, "wy_ws.gl"); + wy_ws.Ga = chk_shared(q, (size_t)Nvh*T, "wy_ws.Ga"); + wy_ws.bs = chk_shared(q, (size_t)Nvh*T, "wy_ws.bs"); + wy_ws.g_one = chk_shared(q, (size_t)Nvh*T, "wy_ws.g_one"); + wy_ws.out_acc = chk_shared(q, (size_t)Nvh*T*Hd, "wy_ws.out_acc"); + wy_ws.GKQ = chk_shared(q, (size_t)Nvh*T*T, "wy_ws.GKQ"); + wy_ws.S0_scaled = chk_shared(q, (size_t)Nvh*Hd*Hd, "wy_ws.S0_scaled"); + wy_ws.u_scaled = chk_shared(q, (size_t)Nvh*T*Hd, "wy_ws.u_scaled"); } /* M1.7b.7 3-kernel FA-split workspace. Sized for `chunk` tokens (the * runtime chunk), TILE_T=128. Kt/Vt sized [NQ] (replicated across GQA -- * see fa_dequant_kt/vt). */ { const int T = chunk, Tt = 128; - fa_ws.Kt = sycl::malloc_shared((size_t)QX_N_HEAD * QX_HEAD_DIM * Tt, q); - fa_ws.Vt = sycl::malloc_shared((size_t)QX_N_HEAD * Tt * QX_HEAD_DIM, q); - fa_ws.S = sycl::malloc_shared((size_t)QX_N_HEAD * T * Tt, q); - fa_ws.P = sycl::malloc_shared((size_t)QX_N_HEAD * T * Tt, q); - fa_ws.O = sycl::malloc_shared((size_t)QX_N_HEAD * T * QX_HEAD_DIM, q); - fa_ws.Oacc = sycl::malloc_shared((size_t)QX_N_HEAD * T * QX_HEAD_DIM, q); - fa_ws.m_state = sycl::malloc_shared((size_t)QX_N_HEAD * T, q); - fa_ws.l_state = sycl::malloc_shared((size_t)QX_N_HEAD * T, q); - fa_ws.Qsep = sycl::malloc_shared((size_t)QX_N_HEAD * T * QX_HEAD_DIM, q); + fa_ws.Kt = chk_shared(q, (size_t)QX_N_HEAD * QX_HEAD_DIM * Tt, "fa_ws.Kt"); + fa_ws.Vt = chk_shared(q, (size_t)QX_N_HEAD * Tt * QX_HEAD_DIM, "fa_ws.Vt"); + fa_ws.S = chk_shared(q, (size_t)QX_N_HEAD * T * Tt, "fa_ws.S"); + fa_ws.P = chk_shared(q, (size_t)QX_N_HEAD * T * Tt, "fa_ws.P"); + fa_ws.O = chk_shared(q, (size_t)QX_N_HEAD * T * QX_HEAD_DIM, "fa_ws.O"); + fa_ws.Oacc = chk_shared(q, (size_t)QX_N_HEAD * T * QX_HEAD_DIM, "fa_ws.Oacc"); + fa_ws.m_state = chk_shared(q, (size_t)QX_N_HEAD * T, "fa_ws.m_state"); + fa_ws.l_state = chk_shared(q, (size_t)QX_N_HEAD * T, "fa_ws.l_state"); + fa_ws.Qsep = chk_shared(q, (size_t)QX_N_HEAD * T * QX_HEAD_DIM, "fa_ws.Qsep"); /* M1.7c fused-FA fp16 K/V shadow (grown on demand by * fa_shadow_ensure). 8192 tokens covers the bench prompt. */ fa_ws.shadow_tok = 8192; - fa_ws.Kh = sycl::malloc_shared((size_t)fa_ws.shadow_tok * QX_KV_DIM, q); - fa_ws.Vh = sycl::malloc_shared((size_t)fa_ws.shadow_tok * QX_KV_DIM, q); + fa_ws.Kh = chk_shared(q, (size_t)fa_ws.shadow_tok * QX_KV_DIM, "fa_ws.Kh"); + fa_ws.Vh = chk_shared(q, (size_t)fa_ws.shadow_tok * QX_KV_DIM, "fa_ws.Vh"); } g_bvnni_d = bvnni_d; g_suma_d = suma_d; g_act_scale_d = act_scale_d; return true; @@ -1641,14 +1644,14 @@ void engine::on_open_slot(int slot) { size_t krow_bytes = (k_ctk == KCT::fp16) ? (size_t)KROW_BYTES_F16 : (size_t)KROW_BYTES_Q8; for (int i = 0; i < n_ssm; i++) { - s.ssm_state[i] = sycl::malloc_shared(QX_SSM_DT_RANK * QX_SSM_STATE * QX_SSM_STATE, q); - s.ssm_conv[i] = sycl::malloc_shared((QX_SSM_CONV_K-1) * QX_SSM_CONVD, q); + s.ssm_state[i] = chk_shared(q, QX_SSM_DT_RANK * QX_SSM_STATE * QX_SSM_STATE, "slot.ssm_state"); + s.ssm_conv[i] = chk_shared(q, (QX_SSM_CONV_K-1) * QX_SSM_CONVD, "slot.ssm_conv"); q.memset(s.ssm_state[i], 0, QX_SSM_DT_RANK*QX_SSM_STATE*QX_SSM_STATE*sizeof(float)); q.memset(s.ssm_conv[i], 0, (QX_SSM_CONV_K-1)*QX_SSM_CONVD*sizeof(float)); } for (int i = 0; i < n_attn; i++) { - s.kv_k[i] = sycl::malloc_shared((size_t)max_ctx * krow_bytes, q); - s.kv_v[i] = sycl::malloc_shared((size_t)max_ctx * VROW_U32, q); + s.kv_k[i] = chk_shared(q, (size_t)max_ctx * krow_bytes, "slot.kv_k"); + s.kv_v[i] = chk_shared(q, (size_t)max_ctx * VROW_U32, "slot.kv_v"); q.memset(s.kv_k[i], 0, (size_t)max_ctx * krow_bytes); q.memset(s.kv_v[i], 0, (size_t)max_ctx*VROW_U32*sizeof(uint32_t)); } diff --git a/src/qxmx_gpu.h b/src/qxmx_gpu.h index ba8225d..f4bd2c9 100644 --- a/src/qxmx_gpu.h +++ b/src/qxmx_gpu.h @@ -20,6 +20,7 @@ #include "qxmx_model.h" #include "qxmx_deltanet.h" #include "qxmx_prefix_cache.h" +#include "qxmx_device.h" namespace qx { @@ -63,18 +64,18 @@ struct fa_split_workspace { inline fa_split_workspace fa_split_alloc(sycl::queue& q, int n_tok_max) { fa_split_workspace w; constexpr int TILE_T = 128; - w.Kt = sycl::malloc_shared((size_t)QX_N_HEAD * QX_HEAD_DIM * TILE_T, q); - w.Vt = sycl::malloc_shared((size_t)QX_N_HEAD * TILE_T * QX_HEAD_DIM, q); - w.S = sycl::malloc_shared((size_t)QX_N_HEAD * n_tok_max * TILE_T, q); - w.P = sycl::malloc_shared((size_t)QX_N_HEAD * n_tok_max * TILE_T, q); - w.O = sycl::malloc_shared((size_t)QX_N_HEAD * n_tok_max * QX_HEAD_DIM, q); - w.Oacc = sycl::malloc_shared((size_t)QX_N_HEAD * n_tok_max * QX_HEAD_DIM, q); - w.m_state = sycl::malloc_shared((size_t)QX_N_HEAD * n_tok_max, q); - w.l_state = sycl::malloc_shared((size_t)QX_N_HEAD * n_tok_max, q); - w.Qsep = sycl::malloc_shared((size_t)QX_N_HEAD * n_tok_max * QX_HEAD_DIM, q); + w.Kt = chk_shared(q, (size_t)QX_N_HEAD * QX_HEAD_DIM * TILE_T, "fa_split.Kt"); + w.Vt = chk_shared(q, (size_t)QX_N_HEAD * TILE_T * QX_HEAD_DIM, "fa_split.Vt"); + w.S = chk_shared(q, (size_t)QX_N_HEAD * n_tok_max * TILE_T, "fa_split.S"); + w.P = chk_shared(q, (size_t)QX_N_HEAD * n_tok_max * TILE_T, "fa_split.P"); + w.O = chk_shared(q, (size_t)QX_N_HEAD * n_tok_max * QX_HEAD_DIM, "fa_split.O"); + w.Oacc = chk_shared(q, (size_t)QX_N_HEAD * n_tok_max * QX_HEAD_DIM, "fa_split.Oacc"); + w.m_state = chk_shared(q, (size_t)QX_N_HEAD * n_tok_max, "fa_split.m_state"); + w.l_state = chk_shared(q, (size_t)QX_N_HEAD * n_tok_max, "fa_split.l_state"); + w.Qsep = chk_shared(q, (size_t)QX_N_HEAD * n_tok_max * QX_HEAD_DIM, "fa_split.Qsep"); w.shadow_tok = 8192; - w.Kh = sycl::malloc_shared((size_t)w.shadow_tok * QX_KV_DIM, q); - w.Vh = sycl::malloc_shared((size_t)w.shadow_tok * QX_KV_DIM, q); + w.Kh = chk_shared(q, (size_t)w.shadow_tok * QX_KV_DIM, "fa_split.Kh"); + w.Vh = chk_shared(q, (size_t)w.shadow_tok * QX_KV_DIM, "fa_split.Vh"); return w; } inline void fa_split_free(sycl::queue& q, fa_split_workspace& w) {