#include <array>
#include <cstdint>
#include <cstdio>
#include <cstring>
#include <string>
#include <optional>
#include <unordered_map>
#include <vector>
#include "fattn-onednn.hpp"
#include "fattn-tile.hpp"
#include "convert.hpp"
#define GGML_SYCL_FA_ONEDNN_MIN_Q 32
bool ggml_sycl_fattn_onednn_binds_kv(const ggml_tensor * K, const ggml_tensor * V) {
if (K->type != GGML_TYPE_F16 || V->type != GGML_TYPE_F16) {
return false;
}
auto bindable = [](const ggml_tensor * t) {
return t->nb[0] == sizeof(sycl::half) && t->nb[1] % sizeof(sycl::half) == 0 &&
t->nb[2] % sizeof(sycl::half) == 0 && t->nb[3] % sizeof(sycl::half) == 0;
};
return bindable(K) && bindable(V);
}
bool ggml_sycl_flash_attn_ext_onednn_supported(const ggml_tensor * dst, bool use_shape_limit) {
#if !GGML_SYCL_DNNL
GGML_UNUSED(dst);
GGML_UNUSED(use_shape_limit);
return false;
#else
if (!g_ggml_sycl_fa_onednn) {
return false;
}
const ggml_tensor * Q = dst->src[0];
const ggml_tensor * K = dst->src[1];
const ggml_tensor * V = dst->src[2];
const ggml_tensor * mask = dst->src[3];
const ggml_tensor * sinks = dst->src[4];
if (K->type != GGML_TYPE_F16 || V->type != GGML_TYPE_F16) {
auto kt = K->type, vt = V->type;
bool k_ok = kt == GGML_TYPE_F32 || kt == GGML_TYPE_Q4_0 || kt == GGML_TYPE_Q4_1 ||
kt == GGML_TYPE_Q5_0 || kt == GGML_TYPE_Q5_1 || kt == GGML_TYPE_Q8_0;
bool v_ok = vt == GGML_TYPE_F32 || vt == GGML_TYPE_Q4_0 || vt == GGML_TYPE_Q4_1 ||
vt == GGML_TYPE_Q5_0 || vt == GGML_TYPE_Q5_1 || vt == GGML_TYPE_Q8_0;
if (!k_ok || !v_ok) {
return false;
}
if (use_shape_limit && (Q->ne[1] < 32 || K->ne[1] < 1024)) {
return false;
}
for (const ggml_tensor * t : {K, V}) {
if (t->type == GGML_TYPE_F16 && t->nb[1] % (t->ne[0] * 2) != 0) {
return false;
}
}
}
const gpu_arch arch = ggml_sycl_info().devices[ggml_sycl_get_device()].hw_info.arch;
bool support_spda = !(arch == gpu_arch::intel_gpu_dg2_g10 ||
arch == gpu_arch::intel_gpu_dg2_g11 ||
arch == gpu_arch::intel_gpu_dg2_g12);
if (!support_spda && K->ne[0] == 64) {
return false;
}
if (g_ggml_sycl_fa_onednn_max_kv > 0 && K->ne[1] > g_ggml_sycl_fa_onednn_max_kv) {
return false;
}
if (!mask || mask->type != GGML_TYPE_F16 || mask->ne[2] != 1 || mask->ne[3] != 1 || sinks) {
return false;
}
float max_bias = 0.0f, logit_softcap = 0.0f;
memcpy(&max_bias, (const float *) dst->op_params + 1, sizeof(float));
memcpy(&logit_softcap, (const float *) dst->op_params + 2, sizeof(float));
if (max_bias != 0.0f || logit_softcap != 0.0f) {
return false;
}
const int64_t d = K->ne[0];
if (V->ne[0] != d || Q->ne[3] != 1) {
return false;
}
if (K->ne[2] == 0 || Q->ne[2] % K->ne[2] != 0) {
return false;
}
if (use_shape_limit && Q->ne[1] < GGML_SYCL_FA_ONEDNN_MIN_Q) {
return false;
}
return true;
#endif
}
#if GGML_SYCL_DNNL
#include "dnnl.hpp"
#include "dnnl_sycl.hpp"
#include "oneapi/dnnl/dnnl_graph.hpp"
using namespace dnnl;
using namespace dnnl::graph;
template <typename src_t>
static void cont_to_f16_sycl(const char * src, sycl::half * dst,
int64_t ne0, int64_t ne1, int64_t ne2, int64_t ne3,
size_t nb1, size_t nb2, size_t nb3, dpct::queue_ptr stream) {
const int64_t n = ne0 * ne1 * ne2 * ne3;
stream->parallel_for(sycl::range<1>(n), [=](sycl::id<1> ix) {
const int64_t gid = ix[0];
int64_t i = gid;
const int64_t i0 = i % ne0; i /= ne0;
const int64_t i1 = i % ne1; i /= ne1;
const int64_t i2 = i % ne2; const int64_t i3 = i / ne2;
const src_t * p = (const src_t *) (src + i1 * nb1 + i2 * nb2 + i3 * nb3) + i0;
dst[gid] = (sycl::half) (*p);
});
}
static void permute_sdpa_out_sycl(const sycl::half * out, float * dst,
int64_t mb, int64_t H, int64_t q, int64_t d, dpct::queue_ptr stream) {
const int64_t n = mb * H * q * d;
stream->parallel_for(sycl::range<1>(n), [=](sycl::id<1> ix) {
const int64_t gid = ix[0];
int64_t i = gid;
const int64_t e = i % d; i /= d;
const int64_t t = i % q; i /= q;
const int64_t h = i % H; const int64_t b = i / H;
dst[e + h * d + t * d * H + b * d * H * q] = (float) out[gid];
});
}
struct sdpa_partition {
compiled_partition cp;
std::vector<logical_tensor> ins;
logical_tensor out;
size_t id_q = 0, id_k = 0, id_v = 0, id_scale = 0, id_mask = 0;
bool ok = false;
};
static sdpa_partition build_sdpa(const engine & eng, int H, int Hkv, int q, int seq, int d,
const std::array<int64_t, 5> & k_str, const std::array<int64_t, 5> & v_str) try {
using ltype = logical_tensor::layout_type;
using dt = logical_tensor::data_type;
using ldims = logical_tensor::dims;
const dt fi = dt::f32, t = dt::f16;
const int rep = H / Hkv;
const ldims q_sz = {1, Hkv, rep, q, d}, kv_sz = {1, Hkv, 1, seq, d}, s_sz = {1, Hkv, rep, q, seq},
sc = {1, 1, 1, 1, 1}, msk = {1, 1, 1, q, seq}, o_sz = {1, Hkv, rep, q, d};
const ldims k_st(k_str.begin(), k_str.end()), v_st(v_str.begin(), v_str.end());
int64_t id = 0;
sdpa_partition E;
auto query = logical_tensor(id++, t, q_sz, ltype::strided);
auto key = logical_tensor(id++, t, kv_sz, k_st);
auto score = logical_tensor(id++, fi, s_sz, ltype::strided);
auto bmm1 = op(id++, op::kind::MatMul, "bmm1");
bmm1.set_attr<bool>(op::attr::transpose_b, true); bmm1.add_inputs({query, key}); bmm1.add_outputs({score});
auto scale = logical_tensor(id++, t, sc, ltype::strided);
auto scaled = logical_tensor(id++, fi, s_sz, ltype::strided);
auto sdiv = op(id++, op::kind::Divide, "scale_div"); sdiv.add_inputs({score, scale}); sdiv.add_outputs({scaled});
auto mask = logical_tensor(id++, t, msk, ltype::strided);
auto masked = logical_tensor(id++, fi, s_sz, ltype::strided);
auto madd = op(id++, op::kind::Add, "mask_add");
madd.add_inputs({scaled, mask}); madd.add_outputs({masked});
auto probs = logical_tensor(id++, t, s_sz, ltype::strided);
auto smax = op(id++, op::kind::SoftMax, "softmax");
smax.set_attr<int64_t>(op::attr::axis, -1);
smax.set_attr<std::string>(op::attr::mode, "inf_as_zero");
smax.add_inputs({masked}); smax.add_outputs({probs});
auto value = logical_tensor(id++, t, kv_sz, v_st);
auto output = logical_tensor(id++, t, o_sz, ltype::strided); auto bmm2 = op(id++, op::kind::MatMul, "bmm2");
bmm2.add_inputs({probs, value}); bmm2.add_outputs({output});
dnnl::graph::graph g(eng.get_kind());
g.add_op(bmm1); g.add_op(sdiv); g.add_op(madd); g.add_op(smax); g.add_op(bmm2);
g.finalize();
auto parts = g.get_partitions();
if (parts.size() != 1 || !parts[0].is_supported()) {
GGML_LOG_WARN("%s: oneDNN did not fuse the SDPA graph; falling back to TILE kernel\n", __func__);
return E; }
E.ins = parts[0].get_input_ports();
E.out = parts[0].get_output_ports()[0];
E.cp = parts[0].compile(E.ins, {E.out}, eng);
E.out = E.cp.query_logical_tensor(E.out.get_id());
E.id_q = query.get_id(); E.id_k = key.get_id(); E.id_v = value.get_id();
E.id_scale = scale.get_id(); E.id_mask = mask.get_id();
E.ok = true;
return E;
}
catch (const std::exception & e) {
GGML_LOG_WARN("%s: oneDNN SDPA partition build failed (%s); falling back to TILE kernel\n", __func__, e.what());
return {};
}
void ggml_sycl_flash_attn_ext_onednn(ggml_backend_sycl_context & ctx, ggml_tensor * dst) try {
const ggml_tensor * Q = dst->src[0];
const ggml_tensor * K = dst->src[1];
const ggml_tensor * V = dst->src[2];
const ggml_tensor * mask = dst->src[3];
const int64_t d = K->ne[0]; const int64_t seq = K->ne[1]; const int64_t Hkv = K->ne[2]; const int64_t H = Q->ne[2]; const int64_t q = Q->ne[1]; const int64_t mb = Q->ne[3];
float kq_scale = 1.0f;
memcpy(&kq_scale, (const float *) dst->op_params + 0, sizeof(float));
dpct::queue_ptr stream = ctx.stream();
dnnl::engine eng = ctx.engine_dnnl(stream);
dnnl::stream strm = ctx.stream_dnnl(stream);
const ggml_sycl_fattn_extra extra = ggml_sycl_fattn_get_extra(dst);
std::optional<ggml_sycl_pool_alloc<sycl::half>> Qf_pool;
sycl::half * Qf_ptr = (sycl::half *) extra.Q_buffer_ptr;
if (!Qf_ptr) {
Qf_pool.emplace(ctx.pool(), (size_t) H * q * d);
Qf_ptr = Qf_pool->get();
}
cont_to_f16_sycl<float>((const char *) Q->data, Qf_ptr, d, q, H, mb, Q->nb[1], Q->nb[2], Q->nb[3], stream);
sycl::half * K_ptr = nullptr;
sycl::half * V_ptr = nullptr;
std::array<int64_t, 5> k_str{ Hkv * seq * d, seq * d, seq * d, d, 1 };
std::array<int64_t, 5> v_str = k_str;
std::optional<ggml_sycl_pool_alloc<sycl::half>> Kf_pool;
std::optional<ggml_sycl_pool_alloc<sycl::half>> Vf_pool;
auto stage_k = [&](size_t n) { if (extra.K_buffer_ptr) { return (sycl::half *) extra.K_buffer_ptr; }
Kf_pool.emplace(ctx.pool(), n); return Kf_pool->get(); };
auto stage_v = [&](size_t n) { if (extra.V_buffer_ptr) { return (sycl::half *) extra.V_buffer_ptr; }
Vf_pool.emplace(ctx.pool(), n); return Vf_pool->get(); };
auto elem_strides = [](const ggml_tensor * t) {
const int64_t s1 = (int64_t) (t->nb[1] / t->nb[0]);
const int64_t s2 = (int64_t) (t->nb[2] / t->nb[0]);
const int64_t s3 = (int64_t) (t->nb[3] / t->nb[0]);
return std::array<int64_t, 5>{ s3, s2, s2, s1, 1 };
};
if (ggml_sycl_fattn_onednn_binds_kv(K, V)) {
K_ptr = (sycl::half *) K->data;
V_ptr = (sycl::half *) V->data;
k_str = elem_strides(K);
v_str = elem_strides(V);
} else if (K->type == GGML_TYPE_F16 && V->type == GGML_TYPE_F16) {
K_ptr = stage_k((size_t) Hkv * seq * d);
V_ptr = stage_v((size_t) Hkv * seq * d);
cont_to_f16_sycl<sycl::half>((const char *) K->data, K_ptr, d, seq, Hkv, mb, K->nb[1], K->nb[2], K->nb[3], stream);
cont_to_f16_sycl<sycl::half>((const char *) V->data, V_ptr, d, seq, Hkv, mb, V->nb[1], V->nb[2], V->nb[3], stream);
} else if (ggml_is_quantized(K->type)) {
K_ptr = stage_k((size_t) ggml_nelements(K));
{
const char * K_data = (const char *)K->data;
const bool k_non_dense = ((int64_t)K->ne[1] * K->nb[1] != K->nb[2]) && K->ne[2] > 1;
const bool k_gemma = k_non_dense &&
((int64_t)K->nb[2] < (int64_t)K->ne[1] * (int64_t)K->nb[1]);
if (ggml_is_contiguously_allocated(K) && !k_non_dense) {
to_fp16_sycl_t to_fp16 = ggml_get_to_fp16_sycl(K->type, dst);
to_fp16(K_data, K_ptr, ggml_nelements(K), stream);
} else {
const size_t bs = ggml_blck_size(K->type);
const size_t ts = ggml_type_size(K->type);
to_fp16_nc_sycl_t to_fp16 = ggml_get_to_fp16_nc_sycl(K->type);
int64_t s01, s02, s03;
if (k_gemma) {
const int64_t blk_per_row = (int64_t)K->ne[0] / bs;
s01 = (int64_t)Hkv * blk_per_row;
s02 = blk_per_row;
s03 = (int64_t)K->ne[1] * s01;
} else {
s01 = (int64_t)K->nb[1] / ts;
s02 = (int64_t)K->nb[2] / ts;
s03 = (int64_t)K->nb[3] / ts;
}
to_fp16(K_data, K_ptr,
K->ne[0], K->ne[1], K->ne[2], K->ne[3],
s01, s02, s03, stream);
}
}
V_ptr = stage_v((size_t) ggml_nelements(V));
{
const char * V_data = (const char *)V->data;
const bool v_non_dense = ((int64_t)V->ne[1] * V->nb[1] != V->nb[2]) && V->ne[2] > 1;
const bool v_gemma = v_non_dense &&
((int64_t)V->nb[2] < (int64_t)V->ne[1] * (int64_t)V->nb[1]);
if (ggml_is_contiguously_allocated(V) && !v_non_dense) {
to_fp16_sycl_t to_fp16 = ggml_get_to_fp16_sycl(V->type, dst);
to_fp16(V_data, V_ptr, ggml_nelements(V), stream);
} else {
const size_t bs = ggml_blck_size(V->type);
const size_t ts = ggml_type_size(V->type);
to_fp16_nc_sycl_t to_fp16 = ggml_get_to_fp16_nc_sycl(V->type);
int64_t s01, s02, s03;
if (v_gemma) {
const int64_t blk_per_row = (int64_t)V->ne[0] / bs;
s01 = (int64_t)V->ne[2] * blk_per_row;
s02 = blk_per_row;
s03 = (int64_t)V->ne[1] * s01;
} else {
s01 = (int64_t)V->nb[1] / ts;
s02 = (int64_t)V->nb[2] / ts;
s03 = (int64_t)V->nb[3] / ts;
}
to_fp16(V_data, V_ptr,
V->ne[0], V->ne[1], V->ne[2], V->ne[3],
s01, s02, s03, stream);
}
}
} else {
K_ptr = stage_k((size_t) ggml_nelements(K));
cont_to_f16_sycl<float>((const char *) K->data, K_ptr, K->ne[0], K->ne[1], K->ne[2], K->ne[3],
K->nb[1], K->nb[2], K->nb[3], stream);
V_ptr = stage_v((size_t) ggml_nelements(V));
cont_to_f16_sycl<float>((const char *) V->data, V_ptr, V->ne[0], V->ne[1], V->ne[2], V->ne[3],
V->nb[1], V->nb[2], V->nb[3], stream);
}
const sycl::half scale_h = (sycl::half) (1.0f / kq_scale);
std::optional<ggml_sycl_pool_alloc<sycl::half>> scbuf;
sycl::half * scale_dev = (sycl::half *) extra.scale_buffer_ptr;
if (!scale_dev) {
scbuf.emplace(ctx.pool(), 1);
scale_dev = scbuf->get();
}
stream->single_task([=]() { *scale_dev = scale_h; });
std::optional<ggml_sycl_pool_alloc<sycl::half>> outf_pool;
sycl::half * outf_ptr = (sycl::half *) extra.out_buffer_ptr;
if (!outf_ptr) {
outf_pool.emplace(ctx.pool(), (size_t) H * q * d);
outf_ptr = outf_pool->get();
}
static std::unordered_map<std::string, sdpa_partition> cache;
char keyb[256];
snprintf(keyb, sizeof(keyb), "%d:%lld:%lld:%lld:%lld:%lld:%lld:%lld:%lld:%lld:%lld:%lld", ggml_sycl_get_device(),
(long long) H, (long long) Hkv, (long long) q, (long long) seq, (long long) d,
(long long) k_str[0], (long long) k_str[1], (long long) k_str[3],
(long long) v_str[0], (long long) v_str[1], (long long) v_str[3]);
auto it = cache.find(keyb);
if (it == cache.end()) {
it = cache.emplace(keyb, build_sdpa(eng, (int) H, (int) Hkv, (int) q, (int) seq, (int) d, k_str, v_str)).first;
}
sdpa_partition & E = it->second;
if (!E.ok) {
ggml_sycl_flash_attn_ext_tile(ctx, dst);
return;
}
auto id2ptr = [&](size_t r) -> void * {
if (r == E.id_q) return Qf_ptr;
if (r == E.id_k) return K_ptr;
if (r == E.id_v) return V_ptr;
if (r == E.id_scale) return scale_dev;
if (r == E.id_mask) return (void *) mask->data;
return nullptr;
};
std::vector<tensor> ti;
ti.reserve(E.ins.size());
for (auto & lt : E.ins) {
ti.emplace_back(lt, eng, id2ptr(lt.get_id()));
}
tensor to(E.out, eng, outf_ptr);
E.cp.execute(strm, ti, {to});
permute_sdpa_out_sycl(outf_ptr, (float *) dst->data, mb, H, q, d, stream);
if (ggml_sycl_info().device_count > 1) {
stream->wait_and_throw();
}
}
catch (const std::exception & e) {
GGML_LOG_WARN("%s: oneDNN SDPA failed (%s); falling back to TILE kernel\n", __func__, e.what());
ggml_sycl_flash_attn_ext_tile(ctx, dst);
}
#endif