@graphty/webgpu-graph-algorithms 0.6.15 → 0.6.16

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.
Files changed (134) hide show
  1. package/README.md +5 -5
  2. package/dist/browser.js +1 -1
  3. package/dist/chunks/{context-oXphO3yj.js → context-VIvatQOo.js} +61 -34
  4. package/dist/chunks/context-VIvatQOo.js.map +1 -0
  5. package/dist/node.js +1 -1
  6. package/dist/src/accelerator.d.ts +2 -1
  7. package/dist/src/accelerator.d.ts.map +1 -1
  8. package/dist/src/accelerator.js +53 -1
  9. package/dist/src/accelerator.js.map +1 -1
  10. package/dist/src/algorithms/all-pairs.d.ts +41 -0
  11. package/dist/src/algorithms/all-pairs.d.ts.map +1 -0
  12. package/dist/src/algorithms/all-pairs.js +181 -0
  13. package/dist/src/algorithms/all-pairs.js.map +1 -0
  14. package/dist/src/algorithms/components.d.ts +9 -1
  15. package/dist/src/algorithms/components.d.ts.map +1 -1
  16. package/dist/src/algorithms/components.js +2 -2
  17. package/dist/src/algorithms/components.js.map +1 -1
  18. package/dist/src/algorithms/label-propagation.d.ts +31 -0
  19. package/dist/src/algorithms/label-propagation.d.ts.map +1 -0
  20. package/dist/src/algorithms/label-propagation.js +254 -0
  21. package/dist/src/algorithms/label-propagation.js.map +1 -0
  22. package/dist/src/algorithms/simple-symmetric.d.ts +88 -0
  23. package/dist/src/algorithms/simple-symmetric.d.ts.map +1 -0
  24. package/dist/src/algorithms/simple-symmetric.js +347 -0
  25. package/dist/src/algorithms/simple-symmetric.js.map +1 -0
  26. package/dist/src/algorithms/triangles.d.ts +34 -0
  27. package/dist/src/algorithms/triangles.d.ts.map +1 -0
  28. package/dist/src/algorithms/triangles.js +203 -0
  29. package/dist/src/algorithms/triangles.js.map +1 -0
  30. package/dist/src/constants.d.ts +45 -0
  31. package/dist/src/constants.d.ts.map +1 -1
  32. package/dist/src/constants.js +45 -0
  33. package/dist/src/constants.js.map +1 -1
  34. package/dist/src/index.d.ts +8 -1
  35. package/dist/src/index.d.ts.map +1 -1
  36. package/dist/src/index.js +7 -1
  37. package/dist/src/index.js.map +1 -1
  38. package/dist/src/kernel/prelude.d.ts.map +1 -1
  39. package/dist/src/kernel/prelude.js +4 -1
  40. package/dist/src/kernel/prelude.js.map +1 -1
  41. package/dist/src/kernels.d.ts +16 -4
  42. package/dist/src/kernels.d.ts.map +1 -1
  43. package/dist/src/kernels.js +223 -3
  44. package/dist/src/kernels.js.map +1 -1
  45. package/dist/src/memory/residency.js +15 -4
  46. package/dist/src/memory/residency.js.map +1 -1
  47. package/dist/src/primitives/coo-to-csr.d.ts +73 -0
  48. package/dist/src/primitives/coo-to-csr.d.ts.map +1 -0
  49. package/dist/src/primitives/coo-to-csr.js +183 -0
  50. package/dist/src/primitives/coo-to-csr.js.map +1 -0
  51. package/dist/src/primitives/group-by-key.d.ts +82 -0
  52. package/dist/src/primitives/group-by-key.d.ts.map +1 -0
  53. package/dist/src/primitives/group-by-key.js +147 -0
  54. package/dist/src/primitives/group-by-key.js.map +1 -0
  55. package/dist/src/types/accelerator.d.ts +10 -2
  56. package/dist/src/types/accelerator.d.ts.map +1 -1
  57. package/dist/src/types/all-pairs.d.ts +35 -0
  58. package/dist/src/types/all-pairs.d.ts.map +1 -0
  59. package/dist/src/types/all-pairs.js +8 -0
  60. package/dist/src/types/all-pairs.js.map +1 -0
  61. package/dist/src/types/community.d.ts +18 -0
  62. package/dist/src/types/community.d.ts.map +1 -0
  63. package/dist/src/types/community.js +5 -0
  64. package/dist/src/types/community.js.map +1 -0
  65. package/dist/src/types/structure.d.ts +27 -0
  66. package/dist/src/types/structure.d.ts.map +1 -0
  67. package/dist/src/types/structure.js +8 -0
  68. package/dist/src/types/structure.js.map +1 -0
  69. package/dist/src/wgsl/apsp-fw.wgsl.d.ts +25 -0
  70. package/dist/src/wgsl/apsp-fw.wgsl.d.ts.map +1 -0
  71. package/dist/src/wgsl/apsp-fw.wgsl.js +113 -0
  72. package/dist/src/wgsl/apsp-fw.wgsl.js.map +1 -0
  73. package/dist/src/wgsl/apsp-init.wgsl.d.ts +12 -0
  74. package/dist/src/wgsl/apsp-init.wgsl.d.ts.map +1 -0
  75. package/dist/src/wgsl/apsp-init.wgsl.js +26 -0
  76. package/dist/src/wgsl/apsp-init.wgsl.js.map +1 -0
  77. package/dist/src/wgsl/coo-emit.wgsl.d.ts +10 -0
  78. package/dist/src/wgsl/coo-emit.wgsl.d.ts.map +1 -0
  79. package/dist/src/wgsl/coo-emit.wgsl.js +33 -0
  80. package/dist/src/wgsl/coo-emit.wgsl.js.map +1 -0
  81. package/dist/src/wgsl/coo-scatter.wgsl.d.ts +15 -0
  82. package/dist/src/wgsl/coo-scatter.wgsl.d.ts.map +1 -0
  83. package/dist/src/wgsl/coo-scatter.wgsl.js +32 -0
  84. package/dist/src/wgsl/coo-scatter.wgsl.js.map +1 -0
  85. package/dist/src/wgsl/group-by-key-row.wgsl.d.ts +26 -0
  86. package/dist/src/wgsl/group-by-key-row.wgsl.d.ts.map +1 -0
  87. package/dist/src/wgsl/group-by-key-row.wgsl.js +146 -0
  88. package/dist/src/wgsl/group-by-key-row.wgsl.js.map +1 -0
  89. package/dist/src/wgsl/lpa-step.wgsl.d.ts +10 -0
  90. package/dist/src/wgsl/lpa-step.wgsl.d.ts.map +1 -0
  91. package/dist/src/wgsl/lpa-step.wgsl.js +35 -0
  92. package/dist/src/wgsl/lpa-step.wgsl.js.map +1 -0
  93. package/dist/src/wgsl/orient-flags.wgsl.d.ts +9 -0
  94. package/dist/src/wgsl/orient-flags.wgsl.d.ts.map +1 -0
  95. package/dist/src/wgsl/orient-flags.wgsl.js +21 -0
  96. package/dist/src/wgsl/orient-flags.wgsl.js.map +1 -0
  97. package/dist/src/wgsl/run-flags.wgsl.d.ts +8 -0
  98. package/dist/src/wgsl/run-flags.wgsl.d.ts.map +1 -0
  99. package/dist/src/wgsl/run-flags.wgsl.js +18 -0
  100. package/dist/src/wgsl/run-flags.wgsl.js.map +1 -0
  101. package/dist/src/wgsl/tri-intersect.wgsl.d.ts +11 -0
  102. package/dist/src/wgsl/tri-intersect.wgsl.d.ts.map +1 -0
  103. package/dist/src/wgsl/tri-intersect.wgsl.js +64 -0
  104. package/dist/src/wgsl/tri-intersect.wgsl.js.map +1 -0
  105. package/dist/webgpu-graph-algorithms.js +1819 -181
  106. package/dist/webgpu-graph-algorithms.js.map +1 -1
  107. package/package.json +2 -2
  108. package/src/accelerator.ts +56 -1
  109. package/src/algorithms/all-pairs.ts +228 -0
  110. package/src/algorithms/components.ts +2 -2
  111. package/src/algorithms/label-propagation.ts +280 -0
  112. package/src/algorithms/simple-symmetric.ts +409 -0
  113. package/src/algorithms/triangles.ts +240 -0
  114. package/src/constants.ts +45 -0
  115. package/src/index.ts +12 -1
  116. package/src/kernel/prelude.ts +6 -0
  117. package/src/kernels.ts +248 -6
  118. package/src/memory/residency.ts +15 -4
  119. package/src/primitives/coo-to-csr.ts +251 -0
  120. package/src/primitives/group-by-key.ts +209 -0
  121. package/src/types/accelerator.ts +10 -2
  122. package/src/types/all-pairs.ts +37 -0
  123. package/src/types/community.ts +18 -0
  124. package/src/types/structure.ts +28 -0
  125. package/src/wgsl/apsp-fw.wgsl.ts +112 -0
  126. package/src/wgsl/apsp-init.wgsl.ts +25 -0
  127. package/src/wgsl/coo-emit.wgsl.ts +32 -0
  128. package/src/wgsl/coo-scatter.wgsl.ts +31 -0
  129. package/src/wgsl/group-by-key-row.wgsl.ts +145 -0
  130. package/src/wgsl/lpa-step.wgsl.ts +34 -0
  131. package/src/wgsl/orient-flags.wgsl.ts +20 -0
  132. package/src/wgsl/run-flags.wgsl.ts +17 -0
  133. package/src/wgsl/tri-intersect.wgsl.ts +63 -0
  134. package/dist/chunks/context-oXphO3yj.js.map +0 -1
@@ -0,0 +1,409 @@
1
+ /**
2
+ * The simple symmetric graph (the P11 plan's PD-4): triangle counting and label propagation are defined over an
3
+ * undirected graph with no parallel edges and no self-loops, and a snapshot may be directed, may carry parallel
4
+ * edges and may carry self-loops. This builds that graph on the device from the snapshot's edge list:
5
+ *
6
+ * 1. `coo-emit` writes two arcs per logical edge -- the edge as declared and its reverse -- and drops a self-loop
7
+ * by writing both of its arcs as `INVALID_INDEX`, which sorts last.
8
+ * 2. Two stable radix sorts order the arcs by (source, target): by target first, then by source, the second sort's
9
+ * stability keeping the first's order among equal sources. Between and after them `coo-emit` re-reads the edge
10
+ * list through the sorted permutation (INDEXED), so no gather kernel is needed.
11
+ * 3. `run-flags` marks the first arc of every run of equal pairs; the scan of the flags counts the runs, which is
12
+ * the number of distinct arcs `U` -- read back once, because every later size depends on it.
13
+ * 4. Three `compact`s keep the first arc of every run (its source, its target and its position), and a segmented
14
+ * sum over the runs merges the weights of parallel arcs in input order (so the merged weight is bitwise
15
+ * reproducible).
16
+ * 5. `cooToCsr` in its sorted-input mode builds the rows, which come out sorted by target.
17
+ *
18
+ * The first submit ends at step 3; steps 4 and 5 are recorded into a batch that is RETURNED OPEN, so the caller
19
+ * records its own first work into the same submit and saves a synchronisation. The caller must check the
20
+ * precondition flag of `cooToCsr` with `assertBuildSorted` after that submit.
21
+ *
22
+ * Every intermediate array is `2 x edgeCount` words and is bound whole, so a snapshot whose arcs exceed one storage
23
+ * binding is refused with `E_TOO_LARGE` before any upload. A weighted build refuses, with `E_UNSUPPORTED`, a pair
24
+ * joined by more than PARALLEL_MERGE_LIMIT parallel edges (see the constant).
25
+ *
26
+ * ponytail: about nine of the arrays are alive at the peak (72 bytes per logical edge); reuse buffers across the
27
+ * steps if a 1M-node / 10M-edge build needs the room.
28
+ */
29
+
30
+ import { type GraphSnapshot } from "@graphty/graph-format";
31
+
32
+ import { PARALLEL_MERGE_LIMIT } from "../constants.js";
33
+ import { type GpuContext } from "../context.js";
34
+ import { WebGpuGraphError } from "../errors.js";
35
+ import { CommandBatch, type ReadbackRequest } from "../kernel/batch.js";
36
+ import { plan1d } from "../kernel/dispatch.js";
37
+ import { type Kernel } from "../kernel/kernel.js";
38
+ import { COO_PARAMS, FILL_PARAMS, kernelSpec } from "../kernels.js";
39
+ import { type CoreBinding } from "../memory/residency.js";
40
+ import { prepareCompact } from "../primitives/compact.js";
41
+ import { prepareCooToCsr } from "../primitives/coo-to-csr.js";
42
+ import { prepareRadixSort, type RadixBits, radixHistBytes } from "../primitives/radix-sort.js";
43
+ import { prepareScan } from "../primitives/scan.js";
44
+ import { prepareSegmentedReduce } from "../primitives/segmented-reduce.js";
45
+ import { type Binding } from "../types/memory.js";
46
+ import { type AlgorithmScope } from "./scope.js";
47
+
48
+ /**
49
+ * The simple symmetric graph on the device.
50
+ * @public
51
+ */
52
+ export interface SimpleGraph {
53
+ readonly n: number;
54
+ /** Arcs: twice the undirected edges of the simple graph. */
55
+ readonly arcCount: number;
56
+ /** `n + 1` words. */
57
+ readonly rowPtr: Binding;
58
+ /** `arcCount` words, each row sorted ascending; null when there are no arcs. */
59
+ readonly colIdx: Binding | null;
60
+ /** `arcCount` f32, the summed weight of the parallel edges each arc merges (1 per edge on an unweighted snapshot); null when there are no arcs or weights were not asked for. */
61
+ readonly weights: Binding | null;
62
+ /** The source of every arc (`arcCount` words, non-decreasing); null when there are no arcs. */
63
+ readonly src: Binding | null;
64
+ }
65
+
66
+ /**
67
+ * The open second batch of a build: record more work into it, submit it, then call `assertBuildSorted`.
68
+ * @public
69
+ */
70
+ export interface SimpleGraphBuild {
71
+ readonly graph: SimpleGraph;
72
+ readonly batch: CommandBatch;
73
+ /** The readback of the sorted-input flag of `cooToCsr`, already scheduled on `batch`. */
74
+ readonly flag: ReadbackRequest;
75
+ }
76
+
77
+ /**
78
+ * The key width that orders node indices below `n` with `INVALID_INDEX` after all of them: the low `bits` of
79
+ * `INVALID_INDEX` are all ones, so every node index must stay below `2^bits - 1`.
80
+ * @param n - the node count
81
+ * @returns 8, 16, 24 or 32
82
+ */
83
+ function sortBitsFor(n: number): RadixBits {
84
+ for (const bits of [8, 16, 24] as const) {
85
+ if (n < 2 ** bits) {
86
+ return bits;
87
+ }
88
+ }
89
+ return 32;
90
+ }
91
+
92
+ /**
93
+ * A whole-binding view of a scratch buffer of `words` u32 (at least one word, so a binding is never zero-length).
94
+ * @param scope - the scope the buffer comes from
95
+ * @param words - the words
96
+ * @param label - the scratch label
97
+ * @returns the binding
98
+ */
99
+ function scratchWords(scope: AlgorithmScope, words: number, label: string): Binding {
100
+ const size = 4 * Math.max(1, words);
101
+ return { buffer: scope.scratch(size, label), offset: 0, size, window: null };
102
+ }
103
+
104
+ /**
105
+ * Throws `E_TOO_LARGE` when a buffer the caller binds whole is larger than one storage binding of the device.
106
+ * @param ctx - the context whose limit applies
107
+ * @param needed - the bytes of the largest whole binding
108
+ * @param path - what does not fit
109
+ * @param algorithm - the algorithm, for the error
110
+ */
111
+ export function assertBindable(ctx: GpuContext, needed: number, path: string, algorithm: string): void {
112
+ const limit = ctx.caps.limits.maxStorageBufferBindingSize;
113
+ if (needed > limit) {
114
+ throw new WebGpuGraphError(
115
+ "E_TOO_LARGE",
116
+ `${algorithm}: ${path} needs ${needed} bytes in one binding, above the device limit of ${limit}`,
117
+ { needed, limit, path, algorithm },
118
+ );
119
+ }
120
+ }
121
+
122
+ /**
123
+ * Throws `E_UNSUPPORTED` when a weighted build would merge more than PARALLEL_MERGE_LIMIT parallel arcs into one: the
124
+ * edges joining one pair, in either direction. Only a pair of two vertices that each touch more than the limit can,
125
+ * so the exact count runs over those pairs alone and costs nothing on an ordinary graph.
126
+ * @param n - the node count
127
+ * @param src - the edge sources
128
+ * @param dst - the edge targets
129
+ * @param algorithm - the algorithm, for the error
130
+ */
131
+ function assertMergeable(n: number, src: ArrayLike<number>, dst: ArrayLike<number>, algorithm: string): void {
132
+ const incident = new Uint32Array(n);
133
+ for (let e = 0; e < src.length; e++) {
134
+ if (src[e] !== dst[e]) {
135
+ incident[src[e]]++;
136
+ incident[dst[e]]++;
137
+ }
138
+ }
139
+ const heavy = new Map<number, number>();
140
+ for (let v = 0; v < n; v++) {
141
+ if (incident[v] > PARALLEL_MERGE_LIMIT) {
142
+ heavy.set(v, heavy.size);
143
+ }
144
+ }
145
+ if (heavy.size < 2) {
146
+ return;
147
+ }
148
+ const pairs = new Uint32Array(heavy.size * heavy.size);
149
+ for (let e = 0; e < src.length; e++) {
150
+ const a = heavy.get(src[e]);
151
+ const b = heavy.get(dst[e]);
152
+ if (a === undefined || b === undefined || a === b) {
153
+ continue;
154
+ }
155
+ const key = Math.min(a, b) * heavy.size + Math.max(a, b);
156
+ if (++pairs[key] > PARALLEL_MERGE_LIMIT) {
157
+ throw new WebGpuGraphError(
158
+ "E_UNSUPPORTED",
159
+ `${algorithm}: nodes ${src[e]} and ${dst[e]} are joined by more than ${PARALLEL_MERGE_LIMIT} parallel edges, more than one weighted merge sums`,
160
+ { feature: `${algorithm}.parallelEdges`, hint: "run it unweighted, or merge the parallel edges first" },
161
+ );
162
+ }
163
+ }
164
+ }
165
+
166
+ /**
167
+ * Throws `E_VALIDATION` when the device found the build's arcs out of source order, which would mean a bug in the
168
+ * sorts: the rows would be scrambled and every intersection over them wrong.
169
+ * @param bytes - the batch's readback bytes
170
+ * @param build - the build whose flag was read
171
+ * @param label - the algorithm, for the error
172
+ */
173
+ export function assertBuildSorted(bytes: ArrayBuffer, build: SimpleGraphBuild, label: string): void {
174
+ if (new Uint32Array(bytes, build.flag.offset, 1)[0] !== 0) {
175
+ throw new WebGpuGraphError("E_VALIDATION", `${label}: the simple graph's arcs reached cooToCsr out of order`, {
176
+ label: `${label}/simple-graph`,
177
+ message: "the sorted-input precondition of cooToCsr failed on the device",
178
+ });
179
+ }
180
+ }
181
+
182
+ /**
183
+ * Builds the simple symmetric graph of `s` (see the file header): submits the first batch and waits for its
184
+ * four-byte readback, then records the rest into a batch it returns open.
185
+ * @param ctx - the context
186
+ * @param s - the snapshot (its edge list is uploaded through the residency, or found there)
187
+ * @param scope - the caller's scope: every buffer of the graph is its scratch
188
+ * @param withWeights - merge and keep the weights (label propagation) or not (triangle counting)
189
+ * @param label - the batch label prefix
190
+ * @returns the graph and the open batch
191
+ */
192
+ export async function buildSimpleSymmetric(
193
+ ctx: GpuContext,
194
+ s: GraphSnapshot,
195
+ scope: AlgorithmScope,
196
+ withWeights: boolean,
197
+ label: string,
198
+ ): Promise<SimpleGraphBuild> {
199
+ const n = s.nodeCount;
200
+ const wg = ctx.workgroupSize;
201
+ const cooToCsr = await prepareCooToCsr(scope);
202
+ const flag = scratchWords(scope, 1, "simple/flag");
203
+ const rowPtr = scratchWords(scope, n + 1, "simple/rowPtr");
204
+ const { edgeCount } = s;
205
+ const list = s.edgeList();
206
+ let selfLoops = 0;
207
+ for (let e = 0; e < edgeCount; e++) {
208
+ if (list.src[e] === list.dst[e]) {
209
+ selfLoops++;
210
+ }
211
+ }
212
+ const valid = 2 * (edgeCount - selfLoops);
213
+ const empty = (): SimpleGraphBuild => {
214
+ const batch = new CommandBatch(ctx, `${label}/simple-graph`);
215
+ cooToCsr.record(batch.pass("simple-graph"), {
216
+ src: flag,
217
+ dst: flag,
218
+ weights: null,
219
+ count: 0,
220
+ n,
221
+ sortedInput: true,
222
+ out: { rowPtr, colIdx: flag, weights: null, flag },
223
+ });
224
+ batch.endPass();
225
+ const graph: SimpleGraph = { n, arcCount: 0, rowPtr, colIdx: null, weights: null, src: null };
226
+ return { graph, batch, flag: batch.readback(flag.buffer, 0, 4) };
227
+ };
228
+ if (valid === 0) {
229
+ return empty();
230
+ }
231
+ const arcs = 2 * edgeCount;
232
+ const tableBytes = Math.max(4, radixHistBytes(arcs, wg));
233
+ assertBindable(ctx, Math.max(4 * (arcs + 1), tableBytes), "the simple graph's arc arrays", label);
234
+ if (withWeights) {
235
+ assertMergeable(n, list.src, list.dst, label);
236
+ }
237
+ // core() is what records (or re-records, after a release) the snapshot in the residency; the build reads only
238
+ // the edge list, so it asks for rowPtr alone
239
+ ctx.residency.core(s, ["rowPtr"]);
240
+ const edges = ctx.residency.view(s, "edgeList");
241
+ const edgeWeights = edges.bindings.weights ?? null;
242
+ const emitPlan = plan1d(arcs, wg, ctx.caps);
243
+ const emit = new Map<boolean, Kernel>();
244
+ for (const indexed of [false, true]) {
245
+ const spec = kernelSpec("coo-emit", { INDEXED: indexed, WEIGHTED: edgeWeights !== null });
246
+ emit.set(indexed, await ctx.pipelines.kernel(spec));
247
+ }
248
+ const fill = await ctx.pipelines.kernel(kernelSpec("fill"));
249
+ const runFlags = await ctx.pipelines.kernel(kernelSpec("run-flags"));
250
+ const sort = await prepareRadixSort(scope);
251
+ const scan = await prepareScan(scope);
252
+ const compact = await prepareCompact(scope);
253
+
254
+ const aSrc = scratchWords(scope, arcs, "simple/aSrc");
255
+ const aDst = scratchWords(scope, arcs, "simple/aDst");
256
+ const aW = scratchWords(scope, arcs, "simple/aW");
257
+ const vals = scratchWords(scope, arcs, "simple/vals");
258
+ const sKeys = scratchWords(scope, arcs, "simple/sortKeys");
259
+ const sVals = scratchWords(scope, arcs, "simple/sortVals");
260
+ const hist: Binding = {
261
+ buffer: scope.scratch(tableBytes, "simple/hist"),
262
+ offset: 0,
263
+ size: tableBytes,
264
+ window: null,
265
+ };
266
+ const offsets: Binding = {
267
+ buffer: scope.scratch(tableBytes, "simple/offsets"),
268
+ offset: 0,
269
+ size: tableBytes,
270
+ window: null,
271
+ };
272
+ const sortedSrc = scratchWords(scope, arcs, "simple/sortedSrc");
273
+ const sortedDst = scratchWords(scope, arcs, "simple/sortedDst");
274
+ const sortedW = scratchWords(scope, arcs, "simple/sortedW");
275
+ const flags = scratchWords(scope, arcs, "simple/flags");
276
+ const runIndex = scratchWords(scope, arcs, "simple/runIndex");
277
+ const dummy = scratchWords(scope, 1, "simple/dummy");
278
+ await ctx.allocator.check();
279
+
280
+ const recordEmit = (pass: GPUComputePassEncoder, order: Binding | null, out: readonly Binding[]): void => {
281
+ const kernel = emit.get(order !== null);
282
+ if (kernel === undefined) {
283
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "buildSimpleSymmetric: no coo-emit variant", {
284
+ argument: "order",
285
+ value: order === null ? "null" : "present",
286
+ });
287
+ }
288
+ const params = scope.params(COO_PARAMS, { count: arcs, pad0: 0, pad1: 0, pad2: 0 });
289
+ const bound = kernel.bind({
290
+ edgeSrc: edges.bindings.src,
291
+ edgeDst: edges.bindings.dst,
292
+ edgeWeight: edgeWeights ?? edges.bindings.src,
293
+ order: order ?? dummy,
294
+ outSrc: out[0],
295
+ outDst: out[1],
296
+ outWeight: out[2],
297
+ P: params.binding,
298
+ });
299
+ kernel.dispatch(pass, bound, emitPlan, [params.offset]);
300
+ };
301
+ const recordIota = (pass: GPUComputePassEncoder, dst: Binding, words: number): void => {
302
+ const params = scope.params(FILL_PARAMS, { count: words, value: 0, mode: 1, pad0: 0 });
303
+ fill.dispatch(pass, fill.bind({ dst, P: params.binding }), plan1d(words, wg, ctx.caps), [params.offset]);
304
+ };
305
+
306
+ // batch 1: emit, sort by target, re-emit in that order, sort by source, re-emit sorted, flag the runs, count them
307
+ const bits = sortBitsFor(n);
308
+ const first = new CommandBatch(ctx, `${label}/simple-sort`);
309
+ const pass = first.pass("simple-sort");
310
+ recordEmit(pass, null, [aSrc, aDst, aW]);
311
+ recordIota(pass, vals, arcs);
312
+ const byTarget = sort.record(pass, aDst, vals, arcs, bits, { keys: sKeys, vals: sVals, hist, offsets });
313
+ recordEmit(pass, byTarget.vals, [aSrc, aDst, aW]);
314
+ const free = byTarget.vals.buffer === vals.buffer ? sVals : vals;
315
+ const bySource = sort.record(pass, aSrc, byTarget.vals, arcs, bits, {
316
+ keys: byTarget.keys.buffer === aDst.buffer ? sKeys : aDst,
317
+ vals: free,
318
+ hist,
319
+ offsets,
320
+ });
321
+ recordEmit(pass, bySource.vals, [sortedSrc, sortedDst, sortedW]);
322
+ const params = scope.params(COO_PARAMS, { count: arcs, pad0: 0, pad1: 0, pad2: 0 });
323
+ runFlags.dispatch(pass, runFlags.bind({ keysA: sortedSrc, keysB: sortedDst, flags, P: params.binding }), emitPlan, [
324
+ params.offset,
325
+ ]);
326
+ const total = scan.record(pass, flags, arcs, runIndex);
327
+ first.endPass();
328
+ const totalRequest = first.readback(total.binding.buffer, total.binding.offset + 4 * total.index, 4);
329
+ scope.flush();
330
+ const back = await first.submit().readback;
331
+ ctx.assertReady();
332
+ const unique = new Uint32Array(back, totalRequest.offset, 1)[0];
333
+ if (unique === 0 || unique > valid) {
334
+ throw new WebGpuGraphError("E_VALIDATION", `${label}: ${unique} distinct arcs from ${valid} valid ones`, {
335
+ label: `${label}/simple-graph`,
336
+ message: "the run count of the simple graph build is outside [1, valid arcs]",
337
+ });
338
+ }
339
+
340
+ // batch 2 (returned open): keep the first arc of every run, merge the weights, build the rows
341
+ const uSrc = scratchWords(scope, arcs, "simple/src");
342
+ const uDst = scratchWords(scope, arcs, "simple/dst");
343
+ const runStart = scratchWords(scope, arcs + 1, "simple/runStart");
344
+ const counts = scratchWords(scope, 1, "simple/count");
345
+ const colIdx = scratchWords(scope, unique, "simple/colIdx");
346
+ const weights = withWeights ? scratchWords(scope, unique, "simple/weights") : null;
347
+ const merged = withWeights ? scratchWords(scope, unique, "simple/merged") : null;
348
+ // the terminator of the last run: every valid arc precedes the dropped ones, so the runs end at `valid`
349
+ ctx.device.queue.writeBuffer(runStart.buffer, 4 * unique, Uint32Array.of(valid));
350
+ const reduce =
351
+ merged === null
352
+ ? null
353
+ : await prepareSegmentedReduce(scope, runCore(runStart, unique, sortedDst, sortedW, valid), {
354
+ op: "sum",
355
+ valueSnippet: "v = weight;",
356
+ tiers: null,
357
+ });
358
+ const second = new CommandBatch(ctx, `${label}/simple-graph`);
359
+ const pass2 = second.pass("simple-graph");
360
+ const positions = vals;
361
+ recordIota(pass2, positions, arcs);
362
+ for (const [queue, out] of [
363
+ [sortedSrc, uSrc],
364
+ [sortedDst, uDst],
365
+ [positions, runStart],
366
+ ] as const) {
367
+ compact.record(pass2, { queue, flags, count: arcs, out, outCount: counts, outIndex: 0 });
368
+ }
369
+ if (reduce !== null && merged !== null) {
370
+ reduce.record(pass2, runCore(runStart, unique, sortedDst, sortedW, valid), merged);
371
+ }
372
+ cooToCsr.record(pass2, {
373
+ src: uSrc,
374
+ dst: uDst,
375
+ weights: merged,
376
+ count: unique,
377
+ n,
378
+ sortedInput: true,
379
+ out: { rowPtr, colIdx, weights, flag },
380
+ });
381
+ second.endPass();
382
+ const graph: SimpleGraph = { n, arcCount: unique, rowPtr, colIdx, weights, src: uSrc };
383
+ return { graph, batch: second, flag: second.readback(flag.buffer, 0, 4) };
384
+ }
385
+
386
+ /**
387
+ * The runs as a CSR the segmented sum walks: row r is run r, its arcs the sorted positions `[runStart[r],
388
+ * runStart[r + 1])`, its weights the sorted weights, so a row's sum is the merged weight in input order.
389
+ * @param runStart - the run starts plus the terminator
390
+ * @param runs - the run count
391
+ * @param sortedDst - the sorted targets (the fold reads a neighbour it ignores)
392
+ * @param sortedW - the sorted weights
393
+ * @param valid - the arcs the runs cover
394
+ * @returns the core binding
395
+ */
396
+ function runCore(runStart: Binding, runs: number, sortedDst: Binding, sortedW: Binding, valid: number): CoreBinding {
397
+ return {
398
+ serial: -1,
399
+ plan: "perArray",
400
+ rowPtr: { ...runStart, size: 4 * (runs + 1) },
401
+ colIdx: { ...sortedDst, size: 4 * valid },
402
+ weights: { ...sortedW, size: 4 * valid },
403
+ arcToEdge: null,
404
+ edgeToArc: null,
405
+ windows: null,
406
+ arcBuffers: null,
407
+ hasWeights: true,
408
+ };
409
+ }
@@ -0,0 +1,240 @@
1
+ /**
2
+ * Triangle counting with the local clustering coefficient and the transitivity (design 8.5, 3.3 line 806, design 17
3
+ * line 5067; the P11 plan's P11-T8). Over the simple symmetric graph of the snapshot (`buildSimpleSymmetric`, so a
4
+ * directed snapshot, parallel edges and self-loops all give the answer of the underlying simple undirected graph):
5
+ *
6
+ * 1. `orient-flags` keeps arc (u, v) when v is above u in the (degree, id) order, one arc per undirected edge;
7
+ * 2. two `compact`s keep the oriented arcs' sources and targets -- in order, so the oriented rows stay sorted -- and
8
+ * `cooToCsr` in its sorted-input mode builds the oriented rows (there are exactly `arcs / 2` of them, so no count
9
+ * is read back);
10
+ * 3. `tri-intersect` intersects the two oriented rows of every oriented arc and adds each triangle to the counts of
11
+ * its three vertices with u32 atomics -- exact, and independent of the schedule.
12
+ *
13
+ * All of it rides in the build's second submit, so a call is two synchronisations whatever the graph.
14
+ *
15
+ * PLAN DECISION (the P11 plan's DEP-P11-C): the coefficient and the transitivity are computed on the host from the
16
+ * per-node counts and the simple graph's `rowPtr`, both read back in the same submit, not by a `tri-coefficient`
17
+ * kernel. f32 division on the device is not correctly rounded (2.5 ULP on the reference card), and the transitivity's
18
+ * denominator -- the sum of `d (d - 1) / 2` -- overflows a u32 at a single node of degree 92,682; in f64 on the host
19
+ * both are exact to the f32 the result stores, and the `rowPtr` read costs the same bytes as the coefficient array a
20
+ * device epilogue would have read back instead.
21
+ */
22
+
23
+ import { type F32, type GraphSnapshot, type U32 } from "@graphty/graph-format";
24
+
25
+ import { type GpuContext } from "../context.js";
26
+ import { WebGpuGraphError } from "../errors.js";
27
+ import { plan1d } from "../kernel/dispatch.js";
28
+ import { COO_PARAMS, FILL_PARAMS, kernelSpec } from "../kernels.js";
29
+ import { prepareCompact } from "../primitives/compact.js";
30
+ import { prepareCooToCsr } from "../primitives/coo-to-csr.js";
31
+ import { assertDeviceComputes } from "../primitives/verify.js";
32
+ import { type Binding } from "../types/memory.js";
33
+ import { type GpuRunOptions } from "../types/run.js";
34
+ import { type GpuTriangleResult } from "../types/structure.js";
35
+ import { algorithmScope } from "./scope.js";
36
+ import { assertBuildSorted, buildSimpleSymmetric } from "./simple-symmetric.js";
37
+
38
+ const ALGORITHM = "triangleCount";
39
+ /** Params slots of the largest batch: the build's two sorts and scans plus the orientation's compactions and build. */
40
+ const RING_SLOTS = 1024;
41
+
42
+ /**
43
+ * Validates `options.dest` for the per-node counts.
44
+ * @param dest - the caller's destination array, if any
45
+ * @param n - the node count
46
+ * @returns the destination as a U32, or null when none was given
47
+ */
48
+ function checkDest(dest: Float32Array | Uint32Array | undefined, n: number): U32 | null {
49
+ if (dest === undefined) {
50
+ return null;
51
+ }
52
+ if (dest instanceof Uint32Array && dest.length === n && dest.buffer instanceof ArrayBuffer) {
53
+ return dest as U32;
54
+ }
55
+ throw new WebGpuGraphError(
56
+ "E_INVALID_ARGUMENT",
57
+ `${ALGORITHM}: dest must be a Uint32Array of length ${n} over an ArrayBuffer`,
58
+ {
59
+ argument: "dest",
60
+ value: `${dest.constructor.name}(${dest.length})`,
61
+ expected: `Uint32Array(${n}) over an ArrayBuffer`,
62
+ },
63
+ );
64
+ }
65
+
66
+ /**
67
+ * The coefficient, the total and the transitivity from the per-node counts and the simple graph's rows, in f64.
68
+ * @param perNode - the triangles of every node
69
+ * @param rowPtr - the simple graph's `n + 1` row offsets
70
+ * @returns the result
71
+ */
72
+ function epilogue(perNode: U32, rowPtr: Uint32Array): GpuTriangleResult {
73
+ const n = perNode.length;
74
+ const coefficient: F32 = new Float32Array(n);
75
+ let sum = 0;
76
+ let triples = 0;
77
+ for (let v = 0; v < n; v++) {
78
+ const d = rowPtr[v + 1] - rowPtr[v];
79
+ const t = perNode[v];
80
+ sum += t;
81
+ if (d >= 2) {
82
+ const pairs = (d * (d - 1)) / 2;
83
+ triples += pairs;
84
+ coefficient[v] = t / pairs;
85
+ }
86
+ }
87
+ return { perNode, total: sum / 3, coefficient, transitivity: triples === 0 ? 0 : sum / triples };
88
+ }
89
+
90
+ /**
91
+ * Counts the triangles of the simple undirected graph underlying `s` on the device (see the file header).
92
+ * @param ctx - the context whose device runs the kernels
93
+ * @param s - the snapshot (its edge list is uploaded through ctx.residency, or found there)
94
+ * @param options - dest (the per-node counts), signal, onProgress
95
+ * @returns the per-node counts, the total, the clustering coefficients and the transitivity
96
+ */
97
+ export async function triangleCount(
98
+ ctx: GpuContext,
99
+ s: GraphSnapshot,
100
+ options?: GpuRunOptions,
101
+ ): Promise<GpuTriangleResult> {
102
+ return await triangleCountWithSearch(ctx, s, 0, options);
103
+ }
104
+
105
+ /**
106
+ * `triangleCount` with the intersection path forced (the tier-agreement tests' seam): 0 chooses per arc, 1 always
107
+ * merges, 2 always binary-searches.
108
+ * @internal
109
+ * @param ctx - the context
110
+ * @param s - the snapshot
111
+ * @param search - the `SEARCH` override of `tri-intersect`
112
+ * @param options - dest, signal, onProgress
113
+ * @returns the result
114
+ */
115
+ export async function triangleCountWithSearch(
116
+ ctx: GpuContext,
117
+ s: GraphSnapshot,
118
+ search: 0 | 1 | 2,
119
+ options?: GpuRunOptions,
120
+ ): Promise<GpuTriangleResult> {
121
+ ctx.assertReady();
122
+ await assertDeviceComputes(ctx);
123
+ const n = s.nodeCount;
124
+ const dest = checkDest(options?.dest, n);
125
+ if (options?.signal?.aborted) {
126
+ throw new WebGpuGraphError("E_ABORTED", `${ALGORITHM}: the signal was aborted before any work started`, {});
127
+ }
128
+ if (n === 0) {
129
+ options?.onProgress?.(1, 1);
130
+ return epilogue(dest ?? new Uint32Array(0), new Uint32Array(1));
131
+ }
132
+ const scope = algorithmScope(ctx, ALGORITHM, RING_SLOTS);
133
+ try {
134
+ const build = await buildSimpleSymmetric(ctx, s, scope, false, ALGORITHM);
135
+ const { graph, batch } = build;
136
+ const wg = ctx.workgroupSize;
137
+ const countsBytes = 4 * n;
138
+ const counts: Binding = {
139
+ buffer: scope.scratch(countsBytes, "counts"),
140
+ offset: 0,
141
+ size: countsBytes,
142
+ window: null,
143
+ };
144
+ const fill = await ctx.pipelines.kernel(kernelSpec("fill"));
145
+ const pass = batch.pass("triangles");
146
+ let orientedFlag: Binding | null = null;
147
+ const zero = scope.params(FILL_PARAMS, { count: n, value: 0, mode: 0, pad0: 0 });
148
+ fill.dispatch(pass, fill.bind({ dst: counts, P: zero.binding }), plan1d(n, wg, ctx.caps), [zero.offset]);
149
+ if (graph.arcCount > 0 && graph.colIdx !== null && graph.src !== null) {
150
+ const oriented = graph.arcCount / 2;
151
+ const orient = await ctx.pipelines.kernel(kernelSpec("orient-flags"));
152
+ const intersect = await ctx.pipelines.kernel(kernelSpec("tri-intersect", { SEARCH: search }));
153
+ const compact = await prepareCompact(scope);
154
+ const cooToCsr = await prepareCooToCsr(scope);
155
+ const words = (count: number, label: string): Binding => ({
156
+ buffer: scope.scratch(4 * count, label),
157
+ offset: 0,
158
+ size: 4 * count,
159
+ window: null,
160
+ });
161
+ const flags = words(graph.arcCount, "orient/flags");
162
+ const oSrc = words(graph.arcCount, "orient/src");
163
+ const oDst = words(graph.arcCount, "orient/dst");
164
+ const oRowPtr = words(n + 1, "orient/rowPtr");
165
+ const oColIdx = words(oriented, "orient/colIdx");
166
+ const scratch = words(1, "orient/count");
167
+ const oFlag = words(1, "orient/flag");
168
+ const arcPlan = plan1d(graph.arcCount, wg, ctx.caps);
169
+ const p1 = scope.params(COO_PARAMS, { count: graph.arcCount, pad0: 0, pad1: 0, pad2: 0 });
170
+ orient.dispatch(
171
+ pass,
172
+ orient.bind({ rowPtr: graph.rowPtr, colIdx: graph.colIdx, src: graph.src, flags, P: p1.binding }),
173
+ arcPlan,
174
+ [p1.offset],
175
+ );
176
+ compact.record(pass, {
177
+ queue: graph.src,
178
+ flags,
179
+ count: graph.arcCount,
180
+ out: oSrc,
181
+ outCount: scratch,
182
+ outIndex: 0,
183
+ });
184
+ compact.record(pass, {
185
+ queue: graph.colIdx,
186
+ flags,
187
+ count: graph.arcCount,
188
+ out: oDst,
189
+ outCount: scratch,
190
+ outIndex: 0,
191
+ });
192
+ cooToCsr.record(pass, {
193
+ src: oSrc,
194
+ dst: oDst,
195
+ weights: null,
196
+ count: oriented,
197
+ n,
198
+ sortedInput: true,
199
+ out: { rowPtr: oRowPtr, colIdx: oColIdx, weights: null, flag: oFlag },
200
+ });
201
+ orientedFlag = oFlag;
202
+ const p2 = scope.params(COO_PARAMS, { count: oriented, pad0: 0, pad1: 0, pad2: 0 });
203
+ intersect.dispatch(
204
+ pass,
205
+ intersect.bind({ rowPtr: oRowPtr, colIdx: oColIdx, src: oSrc, counts, P: p2.binding }),
206
+ plan1d(oriented, wg, ctx.caps),
207
+ [p2.offset],
208
+ );
209
+ }
210
+ batch.endPass();
211
+ const countsRequest = batch.readback(counts.buffer, 0, countsBytes);
212
+ const rowsRequest = batch.readback(graph.rowPtr.buffer, 0, 4 * (n + 1));
213
+ const orientedRequest = orientedFlag === null ? null : batch.readback(orientedFlag.buffer, 0, 4);
214
+ scope.flush();
215
+ const submitted = batch.submit();
216
+ const bytes = await submitted.readback;
217
+ ctx.assertReady();
218
+ if (options?.signal?.aborted) {
219
+ throw new WebGpuGraphError("E_ABORTED", `${ALGORITHM}: the signal was aborted`, { batchId: submitted.id });
220
+ }
221
+ assertBuildSorted(bytes, build, ALGORITHM);
222
+ if (orientedRequest !== null && new Uint32Array(bytes, orientedRequest.offset, 1)[0] !== 0) {
223
+ throw new WebGpuGraphError(
224
+ "E_VALIDATION",
225
+ `${ALGORITHM}: the oriented arcs reached cooToCsr out of order`,
226
+ {
227
+ label: `${ALGORITHM}/oriented`,
228
+ message: "the sorted-input precondition of cooToCsr failed on the device",
229
+ },
230
+ );
231
+ }
232
+ const perNode = dest ?? new Uint32Array(n);
233
+ perNode.set(new Uint32Array(bytes, countsRequest.offset, n));
234
+ const rowPtr = new Uint32Array(bytes, rowsRequest.offset, n + 1);
235
+ options?.onProgress?.(1, 1);
236
+ return epilogue(perNode, rowPtr);
237
+ } finally {
238
+ scope.dispose();
239
+ }
240
+ }