mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-27 21:46:57 +02:00
Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
13baab725f | ||
|
|
4ceb171910 | ||
|
|
73c941b111 | ||
|
|
0f8a414b75 | ||
|
|
f95b0d9539 | ||
|
|
c350a40bbd | ||
|
|
9b421fa946 |
@@ -5131,6 +5131,9 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
|
||||
if (b->type == GGML_TYPE_F16 && a->type != GGML_TYPE_F16) {
|
||||
return false;
|
||||
}
|
||||
if (op->op == GGML_OP_MUL_MAT_ID && ggml_get_op_params_i32(op, 3) == GGML_PREC_F32) {
|
||||
return false;
|
||||
}
|
||||
#ifdef GGML_USE_MUSA
|
||||
const int cc = ggml_cuda_info().devices[dev_ctx->device].cc;
|
||||
if (b->ne[2]*b->ne[3] > 1 && !ggml_is_transposed(a) && !ggml_is_transposed(b)) {
|
||||
|
||||
@@ -1156,14 +1156,18 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm_id(ggml_m
|
||||
|
||||
const bool bc_inp = op->src[0]->ne[0] % 32 != 0;
|
||||
|
||||
// src1 prec [TAG_GGML_PREC]
|
||||
const bool amax = ggml_get_op_params_i32(op, 3) == GGML_PREC_F32;
|
||||
|
||||
snprintf(base, 256, "kernel_mul_mm_id_%s_%s", ggml_type_name(tsrc0), ggml_type_name(tsrc1));
|
||||
snprintf(name, 256, "%s_bci=%d", base, bc_inp);
|
||||
snprintf(name, 256, "%s_bci=%d_amax=%d", base, bc_inp, amax);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
if (!res.pipeline) {
|
||||
ggml_metal_cv_t cv = ggml_metal_cv_init();
|
||||
|
||||
ggml_metal_cv_set_bool(cv, bc_inp, FC_MUL_MM + 0);
|
||||
ggml_metal_cv_set_bool(cv, amax, FC_MUL_MM + 6);
|
||||
|
||||
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
|
||||
|
||||
|
||||
@@ -2719,9 +2719,12 @@ int ggml_metal_op_mul_mat_id(ggml_metal_op_t ctx, int idx) {
|
||||
ggml_metal_buffer_id bid_amax = bid_ids;
|
||||
bid_amax.offs += ggml_metal_op_mul_mat_id_extra_ids(op);
|
||||
|
||||
// src1 prec [TAG_GGML_PREC]
|
||||
const bool use_amax = ggml_get_op_params_i32(op, 3) == GGML_PREC_F32;
|
||||
|
||||
// src1 rescale factors, computed before the matmul
|
||||
// ref: https://github.com/ggml-org/llama.cpp/pull/26223
|
||||
{
|
||||
if (use_amax) {
|
||||
ggml_metal_kargs_mul_mm_id_amax args = {
|
||||
/*.ne00 =*/ ne10,
|
||||
/*.ne01 =*/ ne11,
|
||||
@@ -2779,17 +2782,17 @@ int ggml_metal_op_mul_mat_id(ggml_metal_op_t ctx, int idx) {
|
||||
|
||||
ggml_metal_op_concurrency_reset(ctx);
|
||||
|
||||
{
|
||||
if (use_amax) {
|
||||
auto pipeline = ggml_metal_library_get_pipeline_mul_mm_id_amax(lib);
|
||||
|
||||
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||
ggml_metal_encoder_set_buffer (enc, bid_amax, 0);
|
||||
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, 1, 1, 1, 32, 1, 1);
|
||||
}
|
||||
|
||||
// the next kernel has to wait for the amax data
|
||||
ggml_metal_op_concurrency_reset(ctx);
|
||||
// the next kernel has to wait for the amax data
|
||||
ggml_metal_op_concurrency_reset(ctx);
|
||||
}
|
||||
|
||||
{
|
||||
auto pipeline = ggml_metal_library_get_pipeline_mul_mm_id(lib, op);
|
||||
|
||||
@@ -7,6 +7,7 @@ constant short FC_mul_mm_ne12 [[function_constant(FC_MUL_MM + 2)]];
|
||||
constant short FC_mul_mm_ne13 [[function_constant(FC_MUL_MM + 3)]];
|
||||
constant short FC_mul_mm_r2 [[function_constant(FC_MUL_MM + 4)]];
|
||||
constant short FC_mul_mm_r3 [[function_constant(FC_MUL_MM + 5)]];
|
||||
constant bool FC_mul_mm_id_amax [[function_constant(FC_MUL_MM + 6)]];
|
||||
|
||||
// each block_q contains 16*nl weights
|
||||
#ifdef GGML_METAL_HAS_TENSOR
|
||||
@@ -584,8 +585,8 @@ kernel void kernel_mul_mm_id(
|
||||
const short lb1 = (short) tiitg/NL1; // 0 .. NR1-1, this thread's row of the B tile
|
||||
|
||||
// power-of-two rescaling
|
||||
const float s1_inv = ((device const float *) amax)[0];
|
||||
const float s1_scale = ((device const float *) amax)[1];
|
||||
const float s1_inv = FC_mul_mm_id_amax ? ((device const float *) amax)[0] : 1.0f;
|
||||
const float s1_scale = FC_mul_mm_id_amax ? ((device const float *) amax)[1] : 1.0f;
|
||||
|
||||
#ifndef GGML_METAL_HAS_TENSOR
|
||||
S0_8x8 ma[4];
|
||||
|
||||
+44
-73
@@ -4771,80 +4771,51 @@ static void quantize_row_iq1_m_impl(const float * GGML_RESTRICT x, void * GGML_R
|
||||
// 1: +, -
|
||||
// 2: -, +
|
||||
// 3: -, -
|
||||
for (int i1 = 0; i1 <= block_size; ++i1) {
|
||||
for (int i2 = i1; i2 <= block_size; ++i2) {
|
||||
memset(sumqx, 0, 4*sizeof(float));
|
||||
memset(sumq2, 0, 4*sizeof(float));
|
||||
for (int j = 0; j < i1; ++j) {
|
||||
int i = idx[2*j];
|
||||
if (i < block_size/2) {
|
||||
sumqx[0] += weight[i]*x_p[0]*xb[i];
|
||||
sumqx[1] += weight[i]*x_p[0]*xb[i];
|
||||
sumqx[2] += weight[i]*x_m[0]*xb[i];
|
||||
sumqx[3] += weight[i]*x_m[0]*xb[i];
|
||||
sumq2[0] += weight[i]*x_p[0]*x_p[0];
|
||||
sumq2[1] += weight[i]*x_p[0]*x_p[0];
|
||||
sumq2[2] += weight[i]*x_m[0]*x_m[0];
|
||||
sumq2[3] += weight[i]*x_m[0]*x_m[0];
|
||||
} else {
|
||||
sumqx[0] += weight[i]*x_p[0]*xb[i];
|
||||
sumqx[2] += weight[i]*x_p[0]*xb[i];
|
||||
sumqx[1] += weight[i]*x_m[0]*xb[i];
|
||||
sumqx[3] += weight[i]*x_m[0]*xb[i];
|
||||
sumq2[0] += weight[i]*x_p[0]*x_p[0];
|
||||
sumq2[2] += weight[i]*x_p[0]*x_p[0];
|
||||
sumq2[1] += weight[i]*x_m[0]*x_m[0];
|
||||
sumq2[3] += weight[i]*x_m[0]*x_m[0];
|
||||
// prefix sums are kept per half of the block because each half can use a different sign (x_p or x_m)
|
||||
// since v[0]-v[1] = v[1]-v[2] = -1 for both x_p and x_m, the 3-group sum for a split collapses to T*v[2] - px[i1] - px[i2]
|
||||
{
|
||||
float px[2][IQ1M_BLOCK_SIZE+1];
|
||||
float pw[2][IQ1M_BLOCK_SIZE+1];
|
||||
px[0][0] = px[1][0] = 0;
|
||||
pw[0][0] = pw[1][0] = 0;
|
||||
for (int j = 0; j < block_size; ++j) {
|
||||
const int i = idx[2*j];
|
||||
const int h = i < block_size/2 ? 0 : 1;
|
||||
px[h][j+1] = px[h][j] + weight[i]*xb[i];
|
||||
px[1-h][j+1] = px[1-h][j];
|
||||
pw[h][j+1] = pw[h][j] + weight[i];
|
||||
pw[1-h][j+1] = pw[1-h][j];
|
||||
}
|
||||
const float txs[2] = {px[0][block_size], px[1][block_size]}; // total weight*x per half
|
||||
const float tws[2] = {pw[0][block_size], pw[1][block_size]}; // total weight per half
|
||||
const float p2 = x_p[2], m2 = x_m[2];
|
||||
const float cp1 = x_p[0]*x_p[0] - x_p[1]*x_p[1];
|
||||
const float cp2 = x_p[1]*x_p[1] - x_p[2]*x_p[2];
|
||||
const float cm1 = x_m[0]*x_m[0] - x_m[1]*x_m[1];
|
||||
const float cm2 = x_m[1]*x_m[1] - x_m[2]*x_m[2];
|
||||
for (int i1 = 0; i1 <= block_size; ++i1) {
|
||||
for (int i2 = i1; i2 <= block_size; ++i2) {
|
||||
float qx_p[2], qx_m[2], q2_p[2], q2_m[2];
|
||||
for (int h = 0; h < 2; ++h) {
|
||||
const float sx = px[h][i1] + px[h][i2];
|
||||
qx_p[h] = txs[h]*p2 - sx;
|
||||
qx_m[h] = txs[h]*m2 - sx;
|
||||
q2_p[h] = tws[h]*p2*p2 + pw[h][i1]*cp1 + pw[h][i2]*cp2;
|
||||
q2_m[h] = tws[h]*m2*m2 + pw[h][i1]*cm1 + pw[h][i2]*cm2;
|
||||
}
|
||||
}
|
||||
for (int j = i1; j < i2; ++j) {
|
||||
int i = idx[2*j];
|
||||
if (i < block_size/2) {
|
||||
sumqx[0] += weight[i]*x_p[1]*xb[i];
|
||||
sumqx[1] += weight[i]*x_p[1]*xb[i];
|
||||
sumqx[2] += weight[i]*x_m[1]*xb[i];
|
||||
sumqx[3] += weight[i]*x_m[1]*xb[i];
|
||||
sumq2[0] += weight[i]*x_p[1]*x_p[1];
|
||||
sumq2[1] += weight[i]*x_p[1]*x_p[1];
|
||||
sumq2[2] += weight[i]*x_m[1]*x_m[1];
|
||||
sumq2[3] += weight[i]*x_m[1]*x_m[1];
|
||||
} else {
|
||||
sumqx[0] += weight[i]*x_p[1]*xb[i];
|
||||
sumqx[2] += weight[i]*x_p[1]*xb[i];
|
||||
sumqx[1] += weight[i]*x_m[1]*xb[i];
|
||||
sumqx[3] += weight[i]*x_m[1]*xb[i];
|
||||
sumq2[0] += weight[i]*x_p[1]*x_p[1];
|
||||
sumq2[2] += weight[i]*x_p[1]*x_p[1];
|
||||
sumq2[1] += weight[i]*x_m[1]*x_m[1];
|
||||
sumq2[3] += weight[i]*x_m[1]*x_m[1];
|
||||
}
|
||||
}
|
||||
for (int j = i2; j < block_size; ++j) {
|
||||
int i = idx[2*j];
|
||||
if (i < block_size/2) {
|
||||
sumqx[0] += weight[i]*x_p[2]*xb[i];
|
||||
sumqx[1] += weight[i]*x_p[2]*xb[i];
|
||||
sumqx[2] += weight[i]*x_m[2]*xb[i];
|
||||
sumqx[3] += weight[i]*x_m[2]*xb[i];
|
||||
sumq2[0] += weight[i]*x_p[2]*x_p[2];
|
||||
sumq2[1] += weight[i]*x_p[2]*x_p[2];
|
||||
sumq2[2] += weight[i]*x_m[2]*x_m[2];
|
||||
sumq2[3] += weight[i]*x_m[2]*x_m[2];
|
||||
} else {
|
||||
sumqx[0] += weight[i]*x_p[2]*xb[i];
|
||||
sumqx[2] += weight[i]*x_p[2]*xb[i];
|
||||
sumqx[1] += weight[i]*x_m[2]*xb[i];
|
||||
sumqx[3] += weight[i]*x_m[2]*xb[i];
|
||||
sumq2[0] += weight[i]*x_p[2]*x_p[2];
|
||||
sumq2[2] += weight[i]*x_p[2]*x_p[2];
|
||||
sumq2[1] += weight[i]*x_m[2]*x_m[2];
|
||||
sumq2[3] += weight[i]*x_m[2]*x_m[2];
|
||||
}
|
||||
}
|
||||
for (int k = 0; k < 4; ++k) {
|
||||
if (sumq2[k] > 0 && sumqx[k]*sumqx[k] > best_score*sumq2[k]) {
|
||||
scale = sumqx[k]/sumq2[k]; best_score = scale*sumqx[k];
|
||||
besti1 = i1; besti2 = i2; best_k = k;
|
||||
sumqx[0] = qx_p[0] + qx_p[1];
|
||||
sumqx[1] = qx_p[0] + qx_m[1];
|
||||
sumqx[2] = qx_m[0] + qx_p[1];
|
||||
sumqx[3] = qx_m[0] + qx_m[1];
|
||||
sumq2[0] = q2_p[0] + q2_p[1];
|
||||
sumq2[1] = q2_p[0] + q2_m[1];
|
||||
sumq2[2] = q2_m[0] + q2_p[1];
|
||||
sumq2[3] = q2_m[0] + q2_m[1];
|
||||
for (int k = 0; k < 4; ++k) {
|
||||
if (sumq2[k] > 0 && sumqx[k]*sumqx[k] > best_score*sumq2[k]) {
|
||||
scale = sumqx[k]/sumq2[k]; best_score = scale*sumqx[k];
|
||||
besti1 = i1; besti2 = i2; best_k = k;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1173,6 +1173,10 @@ static void launch_fattn_tile_switch_ncols2(ggml_backend_sycl_context & ctx, ggm
|
||||
launch_fattn_tile_switch_ncols1<DKQ, DV, 16, use_logit_softcap>(ctx, dst);
|
||||
return;
|
||||
}
|
||||
if (use_gqa_opt && gqa_ratio % 8 == 0) {
|
||||
launch_fattn_tile_switch_ncols1<DKQ, DV, 8, use_logit_softcap>(ctx, dst);
|
||||
return;
|
||||
}
|
||||
if (use_gqa_opt && gqa_ratio % 4 == 0) {
|
||||
launch_fattn_tile_switch_ncols1<DKQ, DV, 4, use_logit_softcap>(ctx, dst);
|
||||
return;
|
||||
|
||||
@@ -124,6 +124,24 @@ struct vk_flash_attn_push_constants {
|
||||
|
||||
static_assert(sizeof(vk_flash_attn_push_constants) <= 128, "sizeof(vk_flash_attn_push_constants) must be <= 128");
|
||||
|
||||
struct vk_fa_xe_opt_push_constants {
|
||||
uint32_t kv_seq_len;
|
||||
uint32_t activation_length;
|
||||
uint32_t q_head;
|
||||
uint32_t kv_head;
|
||||
uint32_t qk_ratio;
|
||||
uint32_t qk_sub_groups;
|
||||
uint32_t flag;
|
||||
uint32_t nbkv_tok;
|
||||
uint32_t nbkv_head;
|
||||
uint32_t batch_stride_q;
|
||||
uint32_t batch_stride_k;
|
||||
uint32_t batch_stride_v;
|
||||
uint32_t batch_stride_m;
|
||||
uint32_t batch_stride_o;
|
||||
float softmax_scale;
|
||||
};
|
||||
|
||||
struct vk_op_push_constants {
|
||||
uint32_t KX;
|
||||
uint32_t KY;
|
||||
|
||||
@@ -996,6 +996,7 @@ struct vk_device_struct {
|
||||
bool fa_sparse_compact_use_subgroups;
|
||||
|
||||
vk_pipeline pipeline_flash_attn_split_k_reduce;
|
||||
std::map<std::tuple<uint32_t, uint32_t, uint32_t, uint32_t>, std::pair<vk_pipeline, vk_pipeline>> pipeline_xe_fa_decode_dual_phases;
|
||||
vk_pipeline pipeline_count_experts;
|
||||
|
||||
// [2] is for whether to take n_experts from spec constant (0) or push constant (1)
|
||||
|
||||
@@ -2973,6 +2973,46 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
|
||||
ggml_vk_create_pipeline(device, device->pipeline_matmul_split_k_reduce, "split_k_reduce", split_k_reduce_len, split_k_reduce_data, "main", 2, 2 * sizeof(uint32_t), {256 * 4, 1, 1}, {}, 1);
|
||||
ggml_vk_create_pipeline(device, device->pipeline_flash_attn_split_k_reduce, "fa_split_k_reduce", fa_split_k_reduce_len, fa_split_k_reduce_data, "main", 3, sizeof(vk_op_flash_attn_split_k_reduce_push_constants), {1, device->subgroup_size, 1}, {device->subgroup_size}, 1, true);
|
||||
|
||||
if (device->vendor_id == VK_VENDOR_ID_INTEL && (device->architecture == INTEL_XE2 || (device->architecture == INTEL_XE1 && device->coopmat_support && device->uma))) {
|
||||
auto upper_power_of_2 = [&](uint32_t in) {
|
||||
GGML_ASSERT(in != 0);
|
||||
if (in <= 1) return 1u;
|
||||
uint32_t ret = in - 1;
|
||||
ret |= ret >> 1;
|
||||
ret |= ret >> 2;
|
||||
ret |= ret >> 4;
|
||||
ret |= ret >> 8;
|
||||
ret |= ret >> 16;
|
||||
return ret + 1;
|
||||
};
|
||||
|
||||
uint32_t xe_native_sub_group_size = 16;
|
||||
if (device->architecture == INTEL_XE1) {
|
||||
xe_native_sub_group_size = 8;
|
||||
}
|
||||
|
||||
for (auto& it : device->pipeline_xe_fa_decode_dual_phases) {
|
||||
const uint32_t split_p_chunk = 32;
|
||||
auto HdQk = it.first;
|
||||
auto& pipelines = it.second;
|
||||
uint32_t head_dim_qk = std::get<0>(HdQk);
|
||||
uint32_t head_dim_pv = std::get<1>(HdQk);
|
||||
uint32_t gqa_ratio = std::get<2>(HdQk);
|
||||
uint32_t q_len = std::get<3>(HdQk);
|
||||
const uint32_t out_dim_per_wg = gqa_ratio > 16 ? 8 : 16;
|
||||
uint32_t aligned_q_len = upper_power_of_2(q_len);
|
||||
uint32_t group_sz_ph1 = std::min(std::max(aligned_q_len * xe_native_sub_group_size, 64u), 256u);
|
||||
uint32_t out_per_wg_ph1 = std::min(q_len, 256u / xe_native_sub_group_size);
|
||||
uint32_t aligned_gqa_ratio = upper_power_of_2(gqa_ratio);
|
||||
uint32_t split_p_per_iter_ph2 = 256;
|
||||
uint32_t split_p_per_warp = 16;
|
||||
uint32_t group_sz_ph2 = (split_p_per_iter_ph2 / split_p_per_warp) * xe_native_sub_group_size;
|
||||
uint32_t out_per_wg_ph2 = std::min(std::max(16u / aligned_gqa_ratio, 1u), q_len);
|
||||
ggml_vk_create_pipeline(device, pipelines.first, "xe_fa_decode_ph1", fa_decode_ph1_cm1_len, fa_decode_ph1_cm1_data, "main", 5, sizeof(vk_fa_xe_opt_push_constants), { 1, 32, 1 }, { group_sz_ph1, gqa_ratio, head_dim_qk, xe_native_sub_group_size, split_p_chunk, out_per_wg_ph1 }, 1, false, true, xe_native_sub_group_size);
|
||||
ggml_vk_create_pipeline(device, pipelines.second, "xe_fa_decode_ph2", fa_decode_ph2_cm1_len, fa_decode_ph2_cm1_data, "main", 5, sizeof(vk_fa_xe_opt_push_constants), { 1, 1, 1 }, { group_sz_ph2, gqa_ratio, head_dim_pv, out_per_wg_ph2, xe_native_sub_group_size, split_p_per_iter_ph2, split_p_chunk, out_dim_per_wg }, 1, false, true, xe_native_sub_group_size);
|
||||
}
|
||||
}
|
||||
|
||||
for (auto &it : device->pipeline_fa_mask_opt) {
|
||||
auto BrBc = it.first;
|
||||
ggml_vk_create_pipeline(device, it.second, "fa_mask_opt", fa_mask_opt_len, fa_mask_opt_data, "main", 2, sizeof(vk_op_flash_attn_mask_opt_push_constants), {1, 1, 1}, {128, 128 / device->subgroup_size, BrBc.first, BrBc.second}, 1, true, true, device->subgroup_size);
|
||||
@@ -7899,6 +7939,18 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
|
||||
|
||||
vk_pipeline pipeline = nullptr;
|
||||
|
||||
bool xe_fa_opt = false;
|
||||
bool fa_copy_qstate = false;
|
||||
bool xe_fa_supported_platform =
|
||||
(ctx->device.get()->architecture == INTEL_XE2 && ctx->device.get()->properties.deviceID != 0xFD80 && ctx->device.get()->properties.deviceID != 0xFD81) ||
|
||||
(ctx->device.get()->architecture == INTEL_XE1 && ctx->device.get()->coopmat_support && ctx->device.get()->uma);
|
||||
bool xe_fa_supported_usage = neq0 % 32 == 0 && nev0 % 16 == 0 && q->nb[1] > q->nb[2] && k->nb[1] > k->nb[2] && v->nb[1] > v->nb[2] && mask != nullptr;
|
||||
bool xe_fa_supported_dtype = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_F16 && v->type == GGML_TYPE_F16 && (mask != nullptr && mask->type == GGML_TYPE_F16);
|
||||
std::pair<vk_pipeline, vk_pipeline> xe_fa_pipeline_dual_phases = { nullptr , nullptr };
|
||||
vk_pipeline xe_fa_pipeline = nullptr;
|
||||
size_t size_p = 0;
|
||||
size_t size_group_max = 0;
|
||||
|
||||
{
|
||||
std::lock_guard<std::mutex> guard(ctx->device->compile_mutex);
|
||||
auto &pipelines = ctx->device->pipeline_flash_attn_f32_f16;
|
||||
@@ -7956,6 +8008,37 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
|
||||
// of "align", so recompute split_k based on that.
|
||||
split_kv = ROUNDUP_POW2(std::max(1u, KV / split_k), alignment);
|
||||
split_k = CEIL_DIV(KV, split_kv);
|
||||
xe_fa_opt = xe_fa_supported_platform && xe_fa_supported_usage && xe_fa_supported_dtype;
|
||||
if (xe_fa_opt) {
|
||||
std::lock_guard<std::mutex> guard(ctx->device->compile_mutex);
|
||||
const uint32_t split_p_size = 32;
|
||||
const size_t max_dim = (nek1 + split_p_size - 1) / split_p_size;
|
||||
const size_t p_dim = max_dim * split_p_size;
|
||||
auto& pipelines = ctx->device->pipeline_xe_fa_decode_dual_phases;
|
||||
auto it = pipelines.find({ (uint32_t)neq0, (uint32_t)nev0, qk_ratio, (uint32_t)neq1 });
|
||||
if (it != pipelines.end()) {
|
||||
xe_fa_pipeline_dual_phases = it->second;
|
||||
} else {
|
||||
pipelines[{(uint32_t)neq0, (uint32_t)nev0, qk_ratio, (uint32_t)neq1}] = xe_fa_pipeline_dual_phases = std::make_pair(std::make_shared<vk_pipeline_struct>(), std::make_shared<vk_pipeline_struct>());
|
||||
}
|
||||
|
||||
size_p = neq1 * neq2 * p_dim * neq3 * sizeof(ggml_fp16_t);
|
||||
size_group_max = neq1 * neq2 * max_dim * neq3 * sizeof(float);
|
||||
size_t temp_size = ggml_nelements(q) * sizeof(ggml_fp16_t) + size_p + size_group_max;
|
||||
fa_copy_qstate = true;
|
||||
if (ctx->prealloc_size_x < temp_size) {
|
||||
ctx->prealloc_size_x = temp_size;
|
||||
ggml_vk_preallocate_buffers(ctx, subctx);
|
||||
}
|
||||
|
||||
if (ctx->prealloc_x_need_sync) {
|
||||
ggml_vk_sync_buffers(ctx, subctx);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (xe_fa_opt == true) {
|
||||
use_mask_opt = false;
|
||||
}
|
||||
|
||||
// Reserve space for split_k temporaries. For each split x batch, we need to store the O matrix (D x ne1)
|
||||
@@ -8111,7 +8194,68 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
|
||||
mask_n_head_log2, m0, m1,
|
||||
gqa_ratio, split_kv, split_k };
|
||||
|
||||
if (split_k > 1) {
|
||||
if (xe_fa_opt && split_k > 1) {
|
||||
auto upper_power_of_2 = [&](uint32_t in) {
|
||||
GGML_ASSERT(in != 0);
|
||||
if (in <= 1) return 1u;
|
||||
uint32_t ret = in - 1;
|
||||
ret |= ret >> 1;
|
||||
ret |= ret >> 2;
|
||||
ret |= ret >> 4;
|
||||
ret |= ret >> 8;
|
||||
ret |= ret >> 16;
|
||||
return ret + 1;
|
||||
};
|
||||
auto to_fp16_vk_0 = ggml_vk_get_to_fp16(ctx, q->type);
|
||||
const uint32_t out_dim_per_wg = qk_ratio > 16 ? 8 : 16;
|
||||
size_t x_ne = ggml_nelements(q);
|
||||
size_t temp_buf_offset = 0;
|
||||
uint32_t head_stride_k = uint32_t(nbk2 / ggml_type_size(k->type));
|
||||
uint32_t head_stride_v = uint32_t(nbv2 / ggml_type_size(v->type));
|
||||
uint32_t batch_stride_q = uint32_t(nbq3 / ggml_type_size(q->type));
|
||||
uint32_t batch_stride_k = uint32_t(nbk3 / ggml_type_size(k->type));
|
||||
uint32_t batch_stride_v = uint32_t(nbv3 / ggml_type_size(v->type));
|
||||
uint32_t batch_stride_m = mask ? uint32_t(mask->nb[3] / ggml_type_size(mask->type)) : 0u;
|
||||
uint32_t batch_stride_o = uint32_t(nb3 / ggml_type_size(dst->type));
|
||||
vk_fa_xe_opt_push_constants pc_ph1 = { (uint32_t)nek1, (uint32_t)neq1, (uint32_t)neq2, (uint32_t)nek2, qk_ratio, 1, (sinks != nullptr) ? 1u : 0u, (uint32_t)k_stride, head_stride_k,
|
||||
batch_stride_q, batch_stride_k, batch_stride_v, batch_stride_m, batch_stride_o, scale };
|
||||
vk_fa_xe_opt_push_constants pc_ph2 = pc_ph1;
|
||||
pc_ph2.nbkv_tok = v_stride;
|
||||
pc_ph2.nbkv_head = head_stride_v;
|
||||
vk_subbuffer q_temp_buf = fa_copy_qstate ? ggml_vk_subbuffer(ctx, ctx->prealloc_x, temp_buf_offset) : q_buf;
|
||||
temp_buf_offset += fa_copy_qstate ? x_ne * sizeof(ggml_fp16_t) : 0;
|
||||
vk_subbuffer p_temp_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_x, temp_buf_offset);
|
||||
temp_buf_offset += size_p;
|
||||
vk_subbuffer max_temp_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_x, temp_buf_offset);
|
||||
temp_buf_offset += size_group_max;
|
||||
uint32_t xe_native_sub_group_size = ctx->device.get()->architecture == INTEL_XE1 ? 8 : 16;
|
||||
uint32_t aligned_gqa_ratio = upper_power_of_2(qk_ratio);
|
||||
uint32_t out_per_wg_ph1 = std::min(256u / xe_native_sub_group_size, (uint32_t)neq1);
|
||||
uint32_t out_per_wg_ph2 = std::min(std::max(16u / aligned_gqa_ratio, 1u), (uint32_t)neq1);
|
||||
uint32_t ph1_wg = ((neq1 + out_per_wg_ph1 - 1) / out_per_wg_ph1) * nek2;
|
||||
uint32_t ph2_wg = ((neq1 + out_per_wg_ph2 - 1) / out_per_wg_ph2) * ne0 / out_dim_per_wg;
|
||||
if (fa_copy_qstate) {
|
||||
const std::vector<uint32_t> pc_cpy_fp16 =
|
||||
{ (uint32_t)q->ne[0], (uint32_t)q->ne[1], (uint32_t)q->ne[2], (uint32_t)q->ne[3], (uint32_t)(x_ne) };
|
||||
ggml_vk_sync_buffers(ctx, subctx);
|
||||
ggml_pipeline_request_descriptor_sets(ctx, to_fp16_vk_0, 1);
|
||||
ggml_vk_dispatch_pipeline(ctx, subctx, to_fp16_vk_0, { q_buf, q_temp_buf }, pc_cpy_fp16, { (uint32_t)(x_ne), 1, 1 });
|
||||
}
|
||||
|
||||
ggml_vk_sync_buffers(ctx, subctx);
|
||||
ggml_pipeline_request_descriptor_sets(ctx, xe_fa_pipeline_dual_phases.first, 1);
|
||||
ggml_vk_dispatch_pipeline(ctx, subctx, xe_fa_pipeline_dual_phases.first,
|
||||
{ q_temp_buf, k_buf, mask_buf, p_temp_buf, max_temp_buf },
|
||||
pc_ph1, { (uint32_t)ph1_wg, (uint32_t)nek1, (uint32_t)neq3 });
|
||||
|
||||
ggml_vk_sync_buffers(ctx, subctx);
|
||||
ggml_pipeline_request_descriptor_sets(ctx, xe_fa_pipeline_dual_phases.second, 1);
|
||||
ggml_vk_dispatch_pipeline(ctx, subctx, xe_fa_pipeline_dual_phases.second,
|
||||
{ p_temp_buf, v_buf, max_temp_buf, sinks_buf, dst_buf },
|
||||
pc_ph2, { (uint32_t)ph2_wg, (uint32_t)nev2, (uint32_t)neq3 });
|
||||
|
||||
ctx->prealloc_x_need_sync = true;
|
||||
} else if (split_k > 1) {
|
||||
ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_split_k_reduce, 1);
|
||||
|
||||
if (ctx->prealloc_split_k_need_sync) {
|
||||
@@ -14953,6 +15097,9 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
|
||||
// If there's not enough shared memory for row_ids and the result tile, fallback to CPU
|
||||
return false;
|
||||
}
|
||||
if (ggml_get_op_params_i32(op, 3) == GGML_PREC_F32) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
switch (src0_type) {
|
||||
case GGML_TYPE_F32:
|
||||
|
||||
@@ -0,0 +1,263 @@
|
||||
#version 450
|
||||
|
||||
#extension GL_EXT_control_flow_attributes : enable
|
||||
#extension GL_EXT_shader_16bit_storage : require
|
||||
#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require
|
||||
#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require
|
||||
#extension GL_KHR_memory_scope_semantics : enable
|
||||
#extension GL_KHR_shader_subgroup_basic : enable
|
||||
#extension GL_KHR_shader_subgroup_ballot : enable
|
||||
#extension GL_KHR_shader_subgroup_arithmetic : enable
|
||||
#extension GL_KHR_cooperative_matrix : enable
|
||||
#extension GL_EXT_shared_memory_block : enable
|
||||
|
||||
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
|
||||
|
||||
layout (binding = 0) readonly buffer Q {float16_t qState[];};
|
||||
layout (binding = 1) readonly buffer K_VEC4 {f16vec4 kStateVec4[];};
|
||||
layout (binding = 2) buffer MASK_F16 {float16_t mState_f16[];};
|
||||
layout (binding = 3) buffer P_FP16 {float16_t matP_f16[];};
|
||||
layout (binding = 4) buffer OUT_MAX {float out_max_f32[];};
|
||||
|
||||
layout (push_constant) uniform parameter
|
||||
{
|
||||
uint kvSeqLen;
|
||||
uint activationLength;
|
||||
uint qHead;
|
||||
uint kvHead;
|
||||
uint qkRatio;
|
||||
uint qkSubGroups;
|
||||
uint flag;
|
||||
uint kvStride1;
|
||||
uint kvStride2;
|
||||
uint batchStrideQ;
|
||||
uint batchStrideK;
|
||||
uint batchStrideV;
|
||||
uint batchStrideM;
|
||||
uint batchStrideO;
|
||||
float softMaxScale;
|
||||
} p;
|
||||
|
||||
layout (constant_id = 0) const uint GROUPSIZE = 128;
|
||||
layout (constant_id = 1) const uint GQA_RATIO = 8;
|
||||
layout (constant_id = 2) const uint HEAD_DIM = 128;
|
||||
layout (constant_id = 3) const uint WARPSIZE = 16;
|
||||
layout (constant_id = 4) const uint MATP_REDUCE = 32;
|
||||
layout (constant_id = 5) const uint N_TOK = 1;
|
||||
layout (constant_id = 6) const uint COOP_MAT_P_PER_LOOP = 4;
|
||||
|
||||
#define MAX_HEADS 8
|
||||
|
||||
#define TN WARPSIZE
|
||||
#define TM 8
|
||||
#define TK 16
|
||||
#define SUBGROUP_COUNT (GROUPSIZE / WARPSIZE)
|
||||
#define MATP_PER_LOOP (COOP_MAT_P_PER_LOOP * TM)
|
||||
#define P_LOOP_COUNT (MATP_REDUCE / MATP_PER_LOOP)
|
||||
|
||||
#define COOP_MAT_Q_PER_TOKEN ((GQA_RATIO + TN - 1) / TN)
|
||||
#define COOP_MAT_P_M COOP_MAT_Q_PER_TOKEN
|
||||
#define COOP_MAT_P_N (MATP_REDUCE / TM)
|
||||
#define SLM_PV_SIZE (MATP_REDUCE * COOP_MAT_P_M * TN)
|
||||
#define SLM_MASK_SIZE (N_TOK * MATP_REDUCE)
|
||||
#define SLM_POOL_SIZE_K (MATP_PER_LOOP * HEAD_DIM)
|
||||
#define K_LOAD_PER_LOOP (GROUPSIZE * 4)
|
||||
#define HEAD_DIM_VEC4 (HEAD_DIM / 4)
|
||||
#define SLM_CHUNK_SIZE (TK / 4)
|
||||
#define K_LOAD_LOOPS ((SLM_POOL_SIZE_K + K_LOAD_PER_LOOP - 1) / K_LOAD_PER_LOOP)
|
||||
#define O_COUNT ((GQA_RATIO + SUBGROUP_COUNT - 1) / SUBGROUP_COUNT)
|
||||
|
||||
shared slm_pool_block {
|
||||
float slm_pool_pv[SLM_PV_SIZE + SLM_MASK_SIZE];
|
||||
} slm_pool_f32;
|
||||
|
||||
shared slm_pool_alias_block {
|
||||
float16_t slm_pool_k[SLM_POOL_SIZE_K];
|
||||
} slm_pool_f16;
|
||||
|
||||
void main() {
|
||||
const uint lane = gl_SubgroupInvocationID;
|
||||
const uint kHeadIdx = gl_WorkGroupID.x % p.kvHead;
|
||||
const uint outGroupIdx = gl_WorkGroupID.x / p.kvHead;
|
||||
const uint v = gl_WorkGroupID.y;
|
||||
const uint d = gl_WorkGroupID.z;
|
||||
const uint localLinearId = gl_SubgroupID;
|
||||
const uint wgLane = localLinearId * WARPSIZE + lane;
|
||||
const uint qDim = p.qHead * HEAD_DIM;
|
||||
const uint kvDim = p.kvStride1;
|
||||
const uint maskDim = p.kvSeqLen;
|
||||
const uint maxDim = (p.kvSeqLen + MATP_REDUCE - 1) / MATP_REDUCE;
|
||||
const uint pDim = maxDim * MATP_REDUCE;
|
||||
const uint tokFlatIdx = localLinearId + outGroupIdx * N_TOK;
|
||||
uint offsetBaseQ = min(tokFlatIdx, p.activationLength - 1) * qDim;
|
||||
offsetBaseQ = offsetBaseQ + d * p.batchStrideQ + kHeadIdx * HEAD_DIM * GQA_RATIO;
|
||||
const uint offsetBaseK = (d * p.batchStrideK + (v * MATP_REDUCE) * kvDim + kHeadIdx * p.kvStride2) / 4;
|
||||
uint offsetOut = d * p.qHead * p.activationLength * pDim + v * MATP_REDUCE + kHeadIdx * GQA_RATIO * pDim + (localLinearId * O_COUNT + outGroupIdx * N_TOK * p.qHead) * pDim + lane;
|
||||
uint offsetMax = d * p.qHead * p.activationLength * maxDim + v + kHeadIdx * GQA_RATIO * maxDim + (localLinearId * O_COUNT + outGroupIdx * N_TOK * p.qHead) * maxDim;
|
||||
const uint offsetSlmLoadPv = (localLinearId * O_COUNT * MATP_REDUCE + lane);
|
||||
const uint offsetBaseM = v * MATP_REDUCE + lane;
|
||||
const float fp32Min = uintBitsToFloat(0xFEFFFFFF);
|
||||
|
||||
const uint loopCount = HEAD_DIM / TK;
|
||||
float maskFp32[MATP_REDUCE / WARPSIZE];
|
||||
|
||||
if (tokFlatIdx < p.activationLength) {
|
||||
[[unroll]] for (uint mk = 0; mk < MATP_REDUCE / WARPSIZE; mk++) {
|
||||
const uint maskOffset = mk * WARPSIZE + offsetBaseM;
|
||||
if (maskOffset < maskDim) {
|
||||
maskFp32[mk] = float(mState_f16[d * p.batchStrideM + tokFlatIdx * maskDim + maskOffset]);
|
||||
} else {
|
||||
maskFp32[mk] = fp32Min;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
coopmat<float, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> matP[COOP_MAT_P_M][COOP_MAT_P_N];
|
||||
|
||||
[[unroll]] for (uint mp = 0; mp < COOP_MAT_P_M; mp++) {
|
||||
[[unroll]] for (uint np = 0; np < COOP_MAT_P_N; np++) {
|
||||
matP[mp][np] = coopmat<float, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(0.0f);
|
||||
}
|
||||
}
|
||||
|
||||
[[unroll]] for (uint kLoad = 0; kLoad < K_LOAD_LOOPS; kLoad++) {
|
||||
const uint flatOffset = kLoad * GROUPSIZE + wgLane;
|
||||
const uint kRowIdx = flatOffset / HEAD_DIM_VEC4;
|
||||
const uint kColIdx = flatOffset % HEAD_DIM_VEC4;
|
||||
const uint slmChunkCol = kColIdx % SLM_CHUNK_SIZE;
|
||||
const uint slmChunkRow = kColIdx / SLM_CHUNK_SIZE;
|
||||
const uint offsetK = offsetBaseK + kRowIdx * kvDim / 4 + kColIdx;
|
||||
const uint offsetSlmK = kRowIdx * TK + slmChunkRow * TK * MATP_PER_LOOP + slmChunkCol * 4;
|
||||
slm_pool_f16.slm_pool_k[offsetSlmK + 0] = kStateVec4[offsetK].x;
|
||||
slm_pool_f16.slm_pool_k[offsetSlmK + 1] = kStateVec4[offsetK].y;
|
||||
slm_pool_f16.slm_pool_k[offsetSlmK + 2] = kStateVec4[offsetK].z;
|
||||
slm_pool_f16.slm_pool_k[offsetSlmK + 3] = kStateVec4[offsetK].w;
|
||||
}
|
||||
|
||||
[[unroll]] for (uint pLoop = 0; pLoop < P_LOOP_COUNT; pLoop++) {
|
||||
f16vec4 kTemp[K_LOAD_LOOPS];
|
||||
|
||||
if (pLoop + 1 < P_LOOP_COUNT) {
|
||||
[[unroll]] for (uint kLoad = 0; kLoad < K_LOAD_LOOPS; kLoad++) {
|
||||
const uint flatOffset = kLoad * GROUPSIZE + wgLane;
|
||||
const uint kRowIdx = flatOffset / HEAD_DIM_VEC4 + (pLoop + 1) * MATP_PER_LOOP;
|
||||
const uint kColIdx = flatOffset % HEAD_DIM_VEC4;
|
||||
const uint offsetK = offsetBaseK + kRowIdx * kvDim / 4 + kColIdx;
|
||||
kTemp[kLoad] = kStateVec4[offsetK];
|
||||
}
|
||||
}
|
||||
|
||||
barrier();
|
||||
if (localLinearId < N_TOK) {
|
||||
[[unroll]] for (uint loop = 0; loop < loopCount; loop++) {
|
||||
coopmat<float16_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> matQ[COOP_MAT_P_M];
|
||||
coopmat<float16_t, gl_ScopeSubgroup, TM, TK, gl_MatrixUseA> matK[COOP_MAT_P_PER_LOOP];
|
||||
|
||||
[[unroll]] for (uint mq = 0; mq < COOP_MAT_P_M; mq++) {
|
||||
coopMatLoad(
|
||||
matQ[mq],
|
||||
qState,
|
||||
offsetBaseQ + mq * TN * HEAD_DIM + loop * TK,
|
||||
HEAD_DIM,
|
||||
gl_CooperativeMatrixLayoutColumnMajor);
|
||||
}
|
||||
|
||||
[[unroll]] for (uint np = 0; np < COOP_MAT_P_PER_LOOP; np++) {
|
||||
coopMatLoad(
|
||||
matK[np],
|
||||
slm_pool_f16.slm_pool_k,
|
||||
loop * TK * MATP_PER_LOOP + np * TM * TK,
|
||||
TK,
|
||||
gl_CooperativeMatrixLayoutRowMajor);
|
||||
}
|
||||
|
||||
[[unroll]] for (uint mp = 0; mp < COOP_MAT_P_M; mp++) {
|
||||
[[unroll]] for (uint np = 0; np < COOP_MAT_P_PER_LOOP; np++) {
|
||||
matP[mp][pLoop * COOP_MAT_P_PER_LOOP + np] = coopMatMulAdd(matK[np], matQ[mp], matP[mp][pLoop * COOP_MAT_P_PER_LOOP + np]);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
barrier();
|
||||
|
||||
if (pLoop + 1 < P_LOOP_COUNT) {
|
||||
[[unroll]] for (uint kLoad = 0; kLoad < K_LOAD_LOOPS; kLoad++) {
|
||||
const uint flatOffset = kLoad * GROUPSIZE + wgLane;
|
||||
const uint kRowIdx = flatOffset / HEAD_DIM_VEC4;
|
||||
const uint kColIdx = flatOffset % HEAD_DIM_VEC4;
|
||||
const uint slmChunkCol = kColIdx % SLM_CHUNK_SIZE;
|
||||
const uint slmChunkRow = kColIdx / SLM_CHUNK_SIZE;
|
||||
const uint offsetSlmK = kRowIdx * TK + slmChunkRow * TK * MATP_PER_LOOP + slmChunkCol * 4;
|
||||
slm_pool_f16.slm_pool_k[offsetSlmK + 0] = kTemp[kLoad].x;
|
||||
slm_pool_f16.slm_pool_k[offsetSlmK + 1] = kTemp[kLoad].y;
|
||||
slm_pool_f16.slm_pool_k[offsetSlmK + 2] = kTemp[kLoad].z;
|
||||
slm_pool_f16.slm_pool_k[offsetSlmK + 3] = kTemp[kLoad].w;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
barrier();
|
||||
|
||||
if (tokFlatIdx < p.activationLength) {
|
||||
[[unroll]] for (uint mk = 0; mk < MATP_REDUCE / WARPSIZE; mk++) {
|
||||
slm_pool_f32.slm_pool_pv[SLM_PV_SIZE + localLinearId * MATP_REDUCE + mk * WARPSIZE + lane] = maskFp32[mk];
|
||||
}
|
||||
}
|
||||
|
||||
[[unroll]] for (uint oLoop = 0; oLoop < N_TOK; oLoop++) {
|
||||
if (oLoop + outGroupIdx * N_TOK < p.activationLength) {
|
||||
if (localLinearId == oLoop) {
|
||||
[[unroll]] for (uint mp = 0; mp < COOP_MAT_P_M; mp++) {
|
||||
[[unroll]] for (uint np = 0; np < COOP_MAT_P_N; np++) {
|
||||
coopMatStore(matP[mp][np], slm_pool_f32.slm_pool_pv, mp * MATP_REDUCE * TN + np * TM, MATP_REDUCE, gl_CooperativeMatrixLayoutColumnMajor);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
barrier();
|
||||
|
||||
[[unroll]] for (uint maskIdx = 0; maskIdx < MATP_REDUCE / WARPSIZE; maskIdx++) {
|
||||
maskFp32[maskIdx] = slm_pool_f32.slm_pool_pv[SLM_PV_SIZE + oLoop * MATP_REDUCE + maskIdx * WARPSIZE + lane];
|
||||
}
|
||||
|
||||
float fp32O[O_COUNT][MATP_REDUCE / WARPSIZE];
|
||||
float maxOut[O_COUNT];
|
||||
|
||||
[[unroll]] for (uint oc = 0; oc < O_COUNT; oc++) {
|
||||
[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
|
||||
fp32O[oc][os] = slm_pool_f32.slm_pool_pv[offsetSlmLoadPv + os * WARPSIZE + oc * MATP_REDUCE] * p.softMaxScale;
|
||||
}
|
||||
|
||||
[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
|
||||
fp32O[oc][os] = fp32O[oc][os] + maskFp32[os];
|
||||
}
|
||||
|
||||
float maxTemp = fp32Min;
|
||||
[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
|
||||
maxTemp = max(maxTemp, fp32O[oc][os]);
|
||||
}
|
||||
maxOut[oc] = subgroupMax(maxTemp);
|
||||
[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
|
||||
fp32O[oc][os] = exp(fp32O[oc][os] - maxOut[oc]);
|
||||
}
|
||||
}
|
||||
|
||||
[[unroll]] for (uint oc = 0; oc < O_COUNT; oc++) {
|
||||
if (localLinearId * O_COUNT + oc < GQA_RATIO) {
|
||||
[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
|
||||
matP_f16[offsetOut + oc * pDim + os * WARPSIZE] = float16_t(fp32O[oc][os]);
|
||||
}
|
||||
|
||||
if (lane == 0) {
|
||||
out_max_f32[offsetMax + oc * maxDim] = maxOut[oc];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
offsetOut = offsetOut + p.qHead * pDim;
|
||||
offsetMax = offsetMax + p.qHead * maxDim;
|
||||
barrier();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,408 @@
|
||||
#version 450
|
||||
|
||||
#extension GL_EXT_control_flow_attributes : enable
|
||||
#extension GL_EXT_shader_16bit_storage : require
|
||||
#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require
|
||||
#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require
|
||||
#extension GL_KHR_memory_scope_semantics : enable
|
||||
#extension GL_KHR_shader_subgroup_basic : enable
|
||||
#extension GL_KHR_shader_subgroup_ballot : enable
|
||||
#extension GL_KHR_shader_subgroup_arithmetic : enable
|
||||
#extension GL_KHR_cooperative_matrix : enable
|
||||
#extension GL_EXT_shared_memory_block : enable
|
||||
|
||||
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
|
||||
|
||||
layout (binding = 0) readonly buffer P {f16vec4 pStateVec4[];};
|
||||
layout (binding = 1) readonly buffer V {float16_t vState[];};
|
||||
layout (binding = 1) readonly buffer V_VEC4 {f16vec4 vStateVec4[];};
|
||||
layout (binding = 2) buffer MAX_FP32 {float max_f32[];};
|
||||
layout (binding = 3) buffer SINK_FP32 {float sink_f32[];};
|
||||
layout (binding = 4) buffer OUT_FP32 {float out_f32[];};
|
||||
layout (binding = 4) buffer OUT_VEC4 {vec4 out_f32_vec4[];};
|
||||
layout (binding = 4) buffer OUT_F16 {float16_t out_f16[];};
|
||||
|
||||
layout (push_constant) uniform parameter
|
||||
{
|
||||
uint kvSeqLen;
|
||||
uint activationLength;
|
||||
uint qHead;
|
||||
uint kvHead;
|
||||
uint qkRatio;
|
||||
uint qkSubGroups;
|
||||
uint flag;
|
||||
uint kvStride1;
|
||||
uint kvStride2;
|
||||
uint batchStrideQ;
|
||||
uint batchStrideK;
|
||||
uint batchStrideV;
|
||||
uint batchStrideM;
|
||||
uint batchStrideO;
|
||||
float softMaxScale;
|
||||
} p;
|
||||
|
||||
layout (constant_id = 0) const uint GROUPSIZE = 256;
|
||||
layout (constant_id = 1) const uint GQA_RATIO = 8;
|
||||
layout (constant_id = 2) const uint HEAD_DIM = 128;
|
||||
layout (constant_id = 3) const uint N_TOKS_PER_GROUP = 1;
|
||||
layout (constant_id = 4) const uint WARPSIZE = 16;
|
||||
layout (constant_id = 5) const uint MATP_PER_LOOP = 64;
|
||||
layout (constant_id = 6) const uint MATP_REDUCE = 32;
|
||||
layout (constant_id = 7) const uint WARP_V_DIM = 16;
|
||||
|
||||
#define TN WARPSIZE
|
||||
#define TM 8
|
||||
#define TK 16
|
||||
#define MAT_O_N (WARP_V_DIM / TM)
|
||||
#define MAT_P_M (GQA_RATIO * N_TOKS_PER_GROUP)
|
||||
#define ALIGNED_P_M ((MAT_P_M + WARPSIZE - 1) / WARPSIZE)
|
||||
#define V_HEAD_GROUPS (HEAD_DIM / WARP_V_DIM)
|
||||
|
||||
#define SUBGROUP_COUNT (GROUPSIZE / WARPSIZE)
|
||||
#define SPLIT_P_GROUPS (MATP_PER_LOOP / TK)
|
||||
|
||||
#define SLM_POOL_SIZE_O (SUBGROUP_COUNT * ALIGNED_P_M * TN * MAT_O_N * TM)
|
||||
|
||||
#define P_LOAD_PER_LOOP (GROUPSIZE * 4)
|
||||
#define P_LOAD_LOOPS ((MAT_P_M * MATP_PER_LOOP + P_LOAD_PER_LOOP - 1) / P_LOAD_PER_LOOP)
|
||||
#define SLM_POOL_SIZE_P (P_LOAD_LOOPS * P_LOAD_PER_LOOP)
|
||||
#define SIZE_LOCAL_MAX (MAT_P_M * MATP_PER_LOOP / MATP_REDUCE)
|
||||
#define MAX_LOAD_LOOPS ((SIZE_LOCAL_MAX + GROUPSIZE - 1) / GROUPSIZE)
|
||||
#define SLM_POOL_SIZE_LOCAL_MAX (MAX_LOAD_LOOPS * GROUPSIZE)
|
||||
#define MAX_REDUCE_COUNT ((MAT_P_M + SUBGROUP_COUNT - 1) / SUBGROUP_COUNT)
|
||||
#define GLOBAL_MAX_SIZE (MAX_REDUCE_COUNT * SUBGROUP_COUNT)
|
||||
|
||||
#define SLM_POOL_SIZE_SOFTMAX_SUM (SUBGROUP_COUNT * P_LOAD_LOOPS)
|
||||
|
||||
#define SLM_OFFSET_P (GLOBAL_MAX_SIZE * 2 + SLM_POOL_SIZE_SOFTMAX_SUM * 2 + SLM_POOL_SIZE_LOCAL_MAX * 2 * 2)
|
||||
|
||||
#define SLM_OFFSET_GLOBAL_MAX 0
|
||||
#define SLM_OFFSET_SOFTMAX_SUM (SLM_OFFSET_GLOBAL_MAX + GLOBAL_MAX_SIZE)
|
||||
#define SLM_OFFSET_O (SLM_OFFSET_SOFTMAX_SUM + SLM_POOL_SIZE_SOFTMAX_SUM)
|
||||
#define SLM_OFFSET_LOCAL_MAX (GLOBAL_MAX_SIZE + SLM_POOL_SIZE_SOFTMAX_SUM)
|
||||
|
||||
#define P_REDUCE_VEC4 (MATP_PER_LOOP / 4)
|
||||
#define MAX_PER_LOOP (MATP_PER_LOOP / MATP_REDUCE)
|
||||
#define SLM_MAX_STRIDE (MATP_REDUCE / 4)
|
||||
#define SUB_GROUPS_PER_LINE (MATP_PER_LOOP / WARPSIZE / 4)
|
||||
|
||||
shared slm_pool_block {
|
||||
float slm_pool_o[GLOBAL_MAX_SIZE + SLM_POOL_SIZE_SOFTMAX_SUM + SLM_POOL_SIZE_O];
|
||||
} slm_pool_f32;
|
||||
|
||||
shared slm_pool_alias_block {
|
||||
float16_t slm_pool_pv[GLOBAL_MAX_SIZE * 2 + SLM_POOL_SIZE_SOFTMAX_SUM * 2 + SLM_POOL_SIZE_LOCAL_MAX * 2 * 2 + SLM_POOL_SIZE_P * 2];
|
||||
} slm_pool_alias_f16;
|
||||
|
||||
void main() {
|
||||
const uint lane = gl_SubgroupInvocationID;
|
||||
const uint v = gl_WorkGroupID.y;
|
||||
const uint d = gl_WorkGroupID.z;
|
||||
const uint vWarpIdx = gl_WorkGroupID.x % V_HEAD_GROUPS;
|
||||
const uint outTokIdx = gl_WorkGroupID.x / V_HEAD_GROUPS;
|
||||
const uint localLinearId = gl_SubgroupID;
|
||||
const uint wgLane = localLinearId * WARPSIZE + lane;
|
||||
const uint splitIdx = localLinearId;
|
||||
const uint maxDim = (p.kvSeqLen + MATP_REDUCE - 1) / MATP_REDUCE;
|
||||
const uint pDim = maxDim * MATP_REDUCE;
|
||||
const uint kvDim = p.kvStride1;
|
||||
const uint oDim = p.qHead * HEAD_DIM;
|
||||
const uint offsetBaseP = (d * p.activationLength * p.qHead + v * GQA_RATIO + outTokIdx * N_TOKS_PER_GROUP * p.qHead) * pDim / 4;
|
||||
const uint offsetBaseMax = (d * p.activationLength * p.qHead + v * GQA_RATIO + outTokIdx * N_TOKS_PER_GROUP * p.qHead) * maxDim;
|
||||
const uint offsetBaseV = (d * p.batchStrideV + v * p.kvStride2 + vWarpIdx * WARP_V_DIM + splitIdx * TK * kvDim);
|
||||
const uint offsetSlmP = (SLM_OFFSET_P + wgLane * 4);
|
||||
const float fp32Min = uintBitsToFloat(0xFEFFFFFF);
|
||||
const float fp32Max = uintBitsToFloat(0x7EFFFFFF);
|
||||
uint offsetV = offsetBaseV;
|
||||
|
||||
coopmat<float, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> sums[ALIGNED_P_M][MAT_O_N];
|
||||
f16vec4 pStateTemp[P_LOAD_LOOPS];
|
||||
|
||||
float fp32CompensationP[P_LOAD_LOOPS];
|
||||
|
||||
uint loadRowBase[P_LOAD_LOOPS];
|
||||
uint loadColBase[P_LOAD_LOOPS];
|
||||
float fp32SoftMaxSum[P_LOAD_LOOPS];
|
||||
float fp32GlobalMaxP[P_LOAD_LOOPS];
|
||||
uint maxRowBase[MAX_LOAD_LOOPS];
|
||||
uint maxColBase[MAX_LOAD_LOOPS];
|
||||
uint outOffsets[ALIGNED_P_M];
|
||||
bool outputMask[ALIGNED_P_M];
|
||||
float fp32SinkCoeff[ALIGNED_P_M];
|
||||
|
||||
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
|
||||
const uint flatOffset = pm * WARPSIZE + lane;
|
||||
const uint inGroupTokIdx = flatOffset / GQA_RATIO;
|
||||
const uint inGroupHeadIdx = flatOffset % GQA_RATIO;
|
||||
outputMask[pm] = (N_TOKS_PER_GROUP * outTokIdx + inGroupTokIdx < p.activationLength) && (inGroupHeadIdx < GQA_RATIO) && (inGroupTokIdx < N_TOKS_PER_GROUP);
|
||||
outOffsets[pm] = (inGroupTokIdx * oDim + inGroupHeadIdx * HEAD_DIM) / 4;
|
||||
if ((0x1 & p.flag) != 0) {
|
||||
fp32SinkCoeff[pm] = sink_f32[inGroupHeadIdx + v * GQA_RATIO];
|
||||
}
|
||||
}
|
||||
|
||||
[[unroll]] for (uint maxCount = 0; maxCount < MAX_REDUCE_COUNT; maxCount++) {
|
||||
const uint flatIdx = maxCount * SUBGROUP_COUNT + localLinearId;
|
||||
const uint rowIdx = flatIdx % GQA_RATIO;
|
||||
const uint tokIdx = flatIdx / GQA_RATIO;
|
||||
|
||||
if (tokIdx < N_TOKS_PER_GROUP) {
|
||||
float fp32MaxReduce = fp32Min;
|
||||
const uint maxOffset = offsetBaseMax + (tokIdx * p.qHead + rowIdx) * maxDim;
|
||||
[[unroll]] for (uint maxReduce = 0; maxReduce < (maxDim + WARPSIZE - 1) / WARPSIZE; maxReduce++) {
|
||||
if (maxReduce * WARPSIZE + lane < maxDim) {
|
||||
fp32MaxReduce = max(fp32MaxReduce, max_f32[maxOffset + maxReduce * WARPSIZE + lane]);
|
||||
}
|
||||
}
|
||||
fp32MaxReduce = subgroupMax(fp32MaxReduce);
|
||||
if (lane == 0) {
|
||||
slm_pool_f32.slm_pool_o[SLM_OFFSET_GLOBAL_MAX + maxCount * SUBGROUP_COUNT + localLinearId] = fp32MaxReduce;
|
||||
}
|
||||
} else {
|
||||
if (lane == 0) {
|
||||
slm_pool_f32.slm_pool_o[SLM_OFFSET_GLOBAL_MAX + maxCount * SUBGROUP_COUNT + localLinearId] = fp32Max;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
barrier();
|
||||
|
||||
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
|
||||
const uint flatOffset = (pLoad * GROUPSIZE + wgLane) / P_REDUCE_VEC4;
|
||||
const uint rowIdxFlat = flatOffset % GQA_RATIO;
|
||||
const uint tokenIdxFlat = min(flatOffset / GQA_RATIO, N_TOKS_PER_GROUP - 1);
|
||||
loadColBase[pLoad] = (pLoad * GROUPSIZE + wgLane) % P_REDUCE_VEC4;
|
||||
loadRowBase[pLoad] = (tokenIdxFlat * p.qHead + rowIdxFlat);
|
||||
fp32SoftMaxSum[pLoad] = 0.0f;
|
||||
fp32GlobalMaxP[pLoad] = slm_pool_f32.slm_pool_o[SLM_OFFSET_GLOBAL_MAX + flatOffset];
|
||||
}
|
||||
|
||||
[[unroll]] for (uint maxLoad = 0; maxLoad < MAX_LOAD_LOOPS; maxLoad++) {
|
||||
const uint flatOffset = (maxLoad * GROUPSIZE + wgLane) / MAX_PER_LOOP;
|
||||
const uint rowIdxFlat = flatOffset % GQA_RATIO;
|
||||
const uint tokenIdxFlat = min(flatOffset / GQA_RATIO, N_TOKS_PER_GROUP - 1);
|
||||
maxColBase[maxLoad] = (maxLoad * GROUPSIZE + wgLane) % MAX_PER_LOOP;
|
||||
maxRowBase[maxLoad] = (tokenIdxFlat * p.qHead + rowIdxFlat);
|
||||
}
|
||||
|
||||
[[unroll]] for (uint maxLoad = 0; maxLoad < MAX_LOAD_LOOPS; maxLoad++) {
|
||||
const uint flatMaxOffset = maxRowBase[maxLoad] * maxDim + maxColBase[maxLoad];
|
||||
slm_pool_f32.slm_pool_o[SLM_OFFSET_LOCAL_MAX + maxLoad * GROUPSIZE + wgLane] = max_f32[offsetBaseMax + flatMaxOffset];
|
||||
maxColBase[maxLoad] = maxColBase[maxLoad] + MATP_PER_LOOP / MATP_REDUCE;
|
||||
}
|
||||
|
||||
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
|
||||
const uint flatOffset = loadRowBase[pLoad] * pDim / 4 + loadColBase[pLoad];
|
||||
pStateTemp[pLoad] = pStateVec4[offsetBaseP + flatOffset];
|
||||
}
|
||||
|
||||
[[unroll]] for (uint n = 0; n < ALIGNED_P_M; n++) {
|
||||
[[unroll]] for (uint i = 0; i < MAT_O_N; i++) {
|
||||
sums[n][i] = coopmat<float, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(0.0f);
|
||||
}
|
||||
}
|
||||
|
||||
barrier();
|
||||
|
||||
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
|
||||
const uint maxOffset = (pLoad * GROUPSIZE + wgLane) / SLM_MAX_STRIDE;
|
||||
if (loadColBase[pLoad] < pDim / 4) {
|
||||
fp32CompensationP[pLoad] = slm_pool_f32.slm_pool_o[SLM_OFFSET_LOCAL_MAX + maxOffset];
|
||||
float pTemp[4] = float[4](pStateTemp[pLoad].x, pStateTemp[pLoad].y, pStateTemp[pLoad].z, pStateTemp[pLoad].w);
|
||||
float compTemp = exp(fp32CompensationP[pLoad] - fp32GlobalMaxP[pLoad]);
|
||||
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
|
||||
pTemp[kk] = pTemp[kk] * compTemp;
|
||||
fp32SoftMaxSum[pLoad] = fp32SoftMaxSum[pLoad] + pTemp[kk];
|
||||
slm_pool_alias_f16.slm_pool_pv[offsetSlmP + pLoad * GROUPSIZE * 4 + kk] = float16_t(pTemp[kk]);
|
||||
}
|
||||
} else {
|
||||
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
|
||||
slm_pool_alias_f16.slm_pool_pv[offsetSlmP + pLoad * GROUPSIZE * 4 + kk] = float16_t(0.0f);
|
||||
}
|
||||
}
|
||||
|
||||
loadColBase[pLoad] = loadColBase[pLoad] + P_REDUCE_VEC4;
|
||||
}
|
||||
|
||||
const uint loopCount = (p.kvSeqLen + MATP_PER_LOOP - 1) / MATP_PER_LOOP;
|
||||
|
||||
for (uint loop = 0; loop < loopCount; loop++) {
|
||||
const uint slmPingPongLoad = (loop & 0x1);
|
||||
const uint slmPingPongStore = ((loop + 1) & 0x1);
|
||||
|
||||
if (loop + 1 < loopCount) {
|
||||
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
|
||||
const uint flatOffset = loadRowBase[pLoad] * pDim / 4 + loadColBase[pLoad];
|
||||
pStateTemp[pLoad] = pStateVec4[offsetBaseP + flatOffset];
|
||||
}
|
||||
|
||||
[[unroll]] for (uint maxLoad = 0; maxLoad < MAX_LOAD_LOOPS; maxLoad++) {
|
||||
const uint flatMaxOffset = maxRowBase[maxLoad] * maxDim + maxColBase[maxLoad];
|
||||
slm_pool_f32.slm_pool_o[SLM_OFFSET_LOCAL_MAX + slmPingPongStore * SLM_POOL_SIZE_LOCAL_MAX + maxLoad * GROUPSIZE + wgLane] = max_f32[offsetBaseMax + flatMaxOffset];
|
||||
maxColBase[maxLoad] = maxColBase[maxLoad] + MATP_PER_LOOP / MATP_REDUCE;
|
||||
}
|
||||
}
|
||||
|
||||
barrier();
|
||||
|
||||
{
|
||||
const uint coopMatOffsetP = SLM_OFFSET_P + slmPingPongLoad * SLM_POOL_SIZE_P + splitIdx * TK;
|
||||
coopmat<float16_t, gl_ScopeSubgroup, TM, TK, gl_MatrixUseA> matV[MAT_O_N];
|
||||
[[unroll]] for (uint cc = 0; cc < MAT_O_N; cc++) {
|
||||
coopMatLoad(
|
||||
matV[cc],
|
||||
vState,
|
||||
offsetV + TM * cc,
|
||||
kvDim,
|
||||
gl_CooperativeMatrixLayoutColumnMajor);
|
||||
}
|
||||
[[unroll]] for (uint mo = 0; mo < ALIGNED_P_M; mo++) {
|
||||
coopmat<float16_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> matP;
|
||||
coopMatLoad(
|
||||
matP,
|
||||
slm_pool_alias_f16.slm_pool_pv,
|
||||
coopMatOffsetP + mo * TN * MATP_PER_LOOP,
|
||||
MATP_PER_LOOP,
|
||||
gl_CooperativeMatrixLayoutColumnMajor);
|
||||
|
||||
[[unroll]] for (uint no = 0; no < MAT_O_N; no++) {
|
||||
sums[mo][no] = coopMatMulAdd(matV[no], matP, sums[mo][no]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
offsetV += MATP_PER_LOOP * kvDim;
|
||||
if (loop * MATP_PER_LOOP + splitIdx * TK >= p.kvSeqLen) {
|
||||
offsetV = 0;
|
||||
}
|
||||
if (loop + 1 < loopCount) {
|
||||
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
|
||||
const uint maxOffset = (pLoad * GROUPSIZE + wgLane) / SLM_MAX_STRIDE;
|
||||
if (loadColBase[pLoad] < pDim / 4) {
|
||||
fp32CompensationP[pLoad] = slm_pool_f32.slm_pool_o[SLM_OFFSET_LOCAL_MAX + slmPingPongStore * SLM_POOL_SIZE_LOCAL_MAX + maxOffset];
|
||||
float pTemp[4] = float[4](pStateTemp[pLoad].x, pStateTemp[pLoad].y, pStateTemp[pLoad].z, pStateTemp[pLoad].w);
|
||||
float compTemp = exp(fp32CompensationP[pLoad] - fp32GlobalMaxP[pLoad]);
|
||||
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
|
||||
pTemp[kk] = pTemp[kk] * compTemp;
|
||||
fp32SoftMaxSum[pLoad] = fp32SoftMaxSum[pLoad] + pTemp[kk];
|
||||
slm_pool_alias_f16.slm_pool_pv[offsetSlmP + slmPingPongStore * SLM_POOL_SIZE_P + pLoad * GROUPSIZE * 4 + kk] = float16_t(pTemp[kk]);
|
||||
}
|
||||
} else {
|
||||
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
|
||||
slm_pool_alias_f16.slm_pool_pv[offsetSlmP + slmPingPongStore * SLM_POOL_SIZE_P + pLoad * GROUPSIZE * 4 + kk] = float16_t(0.0f);
|
||||
}
|
||||
}
|
||||
loadColBase[pLoad] = loadColBase[pLoad] + P_REDUCE_VEC4;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
barrier();
|
||||
|
||||
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
|
||||
fp32SoftMaxSum[pLoad] = subgroupAdd(fp32SoftMaxSum[pLoad]);
|
||||
}
|
||||
|
||||
[[unroll]] for (uint mo = 0; mo < ALIGNED_P_M; mo++) {
|
||||
[[unroll]] for (uint no = 0; no < MAT_O_N; no++) {
|
||||
coopMatStore(
|
||||
sums[mo][no],
|
||||
slm_pool_f32.slm_pool_o,
|
||||
SLM_OFFSET_O + mo * TN * WARP_V_DIM + TM * no + localLinearId * ALIGNED_P_M * TN * WARP_V_DIM,
|
||||
WARP_V_DIM,
|
||||
gl_CooperativeMatrixLayoutColumnMajor);
|
||||
}
|
||||
}
|
||||
|
||||
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
|
||||
slm_pool_f32.slm_pool_o[SLM_OFFSET_SOFTMAX_SUM + pLoad * SUBGROUP_COUNT + localLinearId] = fp32SoftMaxSum[pLoad];
|
||||
}
|
||||
|
||||
barrier();
|
||||
|
||||
if (localLinearId == 1) {
|
||||
const uint sumBase = SLM_OFFSET_SOFTMAX_SUM + lane * SUB_GROUPS_PER_LINE;
|
||||
float sumTemp[ALIGNED_P_M][SUB_GROUPS_PER_LINE];
|
||||
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
|
||||
[[unroll]] for (uint reduce = 0; reduce < SUB_GROUPS_PER_LINE; reduce++) {
|
||||
sumTemp[pm][reduce] = slm_pool_f32.slm_pool_o[sumBase + pm * WARPSIZE * SUB_GROUPS_PER_LINE + reduce];
|
||||
}
|
||||
}
|
||||
|
||||
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
|
||||
[[unroll]] for (uint reduce = 1; reduce < SUB_GROUPS_PER_LINE; reduce++) {
|
||||
sumTemp[pm][0] = sumTemp[pm][0] + sumTemp[pm][reduce];
|
||||
}
|
||||
}
|
||||
|
||||
if ((0x1 & p.flag) != 0) {
|
||||
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
|
||||
float fp32GlobalMax = slm_pool_f32.slm_pool_o[SLM_OFFSET_GLOBAL_MAX + pm * WARPSIZE + lane];
|
||||
float sinkCompensation = fp32GlobalMax - fp32SinkCoeff[pm];
|
||||
sinkCompensation = exp(sinkCompensation);
|
||||
float softmaxSumTemp = sumTemp[pm][0] * sinkCompensation;
|
||||
sumTemp[pm][0] = sumTemp[pm][0] + 1.0f / sinkCompensation;
|
||||
sumTemp[pm][0] = 1.0f / sumTemp[pm][0];
|
||||
sinkCompensation = sinkCompensation / (1.0f + softmaxSumTemp);
|
||||
sumTemp[pm][0] = fp32GlobalMax < fp32SinkCoeff[pm] ? sinkCompensation : sumTemp[pm][0];
|
||||
slm_pool_f32.slm_pool_o[SLM_OFFSET_SOFTMAX_SUM + pm * WARPSIZE + lane] = sumTemp[pm][0];
|
||||
}
|
||||
} else {
|
||||
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
|
||||
slm_pool_f32.slm_pool_o[SLM_OFFSET_SOFTMAX_SUM + pm * WARPSIZE + lane] = 1.0f / sumTemp[pm][0];
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
[[unroll]] for (uint reduce = 2; reduce < SPLIT_P_GROUPS; reduce = reduce << 1 ) {
|
||||
const uint stride = (reduce >> 1) * ALIGNED_P_M * TN * MAT_O_N * TM;
|
||||
if ((localLinearId % reduce) == 0) {
|
||||
const uint reduceBase = localLinearId * ALIGNED_P_M * TN * MAT_O_N * TM + SLM_OFFSET_O;
|
||||
float sumTemp0[4];
|
||||
float sumTemp1[4];
|
||||
const uint reduceVec4Count = ALIGNED_P_M * TN * MAT_O_N * TM / 4 / WARPSIZE;
|
||||
[[unroll]] for (uint totalLoads = 0; totalLoads < reduceVec4Count; totalLoads++) {
|
||||
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
|
||||
sumTemp0[kk] = slm_pool_f32.slm_pool_o[reduceBase + totalLoads * 4 * WARPSIZE + 4 * lane + kk];
|
||||
sumTemp1[kk] = slm_pool_f32.slm_pool_o[reduceBase + stride + totalLoads * 4 * WARPSIZE + 4 * lane + kk];
|
||||
}
|
||||
|
||||
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
|
||||
sumTemp0[kk] = sumTemp0[kk] + sumTemp1[kk];
|
||||
}
|
||||
|
||||
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
|
||||
slm_pool_f32.slm_pool_o[reduceBase + totalLoads * 4 * WARPSIZE + 4 * lane + kk] = sumTemp0[kk];
|
||||
}
|
||||
}
|
||||
}
|
||||
barrier();
|
||||
}
|
||||
|
||||
if (localLinearId == 0) {
|
||||
const uint slmBase0 = SLM_OFFSET_O + lane * WARP_V_DIM;
|
||||
const uint slmBase1 = slmBase0 + SPLIT_P_GROUPS / 2 * ALIGNED_P_M * TN * MAT_O_N * TM;
|
||||
|
||||
const uint offsetOutBase = (d * p.batchStrideO + vWarpIdx * WARP_V_DIM + v * GQA_RATIO * HEAD_DIM + outTokIdx * oDim * N_TOKS_PER_GROUP) / 4;
|
||||
float fp32SoftMaxMul[ALIGNED_P_M];
|
||||
float fp32Output[ALIGNED_P_M][4];
|
||||
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
|
||||
fp32SoftMaxMul[pm] = slm_pool_f32.slm_pool_o[SLM_OFFSET_SOFTMAX_SUM + pm * WARPSIZE + lane];
|
||||
}
|
||||
|
||||
[[unroll]] for (uint vg = 0; vg < WARP_V_DIM / 4; vg++) {
|
||||
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
|
||||
[[unroll]] for (uint vc = 0; vc < 4; vc++) {
|
||||
fp32Output[pm][vc] = slm_pool_f32.slm_pool_o[slmBase0 + pm * WARPSIZE * WARP_V_DIM + vg * 4 + vc] * fp32SoftMaxMul[pm];
|
||||
fp32Output[pm][vc] = fp32Output[pm][vc] + slm_pool_f32.slm_pool_o[slmBase1 + pm * WARPSIZE * WARP_V_DIM + vg * 4 + vc] * fp32SoftMaxMul[pm];
|
||||
}
|
||||
}
|
||||
|
||||
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
|
||||
if (outputMask[pm] == true) {
|
||||
out_f32_vec4[offsetOutBase + outOffsets[pm] + vg] = vec4(fp32Output[pm][0], fp32Output[pm][1], fp32Output[pm][2], fp32Output[pm][3]);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -922,6 +922,10 @@ void process_shaders() {
|
||||
string_to_spv("fa_split_k_reduce", "flash_attn_split_k_reduce.comp", {});
|
||||
|
||||
string_to_spv("fa_mask_opt", "flash_attn_mask_opt.comp", {});
|
||||
|
||||
string_to_spv("fa_decode_ph1", "flash_attn_decode_phase_1.comp", {}, true, true, false, false);
|
||||
string_to_spv("fa_decode_ph2", "flash_attn_decode_phase_2.comp", {}, true, true, false, false);
|
||||
|
||||
string_to_spv("fa_sparse_compact", "flash_attn_sparse_compact.comp", {});
|
||||
string_to_spv("fa_sparse_compact_subgroup", "flash_attn_sparse_compact.comp", {{"USE_SUBGROUPS", "1"}});
|
||||
|
||||
|
||||
@@ -4506,6 +4506,9 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const
|
||||
default:
|
||||
break;
|
||||
}
|
||||
if (ggml_get_op_params_i32(op, 3) == GGML_PREC_F32) {
|
||||
supports_op = false;
|
||||
}
|
||||
break;
|
||||
case GGML_OP_FLASH_ATTN_EXT:
|
||||
{
|
||||
|
||||
@@ -160,6 +160,18 @@ llama_context::llama_context(
|
||||
}
|
||||
cparams.ctx_other = params.ctx_other;
|
||||
}
|
||||
|
||||
// the draft reads the target features of target_layer_ids, concatenated
|
||||
if (params.ctx_other) {
|
||||
const auto & hparams_tgt = llama_get_model(params.ctx_other)->hparams;
|
||||
|
||||
const uint64_t n_embd_inp_enc = (uint64_t) model.target_layer_ids.size() * hparams_tgt.n_embd;
|
||||
if (n_embd_inp_enc != hparams.n_embd_inp_enc()) {
|
||||
throw std::runtime_error(model.arch_name() + " draft expects an encoder input width of " +
|
||||
std::to_string(hparams.n_embd_inp_enc()) + ", but the target model gives " +
|
||||
std::to_string(n_embd_inp_enc));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (cparams.rope_scaling_type == LLAMA_ROPE_SCALING_TYPE_UNSPECIFIED) {
|
||||
|
||||
@@ -2302,6 +2302,10 @@ ggml_tensor * llm_graph_context::build_moe_ffn(
|
||||
}
|
||||
|
||||
experts = build_lora_mm_id(down_exps, cur, selected_experts, down_exps_s); // [n_embd, n_expert_used, n_tokens]
|
||||
if (arch == LLM_ARCH_MISTRAL4) {
|
||||
// src1 can exceed F16 range
|
||||
ggml_prec_set_src(experts, GGML_PREC_F32, 1);
|
||||
}
|
||||
cb(experts, "ffn_moe_down", il);
|
||||
|
||||
if (down_exps_s) {
|
||||
|
||||
+16
-1
@@ -23,8 +23,23 @@ void llama_model_dflash::load_arch_hparams(llama_model_loader & ml) {
|
||||
if (!ml.get_arr(LLM_KV_TARGET_LAYERS, target_layer_ids, false)) {
|
||||
throw std::runtime_error("DFlash model requires 'target_layers' in GGUF metadata");
|
||||
}
|
||||
if (target_layer_ids.empty() || target_layer_ids.size() > 64) {
|
||||
throw std::runtime_error("DFlash 'target_layers' must have between 1 and 64 entries, got " +
|
||||
std::to_string(target_layer_ids.size()));
|
||||
}
|
||||
for (const auto & id : target_layer_ids) {
|
||||
if (id < 0) {
|
||||
throw std::runtime_error("DFlash 'target_layers' entries must be non-negative");
|
||||
}
|
||||
}
|
||||
|
||||
hparams.n_embd_inp_enc_impl = (uint32_t) target_layer_ids.size() * hparams.n_embd;
|
||||
// the width sizes the encoder input buffers, so keep it in 64 bits until it is checked
|
||||
const uint64_t n_embd_inp_enc = (uint64_t) target_layer_ids.size() * hparams.n_embd;
|
||||
if (n_embd_inp_enc == 0 || n_embd_inp_enc > (1u << 20)) {
|
||||
throw std::runtime_error("DFlash encoder input width " + std::to_string(n_embd_inp_enc) +
|
||||
" is out of range (max 1048576)");
|
||||
}
|
||||
hparams.n_embd_inp_enc_impl = (uint32_t) n_embd_inp_enc;
|
||||
|
||||
std::string layers;
|
||||
const char * sep = "";
|
||||
|
||||
@@ -92,14 +92,16 @@ void llama_model_qwen4exp::load_arch_hparams(llama_model_loader & ml) {
|
||||
qwen4exp_require_nonzero(ml, LLM_KV_PLE_CONV_KERNEL, hparams.ple_conv_kernel);
|
||||
qwen4exp_require_nonzero(ml, LLM_KV_EMBEDDING_LENGTH_PER_LAYER, hparams.n_embd_per_layer);
|
||||
|
||||
hparams.ple_n_heads = (hparams.ple_ngram_size - 1) * hparams.ple_heads_per_ngram;
|
||||
hparams.ple_head_dim = hparams.n_embd_per_layer;
|
||||
if (hparams.ple_ngram_size < 2 || hparams.ple_ngram_size > LLAMA_MAX_PLE_NGRAM) {
|
||||
throw std::runtime_error(format("PLE n-gram size %u is out of range", hparams.ple_ngram_size));
|
||||
}
|
||||
if (hparams.ple_n_heads == 0 || hparams.ple_n_heads > LLAMA_MAX_PLE_HEADS) {
|
||||
throw std::runtime_error(format("PLE head count %u is out of range", hparams.ple_n_heads));
|
||||
// keep in 64 bits: a uint32 product can wrap and pass the range check below
|
||||
const uint64_t ple_n_heads = (uint64_t)(hparams.ple_ngram_size - 1) * hparams.ple_heads_per_ngram;
|
||||
if (ple_n_heads == 0 || ple_n_heads > LLAMA_MAX_PLE_HEADS) {
|
||||
throw std::runtime_error(format("PLE head count %" PRIu64 " is out of range", ple_n_heads));
|
||||
}
|
||||
hparams.ple_n_heads = (uint32_t) ple_n_heads;
|
||||
hparams.ple_head_dim = hparams.n_embd_per_layer;
|
||||
|
||||
qwen4exp_require_arr_len(ml, LLM_KV_PLE_LAYER_MULTIPLIERS, hparams.ple_ngram_size);
|
||||
qwen4exp_require_arr_len(ml, LLM_KV_PLE_HEAD_OFFSETS, hparams.ple_n_heads);
|
||||
@@ -1083,6 +1085,9 @@ void llm_graph_input_ple::set_input(const llama_ubatch * ubatch) {
|
||||
const int64_t eos = hp.ple_eos_token_id;
|
||||
const int64_t n_prev = n_gram - 1;
|
||||
|
||||
// ple_n_heads is derived from the n-gram geometry, see load_arch_hparams()
|
||||
GGML_ASSERT((n_gram - 1) * per_gram == n_heads && "PLE head count does not match the n-gram geometry");
|
||||
|
||||
std::vector<int32_t> idx(n_heads * n_tokens);
|
||||
|
||||
GGML_ASSERT(mctx != nullptr);
|
||||
|
||||
@@ -5186,6 +5186,11 @@ struct test_mul_mat_id : public test_case {
|
||||
ggml_tensor * out = ggml_mul_mat_id(ctx, as, b, ids);
|
||||
ggml_set_name(out, "out");
|
||||
|
||||
if (amax > 65504.0f) {
|
||||
// src1 exceeds F16 range
|
||||
ggml_prec_set_src(out, GGML_PREC_F32, 1);
|
||||
}
|
||||
|
||||
return out;
|
||||
}
|
||||
|
||||
@@ -10185,11 +10190,10 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
|
||||
}
|
||||
|
||||
// test src1 f16 overflow
|
||||
// TODO: https://github.com/ggml-org/llama.cpp/pull/26223#issuecomment-5585815365
|
||||
//for (int n : {16, 32, 64}) {
|
||||
// test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_Q4_K, GGML_TYPE_F32, 128, 4, false, 4096, n, 2048, 1e5f));
|
||||
// test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_Q8_0, GGML_TYPE_F32, 8, 2, false, 512, n, 256, 1e5f));
|
||||
//}
|
||||
for (int n : {16, 32, 64}) {
|
||||
test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_Q4_K, GGML_TYPE_F32, 128, 4, false, 4096, n, 2048, 1e5f));
|
||||
test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_Q8_0, GGML_TYPE_F32, 8, 2, false, 512, n, 256, 1e5f));
|
||||
}
|
||||
|
||||
for (ggml_type type_a : base_types) {
|
||||
for (ggml_type type_b : {GGML_TYPE_F32 /*, GGML_TYPE_F16 */}) {
|
||||
@@ -11334,6 +11338,9 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
|
||||
// sparse decode at long context
|
||||
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 1}, 49152, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 0));
|
||||
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 1}, 49152, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 2048));
|
||||
// gemma-4-26b-a4b global-attn layers: head_count_kv=2, 16 query heads (gqa_ratio=8)
|
||||
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 2, { 8, 1}, 49152, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 0));
|
||||
test_cases.emplace_back(new test_flash_attn_ext(512, 512, 2, { 8, 1}, 49152, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 2048));
|
||||
test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {16, 1}, 49152, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true, 0));
|
||||
test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {16, 1}, 49152, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true, 2048));
|
||||
|
||||
|
||||
+20
-7
@@ -1408,9 +1408,14 @@ struct clip_model_loader {
|
||||
}
|
||||
|
||||
// Load the vision/audio feature layer indices if they are explicitly provided
|
||||
// NOTE: gguf conversions should standardize the values of the vision feature layer to
|
||||
// be non-negative, since we use -1 to mark values as unset here.
|
||||
// NOTE: gguf conversions should standardize the values of the vision feature layer to be non-negative, since we use -1 to mark values as unset here.
|
||||
get_arr_int(string_format(KEY_FEATURE_LAYERS, prefix), hparams.feature_layers, false);
|
||||
for (const auto & v : hparams.feature_layers) {
|
||||
if (v > (int) hparams.n_layer) {
|
||||
throw std::runtime_error(string_format("%s: feature layer index %d is out of range (n_layer: %d)",
|
||||
__func__, v, hparams.n_layer));
|
||||
}
|
||||
}
|
||||
|
||||
// model-specific params
|
||||
switch (model.proj_type) {
|
||||
@@ -1456,7 +1461,12 @@ struct clip_model_loader {
|
||||
std::vector<int> wa_layer_indexes_vec;
|
||||
get_arr_int(KEY_WIN_ATTN_LAYER_INDEXES, wa_layer_indexes_vec, false);
|
||||
if (!wa_layer_indexes_vec.empty()) {
|
||||
hparams.insert_layer_id = wa_layer_indexes_vec[0];
|
||||
const int insert_lid = wa_layer_indexes_vec[0];
|
||||
if (insert_lid < 0 || insert_lid >= (int) hparams.n_layer) {
|
||||
throw std::runtime_error(string_format("%s: layer index %d is out of range (n_layer: %d)",
|
||||
__func__, insert_lid, hparams.n_layer));
|
||||
}
|
||||
hparams.insert_layer_id = insert_lid;
|
||||
}
|
||||
} break;
|
||||
case PROJECTOR_TYPE_INTERNVL:
|
||||
@@ -3226,6 +3236,7 @@ struct clip_model_loader {
|
||||
model.pos_embed = get_tensor(string_format(TN_SAM_POS_EMBD, "weight"));
|
||||
model.patch_embed_proj_w = get_tensor(string_format(TN_SAM_PATCH_EMBD, "weight"));
|
||||
model.patch_embed_proj_b = get_tensor(string_format(TN_SAM_PATCH_EMBD, "bias"));
|
||||
model.n_sam_layers = hparams.sam_n_layer;
|
||||
model.sam_layers.resize(model.n_sam_layers);
|
||||
for (int il = 0; il < model.n_sam_layers; ++il) {
|
||||
auto & layer = model.sam_layers[il];
|
||||
@@ -4652,8 +4663,10 @@ bool clip_encode(struct clip_ctx * ctx, struct clip_encode_params * params) {
|
||||
// -> https://huggingface.co/HuggingFaceM4/siglip-so400m-14-980-flash-attn2-navit
|
||||
// -> https://huggingface.co/HuggingFaceM4/siglip-so400m-14-980-flash-attn2-navit/blob/d66538faeba44480d0bfaa42145eef26f9423199/modeling_siglip.py#L316
|
||||
std::vector<int32_t> positions(pos_h * pos_w);
|
||||
int bucket_coords_h[1024];
|
||||
int bucket_coords_w[1024];
|
||||
// note: sized by the actual patch counts; a tall/wide image produces more
|
||||
// than 1024 patches per side and a fixed [1024] array would be overrun
|
||||
std::vector<int> bucket_coords_h(pos_h);
|
||||
std::vector<int> bucket_coords_w(pos_w);
|
||||
for (int i = 0; i < pos_h; i++){
|
||||
bucket_coords_h[i] = std::floor(70.0*i/pos_h);
|
||||
}
|
||||
@@ -4696,8 +4709,8 @@ bool clip_encode(struct clip_ctx * ctx, struct clip_encode_params * params) {
|
||||
|
||||
// SigLIP position buckets (same as resampler path)
|
||||
std::vector<int32_t> positions(pos_h * pos_w);
|
||||
int bucket_coords_h[1024];
|
||||
int bucket_coords_w[1024];
|
||||
std::vector<int> bucket_coords_h(pos_h);
|
||||
std::vector<int> bucket_coords_w(pos_w);
|
||||
for (int i = 0; i < pos_h; i++){
|
||||
bucket_coords_h[i] = std::floor(70.0*i/pos_h);
|
||||
}
|
||||
|
||||
@@ -294,8 +294,8 @@ private:
|
||||
support = filter_support * filterscale; // Widen filter when downsampling
|
||||
ksize = static_cast<int>(std::ceil(support)) * 2 + 1; // Total pixels in kernel
|
||||
|
||||
std::vector<double> pre_weights(outSize * ksize); // Temporary weights
|
||||
bounds.resize(outSize * 2);
|
||||
std::vector<double> pre_weights((size_t) outSize * ksize); // Temporary weights
|
||||
bounds.resize((size_t) outSize * 2);
|
||||
|
||||
|
||||
// For each output pixel, compute its filter coefficients
|
||||
@@ -322,20 +322,20 @@ private:
|
||||
for (x = 0; x < xmax; x++) {
|
||||
// Distance from input pixel center to output pixel center in input space
|
||||
double w = resample_filter((x + xmin - center + 0.5) * ss);
|
||||
pre_weights[xx * ksize + x] = w;
|
||||
pre_weights[(size_t) xx * ksize + x] = w;
|
||||
ww += w; // Accumulate for normalization
|
||||
}
|
||||
|
||||
// Normalize weights to sum to 1.0 (preserves brightness)
|
||||
for (x = 0; x < xmax; x++) {
|
||||
if (ww != 0.0) {
|
||||
pre_weights[xx * ksize + x] /= ww;
|
||||
pre_weights[(size_t) xx * ksize + x] /= ww;
|
||||
}
|
||||
}
|
||||
|
||||
// Zero-pad remaining kernel positions
|
||||
for (; x < ksize; x++) {
|
||||
pre_weights[xx * ksize + x] = 0;
|
||||
pre_weights[(size_t) xx * ksize + x] = 0;
|
||||
}
|
||||
|
||||
// Store input pixel range for this output pixel
|
||||
@@ -345,11 +345,11 @@ private:
|
||||
|
||||
// Convert floating-point coefficients to fixed-point integers
|
||||
// Formula: int32 = round(float * 2^PRECISION_BITS)
|
||||
weights.resize(outSize * ksize);
|
||||
weights.resize((size_t) outSize * ksize);
|
||||
|
||||
const double fxp_scale = std::ldexp(1.0, PRECISION_BITS); // 1.0 * 2^PRECISION_BITS
|
||||
|
||||
for (int i = 0; i < outSize * ksize; i++) {
|
||||
for (size_t i = 0; i < (size_t) outSize * ksize; i++) {
|
||||
// Pillow adds +/- 0.5 then truncates toward zero; std::round would round twice
|
||||
const double rounded = pre_weights[i] * fxp_scale + (pre_weights[i] < 0 ? -0.5 : 0.5);
|
||||
weights[i] = static_cast<int32_t>(rounded);
|
||||
@@ -442,6 +442,12 @@ private:
|
||||
const int src_width = img.get_size().width;
|
||||
const int src_height = img.get_size().height;
|
||||
|
||||
// sanity check on the target size
|
||||
if (target_width <= 0 || target_width > 65536 || target_height <= 0 || target_height > 65536) {
|
||||
throw std::runtime_error("resize target " + std::to_string(target_width) + "x" +
|
||||
std::to_string(target_height) + " is out of range (max 65536)");
|
||||
}
|
||||
|
||||
bool need_horizontal = (target_width != src_width);
|
||||
bool need_vertical = (target_height != src_height);
|
||||
|
||||
|
||||
@@ -40,6 +40,10 @@ export const VIDEO_FILE_TYPES = {
|
||||
[FileTypeVideo.OGG]: {
|
||||
extensions: [FileExtensionVideo.OGG],
|
||||
mimeTypes: [MimeTypeVideo.OGG]
|
||||
},
|
||||
[FileTypeVideo.WEBM]: {
|
||||
extensions: [FileExtensionVideo.WEBM],
|
||||
mimeTypes: [MimeTypeVideo.WEBM]
|
||||
}
|
||||
} as const;
|
||||
|
||||
|
||||
@@ -38,7 +38,8 @@ export enum FileTypeAudio {
|
||||
|
||||
export enum FileTypeVideo {
|
||||
MP4 = 'mp4',
|
||||
OGG = 'ogg'
|
||||
OGG = 'ogg',
|
||||
WEBM = 'webm'
|
||||
}
|
||||
|
||||
export enum FileTypePdf {
|
||||
@@ -104,7 +105,8 @@ export enum FileExtensionAudio {
|
||||
|
||||
export enum FileExtensionVideo {
|
||||
MP4 = '.mp4',
|
||||
OGG = '.ogg'
|
||||
OGG = '.ogg',
|
||||
WEBM = '.webm'
|
||||
}
|
||||
|
||||
export enum FileExtensionPdf {
|
||||
@@ -203,7 +205,8 @@ export enum MimeTypeAudio {
|
||||
|
||||
export enum MimeTypeVideo {
|
||||
MP4 = 'video/mp4',
|
||||
OGG = 'video/ogg'
|
||||
OGG = 'video/ogg',
|
||||
WEBM = 'video/webm'
|
||||
}
|
||||
|
||||
export enum MimeTypeImage {
|
||||
|
||||
@@ -51,6 +51,7 @@ export function getFileTypeCategory(mimeType: string): FileTypeCategory | null {
|
||||
// Video
|
||||
case MimeTypeVideo.MP4:
|
||||
case MimeTypeVideo.OGG:
|
||||
case MimeTypeVideo.WEBM:
|
||||
return FileTypeCategory.VIDEO;
|
||||
|
||||
// PDF
|
||||
|
||||
Reference in New Issue
Block a user