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