@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,190 @@
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
+ import { WebGpuGraphError } from "../errors.js";
13
+ import { plan1d } from "../kernel/dispatch.js";
14
+ import { FILL_PARAMS, HIST_PARAMS, kernelSpec } from "../kernels.js";
15
+ import { prepareScan } from "./scan.js";
16
+ /** The largest u32: the largest value `HistParams.count` / `bins` can carry. */
17
+ const U32_MAX = 0xffffffff;
18
+ /**
19
+ * Prepares the histogram pipelines of a scope (compiles `histogram` and `fill` once) so record() is synchronous.
20
+ * @param scope - the caller's scope (device, caps, cache, scratch, params)
21
+ * @returns the planner
22
+ */
23
+ export async function prepareHistogram(scope) {
24
+ return prepareHistogramImpl(scope);
25
+ }
26
+ /**
27
+ * The concrete histogram planner (its recordZero is what the counting sort reuses for the cursor).
28
+ * @param scope - the caller's scope
29
+ * @returns the planner
30
+ */
31
+ async function prepareHistogramImpl(scope) {
32
+ const histogram = await scope.pipelines.kernel(kernelSpec("histogram"));
33
+ const fill = await scope.pipelines.kernel(kernelSpec("fill"));
34
+ return new HistogramPlannerImpl(scope, histogram, fill);
35
+ }
36
+ /**
37
+ * Prepares the counting-sort pipelines of a scope (the histogram's, the scan's and `counting-scatter`) so record()
38
+ * is synchronous. The planner lives exactly as long as the scope (the scan keeps a scratch word of it).
39
+ * @param scope - the caller's scope
40
+ * @returns the planner
41
+ */
42
+ export async function prepareCountingSort(scope) {
43
+ const histogram = await prepareHistogramImpl(scope);
44
+ const scan = await prepareScan(scope);
45
+ const scatter = await scope.pipelines.kernel(kernelSpec("counting-scatter"));
46
+ return new CountingSortPlannerImpl(scope, histogram, scan, scatter);
47
+ }
48
+ /**
49
+ * The E_INVALID_ARGUMENT of a binding shorter than `words` u32.
50
+ * @param name - the argument name
51
+ * @param binding - the binding
52
+ * @param words - the words it must hold
53
+ */
54
+ function checkWords(name, binding, words) {
55
+ if (binding.size < 4 * words) {
56
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", `histogram: ${name} is smaller than 4 x ${words} bytes`, {
57
+ argument: name,
58
+ value: binding.size,
59
+ expected: 4 * words,
60
+ });
61
+ }
62
+ }
63
+ /**
64
+ * The argument checks shared by both record() methods (E_INVALID_ARGUMENT before anything is recorded).
65
+ * @param keys - the keys binding
66
+ * @param count - the key count
67
+ * @param bins - the bin count
68
+ * @param hist - the histogram binding
69
+ */
70
+ function checkHistogramArguments(keys, count, bins, hist) {
71
+ if (!Number.isSafeInteger(count) || count < 0 || count > U32_MAX) {
72
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "histogram: count must be a non-negative integer below 2^32", {
73
+ argument: "count",
74
+ value: count,
75
+ });
76
+ }
77
+ if (!Number.isSafeInteger(bins) || bins < 1 || bins > U32_MAX) {
78
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "histogram: bins must be an integer in [1, 2^32)", {
79
+ argument: "bins",
80
+ value: bins,
81
+ });
82
+ }
83
+ checkWords("keys", keys, count);
84
+ checkWords("hist", hist, bins);
85
+ }
86
+ /** The planner: the histogram and fill kernels over one scope. */
87
+ class HistogramPlannerImpl {
88
+ /**
89
+ * Wraps the resolved kernels; use prepareHistogram().
90
+ * @param scope - the caller's scope
91
+ * @param histogram - the `histogram` kernel
92
+ * @param fill - the `fill` kernel
93
+ */
94
+ constructor(scope, histogram, fill) {
95
+ this.dispatches = 0;
96
+ this.scope = scope;
97
+ this.histogram = histogram;
98
+ this.fill = fill;
99
+ }
100
+ /**
101
+ * Dispatches the last record() issued.
102
+ * @returns the count
103
+ */
104
+ get lastDispatches() {
105
+ return this.dispatches;
106
+ }
107
+ /**
108
+ * Records the fill and the histogram (see the interface).
109
+ * @param pass - the compute pass
110
+ * @param keys - the keys
111
+ * @param count - the key count
112
+ * @param bins - the bin count
113
+ * @param hist - the counts
114
+ */
115
+ record(pass, keys, count, bins, hist) {
116
+ checkHistogramArguments(keys, count, bins, hist);
117
+ this.recordZero(pass, hist, bins);
118
+ this.dispatches = 1;
119
+ if (count === 0) {
120
+ return;
121
+ }
122
+ const params = this.scope.params(HIST_PARAMS, { count, bins, pad0: 0, pad1: 0 });
123
+ const bound = this.histogram.bind({ keys, hist, P: params.binding });
124
+ this.histogram.dispatch(pass, bound, plan1d(count, this.scope.workgroupSize, this.scope.caps), [params.offset]);
125
+ this.dispatches = 2;
126
+ }
127
+ /**
128
+ * Records a `fill` of `words` zeros into `dst` (PD-12: a dispatch inside the pass, never an encoder clear).
129
+ * @param pass - the compute pass
130
+ * @param dst - the words to zero
131
+ * @param words - how many (>= 1)
132
+ */
133
+ recordZero(pass, dst, words) {
134
+ const params = this.scope.params(FILL_PARAMS, { count: words, value: 0, mode: 0, pad0: 0 });
135
+ const bound = this.fill.bind({ dst, P: params.binding });
136
+ this.fill.dispatch(pass, bound, plan1d(words, this.scope.workgroupSize, this.scope.caps), [params.offset]);
137
+ }
138
+ }
139
+ /** The planner: the histogram planner, the scan planner and the scatter kernel over one scope. */
140
+ class CountingSortPlannerImpl {
141
+ /**
142
+ * Wraps the resolved planners and kernel; use prepareCountingSort().
143
+ * @param scope - the caller's scope
144
+ * @param histogram - the histogram planner (also the zeroing of the cursor)
145
+ * @param scan - the scan planner
146
+ * @param scatter - the `counting-scatter` kernel
147
+ */
148
+ constructor(scope, histogram, scan, scatter) {
149
+ this.dispatches = 0;
150
+ this.scope = scope;
151
+ this.histogram = histogram;
152
+ this.scan = scan;
153
+ this.scatter = scatter;
154
+ }
155
+ /**
156
+ * Dispatches the last record() issued.
157
+ * @returns the count
158
+ */
159
+ get lastDispatches() {
160
+ return this.dispatches;
161
+ }
162
+ /**
163
+ * Records the four stages (see the interface).
164
+ * @param pass - the compute pass
165
+ * @param keys - the keys
166
+ * @param count - the key count
167
+ * @param bins - the bin count
168
+ * @param scratch - the histogram and cursor
169
+ * @param outIndex - the sorted indices
170
+ * @param outStart - the bin starts
171
+ */
172
+ record(pass, keys, count, bins, scratch, outIndex, outStart) {
173
+ checkHistogramArguments(keys, count, bins, scratch.hist);
174
+ checkWords("cursor", scratch.cursor, bins);
175
+ checkWords("outIndex", outIndex, count);
176
+ checkWords("outStart", outStart, bins);
177
+ this.histogram.record(pass, keys, count, bins, scratch.hist);
178
+ this.scan.record(pass, scratch.hist, bins, outStart);
179
+ this.histogram.recordZero(pass, scratch.cursor, bins);
180
+ this.dispatches = this.histogram.lastDispatches + this.scan.lastDispatches + 1;
181
+ if (count === 0) {
182
+ return;
183
+ }
184
+ const params = this.scope.params(HIST_PARAMS, { count, bins, pad0: 0, pad1: 0 });
185
+ const bound = this.scatter.bind({ keys, start: outStart, cursor: scratch.cursor, outIndex, P: params.binding });
186
+ this.scatter.dispatch(pass, bound, plan1d(count, this.scope.workgroupSize, this.scope.caps), [params.offset]);
187
+ this.dispatches += 1;
188
+ }
189
+ }
190
+ //# sourceMappingURL=histogram.js.map
@@ -0,0 +1 @@
1
+ {"version":3,"file":"histogram.js","sourceRoot":"","sources":["../../../src/primitives/histogram.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;GAUG;AAEH,OAAO,EAAE,gBAAgB,EAAE,MAAM,cAAc,CAAC;AAChD,OAAO,EAAE,MAAM,EAAE,MAAM,uBAAuB,CAAC;AAE/C,OAAO,EAAE,WAAW,EAAE,WAAW,EAAE,UAAU,EAAE,MAAM,eAAe,CAAC;AAGrE,OAAO,EAAE,WAAW,EAAoB,MAAM,WAAW,CAAC;AAoD1D,gFAAgF;AAChF,MAAM,OAAO,GAAG,UAAU,CAAC;AAE3B;;;;GAIG;AACH,MAAM,CAAC,KAAK,UAAU,gBAAgB,CAAC,KAAkB;IACrD,OAAO,oBAAoB,CAAC,KAAK,CAAC,CAAC;AACvC,CAAC;AAED;;;;GAIG;AACH,KAAK,UAAU,oBAAoB,CAAC,KAAkB;IAClD,MAAM,SAAS,GAAG,MAAM,KAAK,CAAC,SAAS,CAAC,MAAM,CAAC,UAAU,CAAC,WAAW,CAAC,CAAC,CAAC;IACxE,MAAM,IAAI,GAAG,MAAM,KAAK,CAAC,SAAS,CAAC,MAAM,CAAC,UAAU,CAAC,MAAM,CAAC,CAAC,CAAC;IAC9D,OAAO,IAAI,oBAAoB,CAAC,KAAK,EAAE,SAAS,EAAE,IAAI,CAAC,CAAC;AAC5D,CAAC;AAED;;;;;GAKG;AACH,MAAM,CAAC,KAAK,UAAU,mBAAmB,CAAC,KAAkB;IACxD,MAAM,SAAS,GAAG,MAAM,oBAAoB,CAAC,KAAK,CAAC,CAAC;IACpD,MAAM,IAAI,GAAG,MAAM,WAAW,CAAC,KAAK,CAAC,CAAC;IACtC,MAAM,OAAO,GAAG,MAAM,KAAK,CAAC,SAAS,CAAC,MAAM,CAAC,UAAU,CAAC,kBAAkB,CAAC,CAAC,CAAC;IAC7E,OAAO,IAAI,uBAAuB,CAAC,KAAK,EAAE,SAAS,EAAE,IAAI,EAAE,OAAO,CAAC,CAAC;AACxE,CAAC;AAED;;;;;GAKG;AACH,SAAS,UAAU,CAAC,IAAY,EAAE,OAAgB,EAAE,KAAa;IAC7D,IAAI,OAAO,CAAC,IAAI,GAAG,CAAC,GAAG,KAAK,EAAE,CAAC;QAC3B,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,cAAc,IAAI,wBAAwB,KAAK,QAAQ,EAAE;YACtG,QAAQ,EAAE,IAAI;YACd,KAAK,EAAE,OAAO,CAAC,IAAI;YACnB,QAAQ,EAAE,CAAC,GAAG,KAAK;SACtB,CAAC,CAAC;IACP,CAAC;AACL,CAAC;AAED;;;;;;GAMG;AACH,SAAS,uBAAuB,CAAC,IAAa,EAAE,KAAa,EAAE,IAAY,EAAE,IAAa;IACtF,IAAI,CAAC,MAAM,CAAC,aAAa,CAAC,KAAK,CAAC,IAAI,KAAK,GAAG,CAAC,IAAI,KAAK,GAAG,OAAO,EAAE,CAAC;QAC/D,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,4DAA4D,EAAE;YAC3G,QAAQ,EAAE,OAAO;YACjB,KAAK,EAAE,KAAK;SACf,CAAC,CAAC;IACP,CAAC;IACD,IAAI,CAAC,MAAM,CAAC,aAAa,CAAC,IAAI,CAAC,IAAI,IAAI,GAAG,CAAC,IAAI,IAAI,GAAG,OAAO,EAAE,CAAC;QAC5D,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,iDAAiD,EAAE;YAChG,QAAQ,EAAE,MAAM;YAChB,KAAK,EAAE,IAAI;SACd,CAAC,CAAC;IACP,CAAC;IACD,UAAU,CAAC,MAAM,EAAE,IAAI,EAAE,KAAK,CAAC,CAAC;IAChC,UAAU,CAAC,MAAM,EAAE,IAAI,EAAE,IAAI,CAAC,CAAC;AACnC,CAAC;AAED,kEAAkE;AAClE,MAAM,oBAAoB;IAMtB;;;;;OAKG;IACH,YAAY,KAAkB,EAAE,SAAiB,EAAE,IAAY;QARvD,eAAU,GAAG,CAAC,CAAC;QASnB,IAAI,CAAC,KAAK,GAAG,KAAK,CAAC;QACnB,IAAI,CAAC,SAAS,GAAG,SAAS,CAAC;QAC3B,IAAI,CAAC,IAAI,GAAG,IAAI,CAAC;IACrB,CAAC;IAED;;;OAGG;IACH,IAAI,cAAc;QACd,OAAO,IAAI,CAAC,UAAU,CAAC;IAC3B,CAAC;IAED;;;;;;;OAOG;IACH,MAAM,CAAC,IAA2B,EAAE,IAAa,EAAE,KAAa,EAAE,IAAY,EAAE,IAAa;QACzF,uBAAuB,CAAC,IAAI,EAAE,KAAK,EAAE,IAAI,EAAE,IAAI,CAAC,CAAC;QACjD,IAAI,CAAC,UAAU,CAAC,IAAI,EAAE,IAAI,EAAE,IAAI,CAAC,CAAC;QAClC,IAAI,CAAC,UAAU,GAAG,CAAC,CAAC;QACpB,IAAI,KAAK,KAAK,CAAC,EAAE,CAAC;YACd,OAAO;QACX,CAAC;QACD,MAAM,MAAM,GAAG,IAAI,CAAC,KAAK,CAAC,MAAM,CAAC,WAAW,EAAE,EAAE,KAAK,EAAE,IAAI,EAAE,IAAI,EAAE,CAAC,EAAE,IAAI,EAAE,CAAC,EAAE,CAAC,CAAC;QACjF,MAAM,KAAK,GAAG,IAAI,CAAC,SAAS,CAAC,IAAI,CAAC,EAAE,IAAI,EAAE,IAAI,EAAE,CAAC,EAAE,MAAM,CAAC,OAAO,EAAE,CAAC,CAAC;QACrE,IAAI,CAAC,SAAS,CAAC,QAAQ,CAAC,IAAI,EAAE,KAAK,EAAE,MAAM,CAAC,KAAK,EAAE,IAAI,CAAC,KAAK,CAAC,aAAa,EAAE,IAAI,CAAC,KAAK,CAAC,IAAI,CAAC,EAAE,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC;QAChH,IAAI,CAAC,UAAU,GAAG,CAAC,CAAC;IACxB,CAAC;IAED;;;;;OAKG;IACH,UAAU,CAAC,IAA2B,EAAE,GAAY,EAAE,KAAa;QAC/D,MAAM,MAAM,GAAG,IAAI,CAAC,KAAK,CAAC,MAAM,CAAC,WAAW,EAAE,EAAE,KAAK,EAAE,KAAK,EAAE,KAAK,EAAE,CAAC,EAAE,IAAI,EAAE,CAAC,EAAE,IAAI,EAAE,CAAC,EAAE,CAAC,CAAC;QAC5F,MAAM,KAAK,GAAG,IAAI,CAAC,IAAI,CAAC,IAAI,CAAC,EAAE,GAAG,EAAE,CAAC,EAAE,MAAM,CAAC,OAAO,EAAE,CAAC,CAAC;QACzD,IAAI,CAAC,IAAI,CAAC,QAAQ,CAAC,IAAI,EAAE,KAAK,EAAE,MAAM,CAAC,KAAK,EAAE,IAAI,CAAC,KAAK,CAAC,aAAa,EAAE,IAAI,CAAC,KAAK,CAAC,IAAI,CAAC,EAAE,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC;IAC/G,CAAC;CACJ;AAED,kGAAkG;AAClG,MAAM,uBAAuB;IAOzB;;;;;;OAMG;IACH,YAAY,KAAkB,EAAE,SAA+B,EAAE,IAAiB,EAAE,OAAe;QAT3F,eAAU,GAAG,CAAC,CAAC;QAUnB,IAAI,CAAC,KAAK,GAAG,KAAK,CAAC;QACnB,IAAI,CAAC,SAAS,GAAG,SAAS,CAAC;QAC3B,IAAI,CAAC,IAAI,GAAG,IAAI,CAAC;QACjB,IAAI,CAAC,OAAO,GAAG,OAAO,CAAC;IAC3B,CAAC;IAED;;;OAGG;IACH,IAAI,cAAc;QACd,OAAO,IAAI,CAAC,UAAU,CAAC;IAC3B,CAAC;IAED;;;;;;;;;OASG;IACH,MAAM,CACF,IAA2B,EAC3B,IAAa,EACb,KAAa,EACb,IAAY,EACZ,OAA4B,EAC5B,QAAiB,EACjB,QAAiB;QAEjB,uBAAuB,CAAC,IAAI,EAAE,KAAK,EAAE,IAAI,EAAE,OAAO,CAAC,IAAI,CAAC,CAAC;QACzD,UAAU,CAAC,QAAQ,EAAE,OAAO,CAAC,MAAM,EAAE,IAAI,CAAC,CAAC;QAC3C,UAAU,CAAC,UAAU,EAAE,QAAQ,EAAE,KAAK,CAAC,CAAC;QACxC,UAAU,CAAC,UAAU,EAAE,QAAQ,EAAE,IAAI,CAAC,CAAC;QACvC,IAAI,CAAC,SAAS,CAAC,MAAM,CAAC,IAAI,EAAE,IAAI,EAAE,KAAK,EAAE,IAAI,EAAE,OAAO,CAAC,IAAI,CAAC,CAAC;QAC7D,IAAI,CAAC,IAAI,CAAC,MAAM,CAAC,IAAI,EAAE,OAAO,CAAC,IAAI,EAAE,IAAI,EAAE,QAAQ,CAAC,CAAC;QACrD,IAAI,CAAC,SAAS,CAAC,UAAU,CAAC,IAAI,EAAE,OAAO,CAAC,MAAM,EAAE,IAAI,CAAC,CAAC;QACtD,IAAI,CAAC,UAAU,GAAG,IAAI,CAAC,SAAS,CAAC,cAAc,GAAG,IAAI,CAAC,IAAI,CAAC,cAAc,GAAG,CAAC,CAAC;QAC/E,IAAI,KAAK,KAAK,CAAC,EAAE,CAAC;YACd,OAAO;QACX,CAAC;QACD,MAAM,MAAM,GAAG,IAAI,CAAC,KAAK,CAAC,MAAM,CAAC,WAAW,EAAE,EAAE,KAAK,EAAE,IAAI,EAAE,IAAI,EAAE,CAAC,EAAE,IAAI,EAAE,CAAC,EAAE,CAAC,CAAC;QACjF,MAAM,KAAK,GAAG,IAAI,CAAC,OAAO,CAAC,IAAI,CAAC,EAAE,IAAI,EAAE,KAAK,EAAE,QAAQ,EAAE,MAAM,EAAE,OAAO,CAAC,MAAM,EAAE,QAAQ,EAAE,CAAC,EAAE,MAAM,CAAC,OAAO,EAAE,CAAC,CAAC;QAChH,IAAI,CAAC,OAAO,CAAC,QAAQ,CAAC,IAAI,EAAE,KAAK,EAAE,MAAM,CAAC,KAAK,EAAE,IAAI,CAAC,KAAK,CAAC,aAAa,EAAE,IAAI,CAAC,KAAK,CAAC,IAAI,CAAC,EAAE,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC;QAC9G,IAAI,CAAC,UAAU,IAAI,CAAC,CAAC;IACzB,CAAC;CACJ"}
@@ -0,0 +1,75 @@
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
+ import { type Binding } from "../types/memory.js";
22
+ import { type ReduceScope } from "./reduce.js";
23
+ /** The key widths record() accepts: `bits / 8` passes, so the low `bits` of every key order the pairs. */
24
+ export type RadixBits = 8 | 16 | 24 | 32;
25
+ /**
26
+ * The caller's scratch: a second key-value pair the passes ping-pong with, the digit-major histogram table and its
27
+ * scanned twin (each `radixHistBytes` bytes at least).
28
+ */
29
+ export interface RadixSortScratch {
30
+ readonly keys: Binding;
31
+ readonly vals: Binding;
32
+ /** The raw digit-major table `radix-hist` writes (the last pass's stays readable after the sort). */
33
+ readonly hist: Binding;
34
+ /** The exclusive scan of `hist`, what `radix-scatter` reads. */
35
+ readonly offsets: Binding;
36
+ }
37
+ /** Where the sorted pairs landed: the input pair or the scratch pair (record() decides by the pass count). */
38
+ export interface RadixSortResult {
39
+ readonly keys: Binding;
40
+ readonly vals: Binding;
41
+ }
42
+ /** A prepared radix sort (spec 6 row 6): records the pass dispatches of one sort into a pass. */
43
+ export interface RadixSortPlanner {
44
+ /**
45
+ * Records the stable sort of `count` (key, value) pairs by the low `bits` of the key; returns the pair holding the
46
+ * result (the scratch pair after an odd pass count, the input pair after an even one); count 0 records nothing and
47
+ * returns the input pair.
48
+ * @param pass - the compute pass
49
+ * @param keys - the keys (at least 4 x count bytes)
50
+ * @param vals - the values (at least 4 x count bytes)
51
+ * @param count - the pair count (a non-negative integer below 2^32)
52
+ * @param bits - the key width: 8, 16, 24 or 32
53
+ * @param scratch - the second pair, the histogram table and its scanned twin
54
+ * @returns the pair the result lives in
55
+ */
56
+ record(pass: GPUComputePassEncoder, keys: Binding, vals: Binding, count: number, bits: RadixBits, scratch: RadixSortScratch): RadixSortResult;
57
+ /** Dispatches the last record() issued: 0 for count 0, else `passes x (2 + the scan's dispatches over the table)`. */
58
+ readonly lastDispatches: number;
59
+ }
60
+ /**
61
+ * The byte size of the digit-major histogram table of one pass over `count` keys at workgroup size `wg`:
62
+ * `4 x 256 x ceil(count / wg)` (the caller sizes `scratch.hist` by it; 0 for count 0, which records nothing).
63
+ * @param count - the pair count
64
+ * @param wg - the workgroup size (`scope.workgroupSize`)
65
+ * @returns the bytes
66
+ */
67
+ export declare function radixHistBytes(count: number, wg: number): number;
68
+ /**
69
+ * Prepares the two radix pipelines and the scan of a scope (compiles once) so record() is synchronous. The planner
70
+ * lives exactly as long as the scope: never use a planner after its scope's dispose().
71
+ * @param scope - the caller's scope (device, caps, cache, scratch, params)
72
+ * @returns the planner
73
+ */
74
+ export declare function prepareRadixSort(scope: ReduceScope): Promise<RadixSortPlanner>;
75
+ //# sourceMappingURL=radix-sort.d.ts.map
@@ -0,0 +1 @@
1
+ {"version":3,"file":"radix-sort.d.ts","sourceRoot":"","sources":["../../../src/primitives/radix-sort.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;;;;;;;;;;GAmBG;AAOH,OAAO,EAAE,KAAK,OAAO,EAAE,MAAM,oBAAoB,CAAC;AAClD,OAAO,EAAE,KAAK,WAAW,EAAE,MAAM,aAAa,CAAC;AAM/C,0GAA0G;AAC1G,MAAM,MAAM,SAAS,GAAG,CAAC,GAAG,EAAE,GAAG,EAAE,GAAG,EAAE,CAAC;AAEzC;;;GAGG;AACH,MAAM,WAAW,gBAAgB;IAC7B,QAAQ,CAAC,IAAI,EAAE,OAAO,CAAC;IACvB,QAAQ,CAAC,IAAI,EAAE,OAAO,CAAC;IACvB,qGAAqG;IACrG,QAAQ,CAAC,IAAI,EAAE,OAAO,CAAC;IACvB,gEAAgE;IAChE,QAAQ,CAAC,OAAO,EAAE,OAAO,CAAC;CAC7B;AAED,8GAA8G;AAC9G,MAAM,WAAW,eAAe;IAC5B,QAAQ,CAAC,IAAI,EAAE,OAAO,CAAC;IACvB,QAAQ,CAAC,IAAI,EAAE,OAAO,CAAC;CAC1B;AAED,iGAAiG;AACjG,MAAM,WAAW,gBAAgB;IAC7B;;;;;;;;;;;OAWG;IACH,MAAM,CACF,IAAI,EAAE,qBAAqB,EAC3B,IAAI,EAAE,OAAO,EACb,IAAI,EAAE,OAAO,EACb,KAAK,EAAE,MAAM,EACb,IAAI,EAAE,SAAS,EACf,OAAO,EAAE,gBAAgB,GAC1B,eAAe,CAAC;IACnB,sHAAsH;IACtH,QAAQ,CAAC,cAAc,EAAE,MAAM,CAAC;CACnC;AAED;;;;;;GAMG;AACH,wBAAgB,cAAc,CAAC,KAAK,EAAE,MAAM,EAAE,EAAE,EAAE,MAAM,GAAG,MAAM,CAEhE;AAED;;;;;GAKG;AACH,wBAAsB,gBAAgB,CAAC,KAAK,EAAE,WAAW,GAAG,OAAO,CAAC,gBAAgB,CAAC,CAKpF"}
@@ -0,0 +1,168 @@
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
+ import { RADIX_BINS, U32_MAX } from "../constants.js";
22
+ import { WebGpuGraphError } from "../errors.js";
23
+ import { plan1d } from "../kernel/dispatch.js";
24
+ import { kernelSpec, RADIX_PARAMS } from "../kernels.js";
25
+ import { prepareScan } from "./scan.js";
26
+ /** The bits of one pass. */
27
+ const RADIX_DIGIT_BITS = 8;
28
+ /**
29
+ * The byte size of the digit-major histogram table of one pass over `count` keys at workgroup size `wg`:
30
+ * `4 x 256 x ceil(count / wg)` (the caller sizes `scratch.hist` by it; 0 for count 0, which records nothing).
31
+ * @param count - the pair count
32
+ * @param wg - the workgroup size (`scope.workgroupSize`)
33
+ * @returns the bytes
34
+ */
35
+ export function radixHistBytes(count, wg) {
36
+ return 4 * RADIX_BINS * Math.ceil(count / wg);
37
+ }
38
+ /**
39
+ * Prepares the two radix pipelines and the scan of a scope (compiles once) so record() is synchronous. The planner
40
+ * lives exactly as long as the scope: never use a planner after its scope's dispose().
41
+ * @param scope - the caller's scope (device, caps, cache, scratch, params)
42
+ * @returns the planner
43
+ */
44
+ export async function prepareRadixSort(scope) {
45
+ const hist = await scope.pipelines.kernel(kernelSpec("radix-hist"));
46
+ const scatter = await scope.pipelines.kernel(kernelSpec("radix-scatter"));
47
+ const scan = await prepareScan(scope);
48
+ return new RadixSortPlannerImpl(scope, hist, scatter, scan);
49
+ }
50
+ /**
51
+ * One binding-size check of record().
52
+ * @param argument - the argument name
53
+ * @param binding - the binding
54
+ * @param bytes - the bytes it must hold
55
+ */
56
+ function checkBinding(argument, binding, bytes) {
57
+ if (binding.size < bytes) {
58
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", `radixSort: ${argument} is smaller than ${bytes} bytes`, {
59
+ argument,
60
+ value: binding.size,
61
+ expected: bytes,
62
+ });
63
+ }
64
+ }
65
+ /**
66
+ * The argument checks of record() (E_INVALID_ARGUMENT before anything is recorded).
67
+ * @param keys - the keys
68
+ * @param vals - the values
69
+ * @param count - the pair count
70
+ * @param bits - the key width
71
+ * @param scratch - the scratch
72
+ * @param wg - the workgroup size
73
+ */
74
+ function checkRecordArguments(keys, vals, count, bits, scratch, wg) {
75
+ if (bits !== 8 && bits !== 16 && bits !== 24 && bits !== 32) {
76
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "radixSort: bits must be 8, 16, 24 or 32", {
77
+ argument: "bits",
78
+ value: bits,
79
+ expected: [8, 16, 24, 32],
80
+ });
81
+ }
82
+ if (!Number.isSafeInteger(count) || count < 0 || count > U32_MAX) {
83
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "radixSort: count must be a non-negative integer below 2^32", {
84
+ argument: "count",
85
+ value: count,
86
+ });
87
+ }
88
+ const pairBytes = 4 * count;
89
+ checkBinding("keys", keys, pairBytes);
90
+ checkBinding("vals", vals, pairBytes);
91
+ checkBinding("scratch.keys", scratch.keys, pairBytes);
92
+ checkBinding("scratch.vals", scratch.vals, pairBytes);
93
+ const tableBytes = radixHistBytes(count, wg);
94
+ checkBinding("scratch.hist", scratch.hist, tableBytes);
95
+ checkBinding("scratch.offsets", scratch.offsets, tableBytes);
96
+ }
97
+ /** The planner: the two resolved kernels and the scan over one scope. */
98
+ class RadixSortPlannerImpl {
99
+ /**
100
+ * Wraps the resolved kernels; use prepareRadixSort().
101
+ * @param scope - the caller's scope
102
+ * @param hist - the `radix-hist` kernel
103
+ * @param scatter - the `radix-scatter` kernel
104
+ * @param scan - the scan planner of the same scope
105
+ */
106
+ constructor(scope, hist, scatter, scan) {
107
+ this.dispatches = 0;
108
+ this.scope = scope;
109
+ this.hist = hist;
110
+ this.scatter = scatter;
111
+ this.scan = scan;
112
+ }
113
+ /**
114
+ * Dispatches the last record() issued.
115
+ * @returns the count
116
+ */
117
+ get lastDispatches() {
118
+ return this.dispatches;
119
+ }
120
+ /**
121
+ * Records the passes into the pass (see the interface).
122
+ * @param pass - the compute pass
123
+ * @param keys - the keys
124
+ * @param vals - the values
125
+ * @param count - the pair count
126
+ * @param bits - the key width
127
+ * @param scratch - the second pair, the histogram table and its scanned twin
128
+ * @returns the pair the result lives in
129
+ */
130
+ record(pass, keys, vals, count, bits, scratch) {
131
+ const wg = this.scope.workgroupSize;
132
+ checkRecordArguments(keys, vals, count, bits, scratch, wg);
133
+ if (count === 0) {
134
+ this.dispatches = 0;
135
+ return { keys, vals };
136
+ }
137
+ const groups = Math.ceil(count / wg);
138
+ const plan = plan1d(count, wg, this.scope.caps);
139
+ const tableWords = RADIX_BINS * groups;
140
+ const tableBytes = 4 * tableWords;
141
+ const histTable = { ...scratch.hist, size: tableBytes };
142
+ const offsets = { ...scratch.offsets, size: tableBytes };
143
+ let src = { keys, vals };
144
+ let dst = { keys: scratch.keys, vals: scratch.vals };
145
+ let dispatches = 0;
146
+ const passes = bits / RADIX_DIGIT_BITS;
147
+ for (let p = 0; p < passes; p++) {
148
+ const params = this.scope.params(RADIX_PARAMS, { count, shift: RADIX_DIGIT_BITS * p, groups, pad0: 0 });
149
+ const histBound = this.hist.bind({ keys: src.keys, hist: histTable, P: params.binding });
150
+ this.hist.dispatch(pass, histBound, plan, [params.offset]);
151
+ this.scan.record(pass, histTable, tableWords, offsets);
152
+ const scatterBound = this.scatter.bind({
153
+ keys: src.keys,
154
+ vals: src.vals,
155
+ offsets,
156
+ keysOut: dst.keys,
157
+ valsOut: dst.vals,
158
+ P: params.binding,
159
+ });
160
+ this.scatter.dispatch(pass, scatterBound, plan, [params.offset]);
161
+ dispatches += 2 + this.scan.lastDispatches;
162
+ [src, dst] = [dst, src];
163
+ }
164
+ this.dispatches = dispatches;
165
+ return src;
166
+ }
167
+ }
168
+ //# sourceMappingURL=radix-sort.js.map
@@ -0,0 +1 @@
1
+ {"version":3,"file":"radix-sort.js","sourceRoot":"","sources":["../../../src/primitives/radix-sort.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;;;;;;;;;;GAmBG;AAEH,OAAO,EAAE,UAAU,EAAE,OAAO,EAAE,MAAM,iBAAiB,CAAC;AACtD,OAAO,EAAE,gBAAgB,EAAE,MAAM,cAAc,CAAC;AAChD,OAAO,EAAE,MAAM,EAAE,MAAM,uBAAuB,CAAC;AAE/C,OAAO,EAAE,UAAU,EAAE,YAAY,EAAE,MAAM,eAAe,CAAC;AAGzD,OAAO,EAAE,WAAW,EAAoB,MAAM,WAAW,CAAC;AAE1D,4BAA4B;AAC5B,MAAM,gBAAgB,GAAG,CAAC,CAAC;AAkD3B;;;;;;GAMG;AACH,MAAM,UAAU,cAAc,CAAC,KAAa,EAAE,EAAU;IACpD,OAAO,CAAC,GAAG,UAAU,GAAG,IAAI,CAAC,IAAI,CAAC,KAAK,GAAG,EAAE,CAAC,CAAC;AAClD,CAAC;AAED;;;;;GAKG;AACH,MAAM,CAAC,KAAK,UAAU,gBAAgB,CAAC,KAAkB;IACrD,MAAM,IAAI,GAAG,MAAM,KAAK,CAAC,SAAS,CAAC,MAAM,CAAC,UAAU,CAAC,YAAY,CAAC,CAAC,CAAC;IACpE,MAAM,OAAO,GAAG,MAAM,KAAK,CAAC,SAAS,CAAC,MAAM,CAAC,UAAU,CAAC,eAAe,CAAC,CAAC,CAAC;IAC1E,MAAM,IAAI,GAAG,MAAM,WAAW,CAAC,KAAK,CAAC,CAAC;IACtC,OAAO,IAAI,oBAAoB,CAAC,KAAK,EAAE,IAAI,EAAE,OAAO,EAAE,IAAI,CAAC,CAAC;AAChE,CAAC;AAED;;;;;GAKG;AACH,SAAS,YAAY,CAAC,QAAgB,EAAE,OAAgB,EAAE,KAAa;IACnE,IAAI,OAAO,CAAC,IAAI,GAAG,KAAK,EAAE,CAAC;QACvB,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,cAAc,QAAQ,oBAAoB,KAAK,QAAQ,EAAE;YACtG,QAAQ;YACR,KAAK,EAAE,OAAO,CAAC,IAAI;YACnB,QAAQ,EAAE,KAAK;SAClB,CAAC,CAAC;IACP,CAAC;AACL,CAAC;AAED;;;;;;;;GAQG;AACH,SAAS,oBAAoB,CACzB,IAAa,EACb,IAAa,EACb,KAAa,EACb,IAAe,EACf,OAAyB,EACzB,EAAU;IAEV,IAAI,IAAI,KAAK,CAAC,IAAI,IAAI,KAAK,EAAE,IAAI,IAAI,KAAK,EAAE,IAAI,IAAI,KAAK,EAAE,EAAE,CAAC;QAC1D,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,yCAAyC,EAAE;YACxF,QAAQ,EAAE,MAAM;YAChB,KAAK,EAAE,IAAI;YACX,QAAQ,EAAE,CAAC,CAAC,EAAE,EAAE,EAAE,EAAE,EAAE,EAAE,CAAC;SAC5B,CAAC,CAAC;IACP,CAAC;IACD,IAAI,CAAC,MAAM,CAAC,aAAa,CAAC,KAAK,CAAC,IAAI,KAAK,GAAG,CAAC,IAAI,KAAK,GAAG,OAAO,EAAE,CAAC;QAC/D,MAAM,IAAI,gBAAgB,CAAC,oBAAoB,EAAE,4DAA4D,EAAE;YAC3G,QAAQ,EAAE,OAAO;YACjB,KAAK,EAAE,KAAK;SACf,CAAC,CAAC;IACP,CAAC;IACD,MAAM,SAAS,GAAG,CAAC,GAAG,KAAK,CAAC;IAC5B,YAAY,CAAC,MAAM,EAAE,IAAI,EAAE,SAAS,CAAC,CAAC;IACtC,YAAY,CAAC,MAAM,EAAE,IAAI,EAAE,SAAS,CAAC,CAAC;IACtC,YAAY,CAAC,cAAc,EAAE,OAAO,CAAC,IAAI,EAAE,SAAS,CAAC,CAAC;IACtD,YAAY,CAAC,cAAc,EAAE,OAAO,CAAC,IAAI,EAAE,SAAS,CAAC,CAAC;IACtD,MAAM,UAAU,GAAG,cAAc,CAAC,KAAK,EAAE,EAAE,CAAC,CAAC;IAC7C,YAAY,CAAC,cAAc,EAAE,OAAO,CAAC,IAAI,EAAE,UAAU,CAAC,CAAC;IACvD,YAAY,CAAC,iBAAiB,EAAE,OAAO,CAAC,OAAO,EAAE,UAAU,CAAC,CAAC;AACjE,CAAC;AAED,yEAAyE;AACzE,MAAM,oBAAoB;IAOtB;;;;;;OAMG;IACH,YAAY,KAAkB,EAAE,IAAY,EAAE,OAAe,EAAE,IAAiB;QATxE,eAAU,GAAG,CAAC,CAAC;QAUnB,IAAI,CAAC,KAAK,GAAG,KAAK,CAAC;QACnB,IAAI,CAAC,IAAI,GAAG,IAAI,CAAC;QACjB,IAAI,CAAC,OAAO,GAAG,OAAO,CAAC;QACvB,IAAI,CAAC,IAAI,GAAG,IAAI,CAAC;IACrB,CAAC;IAED;;;OAGG;IACH,IAAI,cAAc;QACd,OAAO,IAAI,CAAC,UAAU,CAAC;IAC3B,CAAC;IAED;;;;;;;;;OASG;IACH,MAAM,CACF,IAA2B,EAC3B,IAAa,EACb,IAAa,EACb,KAAa,EACb,IAAe,EACf,OAAyB;QAEzB,MAAM,EAAE,GAAG,IAAI,CAAC,KAAK,CAAC,aAAa,CAAC;QACpC,oBAAoB,CAAC,IAAI,EAAE,IAAI,EAAE,KAAK,EAAE,IAAI,EAAE,OAAO,EAAE,EAAE,CAAC,CAAC;QAC3D,IAAI,KAAK,KAAK,CAAC,EAAE,CAAC;YACd,IAAI,CAAC,UAAU,GAAG,CAAC,CAAC;YACpB,OAAO,EAAE,IAAI,EAAE,IAAI,EAAE,CAAC;QAC1B,CAAC;QACD,MAAM,MAAM,GAAG,IAAI,CAAC,IAAI,CAAC,KAAK,GAAG,EAAE,CAAC,CAAC;QACrC,MAAM,IAAI,GAAG,MAAM,CAAC,KAAK,EAAE,EAAE,EAAE,IAAI,CAAC,KAAK,CAAC,IAAI,CAAC,CAAC;QAChD,MAAM,UAAU,GAAG,UAAU,GAAG,MAAM,CAAC;QACvC,MAAM,UAAU,GAAG,CAAC,GAAG,UAAU,CAAC;QAClC,MAAM,SAAS,GAAY,EAAE,GAAG,OAAO,CAAC,IAAI,EAAE,IAAI,EAAE,UAAU,EAAE,CAAC;QACjE,MAAM,OAAO,GAAY,EAAE,GAAG,OAAO,CAAC,OAAO,EAAE,IAAI,EAAE,UAAU,EAAE,CAAC;QAClE,IAAI,GAAG,GAAoB,EAAE,IAAI,EAAE,IAAI,EAAE,CAAC;QAC1C,IAAI,GAAG,GAAoB,EAAE,IAAI,EAAE,OAAO,CAAC,IAAI,EAAE,IAAI,EAAE,OAAO,CAAC,IAAI,EAAE,CAAC;QACtE,IAAI,UAAU,GAAG,CAAC,CAAC;QACnB,MAAM,MAAM,GAAG,IAAI,GAAG,gBAAgB,CAAC;QACvC,KAAK,IAAI,CAAC,GAAG,CAAC,EAAE,CAAC,GAAG,MAAM,EAAE,CAAC,EAAE,EAAE,CAAC;YAC9B,MAAM,MAAM,GAAG,IAAI,CAAC,KAAK,CAAC,MAAM,CAAC,YAAY,EAAE,EAAE,KAAK,EAAE,KAAK,EAAE,gBAAgB,GAAG,CAAC,EAAE,MAAM,EAAE,IAAI,EAAE,CAAC,EAAE,CAAC,CAAC;YACxG,MAAM,SAAS,GAAG,IAAI,CAAC,IAAI,CAAC,IAAI,CAAC,EAAE,IAAI,EAAE,GAAG,CAAC,IAAI,EAAE,IAAI,EAAE,SAAS,EAAE,CAAC,EAAE,MAAM,CAAC,OAAO,EAAE,CAAC,CAAC;YACzF,IAAI,CAAC,IAAI,CAAC,QAAQ,CAAC,IAAI,EAAE,SAAS,EAAE,IAAI,EAAE,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC;YAC3D,IAAI,CAAC,IAAI,CAAC,MAAM,CAAC,IAAI,EAAE,SAAS,EAAE,UAAU,EAAE,OAAO,CAAC,CAAC;YACvD,MAAM,YAAY,GAAG,IAAI,CAAC,OAAO,CAAC,IAAI,CAAC;gBACnC,IAAI,EAAE,GAAG,CAAC,IAAI;gBACd,IAAI,EAAE,GAAG,CAAC,IAAI;gBACd,OAAO;gBACP,OAAO,EAAE,GAAG,CAAC,IAAI;gBACjB,OAAO,EAAE,GAAG,CAAC,IAAI;gBACjB,CAAC,EAAE,MAAM,CAAC,OAAO;aACpB,CAAC,CAAC;YACH,IAAI,CAAC,OAAO,CAAC,QAAQ,CAAC,IAAI,EAAE,YAAY,EAAE,IAAI,EAAE,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC;YACjE,UAAU,IAAI,CAAC,GAAG,IAAI,CAAC,IAAI,CAAC,cAAc,CAAC;YAC3C,CAAC,GAAG,EAAE,GAAG,CAAC,GAAG,CAAC,GAAG,EAAE,GAAG,CAAC,CAAC;QAC5B,CAAC;QACD,IAAI,CAAC,UAAU,GAAG,UAAU,CAAC;QAC7B,OAAO,GAAG,CAAC;IACf,CAAC;CACJ"}
@@ -0,0 +1,44 @@
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
+ import { type Binding } from "../types/memory.js";
15
+ import { type ReduceScope } from "./reduce.js";
16
+ /** 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()). */
17
+ export interface ScanTotal {
18
+ readonly binding: Binding;
19
+ readonly index: number;
20
+ }
21
+ /** A prepared exclusive scan (spec 6 row 2): records the level dispatches of one scan into a pass. */
22
+ export interface ScanPlanner {
23
+ /**
24
+ * Records the scan of `count` u32 of `src` into `out` (exclusive); returns where the total landed (a word of the
25
+ * planner's scratch, valid until the next record()); nothing for count 0 (the total word is then 0).
26
+ * @param pass - the compute pass
27
+ * @param src - the input words (at least 4 x count bytes)
28
+ * @param count - the word count (a non-negative integer below 2^32)
29
+ * @param out - the output words (at least 4 x count bytes)
30
+ * @returns the total's location
31
+ */
32
+ record(pass: GPUComputePassEncoder, src: Binding, count: number, out: Binding): ScanTotal;
33
+ /** 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). */
34
+ readonly lastDispatches: number;
35
+ }
36
+ /**
37
+ * Prepares the two scan pipelines of a scope (compiles once) so record() is synchronous. The zero word the empty
38
+ * scan's total points at is a scratch of `scope`, so the planner lives exactly as long as the scope: never use a
39
+ * planner after its scope's dispose().
40
+ * @param scope - the caller's scope (device, caps, cache, scratch, params)
41
+ * @returns the planner
42
+ */
43
+ export declare function prepareScan(scope: ReduceScope): Promise<ScanPlanner>;
44
+ //# sourceMappingURL=scan.d.ts.map
@@ -0,0 +1 @@
1
+ {"version":3,"file":"scan.d.ts","sourceRoot":"","sources":["../../../src/primitives/scan.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;;;GAYG;AAMH,OAAO,EAAE,KAAK,OAAO,EAAE,MAAM,oBAAoB,CAAC;AAClD,OAAO,EAAE,KAAK,WAAW,EAAE,MAAM,aAAa,CAAC;AAE/C,4IAA4I;AAC5I,MAAM,WAAW,SAAS;IACtB,QAAQ,CAAC,OAAO,EAAE,OAAO,CAAC;IAC1B,QAAQ,CAAC,KAAK,EAAE,MAAM,CAAC;CAC1B;AAED,sGAAsG;AACtG,MAAM,WAAW,WAAW;IACxB;;;;;;;;OAQG;IACH,MAAM,CAAC,IAAI,EAAE,qBAAqB,EAAE,GAAG,EAAE,OAAO,EAAE,KAAK,EAAE,MAAM,EAAE,GAAG,EAAE,OAAO,GAAG,SAAS,CAAC;IAC1F,kJAAkJ;IAClJ,QAAQ,CAAC,cAAc,EAAE,MAAM,CAAC;CACnC;AAKD;;;;;;GAMG;AACH,wBAAsB,WAAW,CAAC,KAAK,EAAE,WAAW,GAAG,OAAO,CAAC,WAAW,CAAC,CAM1E"}
@@ -0,0 +1,151 @@
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
+ import { WebGpuGraphError } from "../errors.js";
15
+ import { groupsOf, plan1d } from "../kernel/dispatch.js";
16
+ import { kernelSpec, SCAN_PARAMS } from "../kernels.js";
17
+ /** The largest u32: the largest count `ScanParams.count` can carry. */
18
+ const U32_MAX = 0xffffffff;
19
+ /**
20
+ * Prepares the two scan pipelines of a scope (compiles once) so record() is synchronous. The zero word the empty
21
+ * scan's total points at is a scratch of `scope`, so the planner lives exactly as long as the scope: never use a
22
+ * planner after its scope's dispose().
23
+ * @param scope - the caller's scope (device, caps, cache, scratch, params)
24
+ * @returns the planner
25
+ */
26
+ export async function prepareScan(scope) {
27
+ const block = await scope.pipelines.kernel(kernelSpec("scan-block"));
28
+ const add = await scope.pipelines.kernel(kernelSpec("scan-add"));
29
+ const zero = scope.scratch(4, "scan/zero");
30
+ scope.device.queue.writeBuffer(zero, 0, new Uint32Array(1));
31
+ return new ScanPlannerImpl(scope, block, add, { buffer: zero, offset: 0, size: 4, window: null });
32
+ }
33
+ /**
34
+ * The argument checks of record() (E_INVALID_ARGUMENT before anything is recorded).
35
+ * @param src - the input binding
36
+ * @param count - the word count
37
+ * @param out - the output binding
38
+ */
39
+ function checkRecordArguments(src, count, out) {
40
+ if (!Number.isSafeInteger(count) || count < 0 || count > U32_MAX) {
41
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "scan: count must be a non-negative integer below 2^32", {
42
+ argument: "count",
43
+ value: count,
44
+ });
45
+ }
46
+ if (src.size < 4 * count) {
47
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "scan: src is smaller than 4 x count bytes", {
48
+ argument: "src",
49
+ value: src.size,
50
+ expected: 4 * count,
51
+ });
52
+ }
53
+ if (out.size < 4 * count) {
54
+ throw new WebGpuGraphError("E_INVALID_ARGUMENT", "scan: out is smaller than 4 x count bytes", {
55
+ argument: "out",
56
+ value: out.size,
57
+ expected: 4 * count,
58
+ });
59
+ }
60
+ }
61
+ /** The planner: the two resolved kernels over one scope. */
62
+ class ScanPlannerImpl {
63
+ /**
64
+ * Wraps the resolved kernels; use prepareScan().
65
+ * @param scope - the caller's scope
66
+ * @param block - the `scan-block` kernel
67
+ * @param add - the `scan-add` kernel
68
+ * @param zero - a one-word binding holding 0 (the total of the empty scan)
69
+ */
70
+ constructor(scope, block, add, zero) {
71
+ this.dispatches = 0;
72
+ this.scope = scope;
73
+ this.block = block;
74
+ this.add = add;
75
+ this.zero = zero;
76
+ }
77
+ /**
78
+ * Dispatches the last record() issued.
79
+ * @returns the count
80
+ */
81
+ get lastDispatches() {
82
+ return this.dispatches;
83
+ }
84
+ /**
85
+ * Records the levels into the pass (see the interface).
86
+ * @param pass - the compute pass
87
+ * @param src - the input words
88
+ * @param count - the word count
89
+ * @param out - the output words
90
+ * @returns the total's location
91
+ */
92
+ record(pass, src, count, out) {
93
+ checkRecordArguments(src, count, out);
94
+ if (count === 0) {
95
+ this.dispatches = 0;
96
+ return { binding: this.zero, index: 0 };
97
+ }
98
+ const wg = this.scope.workgroupSize;
99
+ const levels = [];
100
+ let input = src;
101
+ let output = out;
102
+ let levelCount = count;
103
+ for (;;) {
104
+ const blocks = Math.ceil(levelCount / wg);
105
+ // Sized by the PADDED workgroup count (reduce's shape, reduce.ts partials-1): above MAX_1D_ITEMS plan1d pads
106
+ // the grid to x * y >= blocks, and every padded workgroup still stores its (zero) block sum at its own index.
107
+ // The next level scans only `blocks` words, so the padded tail is written and never read.
108
+ const plan = plan1d(levelCount, wg, this.scope.caps);
109
+ const groups = groupsOf(plan);
110
+ const sums = this.scratch(groups, `scan/sums${levels.length}`);
111
+ levels.push({ input, output, count: levelCount, plan, sums });
112
+ if (blocks === 1) {
113
+ break;
114
+ }
115
+ input = sums;
116
+ output = this.scratch(groups, `scan/offsets${levels.length}`);
117
+ levelCount = blocks;
118
+ }
119
+ for (const level of levels) {
120
+ const params = this.scope.params(SCAN_PARAMS, { count: level.count, pad0: 0, pad1: 0, pad2: 0 });
121
+ const bound = this.block.bind({
122
+ src: level.input,
123
+ out: level.output,
124
+ blockSums: level.sums,
125
+ P: params.binding,
126
+ });
127
+ this.block.dispatch(pass, bound, level.plan, [params.offset]);
128
+ }
129
+ for (let l = levels.length - 2; l >= 0; l--) {
130
+ const level = levels[l];
131
+ const params = this.scope.params(SCAN_PARAMS, { count: level.count, pad0: 0, pad1: 0, pad2: 0 });
132
+ const bound = this.add.bind({ out: level.output, blockOffsets: levels[l + 1].output, P: params.binding });
133
+ this.add.dispatch(pass, bound, level.plan, [params.offset]);
134
+ }
135
+ this.dispatches = 2 * levels.length - 1;
136
+ const top = levels[levels.length - 1];
137
+ return { binding: top.sums, index: 0 };
138
+ }
139
+ /**
140
+ * A scratch of `words` u32 from the scope, bound whole (never zero-length; the pool rounds the buffer up).
141
+ * @param words - the word count (>= 1)
142
+ * @param label - the scratch label
143
+ * @returns the binding
144
+ */
145
+ scratch(words, label) {
146
+ const size = 4 * words;
147
+ const buffer = this.scope.scratch(size, label);
148
+ return { buffer, offset: 0, size, window: null };
149
+ }
150
+ }
151
+ //# sourceMappingURL=scan.js.map