Xenova HF Staff commited on
Commit
9088f09
·
verified ·
1 Parent(s): 8e0c6a5

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -52,7 +52,7 @@ Default values (overridable per request):
52
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
53
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
54
  - [`test.json`](build/webgpu/test.json) — correctness cases
55
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
56
  - [`norm-row-stats.wgsl.jinja`](build/webgpu/norm-row-stats.wgsl.jinja)
57
  - [`rms-normalization-splitk-normalize.wgsl.jinja`](build/webgpu/rms-normalization-splitk-normalize.wgsl.jinja)
58
  - [`rms-normalization-splitk-partials.wgsl.jinja`](build/webgpu/rms-normalization-splitk-partials.wgsl.jinja)
@@ -61,7 +61,7 @@ Default values (overridable per request):
61
  ## Use with `@huggingface/kernels`
62
 
63
  ```sh
64
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
65
  ```
66
 
67
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
52
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
53
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
54
  - [`test.json`](build/webgpu/test.json) — correctness cases
55
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
56
  - [`norm-row-stats.wgsl.jinja`](build/webgpu/norm-row-stats.wgsl.jinja)
57
  - [`rms-normalization-splitk-normalize.wgsl.jinja`](build/webgpu/rms-normalization-splitk-normalize.wgsl.jinja)
58
  - [`rms-normalization-splitk-partials.wgsl.jinja`](build/webgpu/rms-normalization-splitk-partials.wgsl.jinja)
 
61
  ## Use with `@huggingface/kernels`
62
 
63
  ```sh
64
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
65
  ```
66
 
67
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/manifest.json CHANGED
@@ -42,7 +42,6 @@
42
  "rowWg": "min(normMaxWorkgroup, pow2ceil(max(1, normHidden)))",
43
  "baseOk": "ranks.x >= 1 and sameShape(shapes.y, shapes.x) and ranks.scale >= 0 and ranks.scale <= ranks.x and broadcastable(shapes.scale, shapes.x) and attrs.axis + ranks.x >= 0 and attrs.axis < ranks.x and normHidden > 0 and attrs.stash_type == onnxDtypeCode(\"float32\") and f16Ok(dtypes.T) and f16Ok(dtypes.V)",
44
  "lastAxisOk": "baseOk and (attrs.axis == -1 or attrs.axis == ranks.x - 1)",
45
- "suffixAxisOk": "baseOk and ranks.x >= 2 and not (attrs.axis == -1 or attrs.axis == ranks.x - 1)",
46
  "noStats": "not present.invStdVar",
47
  "statsOk": "present.invStdVar and ranks.invStdVar == ranks.x and sameShape(prefix(shapes.invStdVar, axisNorm), prefix(shapes.x, axisNorm)) and numel(suffix(shapes.invStdVar, axisNorm)) == 1",
48
  "sameDtype": "dtypes.T == dtypes.V",
@@ -51,24 +50,23 @@
51
  "splitFits": "normRows <= tunables.SPLIT_MAX_ROWS and splitCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and splitScratchBytes <= device.limits.maxStorageBufferBindingSize and splitScratchBytes <= device.limits.maxBufferSize"
52
  },
53
  "bindings": {
54
- "x": { "buffer": "read-only-storage", "elementType": "$xElement" },
55
- "scale": { "buffer": "read-only-storage", "elementType": "$ioElement" },
56
- "y": { "buffer": "storage", "elementType": "$ioElement" },
57
  "params": {
58
- "buffer": "uniform",
59
  "struct": [
60
  { "name": "rows", "type": "u32", "value": "normRows" },
61
  { "name": "rowStride", "type": "u32", "value": "normRowStride" }
62
  ]
63
  },
64
- "inv_std_out": { "arg": "invStdVar", "buffer": "storage", "elementType": "f32" },
65
- "partials_2": { "name": "partials", "buffer": "read-only-storage", "elementType": "f32" }
66
  },
67
  "variants": [
68
  {
69
  "id": "last_axis",
70
  "priority": 1,
71
- "when": ["lastAxisOk", "noStats"],
72
  "derive": {
73
  "scalar": "dtypes.V",
74
  "xElement": "dtypes.T",
@@ -87,8 +85,7 @@
87
  "scaleShape": "shapes.scale",
88
  "xRank": "ranks.x",
89
  "scaleRank": "ranks.scale",
90
- "writeStats": false,
91
- "rmsScaleAfterCast": false
92
  },
93
  "bindings": ["x", "scale", "y", "params"],
94
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
@@ -98,7 +95,7 @@
98
  {
99
  "id": "last_axis_stats",
100
  "priority": 2,
101
- "when": ["lastAxisOk", "statsOk"],
102
  "derive": {
103
  "scalar": "dtypes.V",
104
  "xElement": "dtypes.T",
@@ -117,68 +114,7 @@
117
  "scaleShape": "shapes.scale",
118
  "xRank": "ranks.x",
119
  "scaleRank": "ranks.scale",
120
- "writeStats": true,
121
- "rmsScaleAfterCast": false
122
- },
123
- "bindings": ["x", "scale", "y", "inv_std_out", "params"],
124
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
125
- }
126
- ]
127
- },
128
- {
129
- "id": "suffix_axis",
130
- "priority": 10,
131
- "when": ["suffixAxisOk", "noStats"],
132
- "derive": {
133
- "scalar": "dtypes.V",
134
- "xElement": "dtypes.T",
135
- "ioElement": "dtypes.V",
136
- "hiddenSize": "normHidden",
137
- "workgroupSize": "rowWg",
138
- "epsilon": "attrs.epsilon"
139
- },
140
- "passes": [
141
- {
142
- "id": "main",
143
- "name": "SimplifiedLayerNormalization.Row",
144
- "shader": "rms-normalization.wgsl.jinja",
145
- "derive": {
146
- "xShape": "shapes.x",
147
- "scaleShape": "shapes.scale",
148
- "xRank": "ranks.x",
149
- "scaleRank": "ranks.scale",
150
- "writeStats": false,
151
- "rmsScaleAfterCast": false
152
- },
153
- "bindings": ["x", "scale", "y", "params"],
154
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
155
- }
156
- ]
157
- },
158
- {
159
- "id": "suffix_axis_stats",
160
- "priority": 11,
161
- "when": ["suffixAxisOk", "statsOk"],
162
- "derive": {
163
- "scalar": "dtypes.V",
164
- "xElement": "dtypes.T",
165
- "ioElement": "dtypes.V",
166
- "hiddenSize": "normHidden",
167
- "workgroupSize": "rowWg",
168
- "epsilon": "attrs.epsilon"
169
- },
170
- "passes": [
171
- {
172
- "id": "main",
173
- "name": "SimplifiedLayerNormalization.Row",
174
- "shader": "rms-normalization.wgsl.jinja",
175
- "derive": {
176
- "xShape": "shapes.x",
177
- "scaleShape": "shapes.scale",
178
- "xRank": "ranks.x",
179
- "scaleRank": "ranks.scale",
180
- "writeStats": true,
181
- "rmsScaleAfterCast": false
182
  },
183
  "bindings": ["x", "scale", "y", "inv_std_out", "params"],
184
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
@@ -206,7 +142,7 @@
206
  "id": "partials",
207
  "name": "SimplifiedLayerNormalization.SplitKPartials",
208
  "shader": "rms-normalization-splitk-partials.wgsl.jinja",
209
- "bindings": ["x", { "name": "partials", "buffer": "storage", "elementType": "f32" }, "params"],
210
  "dispatch": {
211
  "x": "min(normRows, DISPATCH_FOLD_WIDTH)",
212
  "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)",
@@ -222,12 +158,11 @@
222
  "scaleShape": "shapes.scale",
223
  "xRank": "ranks.x",
224
  "scaleRank": "ranks.scale",
225
- "writeStats": false,
226
- "rmsScaleAfterCast": false,
227
  "normalizeBlocks": "max(split, min(min(device.limits.maxComputeWorkgroupsPerDimension, 65535), ceilDiv(workgroupSize, max(1, normalizeRows)), ceilDiv(hiddenSize, workgroupSize * 4)))",
228
  "normalizeChunk": "ceilDiv(hiddenSize, normalizeBlocks)"
229
  },
230
- "bindings": ["x", "scale", "partials_2", "y", "params"],
231
  "dispatch": {
232
  "x": "min(normRows, DISPATCH_FOLD_WIDTH)",
233
  "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)",
@@ -257,7 +192,7 @@
257
  "id": "partials",
258
  "name": "SimplifiedLayerNormalization.SplitKPartials",
259
  "shader": "rms-normalization-splitk-partials.wgsl.jinja",
260
- "bindings": ["x", { "name": "partials", "buffer": "storage", "elementType": "f32" }, "params"],
261
  "dispatch": {
262
  "x": "min(normRows, DISPATCH_FOLD_WIDTH)",
263
  "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)",
@@ -273,12 +208,11 @@
273
  "scaleShape": "shapes.scale",
274
  "xRank": "ranks.x",
275
  "scaleRank": "ranks.scale",
276
- "writeStats": true,
277
- "rmsScaleAfterCast": false,
278
  "normalizeBlocks": "max(split, min(min(device.limits.maxComputeWorkgroupsPerDimension, 65535), ceilDiv(workgroupSize, max(1, normalizeRows)), ceilDiv(hiddenSize, workgroupSize * 4)))",
279
  "normalizeChunk": "ceilDiv(hiddenSize, normalizeBlocks)"
280
  },
281
- "bindings": ["x", "scale", "partials_2", "y", "inv_std_out", "params"],
282
  "dispatch": {
283
  "x": "min(normRows, DISPATCH_FOLD_WIDTH)",
284
  "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)",
@@ -302,12 +236,8 @@
302
  "name": "SimplifiedLayerNormalization.LastAxisRow",
303
  "shader": "norm-row-stats.wgsl.jinja",
304
  "derive": {
305
- "modeSpec": "\"rms\"",
306
  "vec4": true,
307
- "writeStats": false,
308
- "rmsScaleAfterCast": false,
309
- "scalar": "dtypes.T",
310
- "usesF16Spec": "dtypes.T == \"f16\"",
311
  "hidden": "dim(shapes.x, -1)",
312
  "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1) / 4)))",
313
  "epsilon": "attrs.epsilon",
@@ -316,8 +246,7 @@
316
  "combineSubgroups": "hasSubgroupId"
317
  },
318
  "bindings": ["x", "scale", "y", "params"],
319
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
320
- "subgroupCollectivesWidth": "portable"
321
  }
322
  ]
323
  },
@@ -332,12 +261,8 @@
332
  "name": "SimplifiedLayerNormalization.LastAxisRow",
333
  "shader": "norm-row-stats.wgsl.jinja",
334
  "derive": {
335
- "modeSpec": "\"rms\"",
336
  "vec4": false,
337
- "writeStats": false,
338
- "rmsScaleAfterCast": false,
339
- "scalar": "dtypes.T",
340
- "usesF16Spec": "dtypes.T == \"f16\"",
341
  "hidden": "dim(shapes.x, -1)",
342
  "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1))))",
343
  "epsilon": "attrs.epsilon",
@@ -346,8 +271,7 @@
346
  "combineSubgroups": "hasSubgroupId"
347
  },
348
  "bindings": ["x", "scale", "y", "params"],
349
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
350
- "subgroupCollectivesWidth": "portable"
351
  }
352
  ]
353
  },
@@ -366,12 +290,8 @@
366
  "name": "SimplifiedLayerNormalization.LastAxisRow",
367
  "shader": "norm-row-stats.wgsl.jinja",
368
  "derive": {
369
- "modeSpec": "\"rms\"",
370
  "vec4": true,
371
- "writeStats": true,
372
- "rmsScaleAfterCast": false,
373
- "scalar": "dtypes.T",
374
- "usesF16Spec": "dtypes.T == \"f16\"",
375
  "hidden": "dim(shapes.x, -1)",
376
  "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1) / 4)))",
377
  "epsilon": "attrs.epsilon",
@@ -380,8 +300,7 @@
380
  "combineSubgroups": "hasSubgroupId"
381
  },
382
  "bindings": ["x", "scale", "y", "inv_std_out", "params"],
383
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
384
- "subgroupCollectivesWidth": "portable"
385
  }
386
  ]
387
  },
@@ -396,12 +315,8 @@
396
  "name": "SimplifiedLayerNormalization.LastAxisRow",
397
  "shader": "norm-row-stats.wgsl.jinja",
398
  "derive": {
399
- "modeSpec": "\"rms\"",
400
  "vec4": false,
401
- "writeStats": true,
402
- "rmsScaleAfterCast": false,
403
- "scalar": "dtypes.T",
404
- "usesF16Spec": "dtypes.T == \"f16\"",
405
  "hidden": "dim(shapes.x, -1)",
406
  "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1))))",
407
  "epsilon": "attrs.epsilon",
@@ -410,8 +325,7 @@
410
  "combineSubgroups": "hasSubgroupId"
411
  },
412
  "bindings": ["x", "scale", "y", "inv_std_out", "params"],
413
- "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 },
414
- "subgroupCollectivesWidth": "portable"
415
  }
416
  ]
417
  }
 
42
  "rowWg": "min(normMaxWorkgroup, pow2ceil(max(1, normHidden)))",
43
  "baseOk": "ranks.x >= 1 and sameShape(shapes.y, shapes.x) and ranks.scale >= 0 and ranks.scale <= ranks.x and broadcastable(shapes.scale, shapes.x) and attrs.axis + ranks.x >= 0 and attrs.axis < ranks.x and normHidden > 0 and attrs.stash_type == onnxDtypeCode(\"float32\") and f16Ok(dtypes.T) and f16Ok(dtypes.V)",
44
  "lastAxisOk": "baseOk and (attrs.axis == -1 or attrs.axis == ranks.x - 1)",
 
45
  "noStats": "not present.invStdVar",
46
  "statsOk": "present.invStdVar and ranks.invStdVar == ranks.x and sameShape(prefix(shapes.invStdVar, axisNorm), prefix(shapes.x, axisNorm)) and numel(suffix(shapes.invStdVar, axisNorm)) == 1",
47
  "sameDtype": "dtypes.T == dtypes.V",
 
50
  "splitFits": "normRows <= tunables.SPLIT_MAX_ROWS and splitCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and splitScratchBytes <= device.limits.maxStorageBufferBindingSize and splitScratchBytes <= device.limits.maxBufferSize"
51
  },
52
  "bindings": {
53
+ "x": { "elementType": "$xElement" },
54
+ "scale": { "elementType": "$ioElement" },
55
+ "y": { "elementType": "$ioElement" },
56
  "params": {
 
57
  "struct": [
58
  { "name": "rows", "type": "u32", "value": "normRows" },
59
  { "name": "rowStride", "type": "u32", "value": "normRowStride" }
60
  ]
61
  },
62
+ "inv_std_out": { "arg": "invStdVar", "elementType": "f32" },
63
+ "partials_f32": { "name": "partials", "buffer": "read-only-storage", "elementType": "f32" }
64
  },
65
  "variants": [
66
  {
67
  "id": "last_axis",
68
  "priority": 1,
69
+ "when": ["baseOk", "noStats"],
70
  "derive": {
71
  "scalar": "dtypes.V",
72
  "xElement": "dtypes.T",
 
85
  "scaleShape": "shapes.scale",
86
  "xRank": "ranks.x",
87
  "scaleRank": "ranks.scale",
88
+ "writeStats": "present.invStdVar"
 
89
  },
90
  "bindings": ["x", "scale", "y", "params"],
91
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
95
  {
96
  "id": "last_axis_stats",
97
  "priority": 2,
98
+ "when": ["baseOk", "statsOk"],
99
  "derive": {
100
  "scalar": "dtypes.V",
101
  "xElement": "dtypes.T",
 
114
  "scaleShape": "shapes.scale",
115
  "xRank": "ranks.x",
116
  "scaleRank": "ranks.scale",
117
+ "writeStats": "present.invStdVar"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
118
  },
119
  "bindings": ["x", "scale", "y", "inv_std_out", "params"],
120
  "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
142
  "id": "partials",
143
  "name": "SimplifiedLayerNormalization.SplitKPartials",
144
  "shader": "rms-normalization-splitk-partials.wgsl.jinja",
145
+ "bindings": ["x", { "name": "partials", "elementType": "f32" }, "params"],
146
  "dispatch": {
147
  "x": "min(normRows, DISPATCH_FOLD_WIDTH)",
148
  "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)",
 
158
  "scaleShape": "shapes.scale",
159
  "xRank": "ranks.x",
160
  "scaleRank": "ranks.scale",
161
+ "writeStats": "present.invStdVar",
 
162
  "normalizeBlocks": "max(split, min(min(device.limits.maxComputeWorkgroupsPerDimension, 65535), ceilDiv(workgroupSize, max(1, normalizeRows)), ceilDiv(hiddenSize, workgroupSize * 4)))",
163
  "normalizeChunk": "ceilDiv(hiddenSize, normalizeBlocks)"
164
  },
165
+ "bindings": ["x", "scale", "partials_f32", "y", "params"],
166
  "dispatch": {
167
  "x": "min(normRows, DISPATCH_FOLD_WIDTH)",
168
  "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)",
 
192
  "id": "partials",
193
  "name": "SimplifiedLayerNormalization.SplitKPartials",
194
  "shader": "rms-normalization-splitk-partials.wgsl.jinja",
195
+ "bindings": ["x", { "name": "partials", "elementType": "f32" }, "params"],
196
  "dispatch": {
197
  "x": "min(normRows, DISPATCH_FOLD_WIDTH)",
198
  "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)",
 
208
  "scaleShape": "shapes.scale",
209
  "xRank": "ranks.x",
210
  "scaleRank": "ranks.scale",
211
+ "writeStats": "present.invStdVar",
 
212
  "normalizeBlocks": "max(split, min(min(device.limits.maxComputeWorkgroupsPerDimension, 65535), ceilDiv(workgroupSize, max(1, normalizeRows)), ceilDiv(hiddenSize, workgroupSize * 4)))",
213
  "normalizeChunk": "ceilDiv(hiddenSize, normalizeBlocks)"
214
  },
215
+ "bindings": ["x", "scale", "partials_f32", "y", "inv_std_out", "params"],
216
  "dispatch": {
217
  "x": "min(normRows, DISPATCH_FOLD_WIDTH)",
218
  "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)",
 
236
  "name": "SimplifiedLayerNormalization.LastAxisRow",
237
  "shader": "norm-row-stats.wgsl.jinja",
238
  "derive": {
 
239
  "vec4": true,
240
+ "writeStats": "present.invStdVar",
 
 
 
241
  "hidden": "dim(shapes.x, -1)",
242
  "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1) / 4)))",
243
  "epsilon": "attrs.epsilon",
 
246
  "combineSubgroups": "hasSubgroupId"
247
  },
248
  "bindings": ["x", "scale", "y", "params"],
249
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
250
  }
251
  ]
252
  },
 
261
  "name": "SimplifiedLayerNormalization.LastAxisRow",
262
  "shader": "norm-row-stats.wgsl.jinja",
263
  "derive": {
 
264
  "vec4": false,
265
+ "writeStats": "present.invStdVar",
 
 
 
266
  "hidden": "dim(shapes.x, -1)",
267
  "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1))))",
268
  "epsilon": "attrs.epsilon",
 
271
  "combineSubgroups": "hasSubgroupId"
272
  },
273
  "bindings": ["x", "scale", "y", "params"],
274
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
275
  }
276
  ]
277
  },
 
290
  "name": "SimplifiedLayerNormalization.LastAxisRow",
291
  "shader": "norm-row-stats.wgsl.jinja",
292
  "derive": {
 
293
  "vec4": true,
294
+ "writeStats": "present.invStdVar",
 
 
 
295
  "hidden": "dim(shapes.x, -1)",
296
  "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1) / 4)))",
297
  "epsilon": "attrs.epsilon",
 
300
  "combineSubgroups": "hasSubgroupId"
301
  },
302
  "bindings": ["x", "scale", "y", "inv_std_out", "params"],
303
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
304
  }
305
  ]
306
  },
 
315
  "name": "SimplifiedLayerNormalization.LastAxisRow",
316
  "shader": "norm-row-stats.wgsl.jinja",
317
  "derive": {
 
318
  "vec4": false,
319
+ "writeStats": "present.invStdVar",
 
 
 
320
  "hidden": "dim(shapes.x, -1)",
321
  "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1))))",
322
  "epsilon": "attrs.epsilon",
 
325
  "combineSubgroups": "hasSubgroupId"
326
  },
327
  "bindings": ["x", "scale", "y", "inv_std_out", "params"],
328
+ "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }
 
329
  }
330
  ]
331
  }
build/webgpu/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "ai.onnx.SimplifiedLayerNormalization",
3
- "id": "_ai_onnx_simplifiedlayernormalization_webgpu_b0bbd51",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
@@ -8,22 +8,20 @@
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "oL4NDZsKfggbFJ8moDoDY5cYbkZb/dnRCxB7mA3YuW0=",
11
- "manifest.json": "M3zOOpjNYbDBPlApZytNSUFrdixB+kGxE4X9xKY7yzk=",
12
- "norm-row-stats.wgsl.jinja": "eRBO50QnNhvyqRW/Wdw6rzfJRgVlRT0P6oWVqcDPf5w=",
13
- "rms-normalization-splitk-normalize.wgsl.jinja": "vqx3jngJ7c+zkmcD3Nc7mt4Fm77XXDmtrQt9ULWgHKE=",
14
- "rms-normalization-splitk-partials.wgsl.jinja": "Vs4HNsa9ZbLPsW/ZOunRg64qFmbg7WfJ/uE6gqALmSQ=",
15
- "rms-normalization.wgsl.jinja": "tRDNHNbkHoidOuTQwnH3Jx0msfP74Px5ukRNF9Hh8fw=",
16
- "test.json": "PYhuNCaFTpBU8GJFS/LzpdoBKLbT0praPXzHhs7rg50="
17
  }
18
  },
19
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
20
  "webgpu": {
21
- "manifestSpec": "2.0",
22
  "variants": {
23
  "last_axis": ["rms-normalization.wgsl.jinja"],
24
  "last_axis_stats": ["rms-normalization.wgsl.jinja"],
25
- "suffix_axis": ["rms-normalization.wgsl.jinja"],
26
- "suffix_axis_stats": ["rms-normalization.wgsl.jinja"],
27
  "suffix_axis_splitk": ["rms-normalization-splitk-normalize.wgsl.jinja", "rms-normalization-splitk-partials.wgsl.jinja"],
28
  "suffix_axis_splitk_stats": ["rms-normalization-splitk-normalize.wgsl.jinja", "rms-normalization-splitk-partials.wgsl.jinja"],
29
  "last_axis_row_vec4": ["norm-row-stats.wgsl.jinja"],
 
1
  {
2
  "name": "ai.onnx.SimplifiedLayerNormalization",
3
+ "id": "_ai_onnx_simplifiedlayernormalization_webgpu_828538c",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
 
8
  "algorithm": "sha256",
9
  "files": {
10
  "bench.json": "oL4NDZsKfggbFJ8moDoDY5cYbkZb/dnRCxB7mA3YuW0=",
11
+ "manifest.json": "L4cR4xPJA2ltA6TlL/1TweUi/WgEMnu7OM6SEzT3FGo=",
12
+ "norm-row-stats.wgsl.jinja": "BiuyYX6gDKD2euzLhyzWkE2B6B1rEPBIwm0sgP85vcQ=",
13
+ "rms-normalization-splitk-normalize.wgsl.jinja": "pWbF7PGQt1vPjkJT5OjujTJCXAAtGEFeciCxM4RKkSo=",
14
+ "rms-normalization-splitk-partials.wgsl.jinja": "KnvI9/ZZe4zPBS+v4ieFHw7RTg/2c/MqVHH1Zj3T/F4=",
15
+ "rms-normalization.wgsl.jinja": "2ArCxRtOyP9LH/DdSUUgwrSEiWWzQeIxp2GYBMrI5ww=",
16
+ "test.json": "DpDg1c4o10S6v8mIfKTaOuxumQ0FYOKkIyY6Gs+eF/Q="
17
  }
18
  },
19
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
20
  "webgpu": {
21
+ "manifestSpec": "2.1",
22
  "variants": {
23
  "last_axis": ["rms-normalization.wgsl.jinja"],
24
  "last_axis_stats": ["rms-normalization.wgsl.jinja"],
 
 
25
  "suffix_axis_splitk": ["rms-normalization-splitk-normalize.wgsl.jinja", "rms-normalization-splitk-partials.wgsl.jinja"],
26
  "suffix_axis_splitk_stats": ["rms-normalization-splitk-normalize.wgsl.jinja", "rms-normalization-splitk-partials.wgsl.jinja"],
27
  "last_axis_row_vec4": ["norm-row-stats.wgsl.jinja"],
build/webgpu/norm-row-stats.wgsl.jinja CHANGED
@@ -1,25 +1,7 @@
1
- {% if usesF16Spec %}
2
- enable f16;
3
- {% endif %}
4
- {% set combineSubgroups = combineSubgroups %}
5
- {% set scalarIo = scalarIo if scalarIo is defined else false %}
6
- {% set packedBf16Embedding = packedBf16Embedding if packedBf16Embedding is defined else false %}
7
- {% set writeStats = writeStats if writeStats is defined else false %}
8
- {% set rmsWeightOffset = rmsWeightOffset if rmsWeightOffset is defined else false %}
9
- {% set rmsScaleAfterCast = rmsScaleAfterCast if rmsScaleAfterCast is defined else false %}
10
- {% set rmsResidualAdd = rmsResidualAdd if rmsResidualAdd is defined else false %}
11
- {% set rmsChainNorm = rmsChainNorm if rmsChainNorm is defined else false %}
12
- {% set hiddenPairs = hiddenPairs | default(0) %}
13
- {% set numRows = numRows | default(0) %}
14
- {% set epsilon = epsilon | default("0.0") %}
15
- {% set epsilon2 = epsilon2 | default("0.0") %}
16
- {% if rmsWeightOffset %}
17
- {% set rmsScaleVec = "(vec4<f32>(1.0) + vec4<f32>(scale[i]))" %}
18
- {% set rmsScaleScalar = "(1.0 + f32(scale[i]))" %}
19
- {% else %}
20
  {% set rmsScaleVec = "vec4<f32>(scale[i])" %}
21
  {% set rmsScaleScalar = "f32(scale[i])" %}
22
- {% endif %}
23
  {% set reduceThreadParameters = ", sg_lane: u32, sg_id: u32, num_sg: u32"
24
  if combineSubgroups else ", tid: u32" %}
25
  {% set reduceThreadArguments = ", sg_lane, sg_id, num_sg"
@@ -42,53 +24,8 @@ const HIDDEN: u32 = {{ hidden }}u;
42
  {% if vec4 %}
43
  const HIDDEN_V: u32 = {{ hiddenVec }}u;
44
  {% endif %}
45
- {% if packedBf16Embedding %}
46
- const HIDDEN_PAIRS: u32 = {{ hiddenPairs }}u;
47
- const NUM_ROWS: u32 = {{ numRows }}u;
48
- {% endif %}
49
  const WG: u32 = {{ wg }}u;
50
  const EPSILON: f32 = {{ epsilon }};
51
- {% if rmsChainNorm %}
52
- const EPSILON2: f32 = {{ epsilon2 }};
53
- {% endif %}
54
-
55
- {% if packedBf16Embedding %}
56
- {% if vec4 %}
57
- fn unpack_bf16_pair(word: u32) -> vec2<f32> {
58
- let bits = vec2<u32>(word & 0xffffu, word >> 16u);
59
- return bitcast<vec2<f32>>(bits << vec2<u32>(16u));
60
- }
61
- {% endif %}
62
-
63
- {% if not vec4 %}
64
- fn embedding_scalar(source_row: u32, hidden: u32) -> f32 {
65
- if (source_row >= NUM_ROWS) {
66
- return 0.0;
67
- }
68
- let word = x[source_row * HIDDEN_PAIRS + (hidden >> 1u)];
69
- let bits = select(word & 0xffffu, word >> 16u, (hidden & 1u) != 0u);
70
- return bitcast<f32>(bits << 16u);
71
- }
72
- {% endif %}
73
-
74
- {% if vec4 %}
75
- fn embedding_vec4(source_row: u32, hidden_vec: u32) -> vec4<f32> {
76
- if (source_row >= NUM_ROWS) {
77
- return vec4<f32>(0.0);
78
- }
79
- let base = source_row * HIDDEN_PAIRS + hidden_vec * 2u;
80
- let low = unpack_bf16_pair(x[base]);
81
- let high = unpack_bf16_pair(x[base + 1u]);
82
- return vec4<f32>(low, high);
83
- }
84
- {% endif %}
85
- {% endif %}
86
-
87
- {% if vec4 and scalarIo %}
88
- fn load_vec4(index: u32) -> vec4<f32> {
89
- return vec4<f32>(x[index], x[index + 1u], x[index + 2u], x[index + 3u]);
90
- }
91
- {% endif %}
92
 
93
  {% if combineSubgroups %}
94
  var<workgroup> sg_partials: array<f32, WG>;
@@ -142,41 +79,21 @@ fn main(
142
  return;
143
  }
144
  let tid = lid.x;
145
- {% if packedBf16Embedding %}
146
- let source_row = indices[row];
147
- {% if vec4 %}
148
- let base = row * HIDDEN_V;
149
- {% else %}
150
- let base = row * HIDDEN;
151
- {% endif %}
152
- {% elif vec4 and not scalarIo %}
153
  let base = row * HIDDEN_V;
154
  {% else %}
155
  let base = row * HIDDEN;
156
  {% endif %}
157
 
158
-
159
  var acc = 0.0;
160
  {% if vec4 %}
161
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
162
- {% if packedBf16Embedding %}
163
- let v = embedding_vec4(source_row, i);
164
- embedding_out[base + i] = v;
165
- {% elif scalarIo %}
166
- let v = load_vec4(base + i * 4u);
167
- {% else %}
168
- let v = vec4<f32>(x[base + i]);
169
- {% endif %}
170
  acc = acc + dot(v, v);
171
  }
172
  {% else %}
173
  for (var i = tid; i < HIDDEN; i = i + WG) {
174
- {% if packedBf16Embedding %}
175
- let v = embedding_scalar(source_row, i);
176
- embedding_out[base + i] = v;
177
- {% else %}
178
  let v = f32(x[base + i]);
179
- {% endif %}
180
  acc = acc + v * v;
181
  }
182
  {% endif %}
@@ -190,64 +107,17 @@ fn main(
190
  }
191
  {% endif %}
192
 
193
- {% if rmsChainNorm %}
194
- var acc2 = 0.0;
195
- {% endif %}
196
  {% if vec4 %}
197
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
198
- {% if packedBf16Embedding %}
199
  let idx = base + i;
200
- let v = embedding_vec4(source_row, i);
201
- {% elif scalarIo %}
202
- let idx = base + i * 4u;
203
- let v = load_vec4(idx);
204
- {% else %}
205
- let idx = base + i;
206
- let v = vec4<f32>(x[idx]);
207
- {% endif %}
208
- {% if rmsScaleAfterCast %}
209
- y[idx] = {{ vecType }}(v * inv) * {{ vecType }}({{ rmsScaleVec }});
210
- {% elif rmsChainNorm %}
211
- // fma(a, b, 0.0) rounds the weighted product exactly as the decomposed pair's store does,
212
- // and prevents the compiler from re-contracting it into the residual add.
213
- let hv = y[idx] + fma(v * inv, {{ rmsScaleVec }}, vec4<f32>(0.0));
214
- y[idx] = hv;
215
- acc2 = acc2 + dot(hv, hv);
216
- {% elif rmsResidualAdd %}
217
- // See the chained branch: fma(a, b, 0.0) pins the pre-add rounding of the decomposed pair.
218
- y[idx] = y[idx] + fma(v * inv, {{ rmsScaleVec }}, vec4<f32>(0.0));
219
- {% else %}
220
  y[idx] = {{ vecType }}(v * inv * {{ rmsScaleVec }});
221
- {% endif %}
222
- }
223
- {% if rmsChainNorm %}
224
-
225
- // The chained second norm reads the residual row this loop just stored. This
226
- // barrier completes those stores and any preceding shared-scratch use before
227
- // the next reduction reuses its scratch; each lane then re-reads only the
228
- // elements it wrote itself.
229
- workgroupBarrier();
230
- let total2 = reduce_scalar(acc2{{ reduceThreadArguments }});
231
- let inv2 = inverseSqrt(total2 / f32(HIDDEN) + EPSILON2);
232
- for (var i = tid; i < HIDDEN_V; i = i + WG) {
233
- let idx = base + i;
234
- let hv = vec4<f32>(y[idx]);
235
- normed2[idx] = {{ vecType }}(hv * inv2 * vec4<f32>(scale2[i]));
236
  }
237
- {% endif %}
238
  {% else %}
239
  for (var i = tid; i < HIDDEN; i = i + WG) {
240
  let idx = base + i;
241
- {% if packedBf16Embedding %}
242
- let v = embedding_scalar(source_row, i);
243
- {% else %}
244
  let v = f32(x[idx]);
245
- {% endif %}
246
- {% if rmsScaleAfterCast %}
247
- y[idx] = {{ scalar }}(v * inv) * {{ scalar }}({{ rmsScaleScalar }});
248
- {% else %}
249
  y[idx] = {{ scalar }}(v * inv * {{ rmsScaleScalar }});
250
- {% endif %}
251
  }
252
  {% endif %}
253
  }
 
1
+ {% set scalarIo = false %}
2
+ {% set packedF32 = "vec4<f32>" %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3
  {% set rmsScaleVec = "vec4<f32>(scale[i])" %}
4
  {% set rmsScaleScalar = "f32(scale[i])" %}
 
5
  {% set reduceThreadParameters = ", sg_lane: u32, sg_id: u32, num_sg: u32"
6
  if combineSubgroups else ", tid: u32" %}
7
  {% set reduceThreadArguments = ", sg_lane, sg_id, num_sg"
 
24
  {% if vec4 %}
25
  const HIDDEN_V: u32 = {{ hiddenVec }}u;
26
  {% endif %}
 
 
 
 
27
  const WG: u32 = {{ wg }}u;
28
  const EPSILON: f32 = {{ epsilon }};
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
29
 
30
  {% if combineSubgroups %}
31
  var<workgroup> sg_partials: array<f32, WG>;
 
79
  return;
80
  }
81
  let tid = lid.x;
82
+ {% if vec4 and not scalarIo %}
 
 
 
 
 
 
 
83
  let base = row * HIDDEN_V;
84
  {% else %}
85
  let base = row * HIDDEN;
86
  {% endif %}
87
 
 
88
  var acc = 0.0;
89
  {% if vec4 %}
90
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
91
+ let v = {{ packedF32 }}(x[base + i]);
 
 
 
 
 
 
 
92
  acc = acc + dot(v, v);
93
  }
94
  {% else %}
95
  for (var i = tid; i < HIDDEN; i = i + WG) {
 
 
 
 
96
  let v = f32(x[base + i]);
 
97
  acc = acc + v * v;
98
  }
99
  {% endif %}
 
107
  }
108
  {% endif %}
109
 
 
 
 
110
  {% if vec4 %}
111
  for (var i = tid; i < HIDDEN_V; i = i + WG) {
 
112
  let idx = base + i;
113
+ let v = {{ packedF32 }}(x[idx]);
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
114
  y[idx] = {{ vecType }}(v * inv * {{ rmsScaleVec }});
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
115
  }
 
116
  {% else %}
117
  for (var i = tid; i < HIDDEN; i = i + WG) {
118
  let idx = base + i;
 
 
 
119
  let v = f32(x[idx]);
 
 
 
 
120
  y[idx] = {{ scalar }}(v * inv * {{ rmsScaleScalar }});
 
121
  }
122
  {% endif %}
123
  }
build/webgpu/rms-normalization-splitk-normalize.wgsl.jinja CHANGED
@@ -54,7 +54,6 @@ fn scale_offset({% if scaleRank > 0 %}out_index: u32{% endif %}) -> u32 {
54
  {% endif %}
55
  }
56
 
57
-
58
  var<workgroup> shared_inv: f32;
59
 
60
  @compute @workgroup_size(WG, 1, 1)
 
54
  {% endif %}
55
  }
56
 
 
57
  var<workgroup> shared_inv: f32;
58
 
59
  @compute @workgroup_size(WG, 1, 1)
build/webgpu/rms-normalization-splitk-partials.wgsl.jinja CHANGED
@@ -1,49 +1,15 @@
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
  /* Split-K partial sum-of-squares for tensors with few rows and a large hidden
48
  dimension. A workgroup-per-row kernel exposes too little parallelism in this
49
  regime, so this pass splits each row across SPLIT workgroups
 
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 }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} }
 
 
 
 
 
 
 
 
 
 
 
 
10
  {{ svar }} = {{ svar }} / 2u;
11
  workgroupBarrier();
12
+ }{% endmacro %}
 
 
 
 
 
 
 
 
 
 
13
  /* Split-K partial sum-of-squares for tensors with few rows and a large hidden
14
  dimension. A workgroup-per-row kernel exposes too little parallelism in this
15
  regime, so this pass splits each row across SPLIT workgroups
build/webgpu/rms-normalization.wgsl.jinja CHANGED
@@ -4,8 +4,6 @@ const HIDDEN: u32 = {{ hiddenSize }}u;
4
  const EPSILON: f32 = {{ epsilon }};
5
  const WG: u32 = {{ workgroupSize }}u;
6
 
7
- var<workgroup> partial: array<f32, WG>;
8
-
9
  {% if scaleRank > 0 %}
10
  const X_RANK: u32 = {{ xRank }}u;
11
  const SCALE_RANK: u32 = {{ scaleRank }}u;
@@ -51,70 +49,34 @@ fn scale_offset({% if scaleRank > 0 %}out_index: u32{% endif %}) -> u32 {
51
  {% endif %}
52
  }
53
 
54
-
55
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
56
- {% if op == "max" %}
57
- {{ a }}[{{ idx }}] = max({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);
58
- {%- else %}
59
- {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] + {{ a }}[{{ idx }} + {{ svar }}];
60
- {%- endif %}
61
- {% endmacro %}
62
- {% macro wgsl_tree_fold(arrays, op="add", idx="lid", wg="WORKGROUP_SIZE", svar="stride", typed=false, form="tail", breakInline=false, bodyInline=false, barrierFirst=false) %}
63
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
64
  loop {
65
- {% if form == "head" %}
66
- {% if breakInline %}
67
- if ({{ svar }} == 0u) { break; }
68
- {% else %}
69
  if ({{ svar }} == 0u) {
70
  break;
71
  }
72
- {% endif %}
73
- {% endif %}
74
- {% if bodyInline %}
75
- if ({{ idx }} < {{ svar }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} }
76
- {% else %}
77
  if ({{ idx }} < {{ svar }}) {
78
  {% for a in arrays %}
79
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
80
  {% endfor %}
81
  }
82
- {% endif %}
83
- {% if form == "head" %}
84
- {% if barrierFirst %}
85
- workgroupBarrier();
86
- {{ svar }} = {{ svar }} / 2u;
87
- {% else %}
88
  {{ svar }} = {{ svar }} / 2u;
89
  workgroupBarrier();
90
- {% endif %}
91
- {% else %}
92
- workgroupBarrier();
93
- if ({{ svar }} == 1u) {
94
- break;
95
- }
96
- {{ svar }} = {{ svar }} / 2u;
97
- {% endif %}
98
- }
99
- {%- endmacro %}
100
-
101
  // Reusing partial after this reduction requires a barrier between the read of
102
  // partial[0] and the next write, or the next round can race the prior readers.
103
- {% set trailingBarrier = trailingBarrier is defined and trailingBarrier %}
104
  fn reduce_sum(value: f32, tid: u32) -> f32 {
105
  partial[tid] = value;
106
  workgroupBarrier();
107
  {{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }}
108
- {% if trailingBarrier %}
109
- let total = partial[0];
110
- workgroupBarrier();
111
- return total;
112
- {% else %}
113
  return partial[0];
114
- {% endif %}
115
  }
116
 
117
-
118
  @compute @workgroup_size(WG, 1, 1)
119
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
120
  let row = wg.x + wg.y * params.rowStride;
 
4
  const EPSILON: f32 = {{ epsilon }};
5
  const WG: u32 = {{ workgroupSize }}u;
6
 
 
 
7
  {% if scaleRank > 0 %}
8
  const X_RANK: u32 = {{ xRank }}u;
9
  const SCALE_RANK: u32 = {{ scaleRank }}u;
 
49
  {% endif %}
50
  }
51
 
52
+ var<workgroup> partial: array<f32, WG>;
53
  {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
54
+ {% if op == "max" or op == "min" %}
55
+ {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
56
+ {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %}
57
+ {% 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) %}
 
 
 
58
  var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
59
  loop {
 
 
 
 
60
  if ({{ svar }} == 0u) {
61
  break;
62
  }
 
 
 
 
 
63
  if ({{ idx }} < {{ svar }}) {
64
  {% for a in arrays %}
65
  {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
66
  {% endfor %}
67
  }
 
 
 
 
 
 
68
  {{ svar }} = {{ svar }} / 2u;
69
  workgroupBarrier();
70
+ }{% endmacro %}
 
 
 
 
 
 
 
 
 
 
71
  // Reusing partial after this reduction requires a barrier between the read of
72
  // partial[0] and the next write, or the next round can race the prior readers.
 
73
  fn reduce_sum(value: f32, tid: u32) -> f32 {
74
  partial[tid] = value;
75
  workgroupBarrier();
76
  {{ wgsl_tree_fold(["partial"], idx="tid", wg="WG", form="head") }}
 
 
 
 
 
77
  return partial[0];
 
78
  }
79
 
 
80
  @compute @workgroup_size(WG, 1, 1)
81
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
82
  let row = wg.x + wg.y * params.rowStride;
build/webgpu/test.json CHANGED
@@ -258,7 +258,7 @@
258
  "provenance": {
259
  "source": "onnxruntime/core/providers/cpu/nn/layer_norm_impl.cc",
260
  "test": "SimplifiedLayerNormalization scalar-scale cast boundary",
261
- "notes": "A scalar scale selects the generic path and requires multiplication before the final float16 cast, independently of optimized row handling."
262
  },
263
  "requires": { "features": ["shader-f16"] },
264
  "attrs": { "epsilon": 0.00001, "axis": -1 },
@@ -289,7 +289,7 @@
289
  "provenance": {
290
  "source": "onnxruntime/core/providers/cpu/nn/layer_norm_impl.cc",
291
  "test": "SimplifiedLayerNormalization split-K f16 cast boundary",
292
- "notes": "Forces the split-K shared kernel on the exact scalar-scale boundary and pins the legacy scale-before-output-cast ordering."
293
  },
294
  "requires": { "features": ["shader-f16"] },
295
  "tunables": { "SPLIT_MIN_HIDDEN": 1, "SPLIT_TARGET_ELEMENTS": 1 },
 
258
  "provenance": {
259
  "source": "onnxruntime/core/providers/cpu/nn/layer_norm_impl.cc",
260
  "test": "SimplifiedLayerNormalization scalar-scale cast boundary",
261
+ "notes": "A scalar scale requires multiplication before the final float16 cast, independent of the row width."
262
  },
263
  "requires": { "features": ["shader-f16"] },
264
  "attrs": { "epsilon": 0.00001, "axis": -1 },
 
289
  "provenance": {
290
  "source": "onnxruntime/core/providers/cpu/nn/layer_norm_impl.cc",
291
  "test": "SimplifiedLayerNormalization split-K f16 cast boundary",
292
+ "notes": "A four-value float16 row with a scalar scale checks the legacy scale-before-float16-output-cast ordering."
293
  },
294
  "requires": { "features": ["shader-f16"] },
295
  "tunables": { "SPLIT_MIN_HIDDEN": 1, "SPLIT_TARGET_ELEMENTS": 1 },