diff --git a/src/qxmx_gpu.cpp b/src/qxmx_gpu.cpp index 5bff392..bb63952 100644 --- a/src/qxmx_gpu.cpp +++ b/src/qxmx_gpu.cpp @@ -545,16 +545,212 @@ void gpu_attn_softmax(sycl::queue& q, const float* q_vec, const char* e = std::getenv("QXMX_FOLD"); return e && e[0] == '0'; }(); + static const bool vec_disabled = [] { + const char* e = std::getenv("QXMX_VEC"); + return e && e[0] == '0'; + }(); int S = 1; if (fd_acc && fd_ml && !fd_disabled && fd_chunk > 0) { - /* Folded WGs do GQA heads' work per pass — scale the split target by - * GQA to keep the device saturated (NKV*S folded ~= NQ*S unfolded). */ - int64_t s64 = fold_disabled + /* Folded/vec WGs do GQA heads' work per pass — scale the split target + * by GQA to keep the device saturated (NKV*S folded ~= NQ*S unfolded). */ + int64_t s64 = (fold_disabled && vec_disabled) ? (n_kv + fd_chunk - 1) / fd_chunk : (n_kv * GQA + 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; + /* v3 llama-vec scan (default): FD-in-WG. One WG per (kv-head, head-group + * of JG=3, split); each of the WG's 8 SGs scans its own strided token + * stream with SG-LOCAL online-softmax state — zero barriers in the scan + * (fold kernel: 2/token) and V dequant runs once per token (v1 vec: 8x + * redundant). Lane owns 8 consecutive elems (4 lanes per 32-block): K = + * byte loads (uint32 miscompiles at 2-mod-4 alignment on icpx 2026.1), + * V = 1 uint32. acc[3][8] = 24 regs — sized to avoid the v1 spill + * (VTune: 13KB). Scan ends with an 8-way SG merge through SLM (same + * math as attn_fd_merge, once per WG), then the deferred inverse WHT. + * QXMX_VEC=0 falls back to the GQA-fold kernel. */ + if (!vec_disabled) { + constexpr int NKV = QX_N_HEAD_KV; + constexpr int NSG = HD / VBLK; /* 8 SGs/WG = token streams */ + constexpr int JG = GQA / 2; /* heads per WG */ + static_assert(GQA == 2 * JG, "v3 needs GQA = 2*JG"); + q.submit([&](sycl::handler& h) { + /* sl: ml[NSG*JG][2] | a_s[NSG*JG][HD] | q_s[JG][HD] */ + sycl::local_accessor sl(sycl::range<1>( + (size_t)NSG*JG*2 + (size_t)NSG*JG*HD + (size_t)JG*HD), h); + h.parallel_for(sycl::nd_range<1>((size_t)NKV * 2 * S * HD, HD), [=](sycl::nd_item<1> it) + [[sycl::reqd_sub_group_size(VBLK)]] { + float* ml = &sl[0]; + float* a_s = &sl[NSG * JG * 2]; + float* q_s = &sl[NSG * JG * 2 + NSG * JG * HD]; + int g = it.get_group(0); + int hk = g % NKV; + int hg = (g / NKV) & 1; + int sp = g / (NKV * 2); + int lid = it.get_local_id(0); + int lane = lid & 31; + int sg_id = lid >> 5; + auto sg = it.get_sub_group(); + /* lane -> 8 elems: block lane>>2, offset (lane&3)*8 */ + const int blk_row = hk * BLK_PER_HEAD + (lane >> 2); + const int dim0 = (lane >> 2) * VBLK + (lane & 3) * 8; + const int o_l = (lane & 3) * 8; +#pragma unroll + for (int j = 0; j < JG; j++) + q_s[j * HD + lid] = q_vec[(size_t)(hk * GQA + hg * JG + j) * HD + lid]; + sycl::group_barrier(it.get_group()); + float acc[JG][8], max_s[JG], sum_s[JG]; +#pragma unroll + for (int j = 0; j < JG; j++) { + max_s[j] = -1e30f; sum_s[j] = 0.0f; +#pragma unroll + for (int i = 0; i < 8; i++) acc[j][i] = 0.0f; + } + 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 + sg_id; t < t1; t += NSG) { + /* K dequant: 8 elems per lane (t always valid). */ + float k8[8]; + if constexpr (std::is_same_v) { + const sycl::half* kp = reinterpret_cast(kv_k) + + (size_t)t * kv_dim + hk * HD + dim0; + const uint32_t* k32 = reinterpret_cast(kp); +#pragma unroll + for (int w = 0; w < 4; w++) { + uint32_t u = k32[w]; + k8[2*w] = float(sycl::bit_cast((uint16_t)(u & 0xFFFFu))); + k8[2*w + 1] = float(sycl::bit_cast((uint16_t)(u >> 16))); + } + } else { + const uint8_t* kb = kv_k + (size_t)t * Kct::row_bytes + blk_row * KBLK_BYTES; + uint16_t sbits = (uint16_t)kb[0] | ((uint16_t)kb[1] << 8); + float d = float(sycl::bit_cast(sbits)); +#pragma unroll + for (int i = 0; i < 8; i++) { + uint8_t code = kb[2 + o_l + i]; + if constexpr (std::is_same_v) { + k8[i] = (float)(int8_t)code * d; + } else { /* KFp8 */ + k8[i] = e4m3_to_fp32(code) * d; + } + } + } + /* V dequant once for this SG's own token. */ + const uint32_t* db = kv_v + (size_t)t * VROW_U32 + blk_row * 5; + uint16_t sc_u = (uint16_t)(db[0] & 0xFFFFu); + uint16_t nc_u = (uint16_t)((db[0] >> 16) & 0xFFFFu); + float vsn = float(sycl::bit_cast(sc_u)) + * float(sycl::bit_cast(nc_u)); + /* 32 codes x 4bit = one uint32 per (block, lane). */ + uint32_t w0 = db[1 + (lane & 3)]; + float v8[8]; +#pragma unroll + for (int i = 0; i < 8; i++) + v8[i] = V_LM_LUT[(w0 >> (i * 4)) & 0xFu] * vsn; + /* Dots + SG-local online softmax per head. */ +#pragma unroll + for (int j = 0; j < JG; j++) { + float v = 0.0f; +#pragma unroll + for (int i = 0; i < 8; i++) v += q_s[j * HD + dim0 + i] * k8[i]; +#pragma unroll + for (int off = VBLK >> 1; off > 0; off >>= 1) + v += sycl::permute_group_by_xor(sg, v, off); + float score = v * scale; + float m_new = sycl::fmax(max_s[j], score); + float f = sycl::native::exp(max_s[j] - m_new); + float e = sycl::native::exp(score - m_new); + sum_s[j] = f * sum_s[j] + e; +#pragma unroll + for (int i = 0; i < 8; i++) acc[j][i] = f * acc[j][i] + e * v8[i]; + max_s[j] = m_new; + } + } + /* 8-way SG partial merge through SLM (once per WG). Empty + * SGs carry (m=-1e30, l=0, acc=0) → zero weight. */ + if (lane == 0) +#pragma unroll + for (int j = 0; j < JG; j++) { + ml[(sg_id * JG + j) * 2] = max_s[j]; + ml[(sg_id * JG + j) * 2 + 1] = sum_s[j]; + } + sycl::group_barrier(it.get_group()); + float mmax[JG], lsum[JG]; +#pragma unroll + for (int j = 0; j < JG; j++) { + float m = -1e30f; +#pragma unroll + for (int s = 0; s < NSG; s++) m = sycl::fmax(m, ml[(s * JG + j) * 2]); + mmax[j] = m; + float ls = 0.0f; +#pragma unroll + for (int s = 0; s < NSG; s++) + ls += sycl::native::exp(ml[(s * JG + j) * 2] - m) * ml[(s * JG + j) * 2 + 1]; + lsum[j] = ls; + } +#pragma unroll + for (int j = 0; j < JG; j++) { + float wj = sycl::native::exp(max_s[j] - mmax[j]); +#pragma unroll + for (int i = 0; i < 8; i++) + a_s[(sg_id * JG + j) * HD + dim0 + i] = acc[j][i] * wj; + } + sycl::group_barrier(it.get_group()); + /* Epilogue per head: cross-SG sum, deferred inverse WHT — + * in-lane bits {1,2,4}, cross-lane bits {8,16} — then NRM. */ +#pragma unroll + for (int j = 0; j < JG; j++) { + float a[8]; +#pragma unroll + for (int i = 0; i < 8; i++) a[i] = 0.0f; +#pragma unroll + for (int s = 0; s < NSG; s++) +#pragma unroll + for (int i = 0; i < 8; i++) + a[i] += a_s[(s * JG + j) * HD + dim0 + i]; +#pragma unroll + for (int bit = 1; bit < 8; bit <<= 1) { + float na[8]; +#pragma unroll + for (int i = 0; i < 8; i++) + na[i] = (i & bit) ? a[i ^ bit] - a[i] : a[i] + a[i ^ bit]; +#pragma unroll + for (int i = 0; i < 8; i++) a[i] = na[i]; + } +#pragma unroll + for (int bit = 1; bit <= 2; bit <<= 1) { + bool hi = (lane & bit) != 0; +#pragma unroll + for (int i = 0; i < 8; i++) { + float p = sycl::permute_group_by_xor(sg, a[i], bit); + a[i] = hi ? p - a[i] : a[i] + p; + } + } + int hh = hk * GQA + hg * JG + j; + if (S == 1) { + float inv = lsum[j] > 0.0f ? NRM / lsum[j] : 0.0f; +#pragma unroll + for (int i = 0; i < 8; i++) + out[(size_t)hh * HD + dim0 + i] = a[i] * inv; + } else { +#pragma unroll + for (int i = 0; i < 8; i++) + fd_acc[((size_t)sp * NQ + hh) * HD + dim0 + i] = a[i] * NRM; + } + } + if (S > 1 && lid == 0) { +#pragma unroll + for (int j = 0; j < JG; j++) { + int hh = hk * GQA + hg * JG + j; + fd_ml[((size_t)sp * NQ + hh) * 2] = mmax[j]; + fd_ml[((size_t)sp * NQ + hh) * 2 + 1] = lsum[j]; + } + } + }); + }); + if (S > 1) attn_fd_merge(q, fd_acc, fd_ml, out, S); + return; + } if (!fold_disabled) { constexpr int NKV = QX_N_HEAD_KV; constexpr int NSG = HD / VBLK; /* SGs per WG */