Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
654 changes: 604 additions & 50 deletions ggml/src/ggml-vulkan/ggml-vulkan.cpp

Large diffs are not rendered by default.

167 changes: 167 additions & 0 deletions ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_decode_phase_1.comp
Original file line number Diff line number Diff line change
@@ -0,0 +1,167 @@
#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

layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;

layout (binding = 0) readonly buffer Q {float qState[];};
layout (binding = 0) readonly buffer Q_VEC2 {f16vec2 qStateVec2[];};
layout (binding = 1) readonly buffer K {float16_t kState[];};
layout (binding = 1) readonly buffer K_MAT2X4 {f16mat2x4 kStateMat[];};
layout (binding = 2) buffer MASK {f16vec2 mState[];};
layout (binding = 2) buffer MASK_F16 {float16_t mState_f16[];};
layout (binding = 3) buffer P_FP32 {float matP_f32[];};
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 kvStride;
uint batchStrideQ;
uint batchStrideK;
uint batchStrideV;
uint batchStrideM;
uint batchStrideO;
float softMaxScale;
} p;

layout (constant_id = 0) const uint GROUPSIZE = 32;
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;

#define MAX_HEADS 8

#define MAT_P_COUNT (DP / WARPSIZE)
#define MATP_REDUCE (GROUPSIZE / 2)
#define SUBGROUP_COUNT (GROUPSIZE / WARPSIZE)
#define SUBGROUP_DIVISION (SUBGROUP_COUNT / GQA_RATIO)

#define K_COUNT (WARPSIZE / 8)
#define O_COUNT ((GQA_RATIO + SUBGROUP_COUNT - 1) / SUBGROUP_COUNT)

shared float slm_pool_pv[GQA_RATIO * GROUPSIZE];
shared float slm_pool_max_temp[GQA_RATIO];

void main() {
const uint lane = gl_SubgroupInvocationID;
const uint h = gl_WorkGroupID.x;
const uint v = gl_WorkGroupID.y;
const uint d = gl_WorkGroupID.z;
const uint localLinearId = gl_SubgroupID;
const uint hhq = localLinearId & 0x1;
const uint vvq = localLinearId >> 1;
const uint qDim = p.qHead * HEAD_DIM;
const uint oDim = qDim;
const uint kvDim = p.kvStride;
const uint maskDim = p.kvSeqLen;
const uint offsetBaseQ = (d * p.batchStrideQ + h * HEAD_DIM * GQA_RATIO + hhq * WARPSIZE + lane);
const uint offsetBaseK = (d * p.batchStrideK + (v * MATP_REDUCE + vvq * WARPSIZE) * kvDim + lane * kvDim + h * HEAD_DIM + hhq * K_COUNT * 8) / 8;
const uint offsetOut = d * p.qHead * p.kvSeqLen + v * MATP_REDUCE + h * GQA_RATIO * p.kvSeqLen + localLinearId * O_COUNT * p.kvSeqLen + lane;
const uint offsetMax = d * p.qHead * p.kvSeqLen / MATP_REDUCE + v + h * GQA_RATIO * p.kvSeqLen / MATP_REDUCE + localLinearId * O_COUNT * p.kvSeqLen / MATP_REDUCE;
const uint offsetSlmPv = (localLinearId * WARPSIZE + lane);
const uint offsetSlmLoadPv = (localLinearId * O_COUNT * GROUPSIZE + lane);
const uint offsetBaseM = d * p.batchStrideM + v * MATP_REDUCE + lane;
const float16_t fp16SoftmaxScale = float16_t(p.softMaxScale);
const float fp32Min = uintBitsToFloat(0xFEFFFFFF);
f16mat2x4 matK[K_COUNT];
float matQ[GQA_RATIO];
float fp32P[GQA_RATIO];
float fp32TempO[GROUPSIZE / WARPSIZE];
float fp32O[O_COUNT][MATP_REDUCE / WARPSIZE];
float maxOut[O_COUNT];
float16_t mask[MATP_REDUCE / WARPSIZE];

[[unroll]] for (uint i = 0; i < GQA_RATIO; i++) {
fp32P[i] = 0.0f;
}

[[unroll]] for (uint i = 0; i < MATP_REDUCE / WARPSIZE; i++) {
mask[i] = mState_f16[i * WARPSIZE + offsetBaseM];
}

const uint loopCount = HEAD_DIM / WARPSIZE / 2;
[[unroll]] for (uint loop = 0; loop < loopCount; loop++) {
[[unroll]] for (uint kd = 0; kd < K_COUNT; kd++) {
matK[kd] = kStateMat[offsetBaseK + loop * 2 * K_COUNT + kd];
}

[[unroll]] for (uint qd = 0; qd < GQA_RATIO; qd++) {
matQ[qd] = qState[offsetBaseQ + loop * 2 * WARPSIZE + qd * HEAD_DIM];
}

[[unroll]] for (uint nn = 0; nn < GQA_RATIO; nn++) {
[[unroll]] for (uint kk = 0; kk < K_COUNT; kk++) {
fp32P[nn] = fma(float(matK[kk][0].x), subgroupBroadcast(matQ[nn], 8 * kk + 0), fp32P[nn]);
fp32P[nn] = fma(float(matK[kk][0].y), subgroupBroadcast(matQ[nn], 8 * kk + 1), fp32P[nn]);
fp32P[nn] = fma(float(matK[kk][0].z), subgroupBroadcast(matQ[nn], 8 * kk + 2), fp32P[nn]);
fp32P[nn] = fma(float(matK[kk][0].w), subgroupBroadcast(matQ[nn], 8 * kk + 3), fp32P[nn]);
fp32P[nn] = fma(float(matK[kk][1].x), subgroupBroadcast(matQ[nn], 8 * kk + 4), fp32P[nn]);
fp32P[nn] = fma(float(matK[kk][1].y), subgroupBroadcast(matQ[nn], 8 * kk + 5), fp32P[nn]);
fp32P[nn] = fma(float(matK[kk][1].z), subgroupBroadcast(matQ[nn], 8 * kk + 6), fp32P[nn]);
fp32P[nn] = fma(float(matK[kk][1].w), subgroupBroadcast(matQ[nn], 8 * kk + 7), fp32P[nn]);
}
}
}

[[unroll]] for (uint nn = 0; nn < GQA_RATIO; nn++) {
fp32P[nn] = fp32P[nn] * p.softMaxScale;
}

[[unroll]] for (uint pd = 0; pd < GQA_RATIO; pd++) {
slm_pool_pv[offsetSlmPv + pd * GROUPSIZE] = fp32P[pd];
}

barrier();

if (O_COUNT * localLinearId < GQA_RATIO) {
float maskFp32[MATP_REDUCE / WARPSIZE];
[[unroll]] for (uint i = 0; i < MATP_REDUCE / WARPSIZE; i++) {
maskFp32[i] = float(mask[i]);
}
[[unroll]] for (uint oc = 0; oc < O_COUNT; oc++) {
[[unroll]] for (uint os = 0; os < GROUPSIZE / WARPSIZE; os++) {
fp32TempO[os] = slm_pool_pv[offsetSlmLoadPv + os * WARPSIZE + oc * GROUPSIZE];
}

[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
fp32O[oc][os] = fp32TempO[2 * os] + fp32TempO[2 * os + 1];
}

[[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++) {
[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
matP_f32[offsetOut + oc * p.kvSeqLen + os * WARPSIZE] = fp32O[oc][os];
}
if (lane == 0) {
out_max_f32[offsetMax + oc * p.kvSeqLen / MATP_REDUCE] = maxOut[oc];
}
}
}
}
Loading
Loading