Xenova HF Staff commited on
Commit
c61b7e6
·
verified ·
1 Parent(s): ea7ecc1

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -44,6 +44,13 @@ Default values (overridable per request):
44
  | --- | --- |
45
  | `T` | `float32`, `float16`, `int32` |
46
 
 
 
 
 
 
 
 
47
  ## Device requirements
48
 
49
  Some implementation variants require `subgroups`. These are route-specific capabilities, not package-wide requirements; availability also depends on the request shape and dtype.
@@ -53,7 +60,7 @@ Some implementation variants require `subgroups`. These are route-specific capab
53
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
54
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
55
  - [`test.json`](build/webgpu/test.json) — correctness cases
56
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
57
  - [`reduce-axis-split-reduce.wgsl.jinja`](build/webgpu/reduce-axis-split-reduce.wgsl.jinja)
58
  - [`reduce-axis0-splitk-combine.wgsl.jinja`](build/webgpu/reduce-axis0-splitk-combine.wgsl.jinja)
59
  - [`reduce-axis0-splitk-reduce.wgsl.jinja`](build/webgpu/reduce-axis0-splitk-reduce.wgsl.jinja)
@@ -70,7 +77,7 @@ Some implementation variants require `subgroups`. These are route-specific capab
70
  ## Use with `@huggingface/kernels`
71
 
72
  ```sh
73
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
74
  ```
75
 
76
  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.
@@ -89,7 +96,7 @@ import { getKernel } from "@huggingface/kernels";
89
 
90
  const kernel = await getKernel("webgpu-kernels/ai.onnx.ReduceLogSumExp", { version: 1 });
91
  // Explicit destinations request optional results or supply metadata that cannot be inferred.
92
- const { y } = await kernel({ x: { data: xData, shape: [] } }, {
93
- outputs: { y: { shape: [], dtype: "float32" } },
94
  });
95
  ```
 
44
  | --- | --- |
45
  | `T` | `float32`, `float16`, `int32` |
46
 
47
+ ## Implementation variants
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
+
54
  ## Device requirements
55
 
56
  Some implementation variants require `subgroups`. These are route-specific capabilities, not package-wide requirements; availability also depends on the request shape and dtype.
 
60
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
61
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
62
  - [`test.json`](build/webgpu/test.json) — correctness cases
63
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
64
  - [`reduce-axis-split-reduce.wgsl.jinja`](build/webgpu/reduce-axis-split-reduce.wgsl.jinja)
65
  - [`reduce-axis0-splitk-combine.wgsl.jinja`](build/webgpu/reduce-axis0-splitk-combine.wgsl.jinja)
66
  - [`reduce-axis0-splitk-reduce.wgsl.jinja`](build/webgpu/reduce-axis0-splitk-reduce.wgsl.jinja)
 
77
  ## Use with `@huggingface/kernels`
78
 
79
  ```sh
80
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
81
  ```
82
 
83
  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.
 
96
 
97
  const kernel = await getKernel("webgpu-kernels/ai.onnx.ReduceLogSumExp", { version: 1 });
98
  // Explicit destinations request optional results or supply metadata that cannot be inferred.
99
+ const { y } = await kernel({ x: { data: xData, shape: [3, 2, 2] } }, {
100
+ outputs: { y: { shape: [1, 1, 1], dtype: "float32" } },
101
  });
102
  ```
build/webgpu/bench.json CHANGED
@@ -75,6 +75,32 @@
75
  ]
76
  }
77
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
78
  {
79
  "name": "fullreduce-single-lane-r3-101x101x101-numel-not-mult4",
80
  "preset": "stress",
@@ -104,7 +130,7 @@
104
  "name": "reducelogsumexp-spatial-axes12-f32-16x256x256-low-lane",
105
  "preset": "stress",
106
  "provenance": {
107
- "source": "repository-authored",
108
  "notes": "Low-lane rank-3 contiguous-suffix reduction that verifies the vec4 subgroup reducer and no-subgroup workgroup-tree fallback."
109
  },
110
  "vars": { "batch": 16, "height": 256, "width": 256 },
 
75
  ]
76
  }
77
  },
78
+ {
79
+ "name": "axis0-crossover-32768x512",
80
+ "preset": "smoke",
81
+ "vars": { "rows": 32768, "cols": 512 },
82
+ "attrs": { "axes": [0], "keepdims": 0 },
83
+ "inputs": { "x": { "shape": [32768, 512], "dtype": "float32", "dist": "normal", "scale": 0.2, "seed": 313 } },
84
+ "outputs": { "y": { "shape": [512], "dtype": "float32" } },
85
+ "bench": {
86
+ "metrics": [
87
+ { "type": "bandwidth", "name": "stable max-plus-exp traffic", "value": "2 * args.rows * args.cols * 4" }
88
+ ]
89
+ }
90
+ },
91
+ {
92
+ "name": "axis0-crossover-32768x1024",
93
+ "preset": "smoke",
94
+ "vars": { "rows": 32768, "cols": 1024 },
95
+ "attrs": { "axes": [0], "keepdims": 0 },
96
+ "inputs": { "x": { "shape": [32768, 1024], "dtype": "float32", "dist": "normal", "scale": 0.2, "seed": 314 } },
97
+ "outputs": { "y": { "shape": [1024], "dtype": "float32" } },
98
+ "bench": {
99
+ "metrics": [
100
+ { "type": "bandwidth", "name": "stable max-plus-exp traffic", "value": "2 * args.rows * args.cols * 4" }
101
+ ]
102
+ }
103
+ },
104
  {
105
  "name": "fullreduce-single-lane-r3-101x101x101-numel-not-mult4",
106
  "preset": "stress",
 
130
  "name": "reducelogsumexp-spatial-axes12-f32-16x256x256-low-lane",
131
  "preset": "stress",
132
  "provenance": {
133
+ "source": "synthetic",
134
  "notes": "Low-lane rank-3 contiguous-suffix reduction that verifies the vec4 subgroup reducer and no-subgroup workgroup-tree fallback."
135
  },
136
  "vars": { "batch": 16, "height": 256, "width": 256 },
build/webgpu/manifest.json CHANGED
@@ -34,7 +34,10 @@
34
  "ROW_SERIAL_MAX_COLS": { "default": 1024 }
35
  },
36
  "derive": {
 
37
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
 
 
38
  "reduceWorkgroupSize": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)",
39
  "treeWorkgroupOk": "reduceWorkgroupSize > 0 and pow2ceil(reduceWorkgroupSize) == reduceWorkgroupSize and reduceWorkgroupSize * dtypeBytes(\"float32\") <= device.limits.maxComputeWorkgroupStorageSize",
40
  "subgroupWorkgroupFloor": "min(reduceWorkgroupSize, max(1, device.adapterInfo.subgroupMaxSize))",
@@ -62,127 +65,117 @@
62
  "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)))"
63
  },
64
  "bindings": {
65
- "x": { "buffer": "read-only-storage", "elementType": "$vectorScalar" },
66
- "y": { "buffer": "storage", "elementType": "$T" },
67
- "params": {
68
- "buffer": "uniform",
69
  "struct": [
70
  { "name": "rows", "type": "u32", "value": "numel(shapes.y)" },
71
  { "name": "chunkCount", "type": "u32", "value": "numel(shapes.x) / numel(shapes.y) / tunables.VECTOR_WIDTH" }
72
  ]
73
  },
74
- "x_2": { "name": "x", "buffer": "read-only-storage", "elementType": "$T" },
75
- "params_2": {
76
  "name": "params",
77
- "buffer": "uniform",
78
  "struct": [
79
- { "name": "rows", "type": "u32", "value": "numel(shapes.y)" },
80
- { "name": "cols", "type": "u32", "value": "numel(shapes.x) / numel(shapes.y)" }
81
  ]
82
  },
83
- "params_3": {
84
  "name": "params",
85
- "buffer": "uniform",
86
- "struct": [{ "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }]
87
  },
88
- "params_5": {
89
  "name": "params",
90
- "buffer": "uniform",
91
- "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.y)" }]
92
  },
93
- "params_6": {
94
  "name": "params",
95
- "buffer": "uniform",
96
  "struct": [
97
  { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
98
- { "name": "chunkCount", "type": "u32", "value": "dim(shapes.x, ranks.x - 1) / tunables.VECTOR_WIDTH" }
 
 
 
 
 
 
 
 
 
 
99
  ]
100
  },
101
- "params_7": {
 
 
 
 
 
102
  "name": "params",
103
- "buffer": "uniform",
104
  "struct": [
105
  { "name": "rows", "type": "u32", "value": "1" },
106
  { "name": "cols", "type": "u32", "value": "1" },
107
  { "name": "outCount", "type": "u32", "value": "1" }
108
  ]
109
  },
110
- "params_8": {
111
  "name": "params",
112
- "buffer": "uniform",
113
  "struct": [
114
  { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
115
  { "name": "cols", "type": "u32", "value": "1" },
116
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
117
  ]
118
  },
119
- "params_9": {
120
  "name": "params",
121
- "buffer": "uniform",
122
  "struct": [
123
  { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
124
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" }
125
  ]
126
  },
127
  "partials": { "buffer": "storage", "elementType": "$partialElement" },
128
- "params_10": {
129
  "name": "params",
130
- "buffer": "uniform",
131
  "struct": [
132
  { "name": "axisDim", "type": "u32", "value": "axisSplitDim" },
133
  { "name": "inner", "type": "u32", "value": "axisSplitInner" },
134
  { "name": "outputs", "type": "u32", "value": "axisSplitOutputs" }
135
  ]
136
  },
137
- "partials_2": { "name": "partials", "buffer": "read-only-storage", "elementType": "$partialElement" },
138
- "params_11": {
139
- "name": "params",
140
- "buffer": "uniform",
141
- "struct": [{ "name": "cols", "type": "u32", "value": "axisSplitOutputs" }]
142
- },
143
- "params_12": {
144
  "name": "params",
145
- "buffer": "uniform",
146
  "struct": [
147
  { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
148
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" }
149
  ]
150
  },
151
- "params_13": {
152
- "name": "params",
153
- "buffer": "uniform",
154
- "struct": [{ "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" }]
155
- },
156
- "params_16": {
157
  "name": "params",
158
- "buffer": "uniform",
159
  "struct": [
160
  { "name": "axisDim", "type": "u32", "value": "axisSplitDim" },
161
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
162
  ]
163
  },
164
- "params_17": {
165
  "name": "params",
166
- "buffer": "uniform",
167
  "struct": [
168
- { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
169
- { "name": "chunkCount", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" }
 
170
  ]
171
  },
172
- "params_18": {
173
  "name": "params",
174
- "buffer": "uniform",
175
  "struct": [
176
- { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
177
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" },
178
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
179
  ]
180
  },
181
- "params_19": {
182
  "name": "params",
183
- "buffer": "uniform",
184
  "struct": [
185
- { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" },
 
186
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
187
  ]
188
  }
@@ -203,15 +196,9 @@
203
  "id": "main",
204
  "name": "ReduceLogSumExp.ContiguousSuffixSubgroupVec4",
205
  "shader": "reduce-row-subgroup.wgsl.jinja",
206
- "derive": {
207
- "op": "\"logsumexp\"",
208
- "vec4": true,
209
- "castF32": "dtypes.T == \"f16\"",
210
- "usesF16Spec": "dtypes.T == \"f16\""
211
- },
212
- "bindings": ["x", "y", "params"],
213
- "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 },
214
- "subgroupCollectivesWidth": "portable"
215
  }
216
  ]
217
  },
@@ -229,13 +216,8 @@
229
  "id": "main",
230
  "name": "ReduceLogSumExp.ContiguousSuffixTreeVec4",
231
  "shader": "reduce-row-tree.wgsl.jinja",
232
- "derive": {
233
- "op": "\"logsumexp\"",
234
- "vec4": true,
235
- "castF32": "dtypes.T == \"f16\"",
236
- "usesF16Spec": "dtypes.T == \"f16\""
237
- },
238
- "bindings": ["x", "y", "params"],
239
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
240
  }
241
  ]
@@ -253,8 +235,7 @@
253
  "id": "main",
254
  "name": "ReduceLogSumExp.ContiguousSuffixTree",
255
  "shader": "reduce-row-tree.wgsl.jinja",
256
- "derive": { "op": "\"logsumexp\"", "castF32": "dtypes.T == \"f16\"", "usesF16Spec": "dtypes.T == \"f16\"" },
257
- "bindings": ["x_2", "y", "params_2"],
258
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
259
  }
260
  ]
@@ -270,7 +251,6 @@
270
  "name": "ReduceLogSumExp.MultiAxisRank3",
271
  "shader": "reduce-serial-axis.wgsl.jinja",
272
  "derive": {
273
- "op": "\"logsumexp\"",
274
  "indexing": "\"multiaxis\"",
275
  "rank": 3,
276
  "reduce": ["hasAxis(attrs.axes, 0, 3)", "hasAxis(attrs.axes, 1, 3)", "hasAxis(attrs.axes, 2, 3)"],
@@ -278,11 +258,9 @@
278
  "outputShape": "shapes.y",
279
  "outputRank": "ranks.y",
280
  "keepDims": "attrs.keepdims != 0",
281
- "intMode": "dtypes.T == \"i32\"",
282
- "castF32": "dtypes.T == \"f16\"",
283
- "usesF16Spec": "dtypes.T == \"f16\""
284
  },
285
- "bindings": ["x_2", "y", "params_3"],
286
  "dispatch": {
287
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
288
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -302,7 +280,6 @@
302
  "name": "ReduceLogSumExp.MultiAxisRank4",
303
  "shader": "reduce-serial-axis.wgsl.jinja",
304
  "derive": {
305
- "op": "\"logsumexp\"",
306
  "indexing": "\"multiaxis\"",
307
  "rank": 4,
308
  "reduce": ["hasAxis(attrs.axes, 0, 4)", "hasAxis(attrs.axes, 1, 4)", "hasAxis(attrs.axes, 2, 4)", "hasAxis(attrs.axes, 3, 4)"],
@@ -310,11 +287,9 @@
310
  "outputShape": "shapes.y",
311
  "outputRank": "ranks.y",
312
  "keepDims": "attrs.keepdims != 0",
313
- "intMode": "dtypes.T == \"i32\"",
314
- "castF32": "dtypes.T == \"f16\"",
315
- "usesF16Spec": "dtypes.T == \"f16\""
316
  },
317
- "bindings": ["x_2", "y", "params_3"],
318
  "dispatch": {
319
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
320
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -337,7 +312,7 @@
337
  "shader": "reduce-i32-axes02.wgsl.jinja",
338
  "derive": { "op": "\"logsumexp\"", "workgroupSizeSpec": "axes02WorkgroupSize" },
339
  "bindings": [
340
- "x_2",
341
  "y",
342
  {
343
  "name": "params",
@@ -368,7 +343,7 @@
368
  "name": "ReduceLogSumExp.NoopEmptyAxes",
369
  "shader": "reduce-noop-empty-axes.wgsl.jinja",
370
  "derive": { "op": "\"identity\"" },
371
- "bindings": ["x_2", "y", "params_5"],
372
  "dispatch": {
373
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
374
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -393,14 +368,12 @@
393
  "id": "main",
394
  "name": "ReduceLogSumExp.SubgroupRowsVec4",
395
  "shader": "reduce-row-subgroup-rows.wgsl.jinja",
396
- "derive": { "op": "\"logsumexp\"", "castF32": "dtypes.T == \"f16\"", "usesF16Spec": "dtypes.T == \"f16\"" },
397
- "bindings": ["x", "y", "params_6"],
398
  "dispatch": {
399
  "x": "min(ceilDiv(lastAxisRows, reduceWorkgroupSize / device.adapterInfo.subgroupMaxSize), 65535)",
400
  "y": "ceilDiv(ceilDiv(lastAxisRows, reduceWorkgroupSize / device.adapterInfo.subgroupMaxSize), 65535)",
401
  "z": 1
402
- },
403
- "subgroupCollectivesWidth": "portable"
404
  }
405
  ]
406
  },
@@ -419,13 +392,8 @@
419
  "id": "main",
420
  "name": "ReduceLogSumExp.TreeRowVec4",
421
  "shader": "reduce-row-tree.wgsl.jinja",
422
- "derive": {
423
- "op": "\"logsumexp\"",
424
- "vec4": true,
425
- "castF32": "dtypes.T == \"f16\"",
426
- "usesF16Spec": "dtypes.T == \"f16\""
427
- },
428
- "bindings": ["x", "y", "params_6"],
429
  "dispatch": { "x": "min(lastAxisRows, 65535)", "y": "ceilDiv(lastAxisRows, 65535)", "z": 1 }
430
  }
431
  ]
@@ -441,14 +409,11 @@
441
  "name": "ReduceLogSumExp.Rank0Scalar",
442
  "shader": "reduce-serial-axis.wgsl.jinja",
443
  "derive": {
444
- "op": "\"logsumexp\"",
445
  "indexing": "\"axis2d\"",
446
  "intMode": "dtypes.T == \"i32\"",
447
- "castF32": "dtypes.T == \"f16\"",
448
- "usesF16Spec": "dtypes.T == \"f16\"",
449
  "logicalBool": "tensorDtypes.x == \"bool\""
450
  },
451
- "bindings": ["x_2", "y", "params_7"],
452
  "dispatch": { "x": 1 }
453
  }
454
  ]
@@ -463,14 +428,11 @@
463
  "name": "ReduceLogSumExp.Rank1Axis0",
464
  "shader": "reduce-serial-axis.wgsl.jinja",
465
  "derive": {
466
- "op": "\"logsumexp\"",
467
  "indexing": "\"axis2d\"",
468
  "intMode": "dtypes.T == \"i32\"",
469
- "castF32": "dtypes.T == \"f16\"",
470
- "usesF16Spec": "dtypes.T == \"f16\"",
471
  "logicalBool": "tensorDtypes.x == \"bool\""
472
  },
473
- "bindings": ["x_2", "y", "params_8"],
474
  "dispatch": {
475
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
476
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -490,8 +452,7 @@
490
  "id": "main",
491
  "name": "ReduceLogSumExp.Axis1Parallel",
492
  "shader": "reduce-row-tree.wgsl.jinja",
493
- "derive": { "op": "\"logsumexp\"", "castF32": "dtypes.T == \"f16\"", "usesF16Spec": "dtypes.T == \"f16\"" },
494
- "bindings": ["x_2", "y", "params_9"],
495
  "dispatch": {
496
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
497
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
@@ -516,13 +477,8 @@
516
  "id": "split_reduce",
517
  "name": "ReduceLogSumExp.AxisSplitReduce",
518
  "shader": "reduce-axis-split-reduce.wgsl.jinja",
519
- "derive": {
520
- "op": "\"logsumexp\"",
521
- "splitSpec": "splitCount",
522
- "castF32": "dtypes.T == \"f16\"",
523
- "usesF16Spec": "dtypes.T == \"f16\""
524
- },
525
- "bindings": ["x_2", "partials", "params_10"],
526
  "dispatch": {
527
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), DISPATCH_FOLD_WIDTH)",
528
  "y": "splitCount",
@@ -533,8 +489,8 @@
533
  "id": "combine",
534
  "name": "ReduceLogSumExp.AxisSplitCombine",
535
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
536
- "derive": { "op": "\"logsumexp\"", "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
537
- "bindings": ["partials_2", "y", "params_11"],
538
  "dispatch": {
539
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
540
  "y": "ceilDiv(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
@@ -561,14 +517,8 @@
561
  "id": "split_reduce",
562
  "name": "ReduceLogSumExp.AxisSplitTiledReduce",
563
  "shader": "reduce-axis0-tilecols.wgsl.jinja",
564
- "derive": {
565
- "op": "\"logsumexp\"",
566
- "splitSpec": "splitCount",
567
- "tileCols": "tunables.AXIS_SPLIT_TILE_COLS",
568
- "castF32": "dtypes.T == \"f16\"",
569
- "usesF16Spec": "dtypes.T == \"f16\""
570
- },
571
- "bindings": ["x_2", "partials", "params_10"],
572
  "dispatch": {
573
  "x": "min(ceilDiv((axisSplitOutputs), (tileCols)), DISPATCH_FOLD_WIDTH)",
574
  "y": "splitCount",
@@ -579,8 +529,8 @@
579
  "id": "combine",
580
  "name": "ReduceLogSumExp.AxisSplitCombine",
581
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
582
- "derive": { "op": "\"logsumexp\"", "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
583
- "bindings": ["partials_2", "y", "params_11"],
584
  "dispatch": {
585
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
586
  "y": "ceilDiv(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
@@ -593,6 +543,7 @@
593
  "id": "axis0_splitk",
594
  "priority": 22,
595
  "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"],
 
596
  "derive": {
597
  "splitCount": "axis0SplitCount",
598
  "partialElement": "\"f32\"",
@@ -605,13 +556,8 @@
605
  "id": "split_reduce",
606
  "name": "ReduceLogSumExp.Axis0SplitKReduce",
607
  "shader": "reduce-axis0-splitk-reduce.wgsl.jinja",
608
- "derive": {
609
- "op": "\"logsumexp\"",
610
- "splitSpec": "splitCount",
611
- "castF32": "dtypes.T == \"f16\"",
612
- "usesF16Spec": "dtypes.T == \"f16\""
613
- },
614
- "bindings": ["x_2", "partials", "params_12"],
615
  "dispatch": {
616
  "x": "min(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), DISPATCH_FOLD_WIDTH)",
617
  "y": "splitCount",
@@ -622,8 +568,8 @@
622
  "id": "combine",
623
  "name": "ReduceLogSumExp.Axis0SplitKCombine",
624
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
625
- "derive": { "op": "\"logsumexp\"", "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
626
- "bindings": ["partials_2", "y", "params_13"],
627
  "dispatch": {
628
  "x": "min(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), 65535)",
629
  "y": "ceilDiv(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), 65535)",
@@ -642,8 +588,8 @@
642
  "id": "main",
643
  "name": "ReduceLogSumExp.Axis0TileCols",
644
  "shader": "reduce-axis0-tilecols.wgsl.jinja",
645
- "derive": { "op": "\"logsumexp\"", "castF32": "dtypes.T == \"f16\"", "usesF16Spec": "dtypes.T == \"f16\"" },
646
- "bindings": ["x_2", "y", "params_12"],
647
  "dispatch": {
648
  "x": "min(ceilDiv((dim(shapes.x, 1)), (tileCols)), 65535)",
649
  "y": "ceilDiv(ceilDiv((dim(shapes.x, 1)), (tileCols)), 65535)",
@@ -657,7 +603,6 @@
657
  "priority": 31,
658
  "when": ["flatParallelCovered"],
659
  "derive": {
660
- "scalar": "dtypes.T",
661
  "workgroupSize": "reduceWorkgroupSize",
662
  "flatScalar": "\"vec4<\" ~ dtypes.T ~ \">\" if numel(shapes.x) % tunables.VECTOR_WIDTH == 0 else dtypes.T",
663
  "split": "flatSplitCount"
@@ -668,14 +613,10 @@
668
  "id": "flat_partial",
669
  "name": "ReduceLogSumExp.AllAxesFlatPartial",
670
  "shader": "reduce-flat-partial-logsumexp.wgsl.jinja",
671
- "derive": {
672
- "vec4": "numel(shapes.x) % tunables.VECTOR_WIDTH == 0",
673
- "castF32": "dtypes.T == \"f16\"",
674
- "usesF16Spec": "dtypes.T == \"f16\""
675
- },
676
  "bindings": [
677
  { "arg": "x", "elementType": "$flatScalar" },
678
- { "name": "partials", "buffer": "storage", "elementType": "f32" },
679
  { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "flatItems" }] }
680
  ],
681
  "dispatch": { "x": "flatSplitCount" }
@@ -706,7 +647,6 @@
706
  "name": "ReduceLogSumExp.RankNSingleAxisGeneric",
707
  "shader": "reduce-serial-axis.wgsl.jinja",
708
  "derive": {
709
- "op": "\"logsumexp\"",
710
  "indexing": "\"rankn\"",
711
  "rank": "ranks.x",
712
  "axisSpec": "reduceAxis",
@@ -714,11 +654,9 @@
714
  "outputShape": "shapes.y",
715
  "outputRank": "ranks.y",
716
  "keepDims": "attrs.keepdims != 0",
717
- "intMode": "dtypes.T == \"i32\"",
718
- "castF32": "dtypes.T == \"f16\"",
719
- "usesF16Spec": "dtypes.T == \"f16\""
720
  },
721
- "bindings": ["x_2", "y", "params_16"],
722
  "dispatch": {
723
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
724
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -742,19 +680,13 @@
742
  "id": "main",
743
  "name": "ReduceLogSumExp.SubgroupRowVec4",
744
  "shader": "reduce-row-subgroup.wgsl.jinja",
745
- "derive": {
746
- "op": "\"logsumexp\"",
747
- "vec4": true,
748
- "castF32": "dtypes.T == \"f16\"",
749
- "usesF16Spec": "dtypes.T == \"f16\""
750
- },
751
- "bindings": ["x", "y", "params_6"],
752
  "dispatch": {
753
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
754
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
755
  "z": 1
756
- },
757
- "subgroupCollectivesWidth": "portable"
758
  }
759
  ]
760
  },
@@ -772,19 +704,13 @@
772
  "id": "main",
773
  "name": "ReduceLogSumExp.SubgroupRow",
774
  "shader": "reduce-row-subgroup.wgsl.jinja",
775
- "derive": {
776
- "op": "\"logsumexp\"",
777
- "vec4": false,
778
- "castF32": "dtypes.T == \"f16\"",
779
- "usesF16Spec": "dtypes.T == \"f16\""
780
- },
781
- "bindings": ["x_2", "y", "params_17"],
782
  "dispatch": {
783
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
784
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
785
  "z": 1
786
- },
787
- "subgroupCollectivesWidth": "portable"
788
  }
789
  ]
790
  },
@@ -797,17 +723,10 @@
797
  "passes": [
798
  {
799
  "id": "main",
800
- "name": "axis0",
801
  "shader": "reduce-serial-axis.wgsl.jinja",
802
- "derive": {
803
- "axis": 0,
804
- "op": "\"logsumexp\"",
805
- "indexing": "\"axis2d\"",
806
- "intMode": "dtypes.T == \"i32\"",
807
- "castF32": "dtypes.T == \"f16\"",
808
- "usesF16Spec": "dtypes.T == \"f16\""
809
- },
810
- "bindings": ["x_2", "y", "params_18"],
811
  "dispatch": {
812
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
813
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -824,17 +743,10 @@
824
  "passes": [
825
  {
826
  "id": "main",
827
- "name": "axis1",
828
  "shader": "reduce-serial-axis.wgsl.jinja",
829
- "derive": {
830
- "axis": 1,
831
- "op": "\"logsumexp\"",
832
- "indexing": "\"axis2d\"",
833
- "intMode": "dtypes.T == \"i32\"",
834
- "castF32": "dtypes.T == \"f16\"",
835
- "usesF16Spec": "dtypes.T == \"f16\""
836
- },
837
- "bindings": ["x_2", "y", "params_19"],
838
  "dispatch": {
839
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
840
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -853,25 +765,8 @@
853
  "id": "main",
854
  "name": "ReduceLogSumExp.Rank3AllAxesKeepdims",
855
  "shader": "reduce-serial-axis.wgsl.jinja",
856
- "derive": {
857
- "op": "\"logsumexp\"",
858
- "indexing": "\"axis2d\"",
859
- "intMode": "dtypes.T == \"i32\"",
860
- "castF32": "dtypes.T == \"f16\"",
861
- "usesF16Spec": "dtypes.T == \"f16\""
862
- },
863
- "bindings": [
864
- "x_2",
865
- "y",
866
- {
867
- "name": "params",
868
- "struct": [
869
- { "name": "rows", "type": "u32", "value": "numel(shapes.x)" },
870
- { "name": "cols", "type": "u32", "value": "1" },
871
- { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
872
- ]
873
- }
874
- ],
875
  "dispatch": {
876
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
877
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
@@ -890,25 +785,8 @@
890
  "id": "main",
891
  "name": "ReduceLogSumExp.Rank3AllAxesNoKeepdims",
892
  "shader": "reduce-serial-axis.wgsl.jinja",
893
- "derive": {
894
- "op": "\"logsumexp\"",
895
- "indexing": "\"axis2d\"",
896
- "intMode": "dtypes.T == \"i32\"",
897
- "castF32": "dtypes.T == \"f16\"",
898
- "usesF16Spec": "dtypes.T == \"f16\""
899
- },
900
- "bindings": [
901
- "x_2",
902
- "y",
903
- {
904
- "name": "params",
905
- "struct": [
906
- { "name": "rows", "type": "u32", "value": "numel(shapes.x)" },
907
- { "name": "cols", "type": "u32", "value": "1" },
908
- { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
909
- ]
910
- }
911
- ],
912
  "dispatch": {
913
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
914
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
34
  "ROW_SERIAL_MAX_COLS": { "default": 1024 }
35
  },
36
  "derive": {
37
+ "reduceOp": "\"logsumexp\"",
38
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
39
+ "op": "reduceOp",
40
+ "castF32": "dtypes.T == \"f16\"",
41
  "reduceWorkgroupSize": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)",
42
  "treeWorkgroupOk": "reduceWorkgroupSize > 0 and pow2ceil(reduceWorkgroupSize) == reduceWorkgroupSize and reduceWorkgroupSize * dtypeBytes(\"float32\") <= device.limits.maxComputeWorkgroupStorageSize",
43
  "subgroupWorkgroupFloor": "min(reduceWorkgroupSize, max(1, device.adapterInfo.subgroupMaxSize))",
 
65
  "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)))"
66
  },
67
  "bindings": {
68
+ "params_suffix_vec4": {
69
+ "name": "params",
 
 
70
  "struct": [
71
  { "name": "rows", "type": "u32", "value": "numel(shapes.y)" },
72
  { "name": "chunkCount", "type": "u32", "value": "numel(shapes.x) / numel(shapes.y) / tunables.VECTOR_WIDTH" }
73
  ]
74
  },
75
+ "params_reduce_row_tree": {
 
76
  "name": "params",
 
77
  "struct": [
78
+ { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
79
+ { "name": "chunkCount", "type": "u32", "value": "dim(shapes.x, ranks.x - 1) / tunables.VECTOR_WIDTH" }
80
  ]
81
  },
82
+ "params_axis_split_combine": {
83
  "name": "params",
84
+ "struct": [{ "name": "cols", "type": "u32", "value": "axisSplitOutputs" }]
 
85
  },
86
+ "params_axis0_combine": {
87
  "name": "params",
88
+ "struct": [{ "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" }]
 
89
  },
90
+ "params_rows_chunk_count": {
91
  "name": "params",
 
92
  "struct": [
93
  { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
94
+ { "name": "chunkCount", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" }
95
+ ]
96
+ },
97
+ "x": { "elementType": "$vectorScalar" },
98
+ "y": { "elementType": "$T" },
99
+ "x_t": { "name": "x", "elementType": "$T" },
100
+ "params_main": {
101
+ "name": "params",
102
+ "struct": [
103
+ { "name": "rows", "type": "u32", "value": "numel(shapes.y)" },
104
+ { "name": "cols", "type": "u32", "value": "numel(shapes.x) / numel(shapes.y)" }
105
  ]
106
  },
107
+ "params_out_count": {
108
+ "name": "params",
109
+ "struct": [{ "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }]
110
+ },
111
+ "params_count": { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.y)" }] },
112
+ "params_rank0_scalar": {
113
  "name": "params",
 
114
  "struct": [
115
  { "name": "rows", "type": "u32", "value": "1" },
116
  { "name": "cols", "type": "u32", "value": "1" },
117
  { "name": "outCount", "type": "u32", "value": "1" }
118
  ]
119
  },
120
+ "params_rank1_axis0": {
121
  "name": "params",
 
122
  "struct": [
123
  { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
124
  { "name": "cols", "type": "u32", "value": "1" },
125
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
126
  ]
127
  },
128
+ "params_rows_cols": {
129
  "name": "params",
 
130
  "struct": [
131
  { "name": "rows", "type": "u32", "value": "rows(shapes.x, ranks.x - 1)" },
132
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, ranks.x - 1)" }
133
  ]
134
  },
135
  "partials": { "buffer": "storage", "elementType": "$partialElement" },
136
+ "params_axis_split": {
137
  "name": "params",
 
138
  "struct": [
139
  { "name": "axisDim", "type": "u32", "value": "axisSplitDim" },
140
  { "name": "inner", "type": "u32", "value": "axisSplitInner" },
141
  { "name": "outputs", "type": "u32", "value": "axisSplitOutputs" }
142
  ]
143
  },
144
+ "partials_combine": { "name": "partials", "buffer": "read-only-storage", "elementType": "$partialElement" },
145
+ "params_split_reduce": {
 
 
 
 
 
146
  "name": "params",
 
147
  "struct": [
148
  { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
149
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" }
150
  ]
151
  },
152
+ "params_axis_dim_out_count": {
 
 
 
 
 
153
  "name": "params",
 
154
  "struct": [
155
  { "name": "axisDim", "type": "u32", "value": "axisSplitDim" },
156
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
157
  ]
158
  },
159
+ "params_axis0": {
160
  "name": "params",
 
161
  "struct": [
162
+ { "name": "rows", "type": "u32", "value": "dim(shapes.x, 0)" },
163
+ { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" },
164
+ { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
165
  ]
166
  },
167
+ "params_axis1": {
168
  "name": "params",
 
169
  "struct": [
 
170
  { "name": "cols", "type": "u32", "value": "dim(shapes.x, 1)" },
171
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
172
  ]
173
  },
174
+ "params_all_axes": {
175
  "name": "params",
 
176
  "struct": [
177
+ { "name": "rows", "type": "u32", "value": "numel(shapes.x)" },
178
+ { "name": "cols", "type": "u32", "value": "1" },
179
  { "name": "outCount", "type": "u32", "value": "numel(shapes.y)" }
180
  ]
181
  }
 
196
  "id": "main",
197
  "name": "ReduceLogSumExp.ContiguousSuffixSubgroupVec4",
198
  "shader": "reduce-row-subgroup.wgsl.jinja",
199
+ "derive": { "vec4": true },
200
+ "bindings": ["x", "y", "params_suffix_vec4"],
201
+ "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
 
 
 
 
 
 
202
  }
203
  ]
204
  },
 
216
  "id": "main",
217
  "name": "ReduceLogSumExp.ContiguousSuffixTreeVec4",
218
  "shader": "reduce-row-tree.wgsl.jinja",
219
+ "derive": { "vec4": true },
220
+ "bindings": ["x", "y", "params_suffix_vec4"],
 
 
 
 
 
221
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
222
  }
223
  ]
 
235
  "id": "main",
236
  "name": "ReduceLogSumExp.ContiguousSuffixTree",
237
  "shader": "reduce-row-tree.wgsl.jinja",
238
+ "bindings": ["x_t", "y", "params_main"],
 
239
  "dispatch": { "x": "min(numel(shapes.y), 65535)", "y": "ceilDiv(numel(shapes.y), 65535)", "z": 1 }
240
  }
241
  ]
 
251
  "name": "ReduceLogSumExp.MultiAxisRank3",
252
  "shader": "reduce-serial-axis.wgsl.jinja",
253
  "derive": {
 
254
  "indexing": "\"multiaxis\"",
255
  "rank": 3,
256
  "reduce": ["hasAxis(attrs.axes, 0, 3)", "hasAxis(attrs.axes, 1, 3)", "hasAxis(attrs.axes, 2, 3)"],
 
258
  "outputShape": "shapes.y",
259
  "outputRank": "ranks.y",
260
  "keepDims": "attrs.keepdims != 0",
261
+ "intMode": "dtypes.T == \"i32\""
 
 
262
  },
263
+ "bindings": ["x_t", "y", "params_out_count"],
264
  "dispatch": {
265
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
266
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
280
  "name": "ReduceLogSumExp.MultiAxisRank4",
281
  "shader": "reduce-serial-axis.wgsl.jinja",
282
  "derive": {
 
283
  "indexing": "\"multiaxis\"",
284
  "rank": 4,
285
  "reduce": ["hasAxis(attrs.axes, 0, 4)", "hasAxis(attrs.axes, 1, 4)", "hasAxis(attrs.axes, 2, 4)", "hasAxis(attrs.axes, 3, 4)"],
 
287
  "outputShape": "shapes.y",
288
  "outputRank": "ranks.y",
289
  "keepDims": "attrs.keepdims != 0",
290
+ "intMode": "dtypes.T == \"i32\""
 
 
291
  },
292
+ "bindings": ["x_t", "y", "params_out_count"],
293
  "dispatch": {
294
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
295
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
312
  "shader": "reduce-i32-axes02.wgsl.jinja",
313
  "derive": { "op": "\"logsumexp\"", "workgroupSizeSpec": "axes02WorkgroupSize" },
314
  "bindings": [
315
+ "x_t",
316
  "y",
317
  {
318
  "name": "params",
 
343
  "name": "ReduceLogSumExp.NoopEmptyAxes",
344
  "shader": "reduce-noop-empty-axes.wgsl.jinja",
345
  "derive": { "op": "\"identity\"" },
346
+ "bindings": ["x_t", "y", "params_count"],
347
  "dispatch": {
348
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
349
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
368
  "id": "main",
369
  "name": "ReduceLogSumExp.SubgroupRowsVec4",
370
  "shader": "reduce-row-subgroup-rows.wgsl.jinja",
371
+ "bindings": ["x", "y", "params_reduce_row_tree"],
 
372
  "dispatch": {
373
  "x": "min(ceilDiv(lastAxisRows, reduceWorkgroupSize / device.adapterInfo.subgroupMaxSize), 65535)",
374
  "y": "ceilDiv(ceilDiv(lastAxisRows, reduceWorkgroupSize / device.adapterInfo.subgroupMaxSize), 65535)",
375
  "z": 1
376
+ }
 
377
  }
378
  ]
379
  },
 
392
  "id": "main",
393
  "name": "ReduceLogSumExp.TreeRowVec4",
394
  "shader": "reduce-row-tree.wgsl.jinja",
395
+ "derive": { "vec4": true },
396
+ "bindings": ["x", "y", "params_reduce_row_tree"],
 
 
 
 
 
397
  "dispatch": { "x": "min(lastAxisRows, 65535)", "y": "ceilDiv(lastAxisRows, 65535)", "z": 1 }
398
  }
399
  ]
 
409
  "name": "ReduceLogSumExp.Rank0Scalar",
410
  "shader": "reduce-serial-axis.wgsl.jinja",
411
  "derive": {
 
412
  "indexing": "\"axis2d\"",
413
  "intMode": "dtypes.T == \"i32\"",
 
 
414
  "logicalBool": "tensorDtypes.x == \"bool\""
415
  },
416
+ "bindings": ["x_t", "y", "params_rank0_scalar"],
417
  "dispatch": { "x": 1 }
418
  }
419
  ]
 
428
  "name": "ReduceLogSumExp.Rank1Axis0",
429
  "shader": "reduce-serial-axis.wgsl.jinja",
430
  "derive": {
 
431
  "indexing": "\"axis2d\"",
432
  "intMode": "dtypes.T == \"i32\"",
 
 
433
  "logicalBool": "tensorDtypes.x == \"bool\""
434
  },
435
+ "bindings": ["x_t", "y", "params_rank1_axis0"],
436
  "dispatch": {
437
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
438
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
452
  "id": "main",
453
  "name": "ReduceLogSumExp.Axis1Parallel",
454
  "shader": "reduce-row-tree.wgsl.jinja",
455
+ "bindings": ["x_t", "y", "params_rows_cols"],
 
456
  "dispatch": {
457
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
458
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
 
477
  "id": "split_reduce",
478
  "name": "ReduceLogSumExp.AxisSplitReduce",
479
  "shader": "reduce-axis-split-reduce.wgsl.jinja",
480
+ "derive": { "splitSpec": "splitCount" },
481
+ "bindings": ["x_t", "partials", "params_axis_split"],
 
 
 
 
 
482
  "dispatch": {
483
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), DISPATCH_FOLD_WIDTH)",
484
  "y": "splitCount",
 
489
  "id": "combine",
490
  "name": "ReduceLogSumExp.AxisSplitCombine",
491
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
492
+ "derive": { "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
493
+ "bindings": ["partials_combine", "y", "params_axis_split_combine"],
494
  "dispatch": {
495
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
496
  "y": "ceilDiv(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
 
517
  "id": "split_reduce",
518
  "name": "ReduceLogSumExp.AxisSplitTiledReduce",
519
  "shader": "reduce-axis0-tilecols.wgsl.jinja",
520
+ "derive": { "splitSpec": "splitCount" },
521
+ "bindings": ["x_t", "partials", "params_axis_split"],
 
 
 
 
 
 
522
  "dispatch": {
523
  "x": "min(ceilDiv((axisSplitOutputs), (tileCols)), DISPATCH_FOLD_WIDTH)",
524
  "y": "splitCount",
 
529
  "id": "combine",
530
  "name": "ReduceLogSumExp.AxisSplitCombine",
531
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
532
+ "derive": { "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
533
+ "bindings": ["partials_combine", "y", "params_axis_split_combine"],
534
  "dispatch": {
535
  "x": "min(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
536
  "y": "ceilDiv(ceilDiv((axisSplitOutputs), (reduceWorkgroupSize)), 65535)",
 
543
  "id": "axis0_splitk",
544
  "priority": 22,
545
  "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"],
546
+ "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"],
547
  "derive": {
548
  "splitCount": "axis0SplitCount",
549
  "partialElement": "\"f32\"",
 
556
  "id": "split_reduce",
557
  "name": "ReduceLogSumExp.Axis0SplitKReduce",
558
  "shader": "reduce-axis0-splitk-reduce.wgsl.jinja",
559
+ "derive": { "splitSpec": "splitCount" },
560
+ "bindings": ["x_t", "partials", "params_split_reduce"],
 
 
 
 
 
561
  "dispatch": {
562
  "x": "min(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), DISPATCH_FOLD_WIDTH)",
563
  "y": "splitCount",
 
568
  "id": "combine",
569
  "name": "ReduceLogSumExp.Axis0SplitKCombine",
570
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
571
+ "derive": { "splitSpec": "splitCount", "outputF16": "dtypes.T == \"f16\"" },
572
+ "bindings": ["partials_combine", "y", "params_axis0_combine"],
573
  "dispatch": {
574
  "x": "min(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), 65535)",
575
  "y": "ceilDiv(ceilDiv((dim(shapes.x, 1)), (reduceWorkgroupSize)), 65535)",
 
588
  "id": "main",
589
  "name": "ReduceLogSumExp.Axis0TileCols",
590
  "shader": "reduce-axis0-tilecols.wgsl.jinja",
591
+ "derive": {},
592
+ "bindings": ["x_t", "y", "params_split_reduce"],
593
  "dispatch": {
594
  "x": "min(ceilDiv((dim(shapes.x, 1)), (tileCols)), 65535)",
595
  "y": "ceilDiv(ceilDiv((dim(shapes.x, 1)), (tileCols)), 65535)",
 
603
  "priority": 31,
604
  "when": ["flatParallelCovered"],
605
  "derive": {
 
606
  "workgroupSize": "reduceWorkgroupSize",
607
  "flatScalar": "\"vec4<\" ~ dtypes.T ~ \">\" if numel(shapes.x) % tunables.VECTOR_WIDTH == 0 else dtypes.T",
608
  "split": "flatSplitCount"
 
613
  "id": "flat_partial",
614
  "name": "ReduceLogSumExp.AllAxesFlatPartial",
615
  "shader": "reduce-flat-partial-logsumexp.wgsl.jinja",
616
+ "derive": { "vec4": "numel(shapes.x) % tunables.VECTOR_WIDTH == 0" },
 
 
 
 
617
  "bindings": [
618
  { "arg": "x", "elementType": "$flatScalar" },
619
+ { "name": "partials", "elementType": "f32" },
620
  { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "flatItems" }] }
621
  ],
622
  "dispatch": { "x": "flatSplitCount" }
 
647
  "name": "ReduceLogSumExp.RankNSingleAxisGeneric",
648
  "shader": "reduce-serial-axis.wgsl.jinja",
649
  "derive": {
 
650
  "indexing": "\"rankn\"",
651
  "rank": "ranks.x",
652
  "axisSpec": "reduceAxis",
 
654
  "outputShape": "shapes.y",
655
  "outputRank": "ranks.y",
656
  "keepDims": "attrs.keepdims != 0",
657
+ "intMode": "dtypes.T == \"i32\""
 
 
658
  },
659
+ "bindings": ["x_t", "y", "params_axis_dim_out_count"],
660
  "dispatch": {
661
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
662
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
680
  "id": "main",
681
  "name": "ReduceLogSumExp.SubgroupRowVec4",
682
  "shader": "reduce-row-subgroup.wgsl.jinja",
683
+ "derive": { "vec4": true },
684
+ "bindings": ["x", "y", "params_reduce_row_tree"],
 
 
 
 
 
685
  "dispatch": {
686
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
687
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
688
  "z": 1
689
+ }
 
690
  }
691
  ]
692
  },
 
704
  "id": "main",
705
  "name": "ReduceLogSumExp.SubgroupRow",
706
  "shader": "reduce-row-subgroup.wgsl.jinja",
707
+ "derive": { "vec4": false },
708
+ "bindings": ["x_t", "y", "params_rows_chunk_count"],
 
 
 
 
 
709
  "dispatch": {
710
  "x": "min(rows(shapes.x, ranks.x - 1), 65535)",
711
  "y": "ceilDiv(rows(shapes.x, ranks.x - 1), 65535)",
712
  "z": 1
713
+ }
 
714
  }
715
  ]
716
  },
 
723
  "passes": [
724
  {
725
  "id": "main",
726
+ "name": "ReduceLogSumExp.Axis0",
727
  "shader": "reduce-serial-axis.wgsl.jinja",
728
+ "derive": { "indexing": "\"axis2d\"", "intMode": "dtypes.T == \"i32\"" },
729
+ "bindings": ["x_t", "y", "params_axis0"],
 
 
 
 
 
 
 
730
  "dispatch": {
731
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
732
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
743
  "passes": [
744
  {
745
  "id": "main",
746
+ "name": "ReduceLogSumExp.Axis1",
747
  "shader": "reduce-serial-axis.wgsl.jinja",
748
+ "derive": { "indexing": "\"axis2d\"", "intMode": "dtypes.T == \"i32\"" },
749
+ "bindings": ["x_t", "y", "params_axis1"],
 
 
 
 
 
 
 
750
  "dispatch": {
751
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
752
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
765
  "id": "main",
766
  "name": "ReduceLogSumExp.Rank3AllAxesKeepdims",
767
  "shader": "reduce-serial-axis.wgsl.jinja",
768
+ "derive": { "indexing": "\"axis2d\"", "intMode": "dtypes.T == \"i32\"" },
769
+ "bindings": ["x_t", "y", "params_all_axes"],
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
770
  "dispatch": {
771
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
772
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
 
785
  "id": "main",
786
  "name": "ReduceLogSumExp.Rank3AllAxesNoKeepdims",
787
  "shader": "reduce-serial-axis.wgsl.jinja",
788
+ "derive": { "indexing": "\"axis2d\"", "intMode": "dtypes.T == \"i32\"" },
789
+ "bindings": ["x_t", "y", "params_all_axes"],
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
790
  "dispatch": {
791
  "x": "min(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
792
  "y": "ceilDiv(ceilDiv((numel(shapes.y)), (reduceWorkgroupSize)), 65535)",
build/webgpu/metadata.json CHANGED
@@ -1,32 +1,32 @@
1
  {
2
  "name": "ai.onnx.ReduceLogSumExp",
3
- "id": "_ai_onnx_reducelogsumexp_webgpu_fd6e0c7",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "msizeqfB88GjW8v7QU9KoKyjW3qc8lJJXaCPf6OejeI=",
11
- "manifest.json": "ofeVPiTS0KOqgOnbSIdNK8ULQiV3GSHu3Sxay4FY9U0=",
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": "PjYkEUQJBeG70td3W2xxmexH9x6XNIfeSzXBY47HbaU=",
16
- "reduce-flat-combine-logsumexp.wgsl.jinja": "vwwisgGJGC6hBX/WYMf/OEXKa71pMJKX6v2ggWuXWpM=",
17
- "reduce-flat-partial-logsumexp.wgsl.jinja": "5ZVJnWsR0xWm4Oet3oqP0V7ePeqLea5wBs/rBsy9L20=",
18
  "reduce-i32-axes02.wgsl.jinja": "FsDSFjExTqueGACUuk9bPO7gN5yp8sJHtEXapuxter0=",
19
- "reduce-noop-empty-axes.wgsl.jinja": "NNvXRO0Tt3Mrvk2Wdssndbea61oQN4tj+O08NVA43Ns=",
20
- "reduce-row-subgroup-rows.wgsl.jinja": "76u7rAvFoZZKrFDs2A2jk0vkB0uNrPL6twdOBE9b+v8=",
21
- "reduce-row-subgroup.wgsl.jinja": "2mu9LEsk8HfaLvucBCfcB1/ENpXkD6ELiCtRt+5UqiU=",
22
- "reduce-row-tree.wgsl.jinja": "Bwa5xcI0bTmKXb4r9Cc1bfVbM5rNqqpQVrWWVqcb8xA=",
23
- "reduce-serial-axis.wgsl.jinja": "fvUV9htqzKzt4Pg05pYtRmup/5QIHYaUGQHJZnthXKo=",
24
- "test.json": "QyxQrLhFNbxc+mf7VOuv8RipIY2AlNZEhJZJvTQjuNw="
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"],
 
1
  {
2
  "name": "ai.onnx.ReduceLogSumExp",
3
+ "id": "_ai_onnx_reducelogsumexp_webgpu_4aeeb2f",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "bench.json": "wXvsZOuaZHscg7oIb5NujBZvnlo8B9U6YLfa69+npYg=",
11
+ "manifest.json": "QcxF0HuWIGT+TZaAjJgOcvYBfwAKV98lh+fie3k41dM=",
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": "dyX/c20cWbPRfgqDdOi5XqfhvFtxyQD8KDJJ2gJ3fqc=",
16
+ "reduce-flat-combine-logsumexp.wgsl.jinja": "YxvwTOwRzXW5NTZ4tEtrt6syg32z4ELS2qW2/KKIHp4=",
17
+ "reduce-flat-partial-logsumexp.wgsl.jinja": "6OTGSyRzItrZtgKYi4OC+Ay3j8c/AoCu3W0tU/mKI5A=",
18
  "reduce-i32-axes02.wgsl.jinja": "FsDSFjExTqueGACUuk9bPO7gN5yp8sJHtEXapuxter0=",
19
+ "reduce-noop-empty-axes.wgsl.jinja": "I3eV+jyJXA4ojj9U6RPZnivKP4CD6R4HFmcM0i7v6vA=",
20
+ "reduce-row-subgroup-rows.wgsl.jinja": "5cMGE1dmrcawwHdoxQD8kPSurwA1wNHyoAQfxAOE7vU=",
21
+ "reduce-row-subgroup.wgsl.jinja": "7bRBWIQ7oHOq/E54hGrmXIzLznzVH9oTfYaf1jmBzV8=",
22
+ "reduce-row-tree.wgsl.jinja": "a+iEV/cBkrlsVz6NS0GY/TAsUSX5rO4N14wY1dKfDi0=",
23
+ "reduce-serial-axis.wgsl.jinja": "XZthzOtNawNeG2SXq528l1yO5sVPqTTK3Y9GVhJ46lk=",
24
+ "test.json": "IfasxXu3otrLuI7M2dRKeSFj/wZJcDf5OvKM7a9luNo="
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"],
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-combine-logsumexp.wgsl.jinja CHANGED
@@ -6,9 +6,6 @@
6
  // transcendental merges when the output has exactly one element.
7
  {% set yv = "f16(" if outputF16 else "" %}
8
  {% set vy = ")" if outputF16 else "" %}
9
- {% if outputF16 %}
10
- enable f16;
11
- {% endif %}
12
  {{ env.wgsl.resourceDeclarations }}
13
 
14
  const WG: u32 = {{ workgroupSize }}u;
@@ -25,7 +22,6 @@ fn is_nan_f32(value: f32) -> bool {
25
  let bits = bitcast<u32>(value);
26
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
27
  }
28
-
29
  @compute @workgroup_size(WG)
30
  fn main(@builtin(local_invocation_id) lid: vec3<u32>) {
31
  let tid = lid.x;
@@ -69,15 +65,23 @@ fn main(@builtin(local_invocation_id) lid: vec3<u32>) {
69
  }
70
  workgroupBarrier();
71
 
72
- stride = WG / 2u;
 
 
 
 
 
73
  loop {
74
- if (stride == 0u) { break; }
75
- if (tid < stride) {
76
- wsum[tid] = wsum[tid] + wsum[tid + stride];
 
 
77
  }
78
- stride = stride / 2u;
79
  workgroupBarrier();
80
- }
 
81
 
82
  if (tid == 0u) {
83
  let has_positive_inf = global_max > F32_MAX;
 
6
  // transcendental merges when the output has exactly one element.
7
  {% set yv = "f16(" if outputF16 else "" %}
8
  {% set vy = ")" if outputF16 else "" %}
 
 
 
9
  {{ env.wgsl.resourceDeclarations }}
10
 
11
  const WG: u32 = {{ workgroupSize }}u;
 
22
  let bits = bitcast<u32>(value);
23
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
24
  }
 
25
  @compute @workgroup_size(WG)
26
  fn main(@builtin(local_invocation_id) lid: vec3<u32>) {
27
  let tid = lid.x;
 
65
  }
66
  workgroupBarrier();
67
 
68
+ {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
69
+ {% if op == "max" or op == "min" %}
70
+ {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
71
+ {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %}
72
+ {% 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) %}
73
+ {{ svar }} = {{ wg }} / 2u;
74
  loop {
75
+ if ({{ svar }} == 0u) { break; }
76
+ if ({{ idx }} < {{ svar }}) {
77
+ {% for a in arrays %}
78
+ {{ wgsl_tree_fold_stmt(a, op, idx, svar) }}
79
+ {% endfor %}
80
  }
81
+ {{ svar }} = {{ svar }} / 2u;
82
  workgroupBarrier();
83
+ }{% endmacro %}
84
+ {{ wgsl_tree_fold(["wsum"], idx="tid", wg="WG", form="head", breakInline=true, reuse=true) }}
85
 
86
  if (tid == 0u) {
87
  let has_positive_inf = global_max > F32_MAX;
build/webgpu/reduce-flat-partial-logsumexp.wgsl.jinja CHANGED
@@ -4,13 +4,8 @@
4
  // Workgroups write three partial planes: segment maximum, shifted exponential
5
  // sum, and a NaN marker. The combine pass merges them stably and emits
6
  // globalMaximum + log(sum), preserving NaN and positive infinity.
7
- {% set vec4 = vec4 | default(true) %}
8
- {% set castF32 = castF32 is defined and castF32 %}
9
  {% set xa = "f32(" if castF32 else "" %}
10
  {% set ax = ")" if castF32 else "" %}
11
- {% if usesF16Spec is defined and usesF16Spec %}
12
- enable f16;
13
- {% endif %}
14
  {{ env.wgsl.resourceDeclarations }}
15
 
16
  const WG: u32 = {{ workgroupSize }}u;
@@ -25,7 +20,6 @@ fn is_nan_f32(value: f32) -> bool {
25
  let bits = bitcast<u32>(value);
26
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
27
  }
28
-
29
  @compute @workgroup_size(WG)
30
  fn main(@builtin(global_invocation_id) gid: vec3<u32>,
31
  @builtin(local_invocation_id) lid: vec3<u32>,
@@ -85,13 +79,19 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>,
85
  }
86
  wsum[tid] = local_sum;
87
  workgroupBarrier();
88
- stride = WG / 2u;
 
 
 
 
 
89
  loop {
90
- if (stride == 0u) { break; }
91
- if (tid < stride) { wsum[tid] = wsum[tid] + wsum[tid + stride]; }
92
- stride = stride / 2u;
93
  workgroupBarrier();
94
- }
 
95
 
96
  if (tid == 0u) {
97
  partials[wg.x] = seg_max;
 
4
  // Workgroups write three partial planes: segment maximum, shifted exponential
5
  // sum, and a NaN marker. The combine pass merges them stably and emits
6
  // globalMaximum + log(sum), preserving NaN and positive infinity.
 
 
7
  {% set xa = "f32(" if castF32 else "" %}
8
  {% set ax = ")" if castF32 else "" %}
 
 
 
9
  {{ env.wgsl.resourceDeclarations }}
10
 
11
  const WG: u32 = {{ workgroupSize }}u;
 
20
  let bits = bitcast<u32>(value);
21
  return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
22
  }
 
23
  @compute @workgroup_size(WG)
24
  fn main(@builtin(global_invocation_id) gid: vec3<u32>,
25
  @builtin(local_invocation_id) lid: vec3<u32>,
 
79
  }
80
  wsum[tid] = local_sum;
81
  workgroupBarrier();
82
+ {% macro wgsl_tree_fold_stmt(a, op, idx, svar) %}
83
+ {% if op == "max" or op == "min" %}
84
+ {{ a }}[{{ idx }}] = {{ op }}({{ a }}[{{ idx }}], {{ a }}[{{ idx }} + {{ svar }}]);{% else %}
85
+ {{ a }}[{{ idx }}] = {{ a }}[{{ idx }}] {{ "*" if op == "prod" else "+" }} {{ a }}[{{ idx }} + {{ svar }}];{% endif %}{% endmacro %}
86
+ {% 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) %}
87
+ {{ svar }} = {{ wg }} / 2u;
88
  loop {
89
+ if ({{ svar }} == 0u) { break; }
90
+ if ({{ idx }} < {{ svar }}) { {{ wgsl_tree_fold_stmt(arrays[0], op, idx, svar) }} }
91
+ {{ svar }} = {{ svar }} / 2u;
92
  workgroupBarrier();
93
+ }{% endmacro %}
94
+ {{ wgsl_tree_fold(["wsum"], idx="tid", wg="WG", form="head", breakInline=true, bodyInline=true, reuse=true) }}
95
 
96
  if (tid == 0u) {
97
  partials[wg.x] = seg_max;
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
  y[i] = v;
13
  }
 
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
  y[i] = v;
16
  }
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/test.json CHANGED
@@ -665,6 +665,19 @@
665
  },
666
  "outputs": { "y": { "dtype": "float32", "shape": [2, 2], "tolerance": 0.0001 } }
667
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
668
  {
669
  "name": "axis0_splitk_8192x32",
670
  "attrs": { "axes": [0], "keepdims": 0 },
@@ -989,6 +1002,19 @@
989
  "outputs": {
990
  "y": { "dtype": "float32", "shape": [64], "tolerance": 0.0001, "allowNaN": true, "relTolerance": 0.0001 }
991
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
992
  }
993
  ]
994
  }
 
665
  },
666
  "outputs": { "y": { "dtype": "float32", "shape": [2, 2], "tolerance": 0.0001 } }
667
  },
668
+ {
669
+ "name": "axis0_variable_subgroup_tile_selection_1024x64",
670
+ "attrs": { "axes": [0], "keepdims": 0 },
671
+ "tunables": { "AXIS0_SPLIT_MIN_ROWS": 512, "AXIS0_SPLIT_TARGET_ROWS": 256 },
672
+ "inputs": {
673
+ "x": {
674
+ "dtype": "float32",
675
+ "shape": [1024, 64],
676
+ "data": { "kind": "fillFloat32", "sinStep": 0.17, "cosStep": 0.07, "scale": 0.2 }
677
+ }
678
+ },
679
+ "outputs": { "y": { "dtype": "float32", "shape": [64], "tolerance": 0.0001 } }
680
+ },
681
  {
682
  "name": "axis0_splitk_8192x32",
683
  "attrs": { "axes": [0], "keepdims": 0 },
 
1002
  "outputs": {
1003
  "y": { "dtype": "float32", "shape": [64], "tolerance": 0.0001, "allowNaN": true, "relTolerance": 0.0001 }
1004
  }
1005
+ },
1006
+ {
1007
+ "name": "ort_empty_set_noop_with_empty_axes_identity",
1008
+ "provenance": {
1009
+ "source": "onnxruntime/test/providers/cpu/reduction/reduction_ops_test.cc",
1010
+ "test": "ReductionOpTest.EmptySetNoopWithEmptyAxes (ReduceLogSumExp)",
1011
+ "notes": "Empty-set input combined with noop_with_empty_axes=1: the output must keep the input shape, not a -inf-filled reduced shape."
1012
+ },
1013
+ "attrs": { "keepdims": 1, "noop_with_empty_axes": 1 },
1014
+ "inputs": { "x": { "dtype": "float32", "shape": [2, 0, 4], "data": { "kind": "values", "values": [] } } },
1015
+ "outputs": {
1016
+ "y": { "dtype": "float32", "shape": [2, 0, 4], "tolerance": 0, "data": { "kind": "values", "values": [] } }
1017
+ }
1018
  }
1019
  ]
1020
  }