@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
|
@@ -1,27 +1,29 @@
|
|
|
1
1
|
/**
|
|
2
|
-
* The segmented (per-row) reduction primitive of spec 6 row 3
|
|
3
|
-
*
|
|
4
|
-
*
|
|
5
|
-
* the
|
|
6
|
-
*
|
|
7
|
-
*
|
|
8
|
-
*
|
|
2
|
+
* The segmented (per-row) reduction primitive of spec 6 row 3: the caller's VALUE snippet folded over every CSR
|
|
3
|
+
* row's arcs into `out[row]` (f32), a row with no arcs receiving the identity element. Without tiers it is ONE
|
|
4
|
+
* thread-per-row dispatch (TIER 0, USE_PERM false: the perm slot carries the rowPtr dummy of graphBindings). With
|
|
5
|
+
* the degree tiers of degreeOrder() (P4-T5, PD-6) it is up to three dispatches over the permuted rows: TIER 2, one
|
|
6
|
+
* workgroup per row of degree >= 1024 over [0, hiEnd); TIER 1, 32 lanes per row of degree 32..1023 over [hiEnd,
|
|
7
|
+
* midEnd); TIER 0, one thread per row over [midEnd, n) -- each compiled only when its range is non-empty. The row
|
|
8
|
+
* and arc counts come from the core's binding sizes: the residency binds every array at its exact byte length
|
|
9
|
+
* (contract 3.8), so rowPtr is 4(n + 1) bytes and colIdx 4 x arcCount. On a windowed core (spec 4.2, P4-T7, PD-8)
|
|
10
|
+
* the dispatches repeat per window with `arcBase = w.start`, `arcEnd = w.end` and `accumulate = 1` over an out
|
|
11
|
+
* pre-filled with the identity element (unless the caller asked to accumulate), so a row split across windows
|
|
12
|
+
* combines its partials: the untiered dispatch covers the window's rows [rowFirst, rowLast], while every tier
|
|
13
|
+
* dispatch covers its FULL tier range (permutation positions, not comparable with the window's node rows; a row
|
|
14
|
+
* whose arcs lie outside the window folds nothing and `finish` writes comb(out[i], identity) = out[i]).
|
|
9
15
|
*/
|
|
10
16
|
|
|
11
17
|
import { WebGpuGraphError } from "../errors.js";
|
|
12
18
|
import { plan1d } from "../kernel/dispatch.js";
|
|
13
19
|
import { type Kernel } from "../kernel/kernel.js";
|
|
14
|
-
import { graphBindings, graphOverrides, kernelSpec, RANGE_PARAMS } from "../kernels.js";
|
|
20
|
+
import { FILL_PARAMS, graphBindings, graphOverrides, kernelSpec, RANGE_PARAMS } from "../kernels.js";
|
|
15
21
|
import { type CoreBinding } from "../memory/residency.js";
|
|
16
|
-
import { type Binding } from "../types/memory.js";
|
|
17
|
-
import {
|
|
22
|
+
import { type ArcWindow, type Binding } from "../types/memory.js";
|
|
23
|
+
import { arcCountOf, type DegreeTiers, MID_TIER_LANES, rowCountOf, windowBinding } from "./core-shape.js";
|
|
18
24
|
import { type ReduceOp, type ReduceScope } from "./reduce.js";
|
|
19
25
|
|
|
20
|
-
|
|
21
|
-
export interface DegreeTiers {
|
|
22
|
-
readonly perm: Binding;
|
|
23
|
-
readonly segmentOffsets: readonly [number, number, number, number, number];
|
|
24
|
-
}
|
|
26
|
+
export { type DegreeTiers } from "./core-shape.js";
|
|
25
27
|
|
|
26
28
|
/** Options of segmentedReduce. `valueSnippet` is the Gunrock-style functor: WGSL statements assigning `v` from (row, arc, nbr, weight) (4.5; `nbr` because `target` is a WGSL reserved word). */
|
|
27
29
|
export interface SegmentedReduceOptions {
|
|
@@ -31,10 +33,33 @@ export interface SegmentedReduceOptions {
|
|
|
31
33
|
readonly accumulate?: boolean | undefined;
|
|
32
34
|
}
|
|
33
35
|
|
|
34
|
-
/** A prepared segmented reduce
|
|
36
|
+
/** A prepared segmented reduce: one thread-per-row dispatch without tiers, up to three tier dispatches with them (times the windows of a windowed core, plus its identity fill). */
|
|
35
37
|
export interface SegmentedReducePlanner {
|
|
36
|
-
/** Records
|
|
38
|
+
/** Records the dispatches over rows [0, n) writing out[i] (f32) per row; a row with no arcs gets the identity element. */
|
|
37
39
|
record(pass: GPUComputePassEncoder, core: CoreBinding, out: Binding): void;
|
|
40
|
+
/** Dispatches the last record() issued (0 for n = 0; 1 without tiers; 1 to 3 with them; on a windowed core the identity fill plus that many per window). */
|
|
41
|
+
readonly lastDispatches: number;
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
/** The largest finite f32, 0x1.fffffep+127 (the prelude's F32_MAX): the min identity; its negation the max identity. */
|
|
45
|
+
const F32_MAX = 2 ** 128 - 2 ** 104;
|
|
46
|
+
|
|
47
|
+
/**
|
|
48
|
+
* The u32 bit pattern the `fill` kernel writes so every out word holds the identity element of the operator before
|
|
49
|
+
* the windowed dispatches accumulate into it (PD-8).
|
|
50
|
+
* @param op - the operator
|
|
51
|
+
* @returns the bits of 0 (sum), F32_MAX (min) or -F32_MAX (max)
|
|
52
|
+
*/
|
|
53
|
+
function identityFillWord(op: ReduceOp): number {
|
|
54
|
+
let identity = 0;
|
|
55
|
+
if (op === "min") {
|
|
56
|
+
identity = F32_MAX;
|
|
57
|
+
} else if (op === "max") {
|
|
58
|
+
identity = -F32_MAX;
|
|
59
|
+
}
|
|
60
|
+
const view = new DataView(new ArrayBuffer(4));
|
|
61
|
+
view.setFloat32(0, identity, true);
|
|
62
|
+
return view.getUint32(0, true);
|
|
38
63
|
}
|
|
39
64
|
|
|
40
65
|
/** The identifiers a VALUE snippet may name (contract 3.11, 4.5); every other identifier is rejected textually before compose. */
|
|
@@ -141,36 +166,75 @@ function validateValueSnippet(snippet: string): void {
|
|
|
141
166
|
}
|
|
142
167
|
}
|
|
143
168
|
|
|
144
|
-
/**
|
|
145
|
-
|
|
169
|
+
/** One tier's compiled pipeline and the rows [start, end) of the permutation it covers per record() (a tier whose range is empty is not compiled). */
|
|
170
|
+
interface TierDispatch {
|
|
171
|
+
readonly tier: 0 | 1 | 2;
|
|
172
|
+
readonly kernel: Kernel;
|
|
173
|
+
readonly start: number;
|
|
174
|
+
readonly end: number;
|
|
175
|
+
}
|
|
176
|
+
|
|
177
|
+
/**
|
|
178
|
+
* The tiered planner: without tiers ONE `segmented-reduce` dispatch with TIER 0 over [0, n); with tiers TIER 2 over
|
|
179
|
+
* [0, hiEnd) (one workgroup per row), then TIER 1 over [hiEnd, midEnd) (WG / 32 rows per workgroup), then TIER 0
|
|
180
|
+
* over [midEnd, n), each with its own RangeParams record and the perm binding.
|
|
181
|
+
*/
|
|
182
|
+
class TieredPlanner implements SegmentedReducePlanner {
|
|
146
183
|
private readonly scope: ReduceScope;
|
|
147
|
-
private readonly
|
|
184
|
+
private readonly tiers: readonly TierDispatch[];
|
|
185
|
+
private readonly fill: Kernel;
|
|
186
|
+
private readonly perm: Binding | null;
|
|
148
187
|
private readonly hasWeights: boolean;
|
|
149
188
|
private readonly accumulate: boolean;
|
|
189
|
+
private readonly identityWord: number;
|
|
190
|
+
private dispatches = 0;
|
|
150
191
|
|
|
151
192
|
/**
|
|
152
|
-
* Wraps
|
|
153
|
-
* @param scope - the scope the
|
|
154
|
-
* @param
|
|
155
|
-
* @param
|
|
193
|
+
* Wraps the compiled tier pipelines with the pattern they were compiled for.
|
|
194
|
+
* @param scope - the scope the pipelines were prepared in
|
|
195
|
+
* @param tiers - the compiled tiers in dispatch order (TIER 2, 1, 0; only the non-empty ones)
|
|
196
|
+
* @param fill - the `fill` pipeline of the identity pre-fill of a windowed core
|
|
197
|
+
* @param perm - the degreeOrder permutation binding (null: USE_PERM false, rows are node indices)
|
|
198
|
+
* @param hasWeights - the HAS_WEIGHTS the pipelines were compiled with
|
|
156
199
|
* @param accumulate - whether record() combines into out instead of overwriting
|
|
200
|
+
* @param identityWord - the u32 bits of the operator's identity element (identityFillWord)
|
|
157
201
|
*/
|
|
158
|
-
constructor(
|
|
202
|
+
constructor(
|
|
203
|
+
scope: ReduceScope,
|
|
204
|
+
tiers: readonly TierDispatch[],
|
|
205
|
+
fill: Kernel,
|
|
206
|
+
perm: Binding | null,
|
|
207
|
+
hasWeights: boolean,
|
|
208
|
+
accumulate: boolean,
|
|
209
|
+
identityWord: number,
|
|
210
|
+
) {
|
|
159
211
|
this.scope = scope;
|
|
160
|
-
this.
|
|
212
|
+
this.tiers = tiers;
|
|
213
|
+
this.fill = fill;
|
|
214
|
+
this.perm = perm;
|
|
161
215
|
this.hasWeights = hasWeights;
|
|
162
216
|
this.accumulate = accumulate;
|
|
217
|
+
this.identityWord = identityWord;
|
|
218
|
+
}
|
|
219
|
+
|
|
220
|
+
/**
|
|
221
|
+
* Dispatches the last record() issued.
|
|
222
|
+
* @returns 0 for n = 0, else the number of non-empty tier ranges
|
|
223
|
+
*/
|
|
224
|
+
get lastDispatches(): number {
|
|
225
|
+
return this.dispatches;
|
|
163
226
|
}
|
|
164
227
|
|
|
165
228
|
/**
|
|
166
|
-
* Records the
|
|
167
|
-
* ever created).
|
|
229
|
+
* Records the dispatches: every tier's rows over the arcs [0, arcCount), each with its own params record;
|
|
230
|
+
* nothing for n = 0 (no zero-length binding is ever created). On a windowed core: the identity fill (unless
|
|
231
|
+
* accumulating into the caller's out), then the dispatches once per window (PD-8).
|
|
168
232
|
* @param pass - the pass to record into
|
|
169
|
-
* @param core - a core with the SAME weights pattern as the one prepared (any snapshot
|
|
233
|
+
* @param core - a core with the SAME weights pattern as the one prepared (any snapshot; with tiers, the one the
|
|
234
|
+
* tiers were built from)
|
|
170
235
|
* @param out - at least 4n bytes of f32
|
|
171
236
|
*/
|
|
172
237
|
record(pass: GPUComputePassEncoder, core: CoreBinding, out: Binding): void {
|
|
173
|
-
assertNotWindowed(core, "segmentedReduce");
|
|
174
238
|
if ((core.weights !== null) !== this.hasWeights) {
|
|
175
239
|
throw new WebGpuGraphError(
|
|
176
240
|
"E_INVALID_ARGUMENT",
|
|
@@ -183,10 +247,6 @@ class ThreadPerRowPlanner implements SegmentedReducePlanner {
|
|
|
183
247
|
);
|
|
184
248
|
}
|
|
185
249
|
const n = rowCountOf(core, "segmentedReduce");
|
|
186
|
-
if (n === 0) {
|
|
187
|
-
return;
|
|
188
|
-
}
|
|
189
|
-
const arcCount = core.colIdx === null ? 0 : core.colIdx.size / 4;
|
|
190
250
|
if (out.size < 4 * n) {
|
|
191
251
|
throw new WebGpuGraphError(
|
|
192
252
|
"E_INVALID_ARGUMENT",
|
|
@@ -198,24 +258,96 @@ class ThreadPerRowPlanner implements SegmentedReducePlanner {
|
|
|
198
258
|
},
|
|
199
259
|
);
|
|
200
260
|
}
|
|
201
|
-
|
|
202
|
-
|
|
203
|
-
|
|
204
|
-
|
|
205
|
-
|
|
206
|
-
|
|
207
|
-
|
|
261
|
+
this.dispatches = 0;
|
|
262
|
+
if (n === 0) {
|
|
263
|
+
return;
|
|
264
|
+
}
|
|
265
|
+
if (core.windows === null) {
|
|
266
|
+
this.recordWindow(pass, core, out, n, null, 0);
|
|
267
|
+
return;
|
|
268
|
+
}
|
|
269
|
+
if (!this.accumulate) {
|
|
270
|
+
const params = this.scope.params(FILL_PARAMS, { count: n, value: this.identityWord, mode: 0 });
|
|
271
|
+
const bound = this.fill.bind({ dst: out, P: params.binding });
|
|
272
|
+
this.fill.dispatch(pass, bound, plan1d(n, this.scope.workgroupSize, this.scope.caps), [params.offset]);
|
|
273
|
+
this.dispatches++;
|
|
274
|
+
}
|
|
275
|
+
const { windows } = core;
|
|
276
|
+
windows.forEach((w, k) => {
|
|
277
|
+
const windowed: CoreBinding = {
|
|
278
|
+
...core,
|
|
279
|
+
colIdx: windowBinding(core, "colIdx", w),
|
|
280
|
+
weights: core.weights === null ? null : windowBinding(core, "weights", w),
|
|
281
|
+
};
|
|
282
|
+
this.recordWindow(pass, windowed, out, n, w, k + 1 < windows.length ? windows[k + 1].start : w.end);
|
|
208
283
|
});
|
|
209
|
-
|
|
210
|
-
|
|
284
|
+
}
|
|
285
|
+
|
|
286
|
+
/**
|
|
287
|
+
* The tier dispatches over one arc window (or the whole core when `w` is null): every non-empty tier over its
|
|
288
|
+
* full range, except the untiered TIER 0 dispatch, which covers exactly the window's rows. Consecutive windows
|
|
289
|
+
* overlap by up to ARC_WINDOW_ALIGN - 1 arcs (the next window opens at the aligned-down end of this one), and a
|
|
290
|
+
* tier dispatch visits EVERY row, so the tiers fold [w.start, nextStart) and leave the overlap to the next
|
|
291
|
+
* window; the untiered dispatch keeps [w.start, w.end): its rows [rowFirst, rowLast] are exactly the rows whose
|
|
292
|
+
* arcs end inside this window, and the next window's rows start after them.
|
|
293
|
+
* @param pass - the pass to record into
|
|
294
|
+
* @param core - the core with the window's colIdx / weights bound
|
|
295
|
+
* @param out - the output binding
|
|
296
|
+
* @param n - the row count
|
|
297
|
+
* @param w - the window, or null for the single whole-core dispatch
|
|
298
|
+
* @param nextStart - the next window's first arc (w.end for the last window)
|
|
299
|
+
*/
|
|
300
|
+
private recordWindow(
|
|
301
|
+
pass: GPUComputePassEncoder,
|
|
302
|
+
core: CoreBinding,
|
|
303
|
+
out: Binding,
|
|
304
|
+
n: number,
|
|
305
|
+
w: ArcWindow | null,
|
|
306
|
+
nextStart: number,
|
|
307
|
+
): void {
|
|
308
|
+
const wg = this.scope.workgroupSize;
|
|
309
|
+
for (const { tier, kernel, start, end } of this.tiers) {
|
|
310
|
+
let first = start;
|
|
311
|
+
let last = tier === 0 ? n : end;
|
|
312
|
+
let arcEnd = w === null ? arcCountOf(core) : Math.min(w.end, nextStart);
|
|
313
|
+
if (w !== null && this.perm === null) {
|
|
314
|
+
first = w.rowFirst;
|
|
315
|
+
last = w.rowLast + 1;
|
|
316
|
+
arcEnd = w.end;
|
|
317
|
+
}
|
|
318
|
+
const rows = last - first;
|
|
319
|
+
if (rows <= 0) {
|
|
320
|
+
continue;
|
|
321
|
+
}
|
|
322
|
+
const params = this.scope.params(RANGE_PARAMS, {
|
|
323
|
+
start: first,
|
|
324
|
+
end: last,
|
|
325
|
+
arcBase: w === null ? 0 : w.start,
|
|
326
|
+
arcEnd,
|
|
327
|
+
accumulate: w !== null || this.accumulate ? 1 : 0,
|
|
328
|
+
n,
|
|
329
|
+
});
|
|
330
|
+
const bound = kernel.bind({ ...graphBindings(core, this.perm), out, P: params.binding });
|
|
331
|
+
let rowsPerGroup = wg;
|
|
332
|
+
if (tier === 2) {
|
|
333
|
+
rowsPerGroup = 1;
|
|
334
|
+
} else if (tier === 1) {
|
|
335
|
+
rowsPerGroup = wg / MID_TIER_LANES;
|
|
336
|
+
}
|
|
337
|
+
kernel.dispatch(pass, bound, plan1d(rows, rowsPerGroup, this.scope.caps), [params.offset]);
|
|
338
|
+
this.dispatches++;
|
|
339
|
+
}
|
|
211
340
|
}
|
|
212
341
|
}
|
|
213
342
|
|
|
214
343
|
/**
|
|
215
|
-
* Prepares the
|
|
344
|
+
* Prepares the tier pipelines for a snapshot's dummy pattern (USE_PERM = tiers !== null, HAS_WEIGHTS) and snippet:
|
|
345
|
+
* TIER 0 always; TIER 1 iff a row of degree 32..1023 exists (segmentOffsets[2] > segmentOffsets[1]); TIER 2 iff a
|
|
346
|
+
* row of degree >= 1024 exists (segmentOffsets[1] > 0). The mid tier folds 32 lanes per row, so a device whose
|
|
347
|
+
* workgroup size is below 32 is E_UNSUPPORTED { feature: "segmentedReduce.tiers" }.
|
|
216
348
|
* @param scope - the reduce scope (pipelines, pool, params writer)
|
|
217
|
-
* @param core - the core whose weights pattern selects HAS_WEIGHTS
|
|
218
|
-
* @param options - operator, snippet, tiers (
|
|
349
|
+
* @param core - the core whose weights pattern selects HAS_WEIGHTS
|
|
350
|
+
* @param options - operator, snippet, tiers (null: the single thread-per-row dispatch), accumulate
|
|
219
351
|
* @returns the planner
|
|
220
352
|
*/
|
|
221
353
|
export async function prepareSegmentedReduce(
|
|
@@ -223,16 +355,42 @@ export async function prepareSegmentedReduce(
|
|
|
223
355
|
core: CoreBinding,
|
|
224
356
|
options: SegmentedReduceOptions,
|
|
225
357
|
): Promise<SegmentedReducePlanner> {
|
|
226
|
-
if (options.tiers !== null) {
|
|
227
|
-
throw new WebGpuGraphError("E_UNSUPPORTED", "segmentedReduce: the degree tiers land at P4; pass tiers: null", {
|
|
228
|
-
feature: "segmentedReduce.tiers",
|
|
229
|
-
});
|
|
230
|
-
}
|
|
231
|
-
assertNotWindowed(core, "segmentedReduce");
|
|
232
358
|
const op = opCode(options.op);
|
|
233
359
|
validateValueSnippet(options.valueSnippet);
|
|
234
|
-
const
|
|
235
|
-
const
|
|
236
|
-
|
|
237
|
-
|
|
360
|
+
const { tiers } = options;
|
|
361
|
+
const perm = tiers?.perm ?? null;
|
|
362
|
+
if (tiers !== null && scope.workgroupSize < MID_TIER_LANES) {
|
|
363
|
+
throw new WebGpuGraphError(
|
|
364
|
+
"E_UNSUPPORTED",
|
|
365
|
+
`segmentedReduce: the mid tier folds ${MID_TIER_LANES} lanes per row, more than the workgroup size ${scope.workgroupSize}`,
|
|
366
|
+
{ feature: "segmentedReduce.tiers" },
|
|
367
|
+
);
|
|
368
|
+
}
|
|
369
|
+
const ranges: readonly { readonly tier: 0 | 1 | 2; readonly start: number; readonly end: number }[] =
|
|
370
|
+
tiers === null
|
|
371
|
+
? [{ tier: 0, start: 0, end: rowCountOf(core, "segmentedReduce") }]
|
|
372
|
+
: [
|
|
373
|
+
{ tier: 2, start: 0, end: tiers.segmentOffsets[1] },
|
|
374
|
+
{ tier: 1, start: tiers.segmentOffsets[1], end: tiers.segmentOffsets[2] },
|
|
375
|
+
{ tier: 0, start: tiers.segmentOffsets[2], end: tiers.segmentOffsets[4] },
|
|
376
|
+
];
|
|
377
|
+
const compiled: TierDispatch[] = [];
|
|
378
|
+
for (const range of ranges) {
|
|
379
|
+
if (range.tier !== 0 && range.end <= range.start) {
|
|
380
|
+
continue;
|
|
381
|
+
}
|
|
382
|
+
const overrides = { ...graphOverrides(core, perm), OP: op, TIER: range.tier };
|
|
383
|
+
const spec = kernelSpec("segmented-reduce", overrides, { VALUE: options.valueSnippet });
|
|
384
|
+
compiled.push({ ...range, kernel: await scope.pipelines.kernel(spec) });
|
|
385
|
+
}
|
|
386
|
+
const fill = await scope.pipelines.kernel(kernelSpec("fill"));
|
|
387
|
+
return new TieredPlanner(
|
|
388
|
+
scope,
|
|
389
|
+
compiled,
|
|
390
|
+
fill,
|
|
391
|
+
perm,
|
|
392
|
+
core.weights !== null,
|
|
393
|
+
options.accumulate === true,
|
|
394
|
+
identityFillWord(options.op),
|
|
395
|
+
);
|
|
238
396
|
}
|
package/src/primitives/spmv.ts
CHANGED
|
@@ -1,7 +1,6 @@
|
|
|
1
1
|
/**
|
|
2
|
-
* The pull SpMV primitive of spec 6 row 9 / 8.2
|
|
3
|
-
*
|
|
4
|
-
* `weight * xNorm[nbr]` over its row's in-arcs (Kahan-compensated, f32) and writing
|
|
2
|
+
* The pull SpMV primitive of spec 6 row 9 / 8.2 (P7; M8b plan PD-1 / PD-2): the `spmv-pull` module over the rows of
|
|
3
|
+
* a REVERSE adjacency, each row folding `weight * xNorm[nbr]` over its in-arcs (a two-level f32 sum) and writing
|
|
5
4
|
* `rankOut[v] = beta * pv + alpha * (sum + danglingMass * pv)`, where `pv` is `personalization[v]` under
|
|
6
5
|
* HAS_PERSONALIZATION and the uniform `P.uniformP` otherwise, and `danglingMass` is `partials[0].danglingMass`
|
|
7
6
|
* under USE_DANGLING. PageRank sets alpha to the damping, beta to 1 - alpha and USE_DANGLING; HITS and eigenvector
|
|
@@ -9,20 +8,21 @@
|
|
|
9
8
|
*
|
|
10
9
|
* It is its own registry entry rather than a `segmentedReduce` VALUE snippet (PD-1): the snippet vocabulary is
|
|
11
10
|
* `row, arc, nbr, weight, v` and cannot read `xNorm[nbr]`, and the design's binding table gives the kernel eight
|
|
12
|
-
* storage bindings of its own.
|
|
13
|
-
*
|
|
14
|
-
*
|
|
11
|
+
* storage bindings of its own. Without tiers it is ONE grid-stride dispatch (TIER 0, USE_PERM false: the perm slot
|
|
12
|
+
* carries the rowPtr dummy of graphBindings); with the in-degree tiers of `reverseDegreeOrder()` (P4-T5, PD-6) it
|
|
13
|
+
* is up to three dispatches over the permuted rows -- TIER 2 one workgroup per row over [0, hiEnd), TIER 1 32 lanes
|
|
14
|
+
* per row over [hiEnd, midEnd), TIER 0 grid-stride over [midEnd, n) -- each compiled only when its range is
|
|
15
|
+
* non-empty. The row and arc counts come from the core's binding sizes exactly as segmentedReduce derives them.
|
|
15
16
|
*/
|
|
16
17
|
|
|
17
18
|
import { WebGpuGraphError } from "../errors.js";
|
|
18
|
-
import { planGridStride } from "../kernel/dispatch.js";
|
|
19
|
+
import { plan1d, planGridStride } from "../kernel/dispatch.js";
|
|
19
20
|
import { type Kernel } from "../kernel/kernel.js";
|
|
20
21
|
import { graphBindings, graphOverrides, kernelSpec, SPMV_PARAMS } from "../kernels.js";
|
|
21
22
|
import { type CoreBinding } from "../memory/residency.js";
|
|
22
23
|
import { type Binding } from "../types/memory.js";
|
|
23
|
-
import { arcCountOf, assertNotWindowed, rowCountOf } from "./core-shape.js";
|
|
24
|
+
import { arcCountOf, assertNotWindowed, type DegreeTiers, MID_TIER_LANES, rowCountOf } from "./core-shape.js";
|
|
24
25
|
import { type ReduceScope } from "./reduce.js";
|
|
25
|
-
import { type DegreeTiers } from "./segmented-reduce.js";
|
|
26
26
|
|
|
27
27
|
/** The group-1 bindings of one pull: the pre-scaled input, the output, the personalization (null binds xNorm as the dummy) and the PrPartial block whose header carries danglingMass. */
|
|
28
28
|
export interface SpmvResources {
|
|
@@ -39,7 +39,7 @@ export interface SpmvCoefficients {
|
|
|
39
39
|
readonly uniformP: number;
|
|
40
40
|
}
|
|
41
41
|
|
|
42
|
-
/** Options of prepareSpmvPull: the two variant flags, the weights binding and the tiers (
|
|
42
|
+
/** Options of prepareSpmvPull: the two variant flags, the weights binding and the in-degree tiers (null: one grid-stride dispatch). */
|
|
43
43
|
export interface SpmvPullOptions {
|
|
44
44
|
readonly personalization: boolean;
|
|
45
45
|
readonly dangling: boolean;
|
|
@@ -48,47 +48,74 @@ export interface SpmvPullOptions {
|
|
|
48
48
|
readonly tiers: DegreeTiers | null;
|
|
49
49
|
}
|
|
50
50
|
|
|
51
|
-
/** A prepared pull: records
|
|
51
|
+
/** A prepared pull: records one grid-stride dispatch (no tiers) or up to three tier dispatches over the rows of a reverse core into a pass. */
|
|
52
52
|
export interface SpmvPullPlanner {
|
|
53
|
-
/** Records the
|
|
54
|
-
record(
|
|
55
|
-
|
|
53
|
+
/** Records the dispatches over rows [0, n) of `rev` writing rankOut[v] (f32) per row; nothing for n = 0. */
|
|
54
|
+
record(
|
|
55
|
+
pass: GPUComputePassEncoder,
|
|
56
|
+
rev: CoreBinding,
|
|
57
|
+
resources: SpmvResources,
|
|
58
|
+
coefficients: SpmvCoefficients,
|
|
59
|
+
): void;
|
|
60
|
+
/** Dispatches the last record() issued (0 for n = 0; 1 without tiers; 1 to 3 with them). */
|
|
56
61
|
readonly lastDispatches: number;
|
|
57
62
|
}
|
|
58
63
|
|
|
59
|
-
/**
|
|
64
|
+
/** One tier's compiled pipeline and the rows [start, end) of the permutation it covers per record() (a tier whose range is empty is not compiled). */
|
|
65
|
+
interface TierDispatch {
|
|
66
|
+
readonly tier: 0 | 1 | 2;
|
|
67
|
+
readonly kernel: Kernel;
|
|
68
|
+
readonly start: number;
|
|
69
|
+
readonly end: number;
|
|
70
|
+
}
|
|
71
|
+
|
|
72
|
+
/**
|
|
73
|
+
* The tiered planner: without tiers ONE `spmv-pull` dispatch, grid-stride over every row; with tiers TIER 2 over
|
|
74
|
+
* [0, hiEnd) (one workgroup per row), then TIER 1 over [hiEnd, midEnd) (WG / 32 rows per workgroup), then TIER 0
|
|
75
|
+
* grid-stride over [midEnd, n), each with its own SpmvParams record (`n` is the range END) and the perm binding.
|
|
76
|
+
*/
|
|
60
77
|
class SpmvPullPlannerImpl implements SpmvPullPlanner {
|
|
61
78
|
private readonly scope: ReduceScope;
|
|
62
|
-
private readonly
|
|
79
|
+
private readonly tiers: readonly TierDispatch[];
|
|
80
|
+
private readonly perm: Binding | null;
|
|
63
81
|
private readonly weights: Binding | null | undefined;
|
|
64
82
|
private dispatches = 0;
|
|
65
83
|
|
|
66
84
|
/**
|
|
67
|
-
* Wraps
|
|
68
|
-
* @param scope - the scope the
|
|
69
|
-
* @param
|
|
70
|
-
* @param
|
|
85
|
+
* Wraps the compiled tier pipelines with the choices they were compiled for.
|
|
86
|
+
* @param scope - the scope the pipelines were prepared in
|
|
87
|
+
* @param tiers - the compiled tiers in dispatch order (TIER 2, 1, 0; only the non-empty ones)
|
|
88
|
+
* @param perm - the reverseDegreeOrder permutation binding (null: USE_PERM false, rows are node indices)
|
|
89
|
+
* @param weights - the weights option the pipelines' HAS_WEIGHTS was derived from; record() binds the same way
|
|
71
90
|
*/
|
|
72
|
-
constructor(
|
|
91
|
+
constructor(
|
|
92
|
+
scope: ReduceScope,
|
|
93
|
+
tiers: readonly TierDispatch[],
|
|
94
|
+
perm: Binding | null,
|
|
95
|
+
weights: Binding | null | undefined,
|
|
96
|
+
) {
|
|
73
97
|
this.scope = scope;
|
|
74
|
-
this.
|
|
98
|
+
this.tiers = tiers;
|
|
99
|
+
this.perm = perm;
|
|
75
100
|
this.weights = weights;
|
|
76
101
|
}
|
|
77
102
|
|
|
78
103
|
/**
|
|
79
104
|
* Dispatches the last record() issued.
|
|
80
|
-
* @returns
|
|
105
|
+
* @returns 0 when the last record covered no rows, else the number of non-empty tier ranges
|
|
81
106
|
*/
|
|
82
107
|
get lastDispatches(): number {
|
|
83
108
|
return this.dispatches;
|
|
84
109
|
}
|
|
85
110
|
|
|
86
111
|
/**
|
|
87
|
-
* Records the
|
|
88
|
-
*
|
|
89
|
-
*
|
|
112
|
+
* Records the dispatches: every tier's rows over the arcs [0, arcCount), each with its own params record (TIER 0
|
|
113
|
+
* grid-stride through planGridStride, the tiers through plan1d); nothing for n = 0 (no zero-length binding is
|
|
114
|
+
* ever created). `personalization ?? xNorm` follows the group-0 dummy rule: both slots are storage-ro, so the
|
|
115
|
+
* aliasing check of Kernel.bind does not fire, and HAS_PERSONALIZATION false never reads it.
|
|
90
116
|
* @param pass - the pass to record into
|
|
91
|
-
* @param rev - the reverse core (any snapshot with the weights pattern the planner was prepared for
|
|
117
|
+
* @param rev - the reverse core (any snapshot with the weights pattern the planner was prepared for; with tiers,
|
|
118
|
+
* the one the tiers were built from)
|
|
92
119
|
* @param resources - xNorm, rankOut, personalization, partials
|
|
93
120
|
* @param coefficients - alpha, beta, uniformP
|
|
94
121
|
*/
|
|
@@ -99,39 +126,53 @@ class SpmvPullPlannerImpl implements SpmvPullPlanner {
|
|
|
99
126
|
coefficients: SpmvCoefficients,
|
|
100
127
|
): void {
|
|
101
128
|
const n = rowCountOf(rev, "spmvPull");
|
|
102
|
-
const
|
|
103
|
-
|
|
104
|
-
|
|
105
|
-
|
|
129
|
+
const wg = this.scope.workgroupSize;
|
|
130
|
+
this.dispatches = 0;
|
|
131
|
+
for (const { tier, kernel, start, end } of this.tiers) {
|
|
132
|
+
const last = tier === 0 ? n : end;
|
|
133
|
+
const rows = last - start;
|
|
134
|
+
if (rows <= 0) {
|
|
135
|
+
continue;
|
|
136
|
+
}
|
|
137
|
+
const plan =
|
|
138
|
+
tier === 0
|
|
139
|
+
? planGridStride(rows, wg, this.scope.caps)
|
|
140
|
+
: plan1d(rows, tier === 2 ? 1 : wg / MID_TIER_LANES, this.scope.caps);
|
|
141
|
+
if (plan.x === 0) {
|
|
142
|
+
continue;
|
|
143
|
+
}
|
|
144
|
+
const params = this.scope.params(SPMV_PARAMS, {
|
|
145
|
+
n: last,
|
|
146
|
+
arcBase: 0,
|
|
147
|
+
arcEnd: arcCountOf(rev),
|
|
148
|
+
stride: plan.stride ?? rows,
|
|
149
|
+
alpha: coefficients.alpha,
|
|
150
|
+
beta: coefficients.beta,
|
|
151
|
+
uniformP: coefficients.uniformP,
|
|
152
|
+
start,
|
|
153
|
+
});
|
|
154
|
+
const bound = kernel.bind({
|
|
155
|
+
...graphBindings(rev, this.perm, this.weights),
|
|
156
|
+
xNorm: resources.xNorm,
|
|
157
|
+
rankOut: resources.rankOut,
|
|
158
|
+
personalization: resources.personalization ?? resources.xNorm,
|
|
159
|
+
partials: resources.partials,
|
|
160
|
+
P: params.binding,
|
|
161
|
+
});
|
|
162
|
+
kernel.dispatch(pass, bound, plan, [params.offset]);
|
|
163
|
+
this.dispatches++;
|
|
106
164
|
}
|
|
107
|
-
const params = this.scope.params(SPMV_PARAMS, {
|
|
108
|
-
n,
|
|
109
|
-
arcBase: 0,
|
|
110
|
-
arcEnd: arcCountOf(rev),
|
|
111
|
-
stride: plan.stride ?? n,
|
|
112
|
-
alpha: coefficients.alpha,
|
|
113
|
-
beta: coefficients.beta,
|
|
114
|
-
uniformP: coefficients.uniformP,
|
|
115
|
-
pad0: 0,
|
|
116
|
-
});
|
|
117
|
-
const bound = this.kernel.bind({
|
|
118
|
-
...graphBindings(rev, null, this.weights),
|
|
119
|
-
xNorm: resources.xNorm,
|
|
120
|
-
rankOut: resources.rankOut,
|
|
121
|
-
personalization: resources.personalization ?? resources.xNorm,
|
|
122
|
-
partials: resources.partials,
|
|
123
|
-
P: params.binding,
|
|
124
|
-
});
|
|
125
|
-
this.kernel.dispatch(pass, bound, plan, [params.offset]);
|
|
126
|
-
this.dispatches = 1;
|
|
127
165
|
}
|
|
128
166
|
}
|
|
129
167
|
|
|
130
168
|
/**
|
|
131
|
-
* Prepares the
|
|
169
|
+
* Prepares the pull pipelines for a reverse core's weights pattern and the two variant flags: TIER 0 always
|
|
170
|
+
* (USE_PERM = tiers !== null); TIER 1 iff a row of in-degree 32..1023 exists (segmentOffsets[2] >
|
|
171
|
+
* segmentOffsets[1]); TIER 2 iff a row of in-degree >= 1024 exists (segmentOffsets[1] > 0). The mid tier folds 32
|
|
172
|
+
* lanes per row, so a device whose workgroup size is below 32 is E_UNSUPPORTED { feature: "spmvPull.tiers" }.
|
|
132
173
|
* @param scope - the reduce scope (pipelines, params writer)
|
|
133
|
-
* @param rev - the reverse core whose weights pattern selects HAS_WEIGHTS
|
|
134
|
-
* @param options - personalization, dangling, weights, tiers (
|
|
174
|
+
* @param rev - the reverse core whose weights pattern selects HAS_WEIGHTS
|
|
175
|
+
* @param options - personalization, dangling, weights, tiers (null: the single grid-stride dispatch)
|
|
135
176
|
* @returns the planner
|
|
136
177
|
*/
|
|
137
178
|
export async function prepareSpmvPull(
|
|
@@ -139,17 +180,36 @@ export async function prepareSpmvPull(
|
|
|
139
180
|
rev: CoreBinding,
|
|
140
181
|
options: SpmvPullOptions,
|
|
141
182
|
): Promise<SpmvPullPlanner> {
|
|
142
|
-
|
|
143
|
-
|
|
144
|
-
|
|
183
|
+
assertNotWindowed(rev, "spmvPull");
|
|
184
|
+
const { tiers } = options;
|
|
185
|
+
const perm = tiers?.perm ?? null;
|
|
186
|
+
if (tiers !== null && scope.workgroupSize < MID_TIER_LANES) {
|
|
187
|
+
throw new WebGpuGraphError(
|
|
188
|
+
"E_UNSUPPORTED",
|
|
189
|
+
`spmvPull: the mid tier folds ${MID_TIER_LANES} lanes per row, more than the workgroup size ${scope.workgroupSize}`,
|
|
190
|
+
{ feature: "spmvPull.tiers" },
|
|
191
|
+
);
|
|
192
|
+
}
|
|
193
|
+
const ranges: readonly { readonly tier: 0 | 1 | 2; readonly start: number; readonly end: number }[] =
|
|
194
|
+
tiers === null
|
|
195
|
+
? [{ tier: 0, start: 0, end: rowCountOf(rev, "spmvPull") }]
|
|
196
|
+
: [
|
|
197
|
+
{ tier: 2, start: 0, end: tiers.segmentOffsets[1] },
|
|
198
|
+
{ tier: 1, start: tiers.segmentOffsets[1], end: tiers.segmentOffsets[2] },
|
|
199
|
+
{ tier: 0, start: tiers.segmentOffsets[2], end: tiers.segmentOffsets[4] },
|
|
200
|
+
];
|
|
201
|
+
const compiled: TierDispatch[] = [];
|
|
202
|
+
for (const range of ranges) {
|
|
203
|
+
if (range.tier !== 0 && range.end <= range.start) {
|
|
204
|
+
continue;
|
|
205
|
+
}
|
|
206
|
+
const spec = kernelSpec("spmv-pull", {
|
|
207
|
+
...graphOverrides(rev, perm, options.weights),
|
|
208
|
+
HAS_PERSONALIZATION: options.personalization,
|
|
209
|
+
USE_DANGLING: options.dangling,
|
|
210
|
+
TIER: range.tier,
|
|
145
211
|
});
|
|
212
|
+
compiled.push({ ...range, kernel: await scope.pipelines.kernel(spec) });
|
|
146
213
|
}
|
|
147
|
-
|
|
148
|
-
const spec = kernelSpec("spmv-pull", {
|
|
149
|
-
...graphOverrides(rev, null, options.weights),
|
|
150
|
-
HAS_PERSONALIZATION: options.personalization,
|
|
151
|
-
USE_DANGLING: options.dangling,
|
|
152
|
-
});
|
|
153
|
-
const kernel = await scope.pipelines.kernel(spec);
|
|
154
|
-
return new SpmvPullPlannerImpl(scope, kernel, options.weights);
|
|
214
|
+
return new SpmvPullPlannerImpl(scope, compiled, perm, options.weights);
|
|
155
215
|
}
|