Download build/webgpu/instance-normalization-splitk-combine.wgsl.jinja from webgpu-kernels/ai.onnx.InstanceNormalization: direct link, hf CLI and curl.
- Browser
- Download file 1.83 kB
-
https://huggingface.co/kernels/webgpu-kernels/ai.onnx.InstanceNormalization/resolve/v1/build/webgpu/instance-normalization-splitk-combine.wgsl.jinja
- Command line
-
hf download hf://webgpu-kernels/ai.onnx.InstanceNormalization@v1/build/webgpu/instance-normalization-splitk-combine.wgsl.jinja
-
curl -L -o instance-normalization-splitk-combine.wgsl.jinja https://huggingface.co/kernels/webgpu-kernels/ai.onnx.InstanceNormalization/resolve/v1/build/webgpu/instance-normalization-splitk-combine.wgsl.jinja
1.83 kB
| // Fold SPLIT per-plane (sum, sum-of-squares) partials into mean and inverse | |
| // standard deviation. One thread handles each plane. The partials are centred on | |
| // the plane's first element, so E[y^2] - E[y]^2 keeps the variance a raw second | |
| // moment would cancel away; max(value, 0) guards against negative rounding residue. | |
| {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %} | |
| {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %} | |
| // 2D-folded flat index: gid.y carries the high bits past the dispatch's | |
| // per-axis workgroup fold width. | |
| let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }}; | |
| if ({{ name }} >= {{ bound }}) { | |
| return; | |
| }{% endmacro %} | |
| {{ env.wgsl.resourceDeclarations }} | |
| const SPLIT: u32 = {{ split }}u; | |
| const COMBINE_WG: u32 = {{ combineWorkgroupSize }}u; | |
| @compute @workgroup_size(COMBINE_WG, 1, 1) | |
| fn main(@builtin(global_invocation_id) gid: vec3<u32>) { | |
| {{ flat_index_2d("COMBINE_WG", "plane", "params.planes") }} | |
| var total = 0.0; | |
| var total_sq = 0.0; | |
| let b = plane * SPLIT; | |
| for (var k = 0u; k < SPLIT; k = k + 1u) { | |
| total = total + partials[(b + k) * 2u]; | |
| total_sq = total_sq + partials[(b + k) * 2u + 1u]; | |
| } | |
| let n = f32(params.spatial); | |
| // The partials are accumulated around the plane's first element; undo the shift | |
| // on the mean and leave the variance, which the shift does not change. | |
| {% if vectorizedSpec %} | |
| let shift = f32(input[plane * (params.spatial / 4u)].x); | |
| {% else %} | |
| let shift = f32(input[plane * params.spatial]); | |
| {% endif %} | |
| let centred_mean = total / n; | |
| let mean = shift + centred_mean; | |
| let variance = max(total_sq / n - centred_mean * centred_mean, 0.0); | |
| stats[plane * 2u] = mean; | |
| stats[plane * 2u + 1u] = inverseSqrt(variance + params.epsilon); | |
| } | |