Commit 1aa2954bd for llama.cpp
commit 1aa2954bde90b1cb4d2dca96f90b07d7b155124b
Author: Anant Shrivastava <anant@anantshri.info>
Date: Mon Sep 21 13:37:04 2026 +0530
sycl : coalesce MKL-FA softmax loads instead of one work-item per row (#28918)
* sycl : coalesce MKL-FA softmax loads instead of one work-item per row
* better human readable variable name
diff --git a/ggml/src/ggml-sycl/fattn-mkl.cpp b/ggml/src/ggml-sycl/fattn-mkl.cpp
index 2d164a084..30947b17b 100644
--- a/ggml/src/ggml-sycl/fattn-mkl.cpp
+++ b/ggml/src/ggml-sycl/fattn-mkl.cpp
@@ -110,8 +110,15 @@ static void mkl_fa_init_softmax_state(
// The tile spans absolute rows [q0, q0 + q_rows). Score buffers
// (KQ_f32/S_f16) are indexed RELATIVE to the tile; the persistent state
// (VKQ_accum/KQ_max/KQ_sum) and mask are indexed by ABSOLUTE row.
-// For each row: find local max → rescale previous VKQ_accum →
-// compute exp(s - max) → write S_f16 → update running max/sum.
+// One WORK-GROUP per query row (local size = wg_size): work-items stride
+// over the chunk so adjacent items touch adjacent elements (coalesced),
+// the row max/sum come from group reductions, and the DV-long VKQ
+// rescale is spread across the items. Item 0 is the sole writer of
+// KQ_max/KQ_sum; its writes are ordered after every other item's reads
+// by the second group reduction (a collective). Per-element math is
+// identical to the original one-item-per-row kernel: softcap before
+// mask, native::exp, -1e30 sentinel, half-precision S. Only the float
+// summation order differs (tree vs serial), i.e. last-ulp level.
static void mkl_fa_online_softmax_chunk(
dpct::queue_ptr stream,
float * __restrict KQ_f32,
@@ -126,25 +133,27 @@ static void mkl_fa_online_softmax_chunk(
int64_t mask_row_stride, int mask_n_heads,
float logit_softcap, int64_t wg_size) {
- const int64_t wg = ((q_rows + wg_size - 1) / wg_size) * wg_size;
-
+ // One work-group per query row: exactly q_rows groups of wg_size
+ // items. q_rows * wg_size is already a multiple of wg_size, so unlike
+ // the one-item-per-row kernels there is no round-up / tail guard.
+ const int64_t wg = q_rows * wg_size;
+ const int local_size = (int) wg_size; // stride in the loops below
stream->submit([&](sycl::handler & cgh) {
cgh.parallel_for(sycl::nd_range<1>(wg, wg_size),
[=](sycl::nd_item<1> item) {
- int jc_rel = item.get_global_id(0);
- if (jc_rel >= q_rows) return;
- int jc_abs = q0 + jc_rel;
-
+ const int local_id = (int)item.get_local_id(0);
+ const int row = (int)item.get_group(0); // tile-relative
+ const int jc_abs = q0 + row;
const int gqa_group = jc_abs / n_queries;
const int q_row = jc_abs % n_queries;
-
// Score buffers are tile-local (relative index).
const float * __restrict KQ_row = KQ_f32
- + jc_rel * (int64_t)chunk_size;
+ + row * (int64_t)chunk_size;
+ sycl::half * __restrict S_row = S_f16
+ + row * (int64_t)chunk_size;
// Persistent accumulator is full-sized (absolute index).
float * __restrict vkq = VKQ_accum
+ jc_abs * (int64_t)DV;
-
const sycl::half * mask_h = nullptr;
int64_t m_stride = 0;
if (mask_data) {
@@ -153,10 +162,8 @@ static void mkl_fa_online_softmax_chunk(
mask_h = mask_data + (int64_t)m_head * mask_head_stride;
m_stride = mask_row_stride;
}
-
- // Row-wise local maximum (softcap before mask)
- float local_max = -1e30f;
- for (int i = 0; i < chunk_size; i++) {
+ // Score at chunk offset i — original per-element math.
+ auto score = [&](int i) {
float s = KQ_row[i];
if (logit_softcap != 0.0f) {
s = logit_softcap * sycl::tanh(s);
@@ -165,40 +172,38 @@ static void mkl_fa_online_softmax_chunk(
s += (float)mask_h[q_row * m_stride
+ (chunk_start + i)];
}
+ return s;
+ };
+ // Pass 1: strided (coalesced) row-wise local maximum.
+ float local_max = -1e30f;
+ for (int i = local_id; i < chunk_size; i += local_size) {
+ float s = score(i);
if (s > local_max) local_max = s;
}
-
+ const float final_local_max = sycl::reduce_over_group(
+ item.get_group(), local_max, sycl::maximum<float>());
// Rescale previous accumulator by exp(old_max - new_max)
float old_max = KQ_max[jc_abs];
- float new_max = (old_max > local_max) ? old_max : local_max;
+ float new_max = (old_max > final_local_max) ? old_max : final_local_max;
float rescale = (old_max < -1e29f) ? 1.0f
: sycl::native::exp(old_max - new_max);
-
- for (int v = 0; v < DV; v++) {
+ for (int v = local_id; v < DV; v += local_size) {
vkq[v] *= rescale;
}
-
- // Softmax and write S_f16 (tile-local index)
+ // Pass 2: softmax numerators, strided; S row written once.
float local_sum = 0.0f;
- sycl::half * __restrict S_row = S_f16
- + jc_rel * (int64_t)chunk_size;
-
- for (int i = 0; i < chunk_size; i++) {
- float s = KQ_row[i];
- if (logit_softcap != 0.0f) {
- s = logit_softcap * sycl::tanh(s);
- }
- if (mask_h) {
- s += (float)mask_h[q_row * m_stride
- + (chunk_start + i)];
- }
+ for (int i = local_id; i < chunk_size; i += local_size) {
+ float s = score(i);
float val = sycl::native::exp(s - new_max);
S_row[i] = sycl::half(val);
local_sum += val;
}
-
- KQ_sum[jc_abs] = KQ_sum[jc_abs] * rescale + local_sum;
- KQ_max[jc_abs] = new_max;
+ const float total_sum = sycl::reduce_over_group(
+ item.get_group(), local_sum, sycl::plus<float>());
+ if (local_id == 0) {
+ KQ_sum[jc_abs] = KQ_sum[jc_abs] * rescale + total_sum;
+ KQ_max[jc_abs] = new_max;
+ }
});
});
}