* Make ggml_gated_delta_net take only the initial recurrent state (D, 1, n_seqs) and passes the snapshot count K as an op parameter instead of inferring it from state->ne[1]. Remove the padding hack and copy all emitted snapshots into the recurrent cache with a single strided ggml_cpy * Make GDN changes in all backends. Address review comments. * Fix CI build errors
190 lines
6.7 KiB
Plaintext
190 lines
6.7 KiB
Plaintext
#version 450
|
|
|
|
#extension GL_EXT_control_flow_attributes : require
|
|
#extension GL_KHR_shader_subgroup_basic : enable
|
|
#if USE_SUBGROUP_CLUSTERED
|
|
#extension GL_KHR_shader_subgroup_clustered : enable
|
|
#endif
|
|
#if USE_SUBGROUP_ADD
|
|
#extension GL_KHR_shader_subgroup_arithmetic : enable
|
|
#endif
|
|
|
|
// Caller guarantees valid spec constants: S_V % COLS_PER_WG == 0 and S_V % LANES_PER_COLUMN == 0,
|
|
// so no bounds checking is needed.
|
|
layout(constant_id = 0) const uint S_V = 128;
|
|
layout(constant_id = 1) const uint KDA = 0;
|
|
layout(constant_id = 2) const uint SUBGROUP_SIZE = 32;
|
|
layout(constant_id = 3) const uint LANES_PER_COLUMN = 32;
|
|
|
|
const uint COLS_PER_WG = SUBGROUP_SIZE / LANES_PER_COLUMN;
|
|
const uint ROWS_PER_LANE = S_V / LANES_PER_COLUMN;
|
|
|
|
layout(local_size_x_id = 2, local_size_y = 1, local_size_z = 1) in;
|
|
|
|
layout(push_constant) uniform Parameters {
|
|
uint H;
|
|
uint n_tokens;
|
|
uint n_seqs;
|
|
uint s_off;
|
|
uint sq1, sq2, sq3;
|
|
uint sv1, sv2, sv3;
|
|
uint sb1, sb2, sb3;
|
|
uint neq1, rq3;
|
|
float scale;
|
|
uint K;
|
|
};
|
|
|
|
layout(binding = 0) readonly buffer QBuf { FLOAT_TYPE data_q[]; };
|
|
layout(binding = 1) readonly buffer KBuf { FLOAT_TYPE data_k[]; };
|
|
layout(binding = 2) readonly buffer VBuf { FLOAT_TYPE data_v[]; };
|
|
layout(binding = 3) readonly buffer GBuf { FLOAT_TYPE data_g[]; };
|
|
layout(binding = 4) readonly buffer BetaBuf { FLOAT_TYPE data_beta[]; };
|
|
layout(binding = 5) readonly buffer StateBuf { FLOAT_TYPE data_state[]; };
|
|
layout(binding = 6) buffer DstBuf { FLOAT_TYPE data_dst[]; };
|
|
|
|
#if !USE_SUBGROUP_ADD && !USE_SUBGROUP_CLUSTERED
|
|
shared FLOAT_TYPE temp[SUBGROUP_SIZE];
|
|
|
|
// This does a reduction across groups of LANES_PER_COLUMN
|
|
FLOAT_TYPE reduce_add_shmem(FLOAT_TYPE partial) {
|
|
const uint lane = gl_SubgroupInvocationID;
|
|
temp[lane] = partial;
|
|
barrier();
|
|
[[unroll]] for (uint s = LANES_PER_COLUMN / 2u; s > 0; s >>= 1u) {
|
|
FLOAT_TYPE other = temp[lane ^ s];
|
|
barrier();
|
|
temp[lane] += other;
|
|
barrier();
|
|
}
|
|
const FLOAT_TYPE result = temp[lane];
|
|
barrier();
|
|
return result;
|
|
}
|
|
#endif
|
|
|
|
// clusterSize for subgroupClusteredAdd must be a compile-time constant; branch on spec constant
|
|
FLOAT_TYPE reduce_partial(FLOAT_TYPE partial) {
|
|
switch (LANES_PER_COLUMN) {
|
|
case 1u:
|
|
return partial;
|
|
#if USE_SUBGROUP_CLUSTERED
|
|
// Workaround for GLSL requiring a literal constant for the cluster size.
|
|
// The branches should all fold away.
|
|
case 2u:
|
|
return subgroupClusteredAdd(partial, 2u);
|
|
case 4u:
|
|
return subgroupClusteredAdd(partial, 4u);
|
|
case 8u:
|
|
return subgroupClusteredAdd(partial, 8u);
|
|
case 16u:
|
|
return subgroupClusteredAdd(partial, 16u);
|
|
case 32u:
|
|
return subgroupClusteredAdd(partial, 32u);
|
|
case 64u:
|
|
return subgroupClusteredAdd(partial, 64u);
|
|
#endif
|
|
default:
|
|
#if USE_SUBGROUP_ADD
|
|
return subgroupAdd(partial);
|
|
#else
|
|
return reduce_add_shmem(partial);
|
|
#endif
|
|
}
|
|
}
|
|
|
|
void main() {
|
|
const uint head_id = gl_WorkGroupID.x;
|
|
const uint seq_id = gl_WorkGroupID.y;
|
|
const uint lane = gl_SubgroupInvocationID % LANES_PER_COLUMN;
|
|
const uint col = gl_WorkGroupID.z * COLS_PER_WG + (gl_SubgroupInvocationID / LANES_PER_COLUMN);
|
|
|
|
const uint iq1 = head_id % neq1;
|
|
const uint iq3 = seq_id / rq3;
|
|
|
|
const uint state_size = S_V * S_V;
|
|
// input state holds s0 only [S_v, S_v, H, n_seqs]: per-seq stride is H*D.
|
|
const uint state_in_base = (seq_id * H + head_id) * state_size;
|
|
// output state layout per slot: same per-(seq,head) offset as the single-slot case.
|
|
const uint state_out_base = (seq_id * H + head_id) * state_size;
|
|
const uint state_size_per_snap = state_size * H * n_seqs;
|
|
|
|
FLOAT_TYPE s_shard[ROWS_PER_LANE];
|
|
[[unroll]] for (uint r = 0; r < ROWS_PER_LANE; r++) {
|
|
s_shard[r] = FLOAT_TYPE(data_state[state_in_base + col * S_V + r * LANES_PER_COLUMN + lane]);
|
|
}
|
|
|
|
// snapshot slot mapping: slot 0 = most recent state, slot s = s tokens back.
|
|
// When n_tokens < K, only slots 0..n_tokens-1 are written; older slots are caller-owned.
|
|
|
|
uint attn_off = (seq_id * n_tokens * H + head_id) * S_V;
|
|
|
|
for (uint t = 0; t < n_tokens; t++) {
|
|
const uint q_off = iq3 * sq3 + t * sq2 + iq1 * sq1;
|
|
const uint k_off = q_off;
|
|
const uint v_off = seq_id * sv3 + t * sv2 + head_id * sv1;
|
|
const uint gb_off = seq_id * sb3 + t * sb2 + head_id * sb1;
|
|
const FLOAT_TYPE beta_val = FLOAT_TYPE(data_beta[gb_off]);
|
|
|
|
FLOAT_TYPE k_reg[ROWS_PER_LANE];
|
|
FLOAT_TYPE q_reg[ROWS_PER_LANE];
|
|
[[unroll]] for (uint r = 0; r < ROWS_PER_LANE; r++) {
|
|
const uint i = r * LANES_PER_COLUMN + lane;
|
|
k_reg[r] = FLOAT_TYPE(data_k[k_off + i]);
|
|
q_reg[r] = FLOAT_TYPE(data_q[q_off + i]);
|
|
}
|
|
|
|
FLOAT_TYPE g_exp[ROWS_PER_LANE];
|
|
if (KDA == 0) {
|
|
const FLOAT_TYPE g_val = exp(FLOAT_TYPE(data_g[gb_off]));
|
|
[[unroll]] for (uint r = 0; r < ROWS_PER_LANE; r++) {
|
|
g_exp[r] = g_val;
|
|
}
|
|
} else {
|
|
const uint g_base = gb_off * S_V;
|
|
[[unroll]] for (uint r = 0; r < ROWS_PER_LANE; r++) {
|
|
const uint i = r * LANES_PER_COLUMN + lane;
|
|
g_exp[r] = exp(FLOAT_TYPE(data_g[g_base + i]));
|
|
}
|
|
}
|
|
|
|
const FLOAT_TYPE v_val = FLOAT_TYPE(data_v[v_off + col]);
|
|
|
|
FLOAT_TYPE kv_shard = 0.0;
|
|
[[unroll]] for (uint r = 0; r < ROWS_PER_LANE; r++) {
|
|
kv_shard += g_exp[r] * s_shard[r] * k_reg[r];
|
|
}
|
|
FLOAT_TYPE kv_col = reduce_partial(kv_shard);
|
|
|
|
FLOAT_TYPE delta_col = (v_val - kv_col) * beta_val;
|
|
|
|
FLOAT_TYPE attn_partial = 0.0;
|
|
[[unroll]] for (uint r = 0; r < ROWS_PER_LANE; r++) {
|
|
s_shard[r] = g_exp[r] * s_shard[r] + k_reg[r] * delta_col;
|
|
attn_partial += s_shard[r] * q_reg[r];
|
|
}
|
|
FLOAT_TYPE attn_col = reduce_partial(attn_partial);
|
|
|
|
if (lane == 0) {
|
|
data_dst[attn_off + col] = attn_col * scale;
|
|
}
|
|
|
|
attn_off += S_V * H;
|
|
|
|
if (K > 1u) {
|
|
const int target_slot = int(n_tokens) - 1 - int(t);
|
|
if (target_slot >= 0 && target_slot < int(K)) {
|
|
const uint slot_base = s_off + uint(target_slot) * state_size_per_snap + state_out_base;
|
|
[[unroll]] for (uint r = 0; r < ROWS_PER_LANE; r++) {
|
|
data_dst[slot_base + col * S_V + r * LANES_PER_COLUMN + lane] = s_shard[r];
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
if (K == 1u) {
|
|
[[unroll]] for (uint r = 0; r < ROWS_PER_LANE; r++) {
|
|
data_dst[s_off + state_out_base + col * S_V + r * LANES_PER_COLUMN + lane] = s_shard[r];
|
|
}
|
|
}
|
|
}
|