Xenova HF Staff commited on
Commit
fe742ec
·
verified ·
1 Parent(s): 1577a99

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -48,6 +48,8 @@ Default values (overridable per request):
48
 
49
  One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
50
 
 
 
51
  - `strided_axis_serial` — Flatten a single non-last reduction axis into outer/axis/inner geometry. Compile its strides and loop bound, keep one output per lane and float32 accumulation, and cap the workgroup by the device limits.
52
 
53
  ## Device requirements
@@ -59,7 +61,7 @@ Some implementation variants require `subgroups`. These are route-specific capab
59
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
60
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
61
  - [`test.json`](build/webgpu/test.json) — correctness cases
62
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
63
  - [`reduce-axis-split-reduce.wgsl.jinja`](build/webgpu/reduce-axis-split-reduce.wgsl.jinja)
64
  - [`reduce-axis0-splitk-combine.wgsl.jinja`](build/webgpu/reduce-axis0-splitk-combine.wgsl.jinja)
65
  - [`reduce-axis0-splitk-reduce.wgsl.jinja`](build/webgpu/reduce-axis0-splitk-reduce.wgsl.jinja)
@@ -76,7 +78,7 @@ Some implementation variants require `subgroups`. These are route-specific capab
76
  ## Use with `@huggingface/kernels`
77
 
78
  ```sh
79
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
80
  ```
81
 
82
  Outputs with inferable metadata are allocated automatically. Explicit `outputs` entries request optional results or provide metadata that cannot be inferred from the supplied inputs and attributes.
@@ -95,7 +97,7 @@ import { getKernel } from "@huggingface/kernels";
95
 
96
  const kernel = await getKernel("webgpu-kernels/ai.onnx.ReduceMean", { version: 1 });
97
  // Explicit destinations request optional results or supply metadata that cannot be inferred.
98
- const { y } = await kernel({ x: { data: xData, shape: [] } }, {
99
- outputs: { y: { shape: [], dtype: "float32" } },
100
  });
101
  ```
 
48
 
49
  One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
50
 
51
+ - `subgroup_last_axis_vec4` — Reduces each contiguous last-axis row with subgroup collectives and vec4-packed reads.
52
+ - `subgroup_last_axis` — Reduces each contiguous last-axis row with subgroup collectives and scalar reads for an unaligned row width.
53
  - `strided_axis_serial` — Flatten a single non-last reduction axis into outer/axis/inner geometry. Compile its strides and loop bound, keep one output per lane and float32 accumulation, and cap the workgroup by the device limits.
54
 
55
  ## Device requirements
 
61
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
62
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
63
  - [`test.json`](build/webgpu/test.json) — correctness cases
64
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
65
  - [`reduce-axis-split-reduce.wgsl.jinja`](build/webgpu/reduce-axis-split-reduce.wgsl.jinja)
66
  - [`reduce-axis0-splitk-combine.wgsl.jinja`](build/webgpu/reduce-axis0-splitk-combine.wgsl.jinja)
67
  - [`reduce-axis0-splitk-reduce.wgsl.jinja`](build/webgpu/reduce-axis0-splitk-reduce.wgsl.jinja)
 
78
  ## Use with `@huggingface/kernels`
79
 
80
  ```sh
81
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
82
  ```
83
 
84
  Outputs with inferable metadata are allocated automatically. Explicit `outputs` entries request optional results or provide metadata that cannot be inferred from the supplied inputs and attributes.
 
97
 
98
  const kernel = await getKernel("webgpu-kernels/ai.onnx.ReduceMean", { version: 1 });
99
  // Explicit destinations request optional results or supply metadata that cannot be inferred.
100
+ const { y } = await kernel({ x: { data: xData, shape: [3, 2, 2] } }, {
101
+ outputs: { y: { shape: [1, 1, 1], dtype: "float32" } },
102
  });
103
  ```
build/webgpu/bench.json CHANGED
@@ -130,7 +130,7 @@
130
  "name": "reducemean-spatial-axes23-f32-2x64x256x256-globalpool-pattern",
131
  "preset": "stress",
132
  "provenance": {
133
- "source": "repository-authored",
134
  "notes": "Rank-4 contiguous-suffix reduction that verifies the vec4 subgroup reducer and portable workgroup-tree fallback."
135
  },
136
  "vars": { "batch": 2, "channels": 64, "height": 256, "width": 256 },
 
130
  "name": "reducemean-spatial-axes23-f32-2x64x256x256-globalpool-pattern",
131
  "preset": "stress",
132
  "provenance": {
133
+ "source": "synthetic",
134
  "notes": "Rank-4 contiguous-suffix reduction that verifies the vec4 subgroup reducer and portable workgroup-tree fallback."
135
  },
136
  "vars": { "batch": 2, "channels": 64, "height": 256, "width": 256 },
build/webgpu/manifest.json CHANGED
@@ -36,6 +36,7 @@
36
  "MULTI_AXIS_COOP_MIN_REDUCED": { "default": 256 }
37
  },
38
  "derive": {
 
39
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
40
  "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32",
41
  "reportedNonWave32Adapter": "not wave32Adapter and (has(device.adapterInfo, \"subgroupMinSize\") or has(device.adapterInfo, \"subgroupMaxSize\"))",
@@ -63,149 +64,147 @@
63
  "flatScratchBytes": "flatSplitCount * dtypeBytes(\"float32\")",
64
  "flatPathFits": "treeWorkgroupOk and flatSplitCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and flatScratchBytes <= device.limits.maxStorageBufferBindingSize and flatScratchBytes <= device.limits.maxBufferSize",
65
  "flatParallelCovered": "(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T) and numel(shapes.y) == 1 and numel(shapes.x) >= tunables.FULL_REDUCE_MIN_ELEMENTS and flatPathFits",
66
- "contiguousSuffixParallelCovered": "(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T) and numel(shapes.y) > 0 and numel(shapes.x) % numel(shapes.y) == 0 and numel(shapes.x) / numel(shapes.y) >= tunables.CONTIGUOUS_SUFFIX_MIN_COLS and ((ranks.x == 3 and hasAxis(attrs.axes, 0, 3) == false and hasAxis(attrs.axes, 1, 3) and hasAxis(attrs.axes, 2, 3) and numel(shapes.y) == dim(shapes.x, 0)) or (ranks.x == 4 and hasAxis(attrs.axes, 0, 4) == false and hasAxis(attrs.axes, 1, 4) == false and hasAxis(attrs.axes, 2, 4) and hasAxis(attrs.axes, 3, 4) and numel(shapes.y) == dim(shapes.x, 0) * dim(shapes.x, 1)) or (ranks.x == 4 and hasAxis(attrs.axes, 0, 4) == false and hasAxis(attrs.axes, 1, 4) and hasAxis(attrs.axes, 2, 4) and hasAxis(attrs.axes, 3, 4) and numel(shapes.y) == dim(shapes.x, 0)))"
 
 
67
  },
68
  "bindings": {
69
- "x": { "buffer": "read-only-storage", "elementType": "$vectorScalar" },
70
- "y": { "buffer": "storage", "elementType": "$T" },
71
- "params": {
72
- "buffer": "uniform",
73
  "struct": [
74
  { "name": "rows", "type": "u32", "value": "numel(shapes.y)" },
75
  { "name": "cols", "type": "u32", "value": "numel(shapes.x) / numel(shapes.y)" },
76
  { "name": "chunkCount", "type": "u32", "value": "numel(shapes.x) / numel(shapes.y) / tunables.VECTOR_WIDTH" }
77
  ]
78
  },
79
- "x_2": { "name": "x", "buffer": "read-only-storage", "elementType": "$T" },
80
- "params_2": {
81
  "name": "params",
82
- "buffer": "uniform",
83
  "struct": [
84
- { "name": "rows", "type": "u32", "value": "numel(shapes.y)" },
85
- { "name": "cols", "type": "u32", "value": "numel(shapes.x) / numel(shapes.y)" }
 
86
  ]
87
  },
88
- "params_3": {
89
  "name": "params",
90
- "buffer": "uniform",
91
- "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.y)" }]
 
 
92
  },
93
- "params_4": {
94
  "name": "params",
95
- "buffer": "uniform",
96
- "struct": [{ "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }]
 
 
97
  },
98
- "params_5": {
 
 
 
 
 
 
 
 
 
99
  "name": "params",
100
- "buffer": "uniform",
101
  "struct": [
102
  { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
103
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" },
104
- { "name": "chunkCount", "type": "u32", "value": "dim(shapes.x, ranks.x - 1) / tunables.VECTOR_WIDTH" }
105
  ]
106
  },
107
- "params_6": {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
108
  "name": "params",
109
- "buffer": "uniform",
110
  "struct": [
111
  { "name": "rows", "type": "u32", "value": "1" },
112
  { "name": "cols", "type": "u32", "value": "1" },
113
  { "name": "outCount", "type": "u32", "value": "1" }
114
  ]
115
  },
116
- "params_7": {
117
  "name": "params",
118
- "buffer": "uniform",
119
  "struct": [
120
  { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
121
  { "name": "cols", "type": "u32", "value": "1" },
122
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
123
  ]
124
  },
125
- "params_8": {
126
  "name": "params",
127
- "buffer": "uniform",
128
  "struct": [
129
  { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
130
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" }
131
  ]
132
  },
133
  "partials": { "buffer": "storage", "elementType": "$partialElement" },
134
- "params_9": {
135
  "name": "params",
136
- "buffer": "uniform",
137
  "struct": [
138
  { "name": "axisDim", "type": "u32", "value": "axisSplitDim" },
139
  { "name": "inner", "type": "u32", "value": "axisSplitInner" },
140
  { "name": "outputs", "type": "u32", "value": "axisSplitOutputs" }
141
  ]
142
  },
143
- "partials_2": { "name": "partials", "buffer": "read-only-storage", "elementType": "$partialElement" },
144
- "params_10": {
145
- "name": "params",
146
- "buffer": "uniform",
147
- "struct": [
148
- { "name": "rows", "type": "u32", "value": "axisSplitDim" },
149
- { "name": "cols", "type": "u32", "value": "axisSplitOutputs" }
150
- ]
151
- },
152
- "params_11": {
153
  "name": "params",
154
- "buffer": "uniform",
155
  "struct": [
156
  { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
157
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" }
158
  ]
159
  },
160
- "partials_3": { "name": "partials", "buffer": "storage", "elementType": "f32" },
161
- "params_12": {
162
  "name": "params",
163
- "buffer": "uniform",
164
  "struct": [
165
  { "name": "count4", "type": "u32", "value": "floor(numel(shapes.x) / tunables.VECTOR_WIDTH)" },
166
  { "name": "numel", "type": "u32", "value": "numel(shapes.x)" }
167
  ]
168
  },
169
- "partials_4": { "name": "partials", "buffer": "read-only-storage", "elementType": "f32" },
170
- "params_13": {
171
- "name": "params",
172
- "buffer": "uniform",
173
- "struct": [
174
- { "name": "rows", "type": "u32", "value": "numel(shapes.x)" },
175
- { "name": "cols", "type": "u32", "value": "1" }
176
- ]
177
- },
178
- "params_15": {
179
  "name": "params",
180
- "buffer": "uniform",
181
  "struct": [
182
  { "name": "axisDim", "type": "u32", "value": "axisSplitDim" },
183
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
184
  ]
185
  },
186
- "params_16": {
187
  "name": "params",
188
- "buffer": "uniform",
189
  "struct": [
190
- { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
191
- { "name": "cols", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" },
192
- { "name": "chunkCount", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" }
193
  ]
194
  },
195
- "params_17": {
196
  "name": "params",
197
- "buffer": "uniform",
198
  "struct": [
199
- { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
200
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" },
201
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
202
  ]
203
  },
204
- "params_18": {
205
  "name": "params",
206
- "buffer": "uniform",
207
  "struct": [
208
- { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" },
 
209
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
210
  ]
211
  }
@@ -226,15 +225,9 @@
226
  "id": "main",
227
  "name": "ReduceMean.ContiguousSuffixSubgroupVec4",
228
  "shader": "reduce-row-subgroup.wgsl.jinja",
229
- "derive": {
230
- "op": "\"mean\"",
231
- "vec4": true,
232
- "castF32": "dtypes.T == \"f16\"",
233
- "usesF16Spec": "dtypes.T == \"f16\""
234
- },
235
- "bindings": ["x", "y", "params"],
236
- "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 },
237
- "subgroupCollectivesWidth": "portable"
238
  }
239
  ]
240
  },
@@ -252,13 +245,8 @@
252
  "id": "main",
253
  "name": "ReduceMean.ContiguousSuffixTreeVec4",
254
  "shader": "reduce-row-tree.wgsl.jinja",
255
- "derive": {
256
- "op": "\"mean\"",
257
- "vec4": true,
258
- "castF32": "dtypes.T == \"f16\"",
259
- "usesF16Spec": "dtypes.T == \"f16\""
260
- },
261
- "bindings": ["x", "y", "params"],
262
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
263
  }
264
  ]
@@ -276,8 +264,7 @@
276
  "id": "main",
277
  "name": "ReduceMean.ContiguousSuffixTree",
278
  "shader": "reduce-row-tree.wgsl.jinja",
279
- "derive": { "op": "\"mean\"", "castF32": "dtypes.T == \"f16\"", "usesF16Spec": "dtypes.T == \"f16\"" },
280
- "bindings": ["x_2", "y", "params_2"],
281
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
282
  }
283
  ]
@@ -292,7 +279,7 @@
292
  "id": "main",
293
  "name": "ReduceMean.NoopEmptyAxes",
294
  "shader": "reduce-noop-empty-axes.wgsl.jinja",
295
- "bindings": ["x_2", "y", "params_3"],
296
  "dispatch": {
297
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
298
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -312,7 +299,6 @@
312
  "name": "ReduceMean.MultiAxisRank3Coop",
313
  "shader": "reduce-multi-axis-coop.wgsl.jinja",
314
  "derive": {
315
- "op": "\"mean\"",
316
  "indexing": "\"multiaxis\"",
317
  "rank": 3,
318
  "reduce": ["hasAxis(attrs.axes, 0, 3)", "hasAxis(attrs.axes, 1, 3)", "hasAxis(attrs.axes, 2, 3)"],
@@ -320,11 +306,9 @@
320
  "outputShape": "shapes.y",
321
  "outputRank": "ranks.y",
322
  "keepDims": "attrs.keepdims != 0",
323
- "intMode": "dtypes.T == \"i32\"",
324
- "castF32": "dtypes.T == \"f16\"",
325
- "usesF16Spec": "dtypes.T == \"f16\""
326
  },
327
- "bindings": ["x_2", "y"],
328
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
329
  }
330
  ]
@@ -340,7 +324,6 @@
340
  "name": "ReduceMean.MultiAxisRank4Coop",
341
  "shader": "reduce-multi-axis-coop.wgsl.jinja",
342
  "derive": {
343
- "op": "\"mean\"",
344
  "indexing": "\"multiaxis\"",
345
  "rank": 4,
346
  "reduce": ["hasAxis(attrs.axes, 0, 4)", "hasAxis(attrs.axes, 1, 4)", "hasAxis(attrs.axes, 2, 4)", "hasAxis(attrs.axes, 3, 4)"],
@@ -348,11 +331,9 @@
348
  "outputShape": "shapes.y",
349
  "outputRank": "ranks.y",
350
  "keepDims": "attrs.keepdims != 0",
351
- "intMode": "dtypes.T == \"i32\"",
352
- "castF32": "dtypes.T == \"f16\"",
353
- "usesF16Spec": "dtypes.T == \"f16\""
354
  },
355
- "bindings": ["x_2", "y"],
356
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
357
  }
358
  ]
@@ -368,7 +349,6 @@
368
  "name": "ReduceMean.MultiAxisRank3",
369
  "shader": "reduce-serial-axis.wgsl.jinja",
370
  "derive": {
371
- "op": "\"mean\"",
372
  "indexing": "\"multiaxis\"",
373
  "rank": 3,
374
  "reduce": ["hasAxis(attrs.axes, 0, 3)", "hasAxis(attrs.axes, 1, 3)", "hasAxis(attrs.axes, 2, 3)"],
@@ -376,11 +356,9 @@
376
  "outputShape": "shapes.y",
377
  "outputRank": "ranks.y",
378
  "keepDims": "attrs.keepdims != 0",
379
- "intMode": "dtypes.T == \"i32\"",
380
- "castF32": "dtypes.T == \"f16\"",
381
- "usesF16Spec": "dtypes.T == \"f16\""
382
  },
383
- "bindings": ["x_2", "y", "params_4"],
384
  "dispatch": {
385
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
386
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -400,7 +378,6 @@
400
  "name": "ReduceMean.MultiAxisRank4",
401
  "shader": "reduce-serial-axis.wgsl.jinja",
402
  "derive": {
403
- "op": "\"mean\"",
404
  "indexing": "\"multiaxis\"",
405
  "rank": 4,
406
  "reduce": ["hasAxis(attrs.axes, 0, 4)", "hasAxis(attrs.axes, 1, 4)", "hasAxis(attrs.axes, 2, 4)", "hasAxis(attrs.axes, 3, 4)"],
@@ -408,11 +385,9 @@
408
  "outputShape": "shapes.y",
409
  "outputRank": "ranks.y",
410
  "keepDims": "attrs.keepdims != 0",
411
- "intMode": "dtypes.T == \"i32\"",
412
- "castF32": "dtypes.T == \"f16\"",
413
- "usesF16Spec": "dtypes.T == \"f16\""
414
  },
415
- "bindings": ["x_2", "y", "params_4"],
416
  "dispatch": {
417
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
418
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -437,14 +412,12 @@
437
  "id": "main",
438
  "name": "ReduceMean.SubgroupRowsVec4",
439
  "shader": "reduce-row-subgroup-rows.wgsl.jinja",
440
- "derive": { "op": "\"mean\"", "castF32": "dtypes.T == \"f16\"", "usesF16Spec": "dtypes.T == \"f16\"" },
441
- "bindings": ["x", "y", "params_5"],
442
  "dispatch": {
443
  "x": "min(ceilDiv(lastAxisRows, reduceWorkgroupSize / device.adapterInfo.subgroupMaxSize), 65535)",
444
  "y": "ceilDiv(ceilDiv(lastAxisRows, reduceWorkgroupSize / device.adapterInfo.subgroupMaxSize), 65535)",
445
  "z": 1
446
- },
447
- "subgroupCollectivesWidth": "portable"
448
  }
449
  ]
450
  },
@@ -463,13 +436,8 @@
463
  "id": "main",
464
  "name": "ReduceMean.TreeRowVec4",
465
  "shader": "reduce-row-tree.wgsl.jinja",
466
- "derive": {
467
- "op": "\"mean\"",
468
- "vec4": true,
469
- "castF32": "dtypes.T == \"f16\"",
470
- "usesF16Spec": "dtypes.T == \"f16\""
471
- },
472
- "bindings": ["x", "y", "params_5"],
473
  "dispatch": { "x": "min(lastAxisRows, 65535)", "y": "ceilDiv(lastAxisRows, 65535)", "z": 1 }
474
  }
475
  ]
@@ -485,14 +453,11 @@
485
  "name": "ReduceMean.Rank0Scalar",
486
  "shader": "reduce-serial-axis.wgsl.jinja",
487
  "derive": {
488
- "op": "\"mean\"",
489
  "indexing": "\"axis2d\"",
490
  "intMode": "dtypes.T == \"i32\"",
491
- "castF32": "dtypes.T == \"f16\"",
492
- "usesF16Spec": "dtypes.T == \"f16\"",
493
  "logicalBool": "tensorDtypes.x == \"bool\""
494
  },
495
- "bindings": ["x_2", "y", "params_6"],
496
  "dispatch": { "x": 1 }
497
  }
498
  ]
@@ -507,14 +472,11 @@
507
  "name": "ReduceMean.Rank1Axis0",
508
  "shader": "reduce-serial-axis.wgsl.jinja",
509
  "derive": {
510
- "op": "\"mean\"",
511
  "indexing": "\"axis2d\"",
512
  "intMode": "dtypes.T == \"i32\"",
513
- "castF32": "dtypes.T == \"f16\"",
514
- "usesF16Spec": "dtypes.T == \"f16\"",
515
  "logicalBool": "tensorDtypes.x == \"bool\""
516
  },
517
- "bindings": ["x_2", "y", "params_7"],
518
  "dispatch": {
519
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
520
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -537,8 +499,7 @@
537
  "id": "main",
538
  "name": "ReduceMean.Axis1Parallel",
539
  "shader": "reduce-row-tree.wgsl.jinja",
540
- "derive": { "op": "\"mean\"", "castF32": "dtypes.T == \"f16\"", "usesF16Spec": "dtypes.T == \"f16\"" },
541
- "bindings": ["x_2", "y", "params_8"],
542
  "dispatch": {
543
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
544
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
@@ -563,13 +524,8 @@
563
  "id": "split_reduce",
564
  "name": "ReduceMean.AxisSplitReduce",
565
  "shader": "reduce-axis-split-reduce.wgsl.jinja",
566
- "derive": {
567
- "op": "\"mean\"",
568
- "splitSpec": "splitCount",
569
- "castF32": "dtypes.T == \"f16\"",
570
- "usesF16Spec": "dtypes.T == \"f16\""
571
- },
572
- "bindings": ["x_2", "partials", "params_9"],
573
  "dispatch": {
574
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), DISPATCH_FOLD_WIDTH)",
575
  "y": "splitCount",
@@ -580,8 +536,8 @@
580
  "id": "combine",
581
  "name": "ReduceMean.AxisSplitCombine",
582
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
583
- "derive": { "op": "\"mean\"", "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
584
- "bindings": ["partials_2", "y", "params_10"],
585
  "dispatch": {
586
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
587
  "y": "ceilDiv(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
@@ -608,14 +564,8 @@
608
  "id": "split_reduce",
609
  "name": "ReduceMean.AxisSplitTiledReduce",
610
  "shader": "reduce-axis0-tilecols.wgsl.jinja",
611
- "derive": {
612
- "op": "\"mean\"",
613
- "splitSpec": "splitCount",
614
- "tileCols": "tunables.AXIS_SPLIT_TILE_COLS",
615
- "castF32": "dtypes.T == \"f16\"",
616
- "usesF16Spec": "dtypes.T == \"f16\""
617
- },
618
- "bindings": ["x_2", "partials", "params_9"],
619
  "dispatch": {
620
  "x": "min(ceilDiv((axisSplitOutputs), (tileCols)), DISPATCH_FOLD_WIDTH)",
621
  "y": "splitCount",
@@ -626,8 +576,8 @@
626
  "id": "combine",
627
  "name": "ReduceMean.AxisSplitCombine",
628
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
629
- "derive": { "op": "\"mean\"", "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
630
- "bindings": ["partials_2", "y", "params_10"],
631
  "dispatch": {
632
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
633
  "y": "ceilDiv(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
@@ -640,6 +590,7 @@
640
  "id": "axis0_splitk",
641
  "priority": 22,
642
  "when": ["not flatParallelCovered", "(dtypes.T == \"f32\" or dtypes.T == \"f16\")", "f16Ok(dtypes.T)", "ranks.x == 2", "reduceAxis == 0", "axis0Rows >= tunables.AXIS0_SPLIT_MIN_ROWS", "dim(shapes.x, 1) > 0", "((attrs.keepdims == 0 and ranks.y == 1 and dim(shapes.y, 0) == dim(shapes.x, 1)) or (attrs.keepdims == 1 and ranks.y == 2 and dim(shapes.y, 0) == 1 and dim(shapes.y, 1) == dim(shapes.x, 1)))", "axis0SplitPathFits"],
 
643
  "derive": {
644
  "splitCount": "axis0SplitCount",
645
  "partialElement": "\"f32\"",
@@ -652,13 +603,8 @@
652
  "id": "split_reduce",
653
  "name": "ReduceMean.Axis0SplitKReduce",
654
  "shader": "reduce-axis0-splitk-reduce.wgsl.jinja",
655
- "derive": {
656
- "op": "\"mean\"",
657
- "splitSpec": "splitCount",
658
- "castF32": "dtypes.T == \"f16\"",
659
- "usesF16Spec": "dtypes.T == \"f16\""
660
- },
661
- "bindings": ["x_2", "partials", "params_11"],
662
  "dispatch": {
663
  "x": "min(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), DISPATCH_FOLD_WIDTH)",
664
  "y": "splitCount",
@@ -669,8 +615,8 @@
669
  "id": "combine",
670
  "name": "ReduceMean.Axis0SplitKCombine",
671
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
672
- "derive": { "op": "\"mean\"", "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
673
- "bindings": ["partials_2", "y", "params_11"],
674
  "dispatch": {
675
  "x": "min(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), 65535)",
676
  "y": "ceilDiv(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), 65535)",
@@ -689,13 +635,8 @@
689
  "id": "main",
690
  "name": "ReduceMean.Axis0TileCols",
691
  "shader": "reduce-axis0-tilecols.wgsl.jinja",
692
- "derive": {
693
- "op": "\"mean\"",
694
- "intMode": "dtypes.T == \"i32\"",
695
- "castF32": "dtypes.T == \"f16\"",
696
- "usesF16Spec": "dtypes.T == \"f16\""
697
- },
698
- "bindings": ["x_2", "y", "params_11"],
699
  "dispatch": {
700
  "x": "min(ceilDiv((dim(shapes.x, 1)), (tileCols)), 65535)",
701
  "y": "ceilDiv(ceilDiv((dim(shapes.x, 1)), (tileCols)), 65535)",
@@ -708,23 +649,22 @@
708
  "id": "all_axes_flat",
709
  "priority": 31,
710
  "when": ["flatParallelCovered"],
711
- "derive": { "scalar": "dtypes.T", "workgroupSize": "reduceWorkgroupSize", "split": "flatSplitCount" },
712
  "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[flatSplitCount]" }],
713
  "passes": [
714
  {
715
  "id": "flat_partial",
716
  "name": "ReduceMean.AllAxesFlatPartial",
717
  "shader": "reduce-flat-partial.wgsl.jinja",
718
- "derive": { "op": "\"mean\"", "castF32": "dtypes.T == \"f16\"", "usesF16Spec": "dtypes.T == \"f16\"" },
719
- "bindings": ["x_2", "partials_3", "params_12"],
720
  "dispatch": { "x": "flatSplitCount" }
721
  },
722
  {
723
  "id": "combine",
724
  "name": "ReduceMean.AllAxesFlatCombine",
725
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
726
- "derive": { "op": "\"mean\"", "outputF16": "dtypes.T == \"f16\"" },
727
- "bindings": ["partials_4", "y", "params_13"],
728
  "dispatch": { "x": 1 }
729
  }
730
  ]
@@ -777,7 +717,6 @@
777
  "name": "ReduceMean.RankNSingleAxisGeneric",
778
  "shader": "reduce-serial-axis.wgsl.jinja",
779
  "derive": {
780
- "op": "\"mean\"",
781
  "indexing": "\"rankn\"",
782
  "rank": "ranks.x",
783
  "axisSpec": "reduceAxis",
@@ -785,11 +724,9 @@
785
  "outputShape": "shapes.y",
786
  "outputRank": "ranks.y",
787
  "keepDims": "attrs.keepdims != 0",
788
- "intMode": "dtypes.T == \"i32\"",
789
- "castF32": "dtypes.T == \"f16\"",
790
- "usesF16Spec": "dtypes.T == \"f16\""
791
  },
792
- "bindings": ["x_2", "y", "params_15"],
793
  "dispatch": {
794
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
795
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -813,19 +750,13 @@
813
  "id": "main",
814
  "name": "ReduceMean.SubgroupRowVec4",
815
  "shader": "reduce-row-subgroup.wgsl.jinja",
816
- "derive": {
817
- "op": "\"mean\"",
818
- "vec4": true,
819
- "castF32": "dtypes.T == \"f16\"",
820
- "usesF16Spec": "dtypes.T == \"f16\""
821
- },
822
- "bindings": ["x", "y", "params_5"],
823
  "dispatch": {
824
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
825
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
826
  "z": 1
827
- },
828
- "subgroupCollectivesWidth": "portable"
829
  }
830
  ]
831
  },
@@ -843,19 +774,13 @@
843
  "id": "main",
844
  "name": "ReduceMean.SubgroupRow",
845
  "shader": "reduce-row-subgroup.wgsl.jinja",
846
- "derive": {
847
- "op": "\"mean\"",
848
- "vec4": false,
849
- "castF32": "dtypes.T == \"f16\"",
850
- "usesF16Spec": "dtypes.T == \"f16\""
851
- },
852
- "bindings": ["x_2", "y", "params_16"],
853
  "dispatch": {
854
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
855
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
856
  "z": 1
857
- },
858
- "subgroupCollectivesWidth": "portable"
859
  }
860
  ]
861
  },
@@ -868,17 +793,10 @@
868
  "passes": [
869
  {
870
  "id": "main",
871
- "name": "axis0",
872
  "shader": "reduce-serial-axis.wgsl.jinja",
873
- "derive": {
874
- "axis": 0,
875
- "op": "\"mean\"",
876
- "indexing": "\"axis2d\"",
877
- "intMode": "dtypes.T == \"i32\"",
878
- "castF32": "dtypes.T == \"f16\"",
879
- "usesF16Spec": "dtypes.T == \"f16\""
880
- },
881
- "bindings": ["x_2", "y", "params_17"],
882
  "dispatch": {
883
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
884
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -895,17 +813,10 @@
895
  "passes": [
896
  {
897
  "id": "main",
898
- "name": "axis1",
899
  "shader": "reduce-serial-axis.wgsl.jinja",
900
- "derive": {
901
- "axis": 1,
902
- "op": "\"mean\"",
903
- "indexing": "\"axis2d\"",
904
- "intMode": "dtypes.T == \"i32\"",
905
- "castF32": "dtypes.T == \"f16\"",
906
- "usesF16Spec": "dtypes.T == \"f16\""
907
- },
908
- "bindings": ["x_2", "y", "params_18"],
909
  "dispatch": {
910
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
911
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -915,34 +826,17 @@
915
  ]
916
  },
917
  {
918
- "id": "all_axes_no_keepdims",
919
  "priority": 30,
920
- "when": ["not flatParallelCovered", "f16Ok(dtypes.T)", "ranks.x >= 3", "attrs.keepdims == 0 and ranks.y == 0"],
921
  "derive": { "axis": 0, "scalar": "dtypes.T" },
922
  "passes": [
923
  {
924
  "id": "main",
925
- "name": "ReduceMean.Rank3AllAxesNoKeepdims",
926
  "shader": "reduce-serial-axis.wgsl.jinja",
927
- "derive": {
928
- "op": "\"mean\"",
929
- "indexing": "\"axis2d\"",
930
- "intMode": "dtypes.T == \"i32\"",
931
- "castF32": "dtypes.T == \"f16\"",
932
- "usesF16Spec": "dtypes.T == \"f16\""
933
- },
934
- "bindings": [
935
- "x_2",
936
- "y",
937
- {
938
- "name": "params",
939
- "struct": [
940
- { "name": "rows", "type": "u32", "value": "numel(shapes.x)" },
941
- { "name": "cols", "type": "u32", "value": "1" },
942
- { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
943
- ]
944
- }
945
- ],
946
  "dispatch": {
947
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
948
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -952,34 +846,17 @@
952
  ]
953
  },
954
  {
955
- "id": "all_axes_keepdims",
956
  "priority": 30,
957
- "when": ["not flatParallelCovered", "f16Ok(dtypes.T)", "ranks.x >= 3", "attrs.keepdims == 1 and ranks.y == ranks.x and numel(shapes.y) == 1"],
958
  "derive": { "axis": 0, "scalar": "dtypes.T" },
959
  "passes": [
960
  {
961
  "id": "main",
962
- "name": "ReduceMean.Rank3AllAxesKeepdims",
963
  "shader": "reduce-serial-axis.wgsl.jinja",
964
- "derive": {
965
- "op": "\"mean\"",
966
- "indexing": "\"axis2d\"",
967
- "intMode": "dtypes.T == \"i32\"",
968
- "castF32": "dtypes.T == \"f16\"",
969
- "usesF16Spec": "dtypes.T == \"f16\""
970
- },
971
- "bindings": [
972
- "x_2",
973
- "y",
974
- {
975
- "name": "params",
976
- "struct": [
977
- { "name": "rows", "type": "u32", "value": "numel(shapes.x)" },
978
- { "name": "cols", "type": "u32", "value": "1" },
979
- { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
980
- ]
981
- }
982
- ],
983
  "dispatch": {
984
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
985
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -1003,7 +880,6 @@
1003
  "axisDim": "axisSplitDim",
1004
  "innerSize": "axisSplitInner",
1005
  "usesF16": "dtypes.T == \"f16\"",
1006
- "op": "\"mean\"",
1007
  "tailPaddingElements": "numel(shapes.y) % 2 if dtypes.T == \"f16\" else 0"
1008
  },
1009
  "bindings": [{ "arg": "x", "elementType": "$scalar" }, { "arg": "y", "elementType": "$scalar" }],
 
36
  "MULTI_AXIS_COOP_MIN_REDUCED": { "default": 256 }
37
  },
38
  "derive": {
39
+ "reduceOp": "\"mean\"",
40
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
41
  "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32",
42
  "reportedNonWave32Adapter": "not wave32Adapter and (has(device.adapterInfo, \"subgroupMinSize\") or has(device.adapterInfo, \"subgroupMaxSize\"))",
 
64
  "flatScratchBytes": "flatSplitCount * dtypeBytes(\"float32\")",
65
  "flatPathFits": "treeWorkgroupOk and flatSplitCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and flatScratchBytes <= device.limits.maxStorageBufferBindingSize and flatScratchBytes <= device.limits.maxBufferSize",
66
  "flatParallelCovered": "(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T) and numel(shapes.y) == 1 and numel(shapes.x) >= tunables.FULL_REDUCE_MIN_ELEMENTS and flatPathFits",
67
+ "contiguousSuffixParallelCovered": "(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T) and numel(shapes.y) > 0 and numel(shapes.x) % numel(shapes.y) == 0 and numel(shapes.x) / numel(shapes.y) >= tunables.CONTIGUOUS_SUFFIX_MIN_COLS and ((ranks.x == 3 and hasAxis(attrs.axes, 0, 3) == false and hasAxis(attrs.axes, 1, 3) and hasAxis(attrs.axes, 2, 3) and numel(shapes.y) == dim(shapes.x, 0)) or (ranks.x == 4 and hasAxis(attrs.axes, 0, 4) == false and hasAxis(attrs.axes, 1, 4) == false and hasAxis(attrs.axes, 2, 4) and hasAxis(attrs.axes, 3, 4) and numel(shapes.y) == dim(shapes.x, 0) * dim(shapes.x, 1)) or (ranks.x == 4 and hasAxis(attrs.axes, 0, 4) == false and hasAxis(attrs.axes, 1, 4) and hasAxis(attrs.axes, 2, 4) and hasAxis(attrs.axes, 3, 4) and numel(shapes.y) == dim(shapes.x, 0)))",
68
+ "op": "reduceOp",
69
+ "castF32": "dtypes.T == \"f16\""
70
  },
71
  "bindings": {
72
+ "params_suffix_vec4": {
73
+ "name": "params",
 
 
74
  "struct": [
75
  { "name": "rows", "type": "u32", "value": "numel(shapes.y)" },
76
  { "name": "cols", "type": "u32", "value": "numel(shapes.x) / numel(shapes.y)" },
77
  { "name": "chunkCount", "type": "u32", "value": "numel(shapes.x) / numel(shapes.y) / tunables.VECTOR_WIDTH" }
78
  ]
79
  },
80
+ "params_reduce_row_tree": {
 
81
  "name": "params",
 
82
  "struct": [
83
+ { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
84
+ { "name": "cols", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" },
85
+ { "name": "chunkCount", "type": "u32", "value": "dim(shapes.x, ranks.x - 1) / tunables.VECTOR_WIDTH" }
86
  ]
87
  },
88
+ "params_axis_split_combine": {
89
  "name": "params",
90
+ "struct": [
91
+ { "name": "rows", "type": "u32", "value": "axisSplitDim" },
92
+ { "name": "cols", "type": "u32", "value": "axisSplitOutputs" }
93
+ ]
94
  },
95
+ "params_axis0_combine": {
96
  "name": "params",
97
+ "struct": [
98
+ { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
99
+ { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" }
100
+ ]
101
  },
102
+ "partials_flat": { "name": "partials", "buffer": "storage", "elementType": "f32" },
103
+ "partials_flat_read": { "name": "partials", "buffer": "read-only-storage", "elementType": "f32" },
104
+ "params_all_axes_flat": {
105
+ "name": "params",
106
+ "struct": [
107
+ { "name": "rows", "type": "u32", "value": "numel(shapes.x)" },
108
+ { "name": "cols", "type": "u32", "value": "1" }
109
+ ]
110
+ },
111
+ "params_rows_chunk_count": {
112
  "name": "params",
 
113
  "struct": [
114
  { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
115
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" },
116
+ { "name": "chunkCount", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" }
117
  ]
118
  },
119
+ "x": { "elementType": "$vectorScalar" },
120
+ "y": { "elementType": "$T" },
121
+ "x_t": { "name": "x", "elementType": "$T" },
122
+ "params_main": {
123
+ "name": "params",
124
+ "struct": [
125
+ { "name": "rows", "type": "u32", "value": "numel(shapes.y)" },
126
+ { "name": "cols", "type": "u32", "value": "numel(shapes.x) / numel(shapes.y)" }
127
+ ]
128
+ },
129
+ "params_count": { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.y)" }] },
130
+ "params_out_count": {
131
+ "name": "params",
132
+ "struct": [{ "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }]
133
+ },
134
+ "params_rank0_scalar": {
135
  "name": "params",
 
136
  "struct": [
137
  { "name": "rows", "type": "u32", "value": "1" },
138
  { "name": "cols", "type": "u32", "value": "1" },
139
  { "name": "outCount", "type": "u32", "value": "1" }
140
  ]
141
  },
142
+ "params_rank1_axis0": {
143
  "name": "params",
 
144
  "struct": [
145
  { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
146
  { "name": "cols", "type": "u32", "value": "1" },
147
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
148
  ]
149
  },
150
+ "params_rows_cols": {
151
  "name": "params",
 
152
  "struct": [
153
  { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
154
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" }
155
  ]
156
  },
157
  "partials": { "buffer": "storage", "elementType": "$partialElement" },
158
+ "params_axis_split": {
159
  "name": "params",
 
160
  "struct": [
161
  { "name": "axisDim", "type": "u32", "value": "axisSplitDim" },
162
  { "name": "inner", "type": "u32", "value": "axisSplitInner" },
163
  { "name": "outputs", "type": "u32", "value": "axisSplitOutputs" }
164
  ]
165
  },
166
+ "partials_combine": { "name": "partials", "buffer": "read-only-storage", "elementType": "$partialElement" },
167
+ "params_split_reduce": {
 
 
 
 
 
 
 
 
168
  "name": "params",
 
169
  "struct": [
170
  { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
171
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" }
172
  ]
173
  },
174
+ "params_count4_numel": {
 
175
  "name": "params",
 
176
  "struct": [
177
  { "name": "count4", "type": "u32", "value": "floor(numel(shapes.x) / tunables.VECTOR_WIDTH)" },
178
  { "name": "numel", "type": "u32", "value": "numel(shapes.x)" }
179
  ]
180
  },
181
+ "params_axis_dim_out_count": {
 
 
 
 
 
 
 
 
 
182
  "name": "params",
 
183
  "struct": [
184
  { "name": "axisDim", "type": "u32", "value": "axisSplitDim" },
185
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
186
  ]
187
  },
188
+ "params_axis0": {
189
  "name": "params",
 
190
  "struct": [
191
+ { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
192
+ { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" },
193
+ { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
194
  ]
195
  },
196
+ "params_axis1": {
197
  "name": "params",
 
198
  "struct": [
 
199
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" },
200
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
201
  ]
202
  },
203
+ "params_all_axes": {
204
  "name": "params",
 
205
  "struct": [
206
+ { "name": "rows", "type": "u32", "value": "numel(shapes.x)" },
207
+ { "name": "cols", "type": "u32", "value": "1" },
208
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
209
  ]
210
  }
 
225
  "id": "main",
226
  "name": "ReduceMean.ContiguousSuffixSubgroupVec4",
227
  "shader": "reduce-row-subgroup.wgsl.jinja",
228
+ "derive": { "vec4": true },
229
+ "bindings": ["x", "y", "params_suffix_vec4"],
230
+ "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
 
 
 
 
 
 
231
  }
232
  ]
233
  },
 
245
  "id": "main",
246
  "name": "ReduceMean.ContiguousSuffixTreeVec4",
247
  "shader": "reduce-row-tree.wgsl.jinja",
248
+ "derive": { "vec4": true },
249
+ "bindings": ["x", "y", "params_suffix_vec4"],
 
 
 
 
 
250
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
251
  }
252
  ]
 
264
  "id": "main",
265
  "name": "ReduceMean.ContiguousSuffixTree",
266
  "shader": "reduce-row-tree.wgsl.jinja",
267
+ "bindings": ["x_t", "y", "params_main"],
 
268
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
269
  }
270
  ]
 
279
  "id": "main",
280
  "name": "ReduceMean.NoopEmptyAxes",
281
  "shader": "reduce-noop-empty-axes.wgsl.jinja",
282
+ "bindings": ["x_t", "y", "params_count"],
283
  "dispatch": {
284
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
285
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
299
  "name": "ReduceMean.MultiAxisRank3Coop",
300
  "shader": "reduce-multi-axis-coop.wgsl.jinja",
301
  "derive": {
 
302
  "indexing": "\"multiaxis\"",
303
  "rank": 3,
304
  "reduce": ["hasAxis(attrs.axes, 0, 3)", "hasAxis(attrs.axes, 1, 3)", "hasAxis(attrs.axes, 2, 3)"],
 
306
  "outputShape": "shapes.y",
307
  "outputRank": "ranks.y",
308
  "keepDims": "attrs.keepdims != 0",
309
+ "intMode": "dtypes.T == \"i32\""
 
 
310
  },
311
+ "bindings": ["x_t", "y"],
312
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
313
  }
314
  ]
 
324
  "name": "ReduceMean.MultiAxisRank4Coop",
325
  "shader": "reduce-multi-axis-coop.wgsl.jinja",
326
  "derive": {
 
327
  "indexing": "\"multiaxis\"",
328
  "rank": 4,
329
  "reduce": ["hasAxis(attrs.axes, 0, 4)", "hasAxis(attrs.axes, 1, 4)", "hasAxis(attrs.axes, 2, 4)", "hasAxis(attrs.axes, 3, 4)"],
 
331
  "outputShape": "shapes.y",
332
  "outputRank": "ranks.y",
333
  "keepDims": "attrs.keepdims != 0",
334
+ "intMode": "dtypes.T == \"i32\""
 
 
335
  },
336
+ "bindings": ["x_t", "y"],
337
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
338
  }
339
  ]
 
349
  "name": "ReduceMean.MultiAxisRank3",
350
  "shader": "reduce-serial-axis.wgsl.jinja",
351
  "derive": {
 
352
  "indexing": "\"multiaxis\"",
353
  "rank": 3,
354
  "reduce": ["hasAxis(attrs.axes, 0, 3)", "hasAxis(attrs.axes, 1, 3)", "hasAxis(attrs.axes, 2, 3)"],
 
356
  "outputShape": "shapes.y",
357
  "outputRank": "ranks.y",
358
  "keepDims": "attrs.keepdims != 0",
359
+ "intMode": "dtypes.T == \"i32\""
 
 
360
  },
361
+ "bindings": ["x_t", "y", "params_out_count"],
362
  "dispatch": {
363
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
364
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
378
  "name": "ReduceMean.MultiAxisRank4",
379
  "shader": "reduce-serial-axis.wgsl.jinja",
380
  "derive": {
 
381
  "indexing": "\"multiaxis\"",
382
  "rank": 4,
383
  "reduce": ["hasAxis(attrs.axes, 0, 4)", "hasAxis(attrs.axes, 1, 4)", "hasAxis(attrs.axes, 2, 4)", "hasAxis(attrs.axes, 3, 4)"],
 
385
  "outputShape": "shapes.y",
386
  "outputRank": "ranks.y",
387
  "keepDims": "attrs.keepdims != 0",
388
+ "intMode": "dtypes.T == \"i32\""
 
 
389
  },
390
+ "bindings": ["x_t", "y", "params_out_count"],
391
  "dispatch": {
392
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
393
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
412
  "id": "main",
413
  "name": "ReduceMean.SubgroupRowsVec4",
414
  "shader": "reduce-row-subgroup-rows.wgsl.jinja",
415
+ "bindings": ["x", "y", "params_reduce_row_tree"],
 
416
  "dispatch": {
417
  "x": "min(ceilDiv(lastAxisRows, reduceWorkgroupSize / device.adapterInfo.subgroupMaxSize), 65535)",
418
  "y": "ceilDiv(ceilDiv(lastAxisRows, reduceWorkgroupSize / device.adapterInfo.subgroupMaxSize), 65535)",
419
  "z": 1
420
+ }
 
421
  }
422
  ]
423
  },
 
436
  "id": "main",
437
  "name": "ReduceMean.TreeRowVec4",
438
  "shader": "reduce-row-tree.wgsl.jinja",
439
+ "derive": { "vec4": true },
440
+ "bindings": ["x", "y", "params_reduce_row_tree"],
 
 
 
 
 
441
  "dispatch": { "x": "min(lastAxisRows, 65535)", "y": "ceilDiv(lastAxisRows, 65535)", "z": 1 }
442
  }
443
  ]
 
453
  "name": "ReduceMean.Rank0Scalar",
454
  "shader": "reduce-serial-axis.wgsl.jinja",
455
  "derive": {
 
456
  "indexing": "\"axis2d\"",
457
  "intMode": "dtypes.T == \"i32\"",
 
 
458
  "logicalBool": "tensorDtypes.x == \"bool\""
459
  },
460
+ "bindings": ["x_t", "y", "params_rank0_scalar"],
461
  "dispatch": { "x": 1 }
462
  }
463
  ]
 
472
  "name": "ReduceMean.Rank1Axis0",
473
  "shader": "reduce-serial-axis.wgsl.jinja",
474
  "derive": {
 
475
  "indexing": "\"axis2d\"",
476
  "intMode": "dtypes.T == \"i32\"",
 
 
477
  "logicalBool": "tensorDtypes.x == \"bool\""
478
  },
479
+ "bindings": ["x_t", "y", "params_rank1_axis0"],
480
  "dispatch": {
481
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
482
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
499
  "id": "main",
500
  "name": "ReduceMean.Axis1Parallel",
501
  "shader": "reduce-row-tree.wgsl.jinja",
502
+ "bindings": ["x_t", "y", "params_rows_cols"],
 
503
  "dispatch": {
504
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
505
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
 
524
  "id": "split_reduce",
525
  "name": "ReduceMean.AxisSplitReduce",
526
  "shader": "reduce-axis-split-reduce.wgsl.jinja",
527
+ "derive": { "splitSpec": "splitCount" },
528
+ "bindings": ["x_t", "partials", "params_axis_split"],
 
 
 
 
 
529
  "dispatch": {
530
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), DISPATCH_FOLD_WIDTH)",
531
  "y": "splitCount",
 
536
  "id": "combine",
537
  "name": "ReduceMean.AxisSplitCombine",
538
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
539
+ "derive": { "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
540
+ "bindings": ["partials_combine", "y", "params_axis_split_combine"],
541
  "dispatch": {
542
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
543
  "y": "ceilDiv(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
 
564
  "id": "split_reduce",
565
  "name": "ReduceMean.AxisSplitTiledReduce",
566
  "shader": "reduce-axis0-tilecols.wgsl.jinja",
567
+ "derive": { "splitSpec": "splitCount" },
568
+ "bindings": ["x_t", "partials", "params_axis_split"],
 
 
 
 
 
 
569
  "dispatch": {
570
  "x": "min(ceilDiv((axisSplitOutputs), (tileCols)), DISPATCH_FOLD_WIDTH)",
571
  "y": "splitCount",
 
576
  "id": "combine",
577
  "name": "ReduceMean.AxisSplitCombine",
578
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
579
+ "derive": { "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
580
+ "bindings": ["partials_combine", "y", "params_axis_split_combine"],
581
  "dispatch": {
582
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
583
  "y": "ceilDiv(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
 
590
  "id": "axis0_splitk",
591
  "priority": 22,
592
  "when": ["not flatParallelCovered", "(dtypes.T == \"f32\" or dtypes.T == \"f16\")", "f16Ok(dtypes.T)", "ranks.x == 2", "reduceAxis == 0", "axis0Rows >= tunables.AXIS0_SPLIT_MIN_ROWS", "dim(shapes.x, 1) > 0", "((attrs.keepdims == 0 and ranks.y == 1 and dim(shapes.y, 0) == dim(shapes.x, 1)) or (attrs.keepdims == 1 and ranks.y == 2 and dim(shapes.y, 0) == 1 and dim(shapes.y, 1) == dim(shapes.x, 1)))", "axis0SplitPathFits"],
593
+ "demoteWhen": ["dtypes.T == \"f32\" and device.adapterInfo.subgroupMinSize == 16 and device.adapterInfo.subgroupMaxSize == 32 and axis0Rows >= 2 * tunables.AXIS0_SPLIT_MIN_ROWS and axis0Cols >= tunables.AXIS0_TILE_MIN_COLS and axis0TilePathFits and ceilDiv(axis0Cols, tunables.AXIS0_TILE_COLS) >= axis0SplitCount"],
594
  "derive": {
595
  "splitCount": "axis0SplitCount",
596
  "partialElement": "\"f32\"",
 
603
  "id": "split_reduce",
604
  "name": "ReduceMean.Axis0SplitKReduce",
605
  "shader": "reduce-axis0-splitk-reduce.wgsl.jinja",
606
+ "derive": { "splitSpec": "splitCount" },
607
+ "bindings": ["x_t", "partials", "params_split_reduce"],
 
 
 
 
 
608
  "dispatch": {
609
  "x": "min(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), DISPATCH_FOLD_WIDTH)",
610
  "y": "splitCount",
 
615
  "id": "combine",
616
  "name": "ReduceMean.Axis0SplitKCombine",
617
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
618
+ "derive": { "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
619
+ "bindings": ["partials_combine", "y", "params_axis0_combine"],
620
  "dispatch": {
621
  "x": "min(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), 65535)",
622
  "y": "ceilDiv(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), 65535)",
 
635
  "id": "main",
636
  "name": "ReduceMean.Axis0TileCols",
637
  "shader": "reduce-axis0-tilecols.wgsl.jinja",
638
+ "derive": { "intMode": "dtypes.T == \"i32\"" },
639
+ "bindings": ["x_t", "y", "params_split_reduce"],
 
 
 
 
 
640
  "dispatch": {
641
  "x": "min(ceilDiv((dim(shapes.x, 1)), (tileCols)), 65535)",
642
  "y": "ceilDiv(ceilDiv((dim(shapes.x, 1)), (tileCols)), 65535)",
 
649
  "id": "all_axes_flat",
650
  "priority": 31,
651
  "when": ["flatParallelCovered"],
652
+ "derive": { "workgroupSize": "reduceWorkgroupSize", "split": "flatSplitCount" },
653
  "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[flatSplitCount]" }],
654
  "passes": [
655
  {
656
  "id": "flat_partial",
657
  "name": "ReduceMean.AllAxesFlatPartial",
658
  "shader": "reduce-flat-partial.wgsl.jinja",
659
+ "bindings": ["x_t", "partials_flat", "params_count4_numel"],
 
660
  "dispatch": { "x": "flatSplitCount" }
661
  },
662
  {
663
  "id": "combine",
664
  "name": "ReduceMean.AllAxesFlatCombine",
665
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
666
+ "derive": { "outputF16": "dtypes.T == \"f16\"" },
667
+ "bindings": ["partials_flat_read", "y", "params_all_axes_flat"],
668
  "dispatch": { "x": 1 }
669
  }
670
  ]
 
717
  "name": "ReduceMean.RankNSingleAxisGeneric",
718
  "shader": "reduce-serial-axis.wgsl.jinja",
719
  "derive": {
 
720
  "indexing": "\"rankn\"",
721
  "rank": "ranks.x",
722
  "axisSpec": "reduceAxis",
 
724
  "outputShape": "shapes.y",
725
  "outputRank": "ranks.y",
726
  "keepDims": "attrs.keepdims != 0",
727
+ "intMode": "dtypes.T == \"i32\""
 
 
728
  },
729
+ "bindings": ["x_t", "y", "params_axis_dim_out_count"],
730
  "dispatch": {
731
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
732
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
750
  "id": "main",
751
  "name": "ReduceMean.SubgroupRowVec4",
752
  "shader": "reduce-row-subgroup.wgsl.jinja",
753
+ "derive": { "vec4": true },
754
+ "bindings": ["x", "y", "params_reduce_row_tree"],
 
 
 
 
 
755
  "dispatch": {
756
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
757
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
758
  "z": 1
759
+ }
 
760
  }
761
  ]
762
  },
 
774
  "id": "main",
775
  "name": "ReduceMean.SubgroupRow",
776
  "shader": "reduce-row-subgroup.wgsl.jinja",
777
+ "derive": { "vec4": false },
778
+ "bindings": ["x_t", "y", "params_rows_chunk_count"],
 
 
 
 
 
779
  "dispatch": {
780
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
781
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
782
  "z": 1
783
+ }
 
784
  }
785
  ]
786
  },
 
793
  "passes": [
794
  {
795
  "id": "main",
796
+ "name": "ReduceMean.Axis0",
797
  "shader": "reduce-serial-axis.wgsl.jinja",
798
+ "derive": { "indexing": "\"axis2d\"", "intMode": "dtypes.T == \"i32\"" },
799
+ "bindings": ["x_t", "y", "params_axis0"],
 
 
 
 
 
 
 
800
  "dispatch": {
801
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
802
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
813
  "passes": [
814
  {
815
  "id": "main",
816
+ "name": "ReduceMean.Axis1",
817
  "shader": "reduce-serial-axis.wgsl.jinja",
818
+ "derive": { "indexing": "\"axis2d\"", "intMode": "dtypes.T == \"i32\"" },
819
+ "bindings": ["x_t", "y", "params_axis1"],
 
 
 
 
 
 
 
820
  "dispatch": {
821
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
822
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
826
  ]
827
  },
828
  {
829
+ "id": "all_axes_keepdims",
830
  "priority": 30,
831
+ "when": ["not flatParallelCovered", "f16Ok(dtypes.T)", "ranks.x >= 3", "attrs.keepdims == 1", "ranks.y == ranks.x", "numel(shapes.y) == 1"],
832
  "derive": { "axis": 0, "scalar": "dtypes.T" },
833
  "passes": [
834
  {
835
  "id": "main",
836
+ "name": "ReduceMean.Rank3AllAxesKeepdims",
837
  "shader": "reduce-serial-axis.wgsl.jinja",
838
+ "derive": { "indexing": "\"axis2d\"", "intMode": "dtypes.T == \"i32\"" },
839
+ "bindings": ["x_t", "y", "params_all_axes"],
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
840
  "dispatch": {
841
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
842
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
846
  ]
847
  },
848
  {
849
+ "id": "all_axes_no_keepdims",
850
  "priority": 30,
851
+ "when": ["not flatParallelCovered", "f16Ok(dtypes.T)", "ranks.x >= 3", "attrs.keepdims == 0", "ranks.y == 0"],
852
  "derive": { "axis": 0, "scalar": "dtypes.T" },
853
  "passes": [
854
  {
855
  "id": "main",
856
+ "name": "ReduceMean.Rank3AllAxesNoKeepdims",
857
  "shader": "reduce-serial-axis.wgsl.jinja",
858
+ "derive": { "indexing": "\"axis2d\"", "intMode": "dtypes.T == \"i32\"" },
859
+ "bindings": ["x_t", "y", "params_all_axes"],
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
860
  "dispatch": {
861
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
862
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
880
  "axisDim": "axisSplitDim",
881
  "innerSize": "axisSplitInner",
882
  "usesF16": "dtypes.T == \"f16\"",
 
883
  "tailPaddingElements": "numel(shapes.y) % 2 if dtypes.T == \"f16\" else 0"
884
  },
885
  "bindings": [{ "arg": "x", "elementType": "$scalar" }, { "arg": "y", "elementType": "$scalar" }],
build/webgpu/metadata.json CHANGED
@@ -1,32 +1,32 @@
1
  {
2
  "name": "ai.onnx.ReduceMean",
3
- "id": "_ai_onnx_reducemean_webgpu_3e440cd",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "pPlab/HIaUOReOgmFnA3jXKU0Vvj6XzHEsJ19hU56E8=",
11
- "manifest.json": "DyrQEvKuUXsnVWpCbkUMBAXl+cjdPCJsxOc1nqmHEfA=",
12
- "reduce-axis-split-reduce.wgsl.jinja": "4+ep9xH4pHZOfaZ8abJhA8SW5M4mUhDC+y1CtDz6vjY=",
13
- "reduce-axis0-splitk-combine.wgsl.jinja": "Yz1hjK55R/kndUrw3ugPqgPaOBwKmYupqZoVdJTHO5Q=",
14
- "reduce-axis0-splitk-reduce.wgsl.jinja": "jB2h58emn6rfKhd8ALzqylSDsmBbrrkEOhtIoNhcQM0=",
15
- "reduce-axis0-tilecols.wgsl.jinja": "I2BXFXhK2SWSwjfXB5BwDYyohcke/IPrl5SJE95kb5U=",
16
- "reduce-flat-partial.wgsl.jinja": "+qToL+wFi9QOxvY887aBAEwZK6Xu/eLQkucKvJYNwSk=",
17
- "reduce-multi-axis-coop.wgsl.jinja": "8v5G5dvcOUqPfLilBMg2ku67y2yYBggmsSqlBJuPUSE=",
18
- "reduce-noop-empty-axes.wgsl.jinja": "5bKufrcfYt3BagZ6WxbUvqwKKKVMU7d+SUshgUFWTDI=",
19
- "reduce-row-subgroup-rows.wgsl.jinja": "76u7rAvFoZZKrFDs2A2jk0vkB0uNrPL6twdOBE9b+v8=",
20
- "reduce-row-subgroup.wgsl.jinja": "2mu9LEsk8HfaLvucBCfcB1/ENpXkD6ELiCtRt+5UqiU=",
21
- "reduce-row-tree.wgsl.jinja": "Bwa5xcI0bTmKXb4r9Cc1bfVbM5rNqqpQVrWWVqcb8xA=",
22
- "reduce-serial-axis.wgsl.jinja": "fvUV9htqzKzt4Pg05pYtRmup/5QIHYaUGQHJZnthXKo=",
23
- "reduce-strided-axis.wgsl.jinja": "vbM7nz/UoIWB0tUb/rLmGalraF2JfjV9CT1RxP0Lfu8=",
24
- "test.json": "86DHyI08xIMOwY+dUT7Kjh13d9OFGJ1WEc2OofnRroQ="
25
  }
26
  },
27
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
28
  "webgpu": {
29
- "manifestSpec": "2.0",
30
  "variants": {
31
  "contiguous_suffix_subgroup_vec4": ["reduce-row-subgroup.wgsl.jinja"],
32
  "contiguous_suffix_tree_vec4": ["reduce-row-tree.wgsl.jinja"],
@@ -52,8 +52,8 @@
52
  "subgroup_last_axis": ["reduce-row-subgroup.wgsl.jinja"],
53
  "axis0": ["reduce-serial-axis.wgsl.jinja"],
54
  "axis1": ["reduce-serial-axis.wgsl.jinja"],
55
- "all_axes_no_keepdims": ["reduce-serial-axis.wgsl.jinja"],
56
  "all_axes_keepdims": ["reduce-serial-axis.wgsl.jinja"],
 
57
  "strided_axis_serial": ["reduce-strided-axis.wgsl.jinja"]
58
  }
59
  }
 
1
  {
2
  "name": "ai.onnx.ReduceMean",
3
+ "id": "_ai_onnx_reducemean_webgpu_b918009",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "bench.json": "bhV+fHo4TS6gK0PMa36JEwynO8XXhEGB9DlTDn1mBWk=",
11
+ "manifest.json": "rXpND7uYqiqCxsOFMhBDatls0p7CQ944DFSmMhomX/g=",
12
+ "reduce-axis-split-reduce.wgsl.jinja": "xo4jwtx/MeHIZkh2KU31ZCR+Ak58xOe4Cg3HWrcvdqo=",
13
+ "reduce-axis0-splitk-combine.wgsl.jinja": "bY5/z/a29FortWA0aT28cnkQX6dzLGr1O1fV/ncfxLs=",
14
+ "reduce-axis0-splitk-reduce.wgsl.jinja": "rW4Vltn1kPBqmbVTl7Za5CWycPaUe7IwaqXoWH+72Tw=",
15
+ "reduce-axis0-tilecols.wgsl.jinja": "thqt34BS8Z2qD/Nyd67lSFruVudITEmfo2sNJZUeKQU=",
16
+ "reduce-flat-partial.wgsl.jinja": "gMQg8GcibV+NQadQLAprKk6K3w13GCOdfQvIdCvYDXs=",
17
+ "reduce-multi-axis-coop.wgsl.jinja": "mQldeEk3lvyh67eiW/mfZFaM2gpRc0fUUgpq7PI3SHo=",
18
+ "reduce-noop-empty-axes.wgsl.jinja": "MMZEXHeACtGwvDaH0nAyvMal+tAIjYrh2yxDuLfh1Lw=",
19
+ "reduce-row-subgroup-rows.wgsl.jinja": "5cMGE1dmrcawwHdoxQD8kPSurwA1wNHyoAQfxAOE7vU=",
20
+ "reduce-row-subgroup.wgsl.jinja": "7bRBWIQ7oHOq/E54hGrmXIzLznzVH9oTfYaf1jmBzV8=",
21
+ "reduce-row-tree.wgsl.jinja": "a+iEV/cBkrlsVz6NS0GY/TAsUSX5rO4N14wY1dKfDi0=",
22
+ "reduce-serial-axis.wgsl.jinja": "XZthzOtNawNeG2SXq528l1yO5sVPqTTK3Y9GVhJ46lk=",
23
+ "reduce-strided-axis.wgsl.jinja": "K+yHpuxWwKU75+/4M8BuQ8UU2e+gGKoklPj5I/uH5rU=",
24
+ "test.json": "1X6wSEe7Gq+J/wtGb7LTVh+F0Ifvw1FoCsLwOdgPMFQ="
25
  }
26
  },
27
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
28
  "webgpu": {
29
+ "manifestSpec": "2.1",
30
  "variants": {
31
  "contiguous_suffix_subgroup_vec4": ["reduce-row-subgroup.wgsl.jinja"],
32
  "contiguous_suffix_tree_vec4": ["reduce-row-tree.wgsl.jinja"],
 
52
  "subgroup_last_axis": ["reduce-row-subgroup.wgsl.jinja"],
53
  "axis0": ["reduce-serial-axis.wgsl.jinja"],
54
  "axis1": ["reduce-serial-axis.wgsl.jinja"],
 
55
  "all_axes_keepdims": ["reduce-serial-axis.wgsl.jinja"],
56
+ "all_axes_no_keepdims": ["reduce-serial-axis.wgsl.jinja"],
57
  "strided_axis_serial": ["reduce-strided-axis.wgsl.jinja"]
58
  }
59
  }
build/webgpu/reduce-axis-split-reduce.wgsl.jinja CHANGED
@@ -8,29 +8,16 @@
8
  // Each output segment writes three partial planes: its maximum, the sum of
9
  // exp(x - maximum), and a packed NaN marker.
10
  {% endif %}
11
- {% set castF32 = castF32 is defined and castF32 %}
12
- {% set xa = "f32(" if castF32 else "" %}
13
  {% set ax = ")" if castF32 else "" %}
14
- {% if usesF16Spec is defined and usesF16Spec %}
15
- enable f16;
16
- {% endif %}
17
  {{ env.wgsl.resourceDeclarations }}
18
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
19
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
20
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
21
- {% set logicalBoolIdentity = logicalBool is defined and logicalBool %}
22
  fn {{ name }}() -> {{ scalar }} {
23
- {% if scalar == "i32" %}
24
- return {{ "-2147483647i - 1i" if op == "max" else "2147483647i" }};
25
- {% elif scalar == "u32" %}
26
- return {{ "1u" if logicalBoolIdentity and op == "min" else "0u" if op == "max" else "4294967295u" }};
27
- {% else %}
28
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
29
  return bitcast<f32>(bits);
30
- {% endif %}
31
- }
32
- {%- endmacro %}
33
-
34
 
35
  const WG: u32 = {{ workgroupSize }}u;
36
  const SPLIT: u32 = {{ split }}u;
@@ -42,6 +29,7 @@ fn is_nan_f32(value: f32) -> bool {
42
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
43
  }
44
  {% elif op == "max" or op == "min" %}
 
45
  {{ wgsl_minmax_identity("reduction_identity", op) }}
46
  {% endif %}
47
 
 
8
  // Each output segment writes three partial planes: its maximum, the sum of
9
  // exp(x - maximum), and a packed NaN marker.
10
  {% endif %}
11
+ {% set xa = ("vec" ~ columnVectorWidth ~ "<f32>(" if columnVectorWidth is defined else "f32(") if castF32 else "" %}
 
12
  {% set ax = ")" if castF32 else "" %}
 
 
 
13
  {{ env.wgsl.resourceDeclarations }}
14
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
15
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
16
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
 
17
  fn {{ name }}() -> {{ scalar }} {
 
 
 
 
 
18
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
19
  return bitcast<f32>(bits);
20
+ }{% endmacro %}
 
 
 
21
 
22
  const WG: u32 = {{ workgroupSize }}u;
23
  const SPLIT: u32 = {{ split }}u;
 
29
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
30
  }
31
  {% elif op == "max" or op == "min" %}
32
+
33
  {{ wgsl_minmax_identity("reduction_identity", op) }}
34
  {% endif %}
35
 
build/webgpu/reduce-axis0-splitk-combine.wgsl.jinja CHANGED
@@ -2,37 +2,20 @@
2
  // folds the segment partials and applies the selected reduction's final step.
3
  // Segments are folded in ascending order for deterministic results. This order
4
  // differs from the single-pass reduction but remains within the f32 tolerance.
5
- {% set addBias = addBias is defined and addBias %}
6
- {% set biasCols = biasCols | default(0) %}
7
  {% set intMode = intMode is defined and intMode %}
8
  {% set yv = "f16(" if outputF16 else "" %}
9
  {% set vy = ")" if outputF16 else "" %}
10
- {% if outputF16 %}
11
- enable f16;
12
- {% endif %}
13
  {{ env.wgsl.resourceDeclarations }}
14
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
15
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
16
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
17
- {% set logicalBoolIdentity = logicalBool is defined and logicalBool %}
18
  fn {{ name }}() -> {{ scalar }} {
19
- {% if scalar == "i32" %}
20
- return {{ "-2147483647i - 1i" if op == "max" else "2147483647i" }};
21
- {% elif scalar == "u32" %}
22
- return {{ "1u" if logicalBoolIdentity and op == "min" else "0u" if op == "max" else "4294967295u" }};
23
- {% else %}
24
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
25
  return bitcast<f32>(bits);
26
- {% endif %}
27
- }
28
- {%- endmacro %}
29
-
30
 
31
  const WG: u32 = {{ workgroupSize }}u;
32
  const SPLIT: u32 = {{ split }}u;
33
- {% if addBias %}
34
- const BIAS_COLS: u32 = {{ biasCols }}u;
35
- {% endif %}
36
  {% if op == "logsumexp" %}
37
  const F32_MIN: f32 = -3.4028234663852886e38;
38
  const F32_MAX: f32 = 3.4028234663852886e38;
@@ -48,7 +31,9 @@ fn is_nan_f32(value: f32) -> bool {
48
  @compute @workgroup_size(WG, 1, 1)
49
  fn main(@builtin(global_invocation_id) gid: vec3<u32>,
50
  @builtin(num_workgroups) nwg: vec3<u32>) {
51
- let stride = nwg.x * WG;
 
 
52
  let start = (gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG) + gid.x;
53
  for (var col = start; col < params.cols; col = col + stride) {
54
  {% if op == "logsumexp" %}
@@ -101,9 +86,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>,
101
  total = total + p;
102
  {% endif %}
103
  }
104
- {% if addBias %}
105
- total = total + f32(bias[col % BIAS_COLS]);
106
- {% endif %}
107
  {% if op == "l2" %}
108
  y[col] = {{ yv }}sqrt(total){{ vy }};
109
  {% elif op == "logsum" %}
 
2
  // folds the segment partials and applies the selected reduction's final step.
3
  // Segments are folded in ascending order for deterministic results. This order
4
  // differs from the single-pass reduction but remains within the f32 tolerance.
 
 
5
  {% set intMode = intMode is defined and intMode %}
6
  {% set yv = "f16(" if outputF16 else "" %}
7
  {% set vy = ")" if outputF16 else "" %}
 
 
 
8
  {{ env.wgsl.resourceDeclarations }}
9
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
10
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
11
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
 
12
  fn {{ name }}() -> {{ scalar }} {
 
 
 
 
 
13
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
14
  return bitcast<f32>(bits);
15
+ }{% endmacro %}
 
 
 
16
 
17
  const WG: u32 = {{ workgroupSize }}u;
18
  const SPLIT: u32 = {{ split }}u;
 
 
 
19
  {% if op == "logsumexp" %}
20
  const F32_MIN: f32 = -3.4028234663852886e38;
21
  const F32_MAX: f32 = 3.4028234663852886e38;
 
31
  @compute @workgroup_size(WG, 1, 1)
32
  fn main(@builtin(global_invocation_id) gid: vec3<u32>,
33
  @builtin(num_workgroups) nwg: vec3<u32>) {
34
+ // The start already folds gid.y in, so the stride must span every y row too;
35
+ // an x-only stride would send y = 0 lanes over columns the y >= 1 rows own.
36
+ let stride = nwg.x * nwg.y * WG;
37
  let start = (gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG) + gid.x;
38
  for (var col = start; col < params.cols; col = col + stride) {
39
  {% if op == "logsumexp" %}
 
86
  total = total + p;
87
  {% endif %}
88
  }
 
 
 
89
  {% if op == "l2" %}
90
  y[col] = {{ yv }}sqrt(total){{ vy }};
91
  {% elif op == "logsum" %}
build/webgpu/reduce-axis0-splitk-reduce.wgsl.jinja CHANGED
@@ -2,29 +2,16 @@
2
  // segments increases residency for tall matrices. Each (column, segment)
3
  // invocation reduces one row slice and writes partials[segment * columns +
4
  // column]. Adjacent column threads keep row reads coalesced.
5
- {% set castF32 = castF32 is defined and castF32 %}
6
  {% set xa = "f32(" if castF32 else "" %}
7
  {% set ax = ")" if castF32 else "" %}
8
- {% if usesF16Spec is defined and usesF16Spec %}
9
- enable f16;
10
- {% endif %}
11
  {{ env.wgsl.resourceDeclarations }}
12
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
13
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
14
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
15
- {% set logicalBoolIdentity = logicalBool is defined and logicalBool %}
16
  fn {{ name }}() -> {{ scalar }} {
17
- {% if scalar == "i32" %}
18
- return {{ "-2147483647i - 1i" if op == "max" else "2147483647i" }};
19
- {% elif scalar == "u32" %}
20
- return {{ "1u" if logicalBoolIdentity and op == "min" else "0u" if op == "max" else "4294967295u" }};
21
- {% else %}
22
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
23
  return bitcast<f32>(bits);
24
- {% endif %}
25
- }
26
- {%- endmacro %}
27
-
28
 
29
  const WG: u32 = {{ workgroupSize }}u;
30
  const SPLIT: u32 = {{ split }}u;
 
2
  // segments increases residency for tall matrices. Each (column, segment)
3
  // invocation reduces one row slice and writes partials[segment * columns +
4
  // column]. Adjacent column threads keep row reads coalesced.
 
5
  {% set xa = "f32(" if castF32 else "" %}
6
  {% set ax = ")" if castF32 else "" %}
 
 
 
7
  {{ env.wgsl.resourceDeclarations }}
8
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
9
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
10
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
 
11
  fn {{ name }}() -> {{ scalar }} {
 
 
 
 
 
12
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
13
  return bitcast<f32>(bits);
14
+ }{% endmacro %}
 
 
 
15
 
16
  const WG: u32 = {{ workgroupSize }}u;
17
  const SPLIT: u32 = {{ split }}u;
build/webgpu/reduce-axis0-tilecols.wgsl.jinja CHANGED
@@ -13,7 +13,6 @@
13
  {% set rowEnd = "params.rows" %}
14
  {% set elem = "x[inputBase + row * params.cols + col]" %}
15
  {% endif %}
16
- {% set castF32 = castF32 is defined and castF32 %}
17
  {% set intMode = intMode is defined and intMode %}
18
  {% set scalar = "f32" if castF32 else scalar %}
19
  {% if castF32 %}
@@ -21,9 +20,6 @@
21
  {% endif %}
22
  {% set yv = "f16(" if castF32 else "" %}
23
  {% set vy = ")" if castF32 else "" %}
24
- {% if usesF16Spec is defined and usesF16Spec %}
25
- enable f16;
26
- {% endif %}
27
  {{ env.wgsl.resourceDeclarations }}
28
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
29
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
@@ -38,15 +34,12 @@ fn {{ name }}() -> {{ scalar }} {
38
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
39
  return bitcast<f32>(bits);
40
  {% endif %}
41
- }
42
- {%- endmacro %}
43
-
44
  {% if not splitMode and not intMode and (op == "logsum" or op == "logsumexp") %}
45
  fn negative_infinity() -> f32 {
46
  var bits = 0xff800000u;
47
  return bitcast<f32>(bits);
48
  }
49
-
50
  {% endif %}
51
 
52
  const WG: u32 = {{ workgroupSize }}u;
@@ -69,7 +62,6 @@ fn is_nan_f32(value: f32) -> bool {
69
  let bits = bitcast<u32>(value);
70
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
71
  }
72
-
73
  {% endif %}
74
  @compute @workgroup_size(WG, 1, 1)
75
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
@@ -77,8 +69,9 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid:
77
  let col_lane = tid % TILE_COLS;
78
  let row_lane = tid / TILE_COLS;
79
  {% if splitMode %}
80
- // Narrow outputs: wg.x covers every column tile, wg.y is the axis segment.
81
- let col = wg.x * TILE_COLS + col_lane;
 
82
  let outputIndex = col;
83
  let in_bounds = col < params.outputs;
84
  let seg = wg.y;
 
13
  {% set rowEnd = "params.rows" %}
14
  {% set elem = "x[inputBase + row * params.cols + col]" %}
15
  {% endif %}
 
16
  {% set intMode = intMode is defined and intMode %}
17
  {% set scalar = "f32" if castF32 else scalar %}
18
  {% if castF32 %}
 
20
  {% endif %}
21
  {% set yv = "f16(" if castF32 else "" %}
22
  {% set vy = ")" if castF32 else "" %}
 
 
 
23
  {{ env.wgsl.resourceDeclarations }}
24
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
25
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
 
34
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
35
  return bitcast<f32>(bits);
36
  {% endif %}
37
+ }{% endmacro %}
 
 
38
  {% if not splitMode and not intMode and (op == "logsum" or op == "logsumexp") %}
39
  fn negative_infinity() -> f32 {
40
  var bits = 0xff800000u;
41
  return bitcast<f32>(bits);
42
  }
 
43
  {% endif %}
44
 
45
  const WG: u32 = {{ workgroupSize }}u;
 
62
  let bits = bitcast<u32>(value);
63
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
64
  }
 
65
  {% endif %}
66
  @compute @workgroup_size(WG, 1, 1)
67
  fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
 
69
  let col_lane = tid % TILE_COLS;
70
  let row_lane = tid / TILE_COLS;
71
  {% if splitMode %}
72
+ // Narrow outputs: wg.x (with wg.z carrying its fold overflow) covers every
73
+ // column tile, wg.y is the axis segment.
74
+ let col = (wg.x + wg.z * {{ DISPATCH_FOLD_WIDTH }}u) * TILE_COLS + col_lane;
75
  let outputIndex = col;
76
  let in_bounds = col < params.outputs;
77
  let seg = wg.y;
build/webgpu/reduce-flat-partial.wgsl.jinja CHANGED
@@ -3,33 +3,20 @@
3
  // in registers, reduce within each workgroup, and write one partial each. A
4
  // following combine pass folds the partials and applies the operation finalizer.
5
  //
6
- // Scalar f32 bindings keep arbitrary element counts legal. The grid-stride loop
7
- // manually assembles full vec4 groups from contiguous scalars, and one global
8
- // thread folds the final zero-to-three scalar elements exactly once.
9
  {% set intMode = intMode is defined and intMode %}
10
- {% set castF32 = castF32 is defined and castF32 %}
11
  {% set xa = "f32(" if castF32 else "" %}
12
  {% set ax = ")" if castF32 else "" %}
13
- {% if usesF16Spec is defined and usesF16Spec %}
14
- enable f16;
15
- {% endif %}
16
  {{ env.wgsl.resourceDeclarations }}
17
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
18
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
19
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
20
- {% set logicalBoolIdentity = logicalBool is defined and logicalBool %}
21
  fn {{ name }}() -> {{ scalar }} {
22
- {% if scalar == "i32" %}
23
- return {{ "-2147483647i - 1i" if op == "max" else "2147483647i" }};
24
- {% elif scalar == "u32" %}
25
- return {{ "1u" if logicalBoolIdentity and op == "min" else "0u" if op == "max" else "4294967295u" }};
26
- {% else %}
27
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
28
  return bitcast<f32>(bits);
29
- {% endif %}
30
- }
31
- {%- endmacro %}
32
-
33
 
34
  const WG: u32 = {{ workgroupSize }}u;
35
  var<workgroup> red: array<{{ "i32" if intMode else "f32" }}, WG>;
@@ -100,23 +87,23 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>,
100
  }
101
  red[tid] = acc;
102
  workgroupBarrier();
103
- var stride: u32 = WG / 2u;
 
 
 
 
 
104
  loop {
105
- if (stride == 0u) { break; }
106
- if (tid < stride) {
107
- {% if op == "max" %}
108
- red[tid] = max(red[tid], red[tid + stride]);
109
- {% elif op == "min" %}
110
- red[tid] = min(red[tid], red[tid + stride]);
111
- {% elif op == "prod" %}
112
- red[tid] = red[tid] * red[tid + stride];
113
- {% else %}
114
- red[tid] = red[tid] + red[tid + stride];
115
- {% endif %}
116
  }
117
- stride = stride / 2u;
118
  workgroupBarrier();
119
- }
 
120
  if (tid == 0u) {
121
  partials[wg.x] = red[0];
122
  }
 
3
  // in registers, reduce within each workgroup, and write one partial each. A
4
  // following combine pass folds the partials and applies the operation finalizer.
5
  //
6
+ // Scalar bindings support arbitrary element counts and a zero-to-three element
7
+ // tail. An aligned packed binding loads each group directly, preserving the
8
+ // component accumulation order and widening half storage before arithmetic.
9
  {% set intMode = intMode is defined and intMode %}
 
10
  {% set xa = "f32(" if castF32 else "" %}
11
  {% set ax = ")" if castF32 else "" %}
 
 
 
12
  {{ env.wgsl.resourceDeclarations }}
13
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
14
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
15
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
 
16
  fn {{ name }}() -> {{ scalar }} {
 
 
 
 
 
17
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
18
  return bitcast<f32>(bits);
19
+ }{% endmacro %}
 
 
 
20
 
21
  const WG: u32 = {{ workgroupSize }}u;
22
  var<workgroup> red: array<{{ "i32" if intMode else "f32" }}, WG>;
 
87
  }
88
  red[tid] = acc;
89
  workgroupBarrier();
90
+ {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
91
+ {% if op == "max" or op == "min" %}
92
+ {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
93
+ {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %}
94
+ {% 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) %}
95
+ var {{ svar }}{{ ": u32 " if typed else " " }}= {{ wg }} / 2u;
96
  loop {
97
+ if ({{ svar }} == 0u) { break; }
98
+ if ({{ idx }} < {{ svar }}) {
99
+ {% for a in arrays %}
100
+ {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
101
+ {% endfor %}
 
 
 
 
 
 
102
  }
103
+ {{ svar }} = {{ svar }} / 2u;
104
  workgroupBarrier();
105
+ }{% endmacro %}
106
+ {{ wgsl_tree_fold(["red"], op=op, idx="tid", wg="WG", typed=true, form="head", breakInline=true) }}
107
  if (tid == 0u) {
108
  partials[wg.x] = red[0];
109
  }
build/webgpu/reduce-multi-axis-coop.wgsl.jinja CHANGED
@@ -1,5 +1,4 @@
1
  {% macro reduce_multi_axis_offset(hasReduced) %}
2
-
3
  // One thread per output element walks the Cartesian product of the reduced axes,
4
  // linearized as reduce_linear. Specialized shapes make every input offset a sum
5
  // of coordinate-times-constant terms.
@@ -44,8 +43,7 @@ fn input_offset(out_index: u32{% if hasReduced %}, reduce_linear: u32{% endif %}
44
  {% set src.value = "(" ~ src.value ~ " * " ~ dataShape[a] ~ "u + coord" ~ a ~ ")" %}
45
  {% endfor %}
46
  return {{ src.value }};
47
- }
48
- {%- endmacro %}
49
  // Cooperative multi-axis reduction with one workgroup per output element.
50
  // Threads take strided shares of the reduced index range, then combine their
51
  // partials through a workgroup tree. The shared offset helper maps each reduced
@@ -53,17 +51,13 @@ fn input_offset(out_index: u32{% if hasReduced %}, reduce_linear: u32{% endif %}
53
  //
54
  // The tree changes f32 association relative to a left-to-right fold, so the two
55
  // accumulation orders need not be bit-identical.
56
- {% set castF32 = castF32 is defined and castF32 %}
57
- {% set intMode = intMode is defined and intMode %}
58
  {% set yv = "f16(" if castF32 else "" %}
59
  {% set vy = ")" if castF32 else "" %}
60
- {% if usesF16Spec is defined and usesF16Spec %}
61
- enable f16;
62
- {% endif %}
63
  {{ env.wgsl.resourceDeclarations }}
64
  {% set hasReducedAxis = namespace(value=false) %}
65
  {% for a in range(rank) %}{% if reduce[a] %}{% set hasReducedAxis.value = true %}{% endif %}{% endfor %}
66
- {{- reduce_multi_axis_offset(hasReducedAxis.value) }}
 
67
 
68
  {% set ACC = scalar if intMode else "f32" %}
69
  const WG: u32 = {{ workgroupSize }}u;
@@ -98,6 +92,7 @@ fn main(
98
  return;
99
  }
100
  let lid = lid3.x;
 
101
  if (REDUCED == 0u) {
102
  {% if intMode %}
103
  if (lid == 0u) { y[i] = {{ scalar }}(0); }
@@ -107,13 +102,23 @@ fn main(
107
  return;
108
  }
109
  var acc = {{ ACC }}(0);
 
 
 
110
 
111
  for (var r = lid; r < REDUCED; r = r + WG) {
112
  {% set at = "x[input_offset(i, r)]" %}
113
  {% if castF32 %}
114
  {% set at = "f32(" ~ at ~ ")" %}
115
  {% endif %}
 
 
 
 
 
 
116
  acc = combine(acc, {{ at }});
 
117
  }
118
 
119
  // Fold the per-thread accumulators. WORKGROUP_SIZE is a power of two, and every
@@ -134,9 +139,19 @@ fn main(
134
  return;
135
  }
136
  let total = partials[0];
 
 
 
 
 
 
 
137
  {% if intMode %}
138
  y[i] = total / {{ scalar }}(REDUCED);
139
  {% else %}
140
  y[i] = {{ yv }}total / f32(REDUCED){{ vy }};
141
  {% endif %}
 
 
 
142
  }
 
1
  {% macro reduce_multi_axis_offset(hasReduced) %}
 
2
  // One thread per output element walks the Cartesian product of the reduced axes,
3
  // linearized as reduce_linear. Specialized shapes make every input offset a sum
4
  // of coordinate-times-constant terms.
 
43
  {% set src.value = "(" ~ src.value ~ " * " ~ dataShape[a] ~ "u + coord" ~ a ~ ")" %}
44
  {% endfor %}
45
  return {{ src.value }};
46
+ }{% endmacro %}
 
47
  // Cooperative multi-axis reduction with one workgroup per output element.
48
  // Threads take strided shares of the reduced index range, then combine their
49
  // partials through a workgroup tree. The shared offset helper maps each reduced
 
51
  //
52
  // The tree changes f32 association relative to a left-to-right fold, so the two
53
  // accumulation orders need not be bit-identical.
 
 
54
  {% set yv = "f16(" if castF32 else "" %}
55
  {% set vy = ")" if castF32 else "" %}
 
 
 
56
  {{ env.wgsl.resourceDeclarations }}
57
  {% set hasReducedAxis = namespace(value=false) %}
58
  {% for a in range(rank) %}{% if reduce[a] %}{% set hasReducedAxis.value = true %}{% endif %}{% endfor %}
59
+
60
+ {{ reduce_multi_axis_offset(hasReducedAxis.value) }}
61
 
62
  {% set ACC = scalar if intMode else "f32" %}
63
  const WG: u32 = {{ workgroupSize }}u;
 
92
  return;
93
  }
94
  let lid = lid3.x;
95
+ {% if op == "mean" %}
96
  if (REDUCED == 0u) {
97
  {% if intMode %}
98
  if (lid == 0u) { y[i] = {{ scalar }}(0); }
 
102
  return;
103
  }
104
  var acc = {{ ACC }}(0);
105
+ {% else %}
106
+ var acc = {{ ACC }}(0);
107
+ {% endif %}
108
 
109
  for (var r = lid; r < REDUCED; r = r + WG) {
110
  {% set at = "x[input_offset(i, r)]" %}
111
  {% if castF32 %}
112
  {% set at = "f32(" ~ at ~ ")" %}
113
  {% endif %}
114
+ {% if op == "l1" %}
115
+ acc = combine(acc, abs({{ at }}));
116
+ {% elif op == "l2" or op == "sumsquare" %}
117
+ let value = {{ at }};
118
+ acc = combine(acc, value * value);
119
+ {% else %}
120
  acc = combine(acc, {{ at }});
121
+ {% endif %}
122
  }
123
 
124
  // Fold the per-thread accumulators. WORKGROUP_SIZE is a power of two, and every
 
139
  return;
140
  }
141
  let total = partials[0];
142
+ {% if op == "l2" %}
143
+ {% if intMode %}
144
+ y[i] = {{ scalar }}(sqrt(f32(total)));
145
+ {% else %}
146
+ y[i] = {{ yv }}sqrt(total){{ vy }};
147
+ {% endif %}
148
+ {% elif op == "mean" %}
149
  {% if intMode %}
150
  y[i] = total / {{ scalar }}(REDUCED);
151
  {% else %}
152
  y[i] = {{ yv }}total / f32(REDUCED){{ vy }};
153
  {% endif %}
154
+ {% else %}
155
+ y[i] = {{ yv }}total{{ vy }};
156
+ {% endif %}
157
  }
build/webgpu/reduce-noop-empty-axes.wgsl.jinja CHANGED
@@ -1,13 +1,16 @@
 
 
 
 
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  @compute @workgroup_size({{ reduceWorkgroupSize }})
4
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
5
- // 2D-folded flat index: gid.y carries the high bits past the
6
- // per-axis dispatch fold width (outputs > 16.7M elements).
7
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ reduceWorkgroupSize }}u;
8
- if (i >= params.count) {
9
- return;
10
- }
11
  let v = x[i];
12
  {% if op == "abs" %}
13
  y[i] = abs(v);
 
1
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
2
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
3
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
+ // per-axis workgroup fold width.
5
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
6
+ if ({{ name }} >= {{ bound }}) {
7
+ return;
8
+ }{% endmacro %}
9
  {{ env.wgsl.resourceDeclarations }}
10
 
11
  @compute @workgroup_size({{ reduceWorkgroupSize }})
12
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
13
+ {{ flat_index_2d(reduceWorkgroupSize) }}
 
 
 
 
 
14
  let v = x[i];
15
  {% if op == "abs" %}
16
  y[i] = abs(v);
build/webgpu/reduce-row-subgroup-rows.wgsl.jinja CHANGED
@@ -15,18 +15,12 @@
15
  // The finalizer takes the logarithm of the sum.
16
  {% endif %}
17
  // f16 storage widens before accumulation and narrows only at the final store.
18
- {% set castF32 = castF32 is defined and castF32 %}
19
  {% set isInt = scalar == "i32" or scalar == "u32" %}
20
  {% set accScalar = scalar if isInt else "f32" %}
21
  {% set xv = "vec4<f32>(" if castF32 else "" %}
22
  {% set vx = ")" if castF32 else "" %}
23
  {% set yv = "f16(" if castF32 else "" %}
24
  {% set vy = ")" if castF32 else "" %}
25
- enable subgroups;
26
- {% if usesF16Spec is defined and usesF16Spec %}
27
- enable f16;
28
- {% endif %}
29
- {{ env.wgsl.resourceDeclarations }}
30
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
31
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
32
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
@@ -40,15 +34,15 @@ fn {{ name }}() -> {{ scalar }} {
40
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
41
  return bitcast<f32>(bits);
42
  {% endif %}
43
- }
44
- {%- endmacro %}
45
-
46
 
47
  const WG: u32 = {{ workgroupSize }}u;
48
- {%- if op == "max" or op == "min" %}
 
49
  {{ wgsl_minmax_identity("reduction_identity", op, accScalar) }}
50
- {%- endif %}
51
- {%- if op == "logsumexp" %}
52
 
53
  const F32_MIN: f32 = -3.4028234663852886e38;
54
  const F32_MAX: f32 = 3.4028234663852886e38;
@@ -57,7 +51,7 @@ fn is_nan_f32(value: f32) -> bool {
57
  let bits = bitcast<u32>(value);
58
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
59
  }
60
- {%- endif %}
61
 
62
  @compute @workgroup_size(WG, 1, 1)
63
  fn main(@builtin(workgroup_id) wg: vec3<u32>,
@@ -110,58 +104,58 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>,
110
  let finiteOrInf = select(rowMax + log(sum), rowMax, hasPositiveInf);
111
  y[row] = {{ yv }}select(finiteOrInf, nanValue, hasNan){{ vy }};
112
  }
113
- {%- else %}
114
  {% if op == "max" or op == "min" %}
115
  let INIT: {{ accScalar }} = reduction_identity();
116
- {%- elif op == "prod" %}
117
  let INIT: {{ accScalar }} = {{ "1.0" if not isInt else accScalar ~ "(1)" }};
118
- {%- else %}
119
  let INIT: {{ accScalar }} = {{ "0.0" if not isInt else accScalar ~ "(0)" }};
120
- {%- endif %}
121
  var acc4 = vec4<{{ accScalar }}>(INIT);
122
  for (var c = sgLane; c < chunkLimit; c = c + sgSize) {
123
  let v = {{ xv }}x[base + c]{{ vx }};
124
- {%- if op == "max" %}
125
  acc4 = max(acc4, v);
126
- {%- elif op == "min" %}
127
  acc4 = min(acc4, v);
128
- {%- elif op == "prod" %}
129
  acc4 = acc4 * v;
130
- {%- elif op == "l1" %}
131
  acc4 = acc4 + abs(v);
132
- {%- elif op == "l2" or op == "sumsquare" %}
133
  acc4 = acc4 + v * v;
134
- {%- else %}
135
  acc4 = acc4 + v;
136
- {%- endif %}
137
  }
138
- {%- if op == "max" %}
139
  let acc = max(max(acc4.x, acc4.y), max(acc4.z, acc4.w));
140
  let total = subgroupMax(acc);
141
- {%- elif op == "min" %}
142
  let acc = min(min(acc4.x, acc4.y), min(acc4.z, acc4.w));
143
  let total = subgroupMin(acc);
144
- {%- elif op == "prod" %}
145
  let acc = (acc4.x * acc4.y) * (acc4.z * acc4.w);
146
  let total = subgroupMul(acc);
147
- {%- else %}
148
  let acc = (acc4.x + acc4.y) + (acc4.z + acc4.w);
149
  let total = subgroupAdd(acc);
150
- {%- endif %}
151
  if (rowValid && sgLane == 0u) {
152
- {%- if op == "mean" and isInt %}
153
  y[row] = total / {{ accScalar }}(params.chunkCount * 4u);
154
- {%- elif op == "mean" %}
155
  y[row] = {{ yv }}total / f32(params.chunkCount * 4u){{ vy }};
156
- {%- elif op == "l2" and isInt %}
157
  y[row] = {{ accScalar }}(sqrt(f32(total)));
158
- {%- elif op == "l2" %}
159
  y[row] = {{ yv }}sqrt(total){{ vy }};
160
- {%- elif op == "logsum" %}
161
  y[row] = {{ yv }}log(total){{ vy }};
162
- {%- else %}
163
  y[row] = {{ yv }}total{{ vy }};
164
- {%- endif %}
165
  }
166
- {%- endif %}
167
  }
 
15
  // The finalizer takes the logarithm of the sum.
16
  {% endif %}
17
  // f16 storage widens before accumulation and narrows only at the final store.
 
18
  {% set isInt = scalar == "i32" or scalar == "u32" %}
19
  {% set accScalar = scalar if isInt else "f32" %}
20
  {% set xv = "vec4<f32>(" if castF32 else "" %}
21
  {% set vx = ")" if castF32 else "" %}
22
  {% set yv = "f16(" if castF32 else "" %}
23
  {% set vy = ")" if castF32 else "" %}
 
 
 
 
 
24
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
25
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
26
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
 
34
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
35
  return bitcast<f32>(bits);
36
  {% endif %}
37
+ }{% endmacro %}
38
+ enable subgroups;
39
+ {{ env.wgsl.resourceDeclarations }}
40
 
41
  const WG: u32 = {{ workgroupSize }}u;
42
+ {% if op == "max" or op == "min" %}
43
+
44
  {{ wgsl_minmax_identity("reduction_identity", op, accScalar) }}
45
+ {% elif op == "logsumexp" %}
 
46
 
47
  const F32_MIN: f32 = -3.4028234663852886e38;
48
  const F32_MAX: f32 = 3.4028234663852886e38;
 
51
  let bits = bitcast<u32>(value);
52
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
53
  }
54
+ {% endif %}
55
 
56
  @compute @workgroup_size(WG, 1, 1)
57
  fn main(@builtin(workgroup_id) wg: vec3<u32>,
 
104
  let finiteOrInf = select(rowMax + log(sum), rowMax, hasPositiveInf);
105
  y[row] = {{ yv }}select(finiteOrInf, nanValue, hasNan){{ vy }};
106
  }
107
+ {% else %}
108
  {% if op == "max" or op == "min" %}
109
  let INIT: {{ accScalar }} = reduction_identity();
110
+ {% elif op == "prod" %}
111
  let INIT: {{ accScalar }} = {{ "1.0" if not isInt else accScalar ~ "(1)" }};
112
+ {% else %}
113
  let INIT: {{ accScalar }} = {{ "0.0" if not isInt else accScalar ~ "(0)" }};
114
+ {% endif %}
115
  var acc4 = vec4<{{ accScalar }}>(INIT);
116
  for (var c = sgLane; c < chunkLimit; c = c + sgSize) {
117
  let v = {{ xv }}x[base + c]{{ vx }};
118
+ {% if op == "max" %}
119
  acc4 = max(acc4, v);
120
+ {% elif op == "min" %}
121
  acc4 = min(acc4, v);
122
+ {% elif op == "prod" %}
123
  acc4 = acc4 * v;
124
+ {% elif op == "l1" %}
125
  acc4 = acc4 + abs(v);
126
+ {% elif op == "l2" or op == "sumsquare" %}
127
  acc4 = acc4 + v * v;
128
+ {% else %}
129
  acc4 = acc4 + v;
130
+ {% endif %}
131
  }
132
+ {% if op == "max" %}
133
  let acc = max(max(acc4.x, acc4.y), max(acc4.z, acc4.w));
134
  let total = subgroupMax(acc);
135
+ {% elif op == "min" %}
136
  let acc = min(min(acc4.x, acc4.y), min(acc4.z, acc4.w));
137
  let total = subgroupMin(acc);
138
+ {% elif op == "prod" %}
139
  let acc = (acc4.x * acc4.y) * (acc4.z * acc4.w);
140
  let total = subgroupMul(acc);
141
+ {% else %}
142
  let acc = (acc4.x + acc4.y) + (acc4.z + acc4.w);
143
  let total = subgroupAdd(acc);
144
+ {% endif %}
145
  if (rowValid && sgLane == 0u) {
146
+ {% if op == "mean" and isInt %}
147
  y[row] = total / {{ accScalar }}(params.chunkCount * 4u);
148
+ {% elif op == "mean" %}
149
  y[row] = {{ yv }}total / f32(params.chunkCount * 4u){{ vy }};
150
+ {% elif op == "l2" and isInt %}
151
  y[row] = {{ accScalar }}(sqrt(f32(total)));
152
+ {% elif op == "l2" %}
153
  y[row] = {{ yv }}sqrt(total){{ vy }};
154
+ {% elif op == "logsum" %}
155
  y[row] = {{ yv }}log(total){{ vy }};
156
+ {% else %}
157
  y[row] = {{ yv }}total{{ vy }};
158
+ {% endif %}
159
  }
160
+ {% endif %}
161
  }
build/webgpu/reduce-row-subgroup.wgsl.jinja CHANGED
@@ -15,17 +15,11 @@
15
  // IEEE-754 bit patterns because WGSL rejects infinite constants.
16
  {% endif %}
17
  // f16 storage widens before accumulation and narrows only at the final store.
18
- {% set castF32 = castF32 is defined and castF32 %}
19
  {% set scalar = "f32" if castF32 else scalar %}
20
  {% set xv = ("vec4<f32>(" if vec4 else "f32(") if castF32 else "" %}
21
  {% set vx = ")" if castF32 else "" %}
22
  {% set yv = "f16(" if castF32 else "" %}
23
  {% set vy = ")" if castF32 else "" %}
24
- enable subgroups;
25
- {% if usesF16Spec is defined and usesF16Spec %}
26
- enable f16;
27
- {% endif %}
28
- {{ env.wgsl.resourceDeclarations }}
29
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
30
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
31
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
@@ -39,15 +33,15 @@ fn {{ name }}() -> {{ scalar }} {
39
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
40
  return bitcast<f32>(bits);
41
  {% endif %}
42
- }
43
- {%- endmacro %}
44
-
45
 
46
  const WG: u32 = {{ workgroupSize }}u;
47
- {%- if op == "max" or op == "min" %}
 
48
  {{ wgsl_minmax_identity("reduction_identity", op, scalar) }}
49
- {%- endif %}
50
- {%- if op == "logsumexp" %}
51
 
52
  const F32_MIN: f32 = -3.4028234663852886e38;
53
  const F32_MAX: f32 = 3.4028234663852886e38;
@@ -56,7 +50,7 @@ fn is_nan_f32(value: f32) -> bool {
56
  let bits = bitcast<u32>(value);
57
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
58
  }
59
- {%- endif %}
60
 
61
  var<workgroup> wgPartial: array<{{ scalar }}, WG>;
62
 
@@ -77,20 +71,19 @@ fn {{ name }}(value: {{ scalar }}, sgLid: u32, sgId: u32, numSg: u32) -> {{ scal
77
  workgroupBarrier();
78
  return total;
79
  }
80
- {%- endmacro %}
81
- {%- if op == "max" %}
82
  {{ emit_reduce("reduce_row", "subgroupMax", "total = max(total, wgPartial[i]);") }}
83
- {%- elif op == "min" %}
84
  {{ emit_reduce("reduce_row", "subgroupMin", "total = min(total, wgPartial[i]);") }}
85
- {%- elif op == "prod" %}
86
  {{ emit_reduce("reduce_row", "subgroupMul", "total = total * wgPartial[i];") }}
87
- {%- elif op == "logsumexp" %}
88
  {{ emit_reduce("reduce_row_add", "subgroupAdd", "total = total + wgPartial[i];") }}
89
  {{ emit_reduce("reduce_row_max", "subgroupMax", "total = max(total, wgPartial[i]);") }}
90
- {%- else %}
91
  {{ emit_reduce("reduce_row", "subgroupAdd", "total = total + wgPartial[i];") }}
92
- {%- endif %}
93
-
94
  @compute @workgroup_size(WG, 1, 1)
95
  fn main(@builtin(workgroup_id) wg: vec3<u32>,
96
  @builtin(local_invocation_id) lid: vec3<u32>,
@@ -103,13 +96,13 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>,
103
  }
104
  let tid = lid.x;
105
  let base = row * params.chunkCount;
106
- {%- if op == "logsumexp" %}
107
  var localMax = F32_MIN;
108
  var localNan = 0.0;
109
  var localNanValue = 0.0;
110
  for (var c = tid; c < params.chunkCount; c = c + WG) {
111
  let v = {{ xv }}x[base + c]{{ vx }};
112
- {%- if vec4 %}
113
  {% for comp in ["x", "y", "z", "w"] %}
114
  if (is_nan_f32(v.{{ comp }})) {
115
  localNan = 1.0;
@@ -117,7 +110,7 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>,
117
  } else {
118
  localMax = max(localMax, v.{{ comp }});
119
  }
120
- {%- endfor %}
121
  {% else %}
122
  if (is_nan_f32(v)) {
123
  localNan = 1.0;
@@ -125,7 +118,7 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>,
125
  } else {
126
  localMax = max(localMax, v);
127
  }
128
- {%- endif %}
129
  }
130
  let rowMax = reduce_row_max(localMax, sgLid, sgId, numSg);
131
  let nanCount = reduce_row_add(localNan, sgLid, sgId, numSg);
@@ -135,83 +128,83 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>,
135
  var acc = 0.0;
136
  for (var c = tid; c < params.chunkCount; c = c + WG) {
137
  let v = {{ xv }}x[base + c]{{ vx }};
138
- {%- if vec4 %}
139
  let e = select(exp(v - vec4<f32>(rowMax)), vec4<f32>(0.0), hasPositiveInf || hasNan);
140
  acc = acc + ((e.x + e.y) + (e.z + e.w));
141
- {%- else %}
142
  acc = acc + select(exp(v - rowMax), 0.0, hasPositiveInf || hasNan);
143
- {%- endif %}
144
  }
145
  let sum = reduce_row_add(acc, sgLid, sgId, numSg);
146
  if (tid == 0u) {
147
  let finiteOrInf = select(rowMax + log(sum), rowMax, hasPositiveInf);
148
  y[row] = {{ yv }}select(finiteOrInf, nanValue, hasNan){{ vy }};
149
  }
150
- {%- else %}
151
  {% if op == "max" or op == "min" %}
152
  let INIT: {{ scalar }} = reduction_identity();
153
- {%- elif op == "prod" %}
154
  let INIT: f32 = 1.0;
155
- {%- else %}
156
  let INIT: f32 = 0.0;
157
- {%- endif %}
158
  {% if vec4 %}
159
  var acc4 = vec4<{{ scalar }}>(INIT);
160
  for (var c = tid; c < params.chunkCount; c = c + WG) {
161
  let v = {{ xv }}x[base + c]{{ vx }};
162
- {%- if op == "max" %}
163
  acc4 = max(acc4, v);
164
- {%- elif op == "min" %}
165
  acc4 = min(acc4, v);
166
- {%- elif op == "prod" %}
167
  acc4 = acc4 * v;
168
- {%- elif op == "l1" %}
169
  acc4 = acc4 + abs(v);
170
- {%- elif op == "l2" or op == "sumsquare" %}
171
  acc4 = acc4 + v * v;
172
- {%- else %}
173
  acc4 = acc4 + v;
174
- {%- endif %}
175
  }
176
- {%- if op == "max" %}
177
  let acc = max(max(acc4.x, acc4.y), max(acc4.z, acc4.w));
178
- {%- elif op == "min" %}
179
  let acc = min(min(acc4.x, acc4.y), min(acc4.z, acc4.w));
180
- {%- elif op == "prod" %}
181
  let acc = (acc4.x * acc4.y) * (acc4.z * acc4.w);
182
- {%- else %}
183
  let acc = (acc4.x + acc4.y) + (acc4.z + acc4.w);
184
- {%- endif %}
185
  {% else %}
186
  var acc = INIT;
187
  for (var c = tid; c < params.chunkCount; c = c + WG) {
188
  let v = {{ xv }}x[base + c]{{ vx }};
189
- {%- if op == "max" %}
190
  acc = max(acc, v);
191
- {%- elif op == "min" %}
192
  acc = min(acc, v);
193
- {%- elif op == "prod" %}
194
  acc = acc * v;
195
- {%- elif op == "l1" %}
196
  acc = acc + abs(v);
197
- {%- elif op == "l2" or op == "sumsquare" %}
198
  acc = acc + v * v;
199
- {%- else %}
200
  acc = acc + v;
201
- {%- endif %}
202
  }
203
- {%- endif %}
204
  let total = reduce_row(acc, sgLid, sgId, numSg);
205
  if (tid == 0u) {
206
- {%- if op == "mean" %}
207
  y[row] = {{ yv }}total / f32(params.cols){{ vy }};
208
- {%- elif op == "l2" %}
209
  y[row] = {{ yv }}sqrt(total){{ vy }};
210
- {%- elif op == "logsum" %}
211
  y[row] = {{ yv }}log(total){{ vy }};
212
- {%- else %}
213
  y[row] = {{ yv }}total{{ vy }};
214
- {%- endif %}
215
  }
216
- {%- endif %}
217
  }
 
15
  // IEEE-754 bit patterns because WGSL rejects infinite constants.
16
  {% endif %}
17
  // f16 storage widens before accumulation and narrows only at the final store.
 
18
  {% set scalar = "f32" if castF32 else scalar %}
19
  {% set xv = ("vec4<f32>(" if vec4 else "f32(") if castF32 else "" %}
20
  {% set vx = ")" if castF32 else "" %}
21
  {% set yv = "f16(" if castF32 else "" %}
22
  {% set vy = ")" if castF32 else "" %}
 
 
 
 
 
23
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
24
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
25
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
 
33
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
34
  return bitcast<f32>(bits);
35
  {% endif %}
36
+ }{% endmacro %}
37
+ enable subgroups;
38
+ {{ env.wgsl.resourceDeclarations }}
39
 
40
  const WG: u32 = {{ workgroupSize }}u;
41
+ {% if op == "max" or op == "min" %}
42
+
43
  {{ wgsl_minmax_identity("reduction_identity", op, scalar) }}
44
+ {% elif op == "logsumexp" %}
 
45
 
46
  const F32_MIN: f32 = -3.4028234663852886e38;
47
  const F32_MAX: f32 = 3.4028234663852886e38;
 
50
  let bits = bitcast<u32>(value);
51
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
52
  }
53
+ {% endif %}
54
 
55
  var<workgroup> wgPartial: array<{{ scalar }}, WG>;
56
 
 
71
  workgroupBarrier();
72
  return total;
73
  }
74
+ {% endmacro %}
75
+ {% if op == "max" %}
76
  {{ emit_reduce("reduce_row", "subgroupMax", "total = max(total, wgPartial[i]);") }}
77
+ {% elif op == "min" %}
78
  {{ emit_reduce("reduce_row", "subgroupMin", "total = min(total, wgPartial[i]);") }}
79
+ {% elif op == "prod" %}
80
  {{ emit_reduce("reduce_row", "subgroupMul", "total = total * wgPartial[i];") }}
81
+ {% elif op == "logsumexp" %}
82
  {{ emit_reduce("reduce_row_add", "subgroupAdd", "total = total + wgPartial[i];") }}
83
  {{ emit_reduce("reduce_row_max", "subgroupMax", "total = max(total, wgPartial[i]);") }}
84
+ {% else %}
85
  {{ emit_reduce("reduce_row", "subgroupAdd", "total = total + wgPartial[i];") }}
86
+ {% endif %}
 
87
  @compute @workgroup_size(WG, 1, 1)
88
  fn main(@builtin(workgroup_id) wg: vec3<u32>,
89
  @builtin(local_invocation_id) lid: vec3<u32>,
 
96
  }
97
  let tid = lid.x;
98
  let base = row * params.chunkCount;
99
+ {% if op == "logsumexp" %}
100
  var localMax = F32_MIN;
101
  var localNan = 0.0;
102
  var localNanValue = 0.0;
103
  for (var c = tid; c < params.chunkCount; c = c + WG) {
104
  let v = {{ xv }}x[base + c]{{ vx }};
105
+ {% if vec4 %}
106
  {% for comp in ["x", "y", "z", "w"] %}
107
  if (is_nan_f32(v.{{ comp }})) {
108
  localNan = 1.0;
 
110
  } else {
111
  localMax = max(localMax, v.{{ comp }});
112
  }
113
+ {% endfor %}
114
  {% else %}
115
  if (is_nan_f32(v)) {
116
  localNan = 1.0;
 
118
  } else {
119
  localMax = max(localMax, v);
120
  }
121
+ {% endif %}
122
  }
123
  let rowMax = reduce_row_max(localMax, sgLid, sgId, numSg);
124
  let nanCount = reduce_row_add(localNan, sgLid, sgId, numSg);
 
128
  var acc = 0.0;
129
  for (var c = tid; c < params.chunkCount; c = c + WG) {
130
  let v = {{ xv }}x[base + c]{{ vx }};
131
+ {% if vec4 %}
132
  let e = select(exp(v - vec4<f32>(rowMax)), vec4<f32>(0.0), hasPositiveInf || hasNan);
133
  acc = acc + ((e.x + e.y) + (e.z + e.w));
134
+ {% else %}
135
  acc = acc + select(exp(v - rowMax), 0.0, hasPositiveInf || hasNan);
136
+ {% endif %}
137
  }
138
  let sum = reduce_row_add(acc, sgLid, sgId, numSg);
139
  if (tid == 0u) {
140
  let finiteOrInf = select(rowMax + log(sum), rowMax, hasPositiveInf);
141
  y[row] = {{ yv }}select(finiteOrInf, nanValue, hasNan){{ vy }};
142
  }
143
+ {% else %}
144
  {% if op == "max" or op == "min" %}
145
  let INIT: {{ scalar }} = reduction_identity();
146
+ {% elif op == "prod" %}
147
  let INIT: f32 = 1.0;
148
+ {% else %}
149
  let INIT: f32 = 0.0;
150
+ {% endif %}
151
  {% if vec4 %}
152
  var acc4 = vec4<{{ scalar }}>(INIT);
153
  for (var c = tid; c < params.chunkCount; c = c + WG) {
154
  let v = {{ xv }}x[base + c]{{ vx }};
155
+ {% if op == "max" %}
156
  acc4 = max(acc4, v);
157
+ {% elif op == "min" %}
158
  acc4 = min(acc4, v);
159
+ {% elif op == "prod" %}
160
  acc4 = acc4 * v;
161
+ {% elif op == "l1" %}
162
  acc4 = acc4 + abs(v);
163
+ {% elif op == "l2" or op == "sumsquare" %}
164
  acc4 = acc4 + v * v;
165
+ {% else %}
166
  acc4 = acc4 + v;
167
+ {% endif %}
168
  }
169
+ {% if op == "max" %}
170
  let acc = max(max(acc4.x, acc4.y), max(acc4.z, acc4.w));
171
+ {% elif op == "min" %}
172
  let acc = min(min(acc4.x, acc4.y), min(acc4.z, acc4.w));
173
+ {% elif op == "prod" %}
174
  let acc = (acc4.x * acc4.y) * (acc4.z * acc4.w);
175
+ {% else %}
176
  let acc = (acc4.x + acc4.y) + (acc4.z + acc4.w);
177
+ {% endif %}
178
  {% else %}
179
  var acc = INIT;
180
  for (var c = tid; c < params.chunkCount; c = c + WG) {
181
  let v = {{ xv }}x[base + c]{{ vx }};
182
+ {% if op == "max" %}
183
  acc = max(acc, v);
184
+ {% elif op == "min" %}
185
  acc = min(acc, v);
186
+ {% elif op == "prod" %}
187
  acc = acc * v;
188
+ {% elif op == "l1" %}
189
  acc = acc + abs(v);
190
+ {% elif op == "l2" or op == "sumsquare" %}
191
  acc = acc + v * v;
192
+ {% else %}
193
  acc = acc + v;
194
+ {% endif %}
195
  }
196
+ {% endif %}
197
  let total = reduce_row(acc, sgLid, sgId, numSg);
198
  if (tid == 0u) {
199
+ {% if op == "mean" %}
200
  y[row] = {{ yv }}total / f32(params.cols){{ vy }};
201
+ {% elif op == "l2" %}
202
  y[row] = {{ yv }}sqrt(total){{ vy }};
203
+ {% elif op == "logsum" %}
204
  y[row] = {{ yv }}log(total){{ vy }};
205
+ {% else %}
206
  y[row] = {{ yv }}total{{ vy }};
207
+ {% endif %}
208
  }
209
+ {% endif %}
210
  }
build/webgpu/reduce-row-tree.wgsl.jinja CHANGED
@@ -16,15 +16,11 @@
16
  // then narrows only at the final store.
17
  {% set isVec4 = vec4 is defined and vec4 %}
18
  {% set rowIsEmpty = "params.chunkCount == 0u" if isVec4 else "params.cols == 0u" %}
19
- {% set castF32 = castF32 is defined and castF32 %}
20
  {% set scalar = "f32" if castF32 else scalar %}
21
  {% set xv = ("vec4<f32>(" if vec4 else "f32(") if castF32 else "" %}
22
  {% set vx = ")" if castF32 else "" %}
23
  {% set yv = "f16(" if castF32 else "" %}
24
  {% set vy = ")" if castF32 else "" %}
25
- {% if usesF16Spec is defined and usesF16Spec %}
26
- enable f16;
27
- {% endif %}
28
  {{ env.wgsl.resourceDeclarations }}
29
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
30
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
@@ -39,15 +35,12 @@ fn {{ name }}() -> {{ scalar }} {
39
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
40
  return bitcast<f32>(bits);
41
  {% endif %}
42
- }
43
- {%- endmacro %}
44
-
45
  {% if scalar != "i32" and scalar != "u32" and (op == "logsum" or op == "logsumexp") %}
46
  fn negative_infinity() -> f32 {
47
  var bits = 0xff800000u;
48
  return bitcast<f32>(bits);
49
  }
50
-
51
  {% endif %}
52
 
53
  const WG: u32 = {{ workgroupSize }}u;
@@ -57,8 +50,8 @@ const F32_MIN: f32 = -3.4028234663852886e38;
57
  const F32_MAX: f32 = 3.4028234663852886e38;
58
 
59
  var<workgroup> partial: array<f32, WG>;
60
- {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true) %}
61
- fn {{ name }}(value: f32, tid: u32) -> f32 {
62
  {{ buffer }}[tid] = value;
63
  workgroupBarrier();
64
  // Ceil-halving keeps every lane when the workgroup size is not a power of
@@ -84,20 +77,15 @@ fn {{ name }}(value: f32, tid: u32) -> f32 {
84
  // slot 0 here, so the next call's first store must not run until all lanes have read it.
85
  // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
86
  let reduced = {{ buffer }}[0];
87
- {% if trailingBarrier %}
88
  workgroupBarrier();
89
- {% endif %}
90
  return reduced;
91
  }
92
  {% endmacro %}
93
-
94
  {{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }}
95
  {{ wgsl_tree_reduce_f32("reduce_max", "max", "partial", "WG") }}
96
-
97
  fn is_nan_f32(value: f32) -> bool {
98
  let bits = bitcast<u32>(value);
99
- return (bits & 0x7f800000u) == 0x7f800000u
100
- && (bits & 0x007fffffu) != 0u;
101
  }
102
  {% else %}
103
  {% set is_int = scalar == "i32" or scalar == "u32" %}
@@ -216,13 +204,14 @@ fn main(@builtin(workgroup_id) wg: vec3<u32>,
216
  {% endif %}
217
  return;
218
  }
 
219
  {% elif op == "logsum" %}
220
  if ({{ rowIsEmpty }}) {
221
  if (tid == 0u) { y[row] = {{ yv }}negative_infinity(){{ vy }}; }
222
  return;
223
  }
224
- {% endif %}
225
 
 
226
  {% if vec4 %}
227
  var acc4 = vec4<{{ accType }}>(identity());
228
  for (var col = tid; col < params.chunkCount; col = col + WG) {
 
16
  // then narrows only at the final store.
17
  {% set isVec4 = vec4 is defined and vec4 %}
18
  {% set rowIsEmpty = "params.chunkCount == 0u" if isVec4 else "params.cols == 0u" %}
 
19
  {% set scalar = "f32" if castF32 else scalar %}
20
  {% set xv = ("vec4<f32>(" if vec4 else "f32(") if castF32 else "" %}
21
  {% set vx = ")" if castF32 else "" %}
22
  {% set yv = "f16(" if castF32 else "" %}
23
  {% set vy = ")" if castF32 else "" %}
 
 
 
24
  {{ env.wgsl.resourceDeclarations }}
25
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
26
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
 
35
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
36
  return bitcast<f32>(bits);
37
  {% endif %}
38
+ }{% endmacro %}
 
 
39
  {% if scalar != "i32" and scalar != "u32" and (op == "logsum" or op == "logsumexp") %}
40
  fn negative_infinity() -> f32 {
41
  var bits = 0xff800000u;
42
  return bitcast<f32>(bits);
43
  }
 
44
  {% endif %}
45
 
46
  const WG: u32 = {{ workgroupSize }}u;
 
50
  const F32_MAX: f32 = 3.4028234663852886e38;
51
 
52
  var<workgroup> partial: array<f32, WG>;
53
+ {% macro wgsl_tree_reduce_f32(name, mode, buffer="partial", wg="WG", trailingBarrier=true, valueType="f32") %}
54
+ fn {{ name }}(value: {{ valueType }}, tid: u32) -> {{ valueType }} {
55
  {{ buffer }}[tid] = value;
56
  workgroupBarrier();
57
  // Ceil-halving keeps every lane when the workgroup size is not a power of
 
77
  // slot 0 here, so the next call's first store must not run until all lanes have read it.
78
  // `trailingBarrier=false` is safe only when the buffer is never written again before kernel exit.
79
  let reduced = {{ buffer }}[0];
 
80
  workgroupBarrier();
 
81
  return reduced;
82
  }
83
  {% endmacro %}
 
84
  {{ wgsl_tree_reduce_f32("reduce_sum", "add", "partial", "WG") }}
85
  {{ wgsl_tree_reduce_f32("reduce_max", "max", "partial", "WG") }}
 
86
  fn is_nan_f32(value: f32) -> bool {
87
  let bits = bitcast<u32>(value);
88
+ return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
 
89
  }
90
  {% else %}
91
  {% set is_int = scalar == "i32" or scalar == "u32" %}
 
204
  {% endif %}
205
  return;
206
  }
207
+
208
  {% elif op == "logsum" %}
209
  if ({{ rowIsEmpty }}) {
210
  if (tid == 0u) { y[row] = {{ yv }}negative_infinity(){{ vy }}; }
211
  return;
212
  }
 
213
 
214
+ {% endif %}
215
  {% if vec4 %}
216
  var acc4 = vec4<{{ accType }}>(identity());
217
  for (var col = tid; col < params.chunkCount; col = col + WG) {
build/webgpu/reduce-serial-axis.wgsl.jinja CHANGED
@@ -1,5 +1,4 @@
1
  {% macro reduce_multi_axis_offset(hasReduced) %}
2
-
3
  // One thread per output element walks the Cartesian product of the reduced axes,
4
  // linearized as reduce_linear. Specialized shapes make every input offset a sum
5
  // of coordinate-times-constant terms.
@@ -44,18 +43,20 @@ fn input_offset(out_index: u32{% if hasReduced %}, reduce_linear: u32{% endif %}
44
  {% set src.value = "(" ~ src.value ~ " * " ~ dataShape[a] ~ "u + coord" ~ a ~ ")" %}
45
  {% endfor %}
46
  return {{ src.value }};
47
- }
48
- {%- endmacro %}
49
  // Serial one-thread-per-output reduction for the no-feature tier. f16 storage
50
  // is widened before every accumulation and narrowed only for the final store.
51
- {% set castF32 = castF32 is defined and castF32 %}
52
  {% set logicalBool = logicalBool is defined and logicalBool %}
53
- {% set intMode = intMode is defined and intMode %}
54
  {% set yv = "f16(" if castF32 else "" %}
55
  {% set vy = ")" if castF32 else "" %}
56
- {% if usesF16Spec is defined and usesF16Spec %}
57
- enable f16;
58
- {% endif %}
 
 
 
 
 
59
  {{ env.wgsl.resourceDeclarations }}
60
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
61
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
@@ -70,15 +71,12 @@ fn {{ name }}() -> {{ scalar }} {
70
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
71
  return bitcast<f32>(bits);
72
  {% endif %}
73
- }
74
- {%- endmacro %}
75
-
76
  {% if not intMode and (op == "logsum" or op == "logsumexp") %}
77
  fn negative_infinity() -> f32 {
78
  var bits = 0xff800000u;
79
  return bitcast<f32>(bits);
80
  }
81
-
82
  {% endif %}
83
  {% if op == "max" or op == "min" %}
84
 
@@ -131,7 +129,8 @@ fn input_offset(out_index: u32, reduce_index: u32) -> u32 {
131
  {% if indexing == "multiaxis" %}
132
  {% set hasReducedAxis = namespace(value=false) %}
133
  {% for a in range(rank) %}{% if reduce[a] %}{% set hasReducedAxis.value = true %}{% endif %}{% endfor %}
134
- {{- reduce_multi_axis_offset(hasReducedAxis.value) }}
 
135
  {% endif %}
136
  {% if indexing == "multiaxis" %}
137
  {% set mcount = namespace(value=1) %}
@@ -166,12 +165,7 @@ fn input_offset(out_index: u32, reduce_index: u32) -> u32 {
166
 
167
  @compute @workgroup_size({{ reduceWorkgroupSize }})
168
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
169
- // 2D-folded flat index: gid.y carries the high bits past the
170
- // per-axis dispatch fold width (outputs > 16.7M elements).
171
- let i = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ reduceWorkgroupSize }}u;
172
- if (i >= params.outCount) {
173
- return;
174
- }
175
  {% if op == "logsumexp" %}
176
  {% if intMode %}
177
  // Integer logsumexp widens each element for exp/log, then truncates the result
 
1
  {% macro reduce_multi_axis_offset(hasReduced) %}
 
2
  // One thread per output element walks the Cartesian product of the reduced axes,
3
  // linearized as reduce_linear. Specialized shapes make every input offset a sum
4
  // of coordinate-times-constant terms.
 
43
  {% set src.value = "(" ~ src.value ~ " * " ~ dataShape[a] ~ "u + coord" ~ a ~ ")" %}
44
  {% endfor %}
45
  return {{ src.value }};
46
+ }{% endmacro %}
 
47
  // Serial one-thread-per-output reduction for the no-feature tier. f16 storage
48
  // is widened before every accumulation and narrowed only for the final store.
 
49
  {% set logicalBool = logicalBool is defined and logicalBool %}
 
50
  {% set yv = "f16(" if castF32 else "" %}
51
  {% set vy = ")" if castF32 else "" %}
52
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
53
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
54
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
55
+ // per-axis workgroup fold width.
56
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};
57
+ if ({{ name }} >= {{ bound }}) {
58
+ return;
59
+ }{% endmacro %}
60
  {{ env.wgsl.resourceDeclarations }}
61
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
62
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
 
71
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
72
  return bitcast<f32>(bits);
73
  {% endif %}
74
+ }{% endmacro %}
 
 
75
  {% if not intMode and (op == "logsum" or op == "logsumexp") %}
76
  fn negative_infinity() -> f32 {
77
  var bits = 0xff800000u;
78
  return bitcast<f32>(bits);
79
  }
 
80
  {% endif %}
81
  {% if op == "max" or op == "min" %}
82
 
 
129
  {% if indexing == "multiaxis" %}
130
  {% set hasReducedAxis = namespace(value=false) %}
131
  {% for a in range(rank) %}{% if reduce[a] %}{% set hasReducedAxis.value = true %}{% endif %}{% endfor %}
132
+
133
+ {{ reduce_multi_axis_offset(hasReducedAxis.value) }}
134
  {% endif %}
135
  {% if indexing == "multiaxis" %}
136
  {% set mcount = namespace(value=1) %}
 
165
 
166
  @compute @workgroup_size({{ reduceWorkgroupSize }})
167
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
168
+ {{ flat_index_2d(reduceWorkgroupSize, bound="params.outCount") }}
 
 
 
 
 
169
  {% if op == "logsumexp" %}
170
  {% if intMode %}
171
  // Integer logsumexp widens each element for exp/log, then truncates the result
build/webgpu/reduce-strided-axis.wgsl.jinja CHANGED
@@ -1,6 +1,3 @@
1
- {% if usesF16 %}
2
- enable f16;
3
- {% endif %}
4
  {{ env.wgsl.resourceDeclarations }}
5
 
6
  @compute @workgroup_size({{ workgroupSize }})
@@ -15,7 +12,13 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
15
  var acc = 0.0;
16
  for (var r = 0u; r < {{ axisDim }}u; r++) {
17
  let value = {% if usesF16 %}f32({% endif %}x[base + r * {{ innerSize }}u]{% if usesF16 %}){% endif %};
 
 
 
 
 
18
  acc = acc + value;
 
19
  }
20
- y[i] = {% if usesF16 %}f16({% endif %}acc / {{ axisDim }}.0{% if usesF16 %}){% endif %};
21
  }
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  @compute @workgroup_size({{ workgroupSize }})
 
12
  var acc = 0.0;
13
  for (var r = 0u; r < {{ axisDim }}u; r++) {
14
  let value = {% if usesF16 %}f32({% endif %}x[base + r * {{ innerSize }}u]{% if usesF16 %}){% endif %};
15
+ {% if op == "l2" or op == "sumsquare" %}
16
+ acc = acc + value * value;
17
+ {% elif op == "l1" %}
18
+ acc = acc + abs(value);
19
+ {% else %}
20
  acc = acc + value;
21
+ {% endif %}
22
  }
23
+ y[i] = {% if usesF16 %}f16({% endif %}{% if op == "l2" %}sqrt(acc){% elif op == "mean" %}acc / {{ axisDim }}.0{% else %}acc{% endif %}{% if usesF16 %}){% endif %};
24
  }
build/webgpu/test.json CHANGED
@@ -901,7 +901,7 @@
901
  "outputs": { "y": { "dtype": "float32", "shape": [], "tolerance": 0.0001, "relTolerance": 0.00001 } }
902
  },
903
  {
904
- "name": "axis0_narrow_f32_8192x3_splitk_guard_lock",
905
  "provenance": {
906
  "notes": "An 8,192-by-3 axis-0 mean exercises split-K with a narrow output. Constant ones verify that finalization remains exactly one."
907
  },
@@ -1179,7 +1179,7 @@
1179
  {
1180
  "name": "int32_lastaxis_subgroup_rows_64x1024",
1181
  "provenance": {
1182
- "notes": "Route lock for the integer subgroup-per-row last-axis mean: 64 rows of 1024 int32 values sum in the output type and divide as integers, matching the tree kernel bit for bit; a five-value cycle shifts phase every row so the truncated means differ (11 or 12)."
1183
  },
1184
  "attrs": { "axes": [1], "keepdims": 0 },
1185
  "inputs": {
 
901
  "outputs": { "y": { "dtype": "float32", "shape": [], "tolerance": 0.0001, "relTolerance": 0.00001 } }
902
  },
903
  {
904
+ "name": "axis0_narrow_f32_8192x3_splitk",
905
  "provenance": {
906
  "notes": "An 8,192-by-3 axis-0 mean exercises split-K with a narrow output. Constant ones verify that finalization remains exactly one."
907
  },
 
1179
  {
1180
  "name": "int32_lastaxis_subgroup_rows_64x1024",
1181
  "provenance": {
1182
+ "notes": "Sixty-four rows of 1024 int32 values check integer mean and truncation; a five-value cycle shifts phase so neighboring row results differ (11 or 12)."
1183
  },
1184
  "attrs": { "axes": [1], "keepdims": 0 },
1185
  "inputs": {