@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
@@ -1,15 +1,110 @@
1
1
  /**
2
- * The option and value helpers every force model's resolver and stats decoder share (PD-7): moved verbatim from
3
- * src/layouts/forceatlas2.ts (P3-T2) so the Fruchterman-Reingold and spring-electrical models of P5 neither copy
4
- * them nor import a sibling model. Layout zone; imports errors.ts only.
2
+ * The option and value helpers every force model's resolver and stats decoder share (P5 PD-7): moved verbatim
3
+ * from src/layouts/forceatlas2.ts (P3-T2) so the Fruchterman-Reingold and spring-electrical models of P5 neither
4
+ * copy them nor import a sibling model; and, since P4-T6, the K2 tier binding and dispatch the three models share
5
+ * (`bindAttraction` / `recordAttraction`, P4 PD-7), so every model's bind() / recordIteration() calls one function
6
+ * instead of holding its own copy. Layout zone.
5
7
  */
6
8
 
7
9
  import { WebGpuGraphError } from "../errors.js";
10
+ import { type DispatchPlan, plan1d } from "../kernel/dispatch.js";
11
+ import { type BoundKernel, type Kernel } from "../kernel/kernel.js";
8
12
  import { type UniformValues } from "../kernel/struct-block.js";
13
+ import { graphBindings, kernelSpec } from "../kernels.js";
14
+ import { MID_TIER_LANES } from "../primitives/core-shape.js";
15
+ import { type Binding } from "../types/memory.js";
16
+ import { type ModelResources } from "./force-simulation.js";
9
17
 
10
18
  /** An override record as the kernel layer takes it. */
11
19
  export type Overrides = Readonly<Record<string, number | boolean>>;
12
20
 
21
+ /** The K2 (`fa2-attraction`) dispatches of one iteration: `[kernel, bind group, plan]` in dispatch order TIER 2, TIER 1, TIER 0 (only the tiers whose row range is non-empty). */
22
+ export interface AttractionBound {
23
+ readonly kernels: readonly (readonly [Kernel, BoundKernel, DispatchPlan])[];
24
+ }
25
+
26
+ /** The group-1 / group-2 bindings of K2 every tier dispatch shares: `pos` (vec4f, mass in w), `force` (stride-3 f32) and the Fa2Params slot of the UniformRing. @public the bindings parameter of bindAttraction */
27
+ export interface AttractionBindings {
28
+ readonly pos: Binding;
29
+ readonly force: Binding;
30
+ readonly params: Binding;
31
+ }
32
+
33
+ /**
34
+ * Compiles (through the cache) and binds the `fa2-attraction` pipelines a load needs against the graph group and
35
+ * { pos, force, P } (P4 PD-7): TIER 0 always; TIER 1 when a row of degree 32..1023 exists (`[hiEnd, midEnd)` is
36
+ * non-empty); TIER 2 when a row of degree >= 1024 exists (`hiEnd > 0`). The dispatch plans are one workgroup per
37
+ * row for TIER 2, `WG / 32` rows per workgroup for TIER 1 and one row per thread for TIER 0, over each tier's row
38
+ * count (the K2 body reads its range from `Fa2Params.hiEnd` / `midEnd` / `tierStart` / `tierEnd`). The TIER 1 / 2
39
+ * pipelines compile on the first load that needs them (a one-time cost at that load); a model's `specs()` lists
40
+ * the TIER 0 spec only, because it has no `n` to know which tiers a load needs. A device whose workgroup size is
41
+ * below 32 cannot fold the mid tier: E_UNSUPPORTED { feature: "fa2-attraction.tiers" } when a permutation is bound.
42
+ * @param resources - the load's resources (core, tiers, weights, pipelines, caps)
43
+ * @param k2 - the K2 override record of the model (LINLOG / DISTRIBUTED / LAW plus USE_PERM / HAS_WEIGHTS; TIER is overwritten per dispatch)
44
+ * @param bindings - the group-1 / group-2 bindings shared by every tier
45
+ * @returns the bound dispatches in order TIER 2, 1, 0
46
+ */
47
+ export async function bindAttraction(
48
+ resources: ModelResources,
49
+ k2: Overrides,
50
+ bindings: AttractionBindings,
51
+ ): Promise<AttractionBound> {
52
+ const { n, core, perm, tiers, weights, pipelines, caps } = resources;
53
+ const so = tiers?.segmentOffsets;
54
+ const hiEnd = so?.[1] ?? 0;
55
+ const midEnd = so?.[2] ?? 0;
56
+ const ranges: readonly { readonly tier: 0 | 1 | 2; readonly rows: number }[] = [
57
+ { tier: 2, rows: hiEnd },
58
+ { tier: 1, rows: midEnd - hiEnd },
59
+ { tier: 0, rows: n - midEnd },
60
+ ];
61
+ const group = {
62
+ ...graphBindings(core, perm, weights),
63
+ pos: bindings.pos,
64
+ force: bindings.force,
65
+ P: bindings.params,
66
+ };
67
+ const kernels: (readonly [Kernel, BoundKernel, DispatchPlan])[] = [];
68
+ for (const { tier, rows } of ranges) {
69
+ if (tier !== 0 && rows <= 0) {
70
+ continue;
71
+ }
72
+ // sequential on purpose: PipelineCache.get compiles inside a validation scope, one stack per device
73
+ const kernel = await pipelines.kernel(kernelSpec("fa2-attraction", { ...k2, TIER: tier }));
74
+ const wg = kernel.workgroupSize;
75
+ if (tier !== 0 && wg < MID_TIER_LANES) {
76
+ throw new WebGpuGraphError(
77
+ "E_UNSUPPORTED",
78
+ `fa2-attraction: the mid tier folds ${MID_TIER_LANES} lanes per row, more than the workgroup size ${wg}`,
79
+ { feature: "fa2-attraction.tiers" },
80
+ );
81
+ }
82
+ if (rows <= 0) {
83
+ continue;
84
+ }
85
+ let rowsPerGroup = wg;
86
+ if (tier === 2) {
87
+ rowsPerGroup = 1;
88
+ } else if (tier === 1) {
89
+ rowsPerGroup = wg / MID_TIER_LANES;
90
+ }
91
+ kernels.push([kernel, kernel.bind(group), plan1d(rows, rowsPerGroup, caps)]);
92
+ }
93
+ return { kernels };
94
+ }
95
+
96
+ /**
97
+ * Records the K2 dispatches of one iteration in order TIER 2, TIER 1, TIER 0 with the iteration's params offset.
98
+ * @param pass - the open compute pass
99
+ * @param bound - what bindAttraction produced
100
+ * @param paramsOffset - the UniformRing byte offset of this iteration's Fa2Params
101
+ */
102
+ export function recordAttraction(pass: GPUComputePassEncoder, bound: AttractionBound, paramsOffset: number): void {
103
+ for (const [kernel, group, plan] of bound.kernels) {
104
+ kernel.dispatch(pass, group, plan, [paramsOffset]);
105
+ }
106
+ }
107
+
13
108
  /** Bytes of the stride-3 f32 force arrays per node. */
14
109
  export const FORCE_BYTES_PER_NODE = 12;
15
110
 
@@ -0,0 +1,451 @@
1
+ /**
2
+ * The grid-tier repulsion stage of the force models (spec 7.7, 7.4; P4-T10): G1-G3 through the T8 grid build, G4-G5
3
+ * through the T9 pyramid, then G6 (`grid-far-field`) and G7 (`grid-near-field`, K3's pair law and epilogue) over
4
+ * `plan1d(n)`, followed by K4 (`fa2-speed-finalize`) exactly as `RepulsionExact` records it. The named grid buffers
5
+ * are the model's `BufferSpec`s (`buffers()`, so `inspect(name)` reaches them); the anonymous scratch of the sort
6
+ * and the scan and the static params of the pyramid draw on ONE `Lease` of the context's pool taken at `create()`
7
+ * (the scan planner takes its block-sum levels at prepare time) and released by `dispose()` (PD-11, DEP-P4-D: a
8
+ * per-batch lease would rebuild the G2-G4 bind groups every batch; a model creates one stage per bind(), so a stage
9
+ * is bound once). The stage is the ONE class every model reaches the grid tier through (PD-22): the FR and
10
+ * spring-electrical models pass `LAW` 1 / 2 in their override set, which reaches G6's per-cell law and G7's pair law
11
+ * (P4-T13).
12
+ */
13
+
14
+ import { GRID_HUB_CELL } from "../constants.js";
15
+ import { BufferUsage } from "../device/webgpu-constants.js";
16
+ import { WebGpuGraphError } from "../errors.js";
17
+ import { plan1d } from "../kernel/dispatch.js";
18
+ import { type BoundKernel, INDIRECT_ARGS_STRIDE, type Kernel } from "../kernel/kernel.js";
19
+ import { type PipelineCache } from "../kernel/pipeline-cache.js";
20
+ import { type UniformBlock, type UniformValues } from "../kernel/struct-block.js";
21
+ import { type WgslModuleSpec } from "../kernel/wgsl.js";
22
+ import { kernelSpec } from "../kernels.js";
23
+ import { type BufferPool } from "../memory/buffer-pool.js";
24
+ import { type Lease } from "../memory/lease.js";
25
+ import {
26
+ type GridBuildPlanner,
27
+ type GridBuildStage,
28
+ gridPyramidBytes,
29
+ type GridSpec,
30
+ prepareGridBuild,
31
+ } from "../primitives/grid.js";
32
+ import { type GridPyramidPlanner, preparePyramid } from "../primitives/grid-pyramid.js";
33
+ import { type ReduceScope } from "../primitives/reduce.js";
34
+ import { type PlanCaps } from "../types/context.js";
35
+ import { type Binding } from "../types/memory.js";
36
+ import { type BufferSpec } from "./force-simulation.js";
37
+ import { type RepulsionExactOverrides, type RepulsionExactResources } from "./repulsion-exact.js";
38
+
39
+ /** The stage names the grid tier records, in dispatch order (the `upTo` vocabulary of recordRepulsion). */
40
+ const GRID_STAGES = ["G1", "G2", "G3", "G4", "G5", "G6", "G7"] as const;
41
+
42
+ /** A grid stage name. */
43
+ export type GridStage = (typeof GRID_STAGES)[number];
44
+
45
+ /**
46
+ * The buffers the grid-tier repulsion stage binds: the exact tier's plus the named grid buffers of `buffers()`. The
47
+ * parameter type of RepulsionGrid.bind (knip: exported for the signature, not imported by name).
48
+ * @public
49
+ */
50
+ export interface RepulsionGridResources extends RepulsionExactResources {
51
+ readonly cellKey: Binding;
52
+ readonly cellVal: Binding;
53
+ readonly sortedKey: Binding;
54
+ readonly sortedIdx: Binding;
55
+ readonly cellHist: Binding;
56
+ readonly cellStart: Binding;
57
+ readonly hubList: Binding;
58
+ readonly hubCounters: Binding;
59
+ readonly hubArgs: Binding;
60
+ readonly pyramid: Binding;
61
+ }
62
+
63
+ /**
64
+ * The override set G6 / G7 / K4 compile with: K3's three plus the repulsion law (0 FA2, 1 FR, 2 coulomb; P4-T13,
65
+ * PD-22).
66
+ * @public
67
+ */
68
+ export interface RepulsionGridOverrides extends RepulsionExactOverrides {
69
+ readonly LAW: 0 | 1 | 2;
70
+ }
71
+
72
+ /**
73
+ * What the stage needs of the context to build its scope: the pieces `ModelResources` carries. The parameter type
74
+ * of RepulsionGrid.create (knip: exported for the signature, not imported by name).
75
+ * @public
76
+ */
77
+ export interface RepulsionGridScope {
78
+ readonly device: GPUDevice;
79
+ readonly caps: PlanCaps;
80
+ readonly pipelines: PipelineCache;
81
+ readonly pool: BufferPool;
82
+ }
83
+
84
+ /** The storage usage of every grid buffer. */
85
+ const STORAGE_RW = BufferUsage.STORAGE | BufferUsage.COPY_SRC | BufferUsage.COPY_DST;
86
+
87
+ /** The lease the scope's scratch() and params() draw on: taken at create(), null after dispose(). */
88
+ interface LeaseBox {
89
+ lease: Lease | null;
90
+ }
91
+
92
+ /**
93
+ * The stage's lease, or E_DISPOSED after dispose().
94
+ * @param box - the stage's lease box
95
+ * @returns the lease
96
+ */
97
+ function leaseOf(box: LeaseBox): Lease {
98
+ if (box.lease === null) {
99
+ throw new WebGpuGraphError("E_DISPOSED", "RepulsionGrid: the stage was disposed", { label: "RepulsionGrid" });
100
+ }
101
+ return box.lease;
102
+ }
103
+
104
+ /**
105
+ * A uniform buffer of the lease holding one written record of `block` (the pyramid's static params, PD-11).
106
+ * @param device - the device
107
+ * @param lease - the stage's lease
108
+ * @param block - the uniform block
109
+ * @param values - the values to write
110
+ * @returns the whole-buffer binding and a zero dynamic offset
111
+ */
112
+ function writeParams(
113
+ device: GPUDevice,
114
+ lease: Lease,
115
+ block: UniformBlock,
116
+ values: UniformValues,
117
+ ): { readonly binding: Binding; readonly offset: number } {
118
+ const buffer = lease.uniform(block.byteLength, `grid/${block.name}`);
119
+ const bytes = new ArrayBuffer(block.byteLength);
120
+ block.write(new DataView(bytes), values);
121
+ device.queue.writeBuffer(buffer, 0, bytes);
122
+ return { binding: { buffer, offset: 0, size: block.byteLength, window: null }, offset: 0 };
123
+ }
124
+
125
+ /**
126
+ * The G6 spec of an override set (the law alone).
127
+ * @param overrides - the override set of the stage
128
+ * @returns the spec
129
+ */
130
+ function farFieldSpec(overrides: RepulsionGridOverrides): WgslModuleSpec {
131
+ return kernelSpec("grid-far-field", { LAW: overrides.LAW });
132
+ }
133
+
134
+ /**
135
+ * The G7 spec of an override set (K3's four).
136
+ * @param overrides - the override set of the stage
137
+ * @returns the spec
138
+ */
139
+ function nearFieldSpec(overrides: RepulsionGridOverrides): WgslModuleSpec {
140
+ return kernelSpec("grid-near-field", {
141
+ SWING_MODE: overrides.SWING_MODE,
142
+ STRONG_GRAVITY: overrides.STRONG_GRAVITY,
143
+ GRAVITY_CENTER: overrides.GRAVITY_CENTER,
144
+ LAW: overrides.LAW,
145
+ });
146
+ }
147
+
148
+ /** The three kernels the stage dispatches itself (the planners hold the others). */
149
+ interface FieldKernels {
150
+ readonly far: Kernel;
151
+ readonly near: Kernel;
152
+ readonly speedFinalize: Kernel;
153
+ }
154
+
155
+ /** What bind() produced. */
156
+ interface Bound {
157
+ readonly far: BoundKernel;
158
+ readonly near: BoundKernel;
159
+ readonly speedFinalize: BoundKernel;
160
+ }
161
+
162
+ /** G1-G7 then K4 (spec 7.7, 7.4): the grid tier's repulsion stage. */
163
+ export class RepulsionGrid {
164
+ /** The overrides G6 / G7 / K4 were compiled with (a frozen copy of the argument of create()). */
165
+ readonly overrides: RepulsionGridOverrides;
166
+ /** The grid geometry the stage was prepared for. */
167
+ readonly spec: GridSpec;
168
+
169
+ private readonly caps: PlanCaps;
170
+ private readonly workgroupSize: number;
171
+ private readonly box: LeaseBox;
172
+ private readonly build: GridBuildPlanner;
173
+ private readonly pyramid: GridPyramidPlanner;
174
+ private readonly far: Kernel;
175
+ private readonly near: Kernel;
176
+ private readonly speedFinalize: Kernel;
177
+ private bound: Bound | null = null;
178
+
179
+ /**
180
+ * Holds the planners and kernels; create() is the only caller.
181
+ * @param scope - the context pieces
182
+ * @param workgroupSize - the device's workgroup size (every kernel compiles with it)
183
+ * @param overrides - the override set G6 / G7 / K4 were compiled with
184
+ * @param spec - the grid
185
+ * @param box - the lease box the scope's scratch() and params() read
186
+ * @param build - the T8 planner (G1-G3)
187
+ * @param pyramid - the T9 planner (G4-G5)
188
+ * @param kernels - G6 (`far`), G7 (`near`) and K4 (`speedFinalize`)
189
+ */
190
+ private constructor(
191
+ scope: RepulsionGridScope,
192
+ workgroupSize: number,
193
+ overrides: RepulsionGridOverrides,
194
+ spec: GridSpec,
195
+ box: LeaseBox,
196
+ build: GridBuildPlanner,
197
+ pyramid: GridPyramidPlanner,
198
+ kernels: FieldKernels,
199
+ ) {
200
+ this.caps = scope.caps;
201
+ this.box = box;
202
+ this.workgroupSize = workgroupSize;
203
+ this.overrides = Object.freeze({
204
+ SWING_MODE: overrides.SWING_MODE,
205
+ STRONG_GRAVITY: overrides.STRONG_GRAVITY,
206
+ GRAVITY_CENTER: overrides.GRAVITY_CENTER,
207
+ LAW: overrides.LAW,
208
+ });
209
+ this.spec = spec;
210
+ this.build = build;
211
+ this.pyramid = pyramid;
212
+ this.far = kernels.far;
213
+ this.near = kernels.near;
214
+ this.speedFinalize = kernels.speedFinalize;
215
+ }
216
+
217
+ /**
218
+ * The model-owned buffers of the grid tier (spec 7.3; PD-11): `cellKey` / `cellVal` / `sortedKey` / `sortedIdx`
219
+ * 4n, `cellHist` / `cellStart` 4 (cells + 2) zeroed, `hubList` one word per possible hub cell, `hubArgs` one
220
+ * indirect slot, `pyramid` 16 B per pyramid cell zeroed. `hubCounters` (16 B, zeroed) is the MODEL's on every
221
+ * tier (PD-14: K1 binds it on the exact tier too). n = 0 reports one node's worth of bytes (spec 3.6).
222
+ * @param n - the node count
223
+ * @param spec - the grid
224
+ * @returns the specs
225
+ */
226
+ static buffers(n: number, spec: GridSpec): readonly BufferSpec[] {
227
+ const words = 4 * Math.max(1, n);
228
+ return [
229
+ { name: "cellKey", byteLength: words, usage: STORAGE_RW, zero: false },
230
+ { name: "cellVal", byteLength: words, usage: STORAGE_RW, zero: false },
231
+ { name: "sortedKey", byteLength: words, usage: STORAGE_RW, zero: false },
232
+ { name: "sortedIdx", byteLength: words, usage: STORAGE_RW, zero: false },
233
+ { name: "cellHist", byteLength: 4 * spec.histWords, usage: STORAGE_RW, zero: true },
234
+ { name: "cellStart", byteLength: 4 * spec.histWords, usage: STORAGE_RW, zero: true },
235
+ {
236
+ name: "hubList",
237
+ byteLength: 4 * Math.max(1, Math.ceil(n / GRID_HUB_CELL)),
238
+ usage: STORAGE_RW,
239
+ zero: false,
240
+ },
241
+ { name: "hubArgs", byteLength: INDIRECT_ARGS_STRIDE, usage: STORAGE_RW | BufferUsage.INDIRECT, zero: true },
242
+ { name: "pyramid", byteLength: gridPyramidBytes(spec), usage: STORAGE_RW, zero: true },
243
+ ];
244
+ }
245
+
246
+ /**
247
+ * The module specs of the grid tier under an override set (for warm() and the compile matrix): the build's
248
+ * (G1, the sort path `spec.deterministic` selects, the histogram, the scan, the fill), the pyramid's (G4, the
249
+ * finalize, G4b, G5), G6, G7 and K4.
250
+ * @param overrides - the override set of the stage
251
+ * @param spec - the grid (only its `deterministic` flag matters: the pipeline key carries no geometry)
252
+ * @returns the specs
253
+ */
254
+ static specs(overrides: RepulsionGridOverrides, spec: GridSpec): readonly WgslModuleSpec[] {
255
+ const sort: WgslModuleSpec[] = spec.deterministic
256
+ ? [kernelSpec("radix-hist"), kernelSpec("radix-scatter"), kernelSpec("scan-block"), kernelSpec("scan-add")]
257
+ : [kernelSpec("counting-scatter"), kernelSpec("scan-block"), kernelSpec("scan-add")];
258
+ return [
259
+ kernelSpec("grid-cell-key"),
260
+ ...sort,
261
+ kernelSpec("histogram"),
262
+ kernelSpec("fill"),
263
+ kernelSpec("grid-centroid"),
264
+ kernelSpec("indirect-finalize"),
265
+ kernelSpec("grid-centroid-hub"),
266
+ kernelSpec("grid-downsample"),
267
+ farFieldSpec(overrides),
268
+ nearFieldSpec(overrides),
269
+ kernelSpec("fa2-speed-finalize", { SWING_MODE: overrides.SWING_MODE }),
270
+ ];
271
+ }
272
+
273
+ /**
274
+ * Compiles every kernel of the tier through the cache (sequentially: PipelineCache.get compiles inside a
275
+ * validation scope, one stack per device) over a scope whose scratch and params draw on the stage's lease.
276
+ * @param scope - the context pieces (device, caps, pipelines, pool)
277
+ * @param workgroupSize - the device's workgroup size
278
+ * @param overrides - the SWING_MODE / STRONG_GRAVITY / GRAVITY_CENTER / LAW set of this stage
279
+ * @param spec - the grid
280
+ * @returns the stage, ready for bind()
281
+ */
282
+ static async create(
283
+ scope: RepulsionGridScope,
284
+ workgroupSize: number,
285
+ overrides: RepulsionGridOverrides,
286
+ spec: GridSpec,
287
+ ): Promise<RepulsionGrid> {
288
+ const box: LeaseBox = { lease: scope.pool.lease() };
289
+ const reduceScope: ReduceScope = {
290
+ device: scope.device,
291
+ caps: scope.caps,
292
+ pipelines: scope.pipelines,
293
+ pool: scope.pool,
294
+ workgroupSize,
295
+ scratch: (byteLength: number, label: string): GPUBuffer => leaseOf(box).storage(byteLength, label),
296
+ params: (
297
+ block: UniformBlock,
298
+ values: UniformValues,
299
+ ): { readonly binding: Binding; readonly offset: number } =>
300
+ writeParams(scope.device, leaseOf(box), block, values),
301
+ };
302
+ try {
303
+ const build = await prepareGridBuild(reduceScope, spec);
304
+ const pyramid = await preparePyramid(reduceScope, spec);
305
+ const far = await scope.pipelines.kernel(farFieldSpec(overrides));
306
+ const near = await scope.pipelines.kernel(nearFieldSpec(overrides));
307
+ const speedFinalize = await scope.pipelines.kernel(
308
+ kernelSpec("fa2-speed-finalize", { SWING_MODE: overrides.SWING_MODE }),
309
+ );
310
+ return new RepulsionGrid(scope, workgroupSize, overrides, spec, box, build, pyramid, {
311
+ far,
312
+ near,
313
+ speedFinalize,
314
+ });
315
+ } catch (error) {
316
+ leaseOf(box).release();
317
+ throw error;
318
+ }
319
+ }
320
+
321
+ /**
322
+ * Builds every bind group once per load(): the build's and the pyramid's bind() (which take their scratch and
323
+ * write their static params through the lease), then G6, G7 and K4 against the named buffers (the 3.10.1
324
+ * binding names). A second bind() takes fresh scratch from the same lease and the first bind()'s scratch stays
325
+ * held until dispose(): the one lease also holds the prepare-time allocations (the scan's block-sum levels), so
326
+ * it cannot be released on a rebind (the model creates one stage per bind(), so a rebind never happens).
327
+ * @param resources - the buffers of spec 7.3 plus the grid's
328
+ */
329
+ bind(resources: RepulsionGridResources): void {
330
+ leaseOf(this.box);
331
+ const r = resources;
332
+ this.build.bind({
333
+ pos: r.pos,
334
+ state: r.state,
335
+ params: r.params,
336
+ cellKey: r.cellKey,
337
+ cellVal: r.cellVal,
338
+ sortedKey: r.sortedKey,
339
+ sortedIdx: r.sortedIdx,
340
+ cellHist: r.cellHist,
341
+ cellStart: r.cellStart,
342
+ });
343
+ this.pyramid.bind({
344
+ pos: r.pos,
345
+ params: r.params,
346
+ sortedIdx: r.sortedIdx,
347
+ cellStart: r.cellStart,
348
+ pyramid: r.pyramid,
349
+ hubList: r.hubList,
350
+ hubCounters: r.hubCounters,
351
+ hubArgs: r.hubArgs,
352
+ });
353
+ this.bound = {
354
+ far: this.far.bind({
355
+ pos: r.pos,
356
+ sortedIdx: r.sortedIdx,
357
+ pyramid: r.pyramid,
358
+ S: r.state,
359
+ force: r.force,
360
+ P: r.params,
361
+ }),
362
+ near: this.near.bind({
363
+ pos: r.pos,
364
+ sortedIdx: r.sortedIdx,
365
+ cellStart: r.cellStart,
366
+ S: r.state,
367
+ force: r.force,
368
+ oldForce: r.oldForce,
369
+ fixedMask: r.fixedMask,
370
+ partials: r.partials,
371
+ P: r.params,
372
+ }),
373
+ speedFinalize: this.speedFinalize.bind({ partials: r.partials, S: r.state, T: r.trace, P: r.params }),
374
+ };
375
+ }
376
+
377
+ /**
378
+ * Records G1..G7 for `n` nodes with the params slot's dynamic offset, stopping after `upTo` when given (spec 7.4
379
+ * grid sequence; the inspect() stage split, spec 11.9 item 2): G1-G3 through the build planner, G4-G5 through the
380
+ * pyramid planner, then G6 and G7 over plan1d(n).
381
+ * @param pass - the open compute pass of the batch
382
+ * @param n - the node count (in [1, the capacity of the grid buffers])
383
+ * @param paramsOffset - the dynamic offset of this iteration's Fa2Params slot in the uniform ring
384
+ * @param upTo - the last grid stage to record (default "G7")
385
+ */
386
+ recordRepulsion(pass: GPUComputePassEncoder, n: number, paramsOffset: number, upTo?: GridStage): void {
387
+ const bound = this.requireBound("recordRepulsion");
388
+ const stop = GRID_STAGES.indexOf(upTo ?? "G7");
389
+ let buildStop: GridBuildStage = "G3";
390
+ if (stop === 0) {
391
+ buildStop = "G1";
392
+ } else if (stop === 1) {
393
+ buildStop = "G2";
394
+ }
395
+ this.build.record(pass, n, paramsOffset, buildStop);
396
+ if (stop < 3) {
397
+ return;
398
+ }
399
+ this.pyramid.record(pass, paramsOffset, stop === 3 ? "G4" : "G5");
400
+ if (stop < 5) {
401
+ return;
402
+ }
403
+ const plan = plan1d(n, this.workgroupSize, this.caps);
404
+ this.far.dispatch(pass, bound.far, plan, [paramsOffset]);
405
+ if (stop < 6) {
406
+ return;
407
+ }
408
+ this.near.dispatch(pass, bound.near, plan, [paramsOffset]);
409
+ }
410
+
411
+ /**
412
+ * Records K4 only (one workgroup).
413
+ * @param pass - the open compute pass
414
+ * @param paramsOffset - the dynamic offset of the Fa2Params slot
415
+ */
416
+ recordSpeedFinalize(pass: GPUComputePassEncoder, paramsOffset: number): void {
417
+ const bound = this.requireBound("recordSpeedFinalize");
418
+ const one = plan1d(this.speedFinalize.workgroupSize, this.speedFinalize.workgroupSize, this.caps);
419
+ this.speedFinalize.dispatch(pass, bound.speedFinalize, one, [paramsOffset]);
420
+ }
421
+
422
+ /** Releases the lease (the sort and scan scratch, the static params) and drops the bind groups; idempotent. */
423
+ dispose(): void {
424
+ if (this.box.lease !== null) {
425
+ this.box.lease.release();
426
+ this.box.lease = null;
427
+ }
428
+ for (const kernel of [this.far, this.near, this.speedFinalize]) {
429
+ kernel.invalidate();
430
+ }
431
+ this.bound = null;
432
+ }
433
+
434
+ /**
435
+ * The bound groups, or E_NOT_LOADED when bind() has not run.
436
+ * @param method - the caller's name for the message
437
+ * @returns the bound groups
438
+ */
439
+ private requireBound(method: string): Bound {
440
+ if (this.bound === null) {
441
+ throw new WebGpuGraphError(
442
+ "E_NOT_LOADED",
443
+ `RepulsionGrid.${method}(): bind() has not been called for this stage`,
444
+ {
445
+ state: "unbound",
446
+ },
447
+ );
448
+ }
449
+ return this.bound;
450
+ }
451
+ }