#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]; } }