@graphty/webgpu-graph-algorithms 0.5.1 → 0.6.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +98 -52
- package/dist/browser.js +1 -1
- package/dist/chunks/{context-BR7fx3vR.js → context-BXqgCifx.js} +190 -40
- package/dist/chunks/context-BXqgCifx.js.map +1 -0
- package/dist/node.js +1 -1
- package/dist/src/algorithms/components.d.ts.map +1 -1
- package/dist/src/algorithms/components.js +12 -13
- package/dist/src/algorithms/components.js.map +1 -1
- package/dist/src/algorithms/degree.d.ts +6 -8
- package/dist/src/algorithms/degree.d.ts.map +1 -1
- package/dist/src/algorithms/degree.js +58 -35
- package/dist/src/algorithms/degree.js.map +1 -1
- package/dist/src/algorithms/pagerank.d.ts.map +1 -1
- package/dist/src/algorithms/pagerank.js +16 -14
- package/dist/src/algorithms/pagerank.js.map +1 -1
- package/dist/src/algorithms/power-iteration.d.ts +2 -2
- package/dist/src/algorithms/power-iteration.d.ts.map +1 -1
- package/dist/src/algorithms/power-iteration.js +17 -14
- package/dist/src/algorithms/power-iteration.js.map +1 -1
- package/dist/src/constants.d.ts +38 -8
- package/dist/src/constants.d.ts.map +1 -1
- package/dist/src/constants.js +38 -8
- package/dist/src/constants.js.map +1 -1
- package/dist/src/errors.d.ts +3 -2
- package/dist/src/errors.d.ts.map +1 -1
- package/dist/src/errors.js +2 -1
- package/dist/src/errors.js.map +1 -1
- package/dist/src/index.d.ts +6 -4
- package/dist/src/index.d.ts.map +1 -1
- package/dist/src/index.js +8 -3
- package/dist/src/index.js.map +1 -1
- package/dist/src/kernel/dispatch.d.ts +8 -3
- package/dist/src/kernel/dispatch.d.ts.map +1 -1
- package/dist/src/kernel/dispatch.js +18 -7
- package/dist/src/kernel/dispatch.js.map +1 -1
- package/dist/src/kernel/kernel.d.ts +30 -1
- package/dist/src/kernel/kernel.d.ts.map +1 -1
- package/dist/src/kernel/kernel.js +49 -5
- package/dist/src/kernel/kernel.js.map +1 -1
- package/dist/src/kernel/prelude.d.ts.map +1 -1
- package/dist/src/kernel/prelude.js +6 -1
- package/dist/src/kernel/prelude.js.map +1 -1
- package/dist/src/kernel/profiler.d.ts +15 -3
- package/dist/src/kernel/profiler.d.ts.map +1 -1
- package/dist/src/kernel/profiler.js +27 -4
- package/dist/src/kernel/profiler.js.map +1 -1
- package/dist/src/kernels.d.ts +17 -7
- package/dist/src/kernels.d.ts.map +1 -1
- package/dist/src/kernels.js +323 -16
- package/dist/src/kernels.js.map +1 -1
- package/dist/src/layouts/calibrate.d.ts +51 -0
- package/dist/src/layouts/calibrate.d.ts.map +1 -0
- package/dist/src/layouts/calibrate.js +172 -0
- package/dist/src/layouts/calibrate.js.map +1 -0
- package/dist/src/layouts/force-simulation.d.ts +39 -4
- package/dist/src/layouts/force-simulation.d.ts.map +1 -1
- package/dist/src/layouts/force-simulation.js +71 -19
- package/dist/src/layouts/force-simulation.js.map +1 -1
- package/dist/src/layouts/forceatlas2.d.ts +107 -36
- package/dist/src/layouts/forceatlas2.d.ts.map +1 -1
- package/dist/src/layouts/forceatlas2.js +296 -100
- package/dist/src/layouts/forceatlas2.js.map +1 -1
- package/dist/src/layouts/fruchterman-reingold.d.ts +73 -27
- package/dist/src/layouts/fruchterman-reingold.d.ts.map +1 -1
- package/dist/src/layouts/fruchterman-reingold.js +230 -70
- package/dist/src/layouts/fruchterman-reingold.js.map +1 -1
- package/dist/src/layouts/model-common.d.ts +41 -3
- package/dist/src/layouts/model-common.d.ts.map +1 -1
- package/dist/src/layouts/model-common.js +74 -3
- package/dist/src/layouts/model-common.js.map +1 -1
- package/dist/src/layouts/repulsion-grid.d.ts +152 -0
- package/dist/src/layouts/repulsion-grid.d.ts.map +1 -0
- package/dist/src/layouts/repulsion-grid.js +318 -0
- package/dist/src/layouts/repulsion-grid.js.map +1 -0
- package/dist/src/layouts/spring-electrical.d.ts +75 -30
- package/dist/src/layouts/spring-electrical.d.ts.map +1 -1
- package/dist/src/layouts/spring-electrical.js +231 -74
- package/dist/src/layouts/spring-electrical.js.map +1 -1
- package/dist/src/memory/residency.d.ts +6 -2
- package/dist/src/memory/residency.d.ts.map +1 -1
- package/dist/src/memory/residency.js +84 -14
- package/dist/src/memory/residency.js.map +1 -1
- package/dist/src/primitives/core-shape.d.ts +38 -2
- package/dist/src/primitives/core-shape.d.ts.map +1 -1
- package/dist/src/primitives/core-shape.js +71 -3
- package/dist/src/primitives/core-shape.js.map +1 -1
- package/dist/src/primitives/grid-pyramid.d.ts +71 -0
- package/dist/src/primitives/grid-pyramid.d.ts.map +1 -0
- package/dist/src/primitives/grid-pyramid.js +143 -0
- package/dist/src/primitives/grid-pyramid.js.map +1 -0
- package/dist/src/primitives/grid.d.ts +118 -0
- package/dist/src/primitives/grid.d.ts.map +1 -0
- package/dist/src/primitives/grid.js +225 -0
- package/dist/src/primitives/grid.js.map +1 -0
- package/dist/src/primitives/histogram.d.ts +67 -0
- package/dist/src/primitives/histogram.d.ts.map +1 -0
- package/dist/src/primitives/histogram.js +190 -0
- package/dist/src/primitives/histogram.js.map +1 -0
- package/dist/src/primitives/radix-sort.d.ts +75 -0
- package/dist/src/primitives/radix-sort.d.ts.map +1 -0
- package/dist/src/primitives/radix-sort.js +168 -0
- package/dist/src/primitives/radix-sort.js.map +1 -0
- package/dist/src/primitives/scan.d.ts +44 -0
- package/dist/src/primitives/scan.d.ts.map +1 -0
- package/dist/src/primitives/scan.js +151 -0
- package/dist/src/primitives/scan.js.map +1 -0
- package/dist/src/primitives/segmented-reduce.d.ts +25 -17
- package/dist/src/primitives/segmented-reduce.d.ts.map +1 -1
- package/dist/src/primitives/segmented-reduce.js +166 -47
- package/dist/src/primitives/segmented-reduce.js.map +1 -1
- package/dist/src/primitives/spmv.d.ts +18 -14
- package/dist/src/primitives/spmv.d.ts.map +1 -1
- package/dist/src/primitives/spmv.js +94 -58
- package/dist/src/primitives/spmv.js.map +1 -1
- package/dist/src/primitives/verify.d.ts +49 -0
- package/dist/src/primitives/verify.d.ts.map +1 -0
- package/dist/src/primitives/verify.js +229 -0
- package/dist/src/primitives/verify.js.map +1 -0
- package/dist/src/types/context.d.ts +53 -0
- package/dist/src/types/context.d.ts.map +1 -1
- package/dist/src/types/layout.d.ts +20 -0
- package/dist/src/types/layout.d.ts.map +1 -1
- package/dist/src/wgsl/counting-scatter.wgsl.d.ts +8 -0
- package/dist/src/wgsl/counting-scatter.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/counting-scatter.wgsl.js +17 -0
- package/dist/src/wgsl/counting-scatter.wgsl.js.map +1 -0
- package/dist/src/wgsl/fa2-attraction.wgsl.d.ts +23 -11
- package/dist/src/wgsl/fa2-attraction.wgsl.d.ts.map +1 -1
- package/dist/src/wgsl/fa2-attraction.wgsl.js +98 -20
- package/dist/src/wgsl/fa2-attraction.wgsl.js.map +1 -1
- package/dist/src/wgsl/fa2-stats-finalize.wgsl.d.ts +6 -2
- package/dist/src/wgsl/fa2-stats-finalize.wgsl.d.ts.map +1 -1
- package/dist/src/wgsl/fa2-stats-finalize.wgsl.js +22 -1
- package/dist/src/wgsl/fa2-stats-finalize.wgsl.js.map +1 -1
- package/dist/src/wgsl/grid-cell-key.wgsl.d.ts +8 -0
- package/dist/src/wgsl/grid-cell-key.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/grid-cell-key.wgsl.js +30 -0
- package/dist/src/wgsl/grid-cell-key.wgsl.js.map +1 -0
- package/dist/src/wgsl/grid-centroid-hub.wgsl.d.ts +8 -0
- package/dist/src/wgsl/grid-centroid-hub.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/grid-centroid-hub.wgsl.js +29 -0
- package/dist/src/wgsl/grid-centroid-hub.wgsl.js.map +1 -0
- package/dist/src/wgsl/grid-centroid.wgsl.d.ts +8 -0
- package/dist/src/wgsl/grid-centroid.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/grid-centroid.wgsl.js +29 -0
- package/dist/src/wgsl/grid-centroid.wgsl.js.map +1 -0
- package/dist/src/wgsl/grid-downsample.wgsl.d.ts +7 -0
- package/dist/src/wgsl/grid-downsample.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/grid-downsample.wgsl.js +28 -0
- package/dist/src/wgsl/grid-downsample.wgsl.js.map +1 -0
- package/dist/src/wgsl/grid-far-field.wgsl.d.ts +13 -0
- package/dist/src/wgsl/grid-far-field.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/grid-far-field.wgsl.js +98 -0
- package/dist/src/wgsl/grid-far-field.wgsl.js.map +1 -0
- package/dist/src/wgsl/grid-near-field.wgsl.d.ts +19 -0
- package/dist/src/wgsl/grid-near-field.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/grid-near-field.wgsl.js +129 -0
- package/dist/src/wgsl/grid-near-field.wgsl.js.map +1 -0
- package/dist/src/wgsl/histogram.wgsl.d.ts +7 -0
- package/dist/src/wgsl/histogram.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/histogram.wgsl.js +15 -0
- package/dist/src/wgsl/histogram.wgsl.js.map +1 -0
- package/dist/src/wgsl/indirect-finalize.wgsl.d.ts +8 -0
- package/dist/src/wgsl/indirect-finalize.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/indirect-finalize.wgsl.js +26 -0
- package/dist/src/wgsl/indirect-finalize.wgsl.js.map +1 -0
- package/dist/src/wgsl/radix-hist.wgsl.d.ts +9 -0
- package/dist/src/wgsl/radix-hist.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/radix-hist.wgsl.js +31 -0
- package/dist/src/wgsl/radix-hist.wgsl.js.map +1 -0
- package/dist/src/wgsl/radix-scatter.wgsl.d.ts +9 -0
- package/dist/src/wgsl/radix-scatter.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/radix-scatter.wgsl.js +40 -0
- package/dist/src/wgsl/radix-scatter.wgsl.js.map +1 -0
- package/dist/src/wgsl/scan-add.wgsl.d.ts +6 -0
- package/dist/src/wgsl/scan-add.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/scan-add.wgsl.js +14 -0
- package/dist/src/wgsl/scan-add.wgsl.js.map +1 -0
- package/dist/src/wgsl/scan-block.wgsl.d.ts +8 -0
- package/dist/src/wgsl/scan-block.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/scan-block.wgsl.js +30 -0
- package/dist/src/wgsl/scan-block.wgsl.js.map +1 -0
- package/dist/src/wgsl/segmented-reduce.wgsl.d.ts +22 -8
- package/dist/src/wgsl/segmented-reduce.wgsl.d.ts.map +1 -1
- package/dist/src/wgsl/segmented-reduce.wgsl.js +84 -15
- package/dist/src/wgsl/segmented-reduce.wgsl.js.map +1 -1
- package/dist/src/wgsl/spmv-pull.wgsl.d.ts +22 -11
- package/dist/src/wgsl/spmv-pull.wgsl.d.ts.map +1 -1
- package/dist/src/wgsl/spmv-pull.wgsl.js +110 -36
- package/dist/src/wgsl/spmv-pull.wgsl.js.map +1 -1
- package/dist/tsconfig.build.tsbuildinfo +1 -1
- package/dist/webgpu-graph-algorithms.js +3815 -1003
- package/dist/webgpu-graph-algorithms.js.map +1 -1
- package/package.json +5 -4
- package/src/algorithms/components.ts +12 -16
- package/src/algorithms/degree.ts +58 -43
- package/src/algorithms/pagerank.ts +20 -18
- package/src/algorithms/power-iteration.ts +19 -18
- package/src/constants.ts +38 -8
- package/src/errors.ts +3 -1
- package/src/index.ts +14 -4
- package/src/kernel/dispatch.ts +18 -7
- package/src/kernel/kernel.ts +59 -5
- package/src/kernel/prelude.ts +9 -0
- package/src/kernel/profiler.ts +28 -4
- package/src/kernels.ts +356 -18
- package/src/layouts/calibrate.ts +187 -0
- package/src/layouts/force-simulation.ts +91 -23
- package/src/layouts/forceatlas2.ts +331 -106
- package/src/layouts/fruchterman-reingold.ts +255 -74
- package/src/layouts/model-common.ts +98 -3
- package/src/layouts/repulsion-grid.ts +451 -0
- package/src/layouts/spring-electrical.ts +257 -78
- package/src/memory/residency.ts +126 -20
- package/src/primitives/core-shape.ts +91 -4
- package/src/primitives/grid-pyramid.ts +221 -0
- package/src/primitives/grid.ts +349 -0
- package/src/primitives/histogram.ts +273 -0
- package/src/primitives/radix-sort.ts +246 -0
- package/src/primitives/scan.ts +197 -0
- package/src/primitives/segmented-reduce.ts +214 -56
- package/src/primitives/spmv.ts +125 -65
- package/src/primitives/verify.ts +249 -0
- package/src/types/context.ts +56 -0
- package/src/types/layout.ts +22 -0
- package/src/wgsl/counting-scatter.wgsl.ts +16 -0
- package/src/wgsl/fa2-attraction.wgsl.ts +98 -20
- package/src/wgsl/fa2-stats-finalize.wgsl.ts +22 -1
- package/src/wgsl/grid-cell-key.wgsl.ts +29 -0
- package/src/wgsl/grid-centroid-hub.wgsl.ts +28 -0
- package/src/wgsl/grid-centroid.wgsl.ts +28 -0
- package/src/wgsl/grid-downsample.wgsl.ts +27 -0
- package/src/wgsl/grid-far-field.wgsl.ts +97 -0
- package/src/wgsl/grid-near-field.wgsl.ts +128 -0
- package/src/wgsl/histogram.wgsl.ts +14 -0
- package/src/wgsl/indirect-finalize.wgsl.ts +25 -0
- package/src/wgsl/radix-hist.wgsl.ts +30 -0
- package/src/wgsl/radix-scatter.wgsl.ts +39 -0
- package/src/wgsl/scan-add.wgsl.ts +13 -0
- package/src/wgsl/scan-block.wgsl.ts +29 -0
- package/src/wgsl/segmented-reduce.wgsl.ts +84 -15
- package/src/wgsl/spmv-pull.wgsl.ts +110 -36
- package/dist/chunks/context-BR7fx3vR.js.map +0 -1
|
@@ -0,0 +1,190 @@
|
|
|
1
|
+
/**
|
|
2
|
+
* The `histogram` and `countingSortByKey` primitive drivers (spec 6 row 5; P4-T3, PD-4). `histogram` zeroes `hist`
|
|
3
|
+
* with a `fill` dispatch and then adds one per key in global memory: order-independent, so bitwise deterministic on
|
|
4
|
+
* every adapter. `countingSortByKey` is the histogram, an exclusive scan of it into `outStart`, a zeroed per-bin
|
|
5
|
+
* `cursor`, and the scatter `outIndex[outStart[k] + atomicAdd(&cursor[k], 1)] = i`: the SET of indices inside a bin
|
|
6
|
+
* is fixed, their order follows the schedule (set-deterministic, design 6). Every zeroing is a dispatch inside the
|
|
7
|
+
* caller's pass (PD-12), never an encoder clear.
|
|
8
|
+
*
|
|
9
|
+
* The drivers own no device objects: the caller supplies a ReduceScope (the same record `reduce` takes) and the
|
|
10
|
+
* compute pass to record into. `src/primitives/**` never imports `src/context.ts`.
|
|
11
|
+
*/
|
|
12
|
+
import { WebGpuGraphError } from "../errors.js";
|
|
13
|
+
import { plan1d } from "../kernel/dispatch.js";
|
|
14
|
+
import { FILL_PARAMS, HIST_PARAMS, kernelSpec } from "../kernels.js";
|
|
15
|
+
import { prepareScan } from "./scan.js";
|
|
16
|
+
/** The largest u32: the largest value `HistParams.count` / `bins` can carry. */
|
|
17
|
+
const U32_MAX = 0xffffffff;
|
|
18
|
+
/**
|
|
19
|
+
* Prepares the histogram pipelines of a scope (compiles `histogram` and `fill` once) so record() is synchronous.
|
|
20
|
+
* @param scope - the caller's scope (device, caps, cache, scratch, params)
|
|
21
|
+
* @returns the planner
|
|
22
|
+
*/
|
|
23
|
+
export async function prepareHistogram(scope) {
|
|
24
|
+
return prepareHistogramImpl(scope);
|
|
25
|
+
}
|
|
26
|
+
/**
|
|
27
|
+
* The concrete histogram planner (its recordZero is what the counting sort reuses for the cursor).
|
|
28
|
+
* @param scope - the caller's scope
|
|
29
|
+
* @returns the planner
|
|
30
|
+
*/
|
|
31
|
+
async function prepareHistogramImpl(scope) {
|
|
32
|
+
const histogram = await scope.pipelines.kernel(kernelSpec("histogram"));
|
|
33
|
+
const fill = await scope.pipelines.kernel(kernelSpec("fill"));
|
|
34
|
+
return new HistogramPlannerImpl(scope, histogram, fill);
|
|
35
|
+
}
|
|
36
|
+
/**
|
|
37
|
+
* Prepares the counting-sort pipelines of a scope (the histogram's, the scan's and `counting-scatter`) so record()
|
|
38
|
+
* is synchronous. The planner lives exactly as long as the scope (the scan keeps a scratch word of it).
|
|
39
|
+
* @param scope - the caller's scope
|
|
40
|
+
* @returns the planner
|
|
41
|
+
*/
|
|
42
|
+
export async function prepareCountingSort(scope) {
|
|
43
|
+
const histogram = await prepareHistogramImpl(scope);
|
|
44
|
+
const scan = await prepareScan(scope);
|
|
45
|
+
const scatter = await scope.pipelines.kernel(kernelSpec("counting-scatter"));
|
|
46
|
+
return new CountingSortPlannerImpl(scope, histogram, scan, scatter);
|
|
47
|
+
}
|
|
48
|
+
/**
|
|
49
|
+
* The E_INVALID_ARGUMENT of a binding shorter than `words` u32.
|
|
50
|
+
* @param name - the argument name
|
|
51
|
+
* @param binding - the binding
|
|
52
|
+
* @param words - the words it must hold
|
|
53
|
+
*/
|
|
54
|
+
function checkWords(name, binding, words) {
|
|
55
|
+
if (binding.size < 4 * words) {
|
|
56
|
+
throw new WebGpuGraphError("E_INVALID_ARGUMENT", `histogram: ${name} is smaller than 4 x ${words} bytes`, {
|
|
57
|
+
argument: name,
|
|
58
|
+
value: binding.size,
|
|
59
|
+
expected: 4 * words,
|
|
60
|
+
});
|
|
61
|
+
}
|
|
62
|
+
}
|
|
63
|
+
/**
|
|
64
|
+
* The argument checks shared by both record() methods (E_INVALID_ARGUMENT before anything is recorded).
|
|
65
|
+
* @param keys - the keys binding
|
|
66
|
+
* @param count - the key count
|
|
67
|
+
* @param bins - the bin count
|
|
68
|
+
* @param hist - the histogram binding
|
|
69
|
+
*/
|
|
70
|
+
function checkHistogramArguments(keys, count, bins, hist) {
|
|
71
|
+
if (!Number.isSafeInteger(count) || count < 0 || count > U32_MAX) {
|
|
72
|
+
throw new WebGpuGraphError("E_INVALID_ARGUMENT", "histogram: count must be a non-negative integer below 2^32", {
|
|
73
|
+
argument: "count",
|
|
74
|
+
value: count,
|
|
75
|
+
});
|
|
76
|
+
}
|
|
77
|
+
if (!Number.isSafeInteger(bins) || bins < 1 || bins > U32_MAX) {
|
|
78
|
+
throw new WebGpuGraphError("E_INVALID_ARGUMENT", "histogram: bins must be an integer in [1, 2^32)", {
|
|
79
|
+
argument: "bins",
|
|
80
|
+
value: bins,
|
|
81
|
+
});
|
|
82
|
+
}
|
|
83
|
+
checkWords("keys", keys, count);
|
|
84
|
+
checkWords("hist", hist, bins);
|
|
85
|
+
}
|
|
86
|
+
/** The planner: the histogram and fill kernels over one scope. */
|
|
87
|
+
class HistogramPlannerImpl {
|
|
88
|
+
/**
|
|
89
|
+
* Wraps the resolved kernels; use prepareHistogram().
|
|
90
|
+
* @param scope - the caller's scope
|
|
91
|
+
* @param histogram - the `histogram` kernel
|
|
92
|
+
* @param fill - the `fill` kernel
|
|
93
|
+
*/
|
|
94
|
+
constructor(scope, histogram, fill) {
|
|
95
|
+
this.dispatches = 0;
|
|
96
|
+
this.scope = scope;
|
|
97
|
+
this.histogram = histogram;
|
|
98
|
+
this.fill = fill;
|
|
99
|
+
}
|
|
100
|
+
/**
|
|
101
|
+
* Dispatches the last record() issued.
|
|
102
|
+
* @returns the count
|
|
103
|
+
*/
|
|
104
|
+
get lastDispatches() {
|
|
105
|
+
return this.dispatches;
|
|
106
|
+
}
|
|
107
|
+
/**
|
|
108
|
+
* Records the fill and the histogram (see the interface).
|
|
109
|
+
* @param pass - the compute pass
|
|
110
|
+
* @param keys - the keys
|
|
111
|
+
* @param count - the key count
|
|
112
|
+
* @param bins - the bin count
|
|
113
|
+
* @param hist - the counts
|
|
114
|
+
*/
|
|
115
|
+
record(pass, keys, count, bins, hist) {
|
|
116
|
+
checkHistogramArguments(keys, count, bins, hist);
|
|
117
|
+
this.recordZero(pass, hist, bins);
|
|
118
|
+
this.dispatches = 1;
|
|
119
|
+
if (count === 0) {
|
|
120
|
+
return;
|
|
121
|
+
}
|
|
122
|
+
const params = this.scope.params(HIST_PARAMS, { count, bins, pad0: 0, pad1: 0 });
|
|
123
|
+
const bound = this.histogram.bind({ keys, hist, P: params.binding });
|
|
124
|
+
this.histogram.dispatch(pass, bound, plan1d(count, this.scope.workgroupSize, this.scope.caps), [params.offset]);
|
|
125
|
+
this.dispatches = 2;
|
|
126
|
+
}
|
|
127
|
+
/**
|
|
128
|
+
* Records a `fill` of `words` zeros into `dst` (PD-12: a dispatch inside the pass, never an encoder clear).
|
|
129
|
+
* @param pass - the compute pass
|
|
130
|
+
* @param dst - the words to zero
|
|
131
|
+
* @param words - how many (>= 1)
|
|
132
|
+
*/
|
|
133
|
+
recordZero(pass, dst, words) {
|
|
134
|
+
const params = this.scope.params(FILL_PARAMS, { count: words, value: 0, mode: 0, pad0: 0 });
|
|
135
|
+
const bound = this.fill.bind({ dst, P: params.binding });
|
|
136
|
+
this.fill.dispatch(pass, bound, plan1d(words, this.scope.workgroupSize, this.scope.caps), [params.offset]);
|
|
137
|
+
}
|
|
138
|
+
}
|
|
139
|
+
/** The planner: the histogram planner, the scan planner and the scatter kernel over one scope. */
|
|
140
|
+
class CountingSortPlannerImpl {
|
|
141
|
+
/**
|
|
142
|
+
* Wraps the resolved planners and kernel; use prepareCountingSort().
|
|
143
|
+
* @param scope - the caller's scope
|
|
144
|
+
* @param histogram - the histogram planner (also the zeroing of the cursor)
|
|
145
|
+
* @param scan - the scan planner
|
|
146
|
+
* @param scatter - the `counting-scatter` kernel
|
|
147
|
+
*/
|
|
148
|
+
constructor(scope, histogram, scan, scatter) {
|
|
149
|
+
this.dispatches = 0;
|
|
150
|
+
this.scope = scope;
|
|
151
|
+
this.histogram = histogram;
|
|
152
|
+
this.scan = scan;
|
|
153
|
+
this.scatter = scatter;
|
|
154
|
+
}
|
|
155
|
+
/**
|
|
156
|
+
* Dispatches the last record() issued.
|
|
157
|
+
* @returns the count
|
|
158
|
+
*/
|
|
159
|
+
get lastDispatches() {
|
|
160
|
+
return this.dispatches;
|
|
161
|
+
}
|
|
162
|
+
/**
|
|
163
|
+
* Records the four stages (see the interface).
|
|
164
|
+
* @param pass - the compute pass
|
|
165
|
+
* @param keys - the keys
|
|
166
|
+
* @param count - the key count
|
|
167
|
+
* @param bins - the bin count
|
|
168
|
+
* @param scratch - the histogram and cursor
|
|
169
|
+
* @param outIndex - the sorted indices
|
|
170
|
+
* @param outStart - the bin starts
|
|
171
|
+
*/
|
|
172
|
+
record(pass, keys, count, bins, scratch, outIndex, outStart) {
|
|
173
|
+
checkHistogramArguments(keys, count, bins, scratch.hist);
|
|
174
|
+
checkWords("cursor", scratch.cursor, bins);
|
|
175
|
+
checkWords("outIndex", outIndex, count);
|
|
176
|
+
checkWords("outStart", outStart, bins);
|
|
177
|
+
this.histogram.record(pass, keys, count, bins, scratch.hist);
|
|
178
|
+
this.scan.record(pass, scratch.hist, bins, outStart);
|
|
179
|
+
this.histogram.recordZero(pass, scratch.cursor, bins);
|
|
180
|
+
this.dispatches = this.histogram.lastDispatches + this.scan.lastDispatches + 1;
|
|
181
|
+
if (count === 0) {
|
|
182
|
+
return;
|
|
183
|
+
}
|
|
184
|
+
const params = this.scope.params(HIST_PARAMS, { count, bins, pad0: 0, pad1: 0 });
|
|
185
|
+
const bound = this.scatter.bind({ keys, start: outStart, cursor: scratch.cursor, outIndex, P: params.binding });
|
|
186
|
+
this.scatter.dispatch(pass, bound, plan1d(count, this.scope.workgroupSize, this.scope.caps), [params.offset]);
|
|
187
|
+
this.dispatches += 1;
|
|
188
|
+
}
|
|
189
|
+
}
|
|
190
|
+
//# sourceMappingURL=histogram.js.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"histogram.js","sourceRoot":"","sources":["../../../src/primitives/histogram.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;GAUG;AAEH,OAAO,EAAE,gBAAgB,EAAE,MAAM,cAAc,CAAC;AAChD,OAAO,EAAE,MAAM,EAAE,MAAM,uBAAuB,CAAC;AAE/C,OAAO,EAAE,WAAW,EAAE,WAAW,EAAE,UAAU,EAAE,MAAM,eAAe,CAAC;AAGrE,OAAO,EAAE,WAAW,EAAoB,MAAM,WAAW,CAAC;AAoD1D,gFAAgF;AAChF,MAAM,OAAO,GAAG,UAAU,CAAC;AAE3B;;;;GAIG;AACH,MAAM,CAAC,KAAK,UAAU,gBAAgB,CAAC,KAAkB;IACrD,OAAO,oBAAoB,CAAC,KAAK,CAAC,CAAC;AACvC,CAAC;AAED;;;;GAIG;AACH,KAAK,UAAU,oBAAoB,CAAC,KAAkB;IAClD,MAAM,SAAS,GAAG,MAAM,KAAK,CAAC,SAAS,CAAC,MAAM,CAAC,UAAU,CAAC,WAAW,CAAC,CAAC,CAAC;IACxE,MAAM,IAAI,GAAG,MAAM,KAAK,CAAC,SAAS,CAAC,MAAM,CAAC,UAAU,CAAC,MAAM,CAAC,CAAC,CAAC;IAC9D,OAAO,IAAI,oBAAoB,CAAC,KAAK,EAAE,SAAS,EAAE,IAAI,CAAC,CAAC;AAC5D,CAAC;AAED;;;;;GAKG;AACH,MAAM,CAAC,KAAK,UAAU,mBAAmB,CAAC,KAAkB;IACxD,MAAM,SAAS,GAAG,MAAM,oBAAoB,CAAC,KAAK,CAAC,CAAC;IACpD,MAAM,IAAI,GAAG,MAAM,WAAW,CAAC,KAAK,CAAC,CAAC;IACtC,MAAM,OAAO,GAAG,MAAM,KAAK,CAAC,SAAS,CAAC,MAAM,CAAC,UAAU,CAAC,kBAAkB,CAAC,CAAC,CAAC;IAC7E,OAAO,IAAI,uBAAuB,CAAC,KAAK,EAAE,SAAS,EAAE,IAAI,EAAE,OAAO,CAAC,CAAC;AACxE,CAAC;AAED;;;;;GAKG;AACH,SAAS,UAAU,CAAC,IAAY,EAAE,OAAgB,EAAE,KAAa;IAC7D,IAAI,OAAO,CAAC,IAAI,GAAG,CAAC,GAAG,KAAK,EAAE,CAAC;QAC3B,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,cAAc,IAAI,wBAAwB,KAAK,QAAQ,EAAE;YACtG,QAAQ,EAAE,IAAI;YACd,KAAK,EAAE,OAAO,CAAC,IAAI;YACnB,QAAQ,EAAE,CAAC,GAAG,KAAK;SACtB,CAAC,CAAC;IACP,CAAC;AACL,CAAC;AAED;;;;;;GAMG;AACH,SAAS,uBAAuB,CAAC,IAAa,EAAE,KAAa,EAAE,IAAY,EAAE,IAAa;IACtF,IAAI,CAAC,MAAM,CAAC,aAAa,CAAC,KAAK,CAAC,IAAI,KAAK,GAAG,CAAC,IAAI,KAAK,GAAG,OAAO,EAAE,CAAC;QAC/D,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,4DAA4D,EAAE;YAC3G,QAAQ,EAAE,OAAO;YACjB,KAAK,EAAE,KAAK;SACf,CAAC,CAAC;IACP,CAAC;IACD,IAAI,CAAC,MAAM,CAAC,aAAa,CAAC,IAAI,CAAC,IAAI,IAAI,GAAG,CAAC,IAAI,IAAI,GAAG,OAAO,EAAE,CAAC;QAC5D,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,iDAAiD,EAAE;YAChG,QAAQ,EAAE,MAAM;YAChB,KAAK,EAAE,IAAI;SACd,CAAC,CAAC;IACP,CAAC;IACD,UAAU,CAAC,MAAM,EAAE,IAAI,EAAE,KAAK,CAAC,CAAC;IAChC,UAAU,CAAC,MAAM,EAAE,IAAI,EAAE,IAAI,CAAC,CAAC;AACnC,CAAC;AAED,kEAAkE;AAClE,MAAM,oBAAoB;IAMtB;;;;;OAKG;IACH,YAAY,KAAkB,EAAE,SAAiB,EAAE,IAAY;QARvD,eAAU,GAAG,CAAC,CAAC;QASnB,IAAI,CAAC,KAAK,GAAG,KAAK,CAAC;QACnB,IAAI,CAAC,SAAS,GAAG,SAAS,CAAC;QAC3B,IAAI,CAAC,IAAI,GAAG,IAAI,CAAC;IACrB,CAAC;IAED;;;OAGG;IACH,IAAI,cAAc;QACd,OAAO,IAAI,CAAC,UAAU,CAAC;IAC3B,CAAC;IAED;;;;;;;OAOG;IACH,MAAM,CAAC,IAA2B,EAAE,IAAa,EAAE,KAAa,EAAE,IAAY,EAAE,IAAa;QACzF,uBAAuB,CAAC,IAAI,EAAE,KAAK,EAAE,IAAI,EAAE,IAAI,CAAC,CAAC;QACjD,IAAI,CAAC,UAAU,CAAC,IAAI,EAAE,IAAI,EAAE,IAAI,CAAC,CAAC;QAClC,IAAI,CAAC,UAAU,GAAG,CAAC,CAAC;QACpB,IAAI,KAAK,KAAK,CAAC,EAAE,CAAC;YACd,OAAO;QACX,CAAC;QACD,MAAM,MAAM,GAAG,IAAI,CAAC,KAAK,CAAC,MAAM,CAAC,WAAW,EAAE,EAAE,KAAK,EAAE,IAAI,EAAE,IAAI,EAAE,CAAC,EAAE,IAAI,EAAE,CAAC,EAAE,CAAC,CAAC;QACjF,MAAM,KAAK,GAAG,IAAI,CAAC,SAAS,CAAC,IAAI,CAAC,EAAE,IAAI,EAAE,IAAI,EAAE,CAAC,EAAE,MAAM,CAAC,OAAO,EAAE,CAAC,CAAC;QACrE,IAAI,CAAC,SAAS,CAAC,QAAQ,CAAC,IAAI,EAAE,KAAK,EAAE,MAAM,CAAC,KAAK,EAAE,IAAI,CAAC,KAAK,CAAC,aAAa,EAAE,IAAI,CAAC,KAAK,CAAC,IAAI,CAAC,EAAE,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC;QAChH,IAAI,CAAC,UAAU,GAAG,CAAC,CAAC;IACxB,CAAC;IAED;;;;;OAKG;IACH,UAAU,CAAC,IAA2B,EAAE,GAAY,EAAE,KAAa;QAC/D,MAAM,MAAM,GAAG,IAAI,CAAC,KAAK,CAAC,MAAM,CAAC,WAAW,EAAE,EAAE,KAAK,EAAE,KAAK,EAAE,KAAK,EAAE,CAAC,EAAE,IAAI,EAAE,CAAC,EAAE,IAAI,EAAE,CAAC,EAAE,CAAC,CAAC;QAC5F,MAAM,KAAK,GAAG,IAAI,CAAC,IAAI,CAAC,IAAI,CAAC,EAAE,GAAG,EAAE,CAAC,EAAE,MAAM,CAAC,OAAO,EAAE,CAAC,CAAC;QACzD,IAAI,CAAC,IAAI,CAAC,QAAQ,CAAC,IAAI,EAAE,KAAK,EAAE,MAAM,CAAC,KAAK,EAAE,IAAI,CAAC,KAAK,CAAC,aAAa,EAAE,IAAI,CAAC,KAAK,CAAC,IAAI,CAAC,EAAE,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC;IAC/G,CAAC;CACJ;AAED,kGAAkG;AAClG,MAAM,uBAAuB;IAOzB;;;;;;OAMG;IACH,YAAY,KAAkB,EAAE,SAA+B,EAAE,IAAiB,EAAE,OAAe;QAT3F,eAAU,GAAG,CAAC,CAAC;QAUnB,IAAI,CAAC,KAAK,GAAG,KAAK,CAAC;QACnB,IAAI,CAAC,SAAS,GAAG,SAAS,CAAC;QAC3B,IAAI,CAAC,IAAI,GAAG,IAAI,CAAC;QACjB,IAAI,CAAC,OAAO,GAAG,OAAO,CAAC;IAC3B,CAAC;IAED;;;OAGG;IACH,IAAI,cAAc;QACd,OAAO,IAAI,CAAC,UAAU,CAAC;IAC3B,CAAC;IAED;;;;;;;;;OASG;IACH,MAAM,CACF,IAA2B,EAC3B,IAAa,EACb,KAAa,EACb,IAAY,EACZ,OAA4B,EAC5B,QAAiB,EACjB,QAAiB;QAEjB,uBAAuB,CAAC,IAAI,EAAE,KAAK,EAAE,IAAI,EAAE,OAAO,CAAC,IAAI,CAAC,CAAC;QACzD,UAAU,CAAC,QAAQ,EAAE,OAAO,CAAC,MAAM,EAAE,IAAI,CAAC,CAAC;QAC3C,UAAU,CAAC,UAAU,EAAE,QAAQ,EAAE,KAAK,CAAC,CAAC;QACxC,UAAU,CAAC,UAAU,EAAE,QAAQ,EAAE,IAAI,CAAC,CAAC;QACvC,IAAI,CAAC,SAAS,CAAC,MAAM,CAAC,IAAI,EAAE,IAAI,EAAE,KAAK,EAAE,IAAI,EAAE,OAAO,CAAC,IAAI,CAAC,CAAC;QAC7D,IAAI,CAAC,IAAI,CAAC,MAAM,CAAC,IAAI,EAAE,OAAO,CAAC,IAAI,EAAE,IAAI,EAAE,QAAQ,CAAC,CAAC;QACrD,IAAI,CAAC,SAAS,CAAC,UAAU,CAAC,IAAI,EAAE,OAAO,CAAC,MAAM,EAAE,IAAI,CAAC,CAAC;QACtD,IAAI,CAAC,UAAU,GAAG,IAAI,CAAC,SAAS,CAAC,cAAc,GAAG,IAAI,CAAC,IAAI,CAAC,cAAc,GAAG,CAAC,CAAC;QAC/E,IAAI,KAAK,KAAK,CAAC,EAAE,CAAC;YACd,OAAO;QACX,CAAC;QACD,MAAM,MAAM,GAAG,IAAI,CAAC,KAAK,CAAC,MAAM,CAAC,WAAW,EAAE,EAAE,KAAK,EAAE,IAAI,EAAE,IAAI,EAAE,CAAC,EAAE,IAAI,EAAE,CAAC,EAAE,CAAC,CAAC;QACjF,MAAM,KAAK,GAAG,IAAI,CAAC,OAAO,CAAC,IAAI,CAAC,EAAE,IAAI,EAAE,KAAK,EAAE,QAAQ,EAAE,MAAM,EAAE,OAAO,CAAC,MAAM,EAAE,QAAQ,EAAE,CAAC,EAAE,MAAM,CAAC,OAAO,EAAE,CAAC,CAAC;QAChH,IAAI,CAAC,OAAO,CAAC,QAAQ,CAAC,IAAI,EAAE,KAAK,EAAE,MAAM,CAAC,KAAK,EAAE,IAAI,CAAC,KAAK,CAAC,aAAa,EAAE,IAAI,CAAC,KAAK,CAAC,IAAI,CAAC,EAAE,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC;QAC9G,IAAI,CAAC,UAAU,IAAI,CAAC,CAAC;IACzB,CAAC;CACJ"}
|
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
/**
|
|
2
|
+
* The `radixSort` primitive driver (spec 6 row 6; P4-T4, PD-5): a STABLE LSD key-value sort of u32 words, 8 bits per
|
|
3
|
+
* pass, `bits / 8` passes. Every pass runs `radix-hist` (one 256-bin digit histogram per WG-wide block, stored
|
|
4
|
+
* DIGIT-MAJOR `hist[digit * groups + group]`), `exclusiveScan` over that table (so each word becomes, per digit, the
|
|
5
|
+
* offset of its workgroup in workgroup order) and `radix-scatter` (lane 0 ranks its block's keys serially in index
|
|
6
|
+
* order, every lane writes its pair at `offsets[digit * groups + group] + rank`). The pairs swap after every pass, so an
|
|
7
|
+
* odd pass count (bits 8 / 24) leaves the result in the caller's scratch pair and an even one (16 / 32) in the input
|
|
8
|
+
* pair; record() RETURNS the pair so no caller guesses. u32 arithmetic and a serial ranking make two runs on any two
|
|
9
|
+
* adapters bitwise identical.
|
|
10
|
+
*
|
|
11
|
+
* The scan cannot run in place: `Kernel.bind` (and WebGPU's usage-scope rule) rejects one buffer bound `storage-ro`
|
|
12
|
+
* and `storage` in one dispatch, so the caller supplies TWO tables of `radixHistBytes`: `scratch.hist` (the raw
|
|
13
|
+
* digit-major table `radix-hist` writes; the LAST pass's table stays there) and `scratch.offsets` (its exclusive
|
|
14
|
+
* scan, what `radix-scatter` reads). Both are the caller's so a per-iteration sort (the grid, PD-11) binds the same
|
|
15
|
+
* buffers every iteration and `Kernel.bind`'s identity cache holds.
|
|
16
|
+
*
|
|
17
|
+
* The driver owns no device objects and acquires no scratch of its own: the caller supplies a ReduceScope (the
|
|
18
|
+
* record `reduce` and `exclusiveScan` take; the scan's block sums come from it) and the compute pass to record
|
|
19
|
+
* into. `src/primitives/**` never imports `src/context.ts`.
|
|
20
|
+
*/
|
|
21
|
+
import { type Binding } from "../types/memory.js";
|
|
22
|
+
import { type ReduceScope } from "./reduce.js";
|
|
23
|
+
/** The key widths record() accepts: `bits / 8` passes, so the low `bits` of every key order the pairs. */
|
|
24
|
+
export type RadixBits = 8 | 16 | 24 | 32;
|
|
25
|
+
/**
|
|
26
|
+
* The caller's scratch: a second key-value pair the passes ping-pong with, the digit-major histogram table and its
|
|
27
|
+
* scanned twin (each `radixHistBytes` bytes at least).
|
|
28
|
+
*/
|
|
29
|
+
export interface RadixSortScratch {
|
|
30
|
+
readonly keys: Binding;
|
|
31
|
+
readonly vals: Binding;
|
|
32
|
+
/** The raw digit-major table `radix-hist` writes (the last pass's stays readable after the sort). */
|
|
33
|
+
readonly hist: Binding;
|
|
34
|
+
/** The exclusive scan of `hist`, what `radix-scatter` reads. */
|
|
35
|
+
readonly offsets: Binding;
|
|
36
|
+
}
|
|
37
|
+
/** Where the sorted pairs landed: the input pair or the scratch pair (record() decides by the pass count). */
|
|
38
|
+
export interface RadixSortResult {
|
|
39
|
+
readonly keys: Binding;
|
|
40
|
+
readonly vals: Binding;
|
|
41
|
+
}
|
|
42
|
+
/** A prepared radix sort (spec 6 row 6): records the pass dispatches of one sort into a pass. */
|
|
43
|
+
export interface RadixSortPlanner {
|
|
44
|
+
/**
|
|
45
|
+
* Records the stable sort of `count` (key, value) pairs by the low `bits` of the key; returns the pair holding the
|
|
46
|
+
* result (the scratch pair after an odd pass count, the input pair after an even one); count 0 records nothing and
|
|
47
|
+
* returns the input pair.
|
|
48
|
+
* @param pass - the compute pass
|
|
49
|
+
* @param keys - the keys (at least 4 x count bytes)
|
|
50
|
+
* @param vals - the values (at least 4 x count bytes)
|
|
51
|
+
* @param count - the pair count (a non-negative integer below 2^32)
|
|
52
|
+
* @param bits - the key width: 8, 16, 24 or 32
|
|
53
|
+
* @param scratch - the second pair, the histogram table and its scanned twin
|
|
54
|
+
* @returns the pair the result lives in
|
|
55
|
+
*/
|
|
56
|
+
record(pass: GPUComputePassEncoder, keys: Binding, vals: Binding, count: number, bits: RadixBits, scratch: RadixSortScratch): RadixSortResult;
|
|
57
|
+
/** Dispatches the last record() issued: 0 for count 0, else `passes x (2 + the scan's dispatches over the table)`. */
|
|
58
|
+
readonly lastDispatches: number;
|
|
59
|
+
}
|
|
60
|
+
/**
|
|
61
|
+
* The byte size of the digit-major histogram table of one pass over `count` keys at workgroup size `wg`:
|
|
62
|
+
* `4 x 256 x ceil(count / wg)` (the caller sizes `scratch.hist` by it; 0 for count 0, which records nothing).
|
|
63
|
+
* @param count - the pair count
|
|
64
|
+
* @param wg - the workgroup size (`scope.workgroupSize`)
|
|
65
|
+
* @returns the bytes
|
|
66
|
+
*/
|
|
67
|
+
export declare function radixHistBytes(count: number, wg: number): number;
|
|
68
|
+
/**
|
|
69
|
+
* Prepares the two radix pipelines and the scan of a scope (compiles once) so record() is synchronous. The planner
|
|
70
|
+
* lives exactly as long as the scope: never use a planner after its scope's dispose().
|
|
71
|
+
* @param scope - the caller's scope (device, caps, cache, scratch, params)
|
|
72
|
+
* @returns the planner
|
|
73
|
+
*/
|
|
74
|
+
export declare function prepareRadixSort(scope: ReduceScope): Promise<RadixSortPlanner>;
|
|
75
|
+
//# sourceMappingURL=radix-sort.d.ts.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"radix-sort.d.ts","sourceRoot":"","sources":["../../../src/primitives/radix-sort.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;;;;;;;;;;GAmBG;AAOH,OAAO,EAAE,KAAK,OAAO,EAAE,MAAM,oBAAoB,CAAC;AAClD,OAAO,EAAE,KAAK,WAAW,EAAE,MAAM,aAAa,CAAC;AAM/C,0GAA0G;AAC1G,MAAM,MAAM,SAAS,GAAG,CAAC,GAAG,EAAE,GAAG,EAAE,GAAG,EAAE,CAAC;AAEzC;;;GAGG;AACH,MAAM,WAAW,gBAAgB;IAC7B,QAAQ,CAAC,IAAI,EAAE,OAAO,CAAC;IACvB,QAAQ,CAAC,IAAI,EAAE,OAAO,CAAC;IACvB,qGAAqG;IACrG,QAAQ,CAAC,IAAI,EAAE,OAAO,CAAC;IACvB,gEAAgE;IAChE,QAAQ,CAAC,OAAO,EAAE,OAAO,CAAC;CAC7B;AAED,8GAA8G;AAC9G,MAAM,WAAW,eAAe;IAC5B,QAAQ,CAAC,IAAI,EAAE,OAAO,CAAC;IACvB,QAAQ,CAAC,IAAI,EAAE,OAAO,CAAC;CAC1B;AAED,iGAAiG;AACjG,MAAM,WAAW,gBAAgB;IAC7B;;;;;;;;;;;OAWG;IACH,MAAM,CACF,IAAI,EAAE,qBAAqB,EAC3B,IAAI,EAAE,OAAO,EACb,IAAI,EAAE,OAAO,EACb,KAAK,EAAE,MAAM,EACb,IAAI,EAAE,SAAS,EACf,OAAO,EAAE,gBAAgB,GAC1B,eAAe,CAAC;IACnB,sHAAsH;IACtH,QAAQ,CAAC,cAAc,EAAE,MAAM,CAAC;CACnC;AAED;;;;;;GAMG;AACH,wBAAgB,cAAc,CAAC,KAAK,EAAE,MAAM,EAAE,EAAE,EAAE,MAAM,GAAG,MAAM,CAEhE;AAED;;;;;GAKG;AACH,wBAAsB,gBAAgB,CAAC,KAAK,EAAE,WAAW,GAAG,OAAO,CAAC,gBAAgB,CAAC,CAKpF"}
|
|
@@ -0,0 +1,168 @@
|
|
|
1
|
+
/**
|
|
2
|
+
* The `radixSort` primitive driver (spec 6 row 6; P4-T4, PD-5): a STABLE LSD key-value sort of u32 words, 8 bits per
|
|
3
|
+
* pass, `bits / 8` passes. Every pass runs `radix-hist` (one 256-bin digit histogram per WG-wide block, stored
|
|
4
|
+
* DIGIT-MAJOR `hist[digit * groups + group]`), `exclusiveScan` over that table (so each word becomes, per digit, the
|
|
5
|
+
* offset of its workgroup in workgroup order) and `radix-scatter` (lane 0 ranks its block's keys serially in index
|
|
6
|
+
* order, every lane writes its pair at `offsets[digit * groups + group] + rank`). The pairs swap after every pass, so an
|
|
7
|
+
* odd pass count (bits 8 / 24) leaves the result in the caller's scratch pair and an even one (16 / 32) in the input
|
|
8
|
+
* pair; record() RETURNS the pair so no caller guesses. u32 arithmetic and a serial ranking make two runs on any two
|
|
9
|
+
* adapters bitwise identical.
|
|
10
|
+
*
|
|
11
|
+
* The scan cannot run in place: `Kernel.bind` (and WebGPU's usage-scope rule) rejects one buffer bound `storage-ro`
|
|
12
|
+
* and `storage` in one dispatch, so the caller supplies TWO tables of `radixHistBytes`: `scratch.hist` (the raw
|
|
13
|
+
* digit-major table `radix-hist` writes; the LAST pass's table stays there) and `scratch.offsets` (its exclusive
|
|
14
|
+
* scan, what `radix-scatter` reads). Both are the caller's so a per-iteration sort (the grid, PD-11) binds the same
|
|
15
|
+
* buffers every iteration and `Kernel.bind`'s identity cache holds.
|
|
16
|
+
*
|
|
17
|
+
* The driver owns no device objects and acquires no scratch of its own: the caller supplies a ReduceScope (the
|
|
18
|
+
* record `reduce` and `exclusiveScan` take; the scan's block sums come from it) and the compute pass to record
|
|
19
|
+
* into. `src/primitives/**` never imports `src/context.ts`.
|
|
20
|
+
*/
|
|
21
|
+
import { RADIX_BINS, U32_MAX } from "../constants.js";
|
|
22
|
+
import { WebGpuGraphError } from "../errors.js";
|
|
23
|
+
import { plan1d } from "../kernel/dispatch.js";
|
|
24
|
+
import { kernelSpec, RADIX_PARAMS } from "../kernels.js";
|
|
25
|
+
import { prepareScan } from "./scan.js";
|
|
26
|
+
/** The bits of one pass. */
|
|
27
|
+
const RADIX_DIGIT_BITS = 8;
|
|
28
|
+
/**
|
|
29
|
+
* The byte size of the digit-major histogram table of one pass over `count` keys at workgroup size `wg`:
|
|
30
|
+
* `4 x 256 x ceil(count / wg)` (the caller sizes `scratch.hist` by it; 0 for count 0, which records nothing).
|
|
31
|
+
* @param count - the pair count
|
|
32
|
+
* @param wg - the workgroup size (`scope.workgroupSize`)
|
|
33
|
+
* @returns the bytes
|
|
34
|
+
*/
|
|
35
|
+
export function radixHistBytes(count, wg) {
|
|
36
|
+
return 4 * RADIX_BINS * Math.ceil(count / wg);
|
|
37
|
+
}
|
|
38
|
+
/**
|
|
39
|
+
* Prepares the two radix pipelines and the scan of a scope (compiles once) so record() is synchronous. The planner
|
|
40
|
+
* lives exactly as long as the scope: never use a planner after its scope's dispose().
|
|
41
|
+
* @param scope - the caller's scope (device, caps, cache, scratch, params)
|
|
42
|
+
* @returns the planner
|
|
43
|
+
*/
|
|
44
|
+
export async function prepareRadixSort(scope) {
|
|
45
|
+
const hist = await scope.pipelines.kernel(kernelSpec("radix-hist"));
|
|
46
|
+
const scatter = await scope.pipelines.kernel(kernelSpec("radix-scatter"));
|
|
47
|
+
const scan = await prepareScan(scope);
|
|
48
|
+
return new RadixSortPlannerImpl(scope, hist, scatter, scan);
|
|
49
|
+
}
|
|
50
|
+
/**
|
|
51
|
+
* One binding-size check of record().
|
|
52
|
+
* @param argument - the argument name
|
|
53
|
+
* @param binding - the binding
|
|
54
|
+
* @param bytes - the bytes it must hold
|
|
55
|
+
*/
|
|
56
|
+
function checkBinding(argument, binding, bytes) {
|
|
57
|
+
if (binding.size < bytes) {
|
|
58
|
+
throw new WebGpuGraphError("E_INVALID_ARGUMENT", `radixSort: ${argument} is smaller than ${bytes} bytes`, {
|
|
59
|
+
argument,
|
|
60
|
+
value: binding.size,
|
|
61
|
+
expected: bytes,
|
|
62
|
+
});
|
|
63
|
+
}
|
|
64
|
+
}
|
|
65
|
+
/**
|
|
66
|
+
* The argument checks of record() (E_INVALID_ARGUMENT before anything is recorded).
|
|
67
|
+
* @param keys - the keys
|
|
68
|
+
* @param vals - the values
|
|
69
|
+
* @param count - the pair count
|
|
70
|
+
* @param bits - the key width
|
|
71
|
+
* @param scratch - the scratch
|
|
72
|
+
* @param wg - the workgroup size
|
|
73
|
+
*/
|
|
74
|
+
function checkRecordArguments(keys, vals, count, bits, scratch, wg) {
|
|
75
|
+
if (bits !== 8 && bits !== 16 && bits !== 24 && bits !== 32) {
|
|
76
|
+
throw new WebGpuGraphError("E_INVALID_ARGUMENT", "radixSort: bits must be 8, 16, 24 or 32", {
|
|
77
|
+
argument: "bits",
|
|
78
|
+
value: bits,
|
|
79
|
+
expected: [8, 16, 24, 32],
|
|
80
|
+
});
|
|
81
|
+
}
|
|
82
|
+
if (!Number.isSafeInteger(count) || count < 0 || count > U32_MAX) {
|
|
83
|
+
throw new WebGpuGraphError("E_INVALID_ARGUMENT", "radixSort: count must be a non-negative integer below 2^32", {
|
|
84
|
+
argument: "count",
|
|
85
|
+
value: count,
|
|
86
|
+
});
|
|
87
|
+
}
|
|
88
|
+
const pairBytes = 4 * count;
|
|
89
|
+
checkBinding("keys", keys, pairBytes);
|
|
90
|
+
checkBinding("vals", vals, pairBytes);
|
|
91
|
+
checkBinding("scratch.keys", scratch.keys, pairBytes);
|
|
92
|
+
checkBinding("scratch.vals", scratch.vals, pairBytes);
|
|
93
|
+
const tableBytes = radixHistBytes(count, wg);
|
|
94
|
+
checkBinding("scratch.hist", scratch.hist, tableBytes);
|
|
95
|
+
checkBinding("scratch.offsets", scratch.offsets, tableBytes);
|
|
96
|
+
}
|
|
97
|
+
/** The planner: the two resolved kernels and the scan over one scope. */
|
|
98
|
+
class RadixSortPlannerImpl {
|
|
99
|
+
/**
|
|
100
|
+
* Wraps the resolved kernels; use prepareRadixSort().
|
|
101
|
+
* @param scope - the caller's scope
|
|
102
|
+
* @param hist - the `radix-hist` kernel
|
|
103
|
+
* @param scatter - the `radix-scatter` kernel
|
|
104
|
+
* @param scan - the scan planner of the same scope
|
|
105
|
+
*/
|
|
106
|
+
constructor(scope, hist, scatter, scan) {
|
|
107
|
+
this.dispatches = 0;
|
|
108
|
+
this.scope = scope;
|
|
109
|
+
this.hist = hist;
|
|
110
|
+
this.scatter = scatter;
|
|
111
|
+
this.scan = scan;
|
|
112
|
+
}
|
|
113
|
+
/**
|
|
114
|
+
* Dispatches the last record() issued.
|
|
115
|
+
* @returns the count
|
|
116
|
+
*/
|
|
117
|
+
get lastDispatches() {
|
|
118
|
+
return this.dispatches;
|
|
119
|
+
}
|
|
120
|
+
/**
|
|
121
|
+
* Records the passes into the pass (see the interface).
|
|
122
|
+
* @param pass - the compute pass
|
|
123
|
+
* @param keys - the keys
|
|
124
|
+
* @param vals - the values
|
|
125
|
+
* @param count - the pair count
|
|
126
|
+
* @param bits - the key width
|
|
127
|
+
* @param scratch - the second pair, the histogram table and its scanned twin
|
|
128
|
+
* @returns the pair the result lives in
|
|
129
|
+
*/
|
|
130
|
+
record(pass, keys, vals, count, bits, scratch) {
|
|
131
|
+
const wg = this.scope.workgroupSize;
|
|
132
|
+
checkRecordArguments(keys, vals, count, bits, scratch, wg);
|
|
133
|
+
if (count === 0) {
|
|
134
|
+
this.dispatches = 0;
|
|
135
|
+
return { keys, vals };
|
|
136
|
+
}
|
|
137
|
+
const groups = Math.ceil(count / wg);
|
|
138
|
+
const plan = plan1d(count, wg, this.scope.caps);
|
|
139
|
+
const tableWords = RADIX_BINS * groups;
|
|
140
|
+
const tableBytes = 4 * tableWords;
|
|
141
|
+
const histTable = { ...scratch.hist, size: tableBytes };
|
|
142
|
+
const offsets = { ...scratch.offsets, size: tableBytes };
|
|
143
|
+
let src = { keys, vals };
|
|
144
|
+
let dst = { keys: scratch.keys, vals: scratch.vals };
|
|
145
|
+
let dispatches = 0;
|
|
146
|
+
const passes = bits / RADIX_DIGIT_BITS;
|
|
147
|
+
for (let p = 0; p < passes; p++) {
|
|
148
|
+
const params = this.scope.params(RADIX_PARAMS, { count, shift: RADIX_DIGIT_BITS * p, groups, pad0: 0 });
|
|
149
|
+
const histBound = this.hist.bind({ keys: src.keys, hist: histTable, P: params.binding });
|
|
150
|
+
this.hist.dispatch(pass, histBound, plan, [params.offset]);
|
|
151
|
+
this.scan.record(pass, histTable, tableWords, offsets);
|
|
152
|
+
const scatterBound = this.scatter.bind({
|
|
153
|
+
keys: src.keys,
|
|
154
|
+
vals: src.vals,
|
|
155
|
+
offsets,
|
|
156
|
+
keysOut: dst.keys,
|
|
157
|
+
valsOut: dst.vals,
|
|
158
|
+
P: params.binding,
|
|
159
|
+
});
|
|
160
|
+
this.scatter.dispatch(pass, scatterBound, plan, [params.offset]);
|
|
161
|
+
dispatches += 2 + this.scan.lastDispatches;
|
|
162
|
+
[src, dst] = [dst, src];
|
|
163
|
+
}
|
|
164
|
+
this.dispatches = dispatches;
|
|
165
|
+
return src;
|
|
166
|
+
}
|
|
167
|
+
}
|
|
168
|
+
//# sourceMappingURL=radix-sort.js.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"radix-sort.js","sourceRoot":"","sources":["../../../src/primitives/radix-sort.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;;;;;;;;;;GAmBG;AAEH,OAAO,EAAE,UAAU,EAAE,OAAO,EAAE,MAAM,iBAAiB,CAAC;AACtD,OAAO,EAAE,gBAAgB,EAAE,MAAM,cAAc,CAAC;AAChD,OAAO,EAAE,MAAM,EAAE,MAAM,uBAAuB,CAAC;AAE/C,OAAO,EAAE,UAAU,EAAE,YAAY,EAAE,MAAM,eAAe,CAAC;AAGzD,OAAO,EAAE,WAAW,EAAoB,MAAM,WAAW,CAAC;AAE1D,4BAA4B;AAC5B,MAAM,gBAAgB,GAAG,CAAC,CAAC;AAkD3B;;;;;;GAMG;AACH,MAAM,UAAU,cAAc,CAAC,KAAa,EAAE,EAAU;IACpD,OAAO,CAAC,GAAG,UAAU,GAAG,IAAI,CAAC,IAAI,CAAC,KAAK,GAAG,EAAE,CAAC,CAAC;AAClD,CAAC;AAED;;;;;GAKG;AACH,MAAM,CAAC,KAAK,UAAU,gBAAgB,CAAC,KAAkB;IACrD,MAAM,IAAI,GAAG,MAAM,KAAK,CAAC,SAAS,CAAC,MAAM,CAAC,UAAU,CAAC,YAAY,CAAC,CAAC,CAAC;IACpE,MAAM,OAAO,GAAG,MAAM,KAAK,CAAC,SAAS,CAAC,MAAM,CAAC,UAAU,CAAC,eAAe,CAAC,CAAC,CAAC;IAC1E,MAAM,IAAI,GAAG,MAAM,WAAW,CAAC,KAAK,CAAC,CAAC;IACtC,OAAO,IAAI,oBAAoB,CAAC,KAAK,EAAE,IAAI,EAAE,OAAO,EAAE,IAAI,CAAC,CAAC;AAChE,CAAC;AAED;;;;;GAKG;AACH,SAAS,YAAY,CAAC,QAAgB,EAAE,OAAgB,EAAE,KAAa;IACnE,IAAI,OAAO,CAAC,IAAI,GAAG,KAAK,EAAE,CAAC;QACvB,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,cAAc,QAAQ,oBAAoB,KAAK,QAAQ,EAAE;YACtG,QAAQ;YACR,KAAK,EAAE,OAAO,CAAC,IAAI;YACnB,QAAQ,EAAE,KAAK;SAClB,CAAC,CAAC;IACP,CAAC;AACL,CAAC;AAED;;;;;;;;GAQG;AACH,SAAS,oBAAoB,CACzB,IAAa,EACb,IAAa,EACb,KAAa,EACb,IAAe,EACf,OAAyB,EACzB,EAAU;IAEV,IAAI,IAAI,KAAK,CAAC,IAAI,IAAI,KAAK,EAAE,IAAI,IAAI,KAAK,EAAE,IAAI,IAAI,KAAK,EAAE,EAAE,CAAC;QAC1D,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,yCAAyC,EAAE;YACxF,QAAQ,EAAE,MAAM;YAChB,KAAK,EAAE,IAAI;YACX,QAAQ,EAAE,CAAC,CAAC,EAAE,EAAE,EAAE,EAAE,EAAE,EAAE,CAAC;SAC5B,CAAC,CAAC;IACP,CAAC;IACD,IAAI,CAAC,MAAM,CAAC,aAAa,CAAC,KAAK,CAAC,IAAI,KAAK,GAAG,CAAC,IAAI,KAAK,GAAG,OAAO,EAAE,CAAC;QAC/D,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,4DAA4D,EAAE;YAC3G,QAAQ,EAAE,OAAO;YACjB,KAAK,EAAE,KAAK;SACf,CAAC,CAAC;IACP,CAAC;IACD,MAAM,SAAS,GAAG,CAAC,GAAG,KAAK,CAAC;IAC5B,YAAY,CAAC,MAAM,EAAE,IAAI,EAAE,SAAS,CAAC,CAAC;IACtC,YAAY,CAAC,MAAM,EAAE,IAAI,EAAE,SAAS,CAAC,CAAC;IACtC,YAAY,CAAC,cAAc,EAAE,OAAO,CAAC,IAAI,EAAE,SAAS,CAAC,CAAC;IACtD,YAAY,CAAC,cAAc,EAAE,OAAO,CAAC,IAAI,EAAE,SAAS,CAAC,CAAC;IACtD,MAAM,UAAU,GAAG,cAAc,CAAC,KAAK,EAAE,EAAE,CAAC,CAAC;IAC7C,YAAY,CAAC,cAAc,EAAE,OAAO,CAAC,IAAI,EAAE,UAAU,CAAC,CAAC;IACvD,YAAY,CAAC,iBAAiB,EAAE,OAAO,CAAC,OAAO,EAAE,UAAU,CAAC,CAAC;AACjE,CAAC;AAED,yEAAyE;AACzE,MAAM,oBAAoB;IAOtB;;;;;;OAMG;IACH,YAAY,KAAkB,EAAE,IAAY,EAAE,OAAe,EAAE,IAAiB;QATxE,eAAU,GAAG,CAAC,CAAC;QAUnB,IAAI,CAAC,KAAK,GAAG,KAAK,CAAC;QACnB,IAAI,CAAC,IAAI,GAAG,IAAI,CAAC;QACjB,IAAI,CAAC,OAAO,GAAG,OAAO,CAAC;QACvB,IAAI,CAAC,IAAI,GAAG,IAAI,CAAC;IACrB,CAAC;IAED;;;OAGG;IACH,IAAI,cAAc;QACd,OAAO,IAAI,CAAC,UAAU,CAAC;IAC3B,CAAC;IAED;;;;;;;;;OASG;IACH,MAAM,CACF,IAA2B,EAC3B,IAAa,EACb,IAAa,EACb,KAAa,EACb,IAAe,EACf,OAAyB;QAEzB,MAAM,EAAE,GAAG,IAAI,CAAC,KAAK,CAAC,aAAa,CAAC;QACpC,oBAAoB,CAAC,IAAI,EAAE,IAAI,EAAE,KAAK,EAAE,IAAI,EAAE,OAAO,EAAE,EAAE,CAAC,CAAC;QAC3D,IAAI,KAAK,KAAK,CAAC,EAAE,CAAC;YACd,IAAI,CAAC,UAAU,GAAG,CAAC,CAAC;YACpB,OAAO,EAAE,IAAI,EAAE,IAAI,EAAE,CAAC;QAC1B,CAAC;QACD,MAAM,MAAM,GAAG,IAAI,CAAC,IAAI,CAAC,KAAK,GAAG,EAAE,CAAC,CAAC;QACrC,MAAM,IAAI,GAAG,MAAM,CAAC,KAAK,EAAE,EAAE,EAAE,IAAI,CAAC,KAAK,CAAC,IAAI,CAAC,CAAC;QAChD,MAAM,UAAU,GAAG,UAAU,GAAG,MAAM,CAAC;QACvC,MAAM,UAAU,GAAG,CAAC,GAAG,UAAU,CAAC;QAClC,MAAM,SAAS,GAAY,EAAE,GAAG,OAAO,CAAC,IAAI,EAAE,IAAI,EAAE,UAAU,EAAE,CAAC;QACjE,MAAM,OAAO,GAAY,EAAE,GAAG,OAAO,CAAC,OAAO,EAAE,IAAI,EAAE,UAAU,EAAE,CAAC;QAClE,IAAI,GAAG,GAAoB,EAAE,IAAI,EAAE,IAAI,EAAE,CAAC;QAC1C,IAAI,GAAG,GAAoB,EAAE,IAAI,EAAE,OAAO,CAAC,IAAI,EAAE,IAAI,EAAE,OAAO,CAAC,IAAI,EAAE,CAAC;QACtE,IAAI,UAAU,GAAG,CAAC,CAAC;QACnB,MAAM,MAAM,GAAG,IAAI,GAAG,gBAAgB,CAAC;QACvC,KAAK,IAAI,CAAC,GAAG,CAAC,EAAE,CAAC,GAAG,MAAM,EAAE,CAAC,EAAE,EAAE,CAAC;YAC9B,MAAM,MAAM,GAAG,IAAI,CAAC,KAAK,CAAC,MAAM,CAAC,YAAY,EAAE,EAAE,KAAK,EAAE,KAAK,EAAE,gBAAgB,GAAG,CAAC,EAAE,MAAM,EAAE,IAAI,EAAE,CAAC,EAAE,CAAC,CAAC;YACxG,MAAM,SAAS,GAAG,IAAI,CAAC,IAAI,CAAC,IAAI,CAAC,EAAE,IAAI,EAAE,GAAG,CAAC,IAAI,EAAE,IAAI,EAAE,SAAS,EAAE,CAAC,EAAE,MAAM,CAAC,OAAO,EAAE,CAAC,CAAC;YACzF,IAAI,CAAC,IAAI,CAAC,QAAQ,CAAC,IAAI,EAAE,SAAS,EAAE,IAAI,EAAE,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC;YAC3D,IAAI,CAAC,IAAI,CAAC,MAAM,CAAC,IAAI,EAAE,SAAS,EAAE,UAAU,EAAE,OAAO,CAAC,CAAC;YACvD,MAAM,YAAY,GAAG,IAAI,CAAC,OAAO,CAAC,IAAI,CAAC;gBACnC,IAAI,EAAE,GAAG,CAAC,IAAI;gBACd,IAAI,EAAE,GAAG,CAAC,IAAI;gBACd,OAAO;gBACP,OAAO,EAAE,GAAG,CAAC,IAAI;gBACjB,OAAO,EAAE,GAAG,CAAC,IAAI;gBACjB,CAAC,EAAE,MAAM,CAAC,OAAO;aACpB,CAAC,CAAC;YACH,IAAI,CAAC,OAAO,CAAC,QAAQ,CAAC,IAAI,EAAE,YAAY,EAAE,IAAI,EAAE,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC;YACjE,UAAU,IAAI,CAAC,GAAG,IAAI,CAAC,IAAI,CAAC,cAAc,CAAC;YAC3C,CAAC,GAAG,EAAE,GAAG,CAAC,GAAG,CAAC,GAAG,EAAE,GAAG,CAAC,CAAC;QAC5B,CAAC;QACD,IAAI,CAAC,UAAU,GAAG,UAAU,CAAC;QAC7B,OAAO,GAAG,CAAC;IACf,CAAC;CACJ"}
|
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
/**
|
|
2
|
+
* The `exclusiveScan` primitive driver (spec 6 row 2; P4-T2, PD-3): a reduce-then-scan over u32 words. Level 0 runs
|
|
3
|
+
* `scan-block` over `count` items into `out` (each workgroup an exclusive Hillis-Steele scan of its WG-wide block) and
|
|
4
|
+
* writes one block sum per workgroup; while a level has more than one block, the next level scans its block sums
|
|
5
|
+
* (exclusive, into `offsets_L`) and writes its own block sums, recursively; the top level has ONE block, so its one
|
|
6
|
+
* sums word is the total. Then, top down, `scan-add` adds `offsets_L[g]` to every element of block `g` of level L's
|
|
7
|
+
* output. u32 addition is exact in any order, so two runs on any two adapters are bitwise identical; there is no
|
|
8
|
+
* subgroup variant (DEP-P4-E) and no decoupled look-back.
|
|
9
|
+
*
|
|
10
|
+
* The driver owns no device objects beyond a one-word zero it keeps for the empty scan: the caller supplies a
|
|
11
|
+
* ReduceScope (the same record `reduce` takes) and the compute pass to record into. `src/primitives/**` never imports
|
|
12
|
+
* `src/context.ts`.
|
|
13
|
+
*/
|
|
14
|
+
import { type Binding } from "../types/memory.js";
|
|
15
|
+
import { type ReduceScope } from "./reduce.js";
|
|
16
|
+
/** Where the scan's total landed: one u32 word at `index` of `binding` (a word of the planner's scratch, valid until the next record()). */
|
|
17
|
+
export interface ScanTotal {
|
|
18
|
+
readonly binding: Binding;
|
|
19
|
+
readonly index: number;
|
|
20
|
+
}
|
|
21
|
+
/** A prepared exclusive scan (spec 6 row 2): records the level dispatches of one scan into a pass. */
|
|
22
|
+
export interface ScanPlanner {
|
|
23
|
+
/**
|
|
24
|
+
* Records the scan of `count` u32 of `src` into `out` (exclusive); returns where the total landed (a word of the
|
|
25
|
+
* planner's scratch, valid until the next record()); nothing for count 0 (the total word is then 0).
|
|
26
|
+
* @param pass - the compute pass
|
|
27
|
+
* @param src - the input words (at least 4 x count bytes)
|
|
28
|
+
* @param count - the word count (a non-negative integer below 2^32)
|
|
29
|
+
* @param out - the output words (at least 4 x count bytes)
|
|
30
|
+
* @returns the total's location
|
|
31
|
+
*/
|
|
32
|
+
record(pass: GPUComputePassEncoder, src: Binding, count: number, out: Binding): ScanTotal;
|
|
33
|
+
/** Dispatches the last record() issued: 0 for count 0, else `2 x levels - 1` (one block scan per level, one add-back per level below the top). */
|
|
34
|
+
readonly lastDispatches: number;
|
|
35
|
+
}
|
|
36
|
+
/**
|
|
37
|
+
* Prepares the two scan pipelines of a scope (compiles once) so record() is synchronous. The zero word the empty
|
|
38
|
+
* scan's total points at is a scratch of `scope`, so the planner lives exactly as long as the scope: never use a
|
|
39
|
+
* planner after its scope's dispose().
|
|
40
|
+
* @param scope - the caller's scope (device, caps, cache, scratch, params)
|
|
41
|
+
* @returns the planner
|
|
42
|
+
*/
|
|
43
|
+
export declare function prepareScan(scope: ReduceScope): Promise<ScanPlanner>;
|
|
44
|
+
//# sourceMappingURL=scan.d.ts.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"scan.d.ts","sourceRoot":"","sources":["../../../src/primitives/scan.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;;;GAYG;AAMH,OAAO,EAAE,KAAK,OAAO,EAAE,MAAM,oBAAoB,CAAC;AAClD,OAAO,EAAE,KAAK,WAAW,EAAE,MAAM,aAAa,CAAC;AAE/C,4IAA4I;AAC5I,MAAM,WAAW,SAAS;IACtB,QAAQ,CAAC,OAAO,EAAE,OAAO,CAAC;IAC1B,QAAQ,CAAC,KAAK,EAAE,MAAM,CAAC;CAC1B;AAED,sGAAsG;AACtG,MAAM,WAAW,WAAW;IACxB;;;;;;;;OAQG;IACH,MAAM,CAAC,IAAI,EAAE,qBAAqB,EAAE,GAAG,EAAE,OAAO,EAAE,KAAK,EAAE,MAAM,EAAE,GAAG,EAAE,OAAO,GAAG,SAAS,CAAC;IAC1F,kJAAkJ;IAClJ,QAAQ,CAAC,cAAc,EAAE,MAAM,CAAC;CACnC;AAKD;;;;;;GAMG;AACH,wBAAsB,WAAW,CAAC,KAAK,EAAE,WAAW,GAAG,OAAO,CAAC,WAAW,CAAC,CAM1E"}
|
|
@@ -0,0 +1,151 @@
|
|
|
1
|
+
/**
|
|
2
|
+
* The `exclusiveScan` primitive driver (spec 6 row 2; P4-T2, PD-3): a reduce-then-scan over u32 words. Level 0 runs
|
|
3
|
+
* `scan-block` over `count` items into `out` (each workgroup an exclusive Hillis-Steele scan of its WG-wide block) and
|
|
4
|
+
* writes one block sum per workgroup; while a level has more than one block, the next level scans its block sums
|
|
5
|
+
* (exclusive, into `offsets_L`) and writes its own block sums, recursively; the top level has ONE block, so its one
|
|
6
|
+
* sums word is the total. Then, top down, `scan-add` adds `offsets_L[g]` to every element of block `g` of level L's
|
|
7
|
+
* output. u32 addition is exact in any order, so two runs on any two adapters are bitwise identical; there is no
|
|
8
|
+
* subgroup variant (DEP-P4-E) and no decoupled look-back.
|
|
9
|
+
*
|
|
10
|
+
* The driver owns no device objects beyond a one-word zero it keeps for the empty scan: the caller supplies a
|
|
11
|
+
* ReduceScope (the same record `reduce` takes) and the compute pass to record into. `src/primitives/**` never imports
|
|
12
|
+
* `src/context.ts`.
|
|
13
|
+
*/
|
|
14
|
+
import { WebGpuGraphError } from "../errors.js";
|
|
15
|
+
import { groupsOf, plan1d } from "../kernel/dispatch.js";
|
|
16
|
+
import { kernelSpec, SCAN_PARAMS } from "../kernels.js";
|
|
17
|
+
/** The largest u32: the largest count `ScanParams.count` can carry. */
|
|
18
|
+
const U32_MAX = 0xffffffff;
|
|
19
|
+
/**
|
|
20
|
+
* Prepares the two scan pipelines of a scope (compiles once) so record() is synchronous. The zero word the empty
|
|
21
|
+
* scan's total points at is a scratch of `scope`, so the planner lives exactly as long as the scope: never use a
|
|
22
|
+
* planner after its scope's dispose().
|
|
23
|
+
* @param scope - the caller's scope (device, caps, cache, scratch, params)
|
|
24
|
+
* @returns the planner
|
|
25
|
+
*/
|
|
26
|
+
export async function prepareScan(scope) {
|
|
27
|
+
const block = await scope.pipelines.kernel(kernelSpec("scan-block"));
|
|
28
|
+
const add = await scope.pipelines.kernel(kernelSpec("scan-add"));
|
|
29
|
+
const zero = scope.scratch(4, "scan/zero");
|
|
30
|
+
scope.device.queue.writeBuffer(zero, 0, new Uint32Array(1));
|
|
31
|
+
return new ScanPlannerImpl(scope, block, add, { buffer: zero, offset: 0, size: 4, window: null });
|
|
32
|
+
}
|
|
33
|
+
/**
|
|
34
|
+
* The argument checks of record() (E_INVALID_ARGUMENT before anything is recorded).
|
|
35
|
+
* @param src - the input binding
|
|
36
|
+
* @param count - the word count
|
|
37
|
+
* @param out - the output binding
|
|
38
|
+
*/
|
|
39
|
+
function checkRecordArguments(src, count, out) {
|
|
40
|
+
if (!Number.isSafeInteger(count) || count < 0 || count > U32_MAX) {
|
|
41
|
+
throw new WebGpuGraphError("E_INVALID_ARGUMENT", "scan: count must be a non-negative integer below 2^32", {
|
|
42
|
+
argument: "count",
|
|
43
|
+
value: count,
|
|
44
|
+
});
|
|
45
|
+
}
|
|
46
|
+
if (src.size < 4 * count) {
|
|
47
|
+
throw new WebGpuGraphError("E_INVALID_ARGUMENT", "scan: src is smaller than 4 x count bytes", {
|
|
48
|
+
argument: "src",
|
|
49
|
+
value: src.size,
|
|
50
|
+
expected: 4 * count,
|
|
51
|
+
});
|
|
52
|
+
}
|
|
53
|
+
if (out.size < 4 * count) {
|
|
54
|
+
throw new WebGpuGraphError("E_INVALID_ARGUMENT", "scan: out is smaller than 4 x count bytes", {
|
|
55
|
+
argument: "out",
|
|
56
|
+
value: out.size,
|
|
57
|
+
expected: 4 * count,
|
|
58
|
+
});
|
|
59
|
+
}
|
|
60
|
+
}
|
|
61
|
+
/** The planner: the two resolved kernels over one scope. */
|
|
62
|
+
class ScanPlannerImpl {
|
|
63
|
+
/**
|
|
64
|
+
* Wraps the resolved kernels; use prepareScan().
|
|
65
|
+
* @param scope - the caller's scope
|
|
66
|
+
* @param block - the `scan-block` kernel
|
|
67
|
+
* @param add - the `scan-add` kernel
|
|
68
|
+
* @param zero - a one-word binding holding 0 (the total of the empty scan)
|
|
69
|
+
*/
|
|
70
|
+
constructor(scope, block, add, zero) {
|
|
71
|
+
this.dispatches = 0;
|
|
72
|
+
this.scope = scope;
|
|
73
|
+
this.block = block;
|
|
74
|
+
this.add = add;
|
|
75
|
+
this.zero = zero;
|
|
76
|
+
}
|
|
77
|
+
/**
|
|
78
|
+
* Dispatches the last record() issued.
|
|
79
|
+
* @returns the count
|
|
80
|
+
*/
|
|
81
|
+
get lastDispatches() {
|
|
82
|
+
return this.dispatches;
|
|
83
|
+
}
|
|
84
|
+
/**
|
|
85
|
+
* Records the levels into the pass (see the interface).
|
|
86
|
+
* @param pass - the compute pass
|
|
87
|
+
* @param src - the input words
|
|
88
|
+
* @param count - the word count
|
|
89
|
+
* @param out - the output words
|
|
90
|
+
* @returns the total's location
|
|
91
|
+
*/
|
|
92
|
+
record(pass, src, count, out) {
|
|
93
|
+
checkRecordArguments(src, count, out);
|
|
94
|
+
if (count === 0) {
|
|
95
|
+
this.dispatches = 0;
|
|
96
|
+
return { binding: this.zero, index: 0 };
|
|
97
|
+
}
|
|
98
|
+
const wg = this.scope.workgroupSize;
|
|
99
|
+
const levels = [];
|
|
100
|
+
let input = src;
|
|
101
|
+
let output = out;
|
|
102
|
+
let levelCount = count;
|
|
103
|
+
for (;;) {
|
|
104
|
+
const blocks = Math.ceil(levelCount / wg);
|
|
105
|
+
// Sized by the PADDED workgroup count (reduce's shape, reduce.ts partials-1): above MAX_1D_ITEMS plan1d pads
|
|
106
|
+
// the grid to x * y >= blocks, and every padded workgroup still stores its (zero) block sum at its own index.
|
|
107
|
+
// The next level scans only `blocks` words, so the padded tail is written and never read.
|
|
108
|
+
const plan = plan1d(levelCount, wg, this.scope.caps);
|
|
109
|
+
const groups = groupsOf(plan);
|
|
110
|
+
const sums = this.scratch(groups, `scan/sums${levels.length}`);
|
|
111
|
+
levels.push({ input, output, count: levelCount, plan, sums });
|
|
112
|
+
if (blocks === 1) {
|
|
113
|
+
break;
|
|
114
|
+
}
|
|
115
|
+
input = sums;
|
|
116
|
+
output = this.scratch(groups, `scan/offsets${levels.length}`);
|
|
117
|
+
levelCount = blocks;
|
|
118
|
+
}
|
|
119
|
+
for (const level of levels) {
|
|
120
|
+
const params = this.scope.params(SCAN_PARAMS, { count: level.count, pad0: 0, pad1: 0, pad2: 0 });
|
|
121
|
+
const bound = this.block.bind({
|
|
122
|
+
src: level.input,
|
|
123
|
+
out: level.output,
|
|
124
|
+
blockSums: level.sums,
|
|
125
|
+
P: params.binding,
|
|
126
|
+
});
|
|
127
|
+
this.block.dispatch(pass, bound, level.plan, [params.offset]);
|
|
128
|
+
}
|
|
129
|
+
for (let l = levels.length - 2; l >= 0; l--) {
|
|
130
|
+
const level = levels[l];
|
|
131
|
+
const params = this.scope.params(SCAN_PARAMS, { count: level.count, pad0: 0, pad1: 0, pad2: 0 });
|
|
132
|
+
const bound = this.add.bind({ out: level.output, blockOffsets: levels[l + 1].output, P: params.binding });
|
|
133
|
+
this.add.dispatch(pass, bound, level.plan, [params.offset]);
|
|
134
|
+
}
|
|
135
|
+
this.dispatches = 2 * levels.length - 1;
|
|
136
|
+
const top = levels[levels.length - 1];
|
|
137
|
+
return { binding: top.sums, index: 0 };
|
|
138
|
+
}
|
|
139
|
+
/**
|
|
140
|
+
* A scratch of `words` u32 from the scope, bound whole (never zero-length; the pool rounds the buffer up).
|
|
141
|
+
* @param words - the word count (>= 1)
|
|
142
|
+
* @param label - the scratch label
|
|
143
|
+
* @returns the binding
|
|
144
|
+
*/
|
|
145
|
+
scratch(words, label) {
|
|
146
|
+
const size = 4 * words;
|
|
147
|
+
const buffer = this.scope.scratch(size, label);
|
|
148
|
+
return { buffer, offset: 0, size, window: null };
|
|
149
|
+
}
|
|
150
|
+
}
|
|
151
|
+
//# sourceMappingURL=scan.js.map
|