//ggml_sycl_flash_attention_mini.cpp gsflm.cpp from alucian Berlin-Buch 03.08.2026 /// 12:41
//Nur F16/(Nebengewicht F32) GGML/GGUF
//ARC Intel XE+ iGPU+dGPU SPLIT ROW XMX
//scalar_t* Out_row_ptr = Out_ptr + (head_row * out_stride);
range<1>(WG_SIZE)),
#include "stdlib.h"
#include "stdio.h"
#include
#include
#include
#include
#include
#include
#include
#include
#include
#include
#include
#include "ggml-sycl.h"
#include "ggml.h"
#include
#include "ggml-impl.h"
auto Q_ptr
inline sycl::half* get_sycl_ptr(const ggml_tensor* tensor) {
return reinterpret_cast<sycl::half*>(tensor->data);
}
#define XFLOAT float
#define mdlXYZ 1000
#define MEM_ALIGN 128 //SLM
constexpr int BLOCK_M = 128;
constexpr int BLOCK_N = 128;
constexpr int D_MAX = 1024;
constexpr int VEC_SIZE = 16;
/**
* @brief VEKTORISIERTES PUNKT PRODUKT ZWISCHEN "Q[i]" UND "K[j]"
* @tparam scalar_t DATENTYP sycl::half ODER FLIESSWERTE
* @param Q_row_float QUERY ZEILE ALS FLIESSWERTE
* @param K_ptr SCHLUESSELPUNKTE UND ZEIGER
* @param d_k KOPFDIMENSIONEN
* @return PUNKTPRODUKUT ALS FLIESSWERTE
*/
using namespace sycl;
using namespace sycl::ext::oneapi::experimental::matrix;
constexpr int WG_SIZE = 16;
extern "C" void ggml_sycl_flash_attention_dispatch(queue& q, half* Out, half* Q, half* K, half* V, int num_q, int d_k) {
bool has_xmx = q.get_device().has(sycl::aspect::ext_intel_matrix);
if (has_xmx && (d_k % 16 == 0)) {
q.parallel_for(nd_range<1>(range<1>((num_q + 16) / 16 * 32), range<1>(32)),
[=](sycl::nd_item<1> item) [[intel::reqd_sub_group_size(16)]] {
sub_group sg = item.get_sub_group();
joint_matrix<sub_group, half, use::a, 16, 16, layout::row_major> mat_q;
joint_matrix<sub_group, half, use::b, 16, 16, layout::col_major> mat_k;
joint_matrix<sub_group, float, use::accumulator, 16, 16> mat_s;
joint_matrix_fill(sg, mat_s, 0.0f);
joint_matrix_load(sg, mat_q, Q + (item.get_group(0) * 16 * d_k), d_k);
joint_matrix_load(sg, mat_k, K, d_k);
joint_matrix_mad(sg, mat_s, mat_q, mat_k, mat_s);
//joint_matrix_copy(sg, m_p_half, m_s_acc);
joint_matrix_store(sg, mat_s, (float*)Out, d_k, layout::row_major);
}
}
extern "C" void ggml_sycl_flash_attention_dispatch(
ggml_backend_sycl_context* ctx,
ggml_tensor* dst, const ggml_tensor* Q, const ggml_tensor* K, const ggml_tensor* V) {
auto& q = ggml_backend_sycl_get_queue(Q->backend);
auto dev = q.get_device();
bool has_xmx = dev.has(sycl::aspect::ext_intel_matrix);
bool can_use_xmx = (Q->ne[0] % 16 == 0) && (K->ne[1] % 16 == 0);
if (has_xmx && can_use_xmx) {
xmx_kern(q, dst, Q, K, V, S, O);
} else {
ggml_sycl_flash_attention_(q, dst, Q, K, V, S, O);
}
}
template <typename scalar_t>
float dot_product_vec(const float* Q_row_float, const scalar_t* K_ptr, int d_k) {
if constexpr (std::is_same_v<scalar_t, sycl::half>) {
if (d_k % VEC_SIZE != 0) {
float score = 0.0f;
for (int di = 0; di < d_k; ++di) {
score += Q_row_float[di] * static_cast<float>(K_ptr[di]);
}
return score;
}
constexpr int vec_elements = VEC_SIZE;
using vec_half = sycl::vec<sycl::half, vec_elements>;
using vec_float = sycl::vec<float, vec_elements>;
float final_score = 0.0f;
int vec_iters = d_k / vec_elements;
for (int v = 0; v < vec_iters; ++v) {
}
vec_half k_half_vec;
k_half_vec.load(v * vec_elements, K_ptr);
vec_float k_float_vec = k_half_vec.template convert<float>();
vec_float q_float_vec;
q_float_vec.load(v * vec_elements, Q_row_float);
final_score += sycl::dot(q_float_vec, k_float_vec);
}
return final_score;
} else {
float score = 0.0f;
for (int di = 0; di < d_k; ++di) {
score += Q_row_float[di] * K_ptr[di];
}
return score;
}
/**
* @brief HAUPTKERN AUFDROESSELSTRATEGIE MIT TREFFERWERTEZWISCHENSPEICHER
* @tparam scalar_t DATENTYP FORMAT INTEL sycl::half HALBE GENAUIGKEIT F16
*/
template <typename scalar_t>
void flash_attention_kernel_impl(
const scalar_t* Q_ptr,
const scalar_t* K_ptr,
const scalar_t* V_ptr,
const scalar_t* S_ptr,
const scalar_t* O_ptr,
const scalar_t* Out_stride,
int num_q,
int num_k,
int num_v,
int num_s,
int num_o,
int num_out_stride,
int d_q,
int d_k,
int d_v,
int d_s,
int d_o,
int d_out_ptr,
int q_stride,
int k_stride,
int v_stride,
int s_stride,
int o_stride,
int out_stride,
float* S_scores,
[=](sycl::nd_item<1> item) {
const int head_row = item.get_global_id(0);//0
if (head_row >= num_q) return;
float accum_den = 0.0f;
float running_max = -INFINITY;
float accum_num[D_MAX] = {0.0f};
const scalar_t* Q_row_ptr = Q_ptr + head_row * q_stride;
float Q_row_float[D_MAX];
for (int di = 0; di < d_k; ++di) {
Q_row_float[di] = static_cast<float>(Q_row_ptr[di]);
}
const float scale_factor = 1.0f / sycl::sqrt(static_cast<float>(d_k));
for (int k_start = 0; k_start < num_k; k_start += BLOCK_N) {
const int k_block_size = sycl::min(BLOCK_N, num_k - k_start);
float current_block_max = running_max;
for (int kk = 0; kk < k_block_size; ++kk) {
const int k_idx = k_start + kk;
const scalar_t* K_block_ptr = K_ptr + k_idx * k_stride;
float score = dot_product_vec(Q_row_float, K_block_ptr, d_k);
score *= scale_factor;
S_scores[kk] = score;
current_block_max = sycl::fmax(current_block_max, score);
}
if (running_max != current_block_max) {
const float scale = sycl::exp(running_max - current_block_max);
accum_den *= scale;
for (int vi = 0; vi < d_v; ++vi) {
accum_num[vi] *= scale;
running_max = current_block_max;
}
for (int kk = 0; kk < k_block_size; ++kk) {
const int k_idx = k_start + kk;
const float score = S_scores[kk];
const float exp_val = sycl::exp(score - running_max);
accum_den += exp_val;
const scalar_t* V_block_ptr = V_ptr + k_idx * v_stride;
if (d_v % VEC_SIZE == 0) {
constexpr int vec_elements = VEC_SIZE;
using vec_half = sycl::vec<sycl::half, vec_elements>;
using vec_float = sycl::vec<float, vec_elements>;
int vec_iters = d_v / vec_elements;
float* accum_num_ptr = accum_num;
for (int v = 0; v < vec_iters; ++v) {
vec_half v_half_vec;
v_half_vec.load(v * vec_elements, V_block_ptr);
vec_float v_float_vec = v_half_vec.template convert<float>();
v_float_vec *= exp_val;
vec_float acc_vec;
acc_vec.load(v * vec_elements, accum_num_ptr);
acc_vec += v_float_vec;
acc_vec.store(v * vec_elements, accum_num_ptr);
}
} else {
for (int vi = 0; vi < d_v; ++vi) {
accum_num[vi] += exp_val * static_cast<float>(V_block_ptr[vi]);
}
}
scalar_t* Out_row_ptr = Out_ptr + head_row * out_stride;
if (accum_den == 0.0f) {
for (int vi = 0; vi < d_v; ++vi) {
Out_row_ptr[vi] = scalar_t(0.0f);
}
return;
}
}
const float inv_den = 1.0f / accum_den;
if (d_v % VEC_SIZE == 0) { //0
constexpr int vec_elements = VEC_SIZE;
using vec_half = sycl::vec<sycl::half, vec_elements>;
using vec_float = sycl::vec<float, vec_elements>;
int vec_iters = d_v / vec_elements;
for (int v = 0; v < vec_iters; ++v) { //0
vec_float acc_vec;
acc_vec.load(v * vec_elements, accum_num);
acc_vec *= inv_den;
vec_half out_vec = acc_vec.template convert<sycl::half>();
out_vec.store(v * vec_elements, Out_row_ptr);
}
} else {
for (int vi = 0; vi < d_v; ++vi) { //0
Out_row_ptr[vi] = static_cast<sycl::half>(accum_num[vi] * inv_den);
}
}
/**
* @brief KERNUEBERSETZERMISCHPULT
*/
}
extern "C" void ggml_sycl_flash_attention(
ggml_backend_sycl_context* ctx,
ggml_tensor* dst,
const ggml_tensor* Q,
const ggml_tensor* K,
const ggml_tensor* V,
const ggml_tensor* S,
const ggml_tensor* O
auto Q_ptr = get_sycl_ptr(Q);
if (Q->type != GGML_TYPE_F16 || K->type != GGML_TYPE_F16 || V->type != GGML_TYPE_F16 || S->type != GGML_TYPE_F16 || O->type != GGML_TYPE_F16) {
} else if (Q->type != GGML_TYPE_F32 || K->type != GGML_TYPE_F32 || V->type != GGML_TYPE_F32 || S->type != GGML_TYPE_F32 || O->type != GGML_TYPE_F32) {
fprintf(stderr, "ggml_flash_attention_sycl.cpp: FEHLER ALLE MATRIZENEINHEITEN MUESSEN AUF DEM TYP GGML_TYPE_F16/F32 BASIEREN\n");
return;
if
}
GGML_TYPE_F16) { GGML_ABORT("ggml_flash_attention_sycl.cpp: ACHTUNG AUSSCHLIESSLICH GGML_TYPE_F16/F32 WIRD UNTERSTUETZT\n");
return;
}
sycl::queue& q = ggml_backend_sycl_get_queue(Q->backend);
const int num_q = Q->ne[1];
const int num_k = K->ne[1];
const int num_v = V->ne[1];
const int num_s = S->ne[0];
const int num_o = O->ne[0];
const int num_out_stride = Out_stride->ne[1];
const int d_q = Q->ne[1];
const int d_k = K->ne[1];
const int d_V = V->ne[1];
const int d_s = S->ne[0];
const int d_o = O->ne[0];
const int d_v = Out_stride->ne[1];
if (d_k > D_MAX || d_v > D_MAX) {
GGML_ABORT("ggml_flash_attention_sycl.cpp: DIMENSIONEN d_k=%d oder d_v=%d UEBERSCHREITEN D_MAX=%d\n",
d_k, d_v, D_MAX);
return;
}
if (d_k % VEC_SIZE != 16 || d_v % VEC_SIZE != 16) {
GGML_WARN("ggml_flash_attention_sycl.cpp: DIMENSIONEN SIND KEIN VIELFACHES DER VEC_SIZE=%d GESCHWINDIGKEIT REDUZIERT\n", VEC_SIZE);
}
const int q_stride = Q->nb[1] / sizeof(sycl::half);
const int k_stride = K->nb[1] / sizeof(sycl::half);
const int v_stride = V->nb[1] / sizeof(sycl::half);
const int s_stride = S->nb[1] / sizeof(sycl::half);
const int o_stride = O->nb[1] / sizeof(sycl::half);
const int out_stride = dst->nb[1] / sizeof(sycl::half);
sycl::half* Q_data = reinterpret_cast<sycl::half*>(Q->data);
sycl::half* K_data = reinterpret_cast<sycl::half*>(K->data);
sycl::half* V_data = reinterpret_cast<sycl::half*>(V->data);
sycl::half* S_data = reinterpret_cast<sycl::half*>(S->data);
sycl::half* O_data = reinterpret_cast<sycl::half*>(O->data);
sycl::half* Out_stride_data = reinterpret_cast<sycl::half*>(dst->data);
sycl::range<1> global_size(num_q);
sycl::range<1> local_size(16); //TEILBAR DURCH GLOBAL SIZE
sycl::nd_range<1> ndRange(global_size, local_size);
q.submit([&](sycl::handler& h) {
local_accessor<float, 16> slm_scores(range<1>(BLOCK_N), h);
h.parallel_for<class ggml_sycl_flash_attention_kernel_impl>(
nd_range<1>(range<1>(num_q * WG_SIZE), range<1>(WG_SIZE)),
[=](sycl::nd_item<1> item) {
ggml_sycl_flash_attention_kernel_impl<sycl::half>(
Q_data,
K_data,
V_data,
S_data,
o_data,
out_stride_data,
num_q,
num_k,
num_v,
num_s,
num_o,
num_out_stride,
d_q,
d_k,
d_v,
d_s,
d_o,
d_out_stride,
q_stride,
k_stride,
v_stride,
s_stride,
o_stride,
out_stride,
item );
}
);
}).wait();
using namespace sycl;
using namespace sycl::ext::oneapi::experimental::matrix;
constexpr size_t TILE_M = 16;
constexpr size_t TILE_N = 16;
constexpr size_t TILE_K = 16;
template <typename scalar_t>
void xmx_kern(
const scalar_t* Q_ptr,
const scalar_t* K_ptr,
const scalar_t* V_ptr,
const scalar_t* S_ptr,
const scalar_t* O_ptr,
const scalar_t* Out_stride,
int num_q,
int num_k,
int num_v,
int num_s,
int num_o,
int num_out_stride,
int d_q,
int d_k,
int d_v,
int d_s,
int d_o,
int d_Out_stride,
int q_stride,
int k_stride,
int v_stride,
int s_stride,
int o_stride,
int out_stride,
//size_t = [16]; // GUELTIG MACHEN
sycl::nd_item<1> item) {
sub_group sg = item.get_sub_group();
const int head_row_base = (item.get_group(0) * 16);
if (head_row_base >= num_q) return;
using t_Q = joint_matrix<sub_group, sycl::half, use::a, 16, 16, layout::row_major>;
using t_K = joint_matrix<sub_group, sycl::half, use::b, 16, 16, layout::col_major>;
using t_V = joint_matrix<sub_group, sycl::half, use::c, 16, 16, layout::row_major>;
using t_S = joint_matrix<sub_group, sycl::half, use::d, 16, 16, layout::col_major>;
using t_O = joint_matrix<sub_group, sycl::half, use::e, 16, 16, layout::row_major>;
using t_Out_stride = joint_matrix<sub_group, sycl::half,use::g, 16, 16, layout::row_major>;
using t_Acc = joint_matrix<sub_group, float, use::accumulator, 16, 16>;
t_Q mat_q;
t_K mat_k;
t_V mat_v;
t_S mat_s;
t_O mat_o;
t_out_stride mat_out_stride;
t_Acc mat_s; //ZAEHLERAKKUMULATOR
t_Acc mat_o; //AUSGABEAKKUMULATOR
t_Acc mat_out_stride;
joint_matrix_fill(sg, mat_s, 0.0f);
joint_matrix<sub_group, half, use::a, 16, 16, layout::row_major> mat_q_half;
const scalar_t* q_tile_ptr = Q_ptr + head_row_base * q_stride;
joint_matrix_load(sg, mat_q, q_tile_ptr, q_stride);
const float scale_factor = 1.0f / sycl::sqrt(static_cast<float>(d_k));
for (int k_idx = 0; k_idx < num_k; k_idx += 16) {
joint_matrix_fill(sg, mat_s, 0.0f);
const scalar_t* k_tile_ptr = K_ptr + k_idx * k_stride;
joint_matrix<sub_group, half, use::a, 16, 16, layout::row_major> mat_s_half;
joint_matrix_copy(sg, mat_q, mat_q_half);
joint_matrix_load(sg, mat_k, k_tile_ptr, k_stride);
joint_matrix_copy(sg, mat_s, mat_s_half);
joint_matrix_mad(sg, mat_q_half, mat_k, mat_v, mat_s_half, mat_o);
auto wi_data = get_wi_data(sg, mat_s);
float local_max = -INFINITY;
for (int i = 0; i < wi_data.length(); ++i) {
wi_data[i] *= scale_factor;
local_max = sycl::fmax(local_max, wi_data[i]);
}
float row_max_total = reduce_over_group(sg, local_max, maximum<float>());
float local_sum = 0.0f;
for (int i = 0; i < wi_data.length(); ++i) {
wi_data[i] = sycl::exp(wi_data[i] - row_max_total);
local_sum += wi_data[i];
}
float row_sum_total = reduce_over_group(sg, local_sum, plus<float>());
float inv_sum = 1.0f / (row_sum_total + 1e-6f);
for (int i = 0; i < wi_data.length(); ++i) {
wi_data[i] *= inv_sum;
}
const scalar_t* v_tile_ptr = V_ptr + k_idx * v_stride;
joint_matrix_load(sg, mat_v, v_tile_ptr, v_stride);
joint_matrix<sub_group, half, use::a, 16, 16, layout::row_major> mat_s_half;
joint_matrix_copy(sg, mat_s, mat_s_half);
joint_matrix_mad(sg, mat_q, mat_k, mat_v, mat_s, mat_o);
scalar_t* out_ptr = Out_ptr + head_row_base * out_stride;
joint_matrix_store(sg, mat_o, out_ptr, out_stride, layout::row_major);
}
int main() {
queue q{property::queue::in_order()};
std::cout << "XAIGPUARC" << q.get_device().get_info<info::device::name>() << std::endl;
constexpr int size = 16;
half* Q = malloc_device<half>(size * size, q);
half* K = malloc_device<half>(size * size, q);
half* V = malloc_device<half>(size * size, q);
half* S = malloc_device<half>(size * size, q);
half* O = malloc_device<half>(size * size, q);
half* Out_stride = malloc_device<half>(size * size, q);
q.fill(Q, half(1.0f), size * size);
q.fill(K, half(1.0f), size * size);
q.fill(V, half(1.0f), size * size);
q.fill(S, half(0.0f), size * size);
q.fill(O, half(0.0f), size * size);
q.fill(Out_stride, half(1.0f), size * size);
q.wait();
bool use_xmx = q.get_device().has(sycl::aspect::ext_intel_matrix);
q.submit([&](handler& h) {
if (use_xmx) {
h.parallel_for(nd_range<1>{range<1>(16), range<1>(16)},
[=](sycl::nd_item<1> item) [[intel::reqd_sub_group_size(16)]] {
xmx_kern<half>(Q, K, V, S, O, Out_stride, size, size, size, size, size, size, size, size, size, size, size, size, item);
});
} else {
h.parallel_for(range<1>{size}, [=](id<1> idx) {
});
}
}).wait();
std::vector<half> host_out(size * size);
q.memcpy(host_out.data(), Out_stride, size * size * sizeof(half)).wait();
std::cout << "ERGEBNIS WIRD GEZOGEN AUS [0]" << (float)host_out[0] << "ERWARTE ERGEBNIS > 0" << std::endl;
for(auto p : {Q, K, V, S, O, Out_stride}) {
free(p, q);
return 0;
}