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