Download build/webgpu/norm-skip-row.wgsl.jinja from webgpu-kernels/com.microsoft.SkipSimplifiedLayerNormalization: direct link, hf CLI and curl.
- Browser
- Download file 3.14 kB
-
https://huggingface.co/kernels/webgpu-kernels/com.microsoft.SkipSimplifiedLayerNormalization/resolve/v1/build/webgpu/norm-skip-row.wgsl.jinja
- Command line
-
hf download hf://webgpu-kernels/com.microsoft.SkipSimplifiedLayerNormalization@v1/build/webgpu/norm-skip-row.wgsl.jinja
-
curl -L -o norm-skip-row.wgsl.jinja https://huggingface.co/kernels/webgpu-kernels/com.microsoft.SkipSimplifiedLayerNormalization/resolve/v1/build/webgpu/norm-skip-row.wgsl.jinja
3.14 kB
| /* Normalize residual = input + skip, with an optional bias. Reductions use | |
| * one workgroup per row; closed-form one-element rows use one invocation. */ | |
| {% set degenerateRow = false %} | |
| {{ env.wgsl.resourceDeclarations }} | |
| const HIDDEN: u32 = {{ hiddenSize }}u; | |
| const WG: u32 = {{ workgroupSize }}u; | |
| var<workgroup> partial: array<f32, WG>; | |
| {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true, valueType="f32") %} | |
| fn {{ name }}(value: {{ valueType }}, tid: u32) -> {{ valueType }} { | |
| {{ 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) { | |
| {{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half]; | |
| } | |
| 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]; | |
| workgroupBarrier(); | |
| return reduced; | |
| } | |
| {% endmacro %} | |
| {{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }} | |
| var<workgroup> row_inv: f32; | |
| 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; | |
| } | |
| @compute @workgroup_size(WG, 1, 1) | |
| fn main( | |
| @builtin({{ "global_invocation_id" if degenerateRow else "workgroup_id" }}) {{ "gid" if degenerateRow else "wg" }}: vec3<u32>, | |
| @builtin(local_invocation_id) lid: vec3<u32>) { | |
| // Fold the row grid across workgroups; independent rows also include the | |
| // invocation offset. The bounds guard drops the final dispatch tail. | |
| let row = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u; | |
| if (row >= params.rows) { | |
| return; | |
| } | |
| let tid = lid.x; | |
| // 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); | |
| {% if writeMean is defined and writeMean %} | |
| mean[row] = 0.0; | |
| {% endif %} | |
| {% if writeInvStd is defined and writeInvStd and not (packedStatistics is defined and packedStatistics) %} | |
| inv_std_var[row] = row_inv; | |
| {% endif %} | |
| } | |
| 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])); | |
| } | |
| } | |