Files
armorpaint/base/tools/iris.c/iris_vulkan_res_qkrmsnorm.comp
T

49 lines
1.7 KiB
Plaintext
Raw Normal View History

2026-06-17 17:44:40 +02:00
#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];
}
}
}