{% macro wgsl_tree_fold_stmt(a, op, idx, svar) %} {% if op == "max" %} {{ a }}[{{ idx }}] = max({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]); {%- else %} {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] + {{ a }}[{{ idx }} + {{ svar }}]; {%- endif %} {% endmacro %} {% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false) %} var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u; loop { {% if form == "head" %} {% if breakInline %} if ({{ svar }} == 0u) { break; } {% else %} if ({{ svar }} == 0u) { break; } {% endif %} {% endif %} {% if bodyInline %} if ({{ idx }} < {{ svar }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} } {% else %} if ({{ idx }} < {{ svar }}) { {% for a in arrays %} {{ wgsl_tree_fold_stmt(a, op, idx, svar) }} {% endfor %} } {% endif %} {% if form == "head" %} {% if barrierFirst %} workgroupBarrier(); {{ svar }} = {{ svar }} / 2u; {% else %} {{ svar }} = {{ svar }} / 2u; workgroupBarrier(); {% endif %} {% else %} workgroupBarrier(); if ({{ svar }} == 1u) { break; } {{ svar }} = {{ svar }} / 2u; {% endif %} } {%- endmacro %} /* Split-K partial sum-of-squares for tensors with few rows and a large hidden dimension. A workgroup-per-row kernel exposes too little parallelism in this regime, so this pass splits each row across SPLIT workgroups (row = wg.x, split index = wg.z). Each workgroup accumulates a partial sum-of-squares over its HIDDEN/SPLIT slice and writes one partial to scratch. The normalize pass folds the SPLIT partials per row. Split-K reassociates the f32 sum, so this route is not bit-identical to the unsplit reduction. */ {{ env.wgsl.resourceDeclarations }} const HIDDEN: u32 = {{ hiddenSize }}u; const WG: u32 = {{ workgroupSize }}u; const SPLIT: u32 = {{ split }}u; var red: array; @compute @workgroup_size(WG, 1, 1) fn main(@builtin(workgroup_id) wg: vec3, @builtin(local_invocation_id) lid: vec3) { let row = wg.x + wg.y * params.rowStride; if (row >= params.rows) { return; } let k = wg.z; let tid = lid.x; let chunk = (HIDDEN + SPLIT - 1u) / SPLIT; let start = k * chunk; var end = start + chunk; if (end > HIDDEN) { end = HIDDEN; } let base = row * HIDDEN; var acc = 0.0; var d = start + tid; loop { if (d >= end) { break; } let v = f32(x[base + d]); acc = acc + v * v; d = d + WG; } red[tid] = acc; workgroupBarrier(); {{ wgsl_tree_fold(["red"], idx="tid", wg="WG", typed=true, form="head", breakInline=true, bodyInline=true) }} if (tid == 0u) { partials[row * SPLIT + k] = red[0]; } }