@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,27 +1,29 @@
1
1
  /**
2
- * The segmented (per-row) reduction primitive of spec 6 row 3 in its thread-per-row tier (P2-P3): one invocation per
3
- * CSR row folds the caller's VALUE snippet over the row's arcs into `out[row]` (f32), a row with no arcs receiving
4
- * the identity element. The degree tiers of degreeOrder() (subgroup-per-row for the mid tier, workgroup-per-row for
5
- * the high tier) land at P4; until then `tiers !== null` is E_UNSUPPORTED and USE_PERM is always false (the perm
6
- * slot carries the rowPtr dummy of graphBindings). The row and arc counts come from the core's binding sizes: the
7
- * residency binds every array at its exact byte length (contract 3.8), so rowPtr is 4(n + 1) bytes and colIdx
8
- * 4 x arcCount.
2
+ * The segmented (per-row) reduction primitive of spec 6 row 3: the caller's VALUE snippet folded over every CSR
3
+ * row's arcs into `out[row]` (f32), a row with no arcs receiving the identity element. Without tiers it is ONE
4
+ * thread-per-row dispatch (TIER 0, USE_PERM false: the perm slot carries the rowPtr dummy of graphBindings). With
5
+ * the degree tiers of degreeOrder() (P4-T5, PD-6) it is up to three dispatches over the permuted rows: TIER 2, one
6
+ * workgroup per row of degree >= 1024 over [0, hiEnd); TIER 1, 32 lanes per row of degree 32..1023 over [hiEnd,
7
+ * midEnd); TIER 0, one thread per row over [midEnd, n) -- each compiled only when its range is non-empty. The row
8
+ * and arc counts come from the core's binding sizes: the residency binds every array at its exact byte length
9
+ * (contract 3.8), so rowPtr is 4(n + 1) bytes and colIdx 4 x arcCount. On a windowed core (spec 4.2, P4-T7, PD-8)
10
+ * the dispatches repeat per window with `arcBase = w.start`, `arcEnd = w.end` and `accumulate = 1` over an out
11
+ * pre-filled with the identity element (unless the caller asked to accumulate), so a row split across windows
12
+ * combines its partials: the untiered dispatch covers the window's rows [rowFirst, rowLast], while every tier
13
+ * dispatch covers its FULL tier range (permutation positions, not comparable with the window's node rows; a row
14
+ * whose arcs lie outside the window folds nothing and `finish` writes comb(out[i], identity) = out[i]).
9
15
  */
10
16
 
11
17
  import { WebGpuGraphError } from "../errors.js";
12
18
  import { plan1d } from "../kernel/dispatch.js";
13
19
  import { type Kernel } from "../kernel/kernel.js";
14
- import { graphBindings, graphOverrides, kernelSpec, RANGE_PARAMS } from "../kernels.js";
20
+ import { FILL_PARAMS, graphBindings, graphOverrides, kernelSpec, RANGE_PARAMS } from "../kernels.js";
15
21
  import { type CoreBinding } from "../memory/residency.js";
16
- import { type Binding } from "../types/memory.js";
17
- import { assertNotWindowed, rowCountOf } from "./core-shape.js";
22
+ import { type ArcWindow, type Binding } from "../types/memory.js";
23
+ import { arcCountOf, type DegreeTiers, MID_TIER_LANES, rowCountOf, windowBinding } from "./core-shape.js";
18
24
  import { type ReduceOp, type ReduceScope } from "./reduce.js";
19
25
 
20
- /** The degree tiers of degreeOrder(): the permutation binding and the CPU-side segmentOffsets [0, hiEnd, midEnd, lowEnd, n]. */
21
- export interface DegreeTiers {
22
- readonly perm: Binding;
23
- readonly segmentOffsets: readonly [number, number, number, number, number];
24
- }
26
+ export { type DegreeTiers } from "./core-shape.js";
25
27
 
26
28
  /** Options of segmentedReduce. `valueSnippet` is the Gunrock-style functor: WGSL statements assigning `v` from (row, arc, nbr, weight) (4.5; `nbr` because `target` is a WGSL reserved word). */
27
29
  export interface SegmentedReduceOptions {
@@ -31,10 +33,33 @@ export interface SegmentedReduceOptions {
31
33
  readonly accumulate?: boolean | undefined;
32
34
  }
33
35
 
34
- /** A prepared segmented reduce (P2-P3: the thread-per-row tier only; `tiers !== null` -> E_UNSUPPORTED { feature: "segmentedReduce.tiers" } until P4). */
36
+ /** A prepared segmented reduce: one thread-per-row dispatch without tiers, up to three tier dispatches with them (times the windows of a windowed core, plus its identity fill). */
35
37
  export interface SegmentedReducePlanner {
36
- /** Records one dispatch over rows [0, n) (tiers null) writing out[i] (f32) per row; a row with no arcs gets the identity element. */
38
+ /** Records the dispatches over rows [0, n) writing out[i] (f32) per row; a row with no arcs gets the identity element. */
37
39
  record(pass: GPUComputePassEncoder, core: CoreBinding, out: Binding): void;
40
+ /** Dispatches the last record() issued (0 for n = 0; 1 without tiers; 1 to 3 with them; on a windowed core the identity fill plus that many per window). */
41
+ readonly lastDispatches: number;
42
+ }
43
+
44
+ /** The largest finite f32, 0x1.fffffep+127 (the prelude's F32_MAX): the min identity; its negation the max identity. */
45
+ const F32_MAX = 2 ** 128 - 2 ** 104;
46
+
47
+ /**
48
+ * The u32 bit pattern the `fill` kernel writes so every out word holds the identity element of the operator before
49
+ * the windowed dispatches accumulate into it (PD-8).
50
+ * @param op - the operator
51
+ * @returns the bits of 0 (sum), F32_MAX (min) or -F32_MAX (max)
52
+ */
53
+ function identityFillWord(op: ReduceOp): number {
54
+ let identity = 0;
55
+ if (op === "min") {
56
+ identity = F32_MAX;
57
+ } else if (op === "max") {
58
+ identity = -F32_MAX;
59
+ }
60
+ const view = new DataView(new ArrayBuffer(4));
61
+ view.setFloat32(0, identity, true);
62
+ return view.getUint32(0, true);
38
63
  }
39
64
 
40
65
  /** The identifiers a VALUE snippet may name (contract 3.11, 4.5); every other identifier is rejected textually before compose. */
@@ -141,36 +166,75 @@ function validateValueSnippet(snippet: string): void {
141
166
  }
142
167
  }
143
168
 
144
- /** The thread-per-row planner: ONE `segmented-reduce` dispatch with TIER 0 over every row. */
145
- class ThreadPerRowPlanner implements SegmentedReducePlanner {
169
+ /** One tier's compiled pipeline and the rows [start, end) of the permutation it covers per record() (a tier whose range is empty is not compiled). */
170
+ interface TierDispatch {
171
+ readonly tier: 0 | 1 | 2;
172
+ readonly kernel: Kernel;
173
+ readonly start: number;
174
+ readonly end: number;
175
+ }
176
+
177
+ /**
178
+ * The tiered planner: without tiers ONE `segmented-reduce` dispatch with TIER 0 over [0, n); with tiers TIER 2 over
179
+ * [0, hiEnd) (one workgroup per row), then TIER 1 over [hiEnd, midEnd) (WG / 32 rows per workgroup), then TIER 0
180
+ * over [midEnd, n), each with its own RangeParams record and the perm binding.
181
+ */
182
+ class TieredPlanner implements SegmentedReducePlanner {
146
183
  private readonly scope: ReduceScope;
147
- private readonly kernel: Kernel;
184
+ private readonly tiers: readonly TierDispatch[];
185
+ private readonly fill: Kernel;
186
+ private readonly perm: Binding | null;
148
187
  private readonly hasWeights: boolean;
149
188
  private readonly accumulate: boolean;
189
+ private readonly identityWord: number;
190
+ private dispatches = 0;
150
191
 
151
192
  /**
152
- * Wraps a compiled thread-per-row pipeline with the pattern it was compiled for.
153
- * @param scope - the scope the pipeline was prepared in
154
- * @param kernel - the compiled kernel
155
- * @param hasWeights - the HAS_WEIGHTS the pipeline was compiled with
193
+ * Wraps the compiled tier pipelines with the pattern they were compiled for.
194
+ * @param scope - the scope the pipelines were prepared in
195
+ * @param tiers - the compiled tiers in dispatch order (TIER 2, 1, 0; only the non-empty ones)
196
+ * @param fill - the `fill` pipeline of the identity pre-fill of a windowed core
197
+ * @param perm - the degreeOrder permutation binding (null: USE_PERM false, rows are node indices)
198
+ * @param hasWeights - the HAS_WEIGHTS the pipelines were compiled with
156
199
  * @param accumulate - whether record() combines into out instead of overwriting
200
+ * @param identityWord - the u32 bits of the operator's identity element (identityFillWord)
157
201
  */
158
- constructor(scope: ReduceScope, kernel: Kernel, hasWeights: boolean, accumulate: boolean) {
202
+ constructor(
203
+ scope: ReduceScope,
204
+ tiers: readonly TierDispatch[],
205
+ fill: Kernel,
206
+ perm: Binding | null,
207
+ hasWeights: boolean,
208
+ accumulate: boolean,
209
+ identityWord: number,
210
+ ) {
159
211
  this.scope = scope;
160
- this.kernel = kernel;
212
+ this.tiers = tiers;
213
+ this.fill = fill;
214
+ this.perm = perm;
161
215
  this.hasWeights = hasWeights;
162
216
  this.accumulate = accumulate;
217
+ this.identityWord = identityWord;
218
+ }
219
+
220
+ /**
221
+ * Dispatches the last record() issued.
222
+ * @returns 0 for n = 0, else the number of non-empty tier ranges
223
+ */
224
+ get lastDispatches(): number {
225
+ return this.dispatches;
163
226
  }
164
227
 
165
228
  /**
166
- * Records the dispatch: rows [0, n), arcs [0, arcCount), plan1d(n); nothing for n = 0 (no zero-length binding is
167
- * ever created).
229
+ * Records the dispatches: every tier's rows over the arcs [0, arcCount), each with its own params record;
230
+ * nothing for n = 0 (no zero-length binding is ever created). On a windowed core: the identity fill (unless
231
+ * accumulating into the caller's out), then the dispatches once per window (PD-8).
168
232
  * @param pass - the pass to record into
169
- * @param core - a core with the SAME weights pattern as the one prepared (any snapshot)
233
+ * @param core - a core with the SAME weights pattern as the one prepared (any snapshot; with tiers, the one the
234
+ * tiers were built from)
170
235
  * @param out - at least 4n bytes of f32
171
236
  */
172
237
  record(pass: GPUComputePassEncoder, core: CoreBinding, out: Binding): void {
173
- assertNotWindowed(core, "segmentedReduce");
174
238
  if ((core.weights !== null) !== this.hasWeights) {
175
239
  throw new WebGpuGraphError(
176
240
  "E_INVALID_ARGUMENT",
@@ -183,10 +247,6 @@ class ThreadPerRowPlanner implements SegmentedReducePlanner {
183
247
  );
184
248
  }
185
249
  const n = rowCountOf(core, "segmentedReduce");
186
- if (n === 0) {
187
- return;
188
- }
189
- const arcCount = core.colIdx === null ? 0 : core.colIdx.size / 4;
190
250
  if (out.size < 4 * n) {
191
251
  throw new WebGpuGraphError(
192
252
  "E_INVALID_ARGUMENT",
@@ -198,24 +258,96 @@ class ThreadPerRowPlanner implements SegmentedReducePlanner {
198
258
  },
199
259
  );
200
260
  }
201
- const params = this.scope.params(RANGE_PARAMS, {
202
- start: 0,
203
- end: n,
204
- arcBase: 0,
205
- arcEnd: arcCount,
206
- accumulate: this.accumulate ? 1 : 0,
207
- n,
261
+ this.dispatches = 0;
262
+ if (n === 0) {
263
+ return;
264
+ }
265
+ if (core.windows === null) {
266
+ this.recordWindow(pass, core, out, n, null, 0);
267
+ return;
268
+ }
269
+ if (!this.accumulate) {
270
+ const params = this.scope.params(FILL_PARAMS, { count: n, value: this.identityWord, mode: 0 });
271
+ const bound = this.fill.bind({ dst: out, P: params.binding });
272
+ this.fill.dispatch(pass, bound, plan1d(n, this.scope.workgroupSize, this.scope.caps), [params.offset]);
273
+ this.dispatches++;
274
+ }
275
+ const { windows } = core;
276
+ windows.forEach((w, k) => {
277
+ const windowed: CoreBinding = {
278
+ ...core,
279
+ colIdx: windowBinding(core, "colIdx", w),
280
+ weights: core.weights === null ? null : windowBinding(core, "weights", w),
281
+ };
282
+ this.recordWindow(pass, windowed, out, n, w, k + 1 < windows.length ? windows[k + 1].start : w.end);
208
283
  });
209
- const bound = this.kernel.bind({ ...graphBindings(core, null), out, P: params.binding });
210
- this.kernel.dispatch(pass, bound, plan1d(n, this.scope.workgroupSize, this.scope.caps), [params.offset]);
284
+ }
285
+
286
+ /**
287
+ * The tier dispatches over one arc window (or the whole core when `w` is null): every non-empty tier over its
288
+ * full range, except the untiered TIER 0 dispatch, which covers exactly the window's rows. Consecutive windows
289
+ * overlap by up to ARC_WINDOW_ALIGN - 1 arcs (the next window opens at the aligned-down end of this one), and a
290
+ * tier dispatch visits EVERY row, so the tiers fold [w.start, nextStart) and leave the overlap to the next
291
+ * window; the untiered dispatch keeps [w.start, w.end): its rows [rowFirst, rowLast] are exactly the rows whose
292
+ * arcs end inside this window, and the next window's rows start after them.
293
+ * @param pass - the pass to record into
294
+ * @param core - the core with the window's colIdx / weights bound
295
+ * @param out - the output binding
296
+ * @param n - the row count
297
+ * @param w - the window, or null for the single whole-core dispatch
298
+ * @param nextStart - the next window's first arc (w.end for the last window)
299
+ */
300
+ private recordWindow(
301
+ pass: GPUComputePassEncoder,
302
+ core: CoreBinding,
303
+ out: Binding,
304
+ n: number,
305
+ w: ArcWindow | null,
306
+ nextStart: number,
307
+ ): void {
308
+ const wg = this.scope.workgroupSize;
309
+ for (const { tier, kernel, start, end } of this.tiers) {
310
+ let first = start;
311
+ let last = tier === 0 ? n : end;
312
+ let arcEnd = w === null ? arcCountOf(core) : Math.min(w.end, nextStart);
313
+ if (w !== null && this.perm === null) {
314
+ first = w.rowFirst;
315
+ last = w.rowLast + 1;
316
+ arcEnd = w.end;
317
+ }
318
+ const rows = last - first;
319
+ if (rows <= 0) {
320
+ continue;
321
+ }
322
+ const params = this.scope.params(RANGE_PARAMS, {
323
+ start: first,
324
+ end: last,
325
+ arcBase: w === null ? 0 : w.start,
326
+ arcEnd,
327
+ accumulate: w !== null || this.accumulate ? 1 : 0,
328
+ n,
329
+ });
330
+ const bound = kernel.bind({ ...graphBindings(core, this.perm), out, P: params.binding });
331
+ let rowsPerGroup = wg;
332
+ if (tier === 2) {
333
+ rowsPerGroup = 1;
334
+ } else if (tier === 1) {
335
+ rowsPerGroup = wg / MID_TIER_LANES;
336
+ }
337
+ kernel.dispatch(pass, bound, plan1d(rows, rowsPerGroup, this.scope.caps), [params.offset]);
338
+ this.dispatches++;
339
+ }
211
340
  }
212
341
  }
213
342
 
214
343
  /**
215
- * Prepares the thread-per-row pipeline for a snapshot's dummy pattern (USE_PERM, HAS_WEIGHTS) and snippet.
344
+ * Prepares the tier pipelines for a snapshot's dummy pattern (USE_PERM = tiers !== null, HAS_WEIGHTS) and snippet:
345
+ * TIER 0 always; TIER 1 iff a row of degree 32..1023 exists (segmentOffsets[2] > segmentOffsets[1]); TIER 2 iff a
346
+ * row of degree >= 1024 exists (segmentOffsets[1] > 0). The mid tier folds 32 lanes per row, so a device whose
347
+ * workgroup size is below 32 is E_UNSUPPORTED { feature: "segmentedReduce.tiers" }.
216
348
  * @param scope - the reduce scope (pipelines, pool, params writer)
217
- * @param core - the core whose weights pattern selects HAS_WEIGHTS (USE_PERM is false: no tiers at P2)
218
- * @param options - operator, snippet, tiers (must be null), accumulate
349
+ * @param core - the core whose weights pattern selects HAS_WEIGHTS
350
+ * @param options - operator, snippet, tiers (null: the single thread-per-row dispatch), accumulate
219
351
  * @returns the planner
220
352
  */
221
353
  export async function prepareSegmentedReduce(
@@ -223,16 +355,42 @@ export async function prepareSegmentedReduce(
223
355
  core: CoreBinding,
224
356
  options: SegmentedReduceOptions,
225
357
  ): Promise<SegmentedReducePlanner> {
226
- if (options.tiers !== null) {
227
- throw new WebGpuGraphError("E_UNSUPPORTED", "segmentedReduce: the degree tiers land at P4; pass tiers: null", {
228
- feature: "segmentedReduce.tiers",
229
- });
230
- }
231
- assertNotWindowed(core, "segmentedReduce");
232
358
  const op = opCode(options.op);
233
359
  validateValueSnippet(options.valueSnippet);
234
- const overrides = { ...graphOverrides(core, null), OP: op, TIER: 0 };
235
- const spec = kernelSpec("segmented-reduce", overrides, { VALUE: options.valueSnippet });
236
- const kernel = await scope.pipelines.kernel(spec);
237
- return new ThreadPerRowPlanner(scope, kernel, core.weights !== null, options.accumulate === true);
360
+ const { tiers } = options;
361
+ const perm = tiers?.perm ?? null;
362
+ if (tiers !== null && scope.workgroupSize < MID_TIER_LANES) {
363
+ throw new WebGpuGraphError(
364
+ "E_UNSUPPORTED",
365
+ `segmentedReduce: the mid tier folds ${MID_TIER_LANES} lanes per row, more than the workgroup size ${scope.workgroupSize}`,
366
+ { feature: "segmentedReduce.tiers" },
367
+ );
368
+ }
369
+ const ranges: readonly { readonly tier: 0 | 1 | 2; readonly start: number; readonly end: number }[] =
370
+ tiers === null
371
+ ? [{ tier: 0, start: 0, end: rowCountOf(core, "segmentedReduce") }]
372
+ : [
373
+ { tier: 2, start: 0, end: tiers.segmentOffsets[1] },
374
+ { tier: 1, start: tiers.segmentOffsets[1], end: tiers.segmentOffsets[2] },
375
+ { tier: 0, start: tiers.segmentOffsets[2], end: tiers.segmentOffsets[4] },
376
+ ];
377
+ const compiled: TierDispatch[] = [];
378
+ for (const range of ranges) {
379
+ if (range.tier !== 0 && range.end <= range.start) {
380
+ continue;
381
+ }
382
+ const overrides = { ...graphOverrides(core, perm), OP: op, TIER: range.tier };
383
+ const spec = kernelSpec("segmented-reduce", overrides, { VALUE: options.valueSnippet });
384
+ compiled.push({ ...range, kernel: await scope.pipelines.kernel(spec) });
385
+ }
386
+ const fill = await scope.pipelines.kernel(kernelSpec("fill"));
387
+ return new TieredPlanner(
388
+ scope,
389
+ compiled,
390
+ fill,
391
+ perm,
392
+ core.weights !== null,
393
+ options.accumulate === true,
394
+ identityFillWord(options.op),
395
+ );
238
396
  }
@@ -1,7 +1,6 @@
1
1
  /**
2
- * The pull SpMV primitive of spec 6 row 9 / 8.2 in its thread-per-row tier (P7; M8b plan PD-1 / PD-2): one
3
- * grid-stride dispatch of the `spmv-pull` module over the rows of a REVERSE adjacency, each invocation folding
4
- * `weight * xNorm[nbr]` over its row's in-arcs (Kahan-compensated, f32) and writing
2
+ * The pull SpMV primitive of spec 6 row 9 / 8.2 (P7; M8b plan PD-1 / PD-2): the `spmv-pull` module over the rows of
3
+ * a REVERSE adjacency, each row folding `weight * xNorm[nbr]` over its in-arcs (a two-level f32 sum) and writing
5
4
  * `rankOut[v] = beta * pv + alpha * (sum + danglingMass * pv)`, where `pv` is `personalization[v]` under
6
5
  * HAS_PERSONALIZATION and the uniform `P.uniformP` otherwise, and `danglingMass` is `partials[0].danglingMass`
7
6
  * under USE_DANGLING. PageRank sets alpha to the damping, beta to 1 - alpha and USE_DANGLING; HITS and eigenvector
@@ -9,20 +8,21 @@
9
8
  *
10
9
  * It is its own registry entry rather than a `segmentedReduce` VALUE snippet (PD-1): the snippet vocabulary is
11
10
  * `row, arc, nbr, weight, v` and cannot read `xNorm[nbr]`, and the design's binding table gives the kernel eight
12
- * storage bindings of its own. The in-degree tiers of `reverseDegreeOrder()` land at P4 (PD-2): `tiers !== null` is
13
- * E_UNSUPPORTED and USE_PERM is always false (the perm slot carries the rowPtr dummy of graphBindings). The row and
14
- * arc counts come from the core's binding sizes exactly as segmentedReduce derives them.
11
+ * storage bindings of its own. Without tiers it is ONE grid-stride dispatch (TIER 0, USE_PERM false: the perm slot
12
+ * carries the rowPtr dummy of graphBindings); with the in-degree tiers of `reverseDegreeOrder()` (P4-T5, PD-6) it
13
+ * is up to three dispatches over the permuted rows -- TIER 2 one workgroup per row over [0, hiEnd), TIER 1 32 lanes
14
+ * per row over [hiEnd, midEnd), TIER 0 grid-stride over [midEnd, n) -- each compiled only when its range is
15
+ * non-empty. The row and arc counts come from the core's binding sizes exactly as segmentedReduce derives them.
15
16
  */
16
17
 
17
18
  import { WebGpuGraphError } from "../errors.js";
18
- import { planGridStride } from "../kernel/dispatch.js";
19
+ import { plan1d, planGridStride } from "../kernel/dispatch.js";
19
20
  import { type Kernel } from "../kernel/kernel.js";
20
21
  import { graphBindings, graphOverrides, kernelSpec, SPMV_PARAMS } from "../kernels.js";
21
22
  import { type CoreBinding } from "../memory/residency.js";
22
23
  import { type Binding } from "../types/memory.js";
23
- import { arcCountOf, assertNotWindowed, rowCountOf } from "./core-shape.js";
24
+ import { arcCountOf, assertNotWindowed, type DegreeTiers, MID_TIER_LANES, rowCountOf } from "./core-shape.js";
24
25
  import { type ReduceScope } from "./reduce.js";
25
- import { type DegreeTiers } from "./segmented-reduce.js";
26
26
 
27
27
  /** The group-1 bindings of one pull: the pre-scaled input, the output, the personalization (null binds xNorm as the dummy) and the PrPartial block whose header carries danglingMass. */
28
28
  export interface SpmvResources {
@@ -39,7 +39,7 @@ export interface SpmvCoefficients {
39
39
  readonly uniformP: number;
40
40
  }
41
41
 
42
- /** Options of prepareSpmvPull: the two variant flags, the weights binding and the tiers (must be null at P7). */
42
+ /** Options of prepareSpmvPull: the two variant flags, the weights binding and the in-degree tiers (null: one grid-stride dispatch). */
43
43
  export interface SpmvPullOptions {
44
44
  readonly personalization: boolean;
45
45
  readonly dangling: boolean;
@@ -48,47 +48,74 @@ export interface SpmvPullOptions {
48
48
  readonly tiers: DegreeTiers | null;
49
49
  }
50
50
 
51
- /** A prepared pull: records ONE grid-stride dispatch over the rows of a reverse core into a pass. */
51
+ /** A prepared pull: records one grid-stride dispatch (no tiers) or up to three tier dispatches over the rows of a reverse core into a pass. */
52
52
  export interface SpmvPullPlanner {
53
- /** Records the dispatch over rows [0, n) of `rev` writing rankOut[v] (f32) per row; nothing for n = 0. */
54
- record(pass: GPUComputePassEncoder, rev: CoreBinding, resources: SpmvResources, coefficients: SpmvCoefficients): void;
55
- /** Dispatches the last record() issued (1, or 0 for n = 0). */
53
+ /** Records the dispatches over rows [0, n) of `rev` writing rankOut[v] (f32) per row; nothing for n = 0. */
54
+ record(
55
+ pass: GPUComputePassEncoder,
56
+ rev: CoreBinding,
57
+ resources: SpmvResources,
58
+ coefficients: SpmvCoefficients,
59
+ ): void;
60
+ /** Dispatches the last record() issued (0 for n = 0; 1 without tiers; 1 to 3 with them). */
56
61
  readonly lastDispatches: number;
57
62
  }
58
63
 
59
- /** The thread-per-row planner: ONE `spmv-pull` dispatch, grid-stride over every row. */
64
+ /** One tier's compiled pipeline and the rows [start, end) of the permutation it covers per record() (a tier whose range is empty is not compiled). */
65
+ interface TierDispatch {
66
+ readonly tier: 0 | 1 | 2;
67
+ readonly kernel: Kernel;
68
+ readonly start: number;
69
+ readonly end: number;
70
+ }
71
+
72
+ /**
73
+ * The tiered planner: without tiers ONE `spmv-pull` dispatch, grid-stride over every row; with tiers TIER 2 over
74
+ * [0, hiEnd) (one workgroup per row), then TIER 1 over [hiEnd, midEnd) (WG / 32 rows per workgroup), then TIER 0
75
+ * grid-stride over [midEnd, n), each with its own SpmvParams record (`n` is the range END) and the perm binding.
76
+ */
60
77
  class SpmvPullPlannerImpl implements SpmvPullPlanner {
61
78
  private readonly scope: ReduceScope;
62
- private readonly kernel: Kernel;
79
+ private readonly tiers: readonly TierDispatch[];
80
+ private readonly perm: Binding | null;
63
81
  private readonly weights: Binding | null | undefined;
64
82
  private dispatches = 0;
65
83
 
66
84
  /**
67
- * Wraps a compiled pipeline with the weights choice it was compiled for.
68
- * @param scope - the scope the pipeline was prepared in
69
- * @param kernel - the compiled kernel
70
- * @param weights - the weights option the pipeline's HAS_WEIGHTS was derived from; record() binds the same way
85
+ * Wraps the compiled tier pipelines with the choices they were compiled for.
86
+ * @param scope - the scope the pipelines were prepared in
87
+ * @param tiers - the compiled tiers in dispatch order (TIER 2, 1, 0; only the non-empty ones)
88
+ * @param perm - the reverseDegreeOrder permutation binding (null: USE_PERM false, rows are node indices)
89
+ * @param weights - the weights option the pipelines' HAS_WEIGHTS was derived from; record() binds the same way
71
90
  */
72
- constructor(scope: ReduceScope, kernel: Kernel, weights: Binding | null | undefined) {
91
+ constructor(
92
+ scope: ReduceScope,
93
+ tiers: readonly TierDispatch[],
94
+ perm: Binding | null,
95
+ weights: Binding | null | undefined,
96
+ ) {
73
97
  this.scope = scope;
74
- this.kernel = kernel;
98
+ this.tiers = tiers;
99
+ this.perm = perm;
75
100
  this.weights = weights;
76
101
  }
77
102
 
78
103
  /**
79
104
  * Dispatches the last record() issued.
80
- * @returns 1, or 0 when the last record covered no rows
105
+ * @returns 0 when the last record covered no rows, else the number of non-empty tier ranges
81
106
  */
82
107
  get lastDispatches(): number {
83
108
  return this.dispatches;
84
109
  }
85
110
 
86
111
  /**
87
- * Records the dispatch: rows [0, n), arcs [0, arcCount), planGridStride(n); nothing for n = 0 (no zero-length
88
- * binding is ever created). `personalization ?? xNorm` follows the group-0 dummy rule: both slots are
89
- * storage-ro, so the aliasing check of Kernel.bind does not fire, and HAS_PERSONALIZATION false never reads it.
112
+ * Records the dispatches: every tier's rows over the arcs [0, arcCount), each with its own params record (TIER 0
113
+ * grid-stride through planGridStride, the tiers through plan1d); nothing for n = 0 (no zero-length binding is
114
+ * ever created). `personalization ?? xNorm` follows the group-0 dummy rule: both slots are storage-ro, so the
115
+ * aliasing check of Kernel.bind does not fire, and HAS_PERSONALIZATION false never reads it.
90
116
  * @param pass - the pass to record into
91
- * @param rev - the reverse core (any snapshot with the weights pattern the planner was prepared for)
117
+ * @param rev - the reverse core (any snapshot with the weights pattern the planner was prepared for; with tiers,
118
+ * the one the tiers were built from)
92
119
  * @param resources - xNorm, rankOut, personalization, partials
93
120
  * @param coefficients - alpha, beta, uniformP
94
121
  */
@@ -99,39 +126,53 @@ class SpmvPullPlannerImpl implements SpmvPullPlanner {
99
126
  coefficients: SpmvCoefficients,
100
127
  ): void {
101
128
  const n = rowCountOf(rev, "spmvPull");
102
- const plan = planGridStride(n, this.scope.workgroupSize, this.scope.caps);
103
- if (plan.x === 0) {
104
- this.dispatches = 0;
105
- return;
129
+ const wg = this.scope.workgroupSize;
130
+ this.dispatches = 0;
131
+ for (const { tier, kernel, start, end } of this.tiers) {
132
+ const last = tier === 0 ? n : end;
133
+ const rows = last - start;
134
+ if (rows <= 0) {
135
+ continue;
136
+ }
137
+ const plan =
138
+ tier === 0
139
+ ? planGridStride(rows, wg, this.scope.caps)
140
+ : plan1d(rows, tier === 2 ? 1 : wg / MID_TIER_LANES, this.scope.caps);
141
+ if (plan.x === 0) {
142
+ continue;
143
+ }
144
+ const params = this.scope.params(SPMV_PARAMS, {
145
+ n: last,
146
+ arcBase: 0,
147
+ arcEnd: arcCountOf(rev),
148
+ stride: plan.stride ?? rows,
149
+ alpha: coefficients.alpha,
150
+ beta: coefficients.beta,
151
+ uniformP: coefficients.uniformP,
152
+ start,
153
+ });
154
+ const bound = kernel.bind({
155
+ ...graphBindings(rev, this.perm, this.weights),
156
+ xNorm: resources.xNorm,
157
+ rankOut: resources.rankOut,
158
+ personalization: resources.personalization ?? resources.xNorm,
159
+ partials: resources.partials,
160
+ P: params.binding,
161
+ });
162
+ kernel.dispatch(pass, bound, plan, [params.offset]);
163
+ this.dispatches++;
106
164
  }
107
- const params = this.scope.params(SPMV_PARAMS, {
108
- n,
109
- arcBase: 0,
110
- arcEnd: arcCountOf(rev),
111
- stride: plan.stride ?? n,
112
- alpha: coefficients.alpha,
113
- beta: coefficients.beta,
114
- uniformP: coefficients.uniformP,
115
- pad0: 0,
116
- });
117
- const bound = this.kernel.bind({
118
- ...graphBindings(rev, null, this.weights),
119
- xNorm: resources.xNorm,
120
- rankOut: resources.rankOut,
121
- personalization: resources.personalization ?? resources.xNorm,
122
- partials: resources.partials,
123
- P: params.binding,
124
- });
125
- this.kernel.dispatch(pass, bound, plan, [params.offset]);
126
- this.dispatches = 1;
127
165
  }
128
166
  }
129
167
 
130
168
  /**
131
- * Prepares the thread-per-row pull pipeline for a reverse core's weights pattern and the two variant flags.
169
+ * Prepares the pull pipelines for a reverse core's weights pattern and the two variant flags: TIER 0 always
170
+ * (USE_PERM = tiers !== null); TIER 1 iff a row of in-degree 32..1023 exists (segmentOffsets[2] >
171
+ * segmentOffsets[1]); TIER 2 iff a row of in-degree >= 1024 exists (segmentOffsets[1] > 0). The mid tier folds 32
172
+ * lanes per row, so a device whose workgroup size is below 32 is E_UNSUPPORTED { feature: "spmvPull.tiers" }.
132
173
  * @param scope - the reduce scope (pipelines, params writer)
133
- * @param rev - the reverse core whose weights pattern selects HAS_WEIGHTS (USE_PERM is false: no tiers at P7)
134
- * @param options - personalization, dangling, weights, tiers (must be null)
174
+ * @param rev - the reverse core whose weights pattern selects HAS_WEIGHTS
175
+ * @param options - personalization, dangling, weights, tiers (null: the single grid-stride dispatch)
135
176
  * @returns the planner
136
177
  */
137
178
  export async function prepareSpmvPull(
@@ -139,17 +180,36 @@ export async function prepareSpmvPull(
139
180
  rev: CoreBinding,
140
181
  options: SpmvPullOptions,
141
182
  ): Promise<SpmvPullPlanner> {
142
- if (options.tiers !== null) {
143
- throw new WebGpuGraphError("E_UNSUPPORTED", "spmvPull: the in-degree tiers land at P4; pass tiers: null", {
144
- feature: "spmvPull.tiers",
183
+ assertNotWindowed(rev, "spmvPull");
184
+ const { tiers } = options;
185
+ const perm = tiers?.perm ?? null;
186
+ if (tiers !== null && scope.workgroupSize < MID_TIER_LANES) {
187
+ throw new WebGpuGraphError(
188
+ "E_UNSUPPORTED",
189
+ `spmvPull: the mid tier folds ${MID_TIER_LANES} lanes per row, more than the workgroup size ${scope.workgroupSize}`,
190
+ { feature: "spmvPull.tiers" },
191
+ );
192
+ }
193
+ const ranges: readonly { readonly tier: 0 | 1 | 2; readonly start: number; readonly end: number }[] =
194
+ tiers === null
195
+ ? [{ tier: 0, start: 0, end: rowCountOf(rev, "spmvPull") }]
196
+ : [
197
+ { tier: 2, start: 0, end: tiers.segmentOffsets[1] },
198
+ { tier: 1, start: tiers.segmentOffsets[1], end: tiers.segmentOffsets[2] },
199
+ { tier: 0, start: tiers.segmentOffsets[2], end: tiers.segmentOffsets[4] },
200
+ ];
201
+ const compiled: TierDispatch[] = [];
202
+ for (const range of ranges) {
203
+ if (range.tier !== 0 && range.end <= range.start) {
204
+ continue;
205
+ }
206
+ const spec = kernelSpec("spmv-pull", {
207
+ ...graphOverrides(rev, perm, options.weights),
208
+ HAS_PERSONALIZATION: options.personalization,
209
+ USE_DANGLING: options.dangling,
210
+ TIER: range.tier,
145
211
  });
212
+ compiled.push({ ...range, kernel: await scope.pipelines.kernel(spec) });
146
213
  }
147
- assertNotWindowed(rev, "spmvPull");
148
- const spec = kernelSpec("spmv-pull", {
149
- ...graphOverrides(rev, null, options.weights),
150
- HAS_PERSONALIZATION: options.personalization,
151
- USE_DANGLING: options.dangling,
152
- });
153
- const kernel = await scope.pipelines.kernel(spec);
154
- return new SpmvPullPlannerImpl(scope, kernel, options.weights);
214
+ return new SpmvPullPlannerImpl(scope, compiled, perm, options.weights);
155
215
  }