File size: 1,222 Bytes
e774a34
 
 
 
 
 
 
 
5b47652
 
 
 
 
 
 
 
 
 
 
e774a34
5b47652
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
{% 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;
}