{% 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 shifted_moments: array, WG>; var plane_shift: array; @compute @workgroup_size(WG, 1, 1) fn main( @builtin(workgroup_id) workgroup: vec3, @builtin(num_workgroups) workgroup_count: vec3, @builtin(local_invocation_id) local: vec3 ) { 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(0.0); if (is_active) { for (var i = lane; i < HIDDEN_V4; i = i + LANES) { let value = vec4(x[base + i]); let delta = value - vec4(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(x[base + i]); y[base + i] = {{ vectorScalar }}((value - vec4(mean)) * vec4(affine_scale) + vec4(affine_bias)); } } }