mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-19 17:24:57 +02:00
ggml-metal: hoist QK^T + softmax out of the d-block loop in the FA tensor kernel
For DV > 128 (NBLK > 1), the tensor kernel previously ran the entire QK^T + online-softmax pass once per d block (NBLKx the QK work). This restructures the NBLK == 2 case (dv = 256) so that QK^T + softmax run once per KV chunk and the P tile (f32, in registers) is reused for both PV d blocks (two live PV destination tiles). NBLK == 1 and NBLK == 4 (dv = 512) keep the per-block form: for NBLK == 4, carrying 4 live PV destination tiles would overflow the register file. 256/256, nq=512, kv=20000 (f16, mask): 43.4 ms -> 31.1 ms (0.71x -> 0.83x vs the vec kernel). Full test-backend-ops FLASH_ATTN_EXT suite passes (tensor on and off).
This commit is contained in:
@@ -2417,20 +2417,264 @@ void kernel_flash_attn_ext_tensor_impl(
|
||||
const uint64_t pad_mask_offs = (uint64_t) (args.nb11 + args.nb21) * C * args.ne_12_2 * args.ne_12_3 +
|
||||
2u * C * args.ne31 * ((uint64_t) (iq2 % args.ne32) + (uint64_t) (iq3 % args.ne33) * args.ne32);
|
||||
|
||||
for (int b = 0; b < NBLK; ++b) {
|
||||
auto make_cT_pv = [&]() {
|
||||
auto tVb = tensor(vp, dextents<int, 2>(PVM, C), array<int, 2>{ 1, sv });
|
||||
return mm_pvb.template get_destination_cooperative_tensor<decltype(tVb), decltype(cT_qk0), float>();
|
||||
};
|
||||
auto make_cT_pv = [&]() {
|
||||
auto tVb = tensor(vp, dextents<int, 2>(PVM, C), array<int, 2>{ 1, sv });
|
||||
return mm_pvb.template get_destination_cooperative_tensor<decltype(tVb), decltype(cT_qk0), float>();
|
||||
};
|
||||
|
||||
auto cT_pv = make_cT_pv();
|
||||
if constexpr (NBLK == 2) {
|
||||
// hoisted form: QK^T + online softmax run once per chunk; the P tile
|
||||
// (f32, in registers) is reused for both PV d blocks. (The per-block
|
||||
// form below recomputes QK^T + softmax per d block: NBLKx the QK work.)
|
||||
auto cT_pv0 = make_cT_pv();
|
||||
auto cT_pv1 = make_cT_pv();
|
||||
|
||||
{
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0; i < cT_pv.get_capacity(); ++i) {
|
||||
if (cT_pv.is_valid_element(i)) { cT_pv[i] = 0.0f; }
|
||||
for (uint i = 0; i < cT_pv0.get_capacity(); ++i) {
|
||||
if (cT_pv0.is_valid_element(i)) { cT_pv0[i] = 0.0f; }
|
||||
}
|
||||
}
|
||||
{
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0; i < cT_pv1.get_capacity(); ++i) {
|
||||
if (cT_pv1.is_valid_element(i)) { cT_pv1[i] = 0.0f; }
|
||||
}
|
||||
}
|
||||
for (int j = t; j < QPSG; j += 32) {
|
||||
sh_M[jb + j] = -FLT_MAX / 2;
|
||||
sh_S[jb + j] = 0.0f;
|
||||
sh_alpha[jb + j] = 1.0f;
|
||||
}
|
||||
simdgroup_barrier(mem_flags::mem_threadgroup);
|
||||
|
||||
for (int ic = 0; ic < nchunks; ++ic) {
|
||||
const int k0 = ic*C;
|
||||
const int kc = min((int) C, kv - k0);
|
||||
|
||||
// reset this thread's (q, t) column
|
||||
{
|
||||
#pragma clang loop unroll(full)
|
||||
for (int j = 0; j < QPSG; ++j) {
|
||||
sh_qmax[jqb + j*32 + t] = -FLT_MAX / 2;
|
||||
sh_qsum[jqb + j*32 + t] = 0.0f;
|
||||
}
|
||||
}
|
||||
|
||||
// the last partial chunk is read from the pad buffer (padded with 0);
|
||||
// the pad buffer holds only the last C kv items per (KV head, batch),
|
||||
// so its rows are indexed from 0 (not k0)
|
||||
const bool use_pad = has_kvpad && k0 + C > kv;
|
||||
device half * kp_c = use_pad ? (device half *) (pad + pad_k_offs) : kp + (uint) k0*sk;
|
||||
device half * vp_c = use_pad ? (device half *) (pad + pad_v_offs) : vp + (uint) k0*sv;
|
||||
|
||||
// mask row base for this chunk: real mask (kv contiguous, query
|
||||
// stride nb31) or the pad section (C per query, indexed from 0)
|
||||
device const half * mp = nullptr;
|
||||
int mstride = 0; // in halfs, per local query j
|
||||
if (has_mask) {
|
||||
if (use_pad) {
|
||||
mp = (device const half *) (pad + pad_mask_offs) + (iq1 + sgitg*QPSG) * (int) C;
|
||||
mstride = C;
|
||||
} else {
|
||||
// global query index: iq1 + sgitg*QPSG + j (j is local, 0..QPSG-1)
|
||||
mp = (device const half *) (mask + (uint64_t) k0 * 2 + mbase)
|
||||
+ (uint64_t) (iq1 + sgitg*QPSG) * (args.nb31 / 2);
|
||||
mstride = (int) (args.nb31 / 2);
|
||||
}
|
||||
}
|
||||
|
||||
auto tK = tensor(kp_c, dextents<int, 2>(DK, C), array<int, 2>{ 1, sk });
|
||||
// d blocks: block 0 covers d = 0..PVM-1, block 1 covers d = PVM..2*PVM-1
|
||||
auto tV0 = tensor(vp_c, dextents<int, 2>(PVM, C), array<int, 2>{ 1, sv });
|
||||
auto tV1 = tensor(vp_c + PVM, dextents<int, 2>(PVM, C), array<int, 2>{ 1, sv });
|
||||
|
||||
// ---- QK^T ----
|
||||
auto cT_qk = mm_qk.template get_destination_cooperative_tensor<decltype(tK), decltype(tQ), float>();
|
||||
{
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0; i < cT_qk.get_capacity(); ++i) {
|
||||
if (cT_qk.is_valid_element(i)) { cT_qk[i] = 0.0f; }
|
||||
}
|
||||
}
|
||||
mm_qk.run(tK, tQ, cT_qk);
|
||||
|
||||
// scale in registers; clobber padded kv (idx1 >= kc) to -inf
|
||||
float lmax[QPSG], lsum[QPSG];
|
||||
#pragma clang loop unroll(full)
|
||||
for (int j = 0; j < QPSG; ++j) { lmax[j] = -FLT_MAX / 2; lsum[j] = 0.0f; }
|
||||
{
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0; i < cT_qk.get_capacity(); ++i) {
|
||||
if (!cT_qk.is_valid_element(i)) { continue; }
|
||||
auto idx = cT_qk.get_multidimensional_index(i);
|
||||
const int q = (int) idx[0];
|
||||
float s;
|
||||
if ((int) idx[1] >= kc) {
|
||||
s = -FLT_MAX / 2;
|
||||
} else {
|
||||
s = cT_qk[i]*args.scale;
|
||||
if (has_scap) { s = args.logit_softcap * tanh(s); }
|
||||
if (has_mask) { s += (float) mp[(uint) q * mstride + (uint) idx[1]] * mscale; }
|
||||
}
|
||||
cT_qk[i] = s;
|
||||
if (s > lmax[q]) { lmax[q] = s; }
|
||||
}
|
||||
}
|
||||
|
||||
// partial max -> shared (non-owned queries hold -FLT_MAX/2 -> harmless)
|
||||
{
|
||||
#pragma clang loop unroll(full)
|
||||
for (int j = 0; j < QPSG; ++j) {
|
||||
sh_qmax[jqb + j*32 + t] = lmax[j];
|
||||
}
|
||||
}
|
||||
simdgroup_barrier(mem_flags::mem_threadgroup);
|
||||
|
||||
// M / alpha update: thread t (< QPSG) finalizes query t
|
||||
if (t < QPSG) {
|
||||
const int j = t;
|
||||
float m_new = sh_M[jb + j];
|
||||
#pragma clang loop unroll(full)
|
||||
for (int tt = 0; tt < 32; ++tt) {
|
||||
m_new = max(m_new, sh_qmax[jqb + j*32 + tt]);
|
||||
}
|
||||
sh_alpha[jb + j] = exp(sh_M[jb + j] - m_new);
|
||||
sh_M[jb + j] = m_new;
|
||||
}
|
||||
simdgroup_barrier(mem_flags::mem_threadgroup);
|
||||
|
||||
// ---- exp in registers, partial sums ----
|
||||
{
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0; i < cT_qk.get_capacity(); ++i) {
|
||||
if (!cT_qk.is_valid_element(i)) { continue; }
|
||||
auto idx = cT_qk.get_multidimensional_index(i);
|
||||
const int q = (int) idx[0];
|
||||
const float p = exp(cT_qk[i] - sh_M[jb + q]);
|
||||
cT_qk[i] = p; // P in registers (f32), later used as PV right input
|
||||
lsum[q] += p;
|
||||
}
|
||||
}
|
||||
{
|
||||
#pragma clang loop unroll(full)
|
||||
for (int j = 0; j < QPSG; ++j) {
|
||||
sh_qsum[jqb + j*32 + t] = lsum[j];
|
||||
}
|
||||
}
|
||||
simdgroup_barrier(mem_flags::mem_threadgroup);
|
||||
|
||||
// S update: thread t (< QPSG) finalizes query t
|
||||
if (t < QPSG) {
|
||||
const int j = t;
|
||||
float s_new = 0.0f;
|
||||
#pragma clang loop unroll(full)
|
||||
for (int tt = 0; tt < 32; ++tt) {
|
||||
s_new += sh_qsum[jqb + j*32 + tt];
|
||||
}
|
||||
sh_S[jb + j] = sh_S[jb + j]*sh_alpha[jb + j] + s_new;
|
||||
}
|
||||
simdgroup_barrier(mem_flags::mem_threadgroup);
|
||||
|
||||
// ---- rescale both O accumulators (per element: query = idx0) ----
|
||||
// NOTE: skip the first chunk: alpha_1 = exp(-FLT_MAX/2 - m_new) = 0
|
||||
if (ic > 0) {
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0; i < cT_pv0.get_capacity(); ++i) {
|
||||
if (!cT_pv0.is_valid_element(i)) { continue; }
|
||||
auto idx = cT_pv0.get_multidimensional_index(i);
|
||||
cT_pv0[i] *= sh_alpha[jb + (int) idx[0]];
|
||||
}
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0; i < cT_pv1.get_capacity(); ++i) {
|
||||
if (!cT_pv1.is_valid_element(i)) { continue; }
|
||||
auto idx = cT_pv1.get_multidimensional_index(i);
|
||||
cT_pv1[i] *= sh_alpha[jb + (int) idx[0]];
|
||||
}
|
||||
}
|
||||
|
||||
// ---- PV: P (in registers) as right input, both d blocks ----
|
||||
{
|
||||
auto cT_pr = mm_pvb.template get_right_input_cooperative_tensor<half, float, float>(cT_qk);
|
||||
mm_pvb.run(tV0, cT_pr, cT_pv0);
|
||||
mm_pvb.run(tV1, cT_pr, cT_pv1);
|
||||
}
|
||||
simdgroup_barrier(mem_flags::mem_threadgroup);
|
||||
}
|
||||
|
||||
// ---- final: sinks, O /= S, element-wise write for both d blocks ----
|
||||
{
|
||||
simdgroup_barrier(mem_flags::mem_threadgroup);
|
||||
|
||||
// sinks: a virtual sink position with score sinks[iq2] and no O
|
||||
// contribution: M' = max(M, sink), S' = S*exp(M - M') + exp(sink - M'),
|
||||
// O' = O*exp(M - M'). The per-query O factor is stashed in sh_alpha
|
||||
// (reused; the last-chunk alpha is not needed at this point).
|
||||
if (has_sinks) {
|
||||
const float s_sink = ((device const float *) sinks)[iq2];
|
||||
for (int j = t; j < QPSG; j += 32) {
|
||||
const float m = sh_M[jb + j];
|
||||
const float m2 = max(m, s_sink);
|
||||
sh_alpha[jb + j] = exp(m - m2);
|
||||
sh_S[jb + j] = sh_S[jb + j]*sh_alpha[jb + j] + exp(s_sink - m2);
|
||||
}
|
||||
} else {
|
||||
for (int j = t; j < QPSG; j += 32) {
|
||||
sh_alpha[jb + j] = 1.0f;
|
||||
}
|
||||
}
|
||||
simdgroup_barrier(mem_flags::mem_threadgroup);
|
||||
|
||||
// output layout is (DV, heads, batch, batch3) with DV innermost:
|
||||
// offset = (iq3*ne2*ne1 + iq2 + (iq1 + sgitg*QPSG + j)*ne1)*DV + d
|
||||
// NOTE: the write must use ELEMENT pointer arithmetic on a float*
|
||||
// (cast dst first, then add element offsets). Doing the offset in
|
||||
// BYTES (char* + element*DV) corrupts the coop tile register layout
|
||||
// of the matmul ops above (verified: QK^T (C,QPSG) tile collapses).
|
||||
device float * op = (device float *) dst +
|
||||
(uint64_t) (iq3*args.ne2*args.ne1 + iq2 + (iq1 + sgitg*QPSG)*args.ne1) * DV;
|
||||
{
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0; i < cT_pv0.get_capacity(); ++i) {
|
||||
if (!cT_pv0.is_valid_element(i)) { continue; }
|
||||
auto idx = cT_pv0.get_multidimensional_index(i);
|
||||
const int j = (int) idx[0]; // query
|
||||
const int d = (int) idx[1]; // head dim (block 0)
|
||||
if (j >= 0 && j < QPSG && d >= 0 && d < (int) PVM && iq1 + j < args.ne01) {
|
||||
const float s = sh_S[jb + j];
|
||||
op[(uint64_t) j*args.ne1*DV + (uint) d] = (s == 0.0f) ? 0.0f : cT_pv0[i]*sh_alpha[jb + j]/s;
|
||||
}
|
||||
}
|
||||
}
|
||||
{
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0; i < cT_pv1.get_capacity(); ++i) {
|
||||
if (!cT_pv1.is_valid_element(i)) { continue; }
|
||||
auto idx = cT_pv1.get_multidimensional_index(i);
|
||||
const int j = (int) idx[0]; // query
|
||||
const int d = PVM + (int) idx[1]; // head dim (block 1)
|
||||
if (j >= 0 && j < QPSG && d >= 0 && d < (int) DV && iq1 + j < args.ne01) {
|
||||
const float s = sh_S[jb + j];
|
||||
op[(uint64_t) j*args.ne1*DV + (uint) d] = (s == 0.0f) ? 0.0f : cT_pv1[i]*sh_alpha[jb + j]/s;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
} else {
|
||||
// NBLK != 2 (NBLK == 1, or NBLK == 4 for dv = 512): per-d-block form.
|
||||
// For NBLK == 4, carrying 4 live PV destination tiles would overflow
|
||||
// the register file, so QK^T + softmax are recomputed per d block
|
||||
// (correct, but NBLKx the QK work).
|
||||
for (int b = 0; b < NBLK; ++b) {
|
||||
auto cT_pv = make_cT_pv();
|
||||
|
||||
{
|
||||
#pragma clang loop unroll(full)
|
||||
for (uint i = 0; i < cT_pv.get_capacity(); ++i) {
|
||||
if (cT_pv.is_valid_element(i)) { cT_pv[i] = 0.0f; }
|
||||
}
|
||||
}
|
||||
for (int j = t; j < QPSG; j += 32) {
|
||||
sh_M[jb + j] = -FLT_MAX / 2;
|
||||
sh_S[jb + j] = 0.0f;
|
||||
@@ -2631,6 +2875,7 @@ void kernel_flash_attn_ext_tensor_impl(
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user