Download build/webgpu/norm-skip-row.wgsl.jinja from webgpu-kernels/com.microsoft.SkipLayerNormalization: direct link, hf CLI and curl.
- Browser
- Download file 5.8 kB
-
https://huggingface.co/kernels/webgpu-kernels/com.microsoft.SkipLayerNormalization/resolve/v1/build/webgpu/norm-skip-row.wgsl.jinja
- Command line
-
hf download hf://webgpu-kernels/com.microsoft.SkipLayerNormalization@v1/build/webgpu/norm-skip-row.wgsl.jinja
-
curl -L -o norm-skip-row.wgsl.jinja https://huggingface.co/kernels/webgpu-kernels/com.microsoft.SkipLayerNormalization/resolve/v1/build/webgpu/norm-skip-row.wgsl.jinja
5.8 kB
| {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %} | |
| {% if op == "max" or op == "min" %} | |
| {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %} | |
| {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ 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, reuse=false) %} | |
| var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u; | |
| loop { | |
| if ({{ svar }} == 0u) { | |
| break; | |
| } | |
| if ({{ idx }} < {{ svar }}) { | |
| {% for a in arrays %} | |
| {{ wgsl_tree_fold_stmt(a, op, idx, svar) }} | |
| {% endfor %} | |
| } | |
| {{ svar }} = {{ svar }} / 2u; | |
| workgroupBarrier(); | |
| }{% endmacro %} | |
| /* Normalize residual = input + skip, with an optional bias. Reductions use | |
| * one workgroup per row; closed-form one-element rows use one invocation. */ | |
| {% set degenerateRow = (not simplified) and hiddenSize == 1 %} | |
| {% if useSubgroups and not degenerateRow %} | |
| enable subgroups; | |
| {% endif %} | |
| {{ env.wgsl.resourceDeclarations }} | |
| {% if not degenerateRow or writeResidualSum or (writeMean is defined and writeMean) %} | |
| const HIDDEN: u32 = {{ hiddenSize }}u; | |
| {% endif %} | |
| const WG: u32 = {{ workgroupSize }}u; | |
| {% 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 %} | |
| {% if not degenerateRow or writeResidualSum or (writeMean is defined and writeMean) %} | |
| 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({{ "global_invocation_id" if degenerateRow else "workgroup_id" }}) {{ "gid" if degenerateRow else "wg" }}: 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 %} | |
| ) { | |
| // Fold the row grid across workgroups; independent rows also include the | |
| // invocation offset. The bounds guard drops the final dispatch tail. | |
| {% if degenerateRow %} | |
| let row = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG; | |
| {% else %} | |
| let row = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u; | |
| {% endif %} | |
| if (row >= params.rows) { | |
| return; | |
| } | |
| {% if not degenerateRow %} | |
| let tid = lid.x; | |
| {% endif %} | |
| {% if 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 packedStatistics is defined and packedStatistics %} | |
| row_stats[row] = vec2<f32>(residual_value(row, 0u), row_inv); | |
| {% elif writeMean is defined and writeMean %} | |
| mean[row] = residual_value(row, 0u); | |
| {% endif %} | |
| {% if writeInvStd is defined and writeInvStd and not (packedStatistics is defined and packedStatistics) %} | |
| inv_std_var[row] = row_inv; | |
| {% endif %} | |
| {% 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); | |
| {% if packedStatistics is defined and packedStatistics %} | |
| if (tid == 0u) { row_stats[row] = vec2<f32>(row_mean, row_inv); } | |
| {% elif writeMean is defined and writeMean %} | |
| if (tid == 0u) { mean[row] = row_mean; } | |
| {% endif %} | |
| {% if writeInvStd is defined and writeInvStd and not (packedStatistics is defined and packedStatistics) %} | |
| if (tid == 0u) { inv_std_var[row] = row_inv; } | |
| {% endif %} | |
| 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 %} | |
| } | |