Files
armorpaint/base/tools/iris.c/iris_vulkan_res_adaln.comp
T
2026-06-17 17:44:40 +02:00

47 lines
1.5 KiB
Plaintext

#version 450
/*
* AdaLN: LayerNorm followed by affine modulation.
* norm = (x - mean) / sqrt(var + eps)
* out = (1 + scale[i]) * norm + shift[i]
* shift/scale are per-channel f32 vectors. One workgroup per row.
*/
layout(local_size_x = 256) in;
layout(std430, binding = 0) readonly buffer In { float x[]; };
layout(std430, binding = 1) readonly buffer Shift { float shift[]; };
layout(std430, binding = 2) readonly buffer Scale { float scale[]; };
layout(std430, binding = 3) writeonly buffer Out { float outp[]; };
layout(push_constant) uniform PC { uint rows, hidden; float eps; } pc;
shared float s_sum[256];
shared float s_sumsq[256];
void main() {
uint row = gl_WorkGroupID.x;
uint tid = gl_LocalInvocationID.x;
uint base = row * pc.hidden;
float sum = 0.0, sumsq = 0.0;
for (uint i = tid; i < pc.hidden; i += 256u) {
float v = x[base + i];
sum += v;
sumsq += v * v;
}
s_sum[tid] = sum;
s_sumsq[tid] = sumsq;
barrier();
for (uint s = 128u; s > 0u; s >>= 1u) {
if (tid < s) { s_sum[tid] += s_sum[tid + s]; s_sumsq[tid] += s_sumsq[tid + s]; }
barrier();
}
float mean = s_sum[0] / float(pc.hidden);
float var = s_sumsq[0] / float(pc.hidden) - mean * mean;
float inv = inversesqrt(var + pc.eps);
for (uint i = tid; i < pc.hidden; i += 256u) {
float norm = (x[base + i] - mean) * inv;
outp[base + i] = (1.0 + scale[i]) * norm + shift[i];
}
}