ai.onnx.GroupNormalization / build /webgpu /group-normalization-splitk-partials.wgsl.jinja
Xenova's picture
Xenova HF Staff
sync 91d990483a17
28e88e5 verified
Raw
History Blame
1.06 kB
{{ 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]; }
}