Xenova HF Staff commited on
Commit
e774a34
·
verified ·
1 Parent(s): 2c28f1c

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -59,7 +59,7 @@ Attributes and default values (overridable per request):
59
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
60
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
61
  - [`test.json`](build/webgpu/test.json) — correctness cases
62
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
63
  - [`embed-mask-index.wgsl.jinja`](build/webgpu/embed-mask-index.wgsl.jinja)
64
  - [`embed-normalize.wgsl.jinja`](build/webgpu/embed-normalize.wgsl.jinja)
65
  - [`embed-sum.wgsl.jinja`](build/webgpu/embed-sum.wgsl.jinja)
@@ -67,7 +67,7 @@ Attributes and default values (overridable per request):
67
  ## Use with `@huggingface/kernels`
68
 
69
  ```sh
70
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
71
  ```
72
 
73
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
59
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
60
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
61
  - [`test.json`](build/webgpu/test.json) — correctness cases
62
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
63
  - [`embed-mask-index.wgsl.jinja`](build/webgpu/embed-mask-index.wgsl.jinja)
64
  - [`embed-normalize.wgsl.jinja`](build/webgpu/embed-normalize.wgsl.jinja)
65
  - [`embed-sum.wgsl.jinja`](build/webgpu/embed-sum.wgsl.jinja)
 
67
  ## Use with `@huggingface/kernels`
68
 
69
  ```sh
70
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
71
  ```
72
 
73
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/embed-mask-index.wgsl.jinja CHANGED
@@ -1,3 +1,11 @@
 
 
 
 
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  // com.microsoft.EmbedLayerNormalization, mask_index pass. With a mask, return
@@ -9,10 +17,7 @@ const SEQUENCE: u32 = {{ sequenceLength }}u;
9
 
10
  @compute @workgroup_size({{ maskWorkgroupSize }}, 1, 1)
11
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
12
- let batch = gid.x;
13
- if (batch >= params.batch) {
14
- return;
15
- }
16
  var first_zero: i32 = 0;
17
  {% if hasMask %}
18
  first_zero = i32(SEQUENCE);
 
1
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
2
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
3
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
+ // per-axis workgroup fold width.
5
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
6
+ if ({{ name }} >= {{ bound }}) {
7
+ return;
8
+ }{% endmacro %}
9
  {{ env.wgsl.resourceDeclarations }}
10
 
11
  // com.microsoft.EmbedLayerNormalization, mask_index pass. With a mask, return
 
17
 
18
  @compute @workgroup_size({{ maskWorkgroupSize }}, 1, 1)
19
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
20
+ {{ flat_index_2d(maskWorkgroupSize, "batch", "params.batch") }}
 
 
 
21
  var first_zero: i32 = 0;
22
  {% if hasMask %}
23
  first_zero = i32(SEQUENCE);
build/webgpu/embed-normalize.wgsl.jinja CHANGED
@@ -12,68 +12,32 @@ const WG: u32 = {{ workgroupSize }}u;
12
  var<workgroup> partial: array<f32, WG>;
13
 
14
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
15
- {% if op == "max" %}
16
- {{ a }}[{{ idx }}] = max({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);
17
- {%- else %}
18
- {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] + {{ a }}[{{ idx }} + {{ svar }}];
19
- {%- endif %}
20
- {% endmacro %}
21
- {% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false) %}
22
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
23
  loop {
24
- {% if form == "head" %}
25
- {% if breakInline %}
26
- if ({{ svar }} == 0u) { break; }
27
- {% else %}
28
  if ({{ svar }} == 0u) {
29
  break;
30
  }
31
- {% endif %}
32
- {% endif %}
33
- {% if bodyInline %}
34
- if ({{ idx }} < {{ svar }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} }
35
- {% else %}
36
  if ({{ idx }} < {{ svar }}) {
37
  {% for a in arrays %}
38
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
39
  {% endfor %}
40
  }
41
- {% endif %}
42
- {% if form == "head" %}
43
- {% if barrierFirst %}
44
- workgroupBarrier();
45
- {{ svar }} = {{ svar }} / 2u;
46
- {% else %}
47
  {{ svar }} = {{ svar }} / 2u;
48
  workgroupBarrier();
49
- {% endif %}
50
- {% else %}
51
- workgroupBarrier();
52
- if ({{ svar }} == 1u) {
53
- break;
54
- }
55
- {{ svar }} = {{ svar }} / 2u;
56
- {% endif %}
57
- }
58
- {%- endmacro %}
59
-
60
  // Reusing partial after this reduction requires a barrier between the read of
61
  // partial[0] and the next write, or the next round can race the prior readers.
62
- {% set trailingBarrier = trailingBarrier is defined and trailingBarrier %}
63
  fn reduce_sum(value: f32, tid: u32) -> f32 {
64
  partial[tid] = value;
65
  workgroupBarrier();
66
  {{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }}
67
- {% if trailingBarrier %}
68
- let total = partial[0];
69
- workgroupBarrier();
70
- return total;
71
- {% else %}
72
  return partial[0];
73
- {% endif %}
74
  }
75
 
76
-
77
  @compute @workgroup_size(WG, 1, 1)
78
  fn main(@builtin(workgroup_id) wg: vec3<u32>,
79
  @builtin(local_invocation_id) lid: vec3<u32>) {
 
12
  var<workgroup> partial: array<f32, WG>;
13
 
14
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
15
+ {% if op == "max" or op == "min" %}
16
+ {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
17
+ {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %}
18
+ {% 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) %}
 
 
 
19
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
20
  loop {
 
 
 
 
21
  if ({{ svar }} == 0u) {
22
  break;
23
  }
 
 
 
 
 
24
  if ({{ idx }} < {{ svar }}) {
25
  {% for a in arrays %}
26
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
27
  {% endfor %}
28
  }
 
 
 
 
 
 
29
  {{ svar }} = {{ svar }} / 2u;
30
  workgroupBarrier();
31
+ }{% endmacro %}
 
 
 
 
 
 
 
 
 
 
32
  // Reusing partial after this reduction requires a barrier between the read of
33
  // partial[0] and the next write, or the next round can race the prior readers.
 
34
  fn reduce_sum(value: f32, tid: u32) -> f32 {
35
  partial[tid] = value;
36
  workgroupBarrier();
37
  {{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }}
 
 
 
 
 
38
  return partial[0];
 
39
  }
40
 
 
41
  @compute @workgroup_size(WG, 1, 1)
42
  fn main(@builtin(workgroup_id) wg: vec3<u32>,
43
  @builtin(local_invocation_id) lid: vec3<u32>) {
build/webgpu/manifest.json CHANGED
@@ -75,29 +75,25 @@
75
  },
76
  "when": ["embedContractOk", "dispatchFits"],
77
  "bindings": {
78
- "input_ids": { "arg": "inputIdsT", "buffer": "read-only-storage", "elementType": "i32" },
79
- "word_embedding": { "arg": "wordEmbeddingT", "buffer": "read-only-storage", "elementType": "$aScalar" },
80
- "position_embedding": { "arg": "positionEmbeddingT", "buffer": "read-only-storage", "elementType": "$aScalar" },
81
- "output": { "arg": "outputT", "buffer": "storage", "elementType": "$aScalar" },
82
- "params": { "buffer": "uniform", "struct": [{ "name": "tokens", "type": "u32", "value": "tokens" }] },
83
- "embedding_sum": { "arg": "embeddingSumT", "buffer": "storage", "elementType": "$aScalar" },
84
- "position_ids": { "arg": "positionIdsT", "buffer": "read-only-storage", "elementType": "i32" },
85
- "segment_ids": { "arg": "segmentIdsT", "buffer": "read-only-storage", "elementType": "i32" },
86
- "segment_embedding": { "arg": "segmentEmbeddingT", "buffer": "read-only-storage", "elementType": "$aScalar" },
87
- "gamma": { "arg": "gammaT", "buffer": "read-only-storage", "elementType": "$aScalar", "length": "$HIDDEN_LEN" },
88
- "beta": { "arg": "betaT", "buffer": "read-only-storage", "elementType": "$aScalar", "length": "$HIDDEN_LEN" },
89
- "mask": { "arg": "maskT", "buffer": "read-only-storage", "elementType": "i32" },
90
- "mask_index": { "arg": "maskIndexT", "buffer": "storage", "elementType": "i32" },
91
- "params_2": {
92
- "name": "params",
93
- "buffer": "uniform",
94
- "struct": [{ "name": "batch", "type": "u32", "value": "batchSize" }]
95
- }
96
  },
97
  "variants": [
98
  {
99
  "id": "noseg_nopos_nosum_nomask",
100
- "when": ["not present.segmentEmbeddingT", "not present.positionIdsT", "not present.embeddingSumT", "not present.maskIndexT"],
101
  "passes": [
102
  {
103
  "id": "sum",
@@ -117,7 +113,7 @@
117
  },
118
  {
119
  "id": "noseg_nopos_nosum_mask",
120
- "when": ["not present.segmentEmbeddingT", "not present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "present.maskT"],
121
  "passes": [
122
  {
123
  "id": "sum",
@@ -137,7 +133,7 @@
137
  "id": "maskIndex",
138
  "name": "EmbedLayerNormalization.MaskIndex",
139
  "shader": "embed-mask-index.wgsl.jinja",
140
- "bindings": ["mask", "mask_index", "params_2"],
141
  "dispatch": {
142
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
143
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -147,14 +143,14 @@
147
  ]
148
  },
149
  {
150
- "id": "noseg_nopos_sum_nomask",
151
- "when": ["not present.segmentEmbeddingT", "not present.positionIdsT", "present.embeddingSumT", "not present.maskIndexT"],
152
  "passes": [
153
  {
154
  "id": "sum",
155
  "name": "EmbedLayerNormalization.EmbeddingSum",
156
  "shader": "embed-sum.wgsl.jinja",
157
- "bindings": ["input_ids", "word_embedding", "position_embedding", "output", "embedding_sum", "params"],
158
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
159
  },
160
  {
@@ -163,12 +159,23 @@
163
  "shader": "embed-normalize.wgsl.jinja",
164
  "bindings": ["output", "gamma", "beta", "params"],
165
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
166
  }
167
  ]
168
  },
169
  {
170
- "id": "noseg_nopos_sum_mask",
171
- "when": ["not present.segmentEmbeddingT", "not present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "present.maskT"],
172
  "passes": [
173
  {
174
  "id": "sum",
@@ -183,29 +190,18 @@
183
  "shader": "embed-normalize.wgsl.jinja",
184
  "bindings": ["output", "gamma", "beta", "params"],
185
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
186
- },
187
- {
188
- "id": "maskIndex",
189
- "name": "EmbedLayerNormalization.MaskIndex",
190
- "shader": "embed-mask-index.wgsl.jinja",
191
- "bindings": ["mask", "mask_index", "params_2"],
192
- "dispatch": {
193
- "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
194
- "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
195
- "z": 1
196
- }
197
  }
198
  ]
199
  },
200
  {
201
- "id": "noseg_posids_nosum_nomask",
202
- "when": ["not present.segmentEmbeddingT", "present.positionIdsT", "not present.embeddingSumT", "not present.maskIndexT"],
203
  "passes": [
204
  {
205
  "id": "sum",
206
  "name": "EmbedLayerNormalization.EmbeddingSum",
207
  "shader": "embed-sum.wgsl.jinja",
208
- "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "params"],
209
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
210
  },
211
  {
@@ -214,18 +210,29 @@
214
  "shader": "embed-normalize.wgsl.jinja",
215
  "bindings": ["output", "gamma", "beta", "params"],
216
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
217
  }
218
  ]
219
  },
220
  {
221
- "id": "noseg_posids_nosum_mask",
222
- "when": ["not present.segmentEmbeddingT", "present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "present.maskT"],
223
  "passes": [
224
  {
225
  "id": "sum",
226
  "name": "EmbedLayerNormalization.EmbeddingSum",
227
  "shader": "embed-sum.wgsl.jinja",
228
- "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "params"],
229
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
230
  },
231
  {
@@ -237,9 +244,9 @@
237
  },
238
  {
239
  "id": "maskIndex",
240
- "name": "EmbedLayerNormalization.MaskIndex",
241
  "shader": "embed-mask-index.wgsl.jinja",
242
- "bindings": ["mask", "mask_index", "params_2"],
243
  "dispatch": {
244
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
245
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -249,14 +256,14 @@
249
  ]
250
  },
251
  {
252
- "id": "noseg_posids_sum_nomask",
253
- "when": ["not present.segmentEmbeddingT", "present.positionIdsT", "present.embeddingSumT", "not present.maskIndexT"],
254
  "passes": [
255
  {
256
  "id": "sum",
257
  "name": "EmbedLayerNormalization.EmbeddingSum",
258
  "shader": "embed-sum.wgsl.jinja",
259
- "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "embedding_sum", "params"],
260
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
261
  },
262
  {
@@ -269,14 +276,14 @@
269
  ]
270
  },
271
  {
272
- "id": "noseg_posids_sum_mask",
273
- "when": ["not present.segmentEmbeddingT", "present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "present.maskT"],
274
  "passes": [
275
  {
276
  "id": "sum",
277
  "name": "EmbedLayerNormalization.EmbeddingSum",
278
  "shader": "embed-sum.wgsl.jinja",
279
- "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "embedding_sum", "params"],
280
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
281
  },
282
  {
@@ -290,7 +297,7 @@
290
  "id": "maskIndex",
291
  "name": "EmbedLayerNormalization.MaskIndex",
292
  "shader": "embed-mask-index.wgsl.jinja",
293
- "bindings": ["mask", "mask_index", "params_2"],
294
  "dispatch": {
295
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
296
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -300,34 +307,14 @@
300
  ]
301
  },
302
  {
303
- "id": "seg_nopos_nosum_nomask",
304
- "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "not present.positionIdsT", "not present.embeddingSumT", "not present.maskIndexT"],
305
- "passes": [
306
- {
307
- "id": "sum",
308
- "name": "EmbedLayerNormalization.EmbeddingSum",
309
- "shader": "embed-sum.wgsl.jinja",
310
- "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
311
- "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
312
- },
313
- {
314
- "id": "normalize",
315
- "name": "EmbedLayerNormalization.Normalize",
316
- "shader": "embed-normalize.wgsl.jinja",
317
- "bindings": ["output", "gamma", "beta", "params"],
318
- "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
319
- }
320
- ]
321
- },
322
- {
323
- "id": "seg_nopos_nosum_mask",
324
- "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "not present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "present.maskT"],
325
  "passes": [
326
  {
327
  "id": "sum",
328
  "name": "EmbedLayerNormalization.EmbeddingSum",
329
  "shader": "embed-sum.wgsl.jinja",
330
- "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
331
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
332
  },
333
  {
@@ -339,9 +326,9 @@
339
  },
340
  {
341
  "id": "maskIndex",
342
- "name": "EmbedLayerNormalization.MaskIndex",
343
  "shader": "embed-mask-index.wgsl.jinja",
344
- "bindings": ["mask", "mask_index", "params_2"],
345
  "dispatch": {
346
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
347
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -351,14 +338,14 @@
351
  ]
352
  },
353
  {
354
- "id": "seg_nopos_sum_nomask",
355
- "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "not present.positionIdsT", "present.embeddingSumT", "not present.maskIndexT"],
356
  "passes": [
357
  {
358
  "id": "sum",
359
  "name": "EmbedLayerNormalization.EmbeddingSum",
360
  "shader": "embed-sum.wgsl.jinja",
361
- "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
362
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
363
  },
364
  {
@@ -371,14 +358,14 @@
371
  ]
372
  },
373
  {
374
- "id": "seg_nopos_sum_mask",
375
- "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "not present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "present.maskT"],
376
  "passes": [
377
  {
378
  "id": "sum",
379
  "name": "EmbedLayerNormalization.EmbeddingSum",
380
  "shader": "embed-sum.wgsl.jinja",
381
- "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
382
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
383
  },
384
  {
@@ -392,7 +379,7 @@
392
  "id": "maskIndex",
393
  "name": "EmbedLayerNormalization.MaskIndex",
394
  "shader": "embed-mask-index.wgsl.jinja",
395
- "bindings": ["mask", "mask_index", "params_2"],
396
  "dispatch": {
397
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
398
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -402,14 +389,14 @@
402
  ]
403
  },
404
  {
405
- "id": "seg_posids_nosum_nomask",
406
- "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "present.positionIdsT", "not present.embeddingSumT", "not present.maskIndexT"],
407
  "passes": [
408
  {
409
  "id": "sum",
410
  "name": "EmbedLayerNormalization.EmbeddingSum",
411
  "shader": "embed-sum.wgsl.jinja",
412
- "bindings": ["input_ids", "segment_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
413
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
414
  },
415
  {
@@ -418,18 +405,29 @@
418
  "shader": "embed-normalize.wgsl.jinja",
419
  "bindings": ["output", "gamma", "beta", "params"],
420
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
421
  }
422
  ]
423
  },
424
  {
425
- "id": "seg_posids_nosum_mask",
426
- "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "present.maskT"],
427
  "passes": [
428
  {
429
  "id": "sum",
430
  "name": "EmbedLayerNormalization.EmbeddingSum",
431
  "shader": "embed-sum.wgsl.jinja",
432
- "bindings": ["input_ids", "segment_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
433
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
434
  },
435
  {
@@ -438,29 +436,18 @@
438
  "shader": "embed-normalize.wgsl.jinja",
439
  "bindings": ["output", "gamma", "beta", "params"],
440
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
441
- },
442
- {
443
- "id": "maskIndex",
444
- "name": "EmbedLayerNormalization.MaskIndex",
445
- "shader": "embed-mask-index.wgsl.jinja",
446
- "bindings": ["mask", "mask_index", "params_2"],
447
- "dispatch": {
448
- "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
449
- "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
450
- "z": 1
451
- }
452
  }
453
  ]
454
  },
455
  {
456
- "id": "seg_posids_sum_nomask",
457
- "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "present.positionIdsT", "present.embeddingSumT", "not present.maskIndexT"],
458
  "passes": [
459
  {
460
  "id": "sum",
461
  "name": "EmbedLayerNormalization.EmbeddingSum",
462
  "shader": "embed-sum.wgsl.jinja",
463
- "bindings": ["input_ids", "segment_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
464
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
465
  },
466
  {
@@ -469,18 +456,29 @@
469
  "shader": "embed-normalize.wgsl.jinja",
470
  "bindings": ["output", "gamma", "beta", "params"],
471
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
472
  }
473
  ]
474
  },
475
  {
476
- "id": "seg_posids_sum_mask",
477
- "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "present.maskT"],
478
  "passes": [
479
  {
480
  "id": "sum",
481
  "name": "EmbedLayerNormalization.EmbeddingSum",
482
  "shader": "embed-sum.wgsl.jinja",
483
- "bindings": ["input_ids", "segment_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
484
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
485
  },
486
  {
@@ -492,9 +490,9 @@
492
  },
493
  {
494
  "id": "maskIndex",
495
- "name": "EmbedLayerNormalization.MaskIndex",
496
  "shader": "embed-mask-index.wgsl.jinja",
497
- "bindings": ["mask", "mask_index", "params_2"],
498
  "dispatch": {
499
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
500
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -504,14 +502,14 @@
504
  ]
505
  },
506
  {
507
- "id": "segdefault_nopos_nosum_nomask",
508
- "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT == (\"nopos\" == \"posids\")", "present.embeddingSumT == (\"nosum\" == \"sum\")", "not present.maskIndexT"],
509
  "passes": [
510
  {
511
  "id": "sum",
512
  "name": "EmbedLayerNormalization.EmbeddingSum",
513
  "shader": "embed-sum.wgsl.jinja",
514
- "bindings": ["input_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
515
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
516
  },
517
  {
@@ -524,14 +522,14 @@
524
  ]
525
  },
526
  {
527
- "id": "segdefault_nopos_nosum_mask",
528
- "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT == (\"nopos\" == \"posids\")", "present.embeddingSumT == (\"nosum\" == \"sum\")", "present.maskIndexT", "present.maskT"],
529
  "passes": [
530
  {
531
  "id": "sum",
532
  "name": "EmbedLayerNormalization.EmbeddingSum",
533
  "shader": "embed-sum.wgsl.jinja",
534
- "bindings": ["input_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
535
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
536
  },
537
  {
@@ -545,7 +543,7 @@
545
  "id": "maskIndex",
546
  "name": "EmbedLayerNormalization.MaskIndex",
547
  "shader": "embed-mask-index.wgsl.jinja",
548
- "bindings": ["mask", "mask_index", "params_2"],
549
  "dispatch": {
550
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
551
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -555,8 +553,8 @@
555
  ]
556
  },
557
  {
558
- "id": "segdefault_nopos_sum_nomask",
559
- "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT == (\"nopos\" == \"posids\")", "present.embeddingSumT == (\"sum\" == \"sum\")", "not present.maskIndexT"],
560
  "passes": [
561
  {
562
  "id": "sum",
@@ -571,18 +569,29 @@
571
  "shader": "embed-normalize.wgsl.jinja",
572
  "bindings": ["output", "gamma", "beta", "params"],
573
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
574
  }
575
  ]
576
  },
577
  {
578
- "id": "segdefault_nopos_sum_mask",
579
- "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT == (\"nopos\" == \"posids\")", "present.embeddingSumT == (\"sum\" == \"sum\")", "present.maskIndexT", "present.maskT"],
580
  "passes": [
581
  {
582
  "id": "sum",
583
  "name": "EmbedLayerNormalization.EmbeddingSum",
584
  "shader": "embed-sum.wgsl.jinja",
585
- "bindings": ["input_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
586
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
587
  },
588
  {
@@ -591,23 +600,12 @@
591
  "shader": "embed-normalize.wgsl.jinja",
592
  "bindings": ["output", "gamma", "beta", "params"],
593
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
594
- },
595
- {
596
- "id": "maskIndex",
597
- "name": "EmbedLayerNormalization.MaskIndex",
598
- "shader": "embed-mask-index.wgsl.jinja",
599
- "bindings": ["mask", "mask_index", "params_2"],
600
- "dispatch": {
601
- "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
602
- "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
603
- "z": 1
604
- }
605
  }
606
  ]
607
  },
608
  {
609
- "id": "segdefault_posids_nosum_nomask",
610
- "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT == (\"posids\" == \"posids\")", "present.embeddingSumT == (\"nosum\" == \"sum\")", "not present.maskIndexT"],
611
  "passes": [
612
  {
613
  "id": "sum",
@@ -622,12 +620,23 @@
622
  "shader": "embed-normalize.wgsl.jinja",
623
  "bindings": ["output", "gamma", "beta", "params"],
624
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
625
  }
626
  ]
627
  },
628
  {
629
- "id": "segdefault_posids_nosum_mask",
630
- "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT == (\"posids\" == \"posids\")", "present.embeddingSumT == (\"nosum\" == \"sum\")", "present.maskIndexT", "present.maskT"],
631
  "passes": [
632
  {
633
  "id": "sum",
@@ -645,9 +654,9 @@
645
  },
646
  {
647
  "id": "maskIndex",
648
- "name": "EmbedLayerNormalization.MaskIndex",
649
  "shader": "embed-mask-index.wgsl.jinja",
650
- "bindings": ["mask", "mask_index", "params_2"],
651
  "dispatch": {
652
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
653
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -658,7 +667,7 @@
658
  },
659
  {
660
  "id": "segdefault_posids_sum_nomask",
661
- "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT == (\"posids\" == \"posids\")", "present.embeddingSumT == (\"sum\" == \"sum\")", "not present.maskIndexT"],
662
  "passes": [
663
  {
664
  "id": "sum",
@@ -678,7 +687,7 @@
678
  },
679
  {
680
  "id": "segdefault_posids_sum_mask",
681
- "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT == (\"posids\" == \"posids\")", "present.embeddingSumT == (\"sum\" == \"sum\")", "present.maskIndexT", "present.maskT"],
682
  "passes": [
683
  {
684
  "id": "sum",
@@ -698,7 +707,7 @@
698
  "id": "maskIndex",
699
  "name": "EmbedLayerNormalization.MaskIndex",
700
  "shader": "embed-mask-index.wgsl.jinja",
701
- "bindings": ["mask", "mask_index", "params_2"],
702
  "dispatch": {
703
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
704
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -708,14 +717,14 @@
708
  ]
709
  },
710
  {
711
- "id": "noseg_nopos_nosum_mask_without_input",
712
- "when": ["present.segmentEmbeddingT == (\"noseg\" != \"noseg\")", "present.segmentIdsT == (\"noseg\" == \"seg\")", "present.positionIdsT == (\"nopos\" == \"posids\")", "present.embeddingSumT == (\"nosum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
713
  "passes": [
714
  {
715
  "id": "sum",
716
  "name": "EmbedLayerNormalization.EmbeddingSum",
717
  "shader": "embed-sum.wgsl.jinja",
718
- "bindings": ["input_ids", "word_embedding", "position_embedding", "output", "params"],
719
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
720
  },
721
  {
@@ -729,7 +738,7 @@
729
  "id": "maskIndex",
730
  "name": "EmbedLayerNormalization.ZeroMaskIndex",
731
  "shader": "embed-mask-index.wgsl.jinja",
732
- "bindings": ["mask_index", "params_2"],
733
  "dispatch": {
734
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
735
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -739,14 +748,14 @@
739
  ]
740
  },
741
  {
742
- "id": "noseg_nopos_sum_mask_without_input",
743
- "when": ["present.segmentEmbeddingT == (\"noseg\" != \"noseg\")", "present.segmentIdsT == (\"noseg\" == \"seg\")", "present.positionIdsT == (\"nopos\" == \"posids\")", "present.embeddingSumT == (\"sum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
744
  "passes": [
745
  {
746
  "id": "sum",
747
  "name": "EmbedLayerNormalization.EmbeddingSum",
748
  "shader": "embed-sum.wgsl.jinja",
749
- "bindings": ["input_ids", "word_embedding", "position_embedding", "output", "embedding_sum", "params"],
750
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
751
  },
752
  {
@@ -755,29 +764,18 @@
755
  "shader": "embed-normalize.wgsl.jinja",
756
  "bindings": ["output", "gamma", "beta", "params"],
757
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
758
- },
759
- {
760
- "id": "maskIndex",
761
- "name": "EmbedLayerNormalization.ZeroMaskIndex",
762
- "shader": "embed-mask-index.wgsl.jinja",
763
- "bindings": ["mask_index", "params_2"],
764
- "dispatch": {
765
- "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
766
- "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
767
- "z": 1
768
- }
769
  }
770
  ]
771
  },
772
  {
773
- "id": "noseg_posids_nosum_mask_without_input",
774
- "when": ["present.segmentEmbeddingT == (\"noseg\" != \"noseg\")", "present.segmentIdsT == (\"noseg\" == \"seg\")", "present.positionIdsT == (\"posids\" == \"posids\")", "present.embeddingSumT == (\"nosum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
775
  "passes": [
776
  {
777
  "id": "sum",
778
  "name": "EmbedLayerNormalization.EmbeddingSum",
779
  "shader": "embed-sum.wgsl.jinja",
780
- "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "params"],
781
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
782
  },
783
  {
@@ -789,9 +787,9 @@
789
  },
790
  {
791
  "id": "maskIndex",
792
- "name": "EmbedLayerNormalization.ZeroMaskIndex",
793
  "shader": "embed-mask-index.wgsl.jinja",
794
- "bindings": ["mask_index", "params_2"],
795
  "dispatch": {
796
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
797
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -801,14 +799,14 @@
801
  ]
802
  },
803
  {
804
- "id": "noseg_posids_sum_mask_without_input",
805
- "when": ["present.segmentEmbeddingT == (\"noseg\" != \"noseg\")", "present.segmentIdsT == (\"noseg\" == \"seg\")", "present.positionIdsT == (\"posids\" == \"posids\")", "present.embeddingSumT == (\"sum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
806
  "passes": [
807
  {
808
  "id": "sum",
809
  "name": "EmbedLayerNormalization.EmbeddingSum",
810
  "shader": "embed-sum.wgsl.jinja",
811
- "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "embedding_sum", "params"],
812
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
813
  },
814
  {
@@ -822,7 +820,7 @@
822
  "id": "maskIndex",
823
  "name": "EmbedLayerNormalization.ZeroMaskIndex",
824
  "shader": "embed-mask-index.wgsl.jinja",
825
- "bindings": ["mask_index", "params_2"],
826
  "dispatch": {
827
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
828
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -832,14 +830,14 @@
832
  ]
833
  },
834
  {
835
- "id": "segdefault_nopos_nosum_mask_without_input",
836
- "when": ["present.segmentEmbeddingT == (\"segdefault\" != \"noseg\")", "present.segmentIdsT == (\"segdefault\" == \"seg\")", "present.positionIdsT == (\"nopos\" == \"posids\")", "present.embeddingSumT == (\"nosum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
837
  "passes": [
838
  {
839
  "id": "sum",
840
  "name": "EmbedLayerNormalization.EmbeddingSum",
841
  "shader": "embed-sum.wgsl.jinja",
842
- "bindings": ["input_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
843
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
844
  },
845
  {
@@ -848,29 +846,18 @@
848
  "shader": "embed-normalize.wgsl.jinja",
849
  "bindings": ["output", "gamma", "beta", "params"],
850
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
851
- },
852
- {
853
- "id": "maskIndex",
854
- "name": "EmbedLayerNormalization.ZeroMaskIndex",
855
- "shader": "embed-mask-index.wgsl.jinja",
856
- "bindings": ["mask_index", "params_2"],
857
- "dispatch": {
858
- "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
859
- "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
860
- "z": 1
861
- }
862
  }
863
  ]
864
  },
865
  {
866
- "id": "segdefault_nopos_sum_mask_without_input",
867
- "when": ["present.segmentEmbeddingT == (\"segdefault\" != \"noseg\")", "present.segmentIdsT == (\"segdefault\" == \"seg\")", "present.positionIdsT == (\"nopos\" == \"posids\")", "present.embeddingSumT == (\"sum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
868
  "passes": [
869
  {
870
  "id": "sum",
871
  "name": "EmbedLayerNormalization.EmbeddingSum",
872
  "shader": "embed-sum.wgsl.jinja",
873
- "bindings": ["input_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
874
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
875
  },
876
  {
@@ -882,9 +869,9 @@
882
  },
883
  {
884
  "id": "maskIndex",
885
- "name": "EmbedLayerNormalization.ZeroMaskIndex",
886
  "shader": "embed-mask-index.wgsl.jinja",
887
- "bindings": ["mask_index", "params_2"],
888
  "dispatch": {
889
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
890
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -894,14 +881,14 @@
894
  ]
895
  },
896
  {
897
- "id": "segdefault_posids_nosum_mask_without_input",
898
- "when": ["present.segmentEmbeddingT == (\"segdefault\" != \"noseg\")", "present.segmentIdsT == (\"segdefault\" == \"seg\")", "present.positionIdsT == (\"posids\" == \"posids\")", "present.embeddingSumT == (\"nosum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
899
  "passes": [
900
  {
901
  "id": "sum",
902
  "name": "EmbedLayerNormalization.EmbeddingSum",
903
  "shader": "embed-sum.wgsl.jinja",
904
- "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
905
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
906
  },
907
  {
@@ -915,7 +902,7 @@
915
  "id": "maskIndex",
916
  "name": "EmbedLayerNormalization.ZeroMaskIndex",
917
  "shader": "embed-mask-index.wgsl.jinja",
918
- "bindings": ["mask_index", "params_2"],
919
  "dispatch": {
920
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
921
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -925,14 +912,14 @@
925
  ]
926
  },
927
  {
928
- "id": "segdefault_posids_sum_mask_without_input",
929
- "when": ["present.segmentEmbeddingT == (\"segdefault\" != \"noseg\")", "present.segmentIdsT == (\"segdefault\" == \"seg\")", "present.positionIdsT == (\"posids\" == \"posids\")", "present.embeddingSumT == (\"sum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
930
  "passes": [
931
  {
932
  "id": "sum",
933
  "name": "EmbedLayerNormalization.EmbeddingSum",
934
  "shader": "embed-sum.wgsl.jinja",
935
- "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
936
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
937
  },
938
  {
@@ -941,29 +928,18 @@
941
  "shader": "embed-normalize.wgsl.jinja",
942
  "bindings": ["output", "gamma", "beta", "params"],
943
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
944
- },
945
- {
946
- "id": "maskIndex",
947
- "name": "EmbedLayerNormalization.ZeroMaskIndex",
948
- "shader": "embed-mask-index.wgsl.jinja",
949
- "bindings": ["mask_index", "params_2"],
950
- "dispatch": {
951
- "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
952
- "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
953
- "z": 1
954
- }
955
  }
956
  ]
957
  },
958
  {
959
- "id": "seg_nopos_nosum_mask_without_input",
960
- "when": ["present.segmentEmbeddingT == (\"seg\" != \"noseg\")", "present.segmentIdsT == (\"seg\" == \"seg\")", "present.positionIdsT == (\"nopos\" == \"posids\")", "present.embeddingSumT == (\"nosum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
961
  "passes": [
962
  {
963
  "id": "sum",
964
  "name": "EmbedLayerNormalization.EmbeddingSum",
965
  "shader": "embed-sum.wgsl.jinja",
966
- "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
967
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
968
  },
969
  {
@@ -975,9 +951,9 @@
975
  },
976
  {
977
  "id": "maskIndex",
978
- "name": "EmbedLayerNormalization.ZeroMaskIndex",
979
  "shader": "embed-mask-index.wgsl.jinja",
980
- "bindings": ["mask_index", "params_2"],
981
  "dispatch": {
982
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
983
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -987,14 +963,14 @@
987
  ]
988
  },
989
  {
990
- "id": "seg_nopos_sum_mask_without_input",
991
- "when": ["present.segmentEmbeddingT == (\"seg\" != \"noseg\")", "present.segmentIdsT == (\"seg\" == \"seg\")", "present.positionIdsT == (\"nopos\" == \"posids\")", "present.embeddingSumT == (\"sum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
992
  "passes": [
993
  {
994
  "id": "sum",
995
  "name": "EmbedLayerNormalization.EmbeddingSum",
996
  "shader": "embed-sum.wgsl.jinja",
997
- "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
998
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
999
  },
1000
  {
@@ -1008,7 +984,7 @@
1008
  "id": "maskIndex",
1009
  "name": "EmbedLayerNormalization.ZeroMaskIndex",
1010
  "shader": "embed-mask-index.wgsl.jinja",
1011
- "bindings": ["mask_index", "params_2"],
1012
  "dispatch": {
1013
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
1014
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -1018,14 +994,34 @@
1018
  ]
1019
  },
1020
  {
1021
- "id": "seg_posids_nosum_mask_without_input",
1022
- "when": ["present.segmentEmbeddingT == (\"seg\" != \"noseg\")", "present.segmentIdsT == (\"seg\" == \"seg\")", "present.positionIdsT == (\"posids\" == \"posids\")", "present.embeddingSumT == (\"nosum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
1023
  "passes": [
1024
  {
1025
  "id": "sum",
1026
  "name": "EmbedLayerNormalization.EmbeddingSum",
1027
  "shader": "embed-sum.wgsl.jinja",
1028
- "bindings": ["input_ids", "segment_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1029
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
1030
  },
1031
  {
@@ -1037,9 +1033,9 @@
1037
  },
1038
  {
1039
  "id": "maskIndex",
1040
- "name": "EmbedLayerNormalization.ZeroMaskIndex",
1041
  "shader": "embed-mask-index.wgsl.jinja",
1042
- "bindings": ["mask_index", "params_2"],
1043
  "dispatch": {
1044
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
1045
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
@@ -1050,7 +1046,7 @@
1050
  },
1051
  {
1052
  "id": "seg_posids_sum_mask_without_input",
1053
- "when": ["present.segmentEmbeddingT == (\"seg\" != \"noseg\")", "present.segmentIdsT == (\"seg\" == \"seg\")", "present.positionIdsT == (\"posids\" == \"posids\")", "present.embeddingSumT == (\"sum\" == \"sum\")", "present.maskIndexT", "not present.maskT"],
1054
  "passes": [
1055
  {
1056
  "id": "sum",
@@ -1070,7 +1066,7 @@
1070
  "id": "maskIndex",
1071
  "name": "EmbedLayerNormalization.ZeroMaskIndex",
1072
  "shader": "embed-mask-index.wgsl.jinja",
1073
- "bindings": ["mask_index", "params_2"],
1074
  "dispatch": {
1075
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
1076
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
75
  },
76
  "when": ["embedContractOk", "dispatchFits"],
77
  "bindings": {
78
+ "input_ids": { "arg": "inputIdsT", "elementType": "i32" },
79
+ "word_embedding": { "arg": "wordEmbeddingT", "elementType": "$aScalar" },
80
+ "position_embedding": { "arg": "positionEmbeddingT", "elementType": "$aScalar" },
81
+ "output": { "arg": "outputT", "elementType": "$aScalar" },
82
+ "params": { "struct": [{ "name": "tokens", "type": "u32", "value": "tokens" }] },
83
+ "embedding_sum": { "arg": "embeddingSumT", "elementType": "$aScalar" },
84
+ "position_ids": { "arg": "positionIdsT", "elementType": "i32" },
85
+ "segment_ids": { "arg": "segmentIdsT", "elementType": "i32" },
86
+ "segment_embedding": { "arg": "segmentEmbeddingT", "elementType": "$aScalar" },
87
+ "gamma": { "arg": "gammaT", "elementType": "$aScalar", "length": "$HIDDEN_LEN" },
88
+ "beta": { "arg": "betaT", "elementType": "$aScalar", "length": "$HIDDEN_LEN" },
89
+ "mask": { "arg": "maskT", "elementType": "i32" },
90
+ "mask_index": { "arg": "maskIndexT", "elementType": "i32" },
91
+ "params_batch": { "name": "params", "struct": [{ "name": "batch", "type": "u32", "value": "batchSize" }] }
 
 
 
 
92
  },
93
  "variants": [
94
  {
95
  "id": "noseg_nopos_nosum_nomask",
96
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "not present.embeddingSumT", "not present.maskIndexT"],
97
  "passes": [
98
  {
99
  "id": "sum",
 
113
  },
114
  {
115
  "id": "noseg_nopos_nosum_mask",
116
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "present.maskT"],
117
  "passes": [
118
  {
119
  "id": "sum",
 
133
  "id": "maskIndex",
134
  "name": "EmbedLayerNormalization.MaskIndex",
135
  "shader": "embed-mask-index.wgsl.jinja",
136
+ "bindings": ["mask", "mask_index", "params_batch"],
137
  "dispatch": {
138
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
139
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
143
  ]
144
  },
145
  {
146
+ "id": "noseg_nopos_nosum_mask_without_input",
147
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
148
  "passes": [
149
  {
150
  "id": "sum",
151
  "name": "EmbedLayerNormalization.EmbeddingSum",
152
  "shader": "embed-sum.wgsl.jinja",
153
+ "bindings": ["input_ids", "word_embedding", "position_embedding", "output", "params"],
154
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
155
  },
156
  {
 
159
  "shader": "embed-normalize.wgsl.jinja",
160
  "bindings": ["output", "gamma", "beta", "params"],
161
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
162
+ },
163
+ {
164
+ "id": "maskIndex",
165
+ "name": "EmbedLayerNormalization.ZeroMaskIndex",
166
+ "shader": "embed-mask-index.wgsl.jinja",
167
+ "bindings": ["mask_index", "params_batch"],
168
+ "dispatch": {
169
+ "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
170
+ "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
171
+ "z": 1
172
+ }
173
  }
174
  ]
175
  },
176
  {
177
+ "id": "noseg_nopos_sum_nomask",
178
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "present.embeddingSumT", "not present.maskIndexT"],
179
  "passes": [
180
  {
181
  "id": "sum",
 
190
  "shader": "embed-normalize.wgsl.jinja",
191
  "bindings": ["output", "gamma", "beta", "params"],
192
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
193
  }
194
  ]
195
  },
196
  {
197
+ "id": "noseg_nopos_sum_mask",
198
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "present.maskT"],
199
  "passes": [
200
  {
201
  "id": "sum",
202
  "name": "EmbedLayerNormalization.EmbeddingSum",
203
  "shader": "embed-sum.wgsl.jinja",
204
+ "bindings": ["input_ids", "word_embedding", "position_embedding", "output", "embedding_sum", "params"],
205
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
206
  },
207
  {
 
210
  "shader": "embed-normalize.wgsl.jinja",
211
  "bindings": ["output", "gamma", "beta", "params"],
212
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
213
+ },
214
+ {
215
+ "id": "maskIndex",
216
+ "name": "EmbedLayerNormalization.MaskIndex",
217
+ "shader": "embed-mask-index.wgsl.jinja",
218
+ "bindings": ["mask", "mask_index", "params_batch"],
219
+ "dispatch": {
220
+ "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
221
+ "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
222
+ "z": 1
223
+ }
224
  }
225
  ]
226
  },
227
  {
228
+ "id": "noseg_nopos_sum_mask_without_input",
229
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
230
  "passes": [
231
  {
232
  "id": "sum",
233
  "name": "EmbedLayerNormalization.EmbeddingSum",
234
  "shader": "embed-sum.wgsl.jinja",
235
+ "bindings": ["input_ids", "word_embedding", "position_embedding", "output", "embedding_sum", "params"],
236
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
237
  },
238
  {
 
244
  },
245
  {
246
  "id": "maskIndex",
247
+ "name": "EmbedLayerNormalization.ZeroMaskIndex",
248
  "shader": "embed-mask-index.wgsl.jinja",
249
+ "bindings": ["mask_index", "params_batch"],
250
  "dispatch": {
251
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
252
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
256
  ]
257
  },
258
  {
259
+ "id": "noseg_posids_nosum_nomask",
260
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "not present.embeddingSumT", "not present.maskIndexT"],
261
  "passes": [
262
  {
263
  "id": "sum",
264
  "name": "EmbedLayerNormalization.EmbeddingSum",
265
  "shader": "embed-sum.wgsl.jinja",
266
+ "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "params"],
267
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
268
  },
269
  {
 
276
  ]
277
  },
278
  {
279
+ "id": "noseg_posids_nosum_mask",
280
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "present.maskT"],
281
  "passes": [
282
  {
283
  "id": "sum",
284
  "name": "EmbedLayerNormalization.EmbeddingSum",
285
  "shader": "embed-sum.wgsl.jinja",
286
+ "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "params"],
287
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
288
  },
289
  {
 
297
  "id": "maskIndex",
298
  "name": "EmbedLayerNormalization.MaskIndex",
299
  "shader": "embed-mask-index.wgsl.jinja",
300
+ "bindings": ["mask", "mask_index", "params_batch"],
301
  "dispatch": {
302
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
303
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
307
  ]
308
  },
309
  {
310
+ "id": "noseg_posids_nosum_mask_without_input",
311
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
312
  "passes": [
313
  {
314
  "id": "sum",
315
  "name": "EmbedLayerNormalization.EmbeddingSum",
316
  "shader": "embed-sum.wgsl.jinja",
317
+ "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "params"],
318
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
319
  },
320
  {
 
326
  },
327
  {
328
  "id": "maskIndex",
329
+ "name": "EmbedLayerNormalization.ZeroMaskIndex",
330
  "shader": "embed-mask-index.wgsl.jinja",
331
+ "bindings": ["mask_index", "params_batch"],
332
  "dispatch": {
333
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
334
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
338
  ]
339
  },
340
  {
341
+ "id": "noseg_posids_sum_nomask",
342
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "present.embeddingSumT", "not present.maskIndexT"],
343
  "passes": [
344
  {
345
  "id": "sum",
346
  "name": "EmbedLayerNormalization.EmbeddingSum",
347
  "shader": "embed-sum.wgsl.jinja",
348
+ "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "embedding_sum", "params"],
349
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
350
  },
351
  {
 
358
  ]
359
  },
360
  {
361
+ "id": "noseg_posids_sum_mask",
362
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "present.maskT"],
363
  "passes": [
364
  {
365
  "id": "sum",
366
  "name": "EmbedLayerNormalization.EmbeddingSum",
367
  "shader": "embed-sum.wgsl.jinja",
368
+ "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "embedding_sum", "params"],
369
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
370
  },
371
  {
 
379
  "id": "maskIndex",
380
  "name": "EmbedLayerNormalization.MaskIndex",
381
  "shader": "embed-mask-index.wgsl.jinja",
382
+ "bindings": ["mask", "mask_index", "params_batch"],
383
  "dispatch": {
384
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
385
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
389
  ]
390
  },
391
  {
392
+ "id": "noseg_posids_sum_mask_without_input",
393
+ "when": ["not present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
394
  "passes": [
395
  {
396
  "id": "sum",
397
  "name": "EmbedLayerNormalization.EmbeddingSum",
398
  "shader": "embed-sum.wgsl.jinja",
399
+ "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "output", "embedding_sum", "params"],
400
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
401
  },
402
  {
 
405
  "shader": "embed-normalize.wgsl.jinja",
406
  "bindings": ["output", "gamma", "beta", "params"],
407
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
408
+ },
409
+ {
410
+ "id": "maskIndex",
411
+ "name": "EmbedLayerNormalization.ZeroMaskIndex",
412
+ "shader": "embed-mask-index.wgsl.jinja",
413
+ "bindings": ["mask_index", "params_batch"],
414
+ "dispatch": {
415
+ "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
416
+ "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
417
+ "z": 1
418
+ }
419
  }
420
  ]
421
  },
422
  {
423
+ "id": "segdefault_nopos_nosum_nomask",
424
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "not present.embeddingSumT", "not present.maskIndexT"],
425
  "passes": [
426
  {
427
  "id": "sum",
428
  "name": "EmbedLayerNormalization.EmbeddingSum",
429
  "shader": "embed-sum.wgsl.jinja",
430
+ "bindings": ["input_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
431
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
432
  },
433
  {
 
436
  "shader": "embed-normalize.wgsl.jinja",
437
  "bindings": ["output", "gamma", "beta", "params"],
438
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
439
  }
440
  ]
441
  },
442
  {
443
+ "id": "segdefault_nopos_nosum_mask",
444
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "present.maskT"],
445
  "passes": [
446
  {
447
  "id": "sum",
448
  "name": "EmbedLayerNormalization.EmbeddingSum",
449
  "shader": "embed-sum.wgsl.jinja",
450
+ "bindings": ["input_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
451
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
452
  },
453
  {
 
456
  "shader": "embed-normalize.wgsl.jinja",
457
  "bindings": ["output", "gamma", "beta", "params"],
458
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
459
+ },
460
+ {
461
+ "id": "maskIndex",
462
+ "name": "EmbedLayerNormalization.MaskIndex",
463
+ "shader": "embed-mask-index.wgsl.jinja",
464
+ "bindings": ["mask", "mask_index", "params_batch"],
465
+ "dispatch": {
466
+ "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
467
+ "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
468
+ "z": 1
469
+ }
470
  }
471
  ]
472
  },
473
  {
474
+ "id": "segdefault_nopos_nosum_mask_without_input",
475
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
476
  "passes": [
477
  {
478
  "id": "sum",
479
  "name": "EmbedLayerNormalization.EmbeddingSum",
480
  "shader": "embed-sum.wgsl.jinja",
481
+ "bindings": ["input_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
482
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
483
  },
484
  {
 
490
  },
491
  {
492
  "id": "maskIndex",
493
+ "name": "EmbedLayerNormalization.ZeroMaskIndex",
494
  "shader": "embed-mask-index.wgsl.jinja",
495
+ "bindings": ["mask_index", "params_batch"],
496
  "dispatch": {
497
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
498
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
502
  ]
503
  },
504
  {
505
+ "id": "segdefault_nopos_sum_nomask",
506
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "present.embeddingSumT", "not present.maskIndexT"],
507
  "passes": [
508
  {
509
  "id": "sum",
510
  "name": "EmbedLayerNormalization.EmbeddingSum",
511
  "shader": "embed-sum.wgsl.jinja",
512
+ "bindings": ["input_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
513
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
514
  },
515
  {
 
522
  ]
523
  },
524
  {
525
+ "id": "segdefault_nopos_sum_mask",
526
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "present.maskT"],
527
  "passes": [
528
  {
529
  "id": "sum",
530
  "name": "EmbedLayerNormalization.EmbeddingSum",
531
  "shader": "embed-sum.wgsl.jinja",
532
+ "bindings": ["input_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
533
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
534
  },
535
  {
 
543
  "id": "maskIndex",
544
  "name": "EmbedLayerNormalization.MaskIndex",
545
  "shader": "embed-mask-index.wgsl.jinja",
546
+ "bindings": ["mask", "mask_index", "params_batch"],
547
  "dispatch": {
548
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
549
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
553
  ]
554
  },
555
  {
556
+ "id": "segdefault_nopos_sum_mask_without_input",
557
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "not present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
558
  "passes": [
559
  {
560
  "id": "sum",
 
569
  "shader": "embed-normalize.wgsl.jinja",
570
  "bindings": ["output", "gamma", "beta", "params"],
571
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
572
+ },
573
+ {
574
+ "id": "maskIndex",
575
+ "name": "EmbedLayerNormalization.ZeroMaskIndex",
576
+ "shader": "embed-mask-index.wgsl.jinja",
577
+ "bindings": ["mask_index", "params_batch"],
578
+ "dispatch": {
579
+ "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
580
+ "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
581
+ "z": 1
582
+ }
583
  }
584
  ]
585
  },
586
  {
587
+ "id": "segdefault_posids_nosum_nomask",
588
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "not present.embeddingSumT", "not present.maskIndexT"],
589
  "passes": [
590
  {
591
  "id": "sum",
592
  "name": "EmbedLayerNormalization.EmbeddingSum",
593
  "shader": "embed-sum.wgsl.jinja",
594
+ "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
595
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
596
  },
597
  {
 
600
  "shader": "embed-normalize.wgsl.jinja",
601
  "bindings": ["output", "gamma", "beta", "params"],
602
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
603
  }
604
  ]
605
  },
606
  {
607
+ "id": "segdefault_posids_nosum_mask",
608
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "present.maskT"],
609
  "passes": [
610
  {
611
  "id": "sum",
 
620
  "shader": "embed-normalize.wgsl.jinja",
621
  "bindings": ["output", "gamma", "beta", "params"],
622
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
623
+ },
624
+ {
625
+ "id": "maskIndex",
626
+ "name": "EmbedLayerNormalization.MaskIndex",
627
+ "shader": "embed-mask-index.wgsl.jinja",
628
+ "bindings": ["mask", "mask_index", "params_batch"],
629
+ "dispatch": {
630
+ "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
631
+ "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
632
+ "z": 1
633
+ }
634
  }
635
  ]
636
  },
637
  {
638
+ "id": "segdefault_posids_nosum_mask_without_input",
639
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
640
  "passes": [
641
  {
642
  "id": "sum",
 
654
  },
655
  {
656
  "id": "maskIndex",
657
+ "name": "EmbedLayerNormalization.ZeroMaskIndex",
658
  "shader": "embed-mask-index.wgsl.jinja",
659
+ "bindings": ["mask_index", "params_batch"],
660
  "dispatch": {
661
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
662
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
667
  },
668
  {
669
  "id": "segdefault_posids_sum_nomask",
670
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "present.embeddingSumT", "not present.maskIndexT"],
671
  "passes": [
672
  {
673
  "id": "sum",
 
687
  },
688
  {
689
  "id": "segdefault_posids_sum_mask",
690
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "present.maskT"],
691
  "passes": [
692
  {
693
  "id": "sum",
 
707
  "id": "maskIndex",
708
  "name": "EmbedLayerNormalization.MaskIndex",
709
  "shader": "embed-mask-index.wgsl.jinja",
710
+ "bindings": ["mask", "mask_index", "params_batch"],
711
  "dispatch": {
712
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
713
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
717
  ]
718
  },
719
  {
720
+ "id": "segdefault_posids_sum_mask_without_input",
721
+ "when": ["present.segmentEmbeddingT", "not present.segmentIdsT", "present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
722
  "passes": [
723
  {
724
  "id": "sum",
725
  "name": "EmbedLayerNormalization.EmbeddingSum",
726
  "shader": "embed-sum.wgsl.jinja",
727
+ "bindings": ["input_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
728
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
729
  },
730
  {
 
738
  "id": "maskIndex",
739
  "name": "EmbedLayerNormalization.ZeroMaskIndex",
740
  "shader": "embed-mask-index.wgsl.jinja",
741
+ "bindings": ["mask_index", "params_batch"],
742
  "dispatch": {
743
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
744
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
748
  ]
749
  },
750
  {
751
+ "id": "seg_nopos_nosum_nomask",
752
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "not present.positionIdsT", "not present.embeddingSumT", "not present.maskIndexT"],
753
  "passes": [
754
  {
755
  "id": "sum",
756
  "name": "EmbedLayerNormalization.EmbeddingSum",
757
  "shader": "embed-sum.wgsl.jinja",
758
+ "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
759
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
760
  },
761
  {
 
764
  "shader": "embed-normalize.wgsl.jinja",
765
  "bindings": ["output", "gamma", "beta", "params"],
766
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
767
  }
768
  ]
769
  },
770
  {
771
+ "id": "seg_nopos_nosum_mask",
772
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "not present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "present.maskT"],
773
  "passes": [
774
  {
775
  "id": "sum",
776
  "name": "EmbedLayerNormalization.EmbeddingSum",
777
  "shader": "embed-sum.wgsl.jinja",
778
+ "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
779
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
780
  },
781
  {
 
787
  },
788
  {
789
  "id": "maskIndex",
790
+ "name": "EmbedLayerNormalization.MaskIndex",
791
  "shader": "embed-mask-index.wgsl.jinja",
792
+ "bindings": ["mask", "mask_index", "params_batch"],
793
  "dispatch": {
794
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
795
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
799
  ]
800
  },
801
  {
802
+ "id": "seg_nopos_nosum_mask_without_input",
803
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "not present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
804
  "passes": [
805
  {
806
  "id": "sum",
807
  "name": "EmbedLayerNormalization.EmbeddingSum",
808
  "shader": "embed-sum.wgsl.jinja",
809
+ "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
810
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
811
  },
812
  {
 
820
  "id": "maskIndex",
821
  "name": "EmbedLayerNormalization.ZeroMaskIndex",
822
  "shader": "embed-mask-index.wgsl.jinja",
823
+ "bindings": ["mask_index", "params_batch"],
824
  "dispatch": {
825
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
826
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
830
  ]
831
  },
832
  {
833
+ "id": "seg_nopos_sum_nomask",
834
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "not present.positionIdsT", "present.embeddingSumT", "not present.maskIndexT"],
835
  "passes": [
836
  {
837
  "id": "sum",
838
  "name": "EmbedLayerNormalization.EmbeddingSum",
839
  "shader": "embed-sum.wgsl.jinja",
840
+ "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
841
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
842
  },
843
  {
 
846
  "shader": "embed-normalize.wgsl.jinja",
847
  "bindings": ["output", "gamma", "beta", "params"],
848
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
849
  }
850
  ]
851
  },
852
  {
853
+ "id": "seg_nopos_sum_mask",
854
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "not present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "present.maskT"],
855
  "passes": [
856
  {
857
  "id": "sum",
858
  "name": "EmbedLayerNormalization.EmbeddingSum",
859
  "shader": "embed-sum.wgsl.jinja",
860
+ "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
861
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
862
  },
863
  {
 
869
  },
870
  {
871
  "id": "maskIndex",
872
+ "name": "EmbedLayerNormalization.MaskIndex",
873
  "shader": "embed-mask-index.wgsl.jinja",
874
+ "bindings": ["mask", "mask_index", "params_batch"],
875
  "dispatch": {
876
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
877
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
881
  ]
882
  },
883
  {
884
+ "id": "seg_nopos_sum_mask_without_input",
885
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "not present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
886
  "passes": [
887
  {
888
  "id": "sum",
889
  "name": "EmbedLayerNormalization.EmbeddingSum",
890
  "shader": "embed-sum.wgsl.jinja",
891
+ "bindings": ["input_ids", "segment_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
892
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
893
  },
894
  {
 
902
  "id": "maskIndex",
903
  "name": "EmbedLayerNormalization.ZeroMaskIndex",
904
  "shader": "embed-mask-index.wgsl.jinja",
905
+ "bindings": ["mask_index", "params_batch"],
906
  "dispatch": {
907
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
908
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
912
  ]
913
  },
914
  {
915
+ "id": "seg_posids_nosum_nomask",
916
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "present.positionIdsT", "not present.embeddingSumT", "not present.maskIndexT"],
917
  "passes": [
918
  {
919
  "id": "sum",
920
  "name": "EmbedLayerNormalization.EmbeddingSum",
921
  "shader": "embed-sum.wgsl.jinja",
922
+ "bindings": ["input_ids", "segment_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
923
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
924
  },
925
  {
 
928
  "shader": "embed-normalize.wgsl.jinja",
929
  "bindings": ["output", "gamma", "beta", "params"],
930
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
931
  }
932
  ]
933
  },
934
  {
935
+ "id": "seg_posids_nosum_mask",
936
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "present.maskT"],
937
  "passes": [
938
  {
939
  "id": "sum",
940
  "name": "EmbedLayerNormalization.EmbeddingSum",
941
  "shader": "embed-sum.wgsl.jinja",
942
+ "bindings": ["input_ids", "segment_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
943
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
944
  },
945
  {
 
951
  },
952
  {
953
  "id": "maskIndex",
954
+ "name": "EmbedLayerNormalization.MaskIndex",
955
  "shader": "embed-mask-index.wgsl.jinja",
956
+ "bindings": ["mask", "mask_index", "params_batch"],
957
  "dispatch": {
958
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
959
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
963
  ]
964
  },
965
  {
966
+ "id": "seg_posids_nosum_mask_without_input",
967
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "present.positionIdsT", "not present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
968
  "passes": [
969
  {
970
  "id": "sum",
971
  "name": "EmbedLayerNormalization.EmbeddingSum",
972
  "shader": "embed-sum.wgsl.jinja",
973
+ "bindings": ["input_ids", "segment_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "params"],
974
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
975
  },
976
  {
 
984
  "id": "maskIndex",
985
  "name": "EmbedLayerNormalization.ZeroMaskIndex",
986
  "shader": "embed-mask-index.wgsl.jinja",
987
+ "bindings": ["mask_index", "params_batch"],
988
  "dispatch": {
989
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
990
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
994
  ]
995
  },
996
  {
997
+ "id": "seg_posids_sum_nomask",
998
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "present.positionIdsT", "present.embeddingSumT", "not present.maskIndexT"],
999
  "passes": [
1000
  {
1001
  "id": "sum",
1002
  "name": "EmbedLayerNormalization.EmbeddingSum",
1003
  "shader": "embed-sum.wgsl.jinja",
1004
+ "bindings": ["input_ids", "segment_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
1005
+ "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
1006
+ },
1007
+ {
1008
+ "id": "normalize",
1009
+ "name": "EmbedLayerNormalization.Normalize",
1010
+ "shader": "embed-normalize.wgsl.jinja",
1011
+ "bindings": ["output", "gamma", "beta", "params"],
1012
+ "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
1013
+ }
1014
+ ]
1015
+ },
1016
+ {
1017
+ "id": "seg_posids_sum_mask",
1018
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "present.maskT"],
1019
+ "passes": [
1020
+ {
1021
+ "id": "sum",
1022
+ "name": "EmbedLayerNormalization.EmbeddingSum",
1023
+ "shader": "embed-sum.wgsl.jinja",
1024
+ "bindings": ["input_ids", "segment_ids", "position_ids", "word_embedding", "position_embedding", "segment_embedding", "output", "embedding_sum", "params"],
1025
  "dispatch": { "x": "min(tokens, 65535)", "y": "ceilDiv(tokens, 65535)", "z": 1 }
1026
  },
1027
  {
 
1033
  },
1034
  {
1035
  "id": "maskIndex",
1036
+ "name": "EmbedLayerNormalization.MaskIndex",
1037
  "shader": "embed-mask-index.wgsl.jinja",
1038
+ "bindings": ["mask", "mask_index", "params_batch"],
1039
  "dispatch": {
1040
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
1041
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
 
1046
  },
1047
  {
1048
  "id": "seg_posids_sum_mask_without_input",
1049
+ "when": ["present.segmentEmbeddingT", "present.segmentIdsT", "present.positionIdsT", "present.embeddingSumT", "present.maskIndexT", "not present.maskT"],
1050
  "passes": [
1051
  {
1052
  "id": "sum",
 
1066
  "id": "maskIndex",
1067
  "name": "EmbedLayerNormalization.ZeroMaskIndex",
1068
  "shader": "embed-mask-index.wgsl.jinja",
1069
+ "bindings": ["mask_index", "params_batch"],
1070
  "dispatch": {
1071
  "x": "min(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
1072
  "y": "ceilDiv(ceilDiv((batchSize), (tunables.MASK_WORKGROUP_SIZE)), 65535)",
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "com.microsoft.EmbedLayerNormalization",
3
- "id": "_com_microsoft_embedlayernormalization_webgpu_1487059",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,52 +8,52 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "Hs0qFEnXlwW07Hr/0rT+Bq7eqznmLpk7Zf1zpHE3qvw=",
11
- "embed-mask-index.wgsl.jinja": "7wT7/LzdrOnanlTU5kmneEYpGb6D8v0H9/9mra8eF5Q=",
12
- "embed-normalize.wgsl.jinja": "U2vsL1A7Y88g0DgeHpCON1vrTL7Dzc9Z2+CUmAFtt+0=",
13
  "embed-sum.wgsl.jinja": "RJVXdkHRkytYbkJ8wtcuZpB7TG/ubEOZPmIOYpOXMCw=",
14
- "manifest.json": "ocqCMw8JDV5GTOF6rXRhscDQ+z2G9NjzZDfxlwAznP8=",
15
- "test.json": "D+ZuT3RBcLxlldTtFpYDKm1J2fRt64z2CvSv0vP8wbk="
16
  }
17
  },
18
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
19
  "webgpu": {
20
- "manifestSpec": "2.0",
21
  "variants": {
22
  "noseg_nopos_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
23
  "noseg_nopos_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
24
  "noseg_nopos_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
25
  "noseg_nopos_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
26
  "noseg_posids_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
27
  "noseg_posids_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
28
  "noseg_posids_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
29
  "noseg_posids_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
30
- "seg_nopos_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
31
- "seg_nopos_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
32
- "seg_nopos_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
33
- "seg_nopos_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
34
- "seg_posids_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
35
- "seg_posids_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
36
- "seg_posids_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
37
- "seg_posids_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
38
  "segdefault_nopos_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
39
  "segdefault_nopos_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
40
  "segdefault_nopos_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
41
  "segdefault_nopos_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
42
  "segdefault_posids_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
43
  "segdefault_posids_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
44
  "segdefault_posids_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
45
  "segdefault_posids_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
46
- "noseg_nopos_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
47
- "noseg_nopos_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
48
- "noseg_posids_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
49
- "noseg_posids_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
50
- "segdefault_nopos_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
51
- "segdefault_nopos_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
52
- "segdefault_posids_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
53
  "segdefault_posids_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
 
54
  "seg_nopos_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
 
55
  "seg_nopos_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
 
56
  "seg_posids_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
 
57
  "seg_posids_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"]
58
  }
59
  }
 
1
  {
2
  "name": "com.microsoft.EmbedLayerNormalization",
3
+ "id": "_com_microsoft_embedlayernormalization_webgpu_ce87271",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "Hs0qFEnXlwW07Hr/0rT+Bq7eqznmLpk7Zf1zpHE3qvw=",
11
+ "embed-mask-index.wgsl.jinja": "FFU/TdSLERpF+HHiz3CiNvomfBpvttIIrWk7iilBdIk=",
12
+ "embed-normalize.wgsl.jinja": "3lFGMIL5hwRMv/t+OhHeG7jq7ufLZesM+gDbHRnSTM8=",
13
  "embed-sum.wgsl.jinja": "RJVXdkHRkytYbkJ8wtcuZpB7TG/ubEOZPmIOYpOXMCw=",
14
+ "manifest.json": "9IEMWYQvhNZNM8OKUE3QgPlqUcXTx7wRUwh+CaWjhCs=",
15
+ "test.json": "BqZ4jf6SD9cfRVwpKREW0nEui3Y1aZuVGNEw31hPbrs="
16
  }
17
  },
18
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
19
  "webgpu": {
20
+ "manifestSpec": "2.1",
21
  "variants": {
22
  "noseg_nopos_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
23
  "noseg_nopos_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
24
+ "noseg_nopos_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
25
  "noseg_nopos_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
26
  "noseg_nopos_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
27
+ "noseg_nopos_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
28
  "noseg_posids_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
29
  "noseg_posids_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
30
+ "noseg_posids_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
31
  "noseg_posids_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
32
  "noseg_posids_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
33
+ "noseg_posids_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
 
 
 
 
 
 
34
  "segdefault_nopos_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
35
  "segdefault_nopos_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
36
+ "segdefault_nopos_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
37
  "segdefault_nopos_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
38
  "segdefault_nopos_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
39
+ "segdefault_nopos_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
40
  "segdefault_posids_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
41
  "segdefault_posids_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
42
+ "segdefault_posids_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
43
  "segdefault_posids_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
44
  "segdefault_posids_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
 
 
 
 
 
 
 
45
  "segdefault_posids_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
46
+ "seg_nopos_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
47
+ "seg_nopos_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
48
  "seg_nopos_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
49
+ "seg_nopos_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
50
+ "seg_nopos_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
51
  "seg_nopos_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
52
+ "seg_posids_nosum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
53
+ "seg_posids_nosum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
54
  "seg_posids_nosum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
55
+ "seg_posids_sum_nomask": ["embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
56
+ "seg_posids_sum_mask": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"],
57
  "seg_posids_sum_mask_without_input": ["embed-mask-index.wgsl.jinja", "embed-normalize.wgsl.jinja", "embed-sum.wgsl.jinja"]
58
  }
59
  }
build/webgpu/test.json CHANGED
@@ -891,7 +891,9 @@
891
  },
892
  {
893
  "name": "noseg_nopos_sum_mask_index_omitted",
894
- "provenance": { "notes": "Compact optional-output regression: embedding_sum does not imply mask_index." },
 
 
895
  "inputs": {
896
  "inputIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
897
  "wordEmbeddingT": {
@@ -915,7 +917,7 @@
915
  {
916
  "name": "noseg_posids_sum_mask_index_omitted",
917
  "provenance": {
918
- "notes": "Compact optional-output regression: requesting embedding_sum does not require mask_index."
919
  },
920
  "inputs": {
921
  "inputIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
@@ -940,7 +942,9 @@
940
  },
941
  {
942
  "name": "seg_nopos_nosum_mask_index_omitted",
943
- "provenance": { "notes": "Compact optional-output regression: segment embeddings do not require mask_index." },
 
 
944
  "inputs": {
945
  "inputIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
946
  "segmentIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
@@ -967,7 +971,7 @@
967
  {
968
  "name": "seg_posids_nosum_mask_index_omitted",
969
  "provenance": {
970
- "notes": "Compact optional-output regression: segment and position ids do not require mask_index."
971
  },
972
  "inputs": {
973
  "inputIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
@@ -996,7 +1000,7 @@
996
  {
997
  "name": "seg_posids_sum_mask_index_omitted",
998
  "provenance": {
999
- "notes": "Compact optional-output regression: embedding_sum remains independently requestable with the full embedding input set."
1000
  },
1001
  "inputs": {
1002
  "inputIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
 
891
  },
892
  {
893
  "name": "noseg_nopos_sum_mask_index_omitted",
894
+ "provenance": {
895
+ "notes": "With no segment or explicit position inputs, requesting the embedding_sum output while omitting mask_index checks that embedding_sum doesn't require mask_index to also be produced."
896
+ },
897
  "inputs": {
898
  "inputIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
899
  "wordEmbeddingT": {
 
917
  {
918
  "name": "noseg_posids_sum_mask_index_omitted",
919
  "provenance": {
920
+ "notes": "With explicit position ids but no segment input, requesting the embedding_sum output while omitting mask_index checks that embedding_sum doesn't require mask_index to also be produced."
921
  },
922
  "inputs": {
923
  "inputIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
 
942
  },
943
  {
944
  "name": "seg_nopos_nosum_mask_index_omitted",
945
+ "provenance": {
946
+ "notes": "With segment ids but no explicit position ids and no embedding_sum output requested, omitting mask_index checks that segment embeddings alone don't require mask_index either."
947
+ },
948
  "inputs": {
949
  "inputIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
950
  "segmentIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
 
971
  {
972
  "name": "seg_posids_nosum_mask_index_omitted",
973
  "provenance": {
974
+ "notes": "With both segment and explicit position ids supplied but no embedding_sum output requested, omitting mask_index checks that neither optional input requires mask_index to be produced."
975
  },
976
  "inputs": {
977
  "inputIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },
 
1000
  {
1001
  "name": "seg_posids_sum_mask_index_omitted",
1002
  "provenance": {
1003
+ "notes": "With segment ids, explicit position ids, and the embedding_sum output all supplied together, omitting mask_index checks that this full combination of optional inputs and outputs still doesn't require it."
1004
  },
1005
  "inputs": {
1006
  "inputIdsT": { "dtype": "int32", "shape": [1, 2], "data": { "kind": "values", "values": [0, 1] } },