49 lines
1.7 KiB
Plaintext
49 lines
1.7 KiB
Plaintext
|
|
#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];
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|