{% 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 %}{% set useSubgroups = source.useSubgroups %} {% if source.usesF16 %} enable f16; {% endif %} {% if useSubgroups %} enable subgroups; {% endif %} {{ env.wgsl.resourceDeclarations }} const HIDDEN: u32 = {{ source.hidden }}u; const HIDDEN_V: u32 = {{ source.hiddenVec }}u; const WG: u32 = {{ source.wg }}u; var sg_partials: array, WG>; fn reduce_pair(value: vec2{% if useSubgroups %}, sg_lane: u32, sg_id: u32, num_sg: u32{% else %}, tid: u32{% endif %}) -> vec2 { {% if useSubgroups %} let s = vec2(subgroupAdd(value.x), subgroupAdd(value.y)); if (num_sg == 1u) { return s; } if (sg_lane == 0u) { sg_partials[sg_id] = s; } workgroupBarrier(); var total = vec2(0.0, 0.0); for (var i = 0u; i < num_sg; i = i + 1u) { total = total + sg_partials[i]; } return total; {% else %} // No-subgroup tier: workgroup barrier tree-reduction (WG is a power of two). sg_partials[tid] = value; workgroupBarrier(); {{ wgsl_tree_fold(["sg_partials"], idx="tid", wg="WG", form="head", breakInline=true) }} return sg_partials[0]; {% endif %} } // 4 contiguous residual elements (input[idx] + skip[skip_idx] [+ bias]) at vec4 // index `vi`. skip_idx == idx for the normal (non-broadcast) path; for a skip // that broadcasts across the leading/batch dim uses a folded index. fn residual_value(idx: u32, skip_idx: u32{% if source.hasBias %}, vi: u32{% endif %}) -> vec4 { var value = vec4(input[idx]) + vec4(skip[skip_idx]); {% if source.hasBias %} value = value + vec4(bias[vi]); {% endif %} return value; } @compute @workgroup_size(WG, 1, 1) fn main( @builtin(workgroup_id) wg_id: vec3, @builtin(local_invocation_id) lid: vec3{% if useSubgroups %}, @builtin(subgroup_invocation_id) sg_lane: u32, @builtin(subgroup_id) sg_id: u32, @builtin(num_subgroups) num_sg: u32{% endif %} ) { let row = wg_id.x + wg_id.y * params.rowStride; if (row >= params.rows) { return; } let tid = lid.x; let base = row * HIDDEN_V; {% if source.broadcastSkip %} // skip broadcasts across the batch dim: fold row into [0, skipRows) so every // batch reuses the same skip row (skipRows == params.rows ⇒ identity). let skip_base = (row % params.skipRows) * HIDDEN_V; {% else %} let skip_base = base; {% endif %} let shift = residual_value(base, skip_base{% if source.hasBias %}, 0u{% endif %}).x; var acc = vec2(0.0, 0.0); for (var i = tid; i < HIDDEN_V; i = i + WG) { let v = residual_value(base + i, skip_base + i{% if source.hasBias %}, i{% endif %}); let d = v - vec4(shift); acc.x = acc.x + d.x + d.y + d.z + d.w; acc.y = acc.y + dot(d, d); } let totals = reduce_pair(acc{% if useSubgroups %}, sg_lane, sg_id, num_sg{% else %}, tid{% endif %}); let mean_d = totals.x / f32(HIDDEN); let variance = max(totals.y / f32(HIDDEN) - mean_d * mean_d, 0.0); let row_inv = inverseSqrt(variance + params.epsilon); let row_mean = shift + mean_d; for (var i = tid; i < HIDDEN_V; i = i + WG) { let idx = base + i; let residual = residual_value(idx, skip_base + i{% if source.hasBias %}, i{% endif %}); {% if source.writeResidualSum %} input_skip_bias_sum[idx] = {{ source.vecType }}(residual); {% endif %} var value = (residual - vec4(row_mean)) * row_inv * vec4(gamma[i]); {% if source.hasBeta %} value = value + vec4(beta[i]); {% endif %} output[idx] = {{ source.vecType }}(value); } }