Xenova HF Staff commited on
Commit
5ec57d5
·
verified ·
1 Parent(s): f1a8138

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -30,6 +30,8 @@ See the [ONNX Runtime `SkipSimplifiedLayerNormalization` contrib-operator spec](
30
  | Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
31
  | --- | --- | --- | --- | --- | --- | --- |
32
  | `outputT` | `output` | `T` | same as `inputT` | same as `inputT` | Normalized output tensor with the same shape as `input`. | required |
 
 
33
  | `residualT` | `input_skip_bias_sum` | `T` | same as `inputT` | same as `inputT` | Sum of `input`, `skip`, and optional `bias` before normalization, with the same shape as `input`. | optional |
34
 
35
  ## Attributes
@@ -45,6 +47,24 @@ Default values (overridable per request):
45
  | Variable | Allowed dtypes |
46
  | --- | --- |
47
  | `T` | `float32`, `float16` |
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
48
 
49
  ## Device requirements
50
 
@@ -55,14 +75,14 @@ Some implementation variants require `shader-f16`. These are route-specific capa
55
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
56
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
57
  - [`test.json`](build/webgpu/test.json) — correctness cases
58
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
59
  - [`norm-skip-row-vec4.wgsl.jinja`](build/webgpu/norm-skip-row-vec4.wgsl.jinja)
60
  - [`norm-skip-row.wgsl.jinja`](build/webgpu/norm-skip-row.wgsl.jinja)
61
 
62
  ## Use with `@huggingface/kernels`
63
 
64
  ```sh
65
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
66
  ```
67
 
68
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
30
  | Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
31
  | --- | --- | --- | --- | --- | --- | --- |
32
  | `outputT` | `output` | `T` | same as `inputT` | same as `inputT` | Normalized output tensor with the same shape as `input`. | required |
33
+ | `meanT` | `mean` | `U` | same as `inputT` | derived | Per-row mean; zero for simplified RMS normalization. Shape matches the input with its final axis replaced by one. | optional |
34
+ | `invStdT` | `inv_std_var` | `U` | same as `inputT` | derived | Per-row inverse standard deviation, or inverse RMS for simplified normalization. Shape matches the input with its final axis replaced by one. | optional |
35
  | `residualT` | `input_skip_bias_sum` | `T` | same as `inputT` | same as `inputT` | Sum of `input`, `skip`, and optional `bias` before normalization, with the same shape as `input`. | optional |
36
 
37
  ## Attributes
 
47
  | Variable | Allowed dtypes |
48
  | --- | --- |
49
  | `T` | `float32`, `float16` |
50
+ | `U` | `float32` |
51
+
52
+ ## Implementation variants
53
+
54
+ One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
55
+
56
+ - `stats_mean_plain` — Row normalization returning mean statistics with `plain` optional inputs and outputs.
57
+ - `stats_mean_residual` — Row normalization returning mean statistics with `residual` optional inputs and outputs.
58
+ - `stats_mean_bias` — Row normalization returning mean statistics with `bias` optional inputs and outputs.
59
+ - `stats_mean_bias_residual` — Row normalization returning mean statistics with `bias_residual` optional inputs and outputs.
60
+ - `stats_inv_plain` — Row normalization returning inv statistics with `plain` optional inputs and outputs.
61
+ - `stats_inv_residual` — Row normalization returning inv statistics with `residual` optional inputs and outputs.
62
+ - `stats_inv_bias` — Row normalization returning inv statistics with `bias` optional inputs and outputs.
63
+ - `stats_inv_bias_residual` — Row normalization returning inv statistics with `bias_residual` optional inputs and outputs.
64
+ - `stats_both_plain` — Row normalization returning both statistics with `plain` optional inputs and outputs.
65
+ - `stats_both_residual` — Row normalization returning both statistics with `residual` optional inputs and outputs.
66
+ - `stats_both_bias` — Row normalization returning both statistics with `bias` optional inputs and outputs.
67
+ - `stats_both_bias_residual` — Row normalization returning both statistics with `bias_residual` optional inputs and outputs.
68
 
69
  ## Device requirements
70
 
 
75
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
76
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
77
  - [`test.json`](build/webgpu/test.json) — correctness cases
78
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
79
  - [`norm-skip-row-vec4.wgsl.jinja`](build/webgpu/norm-skip-row-vec4.wgsl.jinja)
80
  - [`norm-skip-row.wgsl.jinja`](build/webgpu/norm-skip-row.wgsl.jinja)
81
 
82
  ## Use with `@huggingface/kernels`
83
 
84
  ```sh
85
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
86
  ```
87
 
88
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/manifest.json CHANGED
@@ -10,6 +10,20 @@
10
  },
11
  "outputs": {
12
  "outputT": { "onnx": "output", "dtype": "T", "rank": "ranks.inputT", "shape": "shapes.inputT" },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
13
  "residualT": {
14
  "onnx": "input_skip_bias_sum",
15
  "dtype": "T",
@@ -19,7 +33,7 @@
19
  }
20
  },
21
  "attributes": { "epsilon": { "default": 9.999999960041972e-13 } },
22
- "typeConstraints": { "T": ["float32", "float16"] },
23
  "tunables": { "MAX_WORKGROUP_SIZE": { "default": 256 } },
24
  "derive": {
25
  "rowCount": "numel(shapes.inputT) / max(1, dim(shapes.inputT, -1))",
@@ -39,6 +53,7 @@
39
  "vec4Aligned": "dim(shapes.inputT, -1) % 4 == 0",
40
  "hasSubgroups": "device.features.has(\"subgroups\") and device.wgslLanguageFeatures.has(\"subgroup_id\")",
41
  "hasF16": "device.features.has(\"shader-f16\")",
 
42
  "no_bias_contract": "not present.biasT",
43
  "f32_bias_contract": "false if not present.biasT else (ranks.biasT == 1 and tensorDtypes.biasT == \"float32\" and dim(shapes.biasT, 0) == hiddenSize)",
44
  "f16_bias_contract": "false if not present.biasT else (ranks.biasT == 1 and tensorDtypes.biasT == \"float16\" and dim(shapes.biasT, 0) == hiddenSize)",
@@ -49,17 +64,17 @@
49
  "f32_no_bias_output_contract": "coreContract and outputOnlyContract and f32MainDtypes and no_bias_contract",
50
  "f32_bias_output_contract": "coreContract and outputOnlyContract and f32MainDtypes and f32_bias_contract",
51
  "f16_no_bias_output_contract": "hasF16 and coreContract and outputOnlyContract and f16MainDtypes and no_bias_contract",
52
- "f16_bias_output_contract": "hasF16 and coreContract and outputOnlyContract and f16MainDtypes and f16_bias_contract"
 
53
  },
54
  "when": ["normResourcesFit", "rowDispatchFits"],
55
  "bindings": {
56
- "input": { "arg": "inputT", "buffer": "read-only-storage", "elementType": "$vectorScalar" },
57
- "skip": { "arg": "skipT", "buffer": "read-only-storage", "elementType": "$vectorScalar" },
58
- "gamma": { "arg": "gammaT", "buffer": "read-only-storage", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" },
59
- "output": { "arg": "outputT", "buffer": "storage", "elementType": "$vectorScalar" },
60
- "input_skip_bias_sum": { "arg": "residualT", "buffer": "storage", "elementType": "$vectorScalar" },
61
  "params": {
62
- "buffer": "uniform",
63
  "struct": [
64
  { "name": "rows", "type": "u32", "value": "rowCount" },
65
  {
@@ -70,61 +85,37 @@
70
  { "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
71
  ]
72
  },
73
- "bias": { "arg": "biasT", "buffer": "read-only-storage", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" },
74
- "input_2": { "arg": "inputT", "name": "input", "buffer": "read-only-storage", "elementType": "$scalar" },
75
- "skip_2": { "arg": "skipT", "name": "skip", "buffer": "read-only-storage", "elementType": "$scalar" },
76
- "gamma_2": {
77
- "arg": "gammaT",
78
- "name": "gamma",
79
- "buffer": "read-only-storage",
80
- "elementType": "$scalar",
81
- "length": "$HIDDEN_LEN"
82
- },
83
- "output_2": { "arg": "outputT", "name": "output", "buffer": "storage", "elementType": "$scalar" },
84
- "input_skip_bias_sum_2": {
85
- "arg": "residualT",
86
- "name": "input_skip_bias_sum",
87
- "buffer": "storage",
88
- "elementType": "$scalar"
89
- },
90
- "params_2": {
91
  "name": "params",
92
- "buffer": "uniform",
93
  "struct": [
94
  { "name": "rows", "type": "u32", "value": "rowCount" },
95
  { "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
96
  ]
97
  },
98
- "bias_2": {
99
- "arg": "biasT",
100
- "name": "bias",
101
- "buffer": "read-only-storage",
102
- "elementType": "$scalar",
103
- "length": "$HIDDEN_LEN"
104
- }
105
  },
106
  "variants": [
107
  {
108
  "id": "no_bias_vec4_f16",
109
  "priority": 21,
110
- "when": ["f16_no_bias_residual_contract", "vec4Aligned"],
111
- "derive": {
112
- "scalar": "\"f16\"",
113
- "vectorScalar": "\"vec4<f16>\"",
114
- "hasBias": "\"no_bias\" == \"bias\"",
115
- "HIDDEN_LEN": "hiddenSize / 4"
116
- },
117
  "passes": [
118
  {
119
  "id": "main",
120
  "name": "SkipSimplifiedLayerNormalization.Vec4",
121
  "shader": "norm-skip-row-vec4.wgsl.jinja",
122
  "derive": {
123
- "simplified": true,
124
- "hasBias": "\"no_bias\" == \"bias\"",
125
- "hasBeta": false,
126
  "writeResidualSum": true,
127
- "usesF16Spec": true,
128
  "hidden": "hiddenSize",
129
  "hiddenVec": "hiddenSize / 4",
130
  "wg": "skipWgVec4",
@@ -132,32 +123,46 @@
132
  "useSubgroups": "hasSubgroups"
133
  },
134
  "bindings": ["input", "skip", "gamma", "output", "input_skip_bias_sum", "params"],
135
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
136
- "subgroupCollectivesWidth": "portable"
137
  }
138
  ]
139
  },
140
  {
141
- "id": "no_bias_vec4",
142
- "priority": 20,
143
- "when": ["f32_no_bias_residual_contract", "vec4Aligned"],
 
144
  "derive": {
145
- "scalar": "\"f32\"",
146
- "vectorScalar": "\"vec4<f32>\"",
147
- "hasBias": "\"no_bias\" == \"bias\"",
148
- "HIDDEN_LEN": "hiddenSize / 4"
 
 
149
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
150
  "passes": [
151
  {
152
  "id": "main",
153
  "name": "SkipSimplifiedLayerNormalization.Vec4",
154
  "shader": "norm-skip-row-vec4.wgsl.jinja",
155
  "derive": {
156
- "simplified": true,
157
- "hasBias": "\"no_bias\" == \"bias\"",
158
- "hasBeta": false,
159
  "writeResidualSum": true,
160
- "usesF16Spec": false,
161
  "hidden": "hiddenSize",
162
  "hiddenVec": "hiddenSize / 4",
163
  "wg": "skipWgVec4",
@@ -165,65 +170,45 @@
165
  "useSubgroups": "hasSubgroups"
166
  },
167
  "bindings": ["input", "skip", "gamma", "output", "input_skip_bias_sum", "params"],
168
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
169
- "subgroupCollectivesWidth": "portable"
170
  }
171
  ]
172
  },
173
  {
174
- "id": "no_bias_output_only_vec4",
175
- "priority": 20,
176
- "when": ["f32_no_bias_output_contract", "vec4Aligned"],
177
  "derive": {
 
 
 
178
  "scalar": "\"f32\"",
179
- "vectorScalar": "\"vec4<f32>\"",
180
- "hasBias": "\"no_bias\" == \"bias\"",
181
- "HIDDEN_LEN": "hiddenSize / 4"
182
  },
183
  "passes": [
184
  {
185
  "id": "main",
186
- "name": "SkipSimplifiedLayerNormalization.Vec4OutputOnly",
187
- "shader": "norm-skip-row-vec4.wgsl.jinja",
188
- "derive": {
189
- "simplified": true,
190
- "hasBias": "\"no_bias\" == \"bias\"",
191
- "hasBeta": false,
192
- "writeResidualSum": false,
193
- "usesF16Spec": false,
194
- "hidden": "hiddenSize",
195
- "hiddenVec": "hiddenSize / 4",
196
- "wg": "skipWgVec4",
197
- "vecType": "\"vec4<f32>\"",
198
- "useSubgroups": "hasSubgroups"
199
- },
200
- "bindings": ["input", "skip", "gamma", "output", "params"],
201
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
202
- "subgroupCollectivesWidth": "portable"
203
  }
204
  ]
205
  },
206
  {
207
  "id": "no_bias_output_only_vec4_f16",
208
  "priority": 21,
209
- "when": ["f16_no_bias_output_contract", "vec4Aligned"],
210
- "derive": {
211
- "scalar": "\"f16\"",
212
- "vectorScalar": "\"vec4<f16>\"",
213
- "hasBias": "\"no_bias\" == \"bias\"",
214
- "HIDDEN_LEN": "hiddenSize / 4"
215
- },
216
  "passes": [
217
  {
218
  "id": "main",
219
  "name": "SkipSimplifiedLayerNormalization.Vec4OutputOnly",
220
  "shader": "norm-skip-row-vec4.wgsl.jinja",
221
  "derive": {
222
- "simplified": true,
223
- "hasBias": "\"no_bias\" == \"bias\"",
224
- "hasBeta": false,
225
  "writeResidualSum": false,
226
- "usesF16Spec": true,
227
  "hidden": "hiddenSize",
228
  "hiddenVec": "hiddenSize / 4",
229
  "wg": "skipWgVec4",
@@ -231,83 +216,53 @@
231
  "useSubgroups": "hasSubgroups"
232
  },
233
  "bindings": ["input", "skip", "gamma", "output", "params"],
234
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
235
- "subgroupCollectivesWidth": "portable"
236
- }
237
- ]
238
- },
239
- {
240
- "id": "no_bias",
241
- "priority": 0,
242
- "when": ["f32_no_bias_residual_contract"],
243
- "derive": {
244
- "simplified": true,
245
- "useSubgroups": false,
246
- "hasBeta": false,
247
- "writeResidualSum": true,
248
- "hasBias": "\"no_bias\" == \"bias\"",
249
- "scalar": "\"f32\"",
250
- "workgroupSize": "skipWg",
251
- "HIDDEN_LEN": "hiddenSize"
252
- },
253
- "passes": [
254
- {
255
- "id": "main",
256
- "name": "SkipSimplifiedLayerNormalization",
257
- "shader": "norm-skip-row.wgsl.jinja",
258
- "bindings": ["input_2", "skip_2", "gamma_2", "output_2", "input_skip_bias_sum_2", "params_2"],
259
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
260
  }
261
  ]
262
  },
263
  {
264
- "id": "no_bias_f16",
265
  "priority": 0,
266
- "when": ["f16_no_bias_residual_contract"],
267
  "requires": { "features": ["shader-f16"] },
268
  "derive": {
269
- "simplified": true,
270
  "useSubgroups": false,
271
- "hasBeta": false,
272
- "writeResidualSum": true,
273
- "hasBias": "\"no_bias\" == \"bias\"",
274
  "scalar": "\"f16\"",
275
- "usesF16": true,
276
  "workgroupSize": "skipWg",
277
  "HIDDEN_LEN": "hiddenSize"
278
  },
279
  "passes": [
280
  {
281
  "id": "main",
282
- "name": "SkipSimplifiedLayerNormalization",
283
  "shader": "norm-skip-row.wgsl.jinja",
284
- "bindings": ["input_2", "skip_2", "gamma_2", "output_2", "input_skip_bias_sum_2", "params_2"],
285
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
286
  }
287
  ]
288
  },
289
  {
290
- "id": "no_bias_output_only_f16",
291
- "priority": 0,
292
- "when": ["f16_no_bias_output_contract"],
293
- "requires": { "features": ["shader-f16"] },
294
- "derive": {
295
- "simplified": true,
296
- "useSubgroups": false,
297
- "hasBeta": false,
298
- "writeResidualSum": false,
299
- "hasBias": "\"no_bias\" == \"bias\"",
300
- "scalar": "\"f16\"",
301
- "usesF16": true,
302
- "workgroupSize": "skipWg",
303
- "HIDDEN_LEN": "hiddenSize"
304
- },
305
  "passes": [
306
  {
307
  "id": "main",
308
- "name": "SkipSimplifiedLayerNormalization.OutputOnly",
309
- "shader": "norm-skip-row.wgsl.jinja",
310
- "bindings": ["input_2", "skip_2", "gamma_2", "output_2", "params_2"],
 
 
 
 
 
 
 
 
 
311
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
312
  }
313
  ]
@@ -315,13 +270,11 @@
315
  {
316
  "id": "no_bias_output_only",
317
  "priority": 0,
318
- "when": ["f32_no_bias_output_contract"],
319
  "derive": {
320
- "simplified": true,
321
  "useSubgroups": false,
322
- "hasBeta": false,
323
  "writeResidualSum": false,
324
- "hasBias": "\"no_bias\" == \"bias\"",
325
  "scalar": "\"f32\"",
326
  "workgroupSize": "skipWg",
327
  "HIDDEN_LEN": "hiddenSize"
@@ -331,7 +284,7 @@
331
  "id": "main",
332
  "name": "SkipSimplifiedLayerNormalization.OutputOnly",
333
  "shader": "norm-skip-row.wgsl.jinja",
334
- "bindings": ["input_2", "skip_2", "gamma_2", "output_2", "params_2"],
335
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
336
  }
337
  ]
@@ -339,24 +292,16 @@
339
  {
340
  "id": "bias_vec4_f16",
341
  "priority": 21,
342
- "when": ["f16_bias_residual_contract", "vec4Aligned"],
343
- "derive": {
344
- "scalar": "\"f16\"",
345
- "vectorScalar": "\"vec4<f16>\"",
346
- "hasBias": "\"bias\" == \"bias\"",
347
- "HIDDEN_LEN": "hiddenSize / 4"
348
- },
349
  "passes": [
350
  {
351
  "id": "main",
352
  "name": "SkipSimplifiedLayerNormalization.Vec4",
353
  "shader": "norm-skip-row-vec4.wgsl.jinja",
354
  "derive": {
355
- "simplified": true,
356
- "hasBias": "\"bias\" == \"bias\"",
357
- "hasBeta": false,
358
  "writeResidualSum": true,
359
- "usesF16Spec": true,
360
  "hidden": "hiddenSize",
361
  "hiddenVec": "hiddenSize / 4",
362
  "wg": "skipWgVec4",
@@ -364,32 +309,46 @@
364
  "useSubgroups": "hasSubgroups"
365
  },
366
  "bindings": ["input", "skip", "gamma", "bias", "output", "input_skip_bias_sum", "params"],
367
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
368
- "subgroupCollectivesWidth": "portable"
369
  }
370
  ]
371
  },
372
  {
373
- "id": "bias_vec4",
374
- "priority": 20,
375
- "when": ["f32_bias_residual_contract", "vec4Aligned"],
 
376
  "derive": {
377
- "scalar": "\"f32\"",
378
- "vectorScalar": "\"vec4<f32>\"",
379
- "hasBias": "\"bias\" == \"bias\"",
380
- "HIDDEN_LEN": "hiddenSize / 4"
 
 
381
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
382
  "passes": [
383
  {
384
  "id": "main",
385
  "name": "SkipSimplifiedLayerNormalization.Vec4",
386
  "shader": "norm-skip-row-vec4.wgsl.jinja",
387
  "derive": {
388
- "simplified": true,
389
- "hasBias": "\"bias\" == \"bias\"",
390
- "hasBeta": false,
391
  "writeResidualSum": true,
392
- "usesF16Spec": false,
393
  "hidden": "hiddenSize",
394
  "hiddenVec": "hiddenSize / 4",
395
  "wg": "skipWgVec4",
@@ -397,87 +356,111 @@
397
  "useSubgroups": "hasSubgroups"
398
  },
399
  "bindings": ["input", "skip", "gamma", "bias", "output", "input_skip_bias_sum", "params"],
400
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
401
- "subgroupCollectivesWidth": "portable"
402
  }
403
  ]
404
  },
405
  {
406
- "id": "bias_output_only_vec4",
407
- "priority": 20,
408
- "when": ["f32_bias_output_contract", "vec4Aligned"],
409
  "derive": {
 
 
 
410
  "scalar": "\"f32\"",
411
- "vectorScalar": "\"vec4<f32>\"",
412
- "hasBias": "\"bias\" == \"bias\"",
413
- "HIDDEN_LEN": "hiddenSize / 4"
414
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
415
  "passes": [
416
  {
417
  "id": "main",
418
  "name": "SkipSimplifiedLayerNormalization.Vec4OutputOnly",
419
  "shader": "norm-skip-row-vec4.wgsl.jinja",
420
  "derive": {
421
- "simplified": true,
422
- "hasBias": "\"bias\" == \"bias\"",
423
- "hasBeta": false,
424
  "writeResidualSum": false,
425
- "usesF16Spec": false,
426
  "hidden": "hiddenSize",
427
  "hiddenVec": "hiddenSize / 4",
428
  "wg": "skipWgVec4",
429
- "vecType": "\"vec4<f32>\"",
430
  "useSubgroups": "hasSubgroups"
431
  },
432
  "bindings": ["input", "skip", "gamma", "bias", "output", "params"],
433
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
434
- "subgroupCollectivesWidth": "portable"
435
  }
436
  ]
437
  },
438
  {
439
- "id": "bias_output_only_vec4_f16",
440
- "priority": 21,
441
- "when": ["f16_bias_output_contract", "vec4Aligned"],
 
442
  "derive": {
 
 
 
443
  "scalar": "\"f16\"",
444
- "vectorScalar": "\"vec4<f16>\"",
445
- "hasBias": "\"bias\" == \"bias\"",
446
- "HIDDEN_LEN": "hiddenSize / 4"
447
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
448
  "passes": [
449
  {
450
  "id": "main",
451
  "name": "SkipSimplifiedLayerNormalization.Vec4OutputOnly",
452
  "shader": "norm-skip-row-vec4.wgsl.jinja",
453
  "derive": {
454
- "simplified": true,
455
- "hasBias": "\"bias\" == \"bias\"",
456
- "hasBeta": false,
457
  "writeResidualSum": false,
458
- "usesF16Spec": true,
459
  "hidden": "hiddenSize",
460
  "hiddenVec": "hiddenSize / 4",
461
  "wg": "skipWgVec4",
462
- "vecType": "\"vec4<f16>\"",
463
  "useSubgroups": "hasSubgroups"
464
  },
465
  "bindings": ["input", "skip", "gamma", "bias", "output", "params"],
466
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
467
- "subgroupCollectivesWidth": "portable"
468
  }
469
  ]
470
  },
471
  {
472
- "id": "bias",
473
  "priority": 0,
474
- "when": ["f32_bias_residual_contract"],
475
  "derive": {
476
- "simplified": true,
477
  "useSubgroups": false,
478
- "hasBeta": false,
479
- "writeResidualSum": true,
480
- "hasBias": "\"bias\" == \"bias\"",
481
  "scalar": "\"f32\"",
482
  "workgroupSize": "skipWg",
483
  "HIDDEN_LEN": "hiddenSize"
@@ -485,85 +468,273 @@
485
  "passes": [
486
  {
487
  "id": "main",
488
- "name": "SkipSimplifiedLayerNormalization",
489
  "shader": "norm-skip-row.wgsl.jinja",
490
- "bindings": ["input_2", "skip_2", "gamma_2", "bias_2", "output_2", "input_skip_bias_sum_2", "params_2"],
491
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
492
  }
493
  ]
494
  },
495
  {
496
- "id": "bias_f16",
497
- "priority": 0,
498
- "when": ["f16_bias_residual_contract"],
499
- "requires": { "features": ["shader-f16"] },
500
  "derive": {
501
- "simplified": true,
 
 
 
 
502
  "useSubgroups": false,
503
- "hasBeta": false,
504
- "writeResidualSum": true,
505
- "hasBias": "\"bias\" == \"bias\"",
506
- "scalar": "\"f16\"",
507
- "usesF16": true,
508
  "workgroupSize": "skipWg",
509
  "HIDDEN_LEN": "hiddenSize"
510
  },
511
  "passes": [
512
  {
513
  "id": "main",
514
- "name": "SkipSimplifiedLayerNormalization",
515
  "shader": "norm-skip-row.wgsl.jinja",
516
- "bindings": ["input_2", "skip_2", "gamma_2", "bias_2", "output_2", "input_skip_bias_sum_2", "params_2"],
517
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
518
  }
519
  ]
520
  },
521
  {
522
- "id": "bias_output_only_f16",
523
- "priority": 0,
524
- "when": ["f16_bias_output_contract"],
525
- "requires": { "features": ["shader-f16"] },
526
  "derive": {
527
- "simplified": true,
 
 
 
 
528
  "useSubgroups": false,
529
- "hasBeta": false,
530
- "writeResidualSum": false,
531
- "hasBias": "\"bias\" == \"bias\"",
532
- "scalar": "\"f16\"",
533
- "usesF16": true,
534
  "workgroupSize": "skipWg",
535
  "HIDDEN_LEN": "hiddenSize"
536
  },
537
  "passes": [
538
  {
539
  "id": "main",
540
- "name": "SkipSimplifiedLayerNormalization.OutputOnly",
541
  "shader": "norm-skip-row.wgsl.jinja",
542
- "bindings": ["input_2", "skip_2", "gamma_2", "bias_2", "output_2", "params_2"],
543
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
544
  }
545
  ]
546
  },
547
  {
548
- "id": "bias_output_only",
549
- "priority": 0,
550
- "when": ["f32_bias_output_contract"],
551
  "derive": {
552
- "simplified": true,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
553
  "useSubgroups": false,
554
- "hasBeta": false,
555
- "writeResidualSum": false,
556
- "hasBias": "\"bias\" == \"bias\"",
557
- "scalar": "\"f32\"",
558
  "workgroupSize": "skipWg",
559
  "HIDDEN_LEN": "hiddenSize"
560
  },
561
  "passes": [
562
  {
563
  "id": "main",
564
- "name": "SkipSimplifiedLayerNormalization.OutputOnly",
565
  "shader": "norm-skip-row.wgsl.jinja",
566
- "bindings": ["input_2", "skip_2", "gamma_2", "bias_2", "output_2", "params_2"],
567
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
568
  }
569
  ]
 
10
  },
11
  "outputs": {
12
  "outputT": { "onnx": "output", "dtype": "T", "rank": "ranks.inputT", "shape": "shapes.inputT" },
13
+ "meanT": {
14
+ "onnx": "mean",
15
+ "dtype": "U",
16
+ "optional": true,
17
+ "rank": "ranks.inputT",
18
+ "shape": "prefix(shapes.inputT, ranks.inputT - 1) + [1]"
19
+ },
20
+ "invStdT": {
21
+ "onnx": "inv_std_var",
22
+ "dtype": "U",
23
+ "optional": true,
24
+ "rank": "ranks.inputT",
25
+ "shape": "prefix(shapes.inputT, ranks.inputT - 1) + [1]"
26
+ },
27
  "residualT": {
28
  "onnx": "input_skip_bias_sum",
29
  "dtype": "T",
 
33
  }
34
  },
35
  "attributes": { "epsilon": { "default": 9.999999960041972e-13 } },
36
+ "typeConstraints": { "T": ["float32", "float16"], "U": ["float32"] },
37
  "tunables": { "MAX_WORKGROUP_SIZE": { "default": 256 } },
38
  "derive": {
39
  "rowCount": "numel(shapes.inputT) / max(1, dim(shapes.inputT, -1))",
 
53
  "vec4Aligned": "dim(shapes.inputT, -1) % 4 == 0",
54
  "hasSubgroups": "device.features.has(\"subgroups\") and device.wgslLanguageFeatures.has(\"subgroup_id\")",
55
  "hasF16": "device.features.has(\"shader-f16\")",
56
+ "statsRequested": "present.meanT or present.invStdT",
57
  "no_bias_contract": "not present.biasT",
58
  "f32_bias_contract": "false if not present.biasT else (ranks.biasT == 1 and tensorDtypes.biasT == \"float32\" and dim(shapes.biasT, 0) == hiddenSize)",
59
  "f16_bias_contract": "false if not present.biasT else (ranks.biasT == 1 and tensorDtypes.biasT == \"float16\" and dim(shapes.biasT, 0) == hiddenSize)",
 
64
  "f32_no_bias_output_contract": "coreContract and outputOnlyContract and f32MainDtypes and no_bias_contract",
65
  "f32_bias_output_contract": "coreContract and outputOnlyContract and f32MainDtypes and f32_bias_contract",
66
  "f16_no_bias_output_contract": "hasF16 and coreContract and outputOnlyContract and f16MainDtypes and no_bias_contract",
67
+ "f16_bias_output_contract": "hasF16 and coreContract and outputOnlyContract and f16MainDtypes and f16_bias_contract",
68
+ "statsContract": "statsRequested and coreContract and f16Ok(dtypes.T) and (not present.biasT or (ranks.biasT == 1 and dim(shapes.biasT, 0) == hiddenSize)) and (not present.residualT or sameShape(shapes.residualT, shapes.inputT))"
69
  },
70
  "when": ["normResourcesFit", "rowDispatchFits"],
71
  "bindings": {
72
+ "input": { "arg": "inputT", "elementType": "$vectorScalar" },
73
+ "skip": { "arg": "skipT", "elementType": "$vectorScalar" },
74
+ "gamma": { "arg": "gammaT", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" },
75
+ "output": { "arg": "outputT", "elementType": "$vectorScalar" },
76
+ "input_skip_bias_sum": { "arg": "residualT", "elementType": "$vectorScalar" },
77
  "params": {
 
78
  "struct": [
79
  { "name": "rows", "type": "u32", "value": "rowCount" },
80
  {
 
85
  { "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
86
  ]
87
  },
88
+ "bias": { "arg": "biasT", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" },
89
+ "input_main": { "arg": "inputT", "name": "input", "elementType": "$scalar" },
90
+ "skip_main": { "arg": "skipT", "name": "skip", "elementType": "$scalar" },
91
+ "gamma_main": { "arg": "gammaT", "name": "gamma", "elementType": "$scalar", "length": "$HIDDEN_LEN" },
92
+ "output_main": { "arg": "outputT", "name": "output", "elementType": "$scalar" },
93
+ "input_skip_bias_sum_main": { "arg": "residualT", "name": "input_skip_bias_sum", "elementType": "$scalar" },
94
+ "params_main": {
 
 
 
 
 
 
 
 
 
 
 
95
  "name": "params",
 
96
  "struct": [
97
  { "name": "rows", "type": "u32", "value": "rowCount" },
98
  { "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
99
  ]
100
  },
101
+ "bias_main": { "arg": "biasT", "name": "bias", "elementType": "$scalar", "length": "$HIDDEN_LEN" },
102
+ "mean": { "arg": "meanT", "elementType": "f32" },
103
+ "inv_std_var": { "arg": "invStdT", "elementType": "f32" }
 
 
 
 
104
  },
105
  "variants": [
106
  {
107
  "id": "no_bias_vec4_f16",
108
  "priority": 21,
109
+ "when": ["f16_no_bias_residual_contract", "vec4Aligned", "not present.meanT and not present.invStdT"],
110
+ "derive": { "vectorScalar": "\"vec4<f16>\"", "HIDDEN_LEN": "hiddenSize / 4" },
 
 
 
 
 
111
  "passes": [
112
  {
113
  "id": "main",
114
  "name": "SkipSimplifiedLayerNormalization.Vec4",
115
  "shader": "norm-skip-row-vec4.wgsl.jinja",
116
  "derive": {
117
+ "hasBias": "present.biasT",
 
 
118
  "writeResidualSum": true,
 
119
  "hidden": "hiddenSize",
120
  "hiddenVec": "hiddenSize / 4",
121
  "wg": "skipWgVec4",
 
123
  "useSubgroups": "hasSubgroups"
124
  },
125
  "bindings": ["input", "skip", "gamma", "output", "input_skip_bias_sum", "params"],
126
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
127
  }
128
  ]
129
  },
130
  {
131
+ "id": "no_bias_f16",
132
+ "priority": 0,
133
+ "when": ["f16_no_bias_residual_contract", "not present.meanT and not present.invStdT"],
134
+ "requires": { "features": ["shader-f16"] },
135
  "derive": {
136
+ "useSubgroups": false,
137
+ "writeResidualSum": true,
138
+ "hasBias": "present.biasT",
139
+ "scalar": "\"f16\"",
140
+ "workgroupSize": "skipWg",
141
+ "HIDDEN_LEN": "hiddenSize"
142
  },
143
+ "passes": [
144
+ {
145
+ "id": "main",
146
+ "name": "SkipSimplifiedLayerNormalization",
147
+ "shader": "norm-skip-row.wgsl.jinja",
148
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "input_skip_bias_sum_main", "params_main"],
149
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
150
+ }
151
+ ]
152
+ },
153
+ {
154
+ "id": "no_bias_vec4",
155
+ "priority": 20,
156
+ "when": ["f32_no_bias_residual_contract", "vec4Aligned", "not present.meanT and not present.invStdT"],
157
+ "derive": { "vectorScalar": "\"vec4<f32>\"", "HIDDEN_LEN": "hiddenSize / 4" },
158
  "passes": [
159
  {
160
  "id": "main",
161
  "name": "SkipSimplifiedLayerNormalization.Vec4",
162
  "shader": "norm-skip-row-vec4.wgsl.jinja",
163
  "derive": {
164
+ "hasBias": "present.biasT",
 
 
165
  "writeResidualSum": true,
 
166
  "hidden": "hiddenSize",
167
  "hiddenVec": "hiddenSize / 4",
168
  "wg": "skipWgVec4",
 
170
  "useSubgroups": "hasSubgroups"
171
  },
172
  "bindings": ["input", "skip", "gamma", "output", "input_skip_bias_sum", "params"],
173
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
174
  }
175
  ]
176
  },
177
  {
178
+ "id": "no_bias",
179
+ "priority": 0,
180
+ "when": ["f32_no_bias_residual_contract", "not present.meanT and not present.invStdT"],
181
  "derive": {
182
+ "useSubgroups": false,
183
+ "writeResidualSum": true,
184
+ "hasBias": "present.biasT",
185
  "scalar": "\"f32\"",
186
+ "workgroupSize": "skipWg",
187
+ "HIDDEN_LEN": "hiddenSize"
 
188
  },
189
  "passes": [
190
  {
191
  "id": "main",
192
+ "name": "SkipSimplifiedLayerNormalization",
193
+ "shader": "norm-skip-row.wgsl.jinja",
194
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "input_skip_bias_sum_main", "params_main"],
195
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
 
 
 
 
 
 
 
 
 
 
 
 
196
  }
197
  ]
198
  },
199
  {
200
  "id": "no_bias_output_only_vec4_f16",
201
  "priority": 21,
202
+ "when": ["f16_no_bias_output_contract", "vec4Aligned", "not present.meanT and not present.invStdT"],
203
+ "derive": { "vectorScalar": "\"vec4<f16>\"", "HIDDEN_LEN": "hiddenSize / 4" },
 
 
 
 
 
204
  "passes": [
205
  {
206
  "id": "main",
207
  "name": "SkipSimplifiedLayerNormalization.Vec4OutputOnly",
208
  "shader": "norm-skip-row-vec4.wgsl.jinja",
209
  "derive": {
210
+ "hasBias": "present.biasT",
 
 
211
  "writeResidualSum": false,
 
212
  "hidden": "hiddenSize",
213
  "hiddenVec": "hiddenSize / 4",
214
  "wg": "skipWgVec4",
 
216
  "useSubgroups": "hasSubgroups"
217
  },
218
  "bindings": ["input", "skip", "gamma", "output", "params"],
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
219
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
220
  }
221
  ]
222
  },
223
  {
224
+ "id": "no_bias_output_only_f16",
225
  "priority": 0,
226
+ "when": ["f16_no_bias_output_contract", "not present.meanT and not present.invStdT"],
227
  "requires": { "features": ["shader-f16"] },
228
  "derive": {
 
229
  "useSubgroups": false,
230
+ "writeResidualSum": false,
231
+ "hasBias": "present.biasT",
 
232
  "scalar": "\"f16\"",
 
233
  "workgroupSize": "skipWg",
234
  "HIDDEN_LEN": "hiddenSize"
235
  },
236
  "passes": [
237
  {
238
  "id": "main",
239
+ "name": "SkipSimplifiedLayerNormalization.OutputOnly",
240
  "shader": "norm-skip-row.wgsl.jinja",
241
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "params_main"],
242
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
243
  }
244
  ]
245
  },
246
  {
247
+ "id": "no_bias_output_only_vec4",
248
+ "priority": 20,
249
+ "when": ["f32_no_bias_output_contract", "vec4Aligned", "not present.meanT and not present.invStdT"],
250
+ "derive": { "vectorScalar": "\"vec4<f32>\"", "HIDDEN_LEN": "hiddenSize / 4" },
 
 
 
 
 
 
 
 
 
 
 
251
  "passes": [
252
  {
253
  "id": "main",
254
+ "name": "SkipSimplifiedLayerNormalization.Vec4OutputOnly",
255
+ "shader": "norm-skip-row-vec4.wgsl.jinja",
256
+ "derive": {
257
+ "hasBias": "present.biasT",
258
+ "writeResidualSum": false,
259
+ "hidden": "hiddenSize",
260
+ "hiddenVec": "hiddenSize / 4",
261
+ "wg": "skipWgVec4",
262
+ "vecType": "\"vec4<f32>\"",
263
+ "useSubgroups": "hasSubgroups"
264
+ },
265
+ "bindings": ["input", "skip", "gamma", "output", "params"],
266
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
267
  }
268
  ]
 
270
  {
271
  "id": "no_bias_output_only",
272
  "priority": 0,
273
+ "when": ["f32_no_bias_output_contract", "not present.meanT and not present.invStdT"],
274
  "derive": {
 
275
  "useSubgroups": false,
 
276
  "writeResidualSum": false,
277
+ "hasBias": "present.biasT",
278
  "scalar": "\"f32\"",
279
  "workgroupSize": "skipWg",
280
  "HIDDEN_LEN": "hiddenSize"
 
284
  "id": "main",
285
  "name": "SkipSimplifiedLayerNormalization.OutputOnly",
286
  "shader": "norm-skip-row.wgsl.jinja",
287
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "params_main"],
288
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
289
  }
290
  ]
 
292
  {
293
  "id": "bias_vec4_f16",
294
  "priority": 21,
295
+ "when": ["f16_bias_residual_contract", "vec4Aligned", "not present.meanT and not present.invStdT"],
296
+ "derive": { "vectorScalar": "\"vec4<f16>\"", "HIDDEN_LEN": "hiddenSize / 4" },
 
 
 
 
 
297
  "passes": [
298
  {
299
  "id": "main",
300
  "name": "SkipSimplifiedLayerNormalization.Vec4",
301
  "shader": "norm-skip-row-vec4.wgsl.jinja",
302
  "derive": {
303
+ "hasBias": "present.biasT",
 
 
304
  "writeResidualSum": true,
 
305
  "hidden": "hiddenSize",
306
  "hiddenVec": "hiddenSize / 4",
307
  "wg": "skipWgVec4",
 
309
  "useSubgroups": "hasSubgroups"
310
  },
311
  "bindings": ["input", "skip", "gamma", "bias", "output", "input_skip_bias_sum", "params"],
312
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
313
  }
314
  ]
315
  },
316
  {
317
+ "id": "bias_f16",
318
+ "priority": 0,
319
+ "when": ["f16_bias_residual_contract", "not present.meanT and not present.invStdT"],
320
+ "requires": { "features": ["shader-f16"] },
321
  "derive": {
322
+ "useSubgroups": false,
323
+ "writeResidualSum": true,
324
+ "hasBias": "present.biasT",
325
+ "scalar": "\"f16\"",
326
+ "workgroupSize": "skipWg",
327
+ "HIDDEN_LEN": "hiddenSize"
328
  },
329
+ "passes": [
330
+ {
331
+ "id": "main",
332
+ "name": "SkipSimplifiedLayerNormalization",
333
+ "shader": "norm-skip-row.wgsl.jinja",
334
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "input_skip_bias_sum_main", "params_main"],
335
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
336
+ }
337
+ ]
338
+ },
339
+ {
340
+ "id": "bias_vec4",
341
+ "priority": 20,
342
+ "when": ["f32_bias_residual_contract", "vec4Aligned", "not present.meanT and not present.invStdT"],
343
+ "derive": { "vectorScalar": "\"vec4<f32>\"", "HIDDEN_LEN": "hiddenSize / 4" },
344
  "passes": [
345
  {
346
  "id": "main",
347
  "name": "SkipSimplifiedLayerNormalization.Vec4",
348
  "shader": "norm-skip-row-vec4.wgsl.jinja",
349
  "derive": {
350
+ "hasBias": "present.biasT",
 
 
351
  "writeResidualSum": true,
 
352
  "hidden": "hiddenSize",
353
  "hiddenVec": "hiddenSize / 4",
354
  "wg": "skipWgVec4",
 
356
  "useSubgroups": "hasSubgroups"
357
  },
358
  "bindings": ["input", "skip", "gamma", "bias", "output", "input_skip_bias_sum", "params"],
359
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
360
  }
361
  ]
362
  },
363
  {
364
+ "id": "bias",
365
+ "priority": 0,
366
+ "when": ["f32_bias_residual_contract", "not present.meanT and not present.invStdT"],
367
  "derive": {
368
+ "useSubgroups": false,
369
+ "writeResidualSum": true,
370
+ "hasBias": "present.biasT",
371
  "scalar": "\"f32\"",
372
+ "workgroupSize": "skipWg",
373
+ "HIDDEN_LEN": "hiddenSize"
 
374
  },
375
+ "passes": [
376
+ {
377
+ "id": "main",
378
+ "name": "SkipSimplifiedLayerNormalization",
379
+ "shader": "norm-skip-row.wgsl.jinja",
380
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "input_skip_bias_sum_main", "params_main"],
381
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
382
+ }
383
+ ]
384
+ },
385
+ {
386
+ "id": "bias_output_only_vec4_f16",
387
+ "priority": 21,
388
+ "when": ["f16_bias_output_contract", "vec4Aligned", "not present.meanT and not present.invStdT"],
389
+ "derive": { "vectorScalar": "\"vec4<f16>\"", "HIDDEN_LEN": "hiddenSize / 4" },
390
  "passes": [
391
  {
392
  "id": "main",
393
  "name": "SkipSimplifiedLayerNormalization.Vec4OutputOnly",
394
  "shader": "norm-skip-row-vec4.wgsl.jinja",
395
  "derive": {
396
+ "hasBias": "present.biasT",
 
 
397
  "writeResidualSum": false,
 
398
  "hidden": "hiddenSize",
399
  "hiddenVec": "hiddenSize / 4",
400
  "wg": "skipWgVec4",
401
+ "vecType": "\"vec4<f16>\"",
402
  "useSubgroups": "hasSubgroups"
403
  },
404
  "bindings": ["input", "skip", "gamma", "bias", "output", "params"],
405
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
406
  }
407
  ]
408
  },
409
  {
410
+ "id": "bias_output_only_f16",
411
+ "priority": 0,
412
+ "when": ["f16_bias_output_contract", "not present.meanT and not present.invStdT"],
413
+ "requires": { "features": ["shader-f16"] },
414
  "derive": {
415
+ "useSubgroups": false,
416
+ "writeResidualSum": false,
417
+ "hasBias": "present.biasT",
418
  "scalar": "\"f16\"",
419
+ "workgroupSize": "skipWg",
420
+ "HIDDEN_LEN": "hiddenSize"
 
421
  },
422
+ "passes": [
423
+ {
424
+ "id": "main",
425
+ "name": "SkipSimplifiedLayerNormalization.OutputOnly",
426
+ "shader": "norm-skip-row.wgsl.jinja",
427
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "params_main"],
428
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
429
+ }
430
+ ]
431
+ },
432
+ {
433
+ "id": "bias_output_only_vec4",
434
+ "priority": 20,
435
+ "when": ["f32_bias_output_contract", "vec4Aligned", "not present.meanT and not present.invStdT"],
436
+ "derive": { "vectorScalar": "\"vec4<f32>\"", "HIDDEN_LEN": "hiddenSize / 4" },
437
  "passes": [
438
  {
439
  "id": "main",
440
  "name": "SkipSimplifiedLayerNormalization.Vec4OutputOnly",
441
  "shader": "norm-skip-row-vec4.wgsl.jinja",
442
  "derive": {
443
+ "hasBias": "present.biasT",
 
 
444
  "writeResidualSum": false,
 
445
  "hidden": "hiddenSize",
446
  "hiddenVec": "hiddenSize / 4",
447
  "wg": "skipWgVec4",
448
+ "vecType": "\"vec4<f32>\"",
449
  "useSubgroups": "hasSubgroups"
450
  },
451
  "bindings": ["input", "skip", "gamma", "bias", "output", "params"],
452
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
453
  }
454
  ]
455
  },
456
  {
457
+ "id": "bias_output_only",
458
  "priority": 0,
459
+ "when": ["f32_bias_output_contract", "not present.meanT and not present.invStdT"],
460
  "derive": {
 
461
  "useSubgroups": false,
462
+ "writeResidualSum": false,
463
+ "hasBias": "present.biasT",
 
464
  "scalar": "\"f32\"",
465
  "workgroupSize": "skipWg",
466
  "HIDDEN_LEN": "hiddenSize"
 
468
  "passes": [
469
  {
470
  "id": "main",
471
+ "name": "SkipSimplifiedLayerNormalization.OutputOnly",
472
  "shader": "norm-skip-row.wgsl.jinja",
473
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "params_main"],
474
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
475
  }
476
  ]
477
  },
478
  {
479
+ "id": "stats_mean_plain",
480
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "not present.biasT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
 
 
481
  "derive": {
482
+ "hasBias": "present.biasT",
483
+ "writeResidualSum": "present.residualT",
484
+ "writeMean": "present.meanT",
485
+ "writeInvStd": "present.invStdT",
486
+ "scalar": "dtypes.T",
487
  "useSubgroups": false,
 
 
 
 
 
488
  "workgroupSize": "skipWg",
489
  "HIDDEN_LEN": "hiddenSize"
490
  },
491
  "passes": [
492
  {
493
  "id": "main",
 
494
  "shader": "norm-skip-row.wgsl.jinja",
495
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "mean", "params_main"],
496
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
497
  }
498
  ]
499
  },
500
  {
501
+ "id": "stats_mean_residual",
502
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "not present.biasT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
 
 
503
  "derive": {
504
+ "hasBias": "present.biasT",
505
+ "writeResidualSum": "present.residualT",
506
+ "writeMean": "present.meanT",
507
+ "writeInvStd": "present.invStdT",
508
+ "scalar": "dtypes.T",
509
  "useSubgroups": false,
 
 
 
 
 
510
  "workgroupSize": "skipWg",
511
  "HIDDEN_LEN": "hiddenSize"
512
  },
513
  "passes": [
514
  {
515
  "id": "main",
 
516
  "shader": "norm-skip-row.wgsl.jinja",
517
+ "bindings": ["input_main", "skip_main", "gamma_main", "input_skip_bias_sum_main", "output_main", "mean", "params_main"],
518
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
519
  }
520
  ]
521
  },
522
  {
523
+ "id": "stats_mean_bias",
524
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "present.biasT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
 
525
  "derive": {
526
+ "hasBias": "present.biasT",
527
+ "writeResidualSum": "present.residualT",
528
+ "writeMean": "present.meanT",
529
+ "writeInvStd": "present.invStdT",
530
+ "scalar": "dtypes.T",
531
+ "useSubgroups": false,
532
+ "workgroupSize": "skipWg",
533
+ "HIDDEN_LEN": "hiddenSize"
534
+ },
535
+ "passes": [
536
+ {
537
+ "id": "main",
538
+ "shader": "norm-skip-row.wgsl.jinja",
539
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "mean", "params_main"],
540
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
541
+ }
542
+ ]
543
+ },
544
+ {
545
+ "id": "stats_mean_bias_residual",
546
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "present.biasT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
547
+ "derive": {
548
+ "hasBias": "present.biasT",
549
+ "writeResidualSum": "present.residualT",
550
+ "writeMean": "present.meanT",
551
+ "writeInvStd": "present.invStdT",
552
+ "scalar": "dtypes.T",
553
+ "useSubgroups": false,
554
+ "workgroupSize": "skipWg",
555
+ "HIDDEN_LEN": "hiddenSize"
556
+ },
557
+ "passes": [
558
+ {
559
+ "id": "main",
560
+ "shader": "norm-skip-row.wgsl.jinja",
561
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "input_skip_bias_sum_main", "output_main", "mean", "params_main"],
562
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
563
+ }
564
+ ]
565
+ },
566
+ {
567
+ "id": "stats_inv_plain",
568
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "not present.biasT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
569
+ "derive": {
570
+ "hasBias": "present.biasT",
571
+ "writeResidualSum": "present.residualT",
572
+ "writeMean": "present.meanT",
573
+ "writeInvStd": "present.invStdT",
574
+ "scalar": "dtypes.T",
575
+ "useSubgroups": false,
576
+ "workgroupSize": "skipWg",
577
+ "HIDDEN_LEN": "hiddenSize"
578
+ },
579
+ "passes": [
580
+ {
581
+ "id": "main",
582
+ "shader": "norm-skip-row.wgsl.jinja",
583
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "inv_std_var", "params_main"],
584
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
585
+ }
586
+ ]
587
+ },
588
+ {
589
+ "id": "stats_inv_residual",
590
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "not present.biasT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
591
+ "derive": {
592
+ "hasBias": "present.biasT",
593
+ "writeResidualSum": "present.residualT",
594
+ "writeMean": "present.meanT",
595
+ "writeInvStd": "present.invStdT",
596
+ "scalar": "dtypes.T",
597
+ "useSubgroups": false,
598
+ "workgroupSize": "skipWg",
599
+ "HIDDEN_LEN": "hiddenSize"
600
+ },
601
+ "passes": [
602
+ {
603
+ "id": "main",
604
+ "shader": "norm-skip-row.wgsl.jinja",
605
+ "bindings": ["input_main", "skip_main", "gamma_main", "input_skip_bias_sum_main", "output_main", "inv_std_var", "params_main"],
606
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
607
+ }
608
+ ]
609
+ },
610
+ {
611
+ "id": "stats_inv_bias",
612
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "present.biasT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
613
+ "derive": {
614
+ "hasBias": "present.biasT",
615
+ "writeResidualSum": "present.residualT",
616
+ "writeMean": "present.meanT",
617
+ "writeInvStd": "present.invStdT",
618
+ "scalar": "dtypes.T",
619
+ "useSubgroups": false,
620
+ "workgroupSize": "skipWg",
621
+ "HIDDEN_LEN": "hiddenSize"
622
+ },
623
+ "passes": [
624
+ {
625
+ "id": "main",
626
+ "shader": "norm-skip-row.wgsl.jinja",
627
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "inv_std_var", "params_main"],
628
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
629
+ }
630
+ ]
631
+ },
632
+ {
633
+ "id": "stats_inv_bias_residual",
634
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "present.biasT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
635
+ "derive": {
636
+ "hasBias": "present.biasT",
637
+ "writeResidualSum": "present.residualT",
638
+ "writeMean": "present.meanT",
639
+ "writeInvStd": "present.invStdT",
640
+ "scalar": "dtypes.T",
641
+ "useSubgroups": false,
642
+ "workgroupSize": "skipWg",
643
+ "HIDDEN_LEN": "hiddenSize"
644
+ },
645
+ "passes": [
646
+ {
647
+ "id": "main",
648
+ "shader": "norm-skip-row.wgsl.jinja",
649
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "input_skip_bias_sum_main", "output_main", "inv_std_var", "params_main"],
650
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
651
+ }
652
+ ]
653
+ },
654
+ {
655
+ "id": "stats_both_plain",
656
+ "when": ["statsContract", "present.meanT", "present.invStdT", "not present.biasT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
657
+ "derive": {
658
+ "hasBias": "present.biasT",
659
+ "writeResidualSum": "present.residualT",
660
+ "writeMean": "present.meanT",
661
+ "writeInvStd": "present.invStdT",
662
+ "scalar": "dtypes.T",
663
+ "useSubgroups": false,
664
+ "workgroupSize": "skipWg",
665
+ "HIDDEN_LEN": "hiddenSize"
666
+ },
667
+ "passes": [
668
+ {
669
+ "id": "main",
670
+ "shader": "norm-skip-row.wgsl.jinja",
671
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "mean", "inv_std_var", "params_main"],
672
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
673
+ }
674
+ ]
675
+ },
676
+ {
677
+ "id": "stats_both_residual",
678
+ "when": ["statsContract", "present.meanT", "present.invStdT", "not present.biasT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
679
+ "derive": {
680
+ "hasBias": "present.biasT",
681
+ "writeResidualSum": "present.residualT",
682
+ "writeMean": "present.meanT",
683
+ "writeInvStd": "present.invStdT",
684
+ "scalar": "dtypes.T",
685
+ "useSubgroups": false,
686
+ "workgroupSize": "skipWg",
687
+ "HIDDEN_LEN": "hiddenSize"
688
+ },
689
+ "passes": [
690
+ {
691
+ "id": "main",
692
+ "shader": "norm-skip-row.wgsl.jinja",
693
+ "bindings": ["input_main", "skip_main", "gamma_main", "input_skip_bias_sum_main", "output_main", "mean", "inv_std_var", "params_main"],
694
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
695
+ }
696
+ ]
697
+ },
698
+ {
699
+ "id": "stats_both_bias",
700
+ "when": ["statsContract", "present.meanT", "present.invStdT", "present.biasT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
701
+ "derive": {
702
+ "hasBias": "present.biasT",
703
+ "writeResidualSum": "present.residualT",
704
+ "writeMean": "present.meanT",
705
+ "writeInvStd": "present.invStdT",
706
+ "scalar": "dtypes.T",
707
+ "useSubgroups": false,
708
+ "workgroupSize": "skipWg",
709
+ "HIDDEN_LEN": "hiddenSize"
710
+ },
711
+ "passes": [
712
+ {
713
+ "id": "main",
714
+ "shader": "norm-skip-row.wgsl.jinja",
715
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "mean", "inv_std_var", "params_main"],
716
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
717
+ }
718
+ ]
719
+ },
720
+ {
721
+ "id": "stats_both_bias_residual",
722
+ "when": ["statsContract", "present.meanT", "present.invStdT", "present.biasT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
723
+ "derive": {
724
+ "hasBias": "present.biasT",
725
+ "writeResidualSum": "present.residualT",
726
+ "writeMean": "present.meanT",
727
+ "writeInvStd": "present.invStdT",
728
+ "scalar": "dtypes.T",
729
  "useSubgroups": false,
 
 
 
 
730
  "workgroupSize": "skipWg",
731
  "HIDDEN_LEN": "hiddenSize"
732
  },
733
  "passes": [
734
  {
735
  "id": "main",
 
736
  "shader": "norm-skip-row.wgsl.jinja",
737
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "input_skip_bias_sum_main", "output_main", "mean", "inv_std_var", "params_main"],
738
  "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
739
  }
740
  ]
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "com.microsoft.SkipSimplifiedLayerNormalization",
3
- "id": "_com_microsoft_skipsimplifiedlayernormalization_webgpu_b8ecf51",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,32 +8,44 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "w8dZhYzUtAD/LB4ByqSaHZ6t6a1LKgH6Fb7sFyVZLY8=",
11
- "manifest.json": "Z3UvZaHGuT5XUJRtn5Cb07il2Jc1HM60BXm8GXxeWyY=",
12
- "norm-skip-row-vec4.wgsl.jinja": "gTtEoczLSil0/2T/beWedsZzghnBsunoUp5zcKKAfpM=",
13
- "norm-skip-row.wgsl.jinja": "HwkrJ+4YZjuFEQ87AvYOJX89/G3JeFc/CfOiU0BcPcU=",
14
- "test.json": "E137wUKjS95iD915YA7yUZ4fppQlh8c/s3b5V7spWas="
15
  }
16
  },
17
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
18
  "webgpu": {
19
- "manifestSpec": "2.0",
20
  "variants": {
21
  "no_bias_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
 
22
  "no_bias_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
23
- "no_bias_output_only_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
24
- "no_bias_output_only_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
25
  "no_bias": ["norm-skip-row.wgsl.jinja"],
26
- "no_bias_f16": ["norm-skip-row.wgsl.jinja"],
27
  "no_bias_output_only_f16": ["norm-skip-row.wgsl.jinja"],
 
28
  "no_bias_output_only": ["norm-skip-row.wgsl.jinja"],
29
  "bias_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
 
30
  "bias_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
31
- "bias_output_only_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
32
- "bias_output_only_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
33
  "bias": ["norm-skip-row.wgsl.jinja"],
34
- "bias_f16": ["norm-skip-row.wgsl.jinja"],
35
  "bias_output_only_f16": ["norm-skip-row.wgsl.jinja"],
36
- "bias_output_only": ["norm-skip-row.wgsl.jinja"]
 
 
 
 
 
 
 
 
 
 
 
 
 
37
  }
38
  }
39
  }
 
1
  {
2
  "name": "com.microsoft.SkipSimplifiedLayerNormalization",
3
+ "id": "_com_microsoft_skipsimplifiedlayernormalization_webgpu_e267849",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "w8dZhYzUtAD/LB4ByqSaHZ6t6a1LKgH6Fb7sFyVZLY8=",
11
+ "manifest.json": "MuitIvpEQ8pKpv+ccs0IBVcHTSFzWPldPT6UoePS1fA=",
12
+ "norm-skip-row-vec4.wgsl.jinja": "1PVXdoiIUaLwxANM6flmxpoiFxpw6nSHnFMRt+JU/zc=",
13
+ "norm-skip-row.wgsl.jinja": "WxlAB77lmRSfhyW/ogBRXaxD9q0B8uNCntJHfRV2pyU=",
14
+ "test.json": "jLL2YcZhKbFjF5HwfPEIRLCqLoKvcS2E+pbwBYTMlak="
15
  }
16
  },
17
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
18
  "webgpu": {
19
+ "manifestSpec": "2.1",
20
  "variants": {
21
  "no_bias_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
22
+ "no_bias_f16": ["norm-skip-row.wgsl.jinja"],
23
  "no_bias_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
 
 
24
  "no_bias": ["norm-skip-row.wgsl.jinja"],
25
+ "no_bias_output_only_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
26
  "no_bias_output_only_f16": ["norm-skip-row.wgsl.jinja"],
27
+ "no_bias_output_only_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
28
  "no_bias_output_only": ["norm-skip-row.wgsl.jinja"],
29
  "bias_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
30
+ "bias_f16": ["norm-skip-row.wgsl.jinja"],
31
  "bias_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
 
 
32
  "bias": ["norm-skip-row.wgsl.jinja"],
33
+ "bias_output_only_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
34
  "bias_output_only_f16": ["norm-skip-row.wgsl.jinja"],
35
+ "bias_output_only_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
36
+ "bias_output_only": ["norm-skip-row.wgsl.jinja"],
37
+ "stats_mean_plain": ["norm-skip-row.wgsl.jinja"],
38
+ "stats_mean_residual": ["norm-skip-row.wgsl.jinja"],
39
+ "stats_mean_bias": ["norm-skip-row.wgsl.jinja"],
40
+ "stats_mean_bias_residual": ["norm-skip-row.wgsl.jinja"],
41
+ "stats_inv_plain": ["norm-skip-row.wgsl.jinja"],
42
+ "stats_inv_residual": ["norm-skip-row.wgsl.jinja"],
43
+ "stats_inv_bias": ["norm-skip-row.wgsl.jinja"],
44
+ "stats_inv_bias_residual": ["norm-skip-row.wgsl.jinja"],
45
+ "stats_both_plain": ["norm-skip-row.wgsl.jinja"],
46
+ "stats_both_residual": ["norm-skip-row.wgsl.jinja"],
47
+ "stats_both_bias": ["norm-skip-row.wgsl.jinja"],
48
+ "stats_both_bias_residual": ["norm-skip-row.wgsl.jinja"]
49
  }
50
  }
51
  }
build/webgpu/norm-skip-row-vec4.wgsl.jinja CHANGED
@@ -1,52 +1,19 @@
1
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
2
- {% if op == "max" %}
3
- {{ a }}[{{ idx }}] = max({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);
4
- {%- else %}
5
- {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] + {{ a }}[{{ idx }} + {{ svar }}];
6
- {%- endif %}
7
- {% endmacro %}
8
- {% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false) %}
9
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
10
  loop {
11
- {% if form == "head" %}
12
- {% if breakInline %}
13
  if ({{ svar }} == 0u) { break; }
14
- {% else %}
15
- if ({{ svar }} == 0u) {
16
- break;
17
- }
18
- {% endif %}
19
- {% endif %}
20
- {% if bodyInline %}
21
- if ({{ idx }} < {{ svar }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} }
22
- {% else %}
23
  if ({{ idx }} < {{ svar }}) {
24
  {% for a in arrays %}
25
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
26
  {% endfor %}
27
  }
28
- {% endif %}
29
- {% if form == "head" %}
30
- {% if barrierFirst %}
31
- workgroupBarrier();
32
  {{ svar }} = {{ svar }} / 2u;
33
- {% else %}
34
- {{ svar }} = {{ svar }} / 2u;
35
- workgroupBarrier();
36
- {% endif %}
37
- {% else %}
38
  workgroupBarrier();
39
- if ({{ svar }} == 1u) {
40
- break;
41
- }
42
- {{ svar }} = {{ svar }} / 2u;
43
- {% endif %}
44
- }
45
- {%- endmacro %}{% set broadcastSkip = broadcastSkip is defined and broadcastSkip %}
46
- {% set useSubgroups = useSubgroups %}
47
- {% if usesF16Spec %}
48
- enable f16;
49
- {% endif %}
50
  {% if useSubgroups %}
51
  enable subgroups;
52
  {% endif %}
@@ -107,14 +74,7 @@ fn main(
107
  }
108
  let tid = lid.x;
109
  let base = row * HIDDEN_V;
110
- {% if broadcastSkip %}
111
- // skip broadcasts across the batch dim: fold row into [0, skipRows) so every
112
- // batch reuses the same skip row (skipRows == params.rows ⇒ identity).
113
- let skip_base = (row % params.skipRows) * HIDDEN_V;
114
- {% else %}
115
  let skip_base = base;
116
- {% endif %}
117
-
118
 
119
  var acc = 0.0;
120
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
 
1
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
2
+ {% if op == "max" or op == "min" %}
3
+ {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
4
+ {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %}
5
+ {% 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) %}
 
 
 
6
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
7
  loop {
 
 
8
  if ({{ svar }} == 0u) { break; }
 
 
 
 
 
 
 
 
 
9
  if ({{ idx }} < {{ svar }}) {
10
  {% for a in arrays %}
11
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
12
  {% endfor %}
13
  }
 
 
 
 
14
  {{ svar }} = {{ svar }} / 2u;
 
 
 
 
 
15
  workgroupBarrier();
16
+ }{% endmacro %}
 
 
 
 
 
 
 
 
 
 
17
  {% if useSubgroups %}
18
  enable subgroups;
19
  {% endif %}
 
74
  }
75
  let tid = lid.x;
76
  let base = row * HIDDEN_V;
 
 
 
 
 
77
  let skip_base = base;
 
 
78
 
79
  var acc = 0.0;
80
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
build/webgpu/norm-skip-row.wgsl.jinja CHANGED
@@ -1,17 +1,14 @@
1
-
2
- /* One workgroup normalizes each row of residual = input + skip, with an
3
- * optional bias. */
4
- {% if usesF16 %}
5
- enable f16;
6
- {% endif %}
7
  {{ env.wgsl.resourceDeclarations }}
8
 
9
  const HIDDEN: u32 = {{ hiddenSize }}u;
10
  const WG: u32 = {{ workgroupSize }}u;
11
 
12
  var<workgroup> partial: array<f32, WG>;
13
- {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true) %}
14
- fn {{ name }}(value: f32, tid: u32) -> f32 {
15
  {{ buffer }}[tid] = value;
16
  workgroupBarrier();
17
  // Ceil-halving keeps every lane when the workgroup size is not a power of
@@ -21,11 +18,7 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
21
  loop {
22
  let half = (n + 1u) / 2u;
23
  if (tid < n - half) {
24
- {% if mode == "max" %}
25
- {{ buffer }}[tid] = max({{ buffer }}[tid], {{ buffer }}[tid + half]);
26
- {% else %}
27
  {{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half];
28
- {% endif %}
29
  }
30
  workgroupBarrier();
31
  n = half;
@@ -37,13 +30,10 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
37
  // slot 0 here, so the next call's first store must not run until all lanes have read it.
38
  // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
39
  let reduced = {{ buffer }}[0];
40
- {% if trailingBarrier %}
41
  workgroupBarrier();
42
- {% endif %}
43
  return reduced;
44
  }
45
  {% endmacro %}
46
-
47
  {{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }}
48
  var<workgroup> row_inv: f32;
49
 
@@ -58,11 +48,10 @@ fn residual_value(row: u32, d: u32) -> f32 {
58
 
59
  @compute @workgroup_size(WG, 1, 1)
60
  fn main(
61
- @builtin(workgroup_id) wg: vec3<u32>,
62
  @builtin(local_invocation_id) lid: vec3<u32>) {
63
- // 2D-folded row index: wg.y carries the high bits past the per-axis dispatch fold width.
64
- // Reduces to wg.x when the dispatch does not fold;
65
- // the row >= params.rows guard drops the over-dispatched tail.
66
  let row = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u;
67
  if (row >= params.rows) {
68
  return;
@@ -79,6 +68,12 @@ fn main(
79
  let sq = reduce_sum(local_sq, tid);
80
  if (tid == 0u) {
81
  row_inv = inverseSqrt(sq / f32(HIDDEN) + params.epsilon);
 
 
 
 
 
 
82
  }
83
  workgroupBarrier();
84
 
 
1
+ /* Normalize residual = input + skip, with an optional bias. Reductions use
2
+ * one workgroup per row; closed-form one-element rows use one invocation. */
3
+ {% set degenerateRow = false %}
 
 
 
4
  {{ env.wgsl.resourceDeclarations }}
5
 
6
  const HIDDEN: u32 = {{ hiddenSize }}u;
7
  const WG: u32 = {{ workgroupSize }}u;
8
 
9
  var<workgroup> partial: array<f32, WG>;
10
+ {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true, valueType="f32") %}
11
+ fn {{ name }}(value: {{ valueType }}, tid: u32) -> {{ valueType }} {
12
  {{ buffer }}[tid] = value;
13
  workgroupBarrier();
14
  // Ceil-halving keeps every lane when the workgroup size is not a power of
 
18
  loop {
19
  let half = (n + 1u) / 2u;
20
  if (tid < n - half) {
 
 
 
21
  {{ buffer }}[tid] = {{ buffer }}[tid] + {{ buffer }}[tid + half];
 
22
  }
23
  workgroupBarrier();
24
  n = half;
 
30
  // slot 0 here, so the next call's first store must not run until all lanes have read it.
31
  // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
32
  let reduced = {{ buffer }}[0];
 
33
  workgroupBarrier();
 
34
  return reduced;
35
  }
36
  {% endmacro %}
 
37
  {{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }}
38
  var<workgroup> row_inv: f32;
39
 
 
48
 
49
  @compute @workgroup_size(WG, 1, 1)
50
  fn main(
51
+ @builtin({{ "global_invocation_id" if degenerateRow else "workgroup_id" }}) {{ "gid" if degenerateRow else "wg" }}: vec3<u32>,
52
  @builtin(local_invocation_id) lid: vec3<u32>) {
53
+ // Fold the row grid across workgroups; independent rows also include the
54
+ // invocation offset. The bounds guard drops the final dispatch tail.
 
55
  let row = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u;
56
  if (row >= params.rows) {
57
  return;
 
68
  let sq = reduce_sum(local_sq, tid);
69
  if (tid == 0u) {
70
  row_inv = inverseSqrt(sq / f32(HIDDEN) + params.epsilon);
71
+ {% if writeMean is defined and writeMean %}
72
+ mean[row] = 0.0;
73
+ {% endif %}
74
+ {% if writeInvStd is defined and writeInvStd and not (packedStatistics is defined and packedStatistics) %}
75
+ inv_std_var[row] = row_inv;
76
+ {% endif %}
77
  }
78
  workgroupBarrier();
79
 
build/webgpu/test.json CHANGED
@@ -856,6 +856,210 @@
856
  "allowNaN": true
857
  }
858
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
859
  }
860
  ]
861
  }
 
856
  "allowNaN": true
857
  }
858
  }
859
+ },
860
+ {
861
+ "name": "stats_mean_plain",
862
+ "attrs": { "epsilon": 0.00001 },
863
+ "inputs": {
864
+ "inputT": { "dtype": "float32", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
865
+ "skipT": { "dtype": "float32", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
866
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
867
+ },
868
+ "outputs": {
869
+ "outputT": { "dtype": "float32", "shape": [1, 1, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
870
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
871
+ }
872
+ },
873
+ {
874
+ "name": "stats_mean_residual",
875
+ "attrs": { "epsilon": 0.00001 },
876
+ "inputs": {
877
+ "inputT": { "dtype": "float32", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
878
+ "skipT": { "dtype": "float32", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
879
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
880
+ },
881
+ "outputs": {
882
+ "outputT": { "dtype": "float32", "shape": [1, 1, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
883
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
884
+ "residualT": { "dtype": "float32", "shape": [1, 1, 5], "tolerance": 0.00001 }
885
+ }
886
+ },
887
+ {
888
+ "name": "stats_mean_bias",
889
+ "attrs": { "epsilon": 0.00001 },
890
+ "inputs": {
891
+ "inputT": { "dtype": "float16", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
892
+ "skipT": { "dtype": "float16", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
893
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
894
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
895
+ },
896
+ "outputs": {
897
+ "outputT": { "dtype": "float16", "shape": [1, 1, 5], "tolerance": 0.002, "relTolerance": 0.002 },
898
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
899
+ }
900
+ },
901
+ {
902
+ "name": "stats_mean_bias_residual",
903
+ "attrs": { "epsilon": 0.00001 },
904
+ "inputs": {
905
+ "inputT": { "dtype": "float16", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
906
+ "skipT": { "dtype": "float16", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
907
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
908
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
909
+ },
910
+ "outputs": {
911
+ "outputT": { "dtype": "float16", "shape": [1, 1, 5], "tolerance": 0.002, "relTolerance": 0.002 },
912
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
913
+ "residualT": { "dtype": "float16", "shape": [1, 1, 5], "tolerance": 0.002 }
914
+ }
915
+ },
916
+ {
917
+ "name": "stats_inv_plain",
918
+ "attrs": { "epsilon": 0.00001 },
919
+ "inputs": {
920
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
921
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
922
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
923
+ },
924
+ "outputs": {
925
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
926
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
927
+ }
928
+ },
929
+ {
930
+ "name": "stats_inv_residual",
931
+ "attrs": { "epsilon": 0.00001 },
932
+ "inputs": {
933
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
934
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
935
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
936
+ },
937
+ "outputs": {
938
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
939
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
940
+ "residualT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001 }
941
+ }
942
+ },
943
+ {
944
+ "name": "stats_inv_bias",
945
+ "attrs": { "epsilon": 0.00001 },
946
+ "inputs": {
947
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
948
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
949
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
950
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
951
+ },
952
+ "outputs": {
953
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
954
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
955
+ }
956
+ },
957
+ {
958
+ "name": "stats_inv_bias_residual",
959
+ "attrs": { "epsilon": 0.00001 },
960
+ "inputs": {
961
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
962
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
963
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
964
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
965
+ },
966
+ "outputs": {
967
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
968
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
969
+ "residualT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002 }
970
+ }
971
+ },
972
+ {
973
+ "name": "stats_both_plain",
974
+ "attrs": { "epsilon": 0.00001 },
975
+ "inputs": {
976
+ "inputT": { "dtype": "float32", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
977
+ "skipT": { "dtype": "float32", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
978
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
979
+ },
980
+ "outputs": {
981
+ "outputT": { "dtype": "float32", "shape": [1, 1, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
982
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
983
+ "invStdT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
984
+ }
985
+ },
986
+ {
987
+ "name": "stats_both_residual",
988
+ "attrs": { "epsilon": 0.00001 },
989
+ "inputs": {
990
+ "inputT": { "dtype": "float32", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
991
+ "skipT": { "dtype": "float32", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
992
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
993
+ },
994
+ "outputs": {
995
+ "outputT": { "dtype": "float32", "shape": [1, 1, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
996
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
997
+ "invStdT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
998
+ "residualT": { "dtype": "float32", "shape": [1, 1, 5], "tolerance": 0.00001 }
999
+ }
1000
+ },
1001
+ {
1002
+ "name": "stats_both_bias",
1003
+ "attrs": { "epsilon": 0.00001 },
1004
+ "inputs": {
1005
+ "inputT": { "dtype": "float16", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1006
+ "skipT": { "dtype": "float16", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1007
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1008
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
1009
+ },
1010
+ "outputs": {
1011
+ "outputT": { "dtype": "float16", "shape": [1, 1, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1012
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1013
+ "invStdT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1014
+ }
1015
+ },
1016
+ {
1017
+ "name": "stats_both_bias_residual",
1018
+ "attrs": { "epsilon": 0.00001 },
1019
+ "inputs": {
1020
+ "inputT": { "dtype": "float16", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1021
+ "skipT": { "dtype": "float16", "shape": [1, 1, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1022
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1023
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
1024
+ },
1025
+ "outputs": {
1026
+ "outputT": { "dtype": "float16", "shape": [1, 1, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1027
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1028
+ "invStdT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1029
+ "residualT": { "dtype": "float16", "shape": [1, 1, 5], "tolerance": 0.002 }
1030
+ }
1031
+ },
1032
+ {
1033
+ "name": "ort_saved_statistics",
1034
+ "provenance": {
1035
+ "source": "onnxruntime/test/contrib_ops/skiplayernorm_op_test.cc",
1036
+ "test": "SkipSimplifiedLayerNormStatistics"
1037
+ },
1038
+ "attrs": { "epsilon": 1e-12 },
1039
+ "inputs": {
1040
+ "inputT": {
1041
+ "dtype": "float32",
1042
+ "shape": [1, 1, 4],
1043
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0] }
1044
+ },
1045
+ "skipT": { "dtype": "float32", "shape": [1, 1, 4], "data": { "kind": "constant", "value": 0.0 } },
1046
+ "gammaT": { "dtype": "float32", "shape": [4], "data": { "kind": "constant", "value": 1.0 } }
1047
+ },
1048
+ "outputs": {
1049
+ "outputT": {
1050
+ "dtype": "float32",
1051
+ "shape": [1, 1, 4],
1052
+ "tolerance": 0.000001,
1053
+ "data": { "kind": "values", "values": [0.3651484, 0.7302967, 1.0954452, 1.4605935] }
1054
+ },
1055
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "data": { "kind": "values", "values": [0.0] } },
1056
+ "invStdT": {
1057
+ "dtype": "float32",
1058
+ "shape": [1, 1, 1],
1059
+ "tolerance": 0.000001,
1060
+ "data": { "kind": "values", "values": [0.3651484] }
1061
+ }
1062
+ }
1063
  }
1064
  ]
1065
  }