Download build/webgpu/embed-mask-index.wgsl.jinja from webgpu-kernels/com.microsoft.EmbedLayerNormalization: direct link, hf CLI and curl.
- Browser
- Download file 1.22 kB
-
https://huggingface.co/kernels/webgpu-kernels/com.microsoft.EmbedLayerNormalization/resolve/v1/build/webgpu/embed-mask-index.wgsl.jinja
- Command line
-
hf download hf://webgpu-kernels/com.microsoft.EmbedLayerNormalization@v1/build/webgpu/embed-mask-index.wgsl.jinja
-
curl -L -o embed-mask-index.wgsl.jinja https://huggingface.co/kernels/webgpu-kernels/com.microsoft.EmbedLayerNormalization/resolve/v1/build/webgpu/embed-mask-index.wgsl.jinja
1.22 kB
| {% 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 }} | |
| // com.microsoft.EmbedLayerNormalization, mask_index pass. With a mask, return | |
| // the first zero position or the sequence length when every position is set. | |
| // Without a mask, initialize the optional output to zero. | |
| {% if hasMask %} | |
| const SEQUENCE: u32 = {{ sequenceLength }}u; | |
| {% endif %} | |
| @compute @workgroup_size({{ maskWorkgroupSize }}, 1, 1) | |
| fn main(@builtin(global_invocation_id) gid: vec3<u32>) { | |
| {{ flat_index_2d(maskWorkgroupSize, "batch", "params.batch") }} | |
| var first_zero: i32 = 0; | |
| {% if hasMask %} | |
| first_zero = i32(SEQUENCE); | |
| let base = batch * SEQUENCE; | |
| for (var s: u32 = 0u; s < SEQUENCE; s = s + 1u) { | |
| if (mask[base + s] == 0) { | |
| first_zero = i32(s); | |
| break; | |
| } | |
| } | |
| {% endif %} | |
| mask_index[batch] = first_zero; | |
| } | |