ai.onnx.InstanceNormalization / build /webgpu /instance-normalization-batched-planes-vec4.wgsl.jinja
Xenova's picture
Xenova HF Staff
sync 2e7068faf55e
0af9165 verified
Raw
History Blame
2.46 kB
{% if usesF16 %}
enable f16;
{% endif %}
{{ env.wgsl.resourceDeclarations }}
// A power-of-two lane cohort reduces one plane while several cohorts share a
// workgroup.
const HIDDEN: u32 = {{ hidden }}u;
const HIDDEN_V4: u32 = {{ hiddenVec }}u;
const CHANNELS: u32 = {{ channels }}u;
const EPSILON: f32 = {{ epsilon }};
const WG: u32 = {{ workgroupSize }}u;
const LANES: u32 = {{ lanesPerPlane }}u;
const PLANES_PER_WG: u32 = {{ planesPerWorkgroup }}u;
var<workgroup> shifted_moments: array<vec2<f32>, WG>;
var<workgroup> plane_shift: array<f32, PLANES_PER_WG>;
@compute @workgroup_size(WG, 1, 1)
fn main(
@builtin(workgroup_id) workgroup: vec3<u32>,
@builtin(num_workgroups) workgroup_count: vec3<u32>,
@builtin(local_invocation_id) local: vec3<u32>
) {
let tid = local.x;
let plane_in_workgroup = tid / LANES;
let lane = tid % LANES;
let group = workgroup.x + workgroup.y * workgroup_count.x;
let row = group * PLANES_PER_WG + plane_in_workgroup;
let is_active = row < params.rows;
let base = row * HIDDEN_V4;
if (is_active && lane == 0u) {
plane_shift[plane_in_workgroup] = f32(x[base].x);
}
workgroupBarrier();
let shift = plane_shift[plane_in_workgroup];
var moments = vec2<f32>(0.0);
if (is_active) {
for (var i = lane; i < HIDDEN_V4; i = i + LANES) {
let value = vec4<f32>(x[base + i]);
let delta = value - vec4<f32>(shift);
moments.x = moments.x + delta.x + delta.y + delta.z + delta.w;
moments.y = moments.y + dot(delta, delta);
}
}
shifted_moments[tid] = moments;
workgroupBarrier();
var stride = LANES / 2u;
loop {
if (stride == 0u) { break; }
if (lane < stride) {
shifted_moments[tid] = shifted_moments[tid] + shifted_moments[tid + stride];
}
stride = stride / 2u;
workgroupBarrier();
}
if (is_active) {
let total = shifted_moments[plane_in_workgroup * LANES];
let mean_delta = total.x / f32(HIDDEN);
let variance = max(total.y / f32(HIDDEN) - mean_delta * mean_delta, 0.0);
let mean = shift + mean_delta;
let inv_std = inverseSqrt(variance + EPSILON);
let channel = row % CHANNELS;
let affine_scale = inv_std * f32(scale[channel]);
let affine_bias = f32(bias[channel]);
for (var i = lane; i < HIDDEN_V4; i = i + LANES) {
let value = vec4<f32>(x[base + i]);
y[base + i] = {{ vectorScalar }}((value - vec4<f32>(mean)) * vec4<f32>(affine_scale) + vec4<f32>(affine_bias));
}
}
}