@graphty/webgpu-graph-algorithms 0.2.1 → 0.3.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +45 -36
- package/dist/browser.js +1 -1
- package/dist/chunks/{context-E6iKaeuJ.js → context-CRbw2Wyo.js} +178 -19
- package/dist/chunks/{context-E6iKaeuJ.js.map → context-CRbw2Wyo.js.map} +1 -1
- package/dist/node.js +1 -1
- package/dist/src/accelerator.d.ts +9 -6
- package/dist/src/accelerator.d.ts.map +1 -1
- package/dist/src/accelerator.js +85 -6
- package/dist/src/accelerator.js.map +1 -1
- package/dist/src/algorithms/components.d.ts +30 -0
- package/dist/src/algorithms/components.d.ts.map +1 -0
- package/dist/src/algorithms/components.js +300 -0
- package/dist/src/algorithms/components.js.map +1 -0
- package/dist/src/algorithms/pagerank.d.ts +39 -0
- package/dist/src/algorithms/pagerank.d.ts.map +1 -0
- package/dist/src/algorithms/pagerank.js +298 -0
- package/dist/src/algorithms/pagerank.js.map +1 -0
- package/dist/src/algorithms/power-iteration.d.ts +109 -0
- package/dist/src/algorithms/power-iteration.d.ts.map +1 -0
- package/dist/src/algorithms/power-iteration.js +206 -0
- package/dist/src/algorithms/power-iteration.js.map +1 -0
- package/dist/src/algorithms/scope.d.ts +26 -0
- package/dist/src/algorithms/scope.d.ts.map +1 -0
- package/dist/src/algorithms/scope.js +41 -0
- package/dist/src/algorithms/scope.js.map +1 -0
- package/dist/src/algorithms/spectral.d.ts +50 -0
- package/dist/src/algorithms/spectral.d.ts.map +1 -0
- package/dist/src/algorithms/spectral.js +247 -0
- package/dist/src/algorithms/spectral.js.map +1 -0
- package/dist/src/index.d.ts +4 -0
- package/dist/src/index.d.ts.map +1 -1
- package/dist/src/index.js +4 -0
- package/dist/src/index.js.map +1 -1
- package/dist/src/kernel/dispatch.d.ts +4 -1
- package/dist/src/kernel/dispatch.d.ts.map +1 -1
- package/dist/src/kernel/dispatch.js +12 -5
- package/dist/src/kernel/dispatch.js.map +1 -1
- package/dist/src/kernels.d.ts +20 -4
- package/dist/src/kernels.d.ts.map +1 -1
- package/dist/src/kernels.js +172 -2
- package/dist/src/kernels.js.map +1 -1
- package/dist/src/memory/residency.d.ts.map +1 -1
- package/dist/src/memory/residency.js +164 -11
- package/dist/src/memory/residency.js.map +1 -1
- package/dist/src/primitives/core-shape.d.ts +41 -0
- package/dist/src/primitives/core-shape.d.ts.map +1 -0
- package/dist/src/primitives/core-shape.js +89 -0
- package/dist/src/primitives/core-shape.js.map +1 -0
- package/dist/src/primitives/segmented-reduce.d.ts.map +1 -1
- package/dist/src/primitives/segmented-reduce.js +4 -30
- package/dist/src/primitives/segmented-reduce.js.map +1 -1
- package/dist/src/primitives/spmv.d.ts +56 -0
- package/dist/src/primitives/spmv.d.ts.map +1 -0
- package/dist/src/primitives/spmv.js +101 -0
- package/dist/src/primitives/spmv.js.map +1 -0
- package/dist/src/types/accelerator.d.ts +14 -2
- package/dist/src/types/accelerator.d.ts.map +1 -1
- package/dist/src/types/algorithms.d.ts +73 -0
- package/dist/src/types/algorithms.d.ts.map +1 -0
- package/dist/src/types/algorithms.js +17 -0
- package/dist/src/types/algorithms.js.map +1 -0
- package/dist/src/wgsl/pr-finalize.wgsl.d.ts +11 -0
- package/dist/src/wgsl/pr-finalize.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/pr-finalize.wgsl.js +36 -0
- package/dist/src/wgsl/pr-finalize.wgsl.js.map +1 -0
- package/dist/src/wgsl/pr-scale.wgsl.d.ts +14 -0
- package/dist/src/wgsl/pr-scale.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/pr-scale.wgsl.js +48 -0
- package/dist/src/wgsl/pr-scale.wgsl.js.map +1 -0
- package/dist/src/wgsl/spmv-pull.wgsl.d.ts +15 -0
- package/dist/src/wgsl/spmv-pull.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/spmv-pull.wgsl.js +47 -0
- package/dist/src/wgsl/spmv-pull.wgsl.js.map +1 -0
- package/dist/src/wgsl/wcc-compress.wgsl.d.ts +9 -0
- package/dist/src/wgsl/wcc-compress.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/wcc-compress.wgsl.js +26 -0
- package/dist/src/wgsl/wcc-compress.wgsl.js.map +1 -0
- package/dist/src/wgsl/wcc-link-edges.wgsl.d.ts +13 -0
- package/dist/src/wgsl/wcc-link-edges.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/wcc-link-edges.wgsl.js +46 -0
- package/dist/src/wgsl/wcc-link-edges.wgsl.js.map +1 -0
- package/dist/src/wgsl/wcc-link-sample.wgsl.d.ts +11 -0
- package/dist/src/wgsl/wcc-link-sample.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/wcc-link-sample.wgsl.js +45 -0
- package/dist/src/wgsl/wcc-link-sample.wgsl.js.map +1 -0
- package/dist/src/wgsl/wcc-sample.wgsl.d.ts +10 -0
- package/dist/src/wgsl/wcc-sample.wgsl.d.ts.map +1 -0
- package/dist/src/wgsl/wcc-sample.wgsl.js +18 -0
- package/dist/src/wgsl/wcc-sample.wgsl.js.map +1 -0
- package/dist/tsconfig.build.tsbuildinfo +1 -1
- package/dist/webgpu-graph-algorithms.js +1550 -29
- package/dist/webgpu-graph-algorithms.js.map +1 -1
- package/package.json +1 -1
- package/src/accelerator.ts +101 -7
- package/src/algorithms/components.ts +348 -0
- package/src/algorithms/pagerank.ts +343 -0
- package/src/algorithms/power-iteration.ts +278 -0
- package/src/algorithms/scope.ts +52 -0
- package/src/algorithms/spectral.ts +300 -0
- package/src/index.ts +18 -0
- package/src/kernel/dispatch.ts +12 -5
- package/src/kernels.ts +206 -5
- package/src/memory/residency.ts +200 -11
- package/src/primitives/core-shape.ts +103 -0
- package/src/primitives/segmented-reduce.ts +4 -36
- package/src/primitives/spmv.ts +155 -0
- package/src/types/accelerator.ts +28 -2
- package/src/types/algorithms.ts +82 -0
- package/src/wgsl/pr-finalize.wgsl.ts +36 -0
- package/src/wgsl/pr-scale.wgsl.ts +48 -0
- package/src/wgsl/spmv-pull.wgsl.ts +47 -0
- package/src/wgsl/wcc-compress.wgsl.ts +26 -0
- package/src/wgsl/wcc-link-edges.wgsl.ts +46 -0
- package/src/wgsl/wcc-link-sample.wgsl.ts +45 -0
- package/src/wgsl/wcc-sample.wgsl.ts +18 -0
|
@@ -0,0 +1,343 @@
|
|
|
1
|
+
/**
|
|
2
|
+
* PageRank and personalized PageRank on the device (spec 8.2, 3.3): a pull over the reverse adjacency with the
|
|
3
|
+
* out-weight normaliser folded into `xNorm` by `pr-scale`, the dangling mass and the L1 delta folded by
|
|
4
|
+
* `pr-finalize` into the header at `partials[0]`, and `spmv-pull` writing the next iterate. Iterations run in
|
|
5
|
+
* batches of PR_BATCH per submit, ONE readback per batch (the header and the iterate together), and the device
|
|
6
|
+
* records `firstConverged` the first time the delta falls below `tolerance * n`, so the reported `iterations` is the
|
|
7
|
+
* first converged iteration and not the batch boundary (spec 9.7).
|
|
8
|
+
*
|
|
9
|
+
* PLAN DECISION PD-7: the ping-pong is TWO buffers (`rankA`, `rankB`) alternated through two cached bind groups;
|
|
10
|
+
* one buffer of two aliased halves would bind (each kernel sees the halves in one access mode) but saves nothing.
|
|
11
|
+
* PLAN DECISION PD-8: `outWeightSum` is call scratch from the Lease, one `segmentedReduce` pass per call.
|
|
12
|
+
* PLAN DECISION PD-9: `pr-scale` at iteration i reads x(i-1) and x(i-2), so its delta is the error of iteration
|
|
13
|
+
* i - 1 and `pr-finalize` records `P.iteration - 1`; the ping-pong supplies `rankPrev` for free because it is the
|
|
14
|
+
* buffer this iteration overwrites, and `pr-scale` runs before `spmv-pull` in the same pass.
|
|
15
|
+
*/
|
|
16
|
+
|
|
17
|
+
import { type F32, type GraphSnapshot } from "@graphty/graph-format";
|
|
18
|
+
|
|
19
|
+
import { U32_MAX } from "../constants.js";
|
|
20
|
+
import { type GpuContext } from "../context.js";
|
|
21
|
+
import { isWebGpuGraphError, WebGpuGraphError } from "../errors.js";
|
|
22
|
+
import { CommandBatch } from "../kernel/batch.js";
|
|
23
|
+
import { groupsOf, plan1d } from "../kernel/dispatch.js";
|
|
24
|
+
import { kernelSpec, PR_PARAMS, PR_PARTIAL } from "../kernels.js";
|
|
25
|
+
import { type ArrayBinding, type CoreBinding } from "../memory/residency.js";
|
|
26
|
+
import { coreOfView } from "../primitives/core-shape.js";
|
|
27
|
+
import { prepareSegmentedReduce } from "../primitives/segmented-reduce.js";
|
|
28
|
+
import { prepareSpmvPull } from "../primitives/spmv.js";
|
|
29
|
+
import { type GpuPageRankResult, type PageRankOptions } from "../types/algorithms.js";
|
|
30
|
+
import { type Binding } from "../types/memory.js";
|
|
31
|
+
import { type GpuRunOptions } from "../types/run.js";
|
|
32
|
+
import { algorithmScope } from "./scope.js";
|
|
33
|
+
|
|
34
|
+
/** Iterations per submit (spec 8.2: k = 8). */
|
|
35
|
+
const PR_BATCH = 8;
|
|
36
|
+
/** Params slots of one batch: per iteration one PrParams (shared by pr-scale and pr-finalize) and one SpmvParams, plus the normaliser's RangeParams; a smaller ring wraps onto a slot the same batch still reads. */
|
|
37
|
+
const RING_SLOTS = 2 * PR_BATCH + 1;
|
|
38
|
+
|
|
39
|
+
/**
|
|
40
|
+
* Validates `options.dest` for a score result of `n` elements.
|
|
41
|
+
* @param dest - the caller's destination array, if any
|
|
42
|
+
* @param n - the node count
|
|
43
|
+
* @param algorithm - the caller's name, for the message
|
|
44
|
+
* @returns the destination as an F32, or null when none was given
|
|
45
|
+
*/
|
|
46
|
+
function checkDest(dest: Float32Array | Uint32Array | undefined, n: number, algorithm: string): F32 | null {
|
|
47
|
+
if (dest === undefined) {
|
|
48
|
+
return null;
|
|
49
|
+
}
|
|
50
|
+
if (dest instanceof Float32Array && dest.length === n && dest.buffer instanceof ArrayBuffer) {
|
|
51
|
+
return dest as F32;
|
|
52
|
+
}
|
|
53
|
+
throw new WebGpuGraphError(
|
|
54
|
+
"E_INVALID_ARGUMENT",
|
|
55
|
+
`${algorithm}: dest must be a Float32Array of length ${n} over an ArrayBuffer`,
|
|
56
|
+
{
|
|
57
|
+
argument: "dest",
|
|
58
|
+
value: `${dest.constructor.name}(${dest.length})`,
|
|
59
|
+
expected: `Float32Array(${n}) over an ArrayBuffer`,
|
|
60
|
+
},
|
|
61
|
+
);
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
/**
|
|
65
|
+
* The resident core; a windowed plan (`E_TOO_LARGE { path: "windowed", algorithm: null }`, spec 3.8) is re-thrown
|
|
66
|
+
* with the algorithm name (3.12). Transcribed from degree.ts, whose coreOf is module-private.
|
|
67
|
+
* @param ctx - the context
|
|
68
|
+
* @param s - the snapshot
|
|
69
|
+
* @param algorithm - the caller's name
|
|
70
|
+
* @returns the core binding
|
|
71
|
+
*/
|
|
72
|
+
function coreOf(ctx: GpuContext, s: GraphSnapshot, algorithm: string): CoreBinding {
|
|
73
|
+
try {
|
|
74
|
+
return ctx.residency.core(s);
|
|
75
|
+
} catch (error: unknown) {
|
|
76
|
+
if (isWebGpuGraphError(error) && error.code === "E_TOO_LARGE" && error.details.path === "windowed") {
|
|
77
|
+
throw new WebGpuGraphError(
|
|
78
|
+
"E_TOO_LARGE",
|
|
79
|
+
`${algorithm}: the arc arrays need a windowed upload, which P1-P3 plan but do not execute`,
|
|
80
|
+
{ ...error.details, algorithm },
|
|
81
|
+
);
|
|
82
|
+
}
|
|
83
|
+
throw error;
|
|
84
|
+
}
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
/**
|
|
88
|
+
* Whole-buffer binding of a scratch buffer over its first `size` bytes.
|
|
89
|
+
* @param buffer - the buffer
|
|
90
|
+
* @param size - the bound byte length
|
|
91
|
+
* @returns the binding
|
|
92
|
+
*/
|
|
93
|
+
function bindingOf(buffer: GPUBuffer, size: number): Binding {
|
|
94
|
+
return { buffer, offset: 0, size, window: null };
|
|
95
|
+
}
|
|
96
|
+
|
|
97
|
+
/**
|
|
98
|
+
* The E_ABORTED error of a signal.
|
|
99
|
+
* @param algorithm - the caller's name
|
|
100
|
+
* @param batchId - the last submitted batch, when one exists
|
|
101
|
+
* @returns the error
|
|
102
|
+
*/
|
|
103
|
+
function aborted(algorithm: string, batchId?: number): WebGpuGraphError {
|
|
104
|
+
return new WebGpuGraphError(
|
|
105
|
+
"E_ABORTED",
|
|
106
|
+
`${algorithm}: the signal was aborted`,
|
|
107
|
+
batchId === undefined ? {} : { batchId },
|
|
108
|
+
);
|
|
109
|
+
}
|
|
110
|
+
|
|
111
|
+
/**
|
|
112
|
+
* The shared driver: `personalization` is the normalised vector of personalizedPageRank or null for the uniform
|
|
113
|
+
* 1 / n of pageRank.
|
|
114
|
+
* @param ctx - the context
|
|
115
|
+
* @param s - the snapshot
|
|
116
|
+
* @param personalization - the normalised personalization (sum 1), or null
|
|
117
|
+
* @param options - the PageRank and run options
|
|
118
|
+
* @param algorithm - the public name, for messages and labels
|
|
119
|
+
* @returns the result
|
|
120
|
+
*/
|
|
121
|
+
async function run(
|
|
122
|
+
ctx: GpuContext,
|
|
123
|
+
s: GraphSnapshot,
|
|
124
|
+
personalization: F32 | null,
|
|
125
|
+
options: (PageRankOptions & GpuRunOptions) | undefined,
|
|
126
|
+
algorithm: string,
|
|
127
|
+
): Promise<GpuPageRankResult> {
|
|
128
|
+
ctx.assertReady();
|
|
129
|
+
const n = s.nodeCount;
|
|
130
|
+
const alpha = options?.dampingFactor ?? 0.85;
|
|
131
|
+
const maxIterations = options?.maxIterations ?? 100;
|
|
132
|
+
const tolerance = options?.tolerance ?? 1e-6;
|
|
133
|
+
const useWeights = options?.weighted !== false;
|
|
134
|
+
if (!Number.isInteger(maxIterations) || maxIterations < 1) {
|
|
135
|
+
throw new WebGpuGraphError("E_INVALID_ARGUMENT", `${algorithm}: maxIterations must be a positive integer`, {
|
|
136
|
+
argument: "maxIterations",
|
|
137
|
+
value: maxIterations,
|
|
138
|
+
expected: "a positive integer",
|
|
139
|
+
});
|
|
140
|
+
}
|
|
141
|
+
const dest = checkDest(options?.dest, n, algorithm);
|
|
142
|
+
if (options?.signal?.aborted) {
|
|
143
|
+
throw aborted(algorithm);
|
|
144
|
+
}
|
|
145
|
+
if (n === 0) {
|
|
146
|
+
options?.onProgress?.(maxIterations, maxIterations);
|
|
147
|
+
return { scores: dest ?? new Float32Array(0), iterations: 0, converged: true, danglingMass: 0, precision: "f32" };
|
|
148
|
+
}
|
|
149
|
+
const core = coreOf(ctx, s, algorithm);
|
|
150
|
+
if (s.arcCount === 0) {
|
|
151
|
+
// every node is dangling: the whole mass is redistributed by pv each iteration, so pv is the fixed point
|
|
152
|
+
const scores = dest ?? new Float32Array(n);
|
|
153
|
+
if (personalization === null) {
|
|
154
|
+
scores.fill(1 / n);
|
|
155
|
+
} else {
|
|
156
|
+
scores.set(personalization);
|
|
157
|
+
}
|
|
158
|
+
options?.onProgress?.(maxIterations, maxIterations);
|
|
159
|
+
return { scores, iterations: 0, converged: true, danglingMass: 1, precision: "f32" };
|
|
160
|
+
}
|
|
161
|
+
const view = ctx.residency.view(s, "reverse");
|
|
162
|
+
const rev = coreOfView(view, view.scalars.arcCount[0]);
|
|
163
|
+
// weighted: false runs the unweighted algorithm on BOTH sides: the normaliser sums 1 per arc and the pull folds 1
|
|
164
|
+
const weights: Binding | null | undefined = useWeights ? undefined : null;
|
|
165
|
+
const weightedCore: CoreBinding = useWeights ? core : { ...core, weights: null, hasWeights: false };
|
|
166
|
+
const weightedRev: CoreBinding = useWeights ? rev : { ...rev, weights: null, hasWeights: false };
|
|
167
|
+
const scope = algorithmScope(ctx, algorithm, RING_SLOTS);
|
|
168
|
+
let uploaded: ArrayBinding | null = null;
|
|
169
|
+
try {
|
|
170
|
+
const bytes = 4 * n;
|
|
171
|
+
const rankA = scope.scratch(bytes, "rankA");
|
|
172
|
+
const rankB = scope.scratch(bytes, "rankB");
|
|
173
|
+
const xNorm = scope.scratch(bytes, "xNorm");
|
|
174
|
+
const outWeightSum = scope.scratch(bytes, "outWeightSum");
|
|
175
|
+
const scalePlan = plan1d(n, ctx.workgroupSize, ctx.caps);
|
|
176
|
+
const groups = groupsOf(scalePlan);
|
|
177
|
+
const partialsBytes = PR_PARTIAL.byteLength * (1 + groups);
|
|
178
|
+
const partials = scope.scratch(partialsBytes, "partials");
|
|
179
|
+
if (personalization !== null) {
|
|
180
|
+
uploaded = ctx.residency.array(personalization, `${algorithm}/personalization`);
|
|
181
|
+
}
|
|
182
|
+
await ctx.allocator.check();
|
|
183
|
+
const normaliser = await prepareSegmentedReduce(scope, weightedCore, {
|
|
184
|
+
op: "sum",
|
|
185
|
+
valueSnippet: "v = weight;",
|
|
186
|
+
tiers: null,
|
|
187
|
+
});
|
|
188
|
+
const pull = await prepareSpmvPull(scope, weightedRev, {
|
|
189
|
+
personalization: personalization !== null,
|
|
190
|
+
dangling: true,
|
|
191
|
+
weights,
|
|
192
|
+
tiers: null,
|
|
193
|
+
});
|
|
194
|
+
const scale = await ctx.pipelines.kernel(kernelSpec("pr-scale", { NORM_MODE: 0 }));
|
|
195
|
+
const finalize = await ctx.pipelines.kernel(kernelSpec("pr-finalize", { NORM_MODE: 0 }));
|
|
196
|
+
const finalizePlan = plan1d(1, ctx.workgroupSize, ctx.caps);
|
|
197
|
+
const { queue } = ctx.device;
|
|
198
|
+
queue.writeBuffer(rankA, 0, new Float32Array(n).fill(1 / n));
|
|
199
|
+
queue.writeBuffer(rankB, 0, new Float32Array(n));
|
|
200
|
+
const header = new ArrayBuffer(PR_PARTIAL.byteLength);
|
|
201
|
+
PR_PARTIAL.write(new DataView(header), { firstConverged: U32_MAX, iteration: 0 });
|
|
202
|
+
queue.writeBuffer(partials, 0, header);
|
|
203
|
+
const rank = [bindingOf(rankA, bytes), bindingOf(rankB, bytes)];
|
|
204
|
+
const xNormBinding = bindingOf(xNorm, bytes);
|
|
205
|
+
const outWeightSumBinding = bindingOf(outWeightSum, bytes);
|
|
206
|
+
const partialsBinding = bindingOf(partials, partialsBytes);
|
|
207
|
+
const coefficients = { alpha, beta: 1 - alpha, uniformP: 1 / n };
|
|
208
|
+
let cur = 0;
|
|
209
|
+
let iterationsRun = 0;
|
|
210
|
+
for (;;) {
|
|
211
|
+
const k = Math.min(PR_BATCH, maxIterations - iterationsRun);
|
|
212
|
+
const batch = new CommandBatch(ctx, algorithm);
|
|
213
|
+
const pass = batch.pass("iterations");
|
|
214
|
+
if (iterationsRun === 0) {
|
|
215
|
+
normaliser.record(pass, weightedCore, outWeightSumBinding);
|
|
216
|
+
}
|
|
217
|
+
for (let i = 0; i < k; i++) {
|
|
218
|
+
const params = scope.params(PR_PARAMS, {
|
|
219
|
+
n,
|
|
220
|
+
groups,
|
|
221
|
+
iteration: iterationsRun + i + 1,
|
|
222
|
+
trackConvergence: 1,
|
|
223
|
+
convergeThreshold: tolerance * n,
|
|
224
|
+
});
|
|
225
|
+
const other = 1 - cur;
|
|
226
|
+
const scaleBound = scale.bind({
|
|
227
|
+
rankIn: rank[cur],
|
|
228
|
+
rankPrev: rank[other],
|
|
229
|
+
outWeightSum: outWeightSumBinding,
|
|
230
|
+
xNorm: xNormBinding,
|
|
231
|
+
partials: partialsBinding,
|
|
232
|
+
P: params.binding,
|
|
233
|
+
});
|
|
234
|
+
scale.dispatch(pass, scaleBound, scalePlan, [params.offset]);
|
|
235
|
+
const finalizeBound = finalize.bind({ partials: partialsBinding, P: params.binding });
|
|
236
|
+
finalize.dispatch(pass, finalizeBound, finalizePlan, [params.offset]);
|
|
237
|
+
pull.record(
|
|
238
|
+
pass,
|
|
239
|
+
weightedRev,
|
|
240
|
+
{
|
|
241
|
+
xNorm: xNormBinding,
|
|
242
|
+
rankOut: rank[other],
|
|
243
|
+
personalization: uploaded?.binding ?? null,
|
|
244
|
+
partials: partialsBinding,
|
|
245
|
+
},
|
|
246
|
+
coefficients,
|
|
247
|
+
);
|
|
248
|
+
cur = other;
|
|
249
|
+
}
|
|
250
|
+
batch.endPass();
|
|
251
|
+
const headerRequest = batch.readback(partials, 0, PR_PARTIAL.byteLength);
|
|
252
|
+
const scoresRequest = batch.readback(rank[cur].buffer, 0, bytes);
|
|
253
|
+
scope.flush();
|
|
254
|
+
const submitted = batch.submit();
|
|
255
|
+
const back = await submitted.readback;
|
|
256
|
+
iterationsRun += k;
|
|
257
|
+
ctx.assertReady();
|
|
258
|
+
if (options?.signal?.aborted) {
|
|
259
|
+
throw aborted(algorithm, submitted.id);
|
|
260
|
+
}
|
|
261
|
+
options?.onProgress?.(iterationsRun, maxIterations);
|
|
262
|
+
const folded = PR_PARTIAL.read(new DataView(back), headerRequest.offset);
|
|
263
|
+
const firstConverged = folded.firstConverged as number;
|
|
264
|
+
const converged = firstConverged !== U32_MAX;
|
|
265
|
+
if (converged || iterationsRun >= maxIterations) {
|
|
266
|
+
const scores = dest ?? new Float32Array(n);
|
|
267
|
+
scores.set(new Float32Array(back, scoresRequest.offset, n));
|
|
268
|
+
if (converged && iterationsRun < maxIterations) {
|
|
269
|
+
options?.onProgress?.(maxIterations, maxIterations);
|
|
270
|
+
}
|
|
271
|
+
return {
|
|
272
|
+
scores,
|
|
273
|
+
iterations: converged ? firstConverged : iterationsRun,
|
|
274
|
+
converged,
|
|
275
|
+
danglingMass: folded.danglingMass as number,
|
|
276
|
+
precision: "f32",
|
|
277
|
+
};
|
|
278
|
+
}
|
|
279
|
+
}
|
|
280
|
+
} finally {
|
|
281
|
+
uploaded?.destroy();
|
|
282
|
+
scope.dispose();
|
|
283
|
+
}
|
|
284
|
+
}
|
|
285
|
+
|
|
286
|
+
/**
|
|
287
|
+
* PageRank on the device (spec 3.3, 8.2): NetworkX semantics, f32 scores, `iterations` the first converged
|
|
288
|
+
* iteration (spec 9.7).
|
|
289
|
+
* @param ctx - the context whose device runs the kernels
|
|
290
|
+
* @param s - the snapshot (uploaded through ctx.residency, or found there)
|
|
291
|
+
* @param options - dampingFactor 0.85 / maxIterations 100 / tolerance 1e-6 / weighted true, plus dest / signal / onProgress
|
|
292
|
+
* @returns the scores, iterations, converged, danglingMass and precision
|
|
293
|
+
*/
|
|
294
|
+
export function pageRank(
|
|
295
|
+
ctx: GpuContext,
|
|
296
|
+
s: GraphSnapshot,
|
|
297
|
+
options?: PageRankOptions & GpuRunOptions,
|
|
298
|
+
): Promise<GpuPageRankResult> {
|
|
299
|
+
return run(ctx, s, null, options, "pageRank");
|
|
300
|
+
}
|
|
301
|
+
|
|
302
|
+
/**
|
|
303
|
+
* Personalized PageRank on the device (spec 3.3, 8.2): the personalization vector replaces the uniform 1 / n in
|
|
304
|
+
* the teleport and the dangling redistribution; it is normalised to sum 1 on the host.
|
|
305
|
+
* @param ctx - the context whose device runs the kernels
|
|
306
|
+
* @param s - the snapshot (uploaded through ctx.residency, or found there)
|
|
307
|
+
* @param personalization - one finite non-negative mass per node, not all zero
|
|
308
|
+
* @param options - dampingFactor 0.85 / maxIterations 100 / tolerance 1e-6 / weighted true, plus dest / signal / onProgress
|
|
309
|
+
* @returns the scores, iterations, converged, danglingMass and precision
|
|
310
|
+
*/
|
|
311
|
+
export async function personalizedPageRank(
|
|
312
|
+
ctx: GpuContext,
|
|
313
|
+
s: GraphSnapshot,
|
|
314
|
+
personalization: F32,
|
|
315
|
+
options?: PageRankOptions & GpuRunOptions,
|
|
316
|
+
): Promise<GpuPageRankResult> {
|
|
317
|
+
const n = s.nodeCount;
|
|
318
|
+
const invalid = (value: unknown, expected: string): WebGpuGraphError =>
|
|
319
|
+
new WebGpuGraphError("E_INVALID_ARGUMENT", `personalizedPageRank: personalization must be ${expected}`, {
|
|
320
|
+
argument: "personalization",
|
|
321
|
+
value,
|
|
322
|
+
expected,
|
|
323
|
+
});
|
|
324
|
+
if (!(personalization instanceof Float32Array) || personalization.length !== n) {
|
|
325
|
+
throw invalid(`${personalization.constructor.name}(${personalization.length})`, `a Float32Array of length ${n}`);
|
|
326
|
+
}
|
|
327
|
+
let total = 0;
|
|
328
|
+
for (let v = 0; v < n; v++) {
|
|
329
|
+
const mass = personalization[v];
|
|
330
|
+
if (!Number.isFinite(mass) || mass < 0) {
|
|
331
|
+
throw invalid(mass, "finite and non-negative in every entry");
|
|
332
|
+
}
|
|
333
|
+
total += mass;
|
|
334
|
+
}
|
|
335
|
+
if (n > 0 && !(total > 0)) {
|
|
336
|
+
throw invalid(total, "a vector whose entries sum to a positive number");
|
|
337
|
+
}
|
|
338
|
+
const normalised = new Float32Array(n);
|
|
339
|
+
for (let v = 0; v < n; v++) {
|
|
340
|
+
normalised[v] = personalization[v] / total;
|
|
341
|
+
}
|
|
342
|
+
return run(ctx, s, normalised, options, "personalizedPageRank");
|
|
343
|
+
}
|
|
@@ -0,0 +1,278 @@
|
|
|
1
|
+
/**
|
|
2
|
+
* The power-iteration driver HITS, eigenvector centrality and Katz centrality share (spec 8.2 lines 2603-2605; M8b
|
|
3
|
+
* plan PD-10): per iteration `pr-scale` at NORM_MODE 1 | 2 | 4 folds the norm term (and the L1 delta) into the
|
|
4
|
+
* per-workgroup partials, `pr-finalize` folds them into the header at `partials[0]` (the square root for L2) and
|
|
5
|
+
* records `firstConverged`, `pr-scale` at NORM_MODE 3 writes `xNorm[u] = x[u] / partials[0].norm` (skipped at
|
|
6
|
+
* mode 4, the identity, where the first scale pass already wrote `xNorm[u] = x[u]`), and `spmv-pull` writes the
|
|
7
|
+
* next iterate. Four dispatches (three for Katz), all on the device, NO readback inside the batch: the scalar
|
|
8
|
+
* normaliser the host cannot supply without a readback is exactly what stays on the device, so these three batch
|
|
9
|
+
* eight iterations per submit like PageRank does.
|
|
10
|
+
*
|
|
11
|
+
* The ping-pong is TWO buffers alternated through two cached bind groups (PD-7 / DEP-M8B-G), and `rankPrev` is the
|
|
12
|
+
* buffer this iteration overwrites: it still holds x(i-2) when `pr-scale` reads it, because the scale runs before
|
|
13
|
+
* the pull in the same pass (PD-9). With an `alternate` adjacency (HITS, spec 8.2: "alternates two pulls") the
|
|
14
|
+
* iterations pull over `adjacency` and `alternate` in turn, and the ring is THREE buffers so the buffer being
|
|
15
|
+
* overwritten holds x(i-3) -- the previous iterate of the same kind -- and the delta stays meaningful. The batch
|
|
16
|
+
* loop duplicates src/algorithms/pagerank.ts, which its own task owns.
|
|
17
|
+
*/
|
|
18
|
+
|
|
19
|
+
import { type F32, type GraphSnapshot } from "@graphty/graph-format";
|
|
20
|
+
|
|
21
|
+
import { U32_MAX } from "../constants.js";
|
|
22
|
+
import { type GpuContext } from "../context.js";
|
|
23
|
+
import { isWebGpuGraphError, WebGpuGraphError } from "../errors.js";
|
|
24
|
+
import { CommandBatch } from "../kernel/batch.js";
|
|
25
|
+
import { groupsOf, plan1d } from "../kernel/dispatch.js";
|
|
26
|
+
import { kernelSpec, PR_PARAMS, PR_PARTIAL } from "../kernels.js";
|
|
27
|
+
import { type CoreBinding } from "../memory/residency.js";
|
|
28
|
+
import { coreOfView } from "../primitives/core-shape.js";
|
|
29
|
+
import { prepareSpmvPull } from "../primitives/spmv.js";
|
|
30
|
+
import { type Binding } from "../types/memory.js";
|
|
31
|
+
import { algorithmScope } from "./scope.js";
|
|
32
|
+
|
|
33
|
+
/** Iterations per submit (spec 8.2: k = 8). */
|
|
34
|
+
const BATCH = 8;
|
|
35
|
+
/** Params slots of one batch: four blocks per iteration times the batch, plus a margin (the plan's `4 * 8 + 8`). */
|
|
36
|
+
const RING_SLOTS = 4 * BATCH + 8;
|
|
37
|
+
|
|
38
|
+
/**
|
|
39
|
+
* What one power iteration run needs. `adjacency` is the CoreBinding the pull walks -- the forward core for hubs
|
|
40
|
+
* and eigenvector, coreOfView(view(s, "reverse")) for authorities and Katz (spec 8.2 lines 2603-2606). Exported
|
|
41
|
+
* for runPowerIteration's signature, not imported by name.
|
|
42
|
+
* @public
|
|
43
|
+
*/
|
|
44
|
+
export interface PowerIterationConfig {
|
|
45
|
+
/**
|
|
46
|
+
* 1 = L1 (sum) normalise, 2 = L2 normalise, 4 = identity (Katz: no normaliser). Never 0 or 3 here: 0 is
|
|
47
|
+
* PageRank's per-node divisor and 3 is the internal scale pass this driver issues itself.
|
|
48
|
+
*/
|
|
49
|
+
readonly normMode: 1 | 2 | 4;
|
|
50
|
+
readonly adjacency: CoreBinding;
|
|
51
|
+
/**
|
|
52
|
+
* Null: every iteration pulls over `adjacency`. Non-null: the odd iterations (1, 3, ...) pull over `adjacency`
|
|
53
|
+
* and the even ones over `alternate`, so x(i) = A_alt * norm(A * norm(x(i-2))) -- one interleaved chain of HITS,
|
|
54
|
+
* whose hubs and authorities are the last iterates of the two chains (forward-first and reverse-first). The
|
|
55
|
+
* ring then has three buffers and `weights` must be null or undefined (each core's own weights).
|
|
56
|
+
*/
|
|
57
|
+
readonly alternate: CoreBinding | null;
|
|
58
|
+
readonly alpha: number;
|
|
59
|
+
readonly beta: number;
|
|
60
|
+
readonly uniformP: number;
|
|
61
|
+
readonly maxIterations: number;
|
|
62
|
+
readonly tolerance: number;
|
|
63
|
+
/** undefined takes the adjacency's weights, null runs UNWEIGHTED on a weighted snapshot (T4 Step 4). */
|
|
64
|
+
readonly weights: Binding | null | undefined;
|
|
65
|
+
readonly label: string;
|
|
66
|
+
readonly signal: AbortSignal | undefined;
|
|
67
|
+
readonly onProgress: ((done: number, total: number) => void) | undefined;
|
|
68
|
+
}
|
|
69
|
+
|
|
70
|
+
/** The iterate, the device-recorded first converged iteration, and the number of iterations actually run. */
|
|
71
|
+
export interface PowerIterationRun {
|
|
72
|
+
readonly scores: F32;
|
|
73
|
+
/** x(m - 1) on an alternating run: the chain's last iterate of the OTHER kind (hub or authority); null otherwise. */
|
|
74
|
+
readonly previous: F32 | null;
|
|
75
|
+
readonly iterations: number;
|
|
76
|
+
readonly converged: boolean;
|
|
77
|
+
readonly iterationsRun: number;
|
|
78
|
+
}
|
|
79
|
+
|
|
80
|
+
/**
|
|
81
|
+
* Validates `options.dest` for a score result of `n` elements.
|
|
82
|
+
* @param dest - the caller's destination array, if any
|
|
83
|
+
* @param n - the node count
|
|
84
|
+
* @param algorithm - the caller's name, for the message
|
|
85
|
+
* @returns the destination as an F32, or null when none was given
|
|
86
|
+
*/
|
|
87
|
+
export function checkDest(dest: Float32Array | Uint32Array | undefined, n: number, algorithm: string): F32 | null {
|
|
88
|
+
if (dest === undefined) {
|
|
89
|
+
return null;
|
|
90
|
+
}
|
|
91
|
+
if (dest instanceof Float32Array && dest.length === n && dest.buffer instanceof ArrayBuffer) {
|
|
92
|
+
return dest as F32;
|
|
93
|
+
}
|
|
94
|
+
throw new WebGpuGraphError(
|
|
95
|
+
"E_INVALID_ARGUMENT",
|
|
96
|
+
`${algorithm}: dest must be a Float32Array of length ${n} over an ArrayBuffer`,
|
|
97
|
+
{
|
|
98
|
+
argument: "dest",
|
|
99
|
+
value: `${dest.constructor.name}(${dest.length})`,
|
|
100
|
+
expected: `Float32Array(${n}) over an ArrayBuffer`,
|
|
101
|
+
},
|
|
102
|
+
);
|
|
103
|
+
}
|
|
104
|
+
|
|
105
|
+
/**
|
|
106
|
+
* The resident core; a windowed plan (`E_TOO_LARGE { path: "windowed", algorithm: null }`, spec 3.8) is re-thrown
|
|
107
|
+
* with the algorithm name (3.12). Transcribed from degree.ts, whose coreOf is module-private.
|
|
108
|
+
* @param ctx - the context
|
|
109
|
+
* @param s - the snapshot
|
|
110
|
+
* @param algorithm - the caller's name
|
|
111
|
+
* @returns the core binding
|
|
112
|
+
*/
|
|
113
|
+
export function coreOf(ctx: GpuContext, s: GraphSnapshot, algorithm: string): CoreBinding {
|
|
114
|
+
try {
|
|
115
|
+
return ctx.residency.core(s);
|
|
116
|
+
} catch (error: unknown) {
|
|
117
|
+
if (isWebGpuGraphError(error) && error.code === "E_TOO_LARGE" && error.details.path === "windowed") {
|
|
118
|
+
throw new WebGpuGraphError(
|
|
119
|
+
"E_TOO_LARGE",
|
|
120
|
+
`${algorithm}: the arc arrays need a windowed upload, which P1-P3 plan but do not execute`,
|
|
121
|
+
{ ...error.details, algorithm },
|
|
122
|
+
);
|
|
123
|
+
}
|
|
124
|
+
throw error;
|
|
125
|
+
}
|
|
126
|
+
}
|
|
127
|
+
|
|
128
|
+
/**
|
|
129
|
+
* The reverse adjacency as the CoreBinding the pull walks (the forward buffers on an undirected snapshot, PD-13).
|
|
130
|
+
* @param ctx - the context
|
|
131
|
+
* @param s - the snapshot (n > 0)
|
|
132
|
+
* @returns the reverse core
|
|
133
|
+
*/
|
|
134
|
+
export function reverseOf(ctx: GpuContext, s: GraphSnapshot): CoreBinding {
|
|
135
|
+
const view = ctx.residency.view(s, "reverse");
|
|
136
|
+
return coreOfView(view, view.scalars.arcCount[0]);
|
|
137
|
+
}
|
|
138
|
+
|
|
139
|
+
/**
|
|
140
|
+
* The E_ABORTED error of a signal.
|
|
141
|
+
* @param algorithm - the caller's name
|
|
142
|
+
* @param batchId - the last submitted batch, when one exists
|
|
143
|
+
* @returns the error
|
|
144
|
+
*/
|
|
145
|
+
export function aborted(algorithm: string, batchId?: number): WebGpuGraphError {
|
|
146
|
+
return new WebGpuGraphError(
|
|
147
|
+
"E_ABORTED",
|
|
148
|
+
`${algorithm}: the signal was aborted`,
|
|
149
|
+
batchId === undefined ? {} : { batchId },
|
|
150
|
+
);
|
|
151
|
+
}
|
|
152
|
+
|
|
153
|
+
/**
|
|
154
|
+
* Whole-buffer binding of a scratch buffer over its first `size` bytes.
|
|
155
|
+
* @param buffer - the buffer
|
|
156
|
+
* @param size - the bound byte length
|
|
157
|
+
* @returns the binding
|
|
158
|
+
*/
|
|
159
|
+
function whole(buffer: GPUBuffer, size: number): Binding {
|
|
160
|
+
return { buffer, offset: 0, size, window: null };
|
|
161
|
+
}
|
|
162
|
+
|
|
163
|
+
/**
|
|
164
|
+
* Runs the power iteration of `config` over `n` nodes from x(0) = 1 / n, in batches of eight iterations with ONE
|
|
165
|
+
* readback per batch (the header and the iterate together), stopping at the first batch whose header records a
|
|
166
|
+
* converged iteration or at `maxIterations`. The returned iterate is RAW: the last `spmv-pull` output, which the
|
|
167
|
+
* caller normalises once on the host. The iterate travels with the header in EVERY batch, not once after the loop:
|
|
168
|
+
* whether a batch terminates the run is known only from its header, and a second readback for the iterate would
|
|
169
|
+
* be a second mapAsync (the plan's case 8 counts exactly one). The cost is 4n bytes (8n on an alternating run) of
|
|
170
|
+
* staging per batch of 8 that the loop then discards when it continues -- 40 MB per 8 iterations at 10M nodes;
|
|
171
|
+
* pagerank.ts pays the same.
|
|
172
|
+
* @param ctx - the context
|
|
173
|
+
* @param n - the node count (> 0)
|
|
174
|
+
* @param config - the recurrence, the adjacency and the run options
|
|
175
|
+
* @returns the raw iterate, the first converged iteration, whether it converged and the iterations actually run
|
|
176
|
+
*/
|
|
177
|
+
export async function runPowerIteration(
|
|
178
|
+
ctx: GpuContext,
|
|
179
|
+
n: number,
|
|
180
|
+
config: PowerIterationConfig,
|
|
181
|
+
): Promise<PowerIterationRun> {
|
|
182
|
+
const scope = algorithmScope(ctx, config.label, RING_SLOTS);
|
|
183
|
+
try {
|
|
184
|
+
const bytes = 4 * n;
|
|
185
|
+
// x(i) lives in ring[i % ring.length]: two buffers for one pull per iteration, three for the alternating pair
|
|
186
|
+
const ring = (config.alternate === null ? ["rankA", "rankB"] : ["rankA", "rankB", "rankC"]).map((label) =>
|
|
187
|
+
whole(scope.scratch(bytes, label), bytes),
|
|
188
|
+
);
|
|
189
|
+
const xNorm = whole(scope.scratch(bytes, "xNorm"), bytes);
|
|
190
|
+
const scalePlan = plan1d(n, ctx.workgroupSize, ctx.caps);
|
|
191
|
+
const groups = groupsOf(scalePlan);
|
|
192
|
+
const partialsBytes = PR_PARTIAL.byteLength * (1 + groups);
|
|
193
|
+
const partialsBuffer = scope.scratch(partialsBytes, "partials");
|
|
194
|
+
const partials = whole(partialsBuffer, partialsBytes);
|
|
195
|
+
await ctx.allocator.check();
|
|
196
|
+
const pullOptions = { personalization: false, dangling: false, weights: config.weights, tiers: null };
|
|
197
|
+
const pulls = [{ core: config.adjacency, pull: await prepareSpmvPull(scope, config.adjacency, pullOptions) }];
|
|
198
|
+
if (config.alternate !== null) {
|
|
199
|
+
pulls.push({ core: config.alternate, pull: await prepareSpmvPull(scope, config.alternate, pullOptions) });
|
|
200
|
+
}
|
|
201
|
+
const scaleNorm = await ctx.pipelines.kernel(kernelSpec("pr-scale", { NORM_MODE: config.normMode }));
|
|
202
|
+
const scaleApply =
|
|
203
|
+
config.normMode === 4 ? null : await ctx.pipelines.kernel(kernelSpec("pr-scale", { NORM_MODE: 3 }));
|
|
204
|
+
const finalize = await ctx.pipelines.kernel(kernelSpec("pr-finalize", { NORM_MODE: config.normMode }));
|
|
205
|
+
const finalizePlan = plan1d(1, ctx.workgroupSize, ctx.caps);
|
|
206
|
+
const { queue } = ctx.device;
|
|
207
|
+
queue.writeBuffer(ring[0].buffer, 0, new Float32Array(n).fill(1 / n));
|
|
208
|
+
for (const slot of ring.slice(1)) {
|
|
209
|
+
queue.writeBuffer(slot.buffer, 0, new Float32Array(n));
|
|
210
|
+
}
|
|
211
|
+
const header = new ArrayBuffer(PR_PARTIAL.byteLength);
|
|
212
|
+
PR_PARTIAL.write(new DataView(header), { firstConverged: U32_MAX, iteration: 0 });
|
|
213
|
+
queue.writeBuffer(partialsBuffer, 0, header);
|
|
214
|
+
const coefficients = { alpha: config.alpha, beta: config.beta, uniformP: config.uniformP };
|
|
215
|
+
let iterationsRun = 0;
|
|
216
|
+
for (;;) {
|
|
217
|
+
const k = Math.min(BATCH, config.maxIterations - iterationsRun);
|
|
218
|
+
const batch = new CommandBatch(ctx, config.label);
|
|
219
|
+
const pass = batch.pass("iterations");
|
|
220
|
+
for (let i = 0; i < k; i++) {
|
|
221
|
+
const iteration = iterationsRun + i + 1;
|
|
222
|
+
const params = scope.params(PR_PARAMS, {
|
|
223
|
+
n,
|
|
224
|
+
groups,
|
|
225
|
+
iteration,
|
|
226
|
+
trackConvergence: 1,
|
|
227
|
+
convergeThreshold: config.tolerance * n,
|
|
228
|
+
});
|
|
229
|
+
const rankIn = ring[(iteration - 1) % ring.length];
|
|
230
|
+
// the slot the pull overwrites still holds x(iteration - ring.length): that is rankPrev (PD-9)
|
|
231
|
+
const rankOut = ring[iteration % ring.length];
|
|
232
|
+
const { core, pull } = pulls[(iteration - 1) % pulls.length];
|
|
233
|
+
// outWeightSum takes rankIn as its dummy: storage-ro, read only under NORM_MODE 0, never compiled here
|
|
234
|
+
const scaleBindings = { rankIn, rankPrev: rankOut, outWeightSum: rankIn, xNorm, partials, P: params.binding };
|
|
235
|
+
scaleNorm.dispatch(pass, scaleNorm.bind(scaleBindings), scalePlan, [params.offset]);
|
|
236
|
+
finalize.dispatch(pass, finalize.bind({ partials, P: params.binding }), finalizePlan, [params.offset]);
|
|
237
|
+
if (scaleApply !== null) {
|
|
238
|
+
scaleApply.dispatch(pass, scaleApply.bind(scaleBindings), scalePlan, [params.offset]);
|
|
239
|
+
}
|
|
240
|
+
pull.record(pass, core, { xNorm, rankOut, personalization: null, partials }, coefficients);
|
|
241
|
+
}
|
|
242
|
+
batch.endPass();
|
|
243
|
+
const headerRequest = batch.readback(partialsBuffer, 0, PR_PARTIAL.byteLength);
|
|
244
|
+
const scoresRequest = batch.readback(ring[(iterationsRun + k) % ring.length].buffer, 0, bytes);
|
|
245
|
+
// x(m - 1) on an alternating run: the other kind, still intact in the third ring slot
|
|
246
|
+
const previousRequest =
|
|
247
|
+
config.alternate === null
|
|
248
|
+
? null
|
|
249
|
+
: batch.readback(ring[(iterationsRun + k - 1) % ring.length].buffer, 0, bytes);
|
|
250
|
+
scope.flush();
|
|
251
|
+
const submitted = batch.submit();
|
|
252
|
+
const back = await submitted.readback; // the ONE mapAsync of the batch (G7 item 5)
|
|
253
|
+
iterationsRun += k;
|
|
254
|
+
ctx.assertReady();
|
|
255
|
+
if (config.signal?.aborted === true) {
|
|
256
|
+
throw aborted(config.label, submitted.id);
|
|
257
|
+
}
|
|
258
|
+
config.onProgress?.(iterationsRun, config.maxIterations);
|
|
259
|
+
const folded = PR_PARTIAL.read(new DataView(back), headerRequest.offset);
|
|
260
|
+
const firstConverged = folded.firstConverged as number;
|
|
261
|
+
const converged = firstConverged !== U32_MAX;
|
|
262
|
+
if (converged || iterationsRun >= config.maxIterations) {
|
|
263
|
+
if (converged && iterationsRun < config.maxIterations) {
|
|
264
|
+
config.onProgress?.(config.maxIterations, config.maxIterations);
|
|
265
|
+
}
|
|
266
|
+
return {
|
|
267
|
+
scores: new Float32Array(back, scoresRequest.offset, n).slice(),
|
|
268
|
+
previous: previousRequest === null ? null : new Float32Array(back, previousRequest.offset, n).slice(),
|
|
269
|
+
iterations: converged ? firstConverged : iterationsRun,
|
|
270
|
+
converged,
|
|
271
|
+
iterationsRun,
|
|
272
|
+
};
|
|
273
|
+
}
|
|
274
|
+
}
|
|
275
|
+
} finally {
|
|
276
|
+
scope.dispose();
|
|
277
|
+
}
|
|
278
|
+
}
|
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
/**
|
|
2
|
+
* The ReduceScope an algorithm driver hands the primitives (contract 3.11), built over a GpuContext: `scratch` comes
|
|
3
|
+
* from ONE Lease (released by `dispose()` in the algorithm's `finally`, spec 4.4) and `params` from a UniformRing
|
|
4
|
+
* slot, so a batch of k iterations writes k parameter blocks into one buffer and dispatches with dynamic offsets
|
|
5
|
+
* instead of allocating a uniform buffer per dispatch (spec 5.3). The src/ twin of the tests' testReduceScope
|
|
6
|
+
* (test/helpers/segmented-reduce.ts). `slots` sizes the ring for the LARGEST batch the caller records: reserve()
|
|
7
|
+
* wraps, so a ring too small silently reuses a slot another dispatch of the same batch still reads.
|
|
8
|
+
*/
|
|
9
|
+
|
|
10
|
+
import { type GpuContext } from "../context.js";
|
|
11
|
+
import { UniformRing } from "../kernel/uniform-ring.js";
|
|
12
|
+
import { type ReduceScope } from "../primitives/reduce.js";
|
|
13
|
+
|
|
14
|
+
/** A ReduceScope over a context plus the two lifecycle calls an algorithm makes: flush() before submit, dispose() in its finally. */
|
|
15
|
+
export interface AlgorithmScope extends ReduceScope {
|
|
16
|
+
/** queue.writeBuffer of the params slots written since the last flush (called before the batch is submitted). */
|
|
17
|
+
flush(): void;
|
|
18
|
+
/** Destroys the ring and releases every scratch buffer of the lease; idempotent. */
|
|
19
|
+
dispose(): void;
|
|
20
|
+
}
|
|
21
|
+
|
|
22
|
+
/**
|
|
23
|
+
* The scope of one algorithm call.
|
|
24
|
+
* @param ctx - the context
|
|
25
|
+
* @param label - the label prefix of the ring and every scratch buffer
|
|
26
|
+
* @param slots - the params slots of the ring: at least the parameter blocks of the largest batch recorded
|
|
27
|
+
* @returns the scope
|
|
28
|
+
*/
|
|
29
|
+
export function algorithmScope(ctx: GpuContext, label: string, slots: number): AlgorithmScope {
|
|
30
|
+
const lease = ctx.pool.lease();
|
|
31
|
+
const ring = new UniformRing(ctx.device, ctx.allocator, slots, `${label}/ring`);
|
|
32
|
+
return {
|
|
33
|
+
device: ctx.device,
|
|
34
|
+
caps: ctx.caps,
|
|
35
|
+
pipelines: ctx.pipelines,
|
|
36
|
+
pool: ctx.pool,
|
|
37
|
+
workgroupSize: ctx.workgroupSize,
|
|
38
|
+
scratch: (byteLength, scratchLabel) => lease.storage(byteLength, `${label}/${scratchLabel}`),
|
|
39
|
+
params(block, values) {
|
|
40
|
+
const slot = ring.reserve(1);
|
|
41
|
+
ring.write(slot, block, values);
|
|
42
|
+
return { binding: ring.binding(block), offset: ring.offsetOf(slot) };
|
|
43
|
+
},
|
|
44
|
+
flush: () => {
|
|
45
|
+
ring.flush();
|
|
46
|
+
},
|
|
47
|
+
dispose(): void {
|
|
48
|
+
ring.destroy();
|
|
49
|
+
lease.release();
|
|
50
|
+
},
|
|
51
|
+
};
|
|
52
|
+
}
|