{% 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) { {{ 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; }