@graphty/webgpu-graph-algorithms 0.5.0 → 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.
Files changed (264) hide show
  1. package/README.md +104 -52
  2. package/dist/browser.js +1 -1
  3. package/dist/chunks/{context-CRbw2Wyo.js → context-BXqgCifx.js} +225 -33
  4. package/dist/chunks/context-BXqgCifx.js.map +1 -0
  5. package/dist/node.js +1 -1
  6. package/dist/src/accelerator.d.ts +12 -10
  7. package/dist/src/accelerator.d.ts.map +1 -1
  8. package/dist/src/accelerator.js +32 -10
  9. package/dist/src/accelerator.js.map +1 -1
  10. package/dist/src/algorithms/components.d.ts.map +1 -1
  11. package/dist/src/algorithms/components.js +12 -13
  12. package/dist/src/algorithms/components.js.map +1 -1
  13. package/dist/src/algorithms/degree.d.ts +6 -8
  14. package/dist/src/algorithms/degree.d.ts.map +1 -1
  15. package/dist/src/algorithms/degree.js +58 -35
  16. package/dist/src/algorithms/degree.js.map +1 -1
  17. package/dist/src/algorithms/pagerank.d.ts.map +1 -1
  18. package/dist/src/algorithms/pagerank.js +16 -14
  19. package/dist/src/algorithms/pagerank.js.map +1 -1
  20. package/dist/src/algorithms/power-iteration.d.ts +2 -2
  21. package/dist/src/algorithms/power-iteration.d.ts.map +1 -1
  22. package/dist/src/algorithms/power-iteration.js +17 -14
  23. package/dist/src/algorithms/power-iteration.js.map +1 -1
  24. package/dist/src/constants.d.ts +85 -8
  25. package/dist/src/constants.d.ts.map +1 -1
  26. package/dist/src/constants.js +85 -8
  27. package/dist/src/constants.js.map +1 -1
  28. package/dist/src/errors.d.ts +3 -2
  29. package/dist/src/errors.d.ts.map +1 -1
  30. package/dist/src/errors.js +2 -1
  31. package/dist/src/errors.js.map +1 -1
  32. package/dist/src/index.d.ts +10 -5
  33. package/dist/src/index.d.ts.map +1 -1
  34. package/dist/src/index.js +14 -5
  35. package/dist/src/index.js.map +1 -1
  36. package/dist/src/kernel/dispatch.d.ts +8 -3
  37. package/dist/src/kernel/dispatch.d.ts.map +1 -1
  38. package/dist/src/kernel/dispatch.js +18 -7
  39. package/dist/src/kernel/dispatch.js.map +1 -1
  40. package/dist/src/kernel/kernel.d.ts +30 -1
  41. package/dist/src/kernel/kernel.d.ts.map +1 -1
  42. package/dist/src/kernel/kernel.js +49 -5
  43. package/dist/src/kernel/kernel.js.map +1 -1
  44. package/dist/src/kernel/prelude.d.ts.map +1 -1
  45. package/dist/src/kernel/prelude.js +9 -1
  46. package/dist/src/kernel/prelude.js.map +1 -1
  47. package/dist/src/kernel/profiler.d.ts +15 -3
  48. package/dist/src/kernel/profiler.d.ts.map +1 -1
  49. package/dist/src/kernel/profiler.js +27 -4
  50. package/dist/src/kernel/profiler.js.map +1 -1
  51. package/dist/src/kernels.d.ts +18 -8
  52. package/dist/src/kernels.d.ts.map +1 -1
  53. package/dist/src/kernels.js +345 -22
  54. package/dist/src/kernels.js.map +1 -1
  55. package/dist/src/layouts/calibrate.d.ts +51 -0
  56. package/dist/src/layouts/calibrate.d.ts.map +1 -0
  57. package/dist/src/layouts/calibrate.js +172 -0
  58. package/dist/src/layouts/calibrate.js.map +1 -0
  59. package/dist/src/layouts/force-simulation.d.ts +42 -5
  60. package/dist/src/layouts/force-simulation.d.ts.map +1 -1
  61. package/dist/src/layouts/force-simulation.js +84 -22
  62. package/dist/src/layouts/force-simulation.js.map +1 -1
  63. package/dist/src/layouts/forceatlas2.d.ts +107 -38
  64. package/dist/src/layouts/forceatlas2.d.ts.map +1 -1
  65. package/dist/src/layouts/forceatlas2.js +297 -290
  66. package/dist/src/layouts/forceatlas2.js.map +1 -1
  67. package/dist/src/layouts/fruchterman-reingold.d.ts +241 -0
  68. package/dist/src/layouts/fruchterman-reingold.d.ts.map +1 -0
  69. package/dist/src/layouts/fruchterman-reingold.js +739 -0
  70. package/dist/src/layouts/fruchterman-reingold.js.map +1 -0
  71. package/dist/src/layouts/model-common.d.ts +140 -0
  72. package/dist/src/layouts/model-common.d.ts.map +1 -0
  73. package/dist/src/layouts/model-common.js +269 -0
  74. package/dist/src/layouts/model-common.js.map +1 -0
  75. package/dist/src/layouts/repulsion-grid.d.ts +152 -0
  76. package/dist/src/layouts/repulsion-grid.d.ts.map +1 -0
  77. package/dist/src/layouts/repulsion-grid.js +318 -0
  78. package/dist/src/layouts/repulsion-grid.js.map +1 -0
  79. package/dist/src/layouts/spring-electrical.d.ts +224 -0
  80. package/dist/src/layouts/spring-electrical.d.ts.map +1 -0
  81. package/dist/src/layouts/spring-electrical.js +665 -0
  82. package/dist/src/layouts/spring-electrical.js.map +1 -0
  83. package/dist/src/memory/residency.d.ts +6 -2
  84. package/dist/src/memory/residency.d.ts.map +1 -1
  85. package/dist/src/memory/residency.js +84 -14
  86. package/dist/src/memory/residency.js.map +1 -1
  87. package/dist/src/primitives/core-shape.d.ts +38 -2
  88. package/dist/src/primitives/core-shape.d.ts.map +1 -1
  89. package/dist/src/primitives/core-shape.js +71 -3
  90. package/dist/src/primitives/core-shape.js.map +1 -1
  91. package/dist/src/primitives/grid-pyramid.d.ts +71 -0
  92. package/dist/src/primitives/grid-pyramid.d.ts.map +1 -0
  93. package/dist/src/primitives/grid-pyramid.js +143 -0
  94. package/dist/src/primitives/grid-pyramid.js.map +1 -0
  95. package/dist/src/primitives/grid.d.ts +118 -0
  96. package/dist/src/primitives/grid.d.ts.map +1 -0
  97. package/dist/src/primitives/grid.js +225 -0
  98. package/dist/src/primitives/grid.js.map +1 -0
  99. package/dist/src/primitives/histogram.d.ts +67 -0
  100. package/dist/src/primitives/histogram.d.ts.map +1 -0
  101. package/dist/src/primitives/histogram.js +190 -0
  102. package/dist/src/primitives/histogram.js.map +1 -0
  103. package/dist/src/primitives/radix-sort.d.ts +75 -0
  104. package/dist/src/primitives/radix-sort.d.ts.map +1 -0
  105. package/dist/src/primitives/radix-sort.js +168 -0
  106. package/dist/src/primitives/radix-sort.js.map +1 -0
  107. package/dist/src/primitives/scan.d.ts +44 -0
  108. package/dist/src/primitives/scan.d.ts.map +1 -0
  109. package/dist/src/primitives/scan.js +151 -0
  110. package/dist/src/primitives/scan.js.map +1 -0
  111. package/dist/src/primitives/segmented-reduce.d.ts +25 -17
  112. package/dist/src/primitives/segmented-reduce.d.ts.map +1 -1
  113. package/dist/src/primitives/segmented-reduce.js +166 -47
  114. package/dist/src/primitives/segmented-reduce.js.map +1 -1
  115. package/dist/src/primitives/spmv.d.ts +18 -14
  116. package/dist/src/primitives/spmv.d.ts.map +1 -1
  117. package/dist/src/primitives/spmv.js +94 -58
  118. package/dist/src/primitives/spmv.js.map +1 -1
  119. package/dist/src/primitives/verify.d.ts +49 -0
  120. package/dist/src/primitives/verify.d.ts.map +1 -0
  121. package/dist/src/primitives/verify.js +229 -0
  122. package/dist/src/primitives/verify.js.map +1 -0
  123. package/dist/src/types/accelerator.d.ts +7 -3
  124. package/dist/src/types/accelerator.d.ts.map +1 -1
  125. package/dist/src/types/context.d.ts +53 -0
  126. package/dist/src/types/context.d.ts.map +1 -1
  127. package/dist/src/types/layout.d.ts +52 -0
  128. package/dist/src/types/layout.d.ts.map +1 -1
  129. package/dist/src/types/options.d.ts +43 -1
  130. package/dist/src/types/options.d.ts.map +1 -1
  131. package/dist/src/wgsl/counting-scatter.wgsl.d.ts +8 -0
  132. package/dist/src/wgsl/counting-scatter.wgsl.d.ts.map +1 -0
  133. package/dist/src/wgsl/counting-scatter.wgsl.js +17 -0
  134. package/dist/src/wgsl/counting-scatter.wgsl.js.map +1 -0
  135. package/dist/src/wgsl/fa2-attraction.wgsl.d.ts +23 -8
  136. package/dist/src/wgsl/fa2-attraction.wgsl.d.ts.map +1 -1
  137. package/dist/src/wgsl/fa2-attraction.wgsl.js +100 -17
  138. package/dist/src/wgsl/fa2-attraction.wgsl.js.map +1 -1
  139. package/dist/src/wgsl/fa2-integrate.wgsl.d.ts +7 -2
  140. package/dist/src/wgsl/fa2-integrate.wgsl.d.ts.map +1 -1
  141. package/dist/src/wgsl/fa2-integrate.wgsl.js +28 -2
  142. package/dist/src/wgsl/fa2-integrate.wgsl.js.map +1 -1
  143. package/dist/src/wgsl/fa2-repulsion-exact.wgsl.d.ts +4 -2
  144. package/dist/src/wgsl/fa2-repulsion-exact.wgsl.d.ts.map +1 -1
  145. package/dist/src/wgsl/fa2-repulsion-exact.wgsl.js +14 -5
  146. package/dist/src/wgsl/fa2-repulsion-exact.wgsl.js.map +1 -1
  147. package/dist/src/wgsl/fa2-stats-finalize.wgsl.d.ts +12 -1
  148. package/dist/src/wgsl/fa2-stats-finalize.wgsl.d.ts.map +1 -1
  149. package/dist/src/wgsl/fa2-stats-finalize.wgsl.js +54 -0
  150. package/dist/src/wgsl/fa2-stats-finalize.wgsl.js.map +1 -1
  151. package/dist/src/wgsl/grid-cell-key.wgsl.d.ts +8 -0
  152. package/dist/src/wgsl/grid-cell-key.wgsl.d.ts.map +1 -0
  153. package/dist/src/wgsl/grid-cell-key.wgsl.js +30 -0
  154. package/dist/src/wgsl/grid-cell-key.wgsl.js.map +1 -0
  155. package/dist/src/wgsl/grid-centroid-hub.wgsl.d.ts +8 -0
  156. package/dist/src/wgsl/grid-centroid-hub.wgsl.d.ts.map +1 -0
  157. package/dist/src/wgsl/grid-centroid-hub.wgsl.js +29 -0
  158. package/dist/src/wgsl/grid-centroid-hub.wgsl.js.map +1 -0
  159. package/dist/src/wgsl/grid-centroid.wgsl.d.ts +8 -0
  160. package/dist/src/wgsl/grid-centroid.wgsl.d.ts.map +1 -0
  161. package/dist/src/wgsl/grid-centroid.wgsl.js +29 -0
  162. package/dist/src/wgsl/grid-centroid.wgsl.js.map +1 -0
  163. package/dist/src/wgsl/grid-downsample.wgsl.d.ts +7 -0
  164. package/dist/src/wgsl/grid-downsample.wgsl.d.ts.map +1 -0
  165. package/dist/src/wgsl/grid-downsample.wgsl.js +28 -0
  166. package/dist/src/wgsl/grid-downsample.wgsl.js.map +1 -0
  167. package/dist/src/wgsl/grid-far-field.wgsl.d.ts +13 -0
  168. package/dist/src/wgsl/grid-far-field.wgsl.d.ts.map +1 -0
  169. package/dist/src/wgsl/grid-far-field.wgsl.js +98 -0
  170. package/dist/src/wgsl/grid-far-field.wgsl.js.map +1 -0
  171. package/dist/src/wgsl/grid-near-field.wgsl.d.ts +19 -0
  172. package/dist/src/wgsl/grid-near-field.wgsl.d.ts.map +1 -0
  173. package/dist/src/wgsl/grid-near-field.wgsl.js +129 -0
  174. package/dist/src/wgsl/grid-near-field.wgsl.js.map +1 -0
  175. package/dist/src/wgsl/histogram.wgsl.d.ts +7 -0
  176. package/dist/src/wgsl/histogram.wgsl.d.ts.map +1 -0
  177. package/dist/src/wgsl/histogram.wgsl.js +15 -0
  178. package/dist/src/wgsl/histogram.wgsl.js.map +1 -0
  179. package/dist/src/wgsl/indirect-finalize.wgsl.d.ts +8 -0
  180. package/dist/src/wgsl/indirect-finalize.wgsl.d.ts.map +1 -0
  181. package/dist/src/wgsl/indirect-finalize.wgsl.js +26 -0
  182. package/dist/src/wgsl/indirect-finalize.wgsl.js.map +1 -0
  183. package/dist/src/wgsl/radix-hist.wgsl.d.ts +9 -0
  184. package/dist/src/wgsl/radix-hist.wgsl.d.ts.map +1 -0
  185. package/dist/src/wgsl/radix-hist.wgsl.js +31 -0
  186. package/dist/src/wgsl/radix-hist.wgsl.js.map +1 -0
  187. package/dist/src/wgsl/radix-scatter.wgsl.d.ts +9 -0
  188. package/dist/src/wgsl/radix-scatter.wgsl.d.ts.map +1 -0
  189. package/dist/src/wgsl/radix-scatter.wgsl.js +40 -0
  190. package/dist/src/wgsl/radix-scatter.wgsl.js.map +1 -0
  191. package/dist/src/wgsl/scan-add.wgsl.d.ts +6 -0
  192. package/dist/src/wgsl/scan-add.wgsl.d.ts.map +1 -0
  193. package/dist/src/wgsl/scan-add.wgsl.js +14 -0
  194. package/dist/src/wgsl/scan-add.wgsl.js.map +1 -0
  195. package/dist/src/wgsl/scan-block.wgsl.d.ts +8 -0
  196. package/dist/src/wgsl/scan-block.wgsl.d.ts.map +1 -0
  197. package/dist/src/wgsl/scan-block.wgsl.js +30 -0
  198. package/dist/src/wgsl/scan-block.wgsl.js.map +1 -0
  199. package/dist/src/wgsl/segmented-reduce.wgsl.d.ts +22 -8
  200. package/dist/src/wgsl/segmented-reduce.wgsl.d.ts.map +1 -1
  201. package/dist/src/wgsl/segmented-reduce.wgsl.js +84 -15
  202. package/dist/src/wgsl/segmented-reduce.wgsl.js.map +1 -1
  203. package/dist/src/wgsl/spmv-pull.wgsl.d.ts +22 -11
  204. package/dist/src/wgsl/spmv-pull.wgsl.d.ts.map +1 -1
  205. package/dist/src/wgsl/spmv-pull.wgsl.js +110 -36
  206. package/dist/src/wgsl/spmv-pull.wgsl.js.map +1 -1
  207. package/dist/tsconfig.build.tsbuildinfo +1 -1
  208. package/dist/webgpu-graph-algorithms.js +5016 -1130
  209. package/dist/webgpu-graph-algorithms.js.map +1 -1
  210. package/package.json +10 -7
  211. package/src/accelerator.ts +46 -12
  212. package/src/algorithms/components.ts +12 -16
  213. package/src/algorithms/degree.ts +58 -43
  214. package/src/algorithms/pagerank.ts +20 -18
  215. package/src/algorithms/power-iteration.ts +19 -18
  216. package/src/constants.ts +108 -8
  217. package/src/errors.ts +3 -1
  218. package/src/index.ts +25 -5
  219. package/src/kernel/dispatch.ts +18 -7
  220. package/src/kernel/kernel.ts +59 -5
  221. package/src/kernel/prelude.ts +15 -0
  222. package/src/kernel/profiler.ts +28 -4
  223. package/src/kernels.ts +378 -24
  224. package/src/layouts/calibrate.ts +187 -0
  225. package/src/layouts/force-simulation.ts +111 -26
  226. package/src/layouts/forceatlas2.ts +346 -324
  227. package/src/layouts/fruchterman-reingold.ts +918 -0
  228. package/src/layouts/model-common.ts +323 -0
  229. package/src/layouts/repulsion-grid.ts +451 -0
  230. package/src/layouts/spring-electrical.ts +845 -0
  231. package/src/memory/residency.ts +126 -20
  232. package/src/primitives/core-shape.ts +91 -4
  233. package/src/primitives/grid-pyramid.ts +221 -0
  234. package/src/primitives/grid.ts +349 -0
  235. package/src/primitives/histogram.ts +273 -0
  236. package/src/primitives/radix-sort.ts +246 -0
  237. package/src/primitives/scan.ts +197 -0
  238. package/src/primitives/segmented-reduce.ts +214 -56
  239. package/src/primitives/spmv.ts +125 -65
  240. package/src/primitives/verify.ts +249 -0
  241. package/src/types/accelerator.ts +15 -3
  242. package/src/types/context.ts +56 -0
  243. package/src/types/layout.ts +58 -0
  244. package/src/types/options.ts +45 -1
  245. package/src/wgsl/counting-scatter.wgsl.ts +16 -0
  246. package/src/wgsl/fa2-attraction.wgsl.ts +100 -17
  247. package/src/wgsl/fa2-integrate.wgsl.ts +28 -2
  248. package/src/wgsl/fa2-repulsion-exact.wgsl.ts +14 -5
  249. package/src/wgsl/fa2-stats-finalize.wgsl.ts +54 -0
  250. package/src/wgsl/grid-cell-key.wgsl.ts +29 -0
  251. package/src/wgsl/grid-centroid-hub.wgsl.ts +28 -0
  252. package/src/wgsl/grid-centroid.wgsl.ts +28 -0
  253. package/src/wgsl/grid-downsample.wgsl.ts +27 -0
  254. package/src/wgsl/grid-far-field.wgsl.ts +97 -0
  255. package/src/wgsl/grid-near-field.wgsl.ts +128 -0
  256. package/src/wgsl/histogram.wgsl.ts +14 -0
  257. package/src/wgsl/indirect-finalize.wgsl.ts +25 -0
  258. package/src/wgsl/radix-hist.wgsl.ts +30 -0
  259. package/src/wgsl/radix-scatter.wgsl.ts +39 -0
  260. package/src/wgsl/scan-add.wgsl.ts +13 -0
  261. package/src/wgsl/scan-block.wgsl.ts +29 -0
  262. package/src/wgsl/segmented-reduce.wgsl.ts +84 -15
  263. package/src/wgsl/spmv-pull.wgsl.ts +110 -36
  264. package/dist/chunks/context-CRbw2Wyo.js.map +0 -1
@@ -0,0 +1,246 @@
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
+
22
+ import { RADIX_BINS, U32_MAX } from "../constants.js";
23
+ import { WebGpuGraphError } from "../errors.js";
24
+ import { plan1d } from "../kernel/dispatch.js";
25
+ import { type Kernel } from "../kernel/kernel.js";
26
+ import { kernelSpec, RADIX_PARAMS } from "../kernels.js";
27
+ import { type Binding } from "../types/memory.js";
28
+ import { type ReduceScope } from "./reduce.js";
29
+ import { prepareScan, type ScanPlanner } from "./scan.js";
30
+
31
+ /** The bits of one pass. */
32
+ const RADIX_DIGIT_BITS = 8;
33
+
34
+ /** The key widths record() accepts: `bits / 8` passes, so the low `bits` of every key order the pairs. */
35
+ export type RadixBits = 8 | 16 | 24 | 32;
36
+
37
+ /**
38
+ * The caller's scratch: a second key-value pair the passes ping-pong with, the digit-major histogram table and its
39
+ * scanned twin (each `radixHistBytes` bytes at least).
40
+ */
41
+ export interface RadixSortScratch {
42
+ readonly keys: Binding;
43
+ readonly vals: Binding;
44
+ /** The raw digit-major table `radix-hist` writes (the last pass's stays readable after the sort). */
45
+ readonly hist: Binding;
46
+ /** The exclusive scan of `hist`, what `radix-scatter` reads. */
47
+ readonly offsets: Binding;
48
+ }
49
+
50
+ /** Where the sorted pairs landed: the input pair or the scratch pair (record() decides by the pass count). */
51
+ export interface RadixSortResult {
52
+ readonly keys: Binding;
53
+ readonly vals: Binding;
54
+ }
55
+
56
+ /** A prepared radix sort (spec 6 row 6): records the pass dispatches of one sort into a pass. */
57
+ export interface RadixSortPlanner {
58
+ /**
59
+ * Records the stable sort of `count` (key, value) pairs by the low `bits` of the key; returns the pair holding the
60
+ * result (the scratch pair after an odd pass count, the input pair after an even one); count 0 records nothing and
61
+ * returns the input pair.
62
+ * @param pass - the compute pass
63
+ * @param keys - the keys (at least 4 x count bytes)
64
+ * @param vals - the values (at least 4 x count bytes)
65
+ * @param count - the pair count (a non-negative integer below 2^32)
66
+ * @param bits - the key width: 8, 16, 24 or 32
67
+ * @param scratch - the second pair, the histogram table and its scanned twin
68
+ * @returns the pair the result lives in
69
+ */
70
+ record(
71
+ pass: GPUComputePassEncoder,
72
+ keys: Binding,
73
+ vals: Binding,
74
+ count: number,
75
+ bits: RadixBits,
76
+ scratch: RadixSortScratch,
77
+ ): RadixSortResult;
78
+ /** Dispatches the last record() issued: 0 for count 0, else `passes x (2 + the scan's dispatches over the table)`. */
79
+ readonly lastDispatches: number;
80
+ }
81
+
82
+ /**
83
+ * The byte size of the digit-major histogram table of one pass over `count` keys at workgroup size `wg`:
84
+ * `4 x 256 x ceil(count / wg)` (the caller sizes `scratch.hist` by it; 0 for count 0, which records nothing).
85
+ * @param count - the pair count
86
+ * @param wg - the workgroup size (`scope.workgroupSize`)
87
+ * @returns the bytes
88
+ */
89
+ export function radixHistBytes(count: number, wg: number): number {
90
+ return 4 * RADIX_BINS * Math.ceil(count / wg);
91
+ }
92
+
93
+ /**
94
+ * Prepares the two radix pipelines and the scan of a scope (compiles once) so record() is synchronous. The planner
95
+ * lives exactly as long as the scope: never use a planner after its scope's dispose().
96
+ * @param scope - the caller's scope (device, caps, cache, scratch, params)
97
+ * @returns the planner
98
+ */
99
+ export async function prepareRadixSort(scope: ReduceScope): Promise<RadixSortPlanner> {
100
+ const hist = await scope.pipelines.kernel(kernelSpec("radix-hist"));
101
+ const scatter = await scope.pipelines.kernel(kernelSpec("radix-scatter"));
102
+ const scan = await prepareScan(scope);
103
+ return new RadixSortPlannerImpl(scope, hist, scatter, scan);
104
+ }
105
+
106
+ /**
107
+ * One binding-size check of record().
108
+ * @param argument - the argument name
109
+ * @param binding - the binding
110
+ * @param bytes - the bytes it must hold
111
+ */
112
+ function checkBinding(argument: string, binding: Binding, bytes: number): void {
113
+ if (binding.size < bytes) {
114
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", `radixSort: ${argument} is smaller than ${bytes} bytes`, {
115
+ argument,
116
+ value: binding.size,
117
+ expected: bytes,
118
+ });
119
+ }
120
+ }
121
+
122
+ /**
123
+ * The argument checks of record() (E_INVALID_ARGUMENT before anything is recorded).
124
+ * @param keys - the keys
125
+ * @param vals - the values
126
+ * @param count - the pair count
127
+ * @param bits - the key width
128
+ * @param scratch - the scratch
129
+ * @param wg - the workgroup size
130
+ */
131
+ function checkRecordArguments(
132
+ keys: Binding,
133
+ vals: Binding,
134
+ count: number,
135
+ bits: RadixBits,
136
+ scratch: RadixSortScratch,
137
+ wg: number,
138
+ ): void {
139
+ if (bits !== 8 && bits !== 16 && bits !== 24 && bits !== 32) {
140
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "radixSort: bits must be 8, 16, 24 or 32", {
141
+ argument: "bits",
142
+ value: bits,
143
+ expected: [8, 16, 24, 32],
144
+ });
145
+ }
146
+ if (!Number.isSafeInteger(count) || count < 0 || count > U32_MAX) {
147
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "radixSort: count must be a non-negative integer below 2^32", {
148
+ argument: "count",
149
+ value: count,
150
+ });
151
+ }
152
+ const pairBytes = 4 * count;
153
+ checkBinding("keys", keys, pairBytes);
154
+ checkBinding("vals", vals, pairBytes);
155
+ checkBinding("scratch.keys", scratch.keys, pairBytes);
156
+ checkBinding("scratch.vals", scratch.vals, pairBytes);
157
+ const tableBytes = radixHistBytes(count, wg);
158
+ checkBinding("scratch.hist", scratch.hist, tableBytes);
159
+ checkBinding("scratch.offsets", scratch.offsets, tableBytes);
160
+ }
161
+
162
+ /** The planner: the two resolved kernels and the scan over one scope. */
163
+ class RadixSortPlannerImpl implements RadixSortPlanner {
164
+ private readonly scope: ReduceScope;
165
+ private readonly hist: Kernel;
166
+ private readonly scatter: Kernel;
167
+ private readonly scan: ScanPlanner;
168
+ private dispatches = 0;
169
+
170
+ /**
171
+ * Wraps the resolved kernels; use prepareRadixSort().
172
+ * @param scope - the caller's scope
173
+ * @param hist - the `radix-hist` kernel
174
+ * @param scatter - the `radix-scatter` kernel
175
+ * @param scan - the scan planner of the same scope
176
+ */
177
+ constructor(scope: ReduceScope, hist: Kernel, scatter: Kernel, scan: ScanPlanner) {
178
+ this.scope = scope;
179
+ this.hist = hist;
180
+ this.scatter = scatter;
181
+ this.scan = scan;
182
+ }
183
+
184
+ /**
185
+ * Dispatches the last record() issued.
186
+ * @returns the count
187
+ */
188
+ get lastDispatches(): number {
189
+ return this.dispatches;
190
+ }
191
+
192
+ /**
193
+ * Records the passes into the pass (see the interface).
194
+ * @param pass - the compute pass
195
+ * @param keys - the keys
196
+ * @param vals - the values
197
+ * @param count - the pair count
198
+ * @param bits - the key width
199
+ * @param scratch - the second pair, the histogram table and its scanned twin
200
+ * @returns the pair the result lives in
201
+ */
202
+ record(
203
+ pass: GPUComputePassEncoder,
204
+ keys: Binding,
205
+ vals: Binding,
206
+ count: number,
207
+ bits: RadixBits,
208
+ scratch: RadixSortScratch,
209
+ ): RadixSortResult {
210
+ const wg = this.scope.workgroupSize;
211
+ checkRecordArguments(keys, vals, count, bits, scratch, wg);
212
+ if (count === 0) {
213
+ this.dispatches = 0;
214
+ return { keys, vals };
215
+ }
216
+ const groups = Math.ceil(count / wg);
217
+ const plan = plan1d(count, wg, this.scope.caps);
218
+ const tableWords = RADIX_BINS * groups;
219
+ const tableBytes = 4 * tableWords;
220
+ const histTable: Binding = { ...scratch.hist, size: tableBytes };
221
+ const offsets: Binding = { ...scratch.offsets, size: tableBytes };
222
+ let src: RadixSortResult = { keys, vals };
223
+ let dst: RadixSortResult = { keys: scratch.keys, vals: scratch.vals };
224
+ let dispatches = 0;
225
+ const passes = bits / RADIX_DIGIT_BITS;
226
+ for (let p = 0; p < passes; p++) {
227
+ const params = this.scope.params(RADIX_PARAMS, { count, shift: RADIX_DIGIT_BITS * p, groups, pad0: 0 });
228
+ const histBound = this.hist.bind({ keys: src.keys, hist: histTable, P: params.binding });
229
+ this.hist.dispatch(pass, histBound, plan, [params.offset]);
230
+ this.scan.record(pass, histTable, tableWords, offsets);
231
+ const scatterBound = this.scatter.bind({
232
+ keys: src.keys,
233
+ vals: src.vals,
234
+ offsets,
235
+ keysOut: dst.keys,
236
+ valsOut: dst.vals,
237
+ P: params.binding,
238
+ });
239
+ this.scatter.dispatch(pass, scatterBound, plan, [params.offset]);
240
+ dispatches += 2 + this.scan.lastDispatches;
241
+ [src, dst] = [dst, src];
242
+ }
243
+ this.dispatches = dispatches;
244
+ return src;
245
+ }
246
+ }
@@ -0,0 +1,197 @@
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
+
15
+ import { WebGpuGraphError } from "../errors.js";
16
+ import { type DispatchPlan, groupsOf, plan1d } from "../kernel/dispatch.js";
17
+ import { type Kernel } from "../kernel/kernel.js";
18
+ import { kernelSpec, SCAN_PARAMS } from "../kernels.js";
19
+ import { type Binding } from "../types/memory.js";
20
+ import { type ReduceScope } from "./reduce.js";
21
+
22
+ /** 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()). */
23
+ export interface ScanTotal {
24
+ readonly binding: Binding;
25
+ readonly index: number;
26
+ }
27
+
28
+ /** A prepared exclusive scan (spec 6 row 2): records the level dispatches of one scan into a pass. */
29
+ export interface ScanPlanner {
30
+ /**
31
+ * Records the scan of `count` u32 of `src` into `out` (exclusive); returns where the total landed (a word of the
32
+ * planner's scratch, valid until the next record()); nothing for count 0 (the total word is then 0).
33
+ * @param pass - the compute pass
34
+ * @param src - the input words (at least 4 x count bytes)
35
+ * @param count - the word count (a non-negative integer below 2^32)
36
+ * @param out - the output words (at least 4 x count bytes)
37
+ * @returns the total's location
38
+ */
39
+ record(pass: GPUComputePassEncoder, src: Binding, count: number, out: Binding): ScanTotal;
40
+ /** 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). */
41
+ readonly lastDispatches: number;
42
+ }
43
+
44
+ /** The largest u32: the largest count `ScanParams.count` can carry. */
45
+ const U32_MAX = 0xffffffff;
46
+
47
+ /**
48
+ * Prepares the two scan pipelines of a scope (compiles once) so record() is synchronous. The zero word the empty
49
+ * scan's total points at is a scratch of `scope`, so the planner lives exactly as long as the scope: never use a
50
+ * planner after its scope's dispose().
51
+ * @param scope - the caller's scope (device, caps, cache, scratch, params)
52
+ * @returns the planner
53
+ */
54
+ export async function prepareScan(scope: ReduceScope): Promise<ScanPlanner> {
55
+ const block = await scope.pipelines.kernel(kernelSpec("scan-block"));
56
+ const add = await scope.pipelines.kernel(kernelSpec("scan-add"));
57
+ const zero = scope.scratch(4, "scan/zero");
58
+ scope.device.queue.writeBuffer(zero, 0, new Uint32Array(1));
59
+ return new ScanPlannerImpl(scope, block, add, { buffer: zero, offset: 0, size: 4, window: null });
60
+ }
61
+
62
+ /**
63
+ * The argument checks of record() (E_INVALID_ARGUMENT before anything is recorded).
64
+ * @param src - the input binding
65
+ * @param count - the word count
66
+ * @param out - the output binding
67
+ */
68
+ function checkRecordArguments(src: Binding, count: number, out: Binding): void {
69
+ if (!Number.isSafeInteger(count) || count < 0 || count > U32_MAX) {
70
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "scan: count must be a non-negative integer below 2^32", {
71
+ argument: "count",
72
+ value: count,
73
+ });
74
+ }
75
+ if (src.size < 4 * count) {
76
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "scan: src is smaller than 4 x count bytes", {
77
+ argument: "src",
78
+ value: src.size,
79
+ expected: 4 * count,
80
+ });
81
+ }
82
+ if (out.size < 4 * count) {
83
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "scan: out is smaller than 4 x count bytes", {
84
+ argument: "out",
85
+ value: out.size,
86
+ expected: 4 * count,
87
+ });
88
+ }
89
+ }
90
+
91
+ /** One level of the recursion: its input words, its exclusive output, its count, its dispatch plan and its per-block sums. */
92
+ interface Level {
93
+ readonly input: Binding;
94
+ readonly output: Binding;
95
+ readonly count: number;
96
+ readonly plan: DispatchPlan;
97
+ readonly sums: Binding;
98
+ }
99
+
100
+ /** The planner: the two resolved kernels over one scope. */
101
+ class ScanPlannerImpl implements ScanPlanner {
102
+ private readonly scope: ReduceScope;
103
+ private readonly block: Kernel;
104
+ private readonly add: Kernel;
105
+ private readonly zero: Binding;
106
+ private dispatches = 0;
107
+
108
+ /**
109
+ * Wraps the resolved kernels; use prepareScan().
110
+ * @param scope - the caller's scope
111
+ * @param block - the `scan-block` kernel
112
+ * @param add - the `scan-add` kernel
113
+ * @param zero - a one-word binding holding 0 (the total of the empty scan)
114
+ */
115
+ constructor(scope: ReduceScope, block: Kernel, add: Kernel, zero: Binding) {
116
+ this.scope = scope;
117
+ this.block = block;
118
+ this.add = add;
119
+ this.zero = zero;
120
+ }
121
+
122
+ /**
123
+ * Dispatches the last record() issued.
124
+ * @returns the count
125
+ */
126
+ get lastDispatches(): number {
127
+ return this.dispatches;
128
+ }
129
+
130
+ /**
131
+ * Records the levels into the pass (see the interface).
132
+ * @param pass - the compute pass
133
+ * @param src - the input words
134
+ * @param count - the word count
135
+ * @param out - the output words
136
+ * @returns the total's location
137
+ */
138
+ record(pass: GPUComputePassEncoder, src: Binding, count: number, out: Binding): ScanTotal {
139
+ checkRecordArguments(src, count, out);
140
+ if (count === 0) {
141
+ this.dispatches = 0;
142
+ return { binding: this.zero, index: 0 };
143
+ }
144
+ const wg = this.scope.workgroupSize;
145
+ const levels: Level[] = [];
146
+ let input = src;
147
+ let output = out;
148
+ let levelCount = count;
149
+ for (;;) {
150
+ const blocks = Math.ceil(levelCount / wg);
151
+ // Sized by the PADDED workgroup count (reduce's shape, reduce.ts partials-1): above MAX_1D_ITEMS plan1d pads
152
+ // the grid to x * y >= blocks, and every padded workgroup still stores its (zero) block sum at its own index.
153
+ // The next level scans only `blocks` words, so the padded tail is written and never read.
154
+ const plan = plan1d(levelCount, wg, this.scope.caps);
155
+ const groups = groupsOf(plan);
156
+ const sums = this.scratch(groups, `scan/sums${levels.length}`);
157
+ levels.push({ input, output, count: levelCount, plan, sums });
158
+ if (blocks === 1) {
159
+ break;
160
+ }
161
+ input = sums;
162
+ output = this.scratch(groups, `scan/offsets${levels.length}`);
163
+ levelCount = blocks;
164
+ }
165
+ for (const level of levels) {
166
+ const params = this.scope.params(SCAN_PARAMS, { count: level.count, pad0: 0, pad1: 0, pad2: 0 });
167
+ const bound = this.block.bind({
168
+ src: level.input,
169
+ out: level.output,
170
+ blockSums: level.sums,
171
+ P: params.binding,
172
+ });
173
+ this.block.dispatch(pass, bound, level.plan, [params.offset]);
174
+ }
175
+ for (let l = levels.length - 2; l >= 0; l--) {
176
+ const level = levels[l];
177
+ const params = this.scope.params(SCAN_PARAMS, { count: level.count, pad0: 0, pad1: 0, pad2: 0 });
178
+ const bound = this.add.bind({ out: level.output, blockOffsets: levels[l + 1].output, P: params.binding });
179
+ this.add.dispatch(pass, bound, level.plan, [params.offset]);
180
+ }
181
+ this.dispatches = 2 * levels.length - 1;
182
+ const top = levels[levels.length - 1];
183
+ return { binding: top.sums, index: 0 };
184
+ }
185
+
186
+ /**
187
+ * A scratch of `words` u32 from the scope, bound whole (never zero-length; the pool rounds the buffer up).
188
+ * @param words - the word count (>= 1)
189
+ * @param label - the scratch label
190
+ * @returns the binding
191
+ */
192
+ private scratch(words: number, label: string): Binding {
193
+ const size = 4 * words;
194
+ const buffer = this.scope.scratch(size, label);
195
+ return { buffer, offset: 0, size, window: null };
196
+ }
197
+ }