tools: add iris.c
This commit is contained in:
@@ -0,0 +1,46 @@
|
||||
#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];
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user