Xenova HF Staff commited on
Commit
05e7ae4
·
verified ·
1 Parent(s): 62736d3

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -12,7 +12,7 @@ tags:
12
 
13
  ## Description
14
 
15
- Fuses skip addition with layer normalization for rank-2 or rank-3 input and a non-empty hidden axis. With exact-shape skip, float32 and float16 output-only paths support optional `beta`; adding `bias` requires `beta`. Returning the residual sum requires `beta`: float32 supports optional `bias` and arbitrary hidden sizes, while float16 requires `bias` and a hidden size divisible by four. Broadcast skip is supported for rank-3 float32 input, required `beta`, no `bias` or residual output, and a hidden size divisible by four. Bfloat16 and training statistics are not implemented.
16
 
17
  See the [ONNX Runtime `SkipLayerNormalization` contrib-operator spec](https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#com.microsoft.SkipLayerNormalization) for the reference semantics.
18
 
@@ -31,6 +31,8 @@ See the [ONNX Runtime `SkipLayerNormalization` contrib-operator spec](https://gi
31
  | Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
32
  | --- | --- | --- | --- | --- | --- | --- |
33
  | `outputT` | `output` | `T` | same as `inputT` | same as `inputT` | Normalized output tensor with the same shape as `input`. | required |
 
 
34
  | `residualT` | `input_skip_bias_sum` | `T` | same as `inputT` | same as `inputT` | Sum of `input`, `skip`, and `bias` (when present) before normalization, with the same shape as `input`. | optional |
35
 
36
  ## Attributes
@@ -46,6 +48,40 @@ Default values (overridable per request):
46
  | Variable | Allowed dtypes |
47
  | --- | --- |
48
  | `T` | `float32`, `float16` |
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
49
 
50
  ## Device requirements
51
 
@@ -56,14 +92,15 @@ Some implementation variants require `shader-f16`. These are route-specific capa
56
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
57
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
58
  - [`test.json`](build/webgpu/test.json) — correctness cases
59
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
60
  - [`norm-skip-row-vec4.wgsl.jinja`](build/webgpu/norm-skip-row-vec4.wgsl.jinja)
61
  - [`norm-skip-row.wgsl.jinja`](build/webgpu/norm-skip-row.wgsl.jinja)
 
62
 
63
  ## Use with `@huggingface/kernels`
64
 
65
  ```sh
66
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
67
  ```
68
 
69
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
12
 
13
  ## Description
14
 
15
+ Fuses skip addition with layer normalization over a non-empty final hidden axis of rank-2 or rank-3 input. Exact-shape skip supports float32 and float16. Optional float32 mean and inverse-standard-deviation outputs expose row statistics. Broadcast skip supports rank-3 float32 input with beta, no bias or residual output, and a hidden size divisible by four. Output-only and residual-only paths have additional beta, bias and alignment requirements stated by their variants. Bfloat16 is not implemented.
16
 
17
  See the [ONNX Runtime `SkipLayerNormalization` contrib-operator spec](https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#com.microsoft.SkipLayerNormalization) for the reference semantics.
18
 
 
31
  | Name | Upstream name | Logical dtype | Rank | Shape | Description | Presence |
32
  | --- | --- | --- | --- | --- | --- | --- |
33
  | `outputT` | `output` | `T` | same as `inputT` | same as `inputT` | Normalized output tensor with the same shape as `input`. | required |
34
+ | `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 |
35
+ | `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 |
36
  | `residualT` | `input_skip_bias_sum` | `T` | same as `inputT` | same as `inputT` | Sum of `input`, `skip`, and `bias` (when present) before normalization, with the same shape as `input`. | optional |
37
 
38
  ## Attributes
 
48
  | Variable | Allowed dtypes |
49
  | --- | --- |
50
  | `T` | `float32`, `float16` |
51
+ | `U` | `float32` |
52
+
53
+ ## Implementation variants
54
+
55
+ One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
56
+
57
+ - `hidden1_f32_no_beta` — Closed-form one-element rows bind only gamma, optional beta, the output, and row parameters. Each invocation writes one row without reading the unused residual inputs.
58
+ - `hidden1_f32_beta` — Closed-form one-element rows bind only gamma, optional beta, the output, and row parameters. Each invocation writes one row without reading the unused residual inputs.
59
+ - `hidden1_f16_no_beta` — Closed-form one-element rows bind only gamma, optional beta, the output, and row parameters. Each invocation writes one row without reading the unused residual inputs.
60
+ - `hidden1_f16_beta` — Closed-form one-element rows bind only gamma, optional beta, the output, and row parameters. Each invocation writes one row without reading the unused residual inputs.
61
+ - `stats_mean_plain` — Row normalization returning mean statistics with `plain` optional inputs and outputs.
62
+ - `stats_mean_residual` — Row normalization returning mean statistics with `residual` optional inputs and outputs.
63
+ - `stats_mean_beta` — Row normalization returning mean statistics with `beta` optional inputs and outputs.
64
+ - `stats_mean_beta_residual` — Row normalization returning mean statistics with `beta_residual` optional inputs and outputs.
65
+ - `stats_mean_bias` — Row normalization returning mean statistics with `bias` optional inputs and outputs.
66
+ - `stats_mean_bias_residual` — Row normalization returning mean statistics with `bias_residual` optional inputs and outputs.
67
+ - `stats_mean_bias_beta` — Row normalization returning mean statistics with `bias_beta` optional inputs and outputs.
68
+ - `stats_mean_bias_beta_residual` — Row normalization returning mean statistics with `bias_beta_residual` optional inputs and outputs.
69
+ - `stats_inv_plain` — Row normalization returning inv statistics with `plain` optional inputs and outputs.
70
+ - `stats_inv_residual` — Row normalization returning inv statistics with `residual` optional inputs and outputs.
71
+ - `stats_inv_beta` — Row normalization returning inv statistics with `beta` optional inputs and outputs.
72
+ - `stats_inv_beta_residual` — Row normalization returning inv statistics with `beta_residual` optional inputs and outputs.
73
+ - `stats_inv_bias` — Row normalization returning inv statistics with `bias` optional inputs and outputs.
74
+ - `stats_inv_bias_residual` — Row normalization returning inv statistics with `bias_residual` optional inputs and outputs.
75
+ - `stats_inv_bias_beta` — Row normalization returning inv statistics with `bias_beta` optional inputs and outputs.
76
+ - `stats_inv_bias_beta_residual` — Row normalization returning inv statistics with `bias_beta_residual` optional inputs and outputs.
77
+ - `stats_both_plain` — Row normalization returning both statistics with `plain` optional inputs and outputs.
78
+ - `stats_both_residual` — Row normalization returning both statistics with `residual` optional inputs and outputs.
79
+ - `stats_both_beta` — Row normalization returning both statistics with `beta` optional inputs and outputs.
80
+ - `stats_both_beta_residual` — Row normalization returning both statistics with `beta_residual` optional inputs and outputs.
81
+ - `stats_both_bias` — Row normalization returning both statistics with `bias` optional inputs and outputs.
82
+ - `stats_both_bias_residual` — Row normalization returning both statistics with `bias_residual` optional inputs and outputs.
83
+ - `stats_both_bias_beta` — Row normalization returning both statistics with `bias_beta` optional inputs and outputs.
84
+ - `stats_both_bias_beta_residual` — Row normalization returning both statistics with `bias_beta_residual` optional inputs and outputs.
85
 
86
  ## Device requirements
87
 
 
92
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
93
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
94
  - [`test.json`](build/webgpu/test.json) — correctness cases
95
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
96
  - [`norm-skip-row-vec4.wgsl.jinja`](build/webgpu/norm-skip-row-vec4.wgsl.jinja)
97
  - [`norm-skip-row.wgsl.jinja`](build/webgpu/norm-skip-row.wgsl.jinja)
98
+ - [`norm-stats-copy.wgsl.jinja`](build/webgpu/norm-stats-copy.wgsl.jinja)
99
 
100
  ## Use with `@huggingface/kernels`
101
 
102
  ```sh
103
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
104
  ```
105
 
106
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/bench.json CHANGED
@@ -240,6 +240,127 @@
240
  "residualT": { "shape": [65535, 1], "dtype": "float32" }
241
  },
242
  "bench": { "metrics": [{ "type": "bandwidth", "value": "4 * (4 * args.rows * args.hidden + 3 * args.hidden)" }] }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
243
  }
244
  ]
245
  }
 
240
  "residualT": { "shape": [65535, 1], "dtype": "float32" }
241
  },
242
  "bench": { "metrics": [{ "type": "bandwidth", "value": "4 * (4 * args.rows * args.hidden + 3 * args.hidden)" }] }
243
+ },
244
+ {
245
+ "name": "independent_float32_rows257_option2",
246
+ "preset": "edge",
247
+ "attrs": { "epsilon": 0.00001 },
248
+ "inputs": {
249
+ "inputT": { "dtype": "float32", "shape": [257, 1], "dist": "normal", "seed": 1301, "scale": 0.2 },
250
+ "skipT": { "dtype": "float32", "shape": [257, 1], "dist": "normal", "seed": 1302, "scale": 0.2 },
251
+ "gammaT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": 1.125 },
252
+ "betaT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": -0.25 },
253
+ "biasT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": 0.125 }
254
+ },
255
+ "outputs": { "outputT": { "dtype": "float32", "shape": [257, 1] } }
256
+ },
257
+ {
258
+ "name": "independent_float32_rows65537_option0",
259
+ "preset": "edge",
260
+ "attrs": { "epsilon": 0.00001 },
261
+ "inputs": {
262
+ "inputT": { "dtype": "float32", "shape": [65537, 1], "dist": "normal", "seed": 1301, "scale": 0.2 },
263
+ "skipT": { "dtype": "float32", "shape": [65537, 1], "dist": "normal", "seed": 1302, "scale": 0.2 },
264
+ "gammaT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": 1.125 }
265
+ },
266
+ "outputs": { "outputT": { "dtype": "float32", "shape": [65537, 1] } }
267
+ },
268
+ {
269
+ "name": "independent_float32_rows65537_option2",
270
+ "preset": "edge",
271
+ "attrs": { "epsilon": 0.00001 },
272
+ "inputs": {
273
+ "inputT": { "dtype": "float32", "shape": [65537, 1], "dist": "normal", "seed": 1301, "scale": 0.2 },
274
+ "skipT": { "dtype": "float32", "shape": [65537, 1], "dist": "normal", "seed": 1302, "scale": 0.2 },
275
+ "gammaT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": 1.125 },
276
+ "betaT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": -0.25 },
277
+ "biasT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": 0.125 }
278
+ },
279
+ "outputs": { "outputT": { "dtype": "float32", "shape": [65537, 1] } }
280
+ },
281
+ {
282
+ "name": "independent_float32_rows65537_option3",
283
+ "preset": "edge",
284
+ "attrs": { "epsilon": 0.00001 },
285
+ "inputs": {
286
+ "inputT": { "dtype": "float32", "shape": [65537, 1], "dist": "normal", "seed": 1301, "scale": 0.2 },
287
+ "skipT": { "dtype": "float32", "shape": [65537, 1], "dist": "normal", "seed": 1302, "scale": 0.2 },
288
+ "gammaT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": 1.125 },
289
+ "betaT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": -0.25 },
290
+ "biasT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": 0.125 }
291
+ },
292
+ "outputs": {
293
+ "outputT": { "dtype": "float32", "shape": [65537, 1] },
294
+ "residualT": { "dtype": "float32", "shape": [65537, 1] }
295
+ }
296
+ },
297
+ {
298
+ "name": "independent_float16_rows257_option2",
299
+ "preset": "edge",
300
+ "attrs": { "epsilon": 0.00001 },
301
+ "inputs": {
302
+ "inputT": { "dtype": "float16", "shape": [257, 1], "dist": "normal", "seed": 1301, "scale": 0.2 },
303
+ "skipT": { "dtype": "float16", "shape": [257, 1], "dist": "normal", "seed": 1302, "scale": 0.2 },
304
+ "gammaT": { "dtype": "float16", "shape": [1], "dist": "constant", "value": 1.125 },
305
+ "betaT": { "dtype": "float16", "shape": [1], "dist": "constant", "value": -0.25 },
306
+ "biasT": { "dtype": "float16", "shape": [1], "dist": "constant", "value": 0.125 }
307
+ },
308
+ "outputs": { "outputT": { "dtype": "float16", "shape": [257, 1] } }
309
+ },
310
+ {
311
+ "name": "independent_float16_rows65537_option0",
312
+ "preset": "edge",
313
+ "attrs": { "epsilon": 0.00001 },
314
+ "inputs": {
315
+ "inputT": { "dtype": "float16", "shape": [65537, 1], "dist": "normal", "seed": 1301, "scale": 0.2 },
316
+ "skipT": { "dtype": "float16", "shape": [65537, 1], "dist": "normal", "seed": 1302, "scale": 0.2 },
317
+ "gammaT": { "dtype": "float16", "shape": [1], "dist": "constant", "value": 1.125 }
318
+ },
319
+ "outputs": { "outputT": { "dtype": "float16", "shape": [65537, 1] } }
320
+ },
321
+ {
322
+ "name": "independent_float16_rows65537_option2",
323
+ "preset": "edge",
324
+ "attrs": { "epsilon": 0.00001 },
325
+ "inputs": {
326
+ "inputT": { "dtype": "float16", "shape": [65537, 1], "dist": "normal", "seed": 1301, "scale": 0.2 },
327
+ "skipT": { "dtype": "float16", "shape": [65537, 1], "dist": "normal", "seed": 1302, "scale": 0.2 },
328
+ "gammaT": { "dtype": "float16", "shape": [1], "dist": "constant", "value": 1.125 },
329
+ "betaT": { "dtype": "float16", "shape": [1], "dist": "constant", "value": -0.25 },
330
+ "biasT": { "dtype": "float16", "shape": [1], "dist": "constant", "value": 0.125 }
331
+ },
332
+ "outputs": { "outputT": { "dtype": "float16", "shape": [65537, 1] } }
333
+ },
334
+ {
335
+ "name": "independent_float32_rows524289_wg8_fold",
336
+ "preset": "edge",
337
+ "attrs": { "epsilon": 0.00001 },
338
+ "inputs": {
339
+ "inputT": { "dtype": "float32", "shape": [524289, 1], "dist": "normal", "seed": 1301, "scale": 0.2 },
340
+ "skipT": { "dtype": "float32", "shape": [524289, 1], "dist": "normal", "seed": 1302, "scale": 0.2 },
341
+ "gammaT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": 1.125 },
342
+ "betaT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": -0.25 },
343
+ "biasT": { "dtype": "float32", "shape": [1], "dist": "constant", "value": 0.125 }
344
+ },
345
+ "outputs": {
346
+ "outputT": { "dtype": "float32", "shape": [524289, 1] },
347
+ "residualT": { "dtype": "float32", "shape": [524289, 1] }
348
+ },
349
+ "tunables": { "MAX_WORKGROUP_SIZE": 8 }
350
+ },
351
+ {
352
+ "name": "independent_float16_rows524289_wg8_fold",
353
+ "preset": "edge",
354
+ "attrs": { "epsilon": 0.00001 },
355
+ "inputs": {
356
+ "inputT": { "dtype": "float16", "shape": [524289, 1], "dist": "normal", "seed": 1301, "scale": 0.2 },
357
+ "skipT": { "dtype": "float16", "shape": [524289, 1], "dist": "normal", "seed": 1302, "scale": 0.2 },
358
+ "gammaT": { "dtype": "float16", "shape": [1], "dist": "constant", "value": 1.125 },
359
+ "betaT": { "dtype": "float16", "shape": [1], "dist": "constant", "value": -0.25 },
360
+ "biasT": { "dtype": "float16", "shape": [1], "dist": "constant", "value": 0.125 }
361
+ },
362
+ "outputs": { "outputT": { "dtype": "float16", "shape": [524289, 1] } },
363
+ "tunables": { "MAX_WORKGROUP_SIZE": 8 }
364
  }
365
  ]
366
  }
build/webgpu/manifest.json CHANGED
@@ -11,6 +11,20 @@
11
  },
12
  "outputs": {
13
  "outputT": { "onnx": "output", "dtype": "T", "rank": "ranks.inputT", "shape": "shapes.inputT" },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
14
  "residualT": {
15
  "onnx": "input_skip_bias_sum",
16
  "dtype": "T",
@@ -20,36 +34,37 @@
20
  }
21
  },
22
  "attributes": { "epsilon": { "default": 9.999999960041972e-13 } },
23
- "typeConstraints": { "T": ["float32", "float16"] },
24
  "tunables": { "MAX_WORKGROUP_SIZE": { "default": 256 } },
25
  "derive": {
26
  "rowCount": "numel(shapes.inputT) / max(1, dim(shapes.inputT, -1))",
27
  "hiddenSize": "dim(shapes.inputT, -1)",
28
- "skipWg": "max(1, min(tunables.MAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX, pow2ceil(hiddenSize)))",
29
  "skipWgVec4": "max(1, min(tunables.MAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX, pow2ceil(ceilDiv(hiddenSize, 4))))",
30
- "portableWideExecution": "not has(device.adapterInfo, \"subgroupMinSize\") or device.adapterInfo.subgroupMinSize >= 32",
31
- "broadcastRows": "dim(shapes.inputT, 0) * dim(shapes.inputT, 1)",
32
- "broadcastHiddenSize": "dim(shapes.inputT, 2)",
33
- "broadcastSkipWgVec4": "max(1, min(tunables.MAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX, pow2ceil(ceilDiv(broadcastHiddenSize, 4))))",
34
  "rowDispatchFits": "rowCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
35
- "broadcastDispatchFits": "broadcastRows <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
36
  "normResourcesFit": "skipWg * 8 <= device.limits.maxComputeWorkgroupStorageSize and skipWgVec4 * 8 <= device.limits.maxComputeWorkgroupStorageSize",
37
- "broadcastResourcesFit": "broadcastSkipWgVec4 * 8 <= device.limits.maxComputeWorkgroupStorageSize",
38
  "epsilonOk": "attrs.epsilon >= 0",
39
  "coreContract": "epsilonOk and (ranks.inputT == 2 or ranks.inputT == 3) and ranks.skipT == ranks.inputT and ranks.gammaT == 1 and ranks.outputT == ranks.inputT and sameShape(shapes.inputT, shapes.skipT) and sameShape(shapes.outputT, shapes.inputT) and dim(shapes.inputT, -1) > 0 and dim(shapes.gammaT, 0) == dim(shapes.inputT, -1)",
40
  "residualOutputContract": "present.residualT and sameShape(shapes.residualT, shapes.inputT)",
41
  "outputOnlyContract": "not present.residualT",
42
- "betaContract": "false if not present.betaT else (ranks.betaT == 1 and dim(shapes.betaT, 0) == dim(shapes.inputT, -1))",
43
- "noBetaContract": "not present.betaT",
44
  "f32MainDtypes": "tensorDtypes.inputT == \"float32\" and tensorDtypes.skipT == \"float32\" and tensorDtypes.gammaT == \"float32\" and tensorDtypes.outputT == \"float32\"",
45
  "f16MainDtypes": "tensorDtypes.inputT == \"float16\" and tensorDtypes.skipT == \"float16\" and tensorDtypes.gammaT == \"float16\" and tensorDtypes.outputT == \"float16\"",
46
  "f32ResidualDtypes": "f32MainDtypes and tensorDtypes.residualT == \"float32\" if present.residualT else false",
47
  "f16ResidualDtypes": "f16MainDtypes and tensorDtypes.residualT == \"float16\" if present.residualT else false",
48
  "vec4Aligned": "dim(shapes.inputT, -1) % 4 == 0",
49
- "broadcastSkipShapeOk": "(ranks.skipT == 2 and dim(shapes.skipT, 0) == dim(shapes.inputT, 1) and dim(shapes.skipT, 1) == dim(shapes.inputT, 2)) or (ranks.skipT == 3 and ((dim(shapes.skipT, 0) == 1 and dim(shapes.skipT, 1) == dim(shapes.inputT, 1) and dim(shapes.skipT, 2) == dim(shapes.inputT, 2)) or sameShape(shapes.skipT, shapes.inputT)))",
50
- "broadcastOutputOnlyContract": "false if ranks.inputT != 3 or not present.betaT else (epsilonOk and not present.biasT and not present.residualT and dim(shapes.inputT, 2) % 4 == 0 and broadcastSkipShapeOk and ranks.gammaT == 1 and ranks.betaT == 1 and ranks.outputT == 3 and tensorDtypes.inputT == \"float32\" and tensorDtypes.skipT == \"float32\" and tensorDtypes.gammaT == \"float32\" and tensorDtypes.betaT == \"float32\" and tensorDtypes.outputT == \"float32\" and dim(shapes.inputT, 2) > 0 and dim(shapes.gammaT, 0) == dim(shapes.inputT, 2) and dim(shapes.betaT, 0) == dim(shapes.inputT, 2) and sameShape(shapes.outputT, shapes.inputT))",
51
  "hasSubgroups": "device.features.has(\"subgroups\") and device.wgslLanguageFeatures.has(\"subgroup_id\")",
52
  "hasF16": "device.features.has(\"shader-f16\")",
 
 
 
 
 
 
 
 
 
 
 
53
  "f32_beta_no_bias_residual_contract": "coreContract and residualOutputContract and betaContract and f32ResidualDtypes and not present.biasT and tensorDtypes.betaT == \"float32\"",
54
  "f32_beta_bias_residual_contract": "false if not present.biasT else (coreContract and residualOutputContract and betaContract and f32ResidualDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float32\" and tensorDtypes.biasT == \"float32\" and dim(shapes.biasT, 0) == hiddenSize)",
55
  "f16_beta_bias_residual_contract": "false if not present.biasT else (hasF16 and coreContract and residualOutputContract and betaContract and f16ResidualDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float16\" and tensorDtypes.biasT == \"float16\" and dim(shapes.biasT, 0) == hiddenSize)",
@@ -58,19 +73,19 @@
58
  "f32_beta_bias_output_only_contract": "false if not present.biasT else (coreContract and outputOnlyContract and betaContract and f32MainDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float32\" and tensorDtypes.biasT == \"float32\" and dim(shapes.biasT, 0) == hiddenSize)",
59
  "f16_no_beta_output_contract": "hasF16 and coreContract and outputOnlyContract and noBetaContract and f16MainDtypes and not present.biasT",
60
  "f16_beta_no_bias_output_only_contract": "hasF16 and coreContract and outputOnlyContract and betaContract and f16MainDtypes and not present.biasT and tensorDtypes.betaT == \"float16\"",
61
- "f16_beta_bias_output_only_contract": "false if not present.biasT else (hasF16 and coreContract and outputOnlyContract and betaContract and f16MainDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float16\" and tensorDtypes.biasT == \"float16\" and dim(shapes.biasT, 0) == hiddenSize)"
 
62
  },
63
  "bindings": {
64
- "input": { "arg": "inputT", "buffer": "read-only-storage", "elementType": "$vectorScalar" },
65
- "skip": { "arg": "skipT", "buffer": "read-only-storage", "elementType": "$vectorScalar" },
66
- "gamma": { "arg": "gammaT", "buffer": "read-only-storage", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" },
67
- "beta": { "arg": "betaT", "buffer": "read-only-storage", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" },
68
- "output": { "arg": "outputT", "buffer": "storage", "elementType": "$vectorScalar" },
69
- "bias": { "arg": "biasT", "buffer": "read-only-storage", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" },
70
- "input_skip_bias_sum": { "arg": "residualT", "buffer": "storage", "elementType": "$vectorScalar" },
71
- "params_2": {
72
  "name": "params",
73
- "buffer": "uniform",
74
  "struct": [
75
  { "name": "rows", "type": "u32", "value": "rowCount" },
76
  {
@@ -81,51 +96,149 @@
81
  { "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
82
  ]
83
  },
84
- "input_2": { "arg": "inputT", "name": "input", "buffer": "read-only-storage", "elementType": "$scalar" },
85
- "skip_2": { "arg": "skipT", "name": "skip", "buffer": "read-only-storage", "elementType": "$scalar" },
86
- "bias_2": {
87
- "arg": "biasT",
88
- "name": "bias",
89
- "buffer": "read-only-storage",
90
- "elementType": "$scalar",
91
- "length": "$HIDDEN_LEN"
92
- },
93
- "gamma_2": {
94
- "arg": "gammaT",
95
- "name": "gamma",
96
- "buffer": "read-only-storage",
97
- "elementType": "$scalar",
98
- "length": "$HIDDEN_LEN"
99
- },
100
- "beta_2": {
101
- "arg": "betaT",
102
- "name": "beta",
103
- "buffer": "read-only-storage",
104
- "elementType": "$scalar",
105
- "length": "$HIDDEN_LEN"
106
- },
107
- "output_2": { "arg": "outputT", "name": "output", "buffer": "storage", "elementType": "$scalar" },
108
- "input_skip_bias_sum_2": {
109
- "arg": "residualT",
110
- "name": "input_skip_bias_sum",
111
- "buffer": "storage",
112
- "elementType": "$scalar"
113
- },
114
- "params_3": {
115
  "name": "params",
116
- "buffer": "uniform",
117
  "struct": [
118
  { "name": "rows", "type": "u32", "value": "rowCount" },
119
  { "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
120
  ]
121
- }
 
 
 
 
 
 
 
 
 
 
122
  },
123
  "variants": [
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
124
  {
125
  "id": "beta_output_only_vec4_broadcast",
126
  "priority": 19,
127
- "when": ["broadcastOutputOnlyContract", "broadcastResourcesFit", "broadcastDispatchFits"],
128
- "derive": { "scalar": "\"f32\"", "vectorScalar": "\"vec4<f32>\"", "HIDDEN_LEN": "broadcastHiddenSize / 4" },
129
  "passes": [
130
  {
131
  "id": "main",
@@ -136,7 +249,6 @@
136
  "hasBias": false,
137
  "hasBeta": true,
138
  "writeResidualSum": false,
139
- "usesF16Spec": false,
140
  "broadcastSkip": true,
141
  "hidden": "broadcastHiddenSize",
142
  "hiddenVec": "broadcastHiddenSize / 4",
@@ -164,22 +276,15 @@
164
  ]
165
  }
166
  ],
167
- "dispatch": { "x": "min(broadcastRows, 65535)", "y": "ceilDiv(broadcastRows, 65535)", "z": 1 },
168
- "subgroupCollectivesWidth": "portable"
169
  }
170
  ]
171
  },
172
  {
173
  "id": "beta_bias_vec4",
174
  "priority": 15,
175
- "when": ["f32_beta_bias_residual_contract", "vec4Aligned", "hasSubgroups or \"bias\" == \"bias\" or portableWideExecution", "normResourcesFit", "rowDispatchFits"],
176
- "derive": {
177
- "scalar": "\"f32\"",
178
- "vectorScalar": "\"vec4<f32>\"",
179
- "hasBias": "\"bias\" == \"bias\"",
180
- "workgroupSize": "skipWg",
181
- "HIDDEN_LEN": "hiddenSize / 4"
182
- },
183
  "passes": [
184
  {
185
  "id": "normalize",
@@ -187,26 +292,23 @@
187
  "shader": "norm-skip-row-vec4.wgsl.jinja",
188
  "derive": {
189
  "simplified": false,
190
- "hasBias": "\"bias\" == \"bias\"",
191
  "hasBeta": true,
192
  "writeResidualSum": true,
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", "bias", "gamma", "beta", "output", "input_skip_bias_sum", "params_2"],
201
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
202
- "subgroupCollectivesWidth": "portable"
203
  }
204
  ]
205
  },
206
  {
207
  "id": "beta_bias_row",
208
  "priority": 5,
209
- "when": ["f32_beta_bias_residual_contract", "normResourcesFit", "rowDispatchFits"],
210
  "derive": {
211
  "simplified": false,
212
  "useSubgroups": "hasSubgroups",
@@ -223,17 +325,77 @@
223
  "name": "SkipLayerNormalization.Row.Normalize",
224
  "shader": "norm-skip-row.wgsl.jinja",
225
  "derive": { "writeResidualSum": true },
226
- "bindings": ["input_2", "skip_2", "bias_2", "gamma_2", "beta_2", "output_2", "input_skip_bias_sum_2", "params_3"],
227
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
228
- "subgroupCollectivesWidth": "portable"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
229
  }
230
  ]
231
  },
232
  {
233
  "id": "beta_bias_vec4_f16",
234
  "priority": 21,
235
- "when": ["f16_beta_bias_residual_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits"],
236
- "derive": { "scalar": "\"f16\"", "vectorScalar": "\"vec4<f16>\"", "HIDDEN_LEN": "hiddenSize / 4" },
237
  "passes": [
238
  {
239
  "id": "main",
@@ -244,24 +406,22 @@
244
  "hasBias": true,
245
  "hasBeta": true,
246
  "writeResidualSum": true,
247
- "usesF16Spec": true,
248
  "hidden": "hiddenSize",
249
  "hiddenVec": "hiddenSize / 4",
250
  "wg": "skipWgVec4",
251
  "vecType": "\"vec4<f16>\"",
252
  "useSubgroups": "hasSubgroups"
253
  },
254
- "bindings": ["input", "skip", "bias", "gamma", "beta", "output", "input_skip_bias_sum", "params_2"],
255
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
256
- "subgroupCollectivesWidth": "portable"
257
  }
258
  ]
259
  },
260
  {
261
  "id": "no_beta_output_only_vec4",
262
  "priority": 20,
263
- "when": ["f32_no_beta_output_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits"],
264
- "derive": { "scalar": "\"f32\"", "vectorScalar": "\"vec4<f32>\"", "HIDDEN_LEN": "hiddenSize / 4" },
265
  "passes": [
266
  {
267
  "id": "main",
@@ -269,30 +429,28 @@
269
  "shader": "norm-skip-row-vec4.wgsl.jinja",
270
  "derive": {
271
  "simplified": false,
272
- "hasBias": false,
273
- "hasBeta": false,
274
  "writeResidualSum": false,
275
- "usesF16Spec": false,
276
  "hidden": "hiddenSize",
277
  "hiddenVec": "hiddenSize / 4",
278
  "wg": "skipWgVec4",
279
  "vecType": "\"vec4<f32>\"",
280
  "useSubgroups": "hasSubgroups"
281
  },
282
- "bindings": ["input", "skip", "gamma", "output", "params_2"],
283
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
284
- "subgroupCollectivesWidth": "portable"
285
  }
286
  ]
287
  },
288
  {
289
  "id": "no_beta_output_only_row",
290
  "priority": 10,
291
- "when": ["f32_no_beta_output_contract", "normResourcesFit", "rowDispatchFits"],
292
  "derive": {
293
  "simplified": false,
294
- "hasBias": false,
295
- "hasBeta": false,
296
  "writeResidualSum": false,
297
  "useSubgroups": "hasSubgroups",
298
  "scalar": "\"f32\"",
@@ -304,17 +462,20 @@
304
  "id": "main",
305
  "name": "SkipLayerNormalization.NoBetaOutputOnly.Row",
306
  "shader": "norm-skip-row.wgsl.jinja",
307
- "bindings": ["input_2", "skip_2", "gamma_2", "output_2", "params_3"],
308
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
309
- "subgroupCollectivesWidth": "portable"
 
 
 
310
  }
311
  ]
312
  },
313
  {
314
  "id": "no_beta_output_only_vec4_f16",
315
  "priority": 20,
316
- "when": ["f16_no_beta_output_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits"],
317
- "derive": { "scalar": "\"f16\"", "vectorScalar": "\"vec4<f16>\"", "HIDDEN_LEN": "hiddenSize / 4" },
318
  "passes": [
319
  {
320
  "id": "main",
@@ -322,34 +483,31 @@
322
  "shader": "norm-skip-row-vec4.wgsl.jinja",
323
  "derive": {
324
  "simplified": false,
325
- "hasBias": false,
326
- "hasBeta": false,
327
  "writeResidualSum": false,
328
- "usesF16Spec": true,
329
  "hidden": "hiddenSize",
330
  "hiddenVec": "hiddenSize / 4",
331
  "wg": "skipWgVec4",
332
  "vecType": "\"vec4<f16>\"",
333
  "useSubgroups": "hasSubgroups"
334
  },
335
- "bindings": ["input", "skip", "gamma", "output", "params_2"],
336
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
337
- "subgroupCollectivesWidth": "portable"
338
  }
339
  ]
340
  },
341
  {
342
  "id": "no_beta_output_only_row_f16",
343
  "priority": 10,
344
- "when": ["f16_no_beta_output_contract", "normResourcesFit", "rowDispatchFits"],
345
  "derive": {
346
  "simplified": false,
347
- "hasBias": false,
348
- "hasBeta": false,
349
  "writeResidualSum": false,
350
  "useSubgroups": "hasSubgroups",
351
  "scalar": "\"f16\"",
352
- "usesF16": true,
353
  "workgroupSize": "skipWg",
354
  "HIDDEN_LEN": "hiddenSize"
355
  },
@@ -358,23 +516,20 @@
358
  "id": "main",
359
  "name": "SkipLayerNormalization.NoBetaOutputOnly.Row.F16",
360
  "shader": "norm-skip-row.wgsl.jinja",
361
- "bindings": ["input_2", "skip_2", "gamma_2", "output_2", "params_3"],
362
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
363
- "subgroupCollectivesWidth": "portable"
 
 
 
364
  }
365
  ]
366
  },
367
  {
368
- "id": "beta_no_bias_vec4",
369
  "priority": 20,
370
- "when": ["f32_beta_no_bias_residual_contract", "vec4Aligned", "hasSubgroups or portableWideExecution", "normResourcesFit", "rowDispatchFits"],
371
- "derive": {
372
- "scalar": "\"f32\"",
373
- "vectorScalar": "\"vec4<f32>\"",
374
- "hasBias": "\"no_bias\" == \"bias\"",
375
- "workgroupSize": "skipWg",
376
- "HIDDEN_LEN": "hiddenSize / 4"
377
- },
378
  "passes": [
379
  {
380
  "id": "main",
@@ -382,32 +537,30 @@
382
  "shader": "norm-skip-row-vec4.wgsl.jinja",
383
  "derive": {
384
  "simplified": false,
385
- "hasBias": "\"no_bias\" == \"bias\"",
386
- "hasBeta": true,
387
- "writeResidualSum": true,
388
- "usesF16Spec": false,
389
  "hidden": "hiddenSize",
390
  "hiddenVec": "hiddenSize / 4",
391
  "wg": "skipWgVec4",
392
  "vecType": "\"vec4<f32>\"",
393
  "useSubgroups": "hasSubgroups"
394
  },
395
- "bindings": ["input", "skip", "gamma", "beta", "output", "input_skip_bias_sum", "params_2"],
396
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
397
- "subgroupCollectivesWidth": "portable"
398
  }
399
  ]
400
  },
401
  {
402
- "id": "beta_no_bias_row",
403
  "priority": 10,
404
- "when": ["f32_beta_no_bias_residual_contract", "normResourcesFit", "rowDispatchFits"],
405
  "derive": {
406
  "simplified": false,
 
 
 
407
  "useSubgroups": "hasSubgroups",
408
- "hasBeta": true,
409
- "writeResidualSum": true,
410
- "hasBias": "\"no_bias\" == \"bias\"",
411
  "scalar": "\"f32\"",
412
  "workgroupSize": "skipWg",
413
  "HIDDEN_LEN": "hiddenSize"
@@ -417,82 +570,74 @@
417
  "id": "main",
418
  "name": "SkipLayerNormalization.Row",
419
  "shader": "norm-skip-row.wgsl.jinja",
420
- "bindings": ["input_2", "skip_2", "gamma_2", "beta_2", "output_2", "input_skip_bias_sum_2", "params_3"],
421
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
422
- "subgroupCollectivesWidth": "portable"
 
 
 
423
  }
424
  ]
425
  },
426
  {
427
- "id": "beta_no_bias_output_only_vec4",
428
  "priority": 20,
429
- "when": ["f32_beta_no_bias_output_only_contract", "vec4Aligned", "hasSubgroups or \"no_bias\" == \"bias\" or portableWideExecution", "normResourcesFit", "rowDispatchFits"],
430
- "derive": {
431
- "scalar": "\"f32\"",
432
- "vectorScalar": "\"vec4<f32>\"",
433
- "hasBias": "\"no_bias\" == \"bias\"",
434
- "workgroupSize": "skipWg",
435
- "HIDDEN_LEN": "hiddenSize / 4"
436
- },
437
  "passes": [
438
  {
439
  "id": "main",
440
- "name": "SkipLayerNormalization.Vec4",
441
  "shader": "norm-skip-row-vec4.wgsl.jinja",
442
  "derive": {
443
  "simplified": false,
444
- "hasBias": "\"no_bias\" == \"bias\"",
445
- "hasBeta": true,
446
  "writeResidualSum": false,
447
- "usesF16Spec": false,
448
  "hidden": "hiddenSize",
449
  "hiddenVec": "hiddenSize / 4",
450
  "wg": "skipWgVec4",
451
- "vecType": "\"vec4<f32>\"",
452
  "useSubgroups": "hasSubgroups"
453
  },
454
- "bindings": ["input", "skip", "gamma", "beta", "output", "params_2"],
455
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
456
- "subgroupCollectivesWidth": "portable"
457
  }
458
  ]
459
  },
460
  {
461
- "id": "beta_no_bias_output_only_row",
462
  "priority": 10,
463
- "when": ["f32_beta_no_bias_output_only_contract", "normResourcesFit", "rowDispatchFits"],
464
  "derive": {
465
  "simplified": false,
466
- "useSubgroups": "hasSubgroups",
467
- "hasBeta": true,
468
  "writeResidualSum": false,
469
- "hasBias": "\"no_bias\" == \"bias\"",
470
- "scalar": "\"f32\"",
471
  "workgroupSize": "skipWg",
472
  "HIDDEN_LEN": "hiddenSize"
473
  },
474
  "passes": [
475
  {
476
  "id": "main",
477
- "name": "SkipLayerNormalization.Row",
478
  "shader": "norm-skip-row.wgsl.jinja",
479
- "bindings": ["input_2", "skip_2", "gamma_2", "beta_2", "output_2", "params_3"],
480
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
481
- "subgroupCollectivesWidth": "portable"
 
 
 
482
  }
483
  ]
484
  },
485
  {
486
  "id": "beta_bias_output_only_vec4",
487
  "priority": 20,
488
- "when": ["f32_beta_bias_output_only_contract", "vec4Aligned", "hasSubgroups or \"bias\" == \"bias\" or portableWideExecution", "normResourcesFit", "rowDispatchFits"],
489
- "derive": {
490
- "scalar": "\"f32\"",
491
- "vectorScalar": "\"vec4<f32>\"",
492
- "hasBias": "\"bias\" == \"bias\"",
493
- "workgroupSize": "skipWg",
494
- "HIDDEN_LEN": "hiddenSize / 4"
495
- },
496
  "passes": [
497
  {
498
  "id": "main",
@@ -500,32 +645,30 @@
500
  "shader": "norm-skip-row-vec4.wgsl.jinja",
501
  "derive": {
502
  "simplified": false,
503
- "hasBias": "\"bias\" == \"bias\"",
504
- "hasBeta": true,
505
  "writeResidualSum": false,
506
- "usesF16Spec": false,
507
  "hidden": "hiddenSize",
508
  "hiddenVec": "hiddenSize / 4",
509
  "wg": "skipWgVec4",
510
  "vecType": "\"vec4<f32>\"",
511
  "useSubgroups": "hasSubgroups"
512
  },
513
- "bindings": ["input", "skip", "gamma", "beta", "bias", "output", "params_2"],
514
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
515
- "subgroupCollectivesWidth": "portable"
516
  }
517
  ]
518
  },
519
  {
520
  "id": "beta_bias_output_only_row",
521
  "priority": 10,
522
- "when": ["f32_beta_bias_output_only_contract", "normResourcesFit", "rowDispatchFits"],
523
  "derive": {
524
  "simplified": false,
525
- "useSubgroups": "hasSubgroups",
526
- "hasBeta": true,
527
  "writeResidualSum": false,
528
- "hasBias": "\"bias\" == \"bias\"",
529
  "scalar": "\"f32\"",
530
  "workgroupSize": "skipWg",
531
  "HIDDEN_LEN": "hiddenSize"
@@ -535,23 +678,20 @@
535
  "id": "main",
536
  "name": "SkipLayerNormalization.Row",
537
  "shader": "norm-skip-row.wgsl.jinja",
538
- "bindings": ["input_2", "skip_2", "bias_2", "gamma_2", "beta_2", "output_2", "params_3"],
539
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
540
- "subgroupCollectivesWidth": "portable"
 
 
 
541
  }
542
  ]
543
  },
544
  {
545
- "id": "beta_no_bias_output_only_vec4_f16",
546
  "priority": 20,
547
- "when": ["f16_beta_no_bias_output_only_contract", "vec4Aligned", "hasSubgroups or \"no_bias\" == \"bias\" or portableWideExecution", "normResourcesFit", "rowDispatchFits"],
548
- "derive": {
549
- "scalar": "\"f16\"",
550
- "vectorScalar": "\"vec4<f16>\"",
551
- "hasBias": "\"no_bias\" == \"bias\"",
552
- "workgroupSize": "skipWg",
553
- "HIDDEN_LEN": "hiddenSize / 4"
554
- },
555
  "passes": [
556
  {
557
  "id": "main",
@@ -559,34 +699,31 @@
559
  "shader": "norm-skip-row-vec4.wgsl.jinja",
560
  "derive": {
561
  "simplified": false,
562
- "hasBias": "\"no_bias\" == \"bias\"",
563
- "hasBeta": true,
564
  "writeResidualSum": false,
565
- "usesF16Spec": true,
566
  "hidden": "hiddenSize",
567
  "hiddenVec": "hiddenSize / 4",
568
  "wg": "skipWgVec4",
569
  "vecType": "\"vec4<f16>\"",
570
  "useSubgroups": "hasSubgroups"
571
  },
572
- "bindings": ["input", "skip", "gamma", "beta", "output", "params_2"],
573
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
574
- "subgroupCollectivesWidth": "portable"
575
  }
576
  ]
577
  },
578
  {
579
- "id": "beta_no_bias_output_only_row_f16",
580
  "priority": 10,
581
- "when": ["f16_beta_no_bias_output_only_contract", "normResourcesFit", "rowDispatchFits"],
582
  "derive": {
583
  "simplified": false,
584
- "useSubgroups": "hasSubgroups",
585
- "hasBeta": true,
586
  "writeResidualSum": false,
587
- "hasBias": "\"no_bias\" == \"bias\"",
588
  "scalar": "\"f16\"",
589
- "usesF16": true,
590
  "workgroupSize": "skipWg",
591
  "HIDDEN_LEN": "hiddenSize"
592
  },
@@ -595,71 +732,822 @@
595
  "id": "main",
596
  "name": "SkipLayerNormalization.Row.F16",
597
  "shader": "norm-skip-row.wgsl.jinja",
598
- "bindings": ["input_2", "skip_2", "gamma_2", "beta_2", "output_2", "params_3"],
599
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
600
- "subgroupCollectivesWidth": "portable"
 
 
 
601
  }
602
  ]
603
  },
604
  {
605
- "id": "beta_bias_output_only_vec4_f16",
606
- "priority": 20,
607
- "when": ["f16_beta_bias_output_only_contract", "vec4Aligned", "hasSubgroups or \"bias\" == \"bias\" or portableWideExecution", "normResourcesFit", "rowDispatchFits"],
608
  "derive": {
609
- "scalar": "\"f16\"",
610
- "vectorScalar": "\"vec4<f16>\"",
611
- "hasBias": "\"bias\" == \"bias\"",
 
 
 
 
 
612
  "workgroupSize": "skipWg",
613
- "HIDDEN_LEN": "hiddenSize / 4"
 
614
  },
615
  "passes": [
616
  {
617
  "id": "main",
618
- "name": "SkipLayerNormalization.Vec4.F16",
619
- "shader": "norm-skip-row-vec4.wgsl.jinja",
620
- "derive": {
621
- "simplified": false,
622
- "hasBias": "\"bias\" == \"bias\"",
623
- "hasBeta": true,
624
- "writeResidualSum": false,
625
- "usesF16Spec": true,
626
- "hidden": "hiddenSize",
627
- "hiddenVec": "hiddenSize / 4",
628
- "wg": "skipWgVec4",
629
- "vecType": "\"vec4<f16>\"",
630
- "useSubgroups": "hasSubgroups"
631
- },
632
- "bindings": ["input", "skip", "gamma", "beta", "bias", "output", "params_2"],
633
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
634
- "subgroupCollectivesWidth": "portable"
635
  }
636
- ]
 
637
  },
638
  {
639
- "id": "beta_bias_output_only_row_f16",
640
- "priority": 10,
641
- "when": ["f16_beta_bias_output_only_contract", "normResourcesFit", "rowDispatchFits"],
642
  "derive": {
643
  "simplified": false,
644
- "useSubgroups": "hasSubgroups",
645
- "hasBeta": true,
646
- "writeResidualSum": false,
647
- "hasBias": "\"bias\" == \"bias\"",
648
- "scalar": "\"f16\"",
649
- "usesF16": true,
 
650
  "workgroupSize": "skipWg",
651
- "HIDDEN_LEN": "hiddenSize"
 
652
  },
653
  "passes": [
654
  {
655
  "id": "main",
656
- "name": "SkipLayerNormalization.Row.F16",
657
  "shader": "norm-skip-row.wgsl.jinja",
658
- "bindings": ["input_2", "skip_2", "bias_2", "gamma_2", "beta_2", "output_2", "params_3"],
659
- "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 },
660
- "subgroupCollectivesWidth": "portable"
 
 
 
661
  }
662
- ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
663
  }
664
  ]
665
  }
 
11
  },
12
  "outputs": {
13
  "outputT": { "onnx": "output", "dtype": "T", "rank": "ranks.inputT", "shape": "shapes.inputT" },
14
+ "meanT": {
15
+ "onnx": "mean",
16
+ "dtype": "U",
17
+ "optional": true,
18
+ "rank": "ranks.inputT",
19
+ "shape": "prefix(shapes.inputT, ranks.inputT - 1) + [1]"
20
+ },
21
+ "invStdT": {
22
+ "onnx": "inv_std_var",
23
+ "dtype": "U",
24
+ "optional": true,
25
+ "rank": "ranks.inputT",
26
+ "shape": "prefix(shapes.inputT, ranks.inputT - 1) + [1]"
27
+ },
28
  "residualT": {
29
  "onnx": "input_skip_bias_sum",
30
  "dtype": "T",
 
34
  }
35
  },
36
  "attributes": { "epsilon": { "default": 9.999999960041972e-13 } },
37
+ "typeConstraints": { "T": ["float32", "float16"], "U": ["float32"] },
38
  "tunables": { "MAX_WORKGROUP_SIZE": { "default": 256 } },
39
  "derive": {
40
  "rowCount": "numel(shapes.inputT) / max(1, dim(shapes.inputT, -1))",
41
  "hiddenSize": "dim(shapes.inputT, -1)",
42
+ "skipWg": "max(1, min(tunables.MAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)) if hiddenSize == 1 else (max(1, min(tunables.MAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX, pow2ceil(hiddenSize))))",
43
  "skipWgVec4": "max(1, min(tunables.MAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX, pow2ceil(ceilDiv(hiddenSize, 4))))",
 
 
 
 
44
  "rowDispatchFits": "rowCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
 
45
  "normResourcesFit": "skipWg * 8 <= device.limits.maxComputeWorkgroupStorageSize and skipWgVec4 * 8 <= device.limits.maxComputeWorkgroupStorageSize",
 
46
  "epsilonOk": "attrs.epsilon >= 0",
47
  "coreContract": "epsilonOk and (ranks.inputT == 2 or ranks.inputT == 3) and ranks.skipT == ranks.inputT and ranks.gammaT == 1 and ranks.outputT == ranks.inputT and sameShape(shapes.inputT, shapes.skipT) and sameShape(shapes.outputT, shapes.inputT) and dim(shapes.inputT, -1) > 0 and dim(shapes.gammaT, 0) == dim(shapes.inputT, -1)",
48
  "residualOutputContract": "present.residualT and sameShape(shapes.residualT, shapes.inputT)",
49
  "outputOnlyContract": "not present.residualT",
 
 
50
  "f32MainDtypes": "tensorDtypes.inputT == \"float32\" and tensorDtypes.skipT == \"float32\" and tensorDtypes.gammaT == \"float32\" and tensorDtypes.outputT == \"float32\"",
51
  "f16MainDtypes": "tensorDtypes.inputT == \"float16\" and tensorDtypes.skipT == \"float16\" and tensorDtypes.gammaT == \"float16\" and tensorDtypes.outputT == \"float16\"",
52
  "f32ResidualDtypes": "f32MainDtypes and tensorDtypes.residualT == \"float32\" if present.residualT else false",
53
  "f16ResidualDtypes": "f16MainDtypes and tensorDtypes.residualT == \"float16\" if present.residualT else false",
54
  "vec4Aligned": "dim(shapes.inputT, -1) % 4 == 0",
 
 
55
  "hasSubgroups": "device.features.has(\"subgroups\") and device.wgslLanguageFeatures.has(\"subgroup_id\")",
56
  "hasF16": "device.features.has(\"shader-f16\")",
57
+ "statsRequested": "present.meanT or present.invStdT",
58
+ "portableWideExecution": "not has(device.adapterInfo, \"subgroupMinSize\") or device.adapterInfo.subgroupMinSize >= 32",
59
+ "broadcastRows": "dim(shapes.inputT, 0) * dim(shapes.inputT, 1)",
60
+ "broadcastHiddenSize": "dim(shapes.inputT, 2)",
61
+ "broadcastSkipWgVec4": "max(1, min(tunables.MAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX, pow2ceil(ceilDiv(broadcastHiddenSize, 4))))",
62
+ "broadcastDispatchFits": "broadcastRows <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
63
+ "broadcastResourcesFit": "broadcastSkipWgVec4 * 8 <= device.limits.maxComputeWorkgroupStorageSize",
64
+ "betaContract": "false if not present.betaT else (ranks.betaT == 1 and dim(shapes.betaT, 0) == dim(shapes.inputT, -1))",
65
+ "noBetaContract": "not present.betaT",
66
+ "broadcastSkipShapeOk": "(ranks.skipT == 2 and dim(shapes.skipT, 0) == dim(shapes.inputT, 1) and dim(shapes.skipT, 1) == dim(shapes.inputT, 2)) or (ranks.skipT == 3 and ((dim(shapes.skipT, 0) == 1 and dim(shapes.skipT, 1) == dim(shapes.inputT, 1) and dim(shapes.skipT, 2) == dim(shapes.inputT, 2)) or sameShape(shapes.skipT, shapes.inputT)))",
67
+ "broadcastOutputOnlyContract": "false if ranks.inputT != 3 or not present.betaT else (epsilonOk and not present.biasT and not present.residualT and dim(shapes.inputT, 2) % 4 == 0 and broadcastSkipShapeOk and ranks.gammaT == 1 and ranks.betaT == 1 and ranks.outputT == 3 and tensorDtypes.inputT == \"float32\" and tensorDtypes.skipT == \"float32\" and tensorDtypes.gammaT == \"float32\" and tensorDtypes.betaT == \"float32\" and tensorDtypes.outputT == \"float32\" and dim(shapes.inputT, 2) > 0 and dim(shapes.gammaT, 0) == dim(shapes.inputT, 2) and dim(shapes.betaT, 0) == dim(shapes.inputT, 2) and sameShape(shapes.outputT, shapes.inputT))",
68
  "f32_beta_no_bias_residual_contract": "coreContract and residualOutputContract and betaContract and f32ResidualDtypes and not present.biasT and tensorDtypes.betaT == \"float32\"",
69
  "f32_beta_bias_residual_contract": "false if not present.biasT else (coreContract and residualOutputContract and betaContract and f32ResidualDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float32\" and tensorDtypes.biasT == \"float32\" and dim(shapes.biasT, 0) == hiddenSize)",
70
  "f16_beta_bias_residual_contract": "false if not present.biasT else (hasF16 and coreContract and residualOutputContract and betaContract and f16ResidualDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float16\" and tensorDtypes.biasT == \"float16\" and dim(shapes.biasT, 0) == hiddenSize)",
 
73
  "f32_beta_bias_output_only_contract": "false if not present.biasT else (coreContract and outputOnlyContract and betaContract and f32MainDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float32\" and tensorDtypes.biasT == \"float32\" and dim(shapes.biasT, 0) == hiddenSize)",
74
  "f16_no_beta_output_contract": "hasF16 and coreContract and outputOnlyContract and noBetaContract and f16MainDtypes and not present.biasT",
75
  "f16_beta_no_bias_output_only_contract": "hasF16 and coreContract and outputOnlyContract and betaContract and f16MainDtypes and not present.biasT and tensorDtypes.betaT == \"float16\"",
76
+ "f16_beta_bias_output_only_contract": "false if not present.biasT else (hasF16 and coreContract and outputOnlyContract and betaContract and f16MainDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float16\" and tensorDtypes.biasT == \"float16\" and dim(shapes.biasT, 0) == hiddenSize)",
77
+ "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)) and (not present.betaT or (ranks.betaT == 1 and dim(shapes.betaT, 0) == hiddenSize))"
78
  },
79
  "bindings": {
80
+ "input": { "arg": "inputT", "elementType": "$vectorScalar" },
81
+ "skip": { "arg": "skipT", "elementType": "$vectorScalar" },
82
+ "gamma": { "arg": "gammaT", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" },
83
+ "beta": { "arg": "betaT", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" },
84
+ "output": { "arg": "outputT", "elementType": "$vectorScalar" },
85
+ "bias": { "arg": "biasT", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" },
86
+ "input_skip_bias_sum": { "arg": "residualT", "elementType": "$vectorScalar" },
87
+ "params_main": {
88
  "name": "params",
 
89
  "struct": [
90
  { "name": "rows", "type": "u32", "value": "rowCount" },
91
  {
 
96
  { "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
97
  ]
98
  },
99
+ "input_main": { "arg": "inputT", "name": "input", "elementType": "$scalar" },
100
+ "skip_main": { "arg": "skipT", "name": "skip", "elementType": "$scalar" },
101
+ "bias_main": { "arg": "biasT", "name": "bias", "elementType": "$scalar", "length": "$HIDDEN_LEN" },
102
+ "gamma_main": { "arg": "gammaT", "name": "gamma", "elementType": "$scalar", "length": "$HIDDEN_LEN" },
103
+ "beta_main": { "arg": "betaT", "name": "beta", "elementType": "$scalar", "length": "$HIDDEN_LEN" },
104
+ "output_main": { "arg": "outputT", "name": "output", "elementType": "$scalar" },
105
+ "input_skip_bias_sum_main": { "arg": "residualT", "name": "input_skip_bias_sum", "elementType": "$scalar" },
106
+ "params__uniform": {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
107
  "name": "params",
 
108
  "struct": [
109
  { "name": "rows", "type": "u32", "value": "rowCount" },
110
  { "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
111
  ]
112
+ },
113
+ "mean": { "arg": "meanT", "elementType": "f32" },
114
+ "inv_std_var": { "arg": "invStdT", "elementType": "f32" },
115
+ "row_stats": { "scratch": "rowStats", "elementType": "vec2<f32>" },
116
+ "row_stats_read": {
117
+ "scratch": "rowStats",
118
+ "name": "row_stats",
119
+ "buffer": "read-only-storage",
120
+ "elementType": "vec2<f32>"
121
+ },
122
+ "params_stats": { "name": "params", "struct": [{ "name": "rows", "type": "u32", "value": "rowCount" }] }
123
  },
124
  "variants": [
125
+ {
126
+ "id": "hidden1_f32_no_beta",
127
+ "priority": 20,
128
+ "when": ["f32_no_beta_output_contract", "hiddenSize == 1", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
129
+ "derive": {
130
+ "simplified": false,
131
+ "hasBias": false,
132
+ "hasBeta": "present.betaT",
133
+ "writeResidualSum": false,
134
+ "useSubgroups": false,
135
+ "scalar": "dtypes.T",
136
+ "workgroupSize": "skipWg",
137
+ "HIDDEN_LEN": "hiddenSize"
138
+ },
139
+ "passes": [
140
+ {
141
+ "id": "main",
142
+ "name": "SkipLayerNormalization.Hidden1",
143
+ "shader": "norm-skip-row.wgsl.jinja",
144
+ "bindings": ["gamma_main", "output_main", "params__uniform"],
145
+ "dispatch": {
146
+ "x": "min(ceilDiv((rowCount), (skipWg)), 65535)",
147
+ "y": "ceilDiv(ceilDiv((rowCount), (skipWg)), 65535)",
148
+ "z": 1
149
+ }
150
+ }
151
+ ]
152
+ },
153
+ {
154
+ "id": "hidden1_f32_beta",
155
+ "priority": 20,
156
+ "when": ["(f32_beta_no_bias_output_only_contract or f32_beta_bias_output_only_contract)", "hiddenSize == 1", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
157
+ "derive": {
158
+ "simplified": false,
159
+ "hasBias": false,
160
+ "hasBeta": "present.betaT",
161
+ "writeResidualSum": false,
162
+ "useSubgroups": false,
163
+ "scalar": "dtypes.T",
164
+ "workgroupSize": "skipWg",
165
+ "HIDDEN_LEN": "hiddenSize"
166
+ },
167
+ "passes": [
168
+ {
169
+ "id": "main",
170
+ "name": "SkipLayerNormalization.Hidden1",
171
+ "shader": "norm-skip-row.wgsl.jinja",
172
+ "bindings": ["gamma_main", "beta_main", "output_main", "params__uniform"],
173
+ "dispatch": {
174
+ "x": "min(ceilDiv((rowCount), (skipWg)), 65535)",
175
+ "y": "ceilDiv(ceilDiv((rowCount), (skipWg)), 65535)",
176
+ "z": 1
177
+ }
178
+ }
179
+ ]
180
+ },
181
+ {
182
+ "id": "hidden1_f16_no_beta",
183
+ "priority": 20,
184
+ "when": ["f16_no_beta_output_contract", "hiddenSize == 1", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
185
+ "derive": {
186
+ "simplified": false,
187
+ "hasBias": false,
188
+ "hasBeta": "present.betaT",
189
+ "writeResidualSum": false,
190
+ "useSubgroups": false,
191
+ "scalar": "dtypes.T",
192
+ "workgroupSize": "skipWg",
193
+ "HIDDEN_LEN": "hiddenSize"
194
+ },
195
+ "passes": [
196
+ {
197
+ "id": "main",
198
+ "name": "SkipLayerNormalization.Hidden1",
199
+ "shader": "norm-skip-row.wgsl.jinja",
200
+ "bindings": ["gamma_main", "output_main", "params__uniform"],
201
+ "dispatch": {
202
+ "x": "min(ceilDiv((rowCount), (skipWg)), 65535)",
203
+ "y": "ceilDiv(ceilDiv((rowCount), (skipWg)), 65535)",
204
+ "z": 1
205
+ }
206
+ }
207
+ ]
208
+ },
209
+ {
210
+ "id": "hidden1_f16_beta",
211
+ "priority": 20,
212
+ "when": ["(f16_beta_no_bias_output_only_contract or f16_beta_bias_output_only_contract)", "hiddenSize == 1", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
213
+ "derive": {
214
+ "simplified": false,
215
+ "hasBias": false,
216
+ "hasBeta": "present.betaT",
217
+ "writeResidualSum": false,
218
+ "useSubgroups": false,
219
+ "scalar": "dtypes.T",
220
+ "workgroupSize": "skipWg",
221
+ "HIDDEN_LEN": "hiddenSize"
222
+ },
223
+ "passes": [
224
+ {
225
+ "id": "main",
226
+ "name": "SkipLayerNormalization.Hidden1",
227
+ "shader": "norm-skip-row.wgsl.jinja",
228
+ "bindings": ["gamma_main", "beta_main", "output_main", "params__uniform"],
229
+ "dispatch": {
230
+ "x": "min(ceilDiv((rowCount), (skipWg)), 65535)",
231
+ "y": "ceilDiv(ceilDiv((rowCount), (skipWg)), 65535)",
232
+ "z": 1
233
+ }
234
+ }
235
+ ]
236
+ },
237
  {
238
  "id": "beta_output_only_vec4_broadcast",
239
  "priority": 19,
240
+ "when": ["broadcastOutputOnlyContract", "broadcastResourcesFit", "broadcastDispatchFits", "not present.meanT and not present.invStdT"],
241
+ "derive": { "vectorScalar": "\"vec4<f32>\"", "HIDDEN_LEN": "broadcastHiddenSize / 4" },
242
  "passes": [
243
  {
244
  "id": "main",
 
249
  "hasBias": false,
250
  "hasBeta": true,
251
  "writeResidualSum": false,
 
252
  "broadcastSkip": true,
253
  "hidden": "broadcastHiddenSize",
254
  "hiddenVec": "broadcastHiddenSize / 4",
 
276
  ]
277
  }
278
  ],
279
+ "dispatch": { "x": "min(broadcastRows, 65535)", "y": "ceilDiv(broadcastRows, 65535)", "z": 1 }
 
280
  }
281
  ]
282
  },
283
  {
284
  "id": "beta_bias_vec4",
285
  "priority": 15,
286
+ "when": ["f32_beta_bias_residual_contract", "vec4Aligned", "hasSubgroups or \"bias\" == \"bias\" or portableWideExecution", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
287
+ "derive": { "vectorScalar": "\"vec4<f32>\"", "hasBias": "\"bias\" == \"bias\"", "HIDDEN_LEN": "hiddenSize / 4" },
 
 
 
 
 
 
288
  "passes": [
289
  {
290
  "id": "normalize",
 
292
  "shader": "norm-skip-row-vec4.wgsl.jinja",
293
  "derive": {
294
  "simplified": false,
 
295
  "hasBeta": true,
296
  "writeResidualSum": true,
 
297
  "hidden": "hiddenSize",
298
  "hiddenVec": "hiddenSize / 4",
299
  "wg": "skipWgVec4",
300
  "vecType": "\"vec4<f32>\"",
301
  "useSubgroups": "hasSubgroups"
302
  },
303
+ "bindings": ["input", "skip", "bias", "gamma", "beta", "output", "input_skip_bias_sum", "params_main"],
304
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
305
  }
306
  ]
307
  },
308
  {
309
  "id": "beta_bias_row",
310
  "priority": 5,
311
+ "when": ["f32_beta_bias_residual_contract", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
312
  "derive": {
313
  "simplified": false,
314
  "useSubgroups": "hasSubgroups",
 
325
  "name": "SkipLayerNormalization.Row.Normalize",
326
  "shader": "norm-skip-row.wgsl.jinja",
327
  "derive": { "writeResidualSum": true },
328
+ "bindings": ["input_main", "skip_main", "bias_main", "gamma_main", "beta_main", "output_main", "input_skip_bias_sum_main", "params__uniform"],
329
+ "dispatch": {
330
+ "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
331
+ "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
332
+ "z": 1
333
+ }
334
+ }
335
+ ]
336
+ },
337
+ {
338
+ "id": "beta_no_bias_vec4",
339
+ "priority": 20,
340
+ "when": ["f32_beta_no_bias_residual_contract", "vec4Aligned", "hasSubgroups or portableWideExecution", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
341
+ "derive": {
342
+ "vectorScalar": "\"vec4<f32>\"",
343
+ "hasBias": "\"no_bias\" == \"bias\"",
344
+ "HIDDEN_LEN": "hiddenSize / 4"
345
+ },
346
+ "passes": [
347
+ {
348
+ "id": "main",
349
+ "name": "SkipLayerNormalization.Vec4",
350
+ "shader": "norm-skip-row-vec4.wgsl.jinja",
351
+ "derive": {
352
+ "simplified": false,
353
+ "hasBeta": true,
354
+ "writeResidualSum": true,
355
+ "hidden": "hiddenSize",
356
+ "hiddenVec": "hiddenSize / 4",
357
+ "wg": "skipWgVec4",
358
+ "vecType": "\"vec4<f32>\"",
359
+ "useSubgroups": "hasSubgroups"
360
+ },
361
+ "bindings": ["input", "skip", "gamma", "beta", "output", "input_skip_bias_sum", "params_main"],
362
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
363
+ }
364
+ ]
365
+ },
366
+ {
367
+ "id": "beta_no_bias_row",
368
+ "priority": 10,
369
+ "when": ["f32_beta_no_bias_residual_contract", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
370
+ "derive": {
371
+ "simplified": false,
372
+ "useSubgroups": "hasSubgroups",
373
+ "hasBeta": true,
374
+ "writeResidualSum": true,
375
+ "hasBias": "\"no_bias\" == \"bias\"",
376
+ "scalar": "\"f32\"",
377
+ "workgroupSize": "skipWg",
378
+ "HIDDEN_LEN": "hiddenSize"
379
+ },
380
+ "passes": [
381
+ {
382
+ "id": "main",
383
+ "name": "SkipLayerNormalization.Row",
384
+ "shader": "norm-skip-row.wgsl.jinja",
385
+ "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "input_skip_bias_sum_main", "params__uniform"],
386
+ "dispatch": {
387
+ "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
388
+ "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
389
+ "z": 1
390
+ }
391
  }
392
  ]
393
  },
394
  {
395
  "id": "beta_bias_vec4_f16",
396
  "priority": 21,
397
+ "when": ["f16_beta_bias_residual_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
398
+ "derive": { "vectorScalar": "\"vec4<f16>\"", "HIDDEN_LEN": "hiddenSize / 4" },
399
  "passes": [
400
  {
401
  "id": "main",
 
406
  "hasBias": true,
407
  "hasBeta": true,
408
  "writeResidualSum": true,
 
409
  "hidden": "hiddenSize",
410
  "hiddenVec": "hiddenSize / 4",
411
  "wg": "skipWgVec4",
412
  "vecType": "\"vec4<f16>\"",
413
  "useSubgroups": "hasSubgroups"
414
  },
415
+ "bindings": ["input", "skip", "bias", "gamma", "beta", "output", "input_skip_bias_sum", "params_main"],
416
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
417
  }
418
  ]
419
  },
420
  {
421
  "id": "no_beta_output_only_vec4",
422
  "priority": 20,
423
+ "when": ["f32_no_beta_output_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
424
+ "derive": { "vectorScalar": "\"vec4<f32>\"", "HIDDEN_LEN": "hiddenSize / 4" },
425
  "passes": [
426
  {
427
  "id": "main",
 
429
  "shader": "norm-skip-row-vec4.wgsl.jinja",
430
  "derive": {
431
  "simplified": false,
432
+ "hasBias": "present.biasT",
433
+ "hasBeta": "present.betaT",
434
  "writeResidualSum": false,
 
435
  "hidden": "hiddenSize",
436
  "hiddenVec": "hiddenSize / 4",
437
  "wg": "skipWgVec4",
438
  "vecType": "\"vec4<f32>\"",
439
  "useSubgroups": "hasSubgroups"
440
  },
441
+ "bindings": ["input", "skip", "gamma", "output", "params_main"],
442
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
443
  }
444
  ]
445
  },
446
  {
447
  "id": "no_beta_output_only_row",
448
  "priority": 10,
449
+ "when": ["f32_no_beta_output_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"],
450
  "derive": {
451
  "simplified": false,
452
+ "hasBias": "present.biasT",
453
+ "hasBeta": "present.betaT",
454
  "writeResidualSum": false,
455
  "useSubgroups": "hasSubgroups",
456
  "scalar": "\"f32\"",
 
462
  "id": "main",
463
  "name": "SkipLayerNormalization.NoBetaOutputOnly.Row",
464
  "shader": "norm-skip-row.wgsl.jinja",
465
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "params__uniform"],
466
+ "dispatch": {
467
+ "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
468
+ "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
469
+ "z": 1
470
+ }
471
  }
472
  ]
473
  },
474
  {
475
  "id": "no_beta_output_only_vec4_f16",
476
  "priority": 20,
477
+ "when": ["f16_no_beta_output_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
478
+ "derive": { "vectorScalar": "\"vec4<f16>\"", "HIDDEN_LEN": "hiddenSize / 4" },
479
  "passes": [
480
  {
481
  "id": "main",
 
483
  "shader": "norm-skip-row-vec4.wgsl.jinja",
484
  "derive": {
485
  "simplified": false,
486
+ "hasBias": "present.biasT",
487
+ "hasBeta": "present.betaT",
488
  "writeResidualSum": false,
 
489
  "hidden": "hiddenSize",
490
  "hiddenVec": "hiddenSize / 4",
491
  "wg": "skipWgVec4",
492
  "vecType": "\"vec4<f16>\"",
493
  "useSubgroups": "hasSubgroups"
494
  },
495
+ "bindings": ["input", "skip", "gamma", "output", "params_main"],
496
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
497
  }
498
  ]
499
  },
500
  {
501
  "id": "no_beta_output_only_row_f16",
502
  "priority": 10,
503
+ "when": ["f16_no_beta_output_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"],
504
  "derive": {
505
  "simplified": false,
506
+ "hasBias": "present.biasT",
507
+ "hasBeta": "present.betaT",
508
  "writeResidualSum": false,
509
  "useSubgroups": "hasSubgroups",
510
  "scalar": "\"f16\"",
 
511
  "workgroupSize": "skipWg",
512
  "HIDDEN_LEN": "hiddenSize"
513
  },
 
516
  "id": "main",
517
  "name": "SkipLayerNormalization.NoBetaOutputOnly.Row.F16",
518
  "shader": "norm-skip-row.wgsl.jinja",
519
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "params__uniform"],
520
+ "dispatch": {
521
+ "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
522
+ "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
523
+ "z": 1
524
+ }
525
  }
526
  ]
527
  },
528
  {
529
+ "id": "beta_no_bias_output_only_vec4",
530
  "priority": 20,
531
+ "when": ["f32_beta_no_bias_output_only_contract", "vec4Aligned", "hasSubgroups or portableWideExecution", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
532
+ "derive": { "vectorScalar": "\"vec4<f32>\"", "HIDDEN_LEN": "hiddenSize / 4" },
 
 
 
 
 
 
533
  "passes": [
534
  {
535
  "id": "main",
 
537
  "shader": "norm-skip-row-vec4.wgsl.jinja",
538
  "derive": {
539
  "simplified": false,
540
+ "hasBias": "present.biasT",
541
+ "hasBeta": "present.betaT",
542
+ "writeResidualSum": false,
 
543
  "hidden": "hiddenSize",
544
  "hiddenVec": "hiddenSize / 4",
545
  "wg": "skipWgVec4",
546
  "vecType": "\"vec4<f32>\"",
547
  "useSubgroups": "hasSubgroups"
548
  },
549
+ "bindings": ["input", "skip", "gamma", "beta", "output", "params_main"],
550
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
551
  }
552
  ]
553
  },
554
  {
555
+ "id": "beta_no_bias_output_only_row",
556
  "priority": 10,
557
+ "when": ["f32_beta_no_bias_output_only_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"],
558
  "derive": {
559
  "simplified": false,
560
+ "hasBias": "present.biasT",
561
+ "hasBeta": "present.betaT",
562
+ "writeResidualSum": false,
563
  "useSubgroups": "hasSubgroups",
 
 
 
564
  "scalar": "\"f32\"",
565
  "workgroupSize": "skipWg",
566
  "HIDDEN_LEN": "hiddenSize"
 
570
  "id": "main",
571
  "name": "SkipLayerNormalization.Row",
572
  "shader": "norm-skip-row.wgsl.jinja",
573
+ "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "params__uniform"],
574
+ "dispatch": {
575
+ "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
576
+ "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
577
+ "z": 1
578
+ }
579
  }
580
  ]
581
  },
582
  {
583
+ "id": "beta_no_bias_output_only_vec4_f16",
584
  "priority": 20,
585
+ "when": ["f16_beta_no_bias_output_only_contract", "vec4Aligned", "hasSubgroups or portableWideExecution", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
586
+ "derive": { "vectorScalar": "\"vec4<f16>\"", "HIDDEN_LEN": "hiddenSize / 4" },
 
 
 
 
 
 
587
  "passes": [
588
  {
589
  "id": "main",
590
+ "name": "SkipLayerNormalization.Vec4.F16",
591
  "shader": "norm-skip-row-vec4.wgsl.jinja",
592
  "derive": {
593
  "simplified": false,
594
+ "hasBias": "present.biasT",
595
+ "hasBeta": "present.betaT",
596
  "writeResidualSum": false,
 
597
  "hidden": "hiddenSize",
598
  "hiddenVec": "hiddenSize / 4",
599
  "wg": "skipWgVec4",
600
+ "vecType": "\"vec4<f16>\"",
601
  "useSubgroups": "hasSubgroups"
602
  },
603
+ "bindings": ["input", "skip", "gamma", "beta", "output", "params_main"],
604
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
605
  }
606
  ]
607
  },
608
  {
609
+ "id": "beta_no_bias_output_only_row_f16",
610
  "priority": 10,
611
+ "when": ["f16_beta_no_bias_output_only_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"],
612
  "derive": {
613
  "simplified": false,
614
+ "hasBias": "present.biasT",
615
+ "hasBeta": "present.betaT",
616
  "writeResidualSum": false,
617
+ "useSubgroups": "hasSubgroups",
618
+ "scalar": "\"f16\"",
619
  "workgroupSize": "skipWg",
620
  "HIDDEN_LEN": "hiddenSize"
621
  },
622
  "passes": [
623
  {
624
  "id": "main",
625
+ "name": "SkipLayerNormalization.Row.F16",
626
  "shader": "norm-skip-row.wgsl.jinja",
627
+ "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "params__uniform"],
628
+ "dispatch": {
629
+ "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
630
+ "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
631
+ "z": 1
632
+ }
633
  }
634
  ]
635
  },
636
  {
637
  "id": "beta_bias_output_only_vec4",
638
  "priority": 20,
639
+ "when": ["f32_beta_bias_output_only_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
640
+ "derive": { "vectorScalar": "\"vec4<f32>\"", "HIDDEN_LEN": "hiddenSize / 4" },
 
 
 
 
 
 
641
  "passes": [
642
  {
643
  "id": "main",
 
645
  "shader": "norm-skip-row-vec4.wgsl.jinja",
646
  "derive": {
647
  "simplified": false,
648
+ "hasBias": "present.biasT",
649
+ "hasBeta": "present.betaT",
650
  "writeResidualSum": false,
 
651
  "hidden": "hiddenSize",
652
  "hiddenVec": "hiddenSize / 4",
653
  "wg": "skipWgVec4",
654
  "vecType": "\"vec4<f32>\"",
655
  "useSubgroups": "hasSubgroups"
656
  },
657
+ "bindings": ["input", "skip", "gamma", "beta", "bias", "output", "params_main"],
658
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
659
  }
660
  ]
661
  },
662
  {
663
  "id": "beta_bias_output_only_row",
664
  "priority": 10,
665
+ "when": ["f32_beta_bias_output_only_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"],
666
  "derive": {
667
  "simplified": false,
668
+ "hasBias": "present.biasT",
669
+ "hasBeta": "present.betaT",
670
  "writeResidualSum": false,
671
+ "useSubgroups": "hasSubgroups",
672
  "scalar": "\"f32\"",
673
  "workgroupSize": "skipWg",
674
  "HIDDEN_LEN": "hiddenSize"
 
678
  "id": "main",
679
  "name": "SkipLayerNormalization.Row",
680
  "shader": "norm-skip-row.wgsl.jinja",
681
+ "bindings": ["input_main", "skip_main", "bias_main", "gamma_main", "beta_main", "output_main", "params__uniform"],
682
+ "dispatch": {
683
+ "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
684
+ "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
685
+ "z": 1
686
+ }
687
  }
688
  ]
689
  },
690
  {
691
+ "id": "beta_bias_output_only_vec4_f16",
692
  "priority": 20,
693
+ "when": ["f16_beta_bias_output_only_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"],
694
+ "derive": { "vectorScalar": "\"vec4<f16>\"", "HIDDEN_LEN": "hiddenSize / 4" },
 
 
 
 
 
 
695
  "passes": [
696
  {
697
  "id": "main",
 
699
  "shader": "norm-skip-row-vec4.wgsl.jinja",
700
  "derive": {
701
  "simplified": false,
702
+ "hasBias": "present.biasT",
703
+ "hasBeta": "present.betaT",
704
  "writeResidualSum": false,
 
705
  "hidden": "hiddenSize",
706
  "hiddenVec": "hiddenSize / 4",
707
  "wg": "skipWgVec4",
708
  "vecType": "\"vec4<f16>\"",
709
  "useSubgroups": "hasSubgroups"
710
  },
711
+ "bindings": ["input", "skip", "gamma", "beta", "bias", "output", "params_main"],
712
+ "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 }
 
713
  }
714
  ]
715
  },
716
  {
717
+ "id": "beta_bias_output_only_row_f16",
718
  "priority": 10,
719
+ "when": ["f16_beta_bias_output_only_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"],
720
  "derive": {
721
  "simplified": false,
722
+ "hasBias": "present.biasT",
723
+ "hasBeta": "present.betaT",
724
  "writeResidualSum": false,
725
+ "useSubgroups": "hasSubgroups",
726
  "scalar": "\"f16\"",
 
727
  "workgroupSize": "skipWg",
728
  "HIDDEN_LEN": "hiddenSize"
729
  },
 
732
  "id": "main",
733
  "name": "SkipLayerNormalization.Row.F16",
734
  "shader": "norm-skip-row.wgsl.jinja",
735
+ "bindings": ["input_main", "skip_main", "bias_main", "gamma_main", "beta_main", "output_main", "params__uniform"],
736
+ "dispatch": {
737
+ "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
738
+ "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)",
739
+ "z": 1
740
+ }
741
  }
742
  ]
743
  },
744
  {
745
+ "id": "stats_mean_plain",
746
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "not present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
 
747
  "derive": {
748
+ "simplified": false,
749
+ "hasBias": "present.biasT",
750
+ "hasBeta": "present.betaT",
751
+ "writeResidualSum": "present.residualT",
752
+ "writeMean": "present.meanT",
753
+ "writeInvStd": "present.invStdT",
754
+ "scalar": "dtypes.T",
755
+ "useSubgroups": false,
756
  "workgroupSize": "skipWg",
757
+ "HIDDEN_LEN": "hiddenSize",
758
+ "packedStatistics": false
759
  },
760
  "passes": [
761
  {
762
  "id": "main",
763
+ "shader": "norm-skip-row.wgsl.jinja",
764
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "mean", "params__uniform"],
765
+ "dispatch": {
766
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
767
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
768
+ "z": 1
769
+ }
 
 
 
 
 
 
 
 
 
 
770
  }
771
+ ],
772
+ "intermediates": []
773
  },
774
  {
775
+ "id": "stats_mean_residual",
776
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "not present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
 
777
  "derive": {
778
  "simplified": false,
779
+ "hasBias": "present.biasT",
780
+ "hasBeta": "present.betaT",
781
+ "writeResidualSum": "present.residualT",
782
+ "writeMean": "present.meanT",
783
+ "writeInvStd": "present.invStdT",
784
+ "scalar": "dtypes.T",
785
+ "useSubgroups": false,
786
  "workgroupSize": "skipWg",
787
+ "HIDDEN_LEN": "hiddenSize",
788
+ "packedStatistics": false
789
  },
790
  "passes": [
791
  {
792
  "id": "main",
 
793
  "shader": "norm-skip-row.wgsl.jinja",
794
+ "bindings": ["input_main", "skip_main", "gamma_main", "input_skip_bias_sum_main", "output_main", "mean", "params__uniform"],
795
+ "dispatch": {
796
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
797
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
798
+ "z": 1
799
+ }
800
  }
801
+ ],
802
+ "intermediates": []
803
+ },
804
+ {
805
+ "id": "stats_mean_beta",
806
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "not present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
807
+ "derive": {
808
+ "simplified": false,
809
+ "hasBias": "present.biasT",
810
+ "hasBeta": "present.betaT",
811
+ "writeResidualSum": "present.residualT",
812
+ "writeMean": "present.meanT",
813
+ "writeInvStd": "present.invStdT",
814
+ "scalar": "dtypes.T",
815
+ "useSubgroups": false,
816
+ "workgroupSize": "skipWg",
817
+ "HIDDEN_LEN": "hiddenSize",
818
+ "packedStatistics": false
819
+ },
820
+ "passes": [
821
+ {
822
+ "id": "main",
823
+ "shader": "norm-skip-row.wgsl.jinja",
824
+ "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "mean", "params__uniform"],
825
+ "dispatch": {
826
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
827
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
828
+ "z": 1
829
+ }
830
+ }
831
+ ],
832
+ "intermediates": []
833
+ },
834
+ {
835
+ "id": "stats_mean_beta_residual",
836
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "not present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
837
+ "derive": {
838
+ "simplified": false,
839
+ "hasBias": "present.biasT",
840
+ "hasBeta": "present.betaT",
841
+ "writeResidualSum": "present.residualT",
842
+ "writeMean": "present.meanT",
843
+ "writeInvStd": "present.invStdT",
844
+ "scalar": "dtypes.T",
845
+ "useSubgroups": false,
846
+ "workgroupSize": "skipWg",
847
+ "HIDDEN_LEN": "hiddenSize",
848
+ "packedStatistics": false
849
+ },
850
+ "passes": [
851
+ {
852
+ "id": "main",
853
+ "shader": "norm-skip-row.wgsl.jinja",
854
+ "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "input_skip_bias_sum_main", "output_main", "mean", "params__uniform"],
855
+ "dispatch": {
856
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
857
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
858
+ "z": 1
859
+ }
860
+ }
861
+ ],
862
+ "intermediates": []
863
+ },
864
+ {
865
+ "id": "stats_mean_bias",
866
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
867
+ "derive": {
868
+ "simplified": false,
869
+ "hasBias": "present.biasT",
870
+ "hasBeta": "present.betaT",
871
+ "writeResidualSum": "present.residualT",
872
+ "writeMean": "present.meanT",
873
+ "writeInvStd": "present.invStdT",
874
+ "scalar": "dtypes.T",
875
+ "useSubgroups": false,
876
+ "workgroupSize": "skipWg",
877
+ "HIDDEN_LEN": "hiddenSize",
878
+ "packedStatistics": false
879
+ },
880
+ "passes": [
881
+ {
882
+ "id": "main",
883
+ "shader": "norm-skip-row.wgsl.jinja",
884
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "mean", "params__uniform"],
885
+ "dispatch": {
886
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
887
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
888
+ "z": 1
889
+ }
890
+ }
891
+ ],
892
+ "intermediates": []
893
+ },
894
+ {
895
+ "id": "stats_mean_bias_residual",
896
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
897
+ "derive": {
898
+ "simplified": false,
899
+ "hasBias": "present.biasT",
900
+ "hasBeta": "present.betaT",
901
+ "writeResidualSum": "present.residualT",
902
+ "writeMean": "present.meanT",
903
+ "writeInvStd": "present.invStdT",
904
+ "scalar": "dtypes.T",
905
+ "useSubgroups": false,
906
+ "workgroupSize": "skipWg",
907
+ "HIDDEN_LEN": "hiddenSize",
908
+ "packedStatistics": false
909
+ },
910
+ "passes": [
911
+ {
912
+ "id": "main",
913
+ "shader": "norm-skip-row.wgsl.jinja",
914
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "input_skip_bias_sum_main", "output_main", "mean", "params__uniform"],
915
+ "dispatch": {
916
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
917
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
918
+ "z": 1
919
+ }
920
+ }
921
+ ],
922
+ "intermediates": []
923
+ },
924
+ {
925
+ "id": "stats_mean_bias_beta",
926
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
927
+ "derive": {
928
+ "simplified": false,
929
+ "hasBias": "present.biasT",
930
+ "hasBeta": "present.betaT",
931
+ "writeResidualSum": "present.residualT",
932
+ "writeMean": "present.meanT",
933
+ "writeInvStd": "present.invStdT",
934
+ "scalar": "dtypes.T",
935
+ "useSubgroups": false,
936
+ "workgroupSize": "skipWg",
937
+ "HIDDEN_LEN": "hiddenSize",
938
+ "packedStatistics": false
939
+ },
940
+ "passes": [
941
+ {
942
+ "id": "main",
943
+ "shader": "norm-skip-row.wgsl.jinja",
944
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "output_main", "mean", "params__uniform"],
945
+ "dispatch": {
946
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
947
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
948
+ "z": 1
949
+ }
950
+ }
951
+ ],
952
+ "intermediates": []
953
+ },
954
+ {
955
+ "id": "stats_mean_bias_beta_residual",
956
+ "when": ["statsContract", "present.meanT", "not present.invStdT", "present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
957
+ "derive": {
958
+ "simplified": false,
959
+ "hasBias": "present.biasT",
960
+ "hasBeta": "present.betaT",
961
+ "writeResidualSum": "present.residualT",
962
+ "writeMean": "present.meanT",
963
+ "writeInvStd": "present.invStdT",
964
+ "scalar": "dtypes.T",
965
+ "useSubgroups": false,
966
+ "workgroupSize": "skipWg",
967
+ "HIDDEN_LEN": "hiddenSize",
968
+ "packedStatistics": false
969
+ },
970
+ "passes": [
971
+ {
972
+ "id": "main",
973
+ "shader": "norm-skip-row.wgsl.jinja",
974
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "input_skip_bias_sum_main", "output_main", "mean", "params__uniform"],
975
+ "dispatch": {
976
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
977
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
978
+ "z": 1
979
+ }
980
+ }
981
+ ],
982
+ "intermediates": []
983
+ },
984
+ {
985
+ "id": "stats_inv_plain",
986
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "not present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
987
+ "derive": {
988
+ "simplified": false,
989
+ "hasBias": "present.biasT",
990
+ "hasBeta": "present.betaT",
991
+ "writeResidualSum": "present.residualT",
992
+ "writeMean": "present.meanT",
993
+ "writeInvStd": "present.invStdT",
994
+ "scalar": "dtypes.T",
995
+ "useSubgroups": false,
996
+ "workgroupSize": "skipWg",
997
+ "HIDDEN_LEN": "hiddenSize",
998
+ "packedStatistics": false
999
+ },
1000
+ "passes": [
1001
+ {
1002
+ "id": "main",
1003
+ "shader": "norm-skip-row.wgsl.jinja",
1004
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "inv_std_var", "params__uniform"],
1005
+ "dispatch": {
1006
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1007
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1008
+ "z": 1
1009
+ }
1010
+ }
1011
+ ],
1012
+ "intermediates": []
1013
+ },
1014
+ {
1015
+ "id": "stats_inv_residual",
1016
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "not present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
1017
+ "derive": {
1018
+ "simplified": false,
1019
+ "hasBias": "present.biasT",
1020
+ "hasBeta": "present.betaT",
1021
+ "writeResidualSum": "present.residualT",
1022
+ "writeMean": "present.meanT",
1023
+ "writeInvStd": "present.invStdT",
1024
+ "scalar": "dtypes.T",
1025
+ "useSubgroups": false,
1026
+ "workgroupSize": "skipWg",
1027
+ "HIDDEN_LEN": "hiddenSize",
1028
+ "packedStatistics": false
1029
+ },
1030
+ "passes": [
1031
+ {
1032
+ "id": "main",
1033
+ "shader": "norm-skip-row.wgsl.jinja",
1034
+ "bindings": ["input_main", "skip_main", "gamma_main", "input_skip_bias_sum_main", "output_main", "inv_std_var", "params__uniform"],
1035
+ "dispatch": {
1036
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1037
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1038
+ "z": 1
1039
+ }
1040
+ }
1041
+ ],
1042
+ "intermediates": []
1043
+ },
1044
+ {
1045
+ "id": "stats_inv_beta",
1046
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "not present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
1047
+ "derive": {
1048
+ "simplified": false,
1049
+ "hasBias": "present.biasT",
1050
+ "hasBeta": "present.betaT",
1051
+ "writeResidualSum": "present.residualT",
1052
+ "writeMean": "present.meanT",
1053
+ "writeInvStd": "present.invStdT",
1054
+ "scalar": "dtypes.T",
1055
+ "useSubgroups": false,
1056
+ "workgroupSize": "skipWg",
1057
+ "HIDDEN_LEN": "hiddenSize",
1058
+ "packedStatistics": false
1059
+ },
1060
+ "passes": [
1061
+ {
1062
+ "id": "main",
1063
+ "shader": "norm-skip-row.wgsl.jinja",
1064
+ "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "inv_std_var", "params__uniform"],
1065
+ "dispatch": {
1066
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1067
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1068
+ "z": 1
1069
+ }
1070
+ }
1071
+ ],
1072
+ "intermediates": []
1073
+ },
1074
+ {
1075
+ "id": "stats_inv_beta_residual",
1076
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "not present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
1077
+ "derive": {
1078
+ "simplified": false,
1079
+ "hasBias": "present.biasT",
1080
+ "hasBeta": "present.betaT",
1081
+ "writeResidualSum": "present.residualT",
1082
+ "writeMean": "present.meanT",
1083
+ "writeInvStd": "present.invStdT",
1084
+ "scalar": "dtypes.T",
1085
+ "useSubgroups": false,
1086
+ "workgroupSize": "skipWg",
1087
+ "HIDDEN_LEN": "hiddenSize",
1088
+ "packedStatistics": false
1089
+ },
1090
+ "passes": [
1091
+ {
1092
+ "id": "main",
1093
+ "shader": "norm-skip-row.wgsl.jinja",
1094
+ "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "input_skip_bias_sum_main", "output_main", "inv_std_var", "params__uniform"],
1095
+ "dispatch": {
1096
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1097
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1098
+ "z": 1
1099
+ }
1100
+ }
1101
+ ],
1102
+ "intermediates": []
1103
+ },
1104
+ {
1105
+ "id": "stats_inv_bias",
1106
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
1107
+ "derive": {
1108
+ "simplified": false,
1109
+ "hasBias": "present.biasT",
1110
+ "hasBeta": "present.betaT",
1111
+ "writeResidualSum": "present.residualT",
1112
+ "writeMean": "present.meanT",
1113
+ "writeInvStd": "present.invStdT",
1114
+ "scalar": "dtypes.T",
1115
+ "useSubgroups": false,
1116
+ "workgroupSize": "skipWg",
1117
+ "HIDDEN_LEN": "hiddenSize",
1118
+ "packedStatistics": false
1119
+ },
1120
+ "passes": [
1121
+ {
1122
+ "id": "main",
1123
+ "shader": "norm-skip-row.wgsl.jinja",
1124
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "inv_std_var", "params__uniform"],
1125
+ "dispatch": {
1126
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1127
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1128
+ "z": 1
1129
+ }
1130
+ }
1131
+ ],
1132
+ "intermediates": []
1133
+ },
1134
+ {
1135
+ "id": "stats_inv_bias_residual",
1136
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
1137
+ "derive": {
1138
+ "simplified": false,
1139
+ "hasBias": "present.biasT",
1140
+ "hasBeta": "present.betaT",
1141
+ "writeResidualSum": "present.residualT",
1142
+ "writeMean": "present.meanT",
1143
+ "writeInvStd": "present.invStdT",
1144
+ "scalar": "dtypes.T",
1145
+ "useSubgroups": false,
1146
+ "workgroupSize": "skipWg",
1147
+ "HIDDEN_LEN": "hiddenSize",
1148
+ "packedStatistics": false
1149
+ },
1150
+ "passes": [
1151
+ {
1152
+ "id": "main",
1153
+ "shader": "norm-skip-row.wgsl.jinja",
1154
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "input_skip_bias_sum_main", "output_main", "inv_std_var", "params__uniform"],
1155
+ "dispatch": {
1156
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1157
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1158
+ "z": 1
1159
+ }
1160
+ }
1161
+ ],
1162
+ "intermediates": []
1163
+ },
1164
+ {
1165
+ "id": "stats_inv_bias_beta",
1166
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
1167
+ "derive": {
1168
+ "simplified": false,
1169
+ "hasBias": "present.biasT",
1170
+ "hasBeta": "present.betaT",
1171
+ "writeResidualSum": "present.residualT",
1172
+ "writeMean": "present.meanT",
1173
+ "writeInvStd": "present.invStdT",
1174
+ "scalar": "dtypes.T",
1175
+ "useSubgroups": false,
1176
+ "workgroupSize": "skipWg",
1177
+ "HIDDEN_LEN": "hiddenSize",
1178
+ "packedStatistics": false
1179
+ },
1180
+ "passes": [
1181
+ {
1182
+ "id": "main",
1183
+ "shader": "norm-skip-row.wgsl.jinja",
1184
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "output_main", "inv_std_var", "params__uniform"],
1185
+ "dispatch": {
1186
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1187
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1188
+ "z": 1
1189
+ }
1190
+ }
1191
+ ],
1192
+ "intermediates": []
1193
+ },
1194
+ {
1195
+ "id": "stats_inv_bias_beta_residual",
1196
+ "when": ["statsContract", "not present.meanT", "present.invStdT", "present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
1197
+ "derive": {
1198
+ "simplified": false,
1199
+ "hasBias": "present.biasT",
1200
+ "hasBeta": "present.betaT",
1201
+ "writeResidualSum": "present.residualT",
1202
+ "writeMean": "present.meanT",
1203
+ "writeInvStd": "present.invStdT",
1204
+ "scalar": "dtypes.T",
1205
+ "useSubgroups": false,
1206
+ "workgroupSize": "skipWg",
1207
+ "HIDDEN_LEN": "hiddenSize",
1208
+ "packedStatistics": false
1209
+ },
1210
+ "passes": [
1211
+ {
1212
+ "id": "main",
1213
+ "shader": "norm-skip-row.wgsl.jinja",
1214
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "input_skip_bias_sum_main", "output_main", "inv_std_var", "params__uniform"],
1215
+ "dispatch": {
1216
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1217
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1218
+ "z": 1
1219
+ }
1220
+ }
1221
+ ],
1222
+ "intermediates": []
1223
+ },
1224
+ {
1225
+ "id": "stats_both_plain",
1226
+ "when": ["statsContract", "present.meanT", "present.invStdT", "not present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
1227
+ "derive": {
1228
+ "simplified": false,
1229
+ "hasBias": "present.biasT",
1230
+ "hasBeta": "present.betaT",
1231
+ "writeResidualSum": "present.residualT",
1232
+ "writeMean": "present.meanT",
1233
+ "writeInvStd": "present.invStdT",
1234
+ "scalar": "dtypes.T",
1235
+ "useSubgroups": false,
1236
+ "workgroupSize": "skipWg",
1237
+ "HIDDEN_LEN": "hiddenSize",
1238
+ "packedStatistics": true
1239
+ },
1240
+ "passes": [
1241
+ {
1242
+ "id": "main",
1243
+ "shader": "norm-skip-row.wgsl.jinja",
1244
+ "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "row_stats", "params__uniform"],
1245
+ "dispatch": {
1246
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1247
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1248
+ "z": 1
1249
+ }
1250
+ },
1251
+ {
1252
+ "id": "statistics",
1253
+ "shader": "norm-stats-copy.wgsl.jinja",
1254
+ "derive": { "workgroupSize": 64 },
1255
+ "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"],
1256
+ "dispatch": {
1257
+ "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)",
1258
+ "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)",
1259
+ "z": 1
1260
+ }
1261
+ }
1262
+ ],
1263
+ "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }]
1264
+ },
1265
+ {
1266
+ "id": "stats_both_residual",
1267
+ "when": ["statsContract", "present.meanT", "present.invStdT", "not present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
1268
+ "derive": {
1269
+ "simplified": false,
1270
+ "hasBias": "present.biasT",
1271
+ "hasBeta": "present.betaT",
1272
+ "writeResidualSum": "present.residualT",
1273
+ "writeMean": "present.meanT",
1274
+ "writeInvStd": "present.invStdT",
1275
+ "scalar": "dtypes.T",
1276
+ "useSubgroups": false,
1277
+ "workgroupSize": "skipWg",
1278
+ "HIDDEN_LEN": "hiddenSize",
1279
+ "packedStatistics": true
1280
+ },
1281
+ "passes": [
1282
+ {
1283
+ "id": "main",
1284
+ "shader": "norm-skip-row.wgsl.jinja",
1285
+ "bindings": ["input_main", "skip_main", "gamma_main", "input_skip_bias_sum_main", "output_main", "row_stats", "params__uniform"],
1286
+ "dispatch": {
1287
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1288
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1289
+ "z": 1
1290
+ }
1291
+ },
1292
+ {
1293
+ "id": "statistics",
1294
+ "shader": "norm-stats-copy.wgsl.jinja",
1295
+ "derive": { "workgroupSize": 64 },
1296
+ "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"],
1297
+ "dispatch": {
1298
+ "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)",
1299
+ "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)",
1300
+ "z": 1
1301
+ }
1302
+ }
1303
+ ],
1304
+ "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }]
1305
+ },
1306
+ {
1307
+ "id": "stats_both_beta",
1308
+ "when": ["statsContract", "present.meanT", "present.invStdT", "not present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
1309
+ "derive": {
1310
+ "simplified": false,
1311
+ "hasBias": "present.biasT",
1312
+ "hasBeta": "present.betaT",
1313
+ "writeResidualSum": "present.residualT",
1314
+ "writeMean": "present.meanT",
1315
+ "writeInvStd": "present.invStdT",
1316
+ "scalar": "dtypes.T",
1317
+ "useSubgroups": false,
1318
+ "workgroupSize": "skipWg",
1319
+ "HIDDEN_LEN": "hiddenSize",
1320
+ "packedStatistics": true
1321
+ },
1322
+ "passes": [
1323
+ {
1324
+ "id": "main",
1325
+ "shader": "norm-skip-row.wgsl.jinja",
1326
+ "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "row_stats", "params__uniform"],
1327
+ "dispatch": {
1328
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1329
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1330
+ "z": 1
1331
+ }
1332
+ },
1333
+ {
1334
+ "id": "statistics",
1335
+ "shader": "norm-stats-copy.wgsl.jinja",
1336
+ "derive": { "workgroupSize": 64 },
1337
+ "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"],
1338
+ "dispatch": {
1339
+ "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)",
1340
+ "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)",
1341
+ "z": 1
1342
+ }
1343
+ }
1344
+ ],
1345
+ "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }]
1346
+ },
1347
+ {
1348
+ "id": "stats_both_beta_residual",
1349
+ "when": ["statsContract", "present.meanT", "present.invStdT", "not present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
1350
+ "derive": {
1351
+ "simplified": false,
1352
+ "hasBias": "present.biasT",
1353
+ "hasBeta": "present.betaT",
1354
+ "writeResidualSum": "present.residualT",
1355
+ "writeMean": "present.meanT",
1356
+ "writeInvStd": "present.invStdT",
1357
+ "scalar": "dtypes.T",
1358
+ "useSubgroups": false,
1359
+ "workgroupSize": "skipWg",
1360
+ "HIDDEN_LEN": "hiddenSize",
1361
+ "packedStatistics": true
1362
+ },
1363
+ "passes": [
1364
+ {
1365
+ "id": "main",
1366
+ "shader": "norm-skip-row.wgsl.jinja",
1367
+ "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "input_skip_bias_sum_main", "output_main", "row_stats", "params__uniform"],
1368
+ "dispatch": {
1369
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1370
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1371
+ "z": 1
1372
+ }
1373
+ },
1374
+ {
1375
+ "id": "statistics",
1376
+ "shader": "norm-stats-copy.wgsl.jinja",
1377
+ "derive": { "workgroupSize": 64 },
1378
+ "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"],
1379
+ "dispatch": {
1380
+ "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)",
1381
+ "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)",
1382
+ "z": 1
1383
+ }
1384
+ }
1385
+ ],
1386
+ "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }]
1387
+ },
1388
+ {
1389
+ "id": "stats_both_bias",
1390
+ "when": ["statsContract", "present.meanT", "present.invStdT", "present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
1391
+ "derive": {
1392
+ "simplified": false,
1393
+ "hasBias": "present.biasT",
1394
+ "hasBeta": "present.betaT",
1395
+ "writeResidualSum": "present.residualT",
1396
+ "writeMean": "present.meanT",
1397
+ "writeInvStd": "present.invStdT",
1398
+ "scalar": "dtypes.T",
1399
+ "useSubgroups": false,
1400
+ "workgroupSize": "skipWg",
1401
+ "HIDDEN_LEN": "hiddenSize",
1402
+ "packedStatistics": true
1403
+ },
1404
+ "passes": [
1405
+ {
1406
+ "id": "main",
1407
+ "shader": "norm-skip-row.wgsl.jinja",
1408
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "row_stats", "params__uniform"],
1409
+ "dispatch": {
1410
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1411
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1412
+ "z": 1
1413
+ }
1414
+ },
1415
+ {
1416
+ "id": "statistics",
1417
+ "shader": "norm-stats-copy.wgsl.jinja",
1418
+ "derive": { "workgroupSize": 64 },
1419
+ "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"],
1420
+ "dispatch": {
1421
+ "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)",
1422
+ "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)",
1423
+ "z": 1
1424
+ }
1425
+ }
1426
+ ],
1427
+ "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }]
1428
+ },
1429
+ {
1430
+ "id": "stats_both_bias_residual",
1431
+ "when": ["statsContract", "present.meanT", "present.invStdT", "present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
1432
+ "derive": {
1433
+ "simplified": false,
1434
+ "hasBias": "present.biasT",
1435
+ "hasBeta": "present.betaT",
1436
+ "writeResidualSum": "present.residualT",
1437
+ "writeMean": "present.meanT",
1438
+ "writeInvStd": "present.invStdT",
1439
+ "scalar": "dtypes.T",
1440
+ "useSubgroups": false,
1441
+ "workgroupSize": "skipWg",
1442
+ "HIDDEN_LEN": "hiddenSize",
1443
+ "packedStatistics": true
1444
+ },
1445
+ "passes": [
1446
+ {
1447
+ "id": "main",
1448
+ "shader": "norm-skip-row.wgsl.jinja",
1449
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "input_skip_bias_sum_main", "output_main", "row_stats", "params__uniform"],
1450
+ "dispatch": {
1451
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1452
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1453
+ "z": 1
1454
+ }
1455
+ },
1456
+ {
1457
+ "id": "statistics",
1458
+ "shader": "norm-stats-copy.wgsl.jinja",
1459
+ "derive": { "workgroupSize": 64 },
1460
+ "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"],
1461
+ "dispatch": {
1462
+ "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)",
1463
+ "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)",
1464
+ "z": 1
1465
+ }
1466
+ }
1467
+ ],
1468
+ "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }]
1469
+ },
1470
+ {
1471
+ "id": "stats_both_bias_beta",
1472
+ "when": ["statsContract", "present.meanT", "present.invStdT", "present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"],
1473
+ "derive": {
1474
+ "simplified": false,
1475
+ "hasBias": "present.biasT",
1476
+ "hasBeta": "present.betaT",
1477
+ "writeResidualSum": "present.residualT",
1478
+ "writeMean": "present.meanT",
1479
+ "writeInvStd": "present.invStdT",
1480
+ "scalar": "dtypes.T",
1481
+ "useSubgroups": false,
1482
+ "workgroupSize": "skipWg",
1483
+ "HIDDEN_LEN": "hiddenSize",
1484
+ "packedStatistics": true
1485
+ },
1486
+ "passes": [
1487
+ {
1488
+ "id": "main",
1489
+ "shader": "norm-skip-row.wgsl.jinja",
1490
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "output_main", "row_stats", "params__uniform"],
1491
+ "dispatch": {
1492
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1493
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1494
+ "z": 1
1495
+ }
1496
+ },
1497
+ {
1498
+ "id": "statistics",
1499
+ "shader": "norm-stats-copy.wgsl.jinja",
1500
+ "derive": { "workgroupSize": 64 },
1501
+ "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"],
1502
+ "dispatch": {
1503
+ "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)",
1504
+ "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)",
1505
+ "z": 1
1506
+ }
1507
+ }
1508
+ ],
1509
+ "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }]
1510
+ },
1511
+ {
1512
+ "id": "stats_both_bias_beta_residual",
1513
+ "when": ["statsContract", "present.meanT", "present.invStdT", "present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"],
1514
+ "derive": {
1515
+ "simplified": false,
1516
+ "hasBias": "present.biasT",
1517
+ "hasBeta": "present.betaT",
1518
+ "writeResidualSum": "present.residualT",
1519
+ "writeMean": "present.meanT",
1520
+ "writeInvStd": "present.invStdT",
1521
+ "scalar": "dtypes.T",
1522
+ "useSubgroups": false,
1523
+ "workgroupSize": "skipWg",
1524
+ "HIDDEN_LEN": "hiddenSize",
1525
+ "packedStatistics": true
1526
+ },
1527
+ "passes": [
1528
+ {
1529
+ "id": "main",
1530
+ "shader": "norm-skip-row.wgsl.jinja",
1531
+ "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "input_skip_bias_sum_main", "output_main", "row_stats", "params__uniform"],
1532
+ "dispatch": {
1533
+ "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1534
+ "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)",
1535
+ "z": 1
1536
+ }
1537
+ },
1538
+ {
1539
+ "id": "statistics",
1540
+ "shader": "norm-stats-copy.wgsl.jinja",
1541
+ "derive": { "workgroupSize": 64 },
1542
+ "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"],
1543
+ "dispatch": {
1544
+ "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)",
1545
+ "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)",
1546
+ "z": 1
1547
+ }
1548
+ }
1549
+ ],
1550
+ "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }]
1551
  }
1552
  ]
1553
  }
build/webgpu/metadata.json CHANGED
@@ -1,41 +1,70 @@
1
  {
2
  "name": "com.microsoft.SkipLayerNormalization",
3
- "id": "_com_microsoft_skiplayernormalization_webgpu_28a934b",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "dZADuGnHvG3kWGPS0BGgmikuwdi16tD/Sk4HbGC9UV0=",
11
- "manifest.json": "G6ZofqPdSuhkAxEuuKIlW4fQ/fCsDwxEzafKwVqGEY8=",
12
- "norm-skip-row-vec4.wgsl.jinja": "bH8L9BYJ/3XQXE2zP8kvZc4wlhAbcAolI4/p0xXskzM=",
13
- "norm-skip-row.wgsl.jinja": "aERQfDDXx6lwKuyahwcemgFllRSq46egf2++7f6bNtA=",
14
- "test.json": "K+KtXhWTSxbRlFwFfszdd1YKz9JnAOuDZ9ebLg7Y8w0="
 
15
  }
16
  },
17
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
18
  "webgpu": {
19
- "manifestSpec": "2.0",
20
  "variants": {
 
 
 
 
21
  "beta_output_only_vec4_broadcast": ["norm-skip-row-vec4.wgsl.jinja"],
22
  "beta_bias_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
23
  "beta_bias_row": ["norm-skip-row.wgsl.jinja"],
 
 
24
  "beta_bias_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
25
  "no_beta_output_only_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
26
  "no_beta_output_only_row": ["norm-skip-row.wgsl.jinja"],
27
  "no_beta_output_only_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
28
  "no_beta_output_only_row_f16": ["norm-skip-row.wgsl.jinja"],
29
- "beta_no_bias_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
30
- "beta_no_bias_row": ["norm-skip-row.wgsl.jinja"],
31
  "beta_no_bias_output_only_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
32
  "beta_no_bias_output_only_row": ["norm-skip-row.wgsl.jinja"],
33
- "beta_bias_output_only_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
34
- "beta_bias_output_only_row": ["norm-skip-row.wgsl.jinja"],
35
  "beta_no_bias_output_only_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
36
  "beta_no_bias_output_only_row_f16": ["norm-skip-row.wgsl.jinja"],
 
 
37
  "beta_bias_output_only_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
38
- "beta_bias_output_only_row_f16": ["norm-skip-row.wgsl.jinja"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
39
  }
40
  }
41
  }
 
1
  {
2
  "name": "com.microsoft.SkipLayerNormalization",
3
+ "id": "_com_microsoft_skiplayernormalization_webgpu_a30f1f2",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "bench.json": "cVSe9BmYyynkTmGaCfti3+3as2tY8s7SiHcDUGZXadQ=",
11
+ "manifest.json": "eBeukaRPaePto8cDuxf/pVi95UUFHU+jR4Tcc4/FFh4=",
12
+ "norm-skip-row-vec4.wgsl.jinja": "hpoa1+W36ADP2AlqvZAoWFHCZPTBrUIlgVIQXE98oTE=",
13
+ "norm-skip-row.wgsl.jinja": "LOw1CZFwolmaeaXpWSTi8ToiXzTfvJDoaapF2N6a0x4=",
14
+ "norm-stats-copy.wgsl.jinja": "4PNrRNFbMWP9csH7RrxSvJuAGSn5aDfZZPw/IZEVmjU=",
15
+ "test.json": "upshAVlCpYOJn3G1M4TGhznjHzx8ybo8uae8zkB932s="
16
  }
17
  },
18
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
19
  "webgpu": {
20
+ "manifestSpec": "2.1",
21
  "variants": {
22
+ "hidden1_f32_no_beta": ["norm-skip-row.wgsl.jinja"],
23
+ "hidden1_f32_beta": ["norm-skip-row.wgsl.jinja"],
24
+ "hidden1_f16_no_beta": ["norm-skip-row.wgsl.jinja"],
25
+ "hidden1_f16_beta": ["norm-skip-row.wgsl.jinja"],
26
  "beta_output_only_vec4_broadcast": ["norm-skip-row-vec4.wgsl.jinja"],
27
  "beta_bias_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
28
  "beta_bias_row": ["norm-skip-row.wgsl.jinja"],
29
+ "beta_no_bias_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
30
+ "beta_no_bias_row": ["norm-skip-row.wgsl.jinja"],
31
  "beta_bias_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
32
  "no_beta_output_only_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
33
  "no_beta_output_only_row": ["norm-skip-row.wgsl.jinja"],
34
  "no_beta_output_only_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
35
  "no_beta_output_only_row_f16": ["norm-skip-row.wgsl.jinja"],
 
 
36
  "beta_no_bias_output_only_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
37
  "beta_no_bias_output_only_row": ["norm-skip-row.wgsl.jinja"],
 
 
38
  "beta_no_bias_output_only_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
39
  "beta_no_bias_output_only_row_f16": ["norm-skip-row.wgsl.jinja"],
40
+ "beta_bias_output_only_vec4": ["norm-skip-row-vec4.wgsl.jinja"],
41
+ "beta_bias_output_only_row": ["norm-skip-row.wgsl.jinja"],
42
  "beta_bias_output_only_vec4_f16": ["norm-skip-row-vec4.wgsl.jinja"],
43
+ "beta_bias_output_only_row_f16": ["norm-skip-row.wgsl.jinja"],
44
+ "stats_mean_plain": ["norm-skip-row.wgsl.jinja"],
45
+ "stats_mean_residual": ["norm-skip-row.wgsl.jinja"],
46
+ "stats_mean_beta": ["norm-skip-row.wgsl.jinja"],
47
+ "stats_mean_beta_residual": ["norm-skip-row.wgsl.jinja"],
48
+ "stats_mean_bias": ["norm-skip-row.wgsl.jinja"],
49
+ "stats_mean_bias_residual": ["norm-skip-row.wgsl.jinja"],
50
+ "stats_mean_bias_beta": ["norm-skip-row.wgsl.jinja"],
51
+ "stats_mean_bias_beta_residual": ["norm-skip-row.wgsl.jinja"],
52
+ "stats_inv_plain": ["norm-skip-row.wgsl.jinja"],
53
+ "stats_inv_residual": ["norm-skip-row.wgsl.jinja"],
54
+ "stats_inv_beta": ["norm-skip-row.wgsl.jinja"],
55
+ "stats_inv_beta_residual": ["norm-skip-row.wgsl.jinja"],
56
+ "stats_inv_bias": ["norm-skip-row.wgsl.jinja"],
57
+ "stats_inv_bias_residual": ["norm-skip-row.wgsl.jinja"],
58
+ "stats_inv_bias_beta": ["norm-skip-row.wgsl.jinja"],
59
+ "stats_inv_bias_beta_residual": ["norm-skip-row.wgsl.jinja"],
60
+ "stats_both_plain": ["norm-skip-row.wgsl.jinja", "norm-stats-copy.wgsl.jinja"],
61
+ "stats_both_residual": ["norm-skip-row.wgsl.jinja", "norm-stats-copy.wgsl.jinja"],
62
+ "stats_both_beta": ["norm-skip-row.wgsl.jinja", "norm-stats-copy.wgsl.jinja"],
63
+ "stats_both_beta_residual": ["norm-skip-row.wgsl.jinja", "norm-stats-copy.wgsl.jinja"],
64
+ "stats_both_bias": ["norm-skip-row.wgsl.jinja", "norm-stats-copy.wgsl.jinja"],
65
+ "stats_both_bias_residual": ["norm-skip-row.wgsl.jinja", "norm-stats-copy.wgsl.jinja"],
66
+ "stats_both_bias_beta": ["norm-skip-row.wgsl.jinja", "norm-stats-copy.wgsl.jinja"],
67
+ "stats_both_bias_beta_residual": ["norm-skip-row.wgsl.jinja", "norm-stats-copy.wgsl.jinja"]
68
  }
69
  }
70
  }
build/webgpu/norm-skip-row-vec4.wgsl.jinja CHANGED
@@ -1,52 +1,20 @@
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 %}
 
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
+ {% set broadcastSkip = broadcastSkip is defined and broadcastSkip %}
 
 
 
 
 
 
 
 
 
 
 
 
18
  {% if useSubgroups %}
19
  enable subgroups;
20
  {% endif %}
build/webgpu/norm-skip-row.wgsl.jinja CHANGED
@@ -1,61 +1,30 @@
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 %}
46
-
47
- /* One workgroup normalizes each row of residual = input + skip, with an
48
- * optional bias. */
49
  {% set degenerateRow = (not simplified) and hiddenSize == 1 %}
50
- {% if usesF16 %}
51
- enable f16;
52
- {% endif %}
53
  {% if useSubgroups and not degenerateRow %}
54
  enable subgroups;
55
  {% endif %}
56
  {{ env.wgsl.resourceDeclarations }}
57
 
58
- {% if not degenerateRow or writeResidualSum %}
59
  const HIDDEN: u32 = {{ hiddenSize }}u;
60
  {% endif %}
61
  const WG: u32 = {{ workgroupSize }}u;
@@ -88,8 +57,8 @@ fn reduce_pair(value: vec2<f32>, tid: u32) -> vec2<f32> {
88
  }
89
  {% endif %}
90
  {% endif %}
 
91
 
92
- {% if not degenerateRow or writeResidualSum %}
93
  fn residual_value(row: u32, d: u32) -> f32 {
94
  let index = row * HIDDEN + d;
95
  var value = f32(input[index]) + f32(skip[index]);
@@ -102,16 +71,19 @@ fn residual_value(row: u32, d: u32) -> f32 {
102
 
103
  @compute @workgroup_size(WG, 1, 1)
104
  fn main(
105
- @builtin(workgroup_id) wg: vec3<u32>{% if not degenerateRow %},
106
  @builtin(local_invocation_id) lid: vec3<u32>{% endif %}{% if useSubgroups and not degenerateRow %},
107
  @builtin(subgroup_invocation_id) sg_lane: u32,
108
  @builtin(subgroup_id) sg_id: u32,
109
  @builtin(num_subgroups) num_sg: u32{% endif %}
110
  ) {
111
- // 2D-folded row index: wg.y carries the high bits past the per-axis dispatch fold width.
112
- // Reduces to wg.x when the dispatch does not fold;
113
- // the row >= params.rows guard drops the over-dispatched tail.
 
 
114
  let row = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u;
 
115
  if (row >= params.rows) {
116
  return;
117
  }
@@ -124,6 +96,14 @@ fn main(
124
  // the variance are exactly zero and the output reduces to beta. The closed
125
  // form avoids computing that zero by subtracting two equal rounded values.
126
  let row_inv = inverseSqrt(params.epsilon);
 
 
 
 
 
 
 
 
127
  {% if writeResidualSum %}
128
  let residual = residual_value(row, 0u);
129
  input_skip_bias_sum[row] = {{ scalar }}(residual);
@@ -152,6 +132,14 @@ fn main(
152
  let row_mean = shift + mean_delta;
153
  let variance = max(totals.y / f32(HIDDEN) - mean_delta * mean_delta, 0.0);
154
  let row_inv = inverseSqrt(variance + params.epsilon);
 
 
 
 
 
 
 
 
155
  for (var d = tid; d < HIDDEN; d = d + WG) {
156
  let index = row * HIDDEN + d;
157
  let residual = residual_value(row, d);
 
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) {
9
  break;
10
  }
 
 
 
 
 
11
  if ({{ idx }} < {{ svar }}) {
12
  {% for a in arrays %}
13
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
14
  {% endfor %}
15
  }
 
 
 
 
16
  {{ svar }} = {{ svar }} / 2u;
 
 
 
 
 
17
  workgroupBarrier();
18
+ }{% endmacro %}
19
+ /* Normalize residual = input + skip, with an optional bias. Reductions use
20
+ * one workgroup per row; closed-form one-element rows use one invocation. */
 
 
 
 
 
 
 
21
  {% set degenerateRow = (not simplified) and hiddenSize == 1 %}
 
 
 
22
  {% if useSubgroups and not degenerateRow %}
23
  enable subgroups;
24
  {% endif %}
25
  {{ env.wgsl.resourceDeclarations }}
26
 
27
+ {% if not degenerateRow or writeResidualSum or (writeMean is defined and writeMean) %}
28
  const HIDDEN: u32 = {{ hiddenSize }}u;
29
  {% endif %}
30
  const WG: u32 = {{ workgroupSize }}u;
 
57
  }
58
  {% endif %}
59
  {% endif %}
60
+ {% if not degenerateRow or writeResidualSum or (writeMean is defined and writeMean) %}
61
 
 
62
  fn residual_value(row: u32, d: u32) -> f32 {
63
  let index = row * HIDDEN + d;
64
  var value = f32(input[index]) + f32(skip[index]);
 
71
 
72
  @compute @workgroup_size(WG, 1, 1)
73
  fn main(
74
+ @builtin({{ "global_invocation_id" if degenerateRow else "workgroup_id" }}) {{ "gid" if degenerateRow else "wg" }}: vec3<u32>{% if not degenerateRow %},
75
  @builtin(local_invocation_id) lid: vec3<u32>{% endif %}{% if useSubgroups and not degenerateRow %},
76
  @builtin(subgroup_invocation_id) sg_lane: u32,
77
  @builtin(subgroup_id) sg_id: u32,
78
  @builtin(num_subgroups) num_sg: u32{% endif %}
79
  ) {
80
+ // Fold the row grid across workgroups; independent rows also include the
81
+ // invocation offset. The bounds guard drops the final dispatch tail.
82
+ {% if degenerateRow %}
83
+ let row = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
84
+ {% else %}
85
  let row = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u;
86
+ {% endif %}
87
  if (row >= params.rows) {
88
  return;
89
  }
 
96
  // the variance are exactly zero and the output reduces to beta. The closed
97
  // form avoids computing that zero by subtracting two equal rounded values.
98
  let row_inv = inverseSqrt(params.epsilon);
99
+ {% if packedStatistics is defined and packedStatistics %}
100
+ row_stats[row] = vec2<f32>(residual_value(row, 0u), row_inv);
101
+ {% elif writeMean is defined and writeMean %}
102
+ mean[row] = residual_value(row, 0u);
103
+ {% endif %}
104
+ {% if writeInvStd is defined and writeInvStd and not (packedStatistics is defined and packedStatistics) %}
105
+ inv_std_var[row] = row_inv;
106
+ {% endif %}
107
  {% if writeResidualSum %}
108
  let residual = residual_value(row, 0u);
109
  input_skip_bias_sum[row] = {{ scalar }}(residual);
 
132
  let row_mean = shift + mean_delta;
133
  let variance = max(totals.y / f32(HIDDEN) - mean_delta * mean_delta, 0.0);
134
  let row_inv = inverseSqrt(variance + params.epsilon);
135
+ {% if packedStatistics is defined and packedStatistics %}
136
+ if (tid == 0u) { row_stats[row] = vec2<f32>(row_mean, row_inv); }
137
+ {% elif writeMean is defined and writeMean %}
138
+ if (tid == 0u) { mean[row] = row_mean; }
139
+ {% endif %}
140
+ {% if writeInvStd is defined and writeInvStd and not (packedStatistics is defined and packedStatistics) %}
141
+ if (tid == 0u) { inv_std_var[row] = row_inv; }
142
+ {% endif %}
143
  for (var d = tid; d < HIDDEN; d = d + WG) {
144
  let index = row * HIDDEN + d;
145
  let residual = residual_value(row, d);
build/webgpu/norm-stats-copy.wgsl.jinja ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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 }}) { return; }{% endmacro %}
7
+ {{ env.wgsl.resourceDeclarations }}
8
+ // Separate optional outputs keep the reduction within eight storage bindings.
9
+ @compute @workgroup_size({{ workgroupSize }})
10
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
11
+ {{ flat_index_2d(workgroupSize, "row", "params.rows", guardInline=true) }}
12
+ let stats = row_stats[row];
13
+ mean[row] = stats.x;
14
+ inv_std_var[row] = stats.y;
15
+ }
build/webgpu/test.json CHANGED
@@ -476,7 +476,7 @@
476
  {
477
  "name": "no_beta_output_only_hidden6_unaligned_row",
478
  "provenance": {
479
- "notes": "Beta and optional outputs are omitted, and hidden size 6 is not divisible by four, selecting the scalar no-beta output-only row path on every tier."
480
  },
481
  "attrs": { "epsilon": 0.00001 },
482
  "inputs": {
@@ -531,7 +531,7 @@
531
  {
532
  "name": "beta_no_bias_output_only_hidden6_unaligned_row",
533
  "provenance": {
534
- "notes": "Beta is present, bias and optional outputs are omitted, and hidden size 6 selects the scalar output-only row path on every tier."
535
  },
536
  "attrs": { "epsilon": 0.00001 },
537
  "inputs": {
@@ -1172,6 +1172,754 @@
1172
  }
1173
  },
1174
  "outputs": { "outputT": { "dtype": "float16", "shape": [3, 8], "tolerance": 0.002 } }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1175
  }
1176
  ]
1177
  }
 
476
  {
477
  "name": "no_beta_output_only_hidden6_unaligned_row",
478
  "provenance": {
479
+ "notes": "Beta and optional outputs are omitted; hidden size six exercises a non-four-aligned row."
480
  },
481
  "attrs": { "epsilon": 0.00001 },
482
  "inputs": {
 
531
  {
532
  "name": "beta_no_bias_output_only_hidden6_unaligned_row",
533
  "provenance": {
534
+ "notes": "Beta is present, bias and optional outputs are omitted, and hidden size 6 leaves a partial four-element storage group."
535
  },
536
  "attrs": { "epsilon": 0.00001 },
537
  "inputs": {
 
1172
  }
1173
  },
1174
  "outputs": { "outputT": { "dtype": "float16", "shape": [3, 8], "tolerance": 0.002 } }
1175
+ },
1176
+ {
1177
+ "name": "independent_rows_float32_257x1_option0_wgdefault_epsilon0.00001",
1178
+ "provenance": {
1179
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1180
+ },
1181
+ "attrs": { "epsilon": 0.00001 },
1182
+ "inputs": {
1183
+ "inputT": { "dtype": "float32", "shape": [257, 1], "data": { "kind": "constant", "value": 1.25 } },
1184
+ "skipT": { "dtype": "float32", "shape": [257, 1], "data": { "kind": "constant", "value": 0.5 } },
1185
+ "gammaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [1.125] } }
1186
+ },
1187
+ "outputs": { "outputT": { "dtype": "float32", "shape": [257, 1], "tolerance": 0, "allowNaN": false } }
1188
+ },
1189
+ {
1190
+ "name": "independent_rows_float32_257x1_option1_wgdefault_epsilon0.00001",
1191
+ "provenance": {
1192
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1193
+ },
1194
+ "attrs": { "epsilon": 0.00001 },
1195
+ "inputs": {
1196
+ "inputT": { "dtype": "float32", "shape": [257, 1], "data": { "kind": "constant", "value": 1.25 } },
1197
+ "skipT": { "dtype": "float32", "shape": [257, 1], "data": { "kind": "constant", "value": 0.5 } },
1198
+ "gammaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1199
+ "betaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [-0.25] } }
1200
+ },
1201
+ "outputs": { "outputT": { "dtype": "float32", "shape": [257, 1], "tolerance": 0, "allowNaN": false } }
1202
+ },
1203
+ {
1204
+ "name": "independent_rows_float32_257x1_option2_wgdefault_epsilon0.00001",
1205
+ "provenance": {
1206
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1207
+ },
1208
+ "attrs": { "epsilon": 0.00001 },
1209
+ "inputs": {
1210
+ "inputT": { "dtype": "float32", "shape": [257, 1], "data": { "kind": "constant", "value": 1.25 } },
1211
+ "skipT": { "dtype": "float32", "shape": [257, 1], "data": { "kind": "constant", "value": 0.5 } },
1212
+ "gammaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1213
+ "betaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1214
+ "biasT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1215
+ },
1216
+ "outputs": { "outputT": { "dtype": "float32", "shape": [257, 1], "tolerance": 0, "allowNaN": false } }
1217
+ },
1218
+ {
1219
+ "name": "independent_rows_float32_3x31x1_option2_wg1_epsilon0.00001",
1220
+ "provenance": {
1221
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1222
+ },
1223
+ "attrs": { "epsilon": 0.00001 },
1224
+ "tunables": { "MAX_WORKGROUP_SIZE": 1 },
1225
+ "inputs": {
1226
+ "inputT": { "dtype": "float32", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 1.25 } },
1227
+ "skipT": { "dtype": "float32", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 0.5 } },
1228
+ "gammaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1229
+ "betaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1230
+ "biasT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1231
+ },
1232
+ "outputs": { "outputT": { "dtype": "float32", "shape": [3, 31, 1], "tolerance": 0, "allowNaN": false } }
1233
+ },
1234
+ {
1235
+ "name": "independent_rows_float32_3x31x1_option2_wg8_epsilon0.00001",
1236
+ "provenance": {
1237
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1238
+ },
1239
+ "attrs": { "epsilon": 0.00001 },
1240
+ "tunables": { "MAX_WORKGROUP_SIZE": 8 },
1241
+ "inputs": {
1242
+ "inputT": { "dtype": "float32", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 1.25 } },
1243
+ "skipT": { "dtype": "float32", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 0.5 } },
1244
+ "gammaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1245
+ "betaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1246
+ "biasT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1247
+ },
1248
+ "outputs": { "outputT": { "dtype": "float32", "shape": [3, 31, 1], "tolerance": 0, "allowNaN": false } }
1249
+ },
1250
+ {
1251
+ "name": "independent_rows_float32_3x31x1_option2_wg64_epsilon0.00001",
1252
+ "provenance": {
1253
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1254
+ },
1255
+ "attrs": { "epsilon": 0.00001 },
1256
+ "tunables": { "MAX_WORKGROUP_SIZE": 64 },
1257
+ "inputs": {
1258
+ "inputT": { "dtype": "float32", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 1.25 } },
1259
+ "skipT": { "dtype": "float32", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 0.5 } },
1260
+ "gammaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1261
+ "betaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1262
+ "biasT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1263
+ },
1264
+ "outputs": { "outputT": { "dtype": "float32", "shape": [3, 31, 1], "tolerance": 0, "allowNaN": false } }
1265
+ },
1266
+ {
1267
+ "name": "independent_rows_float32_3x31x1_option2_wg128_epsilon0.00001",
1268
+ "provenance": {
1269
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1270
+ },
1271
+ "attrs": { "epsilon": 0.00001 },
1272
+ "tunables": { "MAX_WORKGROUP_SIZE": 128 },
1273
+ "inputs": {
1274
+ "inputT": { "dtype": "float32", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 1.25 } },
1275
+ "skipT": { "dtype": "float32", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 0.5 } },
1276
+ "gammaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1277
+ "betaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1278
+ "biasT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1279
+ },
1280
+ "outputs": { "outputT": { "dtype": "float32", "shape": [3, 31, 1], "tolerance": 0, "allowNaN": false } }
1281
+ },
1282
+ {
1283
+ "name": "independent_rows_float32_524289x1_option3_wg8_epsilon0.00001",
1284
+ "provenance": {
1285
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1286
+ },
1287
+ "attrs": { "epsilon": 0.00001 },
1288
+ "tunables": { "MAX_WORKGROUP_SIZE": 8 },
1289
+ "inputs": {
1290
+ "inputT": { "dtype": "float32", "shape": [524289, 1], "data": { "kind": "constant", "value": 1.25 } },
1291
+ "skipT": { "dtype": "float32", "shape": [524289, 1], "data": { "kind": "constant", "value": 0.5 } },
1292
+ "gammaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1293
+ "betaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1294
+ "biasT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1295
+ },
1296
+ "outputs": {
1297
+ "outputT": { "dtype": "float32", "shape": [524289, 1], "tolerance": 0, "allowNaN": false },
1298
+ "residualT": { "dtype": "float32", "shape": [524289, 1], "tolerance": 0 }
1299
+ }
1300
+ },
1301
+ {
1302
+ "name": "independent_rows_float32_257x1_option2_wgdefault_epsilon0",
1303
+ "provenance": {
1304
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1305
+ },
1306
+ "attrs": { "epsilon": 0 },
1307
+ "inputs": {
1308
+ "inputT": { "dtype": "float32", "shape": [257, 1], "data": { "kind": "constant", "value": 1.25 } },
1309
+ "skipT": { "dtype": "float32", "shape": [257, 1], "data": { "kind": "constant", "value": 0.5 } },
1310
+ "gammaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1311
+ "betaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1312
+ "biasT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1313
+ },
1314
+ "outputs": { "outputT": { "dtype": "float32", "shape": [257, 1], "tolerance": 0, "allowNaN": true } }
1315
+ },
1316
+ {
1317
+ "name": "independent_rows_float16_257x1_option0_wgdefault_epsilon0.00001",
1318
+ "provenance": {
1319
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1320
+ },
1321
+ "attrs": { "epsilon": 0.00001 },
1322
+ "inputs": {
1323
+ "inputT": { "dtype": "float16", "shape": [257, 1], "data": { "kind": "constant", "value": 1.25 } },
1324
+ "skipT": { "dtype": "float16", "shape": [257, 1], "data": { "kind": "constant", "value": 0.5 } },
1325
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [1.125] } }
1326
+ },
1327
+ "outputs": { "outputT": { "dtype": "float16", "shape": [257, 1], "tolerance": 0, "allowNaN": false } }
1328
+ },
1329
+ {
1330
+ "name": "independent_rows_float16_257x1_option1_wgdefault_epsilon0.00001",
1331
+ "provenance": {
1332
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1333
+ },
1334
+ "attrs": { "epsilon": 0.00001 },
1335
+ "inputs": {
1336
+ "inputT": { "dtype": "float16", "shape": [257, 1], "data": { "kind": "constant", "value": 1.25 } },
1337
+ "skipT": { "dtype": "float16", "shape": [257, 1], "data": { "kind": "constant", "value": 0.5 } },
1338
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1339
+ "betaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [-0.25] } }
1340
+ },
1341
+ "outputs": { "outputT": { "dtype": "float16", "shape": [257, 1], "tolerance": 0, "allowNaN": false } }
1342
+ },
1343
+ {
1344
+ "name": "independent_rows_float16_257x1_option2_wgdefault_epsilon0.00001",
1345
+ "provenance": {
1346
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1347
+ },
1348
+ "attrs": { "epsilon": 0.00001 },
1349
+ "inputs": {
1350
+ "inputT": { "dtype": "float16", "shape": [257, 1], "data": { "kind": "constant", "value": 1.25 } },
1351
+ "skipT": { "dtype": "float16", "shape": [257, 1], "data": { "kind": "constant", "value": 0.5 } },
1352
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1353
+ "betaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1354
+ "biasT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1355
+ },
1356
+ "outputs": { "outputT": { "dtype": "float16", "shape": [257, 1], "tolerance": 0, "allowNaN": false } }
1357
+ },
1358
+ {
1359
+ "name": "independent_rows_float16_3x31x1_option2_wg1_epsilon0.00001",
1360
+ "provenance": {
1361
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1362
+ },
1363
+ "attrs": { "epsilon": 0.00001 },
1364
+ "tunables": { "MAX_WORKGROUP_SIZE": 1 },
1365
+ "inputs": {
1366
+ "inputT": { "dtype": "float16", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 1.25 } },
1367
+ "skipT": { "dtype": "float16", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 0.5 } },
1368
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1369
+ "betaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1370
+ "biasT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1371
+ },
1372
+ "outputs": { "outputT": { "dtype": "float16", "shape": [3, 31, 1], "tolerance": 0, "allowNaN": false } }
1373
+ },
1374
+ {
1375
+ "name": "independent_rows_float16_3x31x1_option2_wg8_epsilon0.00001",
1376
+ "provenance": {
1377
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1378
+ },
1379
+ "attrs": { "epsilon": 0.00001 },
1380
+ "tunables": { "MAX_WORKGROUP_SIZE": 8 },
1381
+ "inputs": {
1382
+ "inputT": { "dtype": "float16", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 1.25 } },
1383
+ "skipT": { "dtype": "float16", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 0.5 } },
1384
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1385
+ "betaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1386
+ "biasT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1387
+ },
1388
+ "outputs": { "outputT": { "dtype": "float16", "shape": [3, 31, 1], "tolerance": 0, "allowNaN": false } }
1389
+ },
1390
+ {
1391
+ "name": "independent_rows_float16_3x31x1_option2_wg64_epsilon0.00001",
1392
+ "provenance": {
1393
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1394
+ },
1395
+ "attrs": { "epsilon": 0.00001 },
1396
+ "tunables": { "MAX_WORKGROUP_SIZE": 64 },
1397
+ "inputs": {
1398
+ "inputT": { "dtype": "float16", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 1.25 } },
1399
+ "skipT": { "dtype": "float16", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 0.5 } },
1400
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1401
+ "betaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1402
+ "biasT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1403
+ },
1404
+ "outputs": { "outputT": { "dtype": "float16", "shape": [3, 31, 1], "tolerance": 0, "allowNaN": false } }
1405
+ },
1406
+ {
1407
+ "name": "independent_rows_float16_3x31x1_option2_wg128_epsilon0.00001",
1408
+ "provenance": {
1409
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1410
+ },
1411
+ "attrs": { "epsilon": 0.00001 },
1412
+ "tunables": { "MAX_WORKGROUP_SIZE": 128 },
1413
+ "inputs": {
1414
+ "inputT": { "dtype": "float16", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 1.25 } },
1415
+ "skipT": { "dtype": "float16", "shape": [3, 31, 1], "data": { "kind": "constant", "value": 0.5 } },
1416
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1417
+ "betaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1418
+ "biasT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1419
+ },
1420
+ "outputs": { "outputT": { "dtype": "float16", "shape": [3, 31, 1], "tolerance": 0, "allowNaN": false } }
1421
+ },
1422
+ {
1423
+ "name": "independent_rows_float16_524289x1_option2_wg8_epsilon0.00001",
1424
+ "provenance": {
1425
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1426
+ },
1427
+ "attrs": { "epsilon": 0.00001 },
1428
+ "tunables": { "MAX_WORKGROUP_SIZE": 8 },
1429
+ "inputs": {
1430
+ "inputT": { "dtype": "float16", "shape": [524289, 1], "data": { "kind": "constant", "value": 1.25 } },
1431
+ "skipT": { "dtype": "float16", "shape": [524289, 1], "data": { "kind": "constant", "value": 0.5 } },
1432
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1433
+ "betaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1434
+ "biasT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1435
+ },
1436
+ "outputs": { "outputT": { "dtype": "float16", "shape": [524289, 1], "tolerance": 0, "allowNaN": false } }
1437
+ },
1438
+ {
1439
+ "name": "independent_rows_float16_257x1_option2_wgdefault_epsilon0",
1440
+ "provenance": {
1441
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1442
+ },
1443
+ "attrs": { "epsilon": 0 },
1444
+ "inputs": {
1445
+ "inputT": { "dtype": "float16", "shape": [257, 1], "data": { "kind": "constant", "value": 1.25 } },
1446
+ "skipT": { "dtype": "float16", "shape": [257, 1], "data": { "kind": "constant", "value": 0.5 } },
1447
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1448
+ "betaT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1449
+ "biasT": { "dtype": "float16", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1450
+ },
1451
+ "outputs": { "outputT": { "dtype": "float16", "shape": [257, 1], "tolerance": 0, "allowNaN": true } }
1452
+ },
1453
+ {
1454
+ "name": "independent_rows_float32_257x1_option3_wgdefault_epsilon0.00001",
1455
+ "provenance": {
1456
+ "notes": "One-element normalization has an intrinsically uniform result. Uniform inputs pin its closed-form value, workgroup tails, optional outputs, and epsilon behavior; the GPU row-coverage test checks poisoned outputs and varying residuals separately."
1457
+ },
1458
+ "attrs": { "epsilon": 0.00001 },
1459
+ "inputs": {
1460
+ "inputT": { "dtype": "float32", "shape": [257, 1], "data": { "kind": "constant", "value": 1.25 } },
1461
+ "skipT": { "dtype": "float32", "shape": [257, 1], "data": { "kind": "constant", "value": 0.5 } },
1462
+ "gammaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [1.125] } },
1463
+ "betaT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [-0.25] } },
1464
+ "biasT": { "dtype": "float32", "shape": [1], "data": { "kind": "values", "values": [0.125] } }
1465
+ },
1466
+ "outputs": {
1467
+ "outputT": { "dtype": "float32", "shape": [257, 1], "tolerance": 0, "allowNaN": false },
1468
+ "residualT": { "dtype": "float32", "shape": [257, 1], "tolerance": 0 }
1469
+ }
1470
+ },
1471
+ {
1472
+ "name": "ort_skip_layer_norm_large_magnitude_row",
1473
+ "attrs": { "epsilon": 1e-12 },
1474
+ "inputs": {
1475
+ "inputT": {
1476
+ "dtype": "float32",
1477
+ "shape": [1, 4],
1478
+ "data": { "kind": "values", "values": [10000.0, 10001.0, 9999.0, 10000.0] }
1479
+ },
1480
+ "skipT": { "dtype": "float32", "shape": [1, 4], "data": { "kind": "values", "values": [0.0, 0.0, 0.0, 0.0] } },
1481
+ "gammaT": { "dtype": "float32", "shape": [4], "data": { "kind": "values", "values": [1.0, 1.0, 1.0, 1.0] } },
1482
+ "betaT": { "dtype": "float32", "shape": [4], "data": { "kind": "values", "values": [0.0, 0.0, 0.0, 0.0] } }
1483
+ },
1484
+ "outputs": { "outputT": { "dtype": "float32", "shape": [1, 4], "tolerance": 0.0001 } }
1485
+ },
1486
+ {
1487
+ "name": "stats_mean_plain",
1488
+ "attrs": { "epsilon": 0.00001 },
1489
+ "inputs": {
1490
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1491
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1492
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
1493
+ },
1494
+ "outputs": {
1495
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1496
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1497
+ }
1498
+ },
1499
+ {
1500
+ "name": "stats_mean_residual",
1501
+ "attrs": { "epsilon": 0.00001 },
1502
+ "inputs": {
1503
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1504
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1505
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
1506
+ },
1507
+ "outputs": {
1508
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1509
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1510
+ "residualT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001 }
1511
+ }
1512
+ },
1513
+ {
1514
+ "name": "stats_mean_beta",
1515
+ "attrs": { "epsilon": 0.00001 },
1516
+ "inputs": {
1517
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1518
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1519
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1520
+ "betaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1521
+ },
1522
+ "outputs": {
1523
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1524
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1525
+ }
1526
+ },
1527
+ {
1528
+ "name": "stats_mean_beta_residual",
1529
+ "attrs": { "epsilon": 0.00001 },
1530
+ "inputs": {
1531
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1532
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1533
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1534
+ "betaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1535
+ },
1536
+ "outputs": {
1537
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1538
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1539
+ "residualT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001 }
1540
+ }
1541
+ },
1542
+ {
1543
+ "name": "stats_mean_bias",
1544
+ "attrs": { "epsilon": 0.00001 },
1545
+ "inputs": {
1546
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1547
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1548
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1549
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
1550
+ },
1551
+ "outputs": {
1552
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1553
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1554
+ }
1555
+ },
1556
+ {
1557
+ "name": "stats_mean_bias_residual",
1558
+ "attrs": { "epsilon": 0.00001 },
1559
+ "inputs": {
1560
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1561
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1562
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1563
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
1564
+ },
1565
+ "outputs": {
1566
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1567
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1568
+ "residualT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002 }
1569
+ }
1570
+ },
1571
+ {
1572
+ "name": "stats_mean_bias_beta",
1573
+ "attrs": { "epsilon": 0.00001 },
1574
+ "inputs": {
1575
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1576
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1577
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1578
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } },
1579
+ "betaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1580
+ },
1581
+ "outputs": {
1582
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1583
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1584
+ }
1585
+ },
1586
+ {
1587
+ "name": "stats_mean_bias_beta_residual",
1588
+ "attrs": { "epsilon": 0.00001 },
1589
+ "inputs": {
1590
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1591
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1592
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1593
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } },
1594
+ "betaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1595
+ },
1596
+ "outputs": {
1597
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1598
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1599
+ "residualT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002 }
1600
+ }
1601
+ },
1602
+ {
1603
+ "name": "stats_inv_plain",
1604
+ "attrs": { "epsilon": 0.00001 },
1605
+ "inputs": {
1606
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1607
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1608
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
1609
+ },
1610
+ "outputs": {
1611
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1612
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1613
+ }
1614
+ },
1615
+ {
1616
+ "name": "stats_inv_residual",
1617
+ "attrs": { "epsilon": 0.00001 },
1618
+ "inputs": {
1619
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1620
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1621
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
1622
+ },
1623
+ "outputs": {
1624
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1625
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1626
+ "residualT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001 }
1627
+ }
1628
+ },
1629
+ {
1630
+ "name": "stats_inv_beta",
1631
+ "attrs": { "epsilon": 0.00001 },
1632
+ "inputs": {
1633
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1634
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1635
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1636
+ "betaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1637
+ },
1638
+ "outputs": {
1639
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1640
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1641
+ }
1642
+ },
1643
+ {
1644
+ "name": "stats_inv_beta_residual",
1645
+ "attrs": { "epsilon": 0.00001 },
1646
+ "inputs": {
1647
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1648
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1649
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1650
+ "betaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1651
+ },
1652
+ "outputs": {
1653
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1654
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1655
+ "residualT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001 }
1656
+ }
1657
+ },
1658
+ {
1659
+ "name": "stats_inv_bias",
1660
+ "attrs": { "epsilon": 0.00001 },
1661
+ "inputs": {
1662
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1663
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1664
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1665
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
1666
+ },
1667
+ "outputs": {
1668
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1669
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1670
+ }
1671
+ },
1672
+ {
1673
+ "name": "stats_inv_bias_residual",
1674
+ "attrs": { "epsilon": 0.00001 },
1675
+ "inputs": {
1676
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1677
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1678
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1679
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
1680
+ },
1681
+ "outputs": {
1682
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1683
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1684
+ "residualT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002 }
1685
+ }
1686
+ },
1687
+ {
1688
+ "name": "stats_inv_bias_beta",
1689
+ "attrs": { "epsilon": 0.00001 },
1690
+ "inputs": {
1691
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1692
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1693
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1694
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } },
1695
+ "betaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1696
+ },
1697
+ "outputs": {
1698
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1699
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1700
+ }
1701
+ },
1702
+ {
1703
+ "name": "stats_inv_bias_beta_residual",
1704
+ "attrs": { "epsilon": 0.00001 },
1705
+ "inputs": {
1706
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1707
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1708
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1709
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } },
1710
+ "betaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1711
+ },
1712
+ "outputs": {
1713
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1714
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1715
+ "residualT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002 }
1716
+ }
1717
+ },
1718
+ {
1719
+ "name": "stats_both_plain",
1720
+ "attrs": { "epsilon": 0.00001 },
1721
+ "inputs": {
1722
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1723
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1724
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
1725
+ },
1726
+ "outputs": {
1727
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1728
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1729
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1730
+ }
1731
+ },
1732
+ {
1733
+ "name": "stats_both_residual",
1734
+ "attrs": { "epsilon": 0.00001 },
1735
+ "inputs": {
1736
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1737
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1738
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } }
1739
+ },
1740
+ "outputs": {
1741
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1742
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1743
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1744
+ "residualT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001 }
1745
+ }
1746
+ },
1747
+ {
1748
+ "name": "stats_both_beta",
1749
+ "attrs": { "epsilon": 0.00001 },
1750
+ "inputs": {
1751
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1752
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1753
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1754
+ "betaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1755
+ },
1756
+ "outputs": {
1757
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1758
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1759
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1760
+ }
1761
+ },
1762
+ {
1763
+ "name": "stats_both_beta_residual",
1764
+ "attrs": { "epsilon": 0.00001 },
1765
+ "inputs": {
1766
+ "inputT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1767
+ "skipT": { "dtype": "float32", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1768
+ "gammaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1769
+ "betaT": { "dtype": "float32", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1770
+ },
1771
+ "outputs": {
1772
+ "outputT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001, "relTolerance": 0.00001 },
1773
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1774
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1775
+ "residualT": { "dtype": "float32", "shape": [2, 3, 5], "tolerance": 0.00001 }
1776
+ }
1777
+ },
1778
+ {
1779
+ "name": "stats_both_bias",
1780
+ "attrs": { "epsilon": 0.00001 },
1781
+ "inputs": {
1782
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1783
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1784
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1785
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
1786
+ },
1787
+ "outputs": {
1788
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1789
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1790
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1791
+ }
1792
+ },
1793
+ {
1794
+ "name": "stats_both_bias_residual",
1795
+ "attrs": { "epsilon": 0.00001 },
1796
+ "inputs": {
1797
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1798
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1799
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1800
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } }
1801
+ },
1802
+ "outputs": {
1803
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1804
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1805
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1806
+ "residualT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002 }
1807
+ }
1808
+ },
1809
+ {
1810
+ "name": "stats_both_bias_beta",
1811
+ "attrs": { "epsilon": 0.00001 },
1812
+ "inputs": {
1813
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1814
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1815
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1816
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } },
1817
+ "betaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1818
+ },
1819
+ "outputs": {
1820
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1821
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1822
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 }
1823
+ }
1824
+ },
1825
+ {
1826
+ "name": "stats_both_bias_beta_residual",
1827
+ "attrs": { "epsilon": 0.00001 },
1828
+ "inputs": {
1829
+ "inputT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1830
+ "skipT": { "dtype": "float16", "shape": [2, 3, 5], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1831
+ "gammaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1832
+ "biasT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } },
1833
+ "betaT": { "dtype": "float16", "shape": [5], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1834
+ },
1835
+ "outputs": {
1836
+ "outputT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002, "relTolerance": 0.002 },
1837
+ "meanT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1838
+ "invStdT": { "dtype": "float32", "shape": [2, 3, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1839
+ "residualT": { "dtype": "float16", "shape": [2, 3, 5], "tolerance": 0.002 }
1840
+ }
1841
+ },
1842
+ {
1843
+ "name": "ort_saved_statistics",
1844
+ "provenance": {
1845
+ "source": "onnxruntime/test/contrib_ops/skiplayernorm_op_test.cc",
1846
+ "test": "SkipLayerNormStatistics"
1847
+ },
1848
+ "attrs": { "epsilon": 1e-12 },
1849
+ "inputs": {
1850
+ "inputT": {
1851
+ "dtype": "float32",
1852
+ "shape": [1, 1, 4],
1853
+ "data": { "kind": "values", "values": [10000.0, 10001.0, 9999.0, 10000.0] }
1854
+ },
1855
+ "skipT": { "dtype": "float32", "shape": [1, 1, 4], "data": { "kind": "constant", "value": 0.0 } },
1856
+ "gammaT": { "dtype": "float32", "shape": [4], "data": { "kind": "constant", "value": 1.0 } },
1857
+ "betaT": { "dtype": "float32", "shape": [4], "data": { "kind": "constant", "value": 0.0 } }
1858
+ },
1859
+ "outputs": {
1860
+ "outputT": {
1861
+ "dtype": "float32",
1862
+ "shape": [1, 1, 4],
1863
+ "tolerance": 0.000001,
1864
+ "data": { "kind": "values", "values": [0.0, 1.4142135, -1.4142135, 0.0] }
1865
+ },
1866
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "data": { "kind": "values", "values": [10000.0] } },
1867
+ "invStdT": {
1868
+ "dtype": "float32",
1869
+ "shape": [1, 1, 1],
1870
+ "tolerance": 0.000001,
1871
+ "data": { "kind": "values", "values": [1.4142135] }
1872
+ }
1873
+ }
1874
+ },
1875
+ {
1876
+ "name": "stats_single_element_row",
1877
+ "attrs": { "epsilon": 0.00001 },
1878
+ "inputs": {
1879
+ "inputT": { "dtype": "float16", "shape": [1, 1, 1], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1880
+ "skipT": { "dtype": "float16", "shape": [1, 1, 1], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1881
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1882
+ "biasT": { "dtype": "float16", "shape": [1], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } },
1883
+ "betaT": { "dtype": "float16", "shape": [1], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1884
+ },
1885
+ "outputs": {
1886
+ "outputT": { "dtype": "float16", "shape": [1, 1, 1], "tolerance": 0.002, "relTolerance": 0.002 },
1887
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1888
+ "invStdT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1889
+ "residualT": { "dtype": "float16", "shape": [1, 1, 1], "tolerance": 0.002 }
1890
+ }
1891
+ },
1892
+ {
1893
+ "name": "stats_single_element_mean",
1894
+ "attrs": { "epsilon": 0.00001 },
1895
+ "inputs": {
1896
+ "inputT": { "dtype": "float16", "shape": [1, 1, 1], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1897
+ "skipT": { "dtype": "float16", "shape": [1, 1, 1], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1898
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1899
+ "biasT": { "dtype": "float16", "shape": [1], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } },
1900
+ "betaT": { "dtype": "float16", "shape": [1], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1901
+ },
1902
+ "outputs": {
1903
+ "outputT": { "dtype": "float16", "shape": [1, 1, 1], "tolerance": 0.002, "relTolerance": 0.002 },
1904
+ "meanT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1905
+ "residualT": { "dtype": "float16", "shape": [1, 1, 1], "tolerance": 0.002 }
1906
+ }
1907
+ },
1908
+ {
1909
+ "name": "stats_single_element_inverse",
1910
+ "attrs": { "epsilon": 0.00001 },
1911
+ "inputs": {
1912
+ "inputT": { "dtype": "float16", "shape": [1, 1, 1], "data": { "kind": "linspace", "start": -2.0, "end": 3.0 } },
1913
+ "skipT": { "dtype": "float16", "shape": [1, 1, 1], "data": { "kind": "linspace", "start": 0.4, "end": -0.7 } },
1914
+ "gammaT": { "dtype": "float16", "shape": [1], "data": { "kind": "linspace", "start": 0.5, "end": 1.5 } },
1915
+ "biasT": { "dtype": "float16", "shape": [1], "data": { "kind": "linspace", "start": -0.2, "end": 0.3 } },
1916
+ "betaT": { "dtype": "float16", "shape": [1], "data": { "kind": "linspace", "start": 0.2, "end": -0.3 } }
1917
+ },
1918
+ "outputs": {
1919
+ "outputT": { "dtype": "float16", "shape": [1, 1, 1], "tolerance": 0.002, "relTolerance": 0.002 },
1920
+ "invStdT": { "dtype": "float32", "shape": [1, 1, 1], "tolerance": 0.00001, "relTolerance": 0.00001 },
1921
+ "residualT": { "dtype": "float16", "shape": [1, 1, 1], "tolerance": 0.002 }
1922
+ }
1923
  }
1924
  ]
1925
  }