| {% 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 %} |
| |
| /* One workgroup normalizes each row of residual = input + skip, with an |
| * optional bias. */ |
| {% set degenerateRow = (not simplified) and hiddenSize == 1 %} |
| {% if useSubgroups and not degenerateRow %} |
| enable subgroups; |
| {% endif %} |
| {{ env.wgsl.resourceDeclarations }} |
| |
| {% if not degenerateRow or writeResidualSum %} |
| const HIDDEN: u32 = {{ hiddenSize }}u; |
| {% endif %} |
| const WG: u32 = {{ workgroupSize }}u; |
| {% if simplified %} |
| |
| var<workgroup> partial: array<f32, WG>; |
| {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true) %} |
| fn {{ name }}(value: f32, tid: u32) -> f32 { |
| {{ buffer }}[tid] = value; |
| workgroupBarrier(); |
| // Ceil-halving keeps every lane when the workgroup size is not a power of |
| // two. For even n this matches the power-of-two tree order; for odd n, lanes |
| // [0, n-half) fold the upper tail while the middle lane carries forward. |
| var n: u32 = {{ wg }}; |
| loop { |
| let half = (n + 1u) / 2u; |
| if (tid < n - half) { |
| {% if mode == "max" %} |
| {{ buffer }}[tid] = max({{ buffer }}[tid], {{ buffer }}[tid + half]); |
| {% else %} |
| {{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half]; |
| {% endif %} |
| } |
| workgroupBarrier(); |
| n = half; |
| if (n == 1u) { |
| break; |
| } |
| } |
| // The default trailing barrier makes this helper safe for back-to-back calls: every lane reads |
| // slot 0 here, so the next call's first store must not run until all lanes have read it. |
| // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit. |
| let reduced = {{ buffer }}[0]; |
| {% if trailingBarrier %} |
| workgroupBarrier(); |
| {% endif %} |
| return reduced; |
| } |
| {% endmacro %} |
| |
| {{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }} |
| var<workgroup> row_inv: f32; |
| {% else %} |
| {% if not degenerateRow %} |
| |
| var<workgroup> pair_partial: array<vec2<f32>, WG>; |
| |
| {% if useSubgroups %} |
| fn reduce_pair(value: vec2<f32>, sg_lane: u32, sg_id: u32, num_sg: u32) -> vec2<f32> { |
| let s = vec2<f32>(subgroupAdd(value.x), subgroupAdd(value.y)); |
| if (num_sg == 1u) { |
| return s; |
| } |
| if (sg_lane == 0u) { |
| pair_partial[sg_id] = s; |
| } |
| workgroupBarrier(); |
| var total = vec2<f32>(0.0, 0.0); |
| for (var i = 0u; i < num_sg; i = i + 1u) { |
| total = total + pair_partial[i]; |
| } |
| return total; |
| } |
| {% else %} |
| fn reduce_pair(value: vec2<f32>, tid: u32) -> vec2<f32> { |
| pair_partial[tid] = value; |
| workgroupBarrier(); |
| {{ wgsl_tree_fold(["pair_partial"], idx="tid", wg="WG", form="head") }} |
| return pair_partial[0]; |
| } |
| {% endif %} |
| {% endif %} |
| {% endif %} |
| |
| {% if not degenerateRow or writeResidualSum %} |
| fn residual_value(row: u32, d: u32) -> f32 { |
| let index = row * HIDDEN + d; |
| var value = f32(input[index]) + f32(skip[index]); |
| {% if hasBias %} |
| value = value + f32(bias[d]); |
| {% endif %} |
| return value; |
| } |
| {% endif %} |
| |
| @compute @workgroup_size(WG, 1, 1) |
| fn main( |
| @builtin(workgroup_id) wg: vec3<u32>, |
| @builtin(num_workgroups) nwg: vec3<u32>{% if not degenerateRow %}, |
| @builtin(local_invocation_id) lid: vec3<u32>{% endif %}{% if useSubgroups and not degenerateRow %}, |
| @builtin(subgroup_invocation_id) sg_lane: u32, |
| @builtin(subgroup_id) sg_id: u32, |
| @builtin(num_subgroups) num_sg: u32{% endif %} |
| ) { |
| // 2D-folded row index: wg.y carries the high bits past the maxComputeWorkgroupsPerDimension |
| // workgroup-per-dimension dispatch limit. Reduces to wg.x when nwg.y == 1; |
| // the row >= params.rows guard drops the over-dispatched tail. |
| let row = wg.x + wg.y * nwg.x; |
| if (row >= params.rows) { |
| return; |
| } |
| {% if not degenerateRow %} |
| let tid = lid.x; |
| {% endif %} |
| {% if simplified %} |
| |
| // RMS normalization uses one sum-of-squares sweep, without a mean or beta. |
| |
| var local_sq = 0.0; |
| for (var d: u32 = tid; d < HIDDEN; d = d + WG) { |
| let value = residual_value(row, d); |
| local_sq = local_sq + value * value; |
| } |
| let sq = reduce_sum(local_sq, tid); |
| if (tid == 0u) { |
| row_inv = inverseSqrt(sq / f32(HIDDEN) + params.epsilon); |
| } |
| workgroupBarrier(); |
| |
| for (var d: u32 = tid; d < HIDDEN; d = d + WG) { |
| let index = row * HIDDEN + d; |
| let residual = residual_value(row, d); |
| {% if writeResidualSum %} |
| input_skip_bias_sum[index] = {{ scalar }}(residual); |
| {% endif %} |
| output[index] = {{ scalar }}(residual * row_inv * f32(gamma[d])); |
| } |
| {% elif degenerateRow %} |
| |
| // HIDDEN == 1: the row's mean is its only element, so the centered value and |
| // the variance are exactly zero and the output reduces to beta. The closed |
| // form avoids computing that zero by subtracting two equal rounded values. |
| let row_inv = inverseSqrt(params.epsilon); |
| {% if writeResidualSum %} |
| let residual = residual_value(row, 0u); |
| input_skip_bias_sum[row] = {{ scalar }}(residual); |
| {% endif %} |
| // 0.0 * row_inv keeps the IEEE result when epsilon == 0 makes row_inv +Inf. |
| output[row] = {{ scalar }}(0.0 * row_inv * f32(gamma[0]){% if hasBeta %} + f32(beta[0]){% endif %}); |
| {% else %} |
| |
| // Shifted moments: accumulating (x - x[0], (x - x[0])^2) keeps the sums |
| // small for rows with a large common offset; every thread reconstructs the |
| // row mean and variance from the merged pair. |
| let shift = residual_value(row, 0u); |
| var acc = vec2<f32>(0.0, 0.0); |
| for (var d = tid; d < HIDDEN; d = d + WG) { |
| let centered = residual_value(row, d) - shift; |
| acc.x = acc.x + centered; |
| acc.y = acc.y + centered * centered; |
| } |
| |
| {% if useSubgroups %} |
| let totals = reduce_pair(acc, sg_lane, sg_id, num_sg); |
| {% else %} |
| let totals = reduce_pair(acc, tid); |
| {% endif %} |
| let mean_delta = totals.x / f32(HIDDEN); |
| let row_mean = shift + mean_delta; |
| let variance = max(totals.y / f32(HIDDEN) - mean_delta * mean_delta, 0.0); |
| let row_inv = inverseSqrt(variance + params.epsilon); |
| for (var d = tid; d < HIDDEN; d = d + WG) { |
| let index = row * HIDDEN + d; |
| let residual = residual_value(row, d); |
| {% if writeResidualSum %} |
| input_skip_bias_sum[index] = {{ scalar }}(residual); |
| {% endif %} |
| output[index] = {{ scalar }}((residual - row_mean) * row_inv * f32(gamma[d]){% if hasBeta %} + f32(beta[d]){% endif %}); |
| } |
| {% endif %} |
| } |
| |