//ggml_sycl_flash_attention.cpp STANDALONE ALLEINSTELLUNGSFUNKTION fuer XAIGPUARC OpenSource
//SYCL AI PROGRAMM gsflm.cpp from alucian Berlin-Buch 04.08.2026 /// 18:12
//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);
#include
#include
#include
#include
#include
#include
#include
#include
#include
#include
#include
#include
#include "ggml-sycl.h"
#include "ggml.h"
#include "ggml-impl.h"
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 PUNKTPRODUKT ALS FLIESSWERTE
*/
using namespace sycl;
using namespace sycl::ext::oneapi::experimental::matrix;
constexpr int WG_SIZE = 16;
//Orchestrator XMX KERN UMGEBUNGSVORBAU MIT sub_group joint_matrix
//Priorität 1-3: Orchestrator, XMX-Kernel, Vektor-Fallback
extern "C" void ggml_sycl_flash_attention_xmx_vec(queue& q, half* out, half* o_ptr, half* k, half* v, int num_q, int d_k) {
bool has_xmx = sycl_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::row_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_ptr + (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);
});
}
}
auto& sycl_q = ggml_backend_sycl_get_queue(ctx);
auto dev = q.get_device();
auto sg = item.get_sub_group();
const int m = item.get_group(0) * 16;
const int n = item.get_group(1) * 16;
bool has_xmx = dev.has(sycl::aspect::ext_intel_matrix);
bool use_xmx = sycl_q.get_device().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) {
sycl_q.submit([&](sycl::handler& h) {
void xmx_kernel2(queue& q, T* A, T* B, T* C, int M, int N, int K) {
q.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);
sub_group sg = item.get_sub_group();
//Joint Matrix Zeug aendern
joint_matrix<sub_group, half, use::a, 16, 16, layout::row_major> mat_q;
joint_matrix<sub_group, half, use::b, 16, 16, layout::row_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_ptr_half, m_s_acc);
// Rescale & Store Logik hier integrieren
joint_matrix_store(sg, mat_s, (float*)out, d_k, layout::row_major);
h.parallel_for(nd_range<2>({M/16, N/16}, {1, 1}),
[=](nd_item<2> item) {
joint_matrix<sub_group, float, use::accumulator, 16, 16> mat_c;
joint_matrix_fill(sg, mat_c, 0.0f);
// Loop über K-Dimension
for (int k = 0; k < K; k += 16) {
joint_matrix_load(sg, mat_a, a_ptr, A + m*K + k, K);
joint_matrix_load(sg, mat_b, b_ptr, B + k*N + n, N);
joint_matrix_mad(sg, mat_c, mat_a, mat_b, mat_c);
joint_matrix_store(sg, mat_c, c_ptr, C + m*N + n, N, layout::row_major);
//joint matrix zeug hier
});
}).wait();
} else {
//NORMALER FLASH ATTENTION SYCL KERN SKALAR GEFOLGT VON VECTOR
ggml_sycl_flash_attention(ctx, dst, q, k, v);
}
}
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) {
}
ATTENTION
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
*/
//GGML_SYCL_FLASH_ATTENTION SKALARVERION
template <typename scalar_t>
void ggml_sycl_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
* @tparam vec_float
*/
}
//GGML_SYCL_FLASH_ATTENTION VECTORVERSION
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 != 0 || d_v % VEC_SIZE != 0) {
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>
//XMX_KERN
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* out_stride,
int num_q,
int num_k,
int num_v,
int num_s,
int num_out_stride,
int d_q,
int d_k,
int d_v,
int d_s,
int d_out_stride,
int q_stride,
int k_stride,
int v_stride,
int s_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_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_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);
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);
}
//HAUPTFUNKTIONSABLAUF
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, k);
half* v = malloc_device<half>(size * size, v);
half* out_stride = malloc_device<half>(size * size, o);
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(out_stride, half(1.0f), size * size);
q.wait();
bool use_xmx = sycl_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, 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, out_stride}) {
free(p, q);
return 0;
}
´´
RE: Ich habe ueber die Jahre sowohl Windows als auch Linux Mining Programme gesammelt und sortiert.