ai.onnx.InstanceNormalization / build /webgpu /instance-normalization-batched-planes-vec4.wgsl.jinja
| {% 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)); | |
| } | |
| } | |
| } | |