#version 450 /* * Z-Image per-head QK RMSNorm, applied in place to Q and K. * For each (sequence position s, head h): normalize the head_dim slice by its * RMS and scale by a per-dimension weight (shared across heads). * q[s,h,d] = q[s,h,d] / sqrt(mean_d(q[s,h,:]^2) + eps) * q_weight[d] * k[s,h,d] = k[s,h,d] / sqrt(mean_d(k[s,h,:]^2) + eps) * k_weight[d] * One invocation per (s, h); the head_dim loop runs serially. */ layout(local_size_x = 256) in; layout(std430, binding = 0) buffer Q { float q[]; }; layout(std430, binding = 1) buffer K { float k[]; }; layout(std430, binding = 2) readonly buffer QW { float qw[]; }; layout(std430, binding = 3) readonly buffer KW { float kw[]; }; layout(push_constant) uniform PC { uint seq, heads, head_dim; float eps; } pc; void main() { uint stride = gl_NumWorkGroups.x * 256u; uint total = pc.seq * pc.heads; uint hidden = pc.heads * pc.head_dim; for (uint gid = gl_GlobalInvocationID.x; gid < total; gid += stride) { uint s = gid / pc.heads; uint h = gid - s * pc.heads; uint off = s * hidden + h * pc.head_dim; float sq = 0.0; for (uint d = 0u; d < pc.head_dim; d++) { float v = q[off + d]; sq += v * v; } float inv = inversesqrt(sq / float(pc.head_dim) + pc.eps); for (uint d = 0u; d < pc.head_dim; d++) { q[off + d] = q[off + d] * inv * qw[d]; } sq = 0.0; for (uint d = 0u; d < pc.head_dim; d++) { float v = k[off + d]; sq += v * v; } inv = inversesqrt(sq / float(pc.head_dim) + pc.eps); for (uint d = 0u; d < pc.head_dim; d++) { k[off + d] = k[off + d] * inv * kw[d]; } } }