| {{ env.wgsl.resourceDeclarations }} |
| |
| const HIDDEN: u32 = {{ hiddenSize }}u; |
| const WG: u32 = {{ workgroupSize }}u; |
| const SPLIT: u32 = {{ split }}u; |
| var<workgroup> reduction: array<vec2<f32>, WG>; |
| |
| @compute @workgroup_size(WG, 1, 1) |
| fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) { |
| let row = wg.x; |
| let part = wg.z; |
| if (row >= params.rows) { return; } |
| let tid = lid.x; |
| let chunk = (HIDDEN + SPLIT - 1u) / SPLIT; |
| let start = part * chunk; |
| let end = min(start + chunk, HIDDEN); |
| let base = row * HIDDEN; |
| let shift = f32(x[base]); |
| var pair = vec2<f32>(0.0); |
| for (var d = start + tid; d < end; d = d + WG) { |
| let value = f32(x[base + d]) - shift; |
| pair = pair + vec2<f32>(value, value * value); |
| } |
| reduction[tid] = pair; |
| workgroupBarrier(); |
| for (var stride = WG >> 1u; stride > 0u; stride = stride >> 1u) { |
| if (tid < stride) { reduction[tid] = reduction[tid] + reduction[tid + stride]; } |
| workgroupBarrier(); |
| } |
| if (tid == 0u) { partials[row * SPLIT + part] = reduction[0]; } |
| } |
| |