#version 450 #extension GL_EXT_control_flow_attributes : require #extension GL_EXT_shader_16bit_storage : require #extension GL_EXT_shader_explicit_arithmetic_types_float16 : require #extension GL_KHR_shader_subgroup_basic : enable #if USE_SUBGROUP_ADD #extension GL_KHR_shader_subgroup_arithmetic : enable #endif #define BINDING_IDX_K 0u #include "types.glsl" #include "fa_types.glsl" #define FaTypeV FA_TYPE_F32 layout(constant_id = 0) const uint FaTypeK = FA_TYPE_F32; layout(constant_id = 1) const uint FaBlockBytesK = 4; layout(constant_id = 2) const uint SUBGROUP_SIZE = 32; #include "flash_attn_dequant.glsl" // one workgroup computes one output element, one invocation per head element #define HEAD_SIZE 128 layout(local_size_x = HEAD_SIZE, local_size_y = 1, local_size_z = 1) in; layout(binding = 0) readonly buffer QBuf { float q[]; }; layout(binding = 1) readonly buffer KBufF16 { float16_t k_f16[]; }; layout(binding = 1) readonly buffer KBufF32 { float k_f32[]; }; layout(binding = 1) readonly buffer KBufBF16 { uint16_t k_bf16[]; }; layout(binding = 2) readonly buffer WBuf { float weights[]; }; layout(binding = 3) readonly buffer MBuf { float16_t mask[]; }; layout(binding = 4) writeonly buffer DstBuf { float dst[]; }; layout(push_constant) uniform PushConstants { uint n_kv; uint n_heads; uint n_tokens; uint n_streams; uint n_masks; uint dispatch_x; uint q_nb1; uint q_nb2; uint q_nb3; uint k_nb2; uint k_nb3; uint w_nb1; uint w_nb3; uint m_nb1; uint m_nb3; uint d_nb1; uint d_nb3; }; shared float k_row[HEAD_SIZE]; #if USE_SUBGROUP_ADD shared float sg_partials[HEAD_SIZE / SUBGROUP_SIZE]; #else shared float partials[HEAD_SIZE]; #endif void main() { const uint tid = gl_LocalInvocationID.x; const uint output_idx = gl_WorkGroupID.y * dispatch_x + gl_WorkGroupID.x; const uint n_outputs = n_kv * n_tokens * n_streams; if (fa_type_needs_shmem(FaTypeK)) { init_iq_shmem(gl_WorkGroupSize); } if (output_idx >= n_outputs) { return; } const uint ik = output_idx % n_kv; const uint ts = output_idx / n_kv; const uint t = ts % n_tokens; const uint s = ts / n_tokens; const uint k_offset = ik * k_nb2 + s * k_nb3; // k strides come in as bytes, so scale them down to the view being indexed const uint k_block_elems = fa_block_elems(FaTypeK); const uint k_elem_bytes = FaBlockBytesK / k_block_elems; if (FaTypeK == FA_TYPE_F16) { k_row[tid] = float(k_f16[k_offset / k_elem_bytes + tid]); } else if (FaTypeK == FA_TYPE_F32) { k_row[tid] = k_f32[k_offset / k_elem_bytes + tid]; } else if (FaTypeK == FA_TYPE_BF16) { k_row[tid] = bf16_to_fp32(uint(k_bf16[k_offset / k_elem_bytes + tid])); } else if (4 * tid < HEAD_SIZE) { const uint coord = 4 * tid; const uint ib = coord / k_block_elems; const uint iqs = coord % k_block_elems; const vec4 values = dequantize4(ib, iqs, k_offset / FaBlockBytesK, BINDING_IDX_K); k_row[coord + 0] = values.x; k_row[coord + 1] = values.y; k_row[coord + 2] = values.z; k_row[coord + 3] = values.w; } barrier(); const float k_val = k_row[tid]; float score = 0.0; for (uint h = 0; h < n_heads; ++h) { const float prod = q[h * q_nb1 + t * q_nb2 + s * q_nb3 + tid] * k_val; #if USE_SUBGROUP_ADD const float sg_sum = subgroupAdd(prod); if (gl_SubgroupInvocationID == 0) { sg_partials[gl_SubgroupID] = sg_sum; } barrier(); if (tid == 0) { float sum = 0.0; [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) { sum += sg_partials[i]; } score += max(sum, 0.0) * weights[h + t * w_nb1 + s * w_nb3]; } // the reads above must complete before the next iteration overwrites sg_partials barrier(); #else partials[tid] = prod; barrier(); [[unroll]] for (uint stride = HEAD_SIZE / 2; stride > 0; stride >>= 1) { if (tid < stride) { partials[tid] += partials[tid + stride]; } barrier(); } if (tid == 0) { score += max(partials[0], 0.0) * weights[h + t * w_nb1 + s * w_nb3]; } // the read of partials[0] above must complete before the next iteration // overwrites partials[tid] barrier(); #endif } if (tid == 0) { const uint mask_offset = ik + t * m_nb1 + (s % n_masks) * m_nb3; dst[ik + t * d_nb1 + s * d_nb3] = score + float(mask[mask_offset]); } }