@graphty/webgpu-graph-algorithms 0.5.1 → 0.6.1

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 (243) hide show
  1. package/README.md +459 -58
  2. package/dist/browser.js +1 -1
  3. package/dist/chunks/{context-BR7fx3vR.js → context-BXqgCifx.js} +190 -40
  4. package/dist/chunks/context-BXqgCifx.js.map +1 -0
  5. package/dist/node.js +1 -1
  6. package/dist/src/algorithms/components.d.ts.map +1 -1
  7. package/dist/src/algorithms/components.js +12 -13
  8. package/dist/src/algorithms/components.js.map +1 -1
  9. package/dist/src/algorithms/degree.d.ts +6 -8
  10. package/dist/src/algorithms/degree.d.ts.map +1 -1
  11. package/dist/src/algorithms/degree.js +58 -35
  12. package/dist/src/algorithms/degree.js.map +1 -1
  13. package/dist/src/algorithms/pagerank.d.ts.map +1 -1
  14. package/dist/src/algorithms/pagerank.js +16 -14
  15. package/dist/src/algorithms/pagerank.js.map +1 -1
  16. package/dist/src/algorithms/power-iteration.d.ts +2 -2
  17. package/dist/src/algorithms/power-iteration.d.ts.map +1 -1
  18. package/dist/src/algorithms/power-iteration.js +17 -14
  19. package/dist/src/algorithms/power-iteration.js.map +1 -1
  20. package/dist/src/constants.d.ts +38 -8
  21. package/dist/src/constants.d.ts.map +1 -1
  22. package/dist/src/constants.js +38 -8
  23. package/dist/src/constants.js.map +1 -1
  24. package/dist/src/errors.d.ts +3 -2
  25. package/dist/src/errors.d.ts.map +1 -1
  26. package/dist/src/errors.js +2 -1
  27. package/dist/src/errors.js.map +1 -1
  28. package/dist/src/index.d.ts +6 -4
  29. package/dist/src/index.d.ts.map +1 -1
  30. package/dist/src/index.js +8 -3
  31. package/dist/src/index.js.map +1 -1
  32. package/dist/src/kernel/dispatch.d.ts +8 -3
  33. package/dist/src/kernel/dispatch.d.ts.map +1 -1
  34. package/dist/src/kernel/dispatch.js +18 -7
  35. package/dist/src/kernel/dispatch.js.map +1 -1
  36. package/dist/src/kernel/kernel.d.ts +30 -1
  37. package/dist/src/kernel/kernel.d.ts.map +1 -1
  38. package/dist/src/kernel/kernel.js +49 -5
  39. package/dist/src/kernel/kernel.js.map +1 -1
  40. package/dist/src/kernel/prelude.d.ts.map +1 -1
  41. package/dist/src/kernel/prelude.js +6 -1
  42. package/dist/src/kernel/prelude.js.map +1 -1
  43. package/dist/src/kernel/profiler.d.ts +15 -3
  44. package/dist/src/kernel/profiler.d.ts.map +1 -1
  45. package/dist/src/kernel/profiler.js +27 -4
  46. package/dist/src/kernel/profiler.js.map +1 -1
  47. package/dist/src/kernels.d.ts +17 -7
  48. package/dist/src/kernels.d.ts.map +1 -1
  49. package/dist/src/kernels.js +323 -16
  50. package/dist/src/kernels.js.map +1 -1
  51. package/dist/src/layouts/calibrate.d.ts +51 -0
  52. package/dist/src/layouts/calibrate.d.ts.map +1 -0
  53. package/dist/src/layouts/calibrate.js +172 -0
  54. package/dist/src/layouts/calibrate.js.map +1 -0
  55. package/dist/src/layouts/force-simulation.d.ts +39 -4
  56. package/dist/src/layouts/force-simulation.d.ts.map +1 -1
  57. package/dist/src/layouts/force-simulation.js +71 -19
  58. package/dist/src/layouts/force-simulation.js.map +1 -1
  59. package/dist/src/layouts/forceatlas2.d.ts +107 -36
  60. package/dist/src/layouts/forceatlas2.d.ts.map +1 -1
  61. package/dist/src/layouts/forceatlas2.js +296 -100
  62. package/dist/src/layouts/forceatlas2.js.map +1 -1
  63. package/dist/src/layouts/fruchterman-reingold.d.ts +73 -27
  64. package/dist/src/layouts/fruchterman-reingold.d.ts.map +1 -1
  65. package/dist/src/layouts/fruchterman-reingold.js +230 -70
  66. package/dist/src/layouts/fruchterman-reingold.js.map +1 -1
  67. package/dist/src/layouts/model-common.d.ts +41 -3
  68. package/dist/src/layouts/model-common.d.ts.map +1 -1
  69. package/dist/src/layouts/model-common.js +74 -3
  70. package/dist/src/layouts/model-common.js.map +1 -1
  71. package/dist/src/layouts/repulsion-grid.d.ts +152 -0
  72. package/dist/src/layouts/repulsion-grid.d.ts.map +1 -0
  73. package/dist/src/layouts/repulsion-grid.js +318 -0
  74. package/dist/src/layouts/repulsion-grid.js.map +1 -0
  75. package/dist/src/layouts/spring-electrical.d.ts +75 -30
  76. package/dist/src/layouts/spring-electrical.d.ts.map +1 -1
  77. package/dist/src/layouts/spring-electrical.js +231 -74
  78. package/dist/src/layouts/spring-electrical.js.map +1 -1
  79. package/dist/src/memory/residency.d.ts +6 -2
  80. package/dist/src/memory/residency.d.ts.map +1 -1
  81. package/dist/src/memory/residency.js +84 -14
  82. package/dist/src/memory/residency.js.map +1 -1
  83. package/dist/src/primitives/core-shape.d.ts +38 -2
  84. package/dist/src/primitives/core-shape.d.ts.map +1 -1
  85. package/dist/src/primitives/core-shape.js +71 -3
  86. package/dist/src/primitives/core-shape.js.map +1 -1
  87. package/dist/src/primitives/grid-pyramid.d.ts +71 -0
  88. package/dist/src/primitives/grid-pyramid.d.ts.map +1 -0
  89. package/dist/src/primitives/grid-pyramid.js +143 -0
  90. package/dist/src/primitives/grid-pyramid.js.map +1 -0
  91. package/dist/src/primitives/grid.d.ts +118 -0
  92. package/dist/src/primitives/grid.d.ts.map +1 -0
  93. package/dist/src/primitives/grid.js +225 -0
  94. package/dist/src/primitives/grid.js.map +1 -0
  95. package/dist/src/primitives/histogram.d.ts +67 -0
  96. package/dist/src/primitives/histogram.d.ts.map +1 -0
  97. package/dist/src/primitives/histogram.js +190 -0
  98. package/dist/src/primitives/histogram.js.map +1 -0
  99. package/dist/src/primitives/radix-sort.d.ts +75 -0
  100. package/dist/src/primitives/radix-sort.d.ts.map +1 -0
  101. package/dist/src/primitives/radix-sort.js +168 -0
  102. package/dist/src/primitives/radix-sort.js.map +1 -0
  103. package/dist/src/primitives/scan.d.ts +44 -0
  104. package/dist/src/primitives/scan.d.ts.map +1 -0
  105. package/dist/src/primitives/scan.js +151 -0
  106. package/dist/src/primitives/scan.js.map +1 -0
  107. package/dist/src/primitives/segmented-reduce.d.ts +25 -17
  108. package/dist/src/primitives/segmented-reduce.d.ts.map +1 -1
  109. package/dist/src/primitives/segmented-reduce.js +166 -47
  110. package/dist/src/primitives/segmented-reduce.js.map +1 -1
  111. package/dist/src/primitives/spmv.d.ts +18 -14
  112. package/dist/src/primitives/spmv.d.ts.map +1 -1
  113. package/dist/src/primitives/spmv.js +94 -58
  114. package/dist/src/primitives/spmv.js.map +1 -1
  115. package/dist/src/primitives/verify.d.ts +49 -0
  116. package/dist/src/primitives/verify.d.ts.map +1 -0
  117. package/dist/src/primitives/verify.js +229 -0
  118. package/dist/src/primitives/verify.js.map +1 -0
  119. package/dist/src/types/context.d.ts +53 -0
  120. package/dist/src/types/context.d.ts.map +1 -1
  121. package/dist/src/types/layout.d.ts +20 -0
  122. package/dist/src/types/layout.d.ts.map +1 -1
  123. package/dist/src/wgsl/counting-scatter.wgsl.d.ts +8 -0
  124. package/dist/src/wgsl/counting-scatter.wgsl.d.ts.map +1 -0
  125. package/dist/src/wgsl/counting-scatter.wgsl.js +17 -0
  126. package/dist/src/wgsl/counting-scatter.wgsl.js.map +1 -0
  127. package/dist/src/wgsl/fa2-attraction.wgsl.d.ts +23 -11
  128. package/dist/src/wgsl/fa2-attraction.wgsl.d.ts.map +1 -1
  129. package/dist/src/wgsl/fa2-attraction.wgsl.js +98 -20
  130. package/dist/src/wgsl/fa2-attraction.wgsl.js.map +1 -1
  131. package/dist/src/wgsl/fa2-stats-finalize.wgsl.d.ts +6 -2
  132. package/dist/src/wgsl/fa2-stats-finalize.wgsl.d.ts.map +1 -1
  133. package/dist/src/wgsl/fa2-stats-finalize.wgsl.js +22 -1
  134. package/dist/src/wgsl/fa2-stats-finalize.wgsl.js.map +1 -1
  135. package/dist/src/wgsl/grid-cell-key.wgsl.d.ts +8 -0
  136. package/dist/src/wgsl/grid-cell-key.wgsl.d.ts.map +1 -0
  137. package/dist/src/wgsl/grid-cell-key.wgsl.js +30 -0
  138. package/dist/src/wgsl/grid-cell-key.wgsl.js.map +1 -0
  139. package/dist/src/wgsl/grid-centroid-hub.wgsl.d.ts +8 -0
  140. package/dist/src/wgsl/grid-centroid-hub.wgsl.d.ts.map +1 -0
  141. package/dist/src/wgsl/grid-centroid-hub.wgsl.js +29 -0
  142. package/dist/src/wgsl/grid-centroid-hub.wgsl.js.map +1 -0
  143. package/dist/src/wgsl/grid-centroid.wgsl.d.ts +8 -0
  144. package/dist/src/wgsl/grid-centroid.wgsl.d.ts.map +1 -0
  145. package/dist/src/wgsl/grid-centroid.wgsl.js +29 -0
  146. package/dist/src/wgsl/grid-centroid.wgsl.js.map +1 -0
  147. package/dist/src/wgsl/grid-downsample.wgsl.d.ts +7 -0
  148. package/dist/src/wgsl/grid-downsample.wgsl.d.ts.map +1 -0
  149. package/dist/src/wgsl/grid-downsample.wgsl.js +28 -0
  150. package/dist/src/wgsl/grid-downsample.wgsl.js.map +1 -0
  151. package/dist/src/wgsl/grid-far-field.wgsl.d.ts +13 -0
  152. package/dist/src/wgsl/grid-far-field.wgsl.d.ts.map +1 -0
  153. package/dist/src/wgsl/grid-far-field.wgsl.js +98 -0
  154. package/dist/src/wgsl/grid-far-field.wgsl.js.map +1 -0
  155. package/dist/src/wgsl/grid-near-field.wgsl.d.ts +19 -0
  156. package/dist/src/wgsl/grid-near-field.wgsl.d.ts.map +1 -0
  157. package/dist/src/wgsl/grid-near-field.wgsl.js +129 -0
  158. package/dist/src/wgsl/grid-near-field.wgsl.js.map +1 -0
  159. package/dist/src/wgsl/histogram.wgsl.d.ts +7 -0
  160. package/dist/src/wgsl/histogram.wgsl.d.ts.map +1 -0
  161. package/dist/src/wgsl/histogram.wgsl.js +15 -0
  162. package/dist/src/wgsl/histogram.wgsl.js.map +1 -0
  163. package/dist/src/wgsl/indirect-finalize.wgsl.d.ts +8 -0
  164. package/dist/src/wgsl/indirect-finalize.wgsl.d.ts.map +1 -0
  165. package/dist/src/wgsl/indirect-finalize.wgsl.js +26 -0
  166. package/dist/src/wgsl/indirect-finalize.wgsl.js.map +1 -0
  167. package/dist/src/wgsl/radix-hist.wgsl.d.ts +9 -0
  168. package/dist/src/wgsl/radix-hist.wgsl.d.ts.map +1 -0
  169. package/dist/src/wgsl/radix-hist.wgsl.js +31 -0
  170. package/dist/src/wgsl/radix-hist.wgsl.js.map +1 -0
  171. package/dist/src/wgsl/radix-scatter.wgsl.d.ts +9 -0
  172. package/dist/src/wgsl/radix-scatter.wgsl.d.ts.map +1 -0
  173. package/dist/src/wgsl/radix-scatter.wgsl.js +40 -0
  174. package/dist/src/wgsl/radix-scatter.wgsl.js.map +1 -0
  175. package/dist/src/wgsl/scan-add.wgsl.d.ts +6 -0
  176. package/dist/src/wgsl/scan-add.wgsl.d.ts.map +1 -0
  177. package/dist/src/wgsl/scan-add.wgsl.js +14 -0
  178. package/dist/src/wgsl/scan-add.wgsl.js.map +1 -0
  179. package/dist/src/wgsl/scan-block.wgsl.d.ts +8 -0
  180. package/dist/src/wgsl/scan-block.wgsl.d.ts.map +1 -0
  181. package/dist/src/wgsl/scan-block.wgsl.js +30 -0
  182. package/dist/src/wgsl/scan-block.wgsl.js.map +1 -0
  183. package/dist/src/wgsl/segmented-reduce.wgsl.d.ts +22 -8
  184. package/dist/src/wgsl/segmented-reduce.wgsl.d.ts.map +1 -1
  185. package/dist/src/wgsl/segmented-reduce.wgsl.js +84 -15
  186. package/dist/src/wgsl/segmented-reduce.wgsl.js.map +1 -1
  187. package/dist/src/wgsl/spmv-pull.wgsl.d.ts +22 -11
  188. package/dist/src/wgsl/spmv-pull.wgsl.d.ts.map +1 -1
  189. package/dist/src/wgsl/spmv-pull.wgsl.js +110 -36
  190. package/dist/src/wgsl/spmv-pull.wgsl.js.map +1 -1
  191. package/dist/tsconfig.build.tsbuildinfo +1 -1
  192. package/dist/webgpu-graph-algorithms.js +3815 -1003
  193. package/dist/webgpu-graph-algorithms.js.map +1 -1
  194. package/package.json +9 -8
  195. package/src/algorithms/components.ts +12 -16
  196. package/src/algorithms/degree.ts +58 -43
  197. package/src/algorithms/pagerank.ts +20 -18
  198. package/src/algorithms/power-iteration.ts +19 -18
  199. package/src/constants.ts +38 -8
  200. package/src/errors.ts +3 -1
  201. package/src/index.ts +14 -4
  202. package/src/kernel/dispatch.ts +18 -7
  203. package/src/kernel/kernel.ts +59 -5
  204. package/src/kernel/prelude.ts +9 -0
  205. package/src/kernel/profiler.ts +28 -4
  206. package/src/kernels.ts +356 -18
  207. package/src/layouts/calibrate.ts +187 -0
  208. package/src/layouts/force-simulation.ts +91 -23
  209. package/src/layouts/forceatlas2.ts +331 -106
  210. package/src/layouts/fruchterman-reingold.ts +255 -74
  211. package/src/layouts/model-common.ts +98 -3
  212. package/src/layouts/repulsion-grid.ts +451 -0
  213. package/src/layouts/spring-electrical.ts +257 -78
  214. package/src/memory/residency.ts +126 -20
  215. package/src/primitives/core-shape.ts +91 -4
  216. package/src/primitives/grid-pyramid.ts +221 -0
  217. package/src/primitives/grid.ts +349 -0
  218. package/src/primitives/histogram.ts +273 -0
  219. package/src/primitives/radix-sort.ts +246 -0
  220. package/src/primitives/scan.ts +197 -0
  221. package/src/primitives/segmented-reduce.ts +214 -56
  222. package/src/primitives/spmv.ts +125 -65
  223. package/src/primitives/verify.ts +249 -0
  224. package/src/types/context.ts +56 -0
  225. package/src/types/layout.ts +22 -0
  226. package/src/wgsl/counting-scatter.wgsl.ts +16 -0
  227. package/src/wgsl/fa2-attraction.wgsl.ts +98 -20
  228. package/src/wgsl/fa2-stats-finalize.wgsl.ts +22 -1
  229. package/src/wgsl/grid-cell-key.wgsl.ts +29 -0
  230. package/src/wgsl/grid-centroid-hub.wgsl.ts +28 -0
  231. package/src/wgsl/grid-centroid.wgsl.ts +28 -0
  232. package/src/wgsl/grid-downsample.wgsl.ts +27 -0
  233. package/src/wgsl/grid-far-field.wgsl.ts +97 -0
  234. package/src/wgsl/grid-near-field.wgsl.ts +128 -0
  235. package/src/wgsl/histogram.wgsl.ts +14 -0
  236. package/src/wgsl/indirect-finalize.wgsl.ts +25 -0
  237. package/src/wgsl/radix-hist.wgsl.ts +30 -0
  238. package/src/wgsl/radix-scatter.wgsl.ts +39 -0
  239. package/src/wgsl/scan-add.wgsl.ts +13 -0
  240. package/src/wgsl/scan-block.wgsl.ts +29 -0
  241. package/src/wgsl/segmented-reduce.wgsl.ts +84 -15
  242. package/src/wgsl/spmv-pull.wgsl.ts +110 -36
  243. package/dist/chunks/context-BR7fx3vR.js.map +0 -1
@@ -0,0 +1,349 @@
1
+ /**
2
+ * The grid build (spec 6 row 12, 7.7 G1-G3; P4-T8): `GridSpec` on the host (PD-9) and the planner that records, into
3
+ * the caller's pass, the cell keys (G1, `grid-cell-key`), the stable sort by key (G2: `radixSort` at GRID_SORT_BITS,
4
+ * or `countingSortByKey` when the caller asks for the set-deterministic path) and the per-cell histogram with its
5
+ * exclusive scan (G3: the `histogram` kernel over `cellKey` and the `scan` of it; DEP-P4-I names no grid-specific
6
+ * id). `cellHist` and `cellStart` hold `cells + 2` words: every real cell, the outside pseudo-cell at index `cells`
7
+ * and one more so `cellStart[cells + 1] === n` closes the last range. Every zeroing is a `fill` dispatch inside the
8
+ * pass (PD-12), never an encoder clear.
9
+ *
10
+ * The named grid buffers (`cellKey`, `cellVal`, `sortedKey`, `sortedIdx`, `cellHist`, `cellStart`) are the caller's
11
+ * (the model's `BufferSpec`s, so `inspect(name)` reaches them); the anonymous scratch of the sort (the digit-major
12
+ * table and its scan, the counting-sort cursor) is taken from the scope ONCE at bind() so the sub-kernels bind the
13
+ * same buffers every iteration and `Kernel.bind`'s identity cache holds (PD-11). `src/primitives/**` never imports
14
+ * `src/context.ts`.
15
+ */
16
+
17
+ import { GRID_COARSEST_SIDE, GRID_MIN_SIDE, GRID_SORT_BITS } from "../constants.js";
18
+ import { WebGpuGraphError } from "../errors.js";
19
+ import { plan1d } from "../kernel/dispatch.js";
20
+ import { type Kernel } from "../kernel/kernel.js";
21
+ import { kernelSpec } from "../kernels.js";
22
+ import { type ResolvedLayoutTuning } from "../types/layout.js";
23
+ import { type Binding } from "../types/memory.js";
24
+ import {
25
+ type CountingSortPlanner,
26
+ type HistogramPlanner,
27
+ prepareCountingSort,
28
+ prepareHistogram,
29
+ } from "./histogram.js";
30
+ import { prepareRadixSort, radixHistBytes, type RadixSortPlanner } from "./radix-sort.js";
31
+ import { type ReduceScope } from "./reduce.js";
32
+ import { prepareScan, type ScanPlanner } from "./scan.js";
33
+
34
+ /** The grid's geometry (spec 7.7 geometry table; PD-9), computed on the host once per load. */
35
+ export interface GridSpec {
36
+ /** 2 or 3. */
37
+ readonly dim: 2 | 3;
38
+ /** The finest side `G` per axis: a power of two in [GRID_MIN_SIDE, floorPow2(gridMax)]. */
39
+ readonly g: number;
40
+ /** `log2(G / GRID_COARSEST_SIDE) + 1`. */
41
+ readonly levels: number;
42
+ /** `G^dim` finest cells; the outside pseudo-cell is index `cells`. */
43
+ readonly cells: number;
44
+ /** `cells + 2`: the length of `cellHist` / `cellStart`. */
45
+ readonly histWords: number;
46
+ /** The first cell of every level inside the pyramid: `levelOffsets[0] = 0`, level 0 holds `cells + 1` (the pseudo-cell), level L `(G / 2^L)^dim`. */
47
+ readonly levelOffsets: readonly number[];
48
+ /** Every level's cells together: `levelOffsets[levels - 1] + GRID_COARSEST_SIDE^dim`. */
49
+ readonly pyramidCells: number;
50
+ /** Whether G2 is the stable radix sort (true) or the set-deterministic counting sort (false). */
51
+ readonly deterministic: boolean;
52
+ }
53
+
54
+ /**
55
+ * The smallest power of two >= x (by doubling; 1 for x <= 1).
56
+ * @param x - a non-negative number
57
+ * @returns the power of two
58
+ */
59
+ function nextPow2(x: number): number {
60
+ let p = 1;
61
+ while (p < x) {
62
+ p *= 2;
63
+ }
64
+ return p;
65
+ }
66
+
67
+ /**
68
+ * The largest power of two <= x (by doubling; 1 for x < 2).
69
+ * @param x - a number >= 1
70
+ * @returns the power of two
71
+ */
72
+ function floorPow2(x: number): number {
73
+ let p = 1;
74
+ while (p * 2 <= x) {
75
+ p *= 2;
76
+ }
77
+ return p;
78
+ }
79
+
80
+ /**
81
+ * The grid of `n` nodes in `dim` dimensions under the tuning (spec 7.7 geometry table; PD-9): `G = clamp(nextPow2(2 *
82
+ * ceil(n^(1 / dim))), GRID_MIN_SIDE, floorPow2(gridMax))` where `gridMax` is `gridMax2D` or `gridMax3D`, rounded DOWN
83
+ * to a power of two so every level's side is an integer (512 and 128 stay; 100 becomes 64); `levels = log2(G /
84
+ * GRID_COARSEST_SIDE) + 1`. At the caps: 349,521 pyramid cells in 2D, 2,396,737 in 3D (the design's counts plus the
85
+ * pseudo-cell).
86
+ * @param n - the node count (>= 0)
87
+ * @param dim - 2 or 3
88
+ * @param tuning - the resolved layout tuning (`gridMax2D`, `gridMax3D`, `deterministic`)
89
+ * @returns the spec
90
+ */
91
+ export function gridSpecFor(
92
+ n: number,
93
+ dim: 2 | 3,
94
+ tuning: Pick<ResolvedLayoutTuning, "gridMax2D" | "gridMax3D" | "deterministic">,
95
+ ): GridSpec {
96
+ const gridMax = dim === 3 ? tuning.gridMax3D : tuning.gridMax2D;
97
+ const side = dim === 3 ? Math.cbrt(n) : Math.sqrt(n);
98
+ const cap = Math.max(GRID_MIN_SIDE, floorPow2(Math.max(1, gridMax)));
99
+ const g = Math.min(cap, Math.max(GRID_MIN_SIDE, nextPow2(2 * Math.ceil(side))));
100
+ let levels = 1;
101
+ for (let s = g; s > GRID_COARSEST_SIDE; s /= 2) {
102
+ levels++;
103
+ }
104
+ const cells = g ** dim;
105
+ const levelOffsets: number[] = [0];
106
+ let s = g;
107
+ for (let level = 0; level + 1 < levels; level++) {
108
+ levelOffsets.push(levelOffsets[level] + s ** dim + (level === 0 ? 1 : 0));
109
+ s /= 2;
110
+ }
111
+ return {
112
+ dim,
113
+ g,
114
+ levels,
115
+ cells,
116
+ histWords: cells + 2,
117
+ levelOffsets: Object.freeze(levelOffsets),
118
+ pyramidCells: levelOffsets[levels - 1] + GRID_COARSEST_SIDE ** dim,
119
+ deterministic: tuning.deterministic,
120
+ };
121
+ }
122
+
123
+ /**
124
+ * The bytes of the pyramid (spec 7.7: 16 B per cell, every level, the pseudo-cell included): 38,347,792 at the 3D cap.
125
+ * @param spec - the grid
126
+ * @returns the byte length
127
+ */
128
+ export function gridPyramidBytes(spec: GridSpec): number {
129
+ return 16 * spec.pyramidCells;
130
+ }
131
+
132
+ /**
133
+ * The buffers the build reads and writes (the model's named buffers; `state` and `params` are the blocks G1 reads).
134
+ * The parameter type of GridBuildPlanner.bind (knip: exported for the signature, not imported by name).
135
+ * @public
136
+ */
137
+ export interface GridBuildBindings {
138
+ /** `array<vec4f>` positions (xyz, mass). */
139
+ readonly pos: Binding;
140
+ /** The `Fa2State` block (`gridMin`, `invCellSize`), read-only here. */
141
+ readonly state: Binding;
142
+ /** The `Fa2Params` uniform ring (`n`, `dim`, `gridMax`); `record()` takes the iteration's dynamic offset. */
143
+ readonly params: Binding;
144
+ /**
145
+ * `n` words: G1's keys, node-indexed until G2. On the deterministic path the radix sort ping-pongs through this
146
+ * pair (PD-5: three passes, the even one writes the input pair), so after G2 `cellKey` / `cellVal` hold the
147
+ * sort's last even-pass intermediate -- a permutation of the keys, which is why the histogram recorded after it
148
+ * counts the same multiset; the node-indexed keys of an iteration are read through `upTo: "G1"`. The counting
149
+ * path never writes them.
150
+ */
151
+ readonly cellKey: Binding;
152
+ /** `n` words: G1's values (`i`); the sort's working pair with `cellKey`. */
153
+ readonly cellVal: Binding;
154
+ /** `n` words: the sorted keys (the radix path's result pair; unused by the counting path). */
155
+ readonly sortedKey: Binding;
156
+ /** `n` words: the sorted node indices. */
157
+ readonly sortedIdx: Binding;
158
+ /** `cells + 2` words: the per-cell counts. */
159
+ readonly cellHist: Binding;
160
+ /** `cells + 2` words: the exclusive scan of `cellHist`. */
161
+ readonly cellStart: Binding;
162
+ }
163
+
164
+ /** Where `record()` stops: after the keys (G1), after the sort (G2) or after the histogram and its scan (G3, the default). */
165
+ export type GridBuildStage = "G1" | "G2" | "G3";
166
+
167
+ /** A prepared grid build (spec 7.7 G1-G3): binds once per load, records the stages of one iteration into a pass. */
168
+ export interface GridBuildPlanner {
169
+ /**
170
+ * Binds the named buffers and takes the sort's scratch from the scope, sized by `cellKey` (its word count is the
171
+ * node capacity; PD-11); called once per load (a second call rebinds and takes fresh scratch, so it belongs to a
172
+ * reload, never to an iteration).
173
+ * @param bindings - the buffers
174
+ */
175
+ bind(bindings: GridBuildBindings): void;
176
+ /**
177
+ * Records G1, G2 and G3 for `n` nodes at the `Fa2Params` slot `paramsOffset`; `upTo` stops after the named
178
+ * stage (the counting path has no separate G2 stop: its sort and histogram are one sequence, so "G2" runs it all).
179
+ * @param pass - the compute pass
180
+ * @param n - the node count (in [1, the capacity]; the caller never records a grid iteration for an empty graph)
181
+ * @param paramsOffset - the dynamic offset of this iteration's `Fa2Params`
182
+ * @param upTo - the last stage to record (default "G3")
183
+ */
184
+ record(pass: GPUComputePassEncoder, n: number, paramsOffset: number, upTo?: GridBuildStage): void;
185
+ /** Dispatches the last record() issued. */
186
+ readonly lastDispatches: number;
187
+ }
188
+
189
+ /** The G2 / G3 planners: the stable radix path (PD-5) or the set-deterministic counting path. */
190
+ type SortPath =
191
+ | {
192
+ readonly kind: "radix";
193
+ readonly radix: RadixSortPlanner;
194
+ readonly histogram: HistogramPlanner;
195
+ readonly scan: ScanPlanner;
196
+ }
197
+ | { readonly kind: "counting"; readonly counting: CountingSortPlanner };
198
+
199
+ /**
200
+ * Prepares the grid build's pipelines over a scope (G1 and the sort / histogram / scan planners it composes; compiles
201
+ * once) so bind() and record() are synchronous. The planner lives exactly as long as the scope.
202
+ * @param scope - the caller's scope (device, caps, cache, scratch, params)
203
+ * @param spec - the grid
204
+ * @returns the planner
205
+ */
206
+ export async function prepareGridBuild(scope: ReduceScope, spec: GridSpec): Promise<GridBuildPlanner> {
207
+ const cellKey = await scope.pipelines.kernel(kernelSpec("grid-cell-key"));
208
+ const path: SortPath = spec.deterministic
209
+ ? {
210
+ kind: "radix",
211
+ radix: await prepareRadixSort(scope),
212
+ histogram: await prepareHistogram(scope),
213
+ scan: await prepareScan(scope),
214
+ }
215
+ : { kind: "counting", counting: await prepareCountingSort(scope) };
216
+ return new GridBuildPlannerImpl(scope, spec, cellKey, path);
217
+ }
218
+
219
+ /** What bind() prepared: the buffers, the node capacity and the sort's scratch (the radix table and its scan, or the counting-sort cursor and an unused twin). */
220
+ interface Bound {
221
+ readonly bindings: GridBuildBindings;
222
+ readonly capacity: number;
223
+ readonly scratchA: Binding;
224
+ readonly scratchB: Binding;
225
+ }
226
+
227
+ /** The planner: G1 and the composed sort / histogram / scan over one scope. */
228
+ class GridBuildPlannerImpl implements GridBuildPlanner {
229
+ private readonly scope: ReduceScope;
230
+ private readonly spec: GridSpec;
231
+ private readonly cellKey: Kernel;
232
+ private readonly path: SortPath;
233
+ private bound: Bound | null = null;
234
+ private dispatches = 0;
235
+
236
+ /**
237
+ * Wraps the resolved kernel and planners; use prepareGridBuild().
238
+ * @param scope - the caller's scope
239
+ * @param spec - the grid
240
+ * @param cellKey - the `grid-cell-key` kernel
241
+ * @param path - the sort path
242
+ */
243
+ constructor(scope: ReduceScope, spec: GridSpec, cellKey: Kernel, path: SortPath) {
244
+ this.scope = scope;
245
+ this.spec = spec;
246
+ this.cellKey = cellKey;
247
+ this.path = path;
248
+ }
249
+
250
+ /**
251
+ * Dispatches the last record() issued.
252
+ * @returns the count
253
+ */
254
+ get lastDispatches(): number {
255
+ return this.dispatches;
256
+ }
257
+
258
+ /**
259
+ * Binds the buffers and takes the sort's scratch (see the interface).
260
+ * @param bindings - the buffers
261
+ */
262
+ bind(bindings: GridBuildBindings): void {
263
+ const capacity = Math.floor(bindings.cellKey.size / 4);
264
+ if (capacity < 1) {
265
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "gridBuild: cellKey must hold at least one word", {
266
+ argument: "cellKey",
267
+ value: bindings.cellKey.size,
268
+ expected: 4,
269
+ });
270
+ }
271
+ const words =
272
+ this.path.kind === "radix" ? radixHistBytes(capacity, this.scope.workgroupSize) / 4 : this.spec.histWords;
273
+ this.bound = {
274
+ bindings,
275
+ capacity,
276
+ scratchA: this.scratch(words, "grid/sort-scratch-a"),
277
+ scratchB: this.scratch(words, "grid/sort-scratch-b"),
278
+ };
279
+ }
280
+
281
+ /**
282
+ * Records the stages (see the interface).
283
+ * @param pass - the compute pass
284
+ * @param n - the node count
285
+ * @param paramsOffset - the `Fa2Params` dynamic offset
286
+ * @param upTo - the last stage
287
+ */
288
+ record(pass: GPUComputePassEncoder, n: number, paramsOffset: number, upTo?: GridBuildStage): void {
289
+ const { bound } = this;
290
+ if (bound === null) {
291
+ throw new WebGpuGraphError("E_NOT_LOADED", "gridBuild: record() before bind()", { argument: "bind" });
292
+ }
293
+ if (!Number.isSafeInteger(n) || n < 1 || n > bound.capacity) {
294
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "gridBuild: n must be an integer in [1, the capacity]", {
295
+ argument: "n",
296
+ value: n,
297
+ expected: bound.capacity,
298
+ });
299
+ }
300
+ const stop = upTo ?? "G3";
301
+ const b = bound.bindings;
302
+ const keyBound = this.cellKey.bind({
303
+ pos: b.pos,
304
+ S: b.state,
305
+ cellKey: b.cellKey,
306
+ cellVal: b.cellVal,
307
+ P: b.params,
308
+ });
309
+ this.cellKey.dispatch(pass, keyBound, plan1d(n, this.scope.workgroupSize, this.scope.caps), [paramsOffset]);
310
+ this.dispatches = 1;
311
+ if (stop === "G1") {
312
+ return;
313
+ }
314
+ const { histWords: bins } = this.spec;
315
+ if (this.path.kind === "counting") {
316
+ const { counting } = this.path;
317
+ const scratch = { hist: b.cellHist, cursor: bound.scratchA };
318
+ counting.record(pass, b.cellKey, n, bins, scratch, b.sortedIdx, b.cellStart);
319
+ this.dispatches += counting.lastDispatches;
320
+ return;
321
+ }
322
+ const { radix, histogram, scan } = this.path;
323
+ radix.record(pass, b.cellKey, b.cellVal, n, GRID_SORT_BITS, {
324
+ keys: b.sortedKey,
325
+ vals: b.sortedIdx,
326
+ hist: bound.scratchA,
327
+ offsets: bound.scratchB,
328
+ });
329
+ this.dispatches += radix.lastDispatches;
330
+ if (stop === "G2") {
331
+ return;
332
+ }
333
+ histogram.record(pass, b.cellKey, n, bins, b.cellHist);
334
+ scan.record(pass, b.cellHist, bins, b.cellStart);
335
+ this.dispatches += histogram.lastDispatches + scan.lastDispatches;
336
+ }
337
+
338
+ /**
339
+ * A scratch of `words` u32 from the scope, bound whole.
340
+ * @param words - the word count (>= 1)
341
+ * @param label - the scratch label
342
+ * @returns the binding
343
+ */
344
+ private scratch(words: number, label: string): Binding {
345
+ const size = 4 * Math.max(1, words);
346
+ const buffer = this.scope.scratch(size, label);
347
+ return { buffer, offset: 0, size, window: null };
348
+ }
349
+ }
@@ -0,0 +1,273 @@
1
+ /**
2
+ * The `histogram` and `countingSortByKey` primitive drivers (spec 6 row 5; P4-T3, PD-4). `histogram` zeroes `hist`
3
+ * with a `fill` dispatch and then adds one per key in global memory: order-independent, so bitwise deterministic on
4
+ * every adapter. `countingSortByKey` is the histogram, an exclusive scan of it into `outStart`, a zeroed per-bin
5
+ * `cursor`, and the scatter `outIndex[outStart[k] + atomicAdd(&cursor[k], 1)] = i`: the SET of indices inside a bin
6
+ * is fixed, their order follows the schedule (set-deterministic, design 6). Every zeroing is a dispatch inside the
7
+ * caller's pass (PD-12), never an encoder clear.
8
+ *
9
+ * The drivers own no device objects: the caller supplies a ReduceScope (the same record `reduce` takes) and the
10
+ * compute pass to record into. `src/primitives/**` never imports `src/context.ts`.
11
+ */
12
+
13
+ import { WebGpuGraphError } from "../errors.js";
14
+ import { plan1d } from "../kernel/dispatch.js";
15
+ import { type Kernel } from "../kernel/kernel.js";
16
+ import { FILL_PARAMS, HIST_PARAMS, kernelSpec } from "../kernels.js";
17
+ import { type Binding } from "../types/memory.js";
18
+ import { type ReduceScope } from "./reduce.js";
19
+ import { prepareScan, type ScanPlanner } from "./scan.js";
20
+
21
+ /** A prepared histogram (spec 6 row 5): records the fill and the one counting dispatch into a pass. */
22
+ export interface HistogramPlanner {
23
+ /**
24
+ * Records the histogram of `count` u32 keys of `keys` over `bins` bins into `hist` (zeroed first by a `fill`
25
+ * dispatch); a key >= bins is not counted. For count 0 only the fill is recorded, so `hist` is all zero.
26
+ * @param pass - the compute pass
27
+ * @param keys - the keys (at least 4 x count bytes)
28
+ * @param count - the key count (a non-negative integer below 2^32)
29
+ * @param bins - the bin count (an integer in [1, 2^32))
30
+ * @param hist - the counts (at least 4 x bins bytes)
31
+ */
32
+ record(pass: GPUComputePassEncoder, keys: Binding, count: number, bins: number, hist: Binding): void;
33
+ /** Dispatches the last record() issued: 1 (the fill) for count 0, else 2. */
34
+ readonly lastDispatches: number;
35
+ }
36
+
37
+ /** The two per-bin scratch arrays of a counting sort: the histogram and the scatter cursor, each at least 4 x bins bytes. */
38
+ interface CountingSortScratch {
39
+ readonly hist: Binding;
40
+ readonly cursor: Binding;
41
+ }
42
+
43
+ /** A prepared counting sort by key (spec 6 row 5): records the histogram, the scan, the cursor fill and the scatter into a pass. */
44
+ export interface CountingSortPlanner {
45
+ /**
46
+ * Records the counting sort of `count` keys of `keys` (every key < bins) into `outIndex` (the indices in key
47
+ * order; the order inside a bin follows the schedule) and `outStart` (the exclusive scan of the histogram, so bin
48
+ * k holds `outIndex[outStart[k] .. outStart[k + 1])` when outStart has bins + 1 words, as the grid's does).
49
+ * For count 0 the two fills and the scan run and `outStart` is all zero.
50
+ * @param pass - the compute pass
51
+ * @param keys - the keys (at least 4 x count bytes)
52
+ * @param count - the key count (a non-negative integer below 2^32)
53
+ * @param bins - the bin count (an integer in [1, 2^32))
54
+ * @param scratch - the per-bin histogram and cursor (each at least 4 x bins bytes)
55
+ * @param outIndex - the sorted indices (at least 4 x count bytes)
56
+ * @param outStart - the bin starts (at least 4 x bins bytes)
57
+ */
58
+ record(
59
+ pass: GPUComputePassEncoder,
60
+ keys: Binding,
61
+ count: number,
62
+ bins: number,
63
+ scratch: CountingSortScratch,
64
+ outIndex: Binding,
65
+ outStart: Binding,
66
+ ): void;
67
+ /** Dispatches the last record() issued: the histogram's + the scan's + 1 (the cursor fill) + 1 (the scatter, count > 0 only). */
68
+ readonly lastDispatches: number;
69
+ }
70
+
71
+ /** The largest u32: the largest value `HistParams.count` / `bins` can carry. */
72
+ const U32_MAX = 0xffffffff;
73
+
74
+ /**
75
+ * Prepares the histogram pipelines of a scope (compiles `histogram` and `fill` once) so record() is synchronous.
76
+ * @param scope - the caller's scope (device, caps, cache, scratch, params)
77
+ * @returns the planner
78
+ */
79
+ export async function prepareHistogram(scope: ReduceScope): Promise<HistogramPlanner> {
80
+ return prepareHistogramImpl(scope);
81
+ }
82
+
83
+ /**
84
+ * The concrete histogram planner (its recordZero is what the counting sort reuses for the cursor).
85
+ * @param scope - the caller's scope
86
+ * @returns the planner
87
+ */
88
+ async function prepareHistogramImpl(scope: ReduceScope): Promise<HistogramPlannerImpl> {
89
+ const histogram = await scope.pipelines.kernel(kernelSpec("histogram"));
90
+ const fill = await scope.pipelines.kernel(kernelSpec("fill"));
91
+ return new HistogramPlannerImpl(scope, histogram, fill);
92
+ }
93
+
94
+ /**
95
+ * Prepares the counting-sort pipelines of a scope (the histogram's, the scan's and `counting-scatter`) so record()
96
+ * is synchronous. The planner lives exactly as long as the scope (the scan keeps a scratch word of it).
97
+ * @param scope - the caller's scope
98
+ * @returns the planner
99
+ */
100
+ export async function prepareCountingSort(scope: ReduceScope): Promise<CountingSortPlanner> {
101
+ const histogram = await prepareHistogramImpl(scope);
102
+ const scan = await prepareScan(scope);
103
+ const scatter = await scope.pipelines.kernel(kernelSpec("counting-scatter"));
104
+ return new CountingSortPlannerImpl(scope, histogram, scan, scatter);
105
+ }
106
+
107
+ /**
108
+ * The E_INVALID_ARGUMENT of a binding shorter than `words` u32.
109
+ * @param name - the argument name
110
+ * @param binding - the binding
111
+ * @param words - the words it must hold
112
+ */
113
+ function checkWords(name: string, binding: Binding, words: number): void {
114
+ if (binding.size < 4 * words) {
115
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", `histogram: ${name} is smaller than 4 x ${words} bytes`, {
116
+ argument: name,
117
+ value: binding.size,
118
+ expected: 4 * words,
119
+ });
120
+ }
121
+ }
122
+
123
+ /**
124
+ * The argument checks shared by both record() methods (E_INVALID_ARGUMENT before anything is recorded).
125
+ * @param keys - the keys binding
126
+ * @param count - the key count
127
+ * @param bins - the bin count
128
+ * @param hist - the histogram binding
129
+ */
130
+ function checkHistogramArguments(keys: Binding, count: number, bins: number, hist: Binding): void {
131
+ if (!Number.isSafeInteger(count) || count < 0 || count > U32_MAX) {
132
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "histogram: count must be a non-negative integer below 2^32", {
133
+ argument: "count",
134
+ value: count,
135
+ });
136
+ }
137
+ if (!Number.isSafeInteger(bins) || bins < 1 || bins > U32_MAX) {
138
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "histogram: bins must be an integer in [1, 2^32)", {
139
+ argument: "bins",
140
+ value: bins,
141
+ });
142
+ }
143
+ checkWords("keys", keys, count);
144
+ checkWords("hist", hist, bins);
145
+ }
146
+
147
+ /** The planner: the histogram and fill kernels over one scope. */
148
+ class HistogramPlannerImpl implements HistogramPlanner {
149
+ private readonly scope: ReduceScope;
150
+ private readonly histogram: Kernel;
151
+ private readonly fill: Kernel;
152
+ private dispatches = 0;
153
+
154
+ /**
155
+ * Wraps the resolved kernels; use prepareHistogram().
156
+ * @param scope - the caller's scope
157
+ * @param histogram - the `histogram` kernel
158
+ * @param fill - the `fill` kernel
159
+ */
160
+ constructor(scope: ReduceScope, histogram: Kernel, fill: Kernel) {
161
+ this.scope = scope;
162
+ this.histogram = histogram;
163
+ this.fill = fill;
164
+ }
165
+
166
+ /**
167
+ * Dispatches the last record() issued.
168
+ * @returns the count
169
+ */
170
+ get lastDispatches(): number {
171
+ return this.dispatches;
172
+ }
173
+
174
+ /**
175
+ * Records the fill and the histogram (see the interface).
176
+ * @param pass - the compute pass
177
+ * @param keys - the keys
178
+ * @param count - the key count
179
+ * @param bins - the bin count
180
+ * @param hist - the counts
181
+ */
182
+ record(pass: GPUComputePassEncoder, keys: Binding, count: number, bins: number, hist: Binding): void {
183
+ checkHistogramArguments(keys, count, bins, hist);
184
+ this.recordZero(pass, hist, bins);
185
+ this.dispatches = 1;
186
+ if (count === 0) {
187
+ return;
188
+ }
189
+ const params = this.scope.params(HIST_PARAMS, { count, bins, pad0: 0, pad1: 0 });
190
+ const bound = this.histogram.bind({ keys, hist, P: params.binding });
191
+ this.histogram.dispatch(pass, bound, plan1d(count, this.scope.workgroupSize, this.scope.caps), [params.offset]);
192
+ this.dispatches = 2;
193
+ }
194
+
195
+ /**
196
+ * Records a `fill` of `words` zeros into `dst` (PD-12: a dispatch inside the pass, never an encoder clear).
197
+ * @param pass - the compute pass
198
+ * @param dst - the words to zero
199
+ * @param words - how many (>= 1)
200
+ */
201
+ recordZero(pass: GPUComputePassEncoder, dst: Binding, words: number): void {
202
+ const params = this.scope.params(FILL_PARAMS, { count: words, value: 0, mode: 0, pad0: 0 });
203
+ const bound = this.fill.bind({ dst, P: params.binding });
204
+ this.fill.dispatch(pass, bound, plan1d(words, this.scope.workgroupSize, this.scope.caps), [params.offset]);
205
+ }
206
+ }
207
+
208
+ /** The planner: the histogram planner, the scan planner and the scatter kernel over one scope. */
209
+ class CountingSortPlannerImpl implements CountingSortPlanner {
210
+ private readonly scope: ReduceScope;
211
+ private readonly histogram: HistogramPlannerImpl;
212
+ private readonly scan: ScanPlanner;
213
+ private readonly scatter: Kernel;
214
+ private dispatches = 0;
215
+
216
+ /**
217
+ * Wraps the resolved planners and kernel; use prepareCountingSort().
218
+ * @param scope - the caller's scope
219
+ * @param histogram - the histogram planner (also the zeroing of the cursor)
220
+ * @param scan - the scan planner
221
+ * @param scatter - the `counting-scatter` kernel
222
+ */
223
+ constructor(scope: ReduceScope, histogram: HistogramPlannerImpl, scan: ScanPlanner, scatter: Kernel) {
224
+ this.scope = scope;
225
+ this.histogram = histogram;
226
+ this.scan = scan;
227
+ this.scatter = scatter;
228
+ }
229
+
230
+ /**
231
+ * Dispatches the last record() issued.
232
+ * @returns the count
233
+ */
234
+ get lastDispatches(): number {
235
+ return this.dispatches;
236
+ }
237
+
238
+ /**
239
+ * Records the four stages (see the interface).
240
+ * @param pass - the compute pass
241
+ * @param keys - the keys
242
+ * @param count - the key count
243
+ * @param bins - the bin count
244
+ * @param scratch - the histogram and cursor
245
+ * @param outIndex - the sorted indices
246
+ * @param outStart - the bin starts
247
+ */
248
+ record(
249
+ pass: GPUComputePassEncoder,
250
+ keys: Binding,
251
+ count: number,
252
+ bins: number,
253
+ scratch: CountingSortScratch,
254
+ outIndex: Binding,
255
+ outStart: Binding,
256
+ ): void {
257
+ checkHistogramArguments(keys, count, bins, scratch.hist);
258
+ checkWords("cursor", scratch.cursor, bins);
259
+ checkWords("outIndex", outIndex, count);
260
+ checkWords("outStart", outStart, bins);
261
+ this.histogram.record(pass, keys, count, bins, scratch.hist);
262
+ this.scan.record(pass, scratch.hist, bins, outStart);
263
+ this.histogram.recordZero(pass, scratch.cursor, bins);
264
+ this.dispatches = this.histogram.lastDispatches + this.scan.lastDispatches + 1;
265
+ if (count === 0) {
266
+ return;
267
+ }
268
+ const params = this.scope.params(HIST_PARAMS, { count, bins, pad0: 0, pad1: 0 });
269
+ const bound = this.scatter.bind({ keys, start: outStart, cursor: scratch.cursor, outIndex, P: params.binding });
270
+ this.scatter.dispatch(pass, bound, plan1d(count, this.scope.workgroupSize, this.scope.caps), [params.offset]);
271
+ this.dispatches += 1;
272
+ }
273
+ }