vulkan backend ops: implemented GATED_LINEAR_ATTN (#25601)
* vulkan : add GATED_LINEAR_ATTN op * docs : update Vulkan ops * vulkan : remove unused GLA spec constant * Updated ops.md * ops.md update
This commit is contained in:
@@ -0,0 +1,82 @@
|
||||
#version 450
|
||||
|
||||
#extension GL_EXT_control_flow_attributes : require
|
||||
|
||||
#define BLOCK_SIZE 64
|
||||
layout(local_size_x = BLOCK_SIZE, local_size_y = 1, local_size_z = 1) in;
|
||||
|
||||
layout(push_constant) uniform Parameters {
|
||||
uint B;
|
||||
uint T;
|
||||
uint C;
|
||||
uint H;
|
||||
float scale;
|
||||
};
|
||||
|
||||
layout(binding = 0) readonly buffer KBuf { A_TYPE k[]; };
|
||||
layout(binding = 1) readonly buffer VBuf { A_TYPE v[]; };
|
||||
layout(binding = 2) readonly buffer QBuf { A_TYPE q[]; };
|
||||
layout(binding = 3) readonly buffer GBuf { A_TYPE g[]; };
|
||||
layout(binding = 4) readonly buffer StateBuf { A_TYPE state_in[]; };
|
||||
layout(binding = 5) buffer DstBuf { A_TYPE dst[]; };
|
||||
|
||||
shared A_TYPE _k[BLOCK_SIZE], _q[BLOCK_SIZE], _g[BLOCK_SIZE];
|
||||
|
||||
void main() {
|
||||
const uint head_size = BLOCK_SIZE;
|
||||
const uint batch_id = gl_WorkGroupID.x / H;
|
||||
const uint head_id = gl_WorkGroupID.x % H;
|
||||
const uint tid = gl_LocalInvocationID.x;
|
||||
|
||||
const uint state_size = C * head_size;
|
||||
const uint n_seq_tokens = T / B;
|
||||
|
||||
if (batch_id >= B || head_id >= H) {
|
||||
return;
|
||||
}
|
||||
|
||||
// state[i] holds column tid of this head's state matrix: S[i][tid]
|
||||
A_TYPE state[BLOCK_SIZE];
|
||||
[[unroll]] for (uint i = 0; i < head_size; i++) {
|
||||
state[i] = state_in[batch_id * state_size + head_id * head_size * head_size
|
||||
+ i * head_size + tid];
|
||||
}
|
||||
|
||||
const uint start_t = batch_id * n_seq_tokens * C + head_id * head_size + tid;
|
||||
const uint end_t = (batch_id + 1) * n_seq_tokens * C + head_id * head_size + tid;
|
||||
|
||||
for (uint t = start_t; t < end_t; t += C) {
|
||||
barrier();
|
||||
_k[tid] = k[t];
|
||||
_q[tid] = q[t];
|
||||
_g[tid] = g[t];
|
||||
barrier();
|
||||
|
||||
const A_TYPE v_val = v[t];
|
||||
A_TYPE y = 0.0;
|
||||
|
||||
[[unroll]] for (uint i = 0; i < head_size; i += 4) {
|
||||
vec4 k_vec = vec4(_k[i], _k[i+1], _k[i+2], _k[i+3]);
|
||||
vec4 q_vec = vec4(_q[i], _q[i+1], _q[i+2], _q[i+3]);
|
||||
vec4 g_vec = vec4(_g[i], _g[i+1], _g[i+2], _g[i+3]);
|
||||
vec4 s_vec = vec4(state[i], state[i+1], state[i+2], state[i+3]);
|
||||
|
||||
vec4 kv = k_vec * v_val;
|
||||
|
||||
s_vec = s_vec * g_vec + kv;
|
||||
y += dot(q_vec, s_vec);
|
||||
|
||||
state[i] = s_vec.x;
|
||||
state[i+1] = s_vec.y;
|
||||
state[i+2] = s_vec.z;
|
||||
state[i+3] = s_vec.w;
|
||||
}
|
||||
|
||||
dst[t] = y * scale;
|
||||
}
|
||||
|
||||
[[unroll]] for (uint i = 0; i < head_size; i++) {
|
||||
dst[T * C + batch_id * state_size + head_id * head_size * head_size
|
||||
+ i * head_size + tid] = state[i];
|
||||
}
|
||||
}
|
||||
@@ -1057,6 +1057,8 @@ void process_shaders() {
|
||||
|
||||
string_to_spv("rwkv_wkv6_f32", "wkv6.comp", merge_maps(base_dict, {{"A_TYPE", "float"}}));
|
||||
|
||||
string_to_spv("gated_linear_attn_f32", "gla.comp", merge_maps(base_dict, {{"A_TYPE", "float"}}));
|
||||
|
||||
string_to_spv("rwkv_wkv7_f32", "wkv7.comp", merge_maps(base_dict, {{"A_TYPE", "float"}}));
|
||||
|
||||
string_to_spv("gated_delta_net_f32", "gated_delta_net.comp", merge_maps(base_dict, {{"FLOAT_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}, {"USE_SUBGROUP_CLUSTERED", "1"}}));
|
||||
|
||||
Reference in New Issue
Block a user