Download build/webgpu/norm-skip-row-vec4.wgsl.jinja from webgpu-kernels/com.microsoft.SkipSimplifiedLayerNormalization: direct link, hf CLI and curl.
- Browser
- Download file 3.36 kB
-
https://huggingface.co/kernels/webgpu-kernels/com.microsoft.SkipSimplifiedLayerNormalization/resolve/v1/build/webgpu/norm-skip-row-vec4.wgsl.jinja
- Command line
-
hf download hf://webgpu-kernels/com.microsoft.SkipSimplifiedLayerNormalization@v1/build/webgpu/norm-skip-row-vec4.wgsl.jinja
-
curl -L -o norm-skip-row-vec4.wgsl.jinja https://huggingface.co/kernels/webgpu-kernels/com.microsoft.SkipSimplifiedLayerNormalization/resolve/v1/build/webgpu/norm-skip-row-vec4.wgsl.jinja
3.36 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 %} | |
| {% if useSubgroups %} | |
| enable subgroups; | |
| {% endif %} | |
| {{ env.wgsl.resourceDeclarations }} | |
| const HIDDEN: u32 = {{ hidden }}u; | |
| const HIDDEN_V: u32 = {{ hiddenVec }}u; | |
| const WG: u32 = {{ wg }}u; | |
| var<workgroup> sg_partials: array<f32, WG>; | |
| fn reduce_scalar(value: f32{% if useSubgroups %}, sg_lane: u32, sg_id: u32, num_sg: u32{% else %}, tid: u32{% endif %}) -> f32 { | |
| {% if useSubgroups %} | |
| let s = subgroupAdd(value); | |
| if (num_sg == 1u) { | |
| return s; | |
| } | |
| if (sg_lane == 0u) { | |
| sg_partials[sg_id] = s; | |
| } | |
| workgroupBarrier(); | |
| var total = 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 hasBias %}, vi: u32{% endif %}) -> vec4<f32> { | |
| var value = vec4<f32>(input[idx]) + vec4<f32>(skip[skip_idx]); | |
| {% if hasBias %} | |
| value = value + vec4<f32>(bias[vi]); | |
| {% endif %} | |
| return value; | |
| } | |
| @compute @workgroup_size(WG, 1, 1) | |
| fn main( | |
| @builtin(workgroup_id) wg_id: vec3<u32>, | |
| @builtin(local_invocation_id) lid: vec3<u32>{% 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; | |
| let skip_base = base; | |
| var acc = 0.0; | |
| for (var i = tid; i < HIDDEN_V; i = i + WG) { | |
| let v = residual_value(base + i, skip_base + i{% if hasBias %}, i{% endif %}); | |
| acc = acc + dot(v, v); | |
| } | |
| let total = reduce_scalar(acc{% if useSubgroups %}, sg_lane, sg_id, num_sg{% else %}, tid{% endif %}); | |
| let row_inv = inverseSqrt(total / f32(HIDDEN) + params.epsilon); | |
| for (var i = tid; i < HIDDEN_V; i = i + WG) { | |
| let idx = base + i; | |
| let residual = residual_value(idx, skip_base + i{% if hasBias %}, i{% endif %}); | |
| {% if writeResidualSum %} | |
| input_skip_bias_sum[idx] = {{ vecType }}(residual); | |
| {% endif %} | |
| output[idx] = {{ vecType }}(residual * row_inv * vec4<f32>(gamma[i])); | |
| } | |
| } | |