#version 450
// Fused GDN gating, mirroring cuda/gdn.cu fused_gdn_gating_kernel.
// beta = sigmoid(b)
// g = -exp(a_log) * softplus(a + dt_bias)
// a_log and dt_bias are per-head (indexed by idx % num_heads), broadcast over
// the leading batch*seq dims. One invocation per element.
//
// b, a, beta_out, g_out: [total_elements] a_log, dt_bias: [num_heads]
// All f32 on this backend (f16/bf16 tensors are stored as f32).
layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in;
layout(set = 0, binding = 0) readonly buffer B { float b_in[]; };
layout(set = 0, binding = 1) readonly buffer A { float a_in[]; };
layout(set = 0, binding = 2) readonly buffer ALog { float a_log[]; };
layout(set = 0, binding = 3) readonly buffer DtBias { float dt_bias[]; };
layout(set = 0, binding = 4) writeonly buffer BetaOut { float beta_out[]; };
layout(set = 0, binding = 5) writeonly buffer GOut { float g_out[]; };
layout(push_constant) uniform Pc { uint total_elements; uint num_heads; };
void main() {
uint idx = gl_GlobalInvocationID.x;
if (idx >= total_elements) { return; }
uint head_idx = idx % num_heads;
float b_val = b_in[idx];
float beta_val = 1.0 / (1.0 + exp(-b_val));
float a_val = a_in[idx];
float a_log_val = a_log[head_idx];
float dt_bias_val = dt_bias[head_idx];
float sp_input = a_val + dt_bias_val;
float softplus_val = log(1.0 + exp(sp_input));
float g_val = -exp(a_log_val) * softplus_val;
beta_out[idx] = beta_val;
g_out[idx] = g_val;
}