| {% 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<workgroup> red: array<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 + 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]; |
| } |
| } |
| |