wgblas 0.1.2 → 1.0.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.
Files changed (55) hide show
  1. package/README.md +3 -0
  2. package/dist/wgblas.browser.js +1078 -37
  3. package/index.d.mts +6 -0
  4. package/index.mjs +6 -0
  5. package/package.json +32 -1
  6. package/src/classes/GpuMatrix.d.mts +85 -0
  7. package/src/classes/GpuMatrix.mjs +91 -0
  8. package/src/classes/GpuVector.d.mts +11 -7
  9. package/src/classes/GpuVector.mjs +31 -5
  10. package/src/dasum/dasum.d.mts +48 -0
  11. package/src/dasum/dasum.mjs +121 -0
  12. package/src/init.mjs +3 -1
  13. package/src/isamax/isamax.mjs +66 -67
  14. package/src/random/random.d.mts +37 -0
  15. package/src/random/random.mjs +17 -0
  16. package/src/sasum/sasum.mjs +55 -51
  17. package/src/saxpy/saxpy.mjs +48 -35
  18. package/src/scopy/scopy.mjs +43 -32
  19. package/src/sdot/sdot.mjs +60 -55
  20. package/src/sgemv/sgemv.d.mts +118 -0
  21. package/src/sgemv/sgemv.mjs +141 -0
  22. package/src/shaders/browser-shaders.mjs +20 -0
  23. package/src/shaders/dasum.wgsl +98 -0
  24. package/src/shaders/f64add.wgsl +281 -0
  25. package/src/shaders/isamax.wgsl +32 -9
  26. package/src/shaders/reduction/sumF64.wgsl +49 -0
  27. package/src/shaders/sasum.wgsl +18 -4
  28. package/src/shaders/sdot.wgsl +18 -4
  29. package/src/shaders/sgemv_n.wgsl +75 -0
  30. package/src/shaders/sgemv_t.wgsl +65 -0
  31. package/src/shaders/snrm2.wgsl +22 -4
  32. package/src/shaders/ssymv.wgsl +69 -0
  33. package/src/shaders/strmv.wgsl +103 -0
  34. package/src/shaders/strsv_apply_inverse.wgsl +47 -0
  35. package/src/shaders/strsv_invert_block.wgsl +109 -0
  36. package/src/shaders/strsv_update.wgsl +75 -0
  37. package/src/snrm2/snrm2.mjs +56 -52
  38. package/src/srot/srot.mjs +57 -41
  39. package/src/srotm/srotm.mjs +54 -38
  40. package/src/sscal/sscal.mjs +43 -32
  41. package/src/sswap/sswap.mjs +49 -34
  42. package/src/ssymv/ssymv.d.mts +109 -0
  43. package/src/ssymv/ssymv.mjs +130 -0
  44. package/src/strmv/strmv.d.mts +109 -0
  45. package/src/strmv/strmv.mjs +132 -0
  46. package/src/strsv/strsv.d.mts +98 -0
  47. package/src/strsv/strsv.mjs +212 -0
  48. package/src/util/benchmark.mjs +1 -1
  49. package/src/util/bindgroup.mjs +14 -10
  50. package/src/util/buffer.mjs +7 -2
  51. package/src/util/compute.mjs +41 -15
  52. package/src/util/f64pack.mjs +152 -0
  53. package/src/util/pipeline.mjs +32 -17
  54. package/src/util/result.mjs +8 -4
  55. package/src/util/workgroup.mjs +10 -10
@@ -0,0 +1,212 @@
1
+ import {
2
+ uploadBuffer,
3
+ createParamsBuffer,
4
+ createStorageBuffer,
5
+ stageReadback,
6
+ destroyBuffers,
7
+ } from "../util/buffer.mjs";
8
+ import { createBindGroup } from "../util/bindgroup.mjs";
9
+ import { beginTimedEncoder, encodePass, submit } from "../util/compute.mjs";
10
+ import { extractResult } from "../util/result.mjs";
11
+ import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
12
+ import { getPipeline } from "../util/pipeline.mjs";
13
+ import { GpuVector } from "../classes/GpuVector.mjs";
14
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
15
+
16
+ // Blocked triangular solve, explicit-inversion variant: a pre-pass inverts
17
+ // every diagonal block in parallel (strsv_invert_block.wgsl), then each
18
+ // block is solved via a dense matvec against its inverse (strsv_apply_inverse.wgsl)
19
+ // instead of a barrier-per-row substitution. strsv_update.wgsl (propagating
20
+ // a solved block onto remaining rows) is unchanged.
21
+ const BLOCK_SIZE = 64;
22
+
23
+ // Packs one small uniform struct per block into a single shared buffer
24
+ // (at `blockIndex * stride`) instead of allocating numBlocks separate tiny
25
+ // buffers — measured CPU-side overhead (buffer/bind-group setup, not actual
26
+ // GPU compute) was 66-83% of total wall-clock time before this, dominated
27
+ // by the O(numBlocks) createBuffer/writeBuffer calls the old per-block loop
28
+ // made. `stride` must be a multiple of the device's uniform offset alignment
29
+ // (device.limits.minUniformBufferOffsetAlignment) so each block's slot is a
30
+ // valid fixed-offset binding.
31
+ function packBlockParams(numBlocks, stride, fieldsPerBlock) {
32
+ const data = new ArrayBuffer(numBlocks * stride);
33
+ const view = new DataView(data);
34
+ for (let blockIndex = 0; blockIndex < numBlocks; blockIndex++) {
35
+ const fields = fieldsPerBlock(blockIndex);
36
+ const base = blockIndex * stride;
37
+ fields.forEach((value, i) => view.setUint32(base + i * 4, value, true));
38
+ }
39
+ return data;
40
+ }
41
+
42
+ function createSharedParamsBuffer(device, data, label) {
43
+ const buffer = device.createBuffer({
44
+ label,
45
+ size: data.byteLength,
46
+ usage: GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST,
47
+ });
48
+ device.queue.writeBuffer(buffer, 0, data);
49
+ return buffer;
50
+ }
51
+
52
+ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx) {
53
+ const xIsGpu = x instanceof GpuVector;
54
+ const AIsGpu = A instanceof GpuMatrix;
55
+ const isLower = uplo === "lower";
56
+ const isNoTrans = trans === "no-transpose";
57
+ const isUnit = diag === "unit";
58
+
59
+ if (!(device instanceof GPUDevice))
60
+ throw new Error("device must be a GPUDevice.");
61
+ if (!isLower && uplo !== "upper")
62
+ throw new Error("uplo must be 'lower' or 'upper'.");
63
+ if (!isNoTrans && trans !== "transpose")
64
+ throw new Error("trans must be 'no-transpose' or 'transpose'.");
65
+ if (!isUnit && diag !== "non-unit")
66
+ throw new Error("diag must be 'unit' or 'non-unit'.");
67
+ if (!Number.isInteger(n) || !Number.isInteger(incx) || !Number.isInteger(lda))
68
+ throw new Error("n, incx, and lda must be integers.");
69
+ if (incx <= 0) throw new Error("incx must be positive.");
70
+ if (lda < n) throw new Error("lda must be >= n.");
71
+ if (!AIsGpu && !(A instanceof Float32Array))
72
+ throw new Error("A must be a Float32Array or GpuMatrix.");
73
+ if (!xIsGpu && !(x instanceof Float32Array))
74
+ throw new Error("x must be a Float32Array or GpuVector.");
75
+ if (xIsGpu && !AIsGpu)
76
+ throw new Error("A must be a GpuMatrix when x is a GpuVector.");
77
+ if (AIsGpu && lda !== A.lda)
78
+ throw new Error("lda must match A.lda when A is a GpuMatrix.");
79
+ if (AIsGpu && (A.rows < n || A.cols < n))
80
+ throw new Error("A is too small for the given n.");
81
+ if (n < 0) throw new Error("n must be non-negative.");
82
+ if (n === 0) return xIsGpu ? {} : { x };
83
+
84
+ if (!AIsGpu && A.length < (n - 1) * lda + n)
85
+ throw new Error(
86
+ "A does not have enough elements for the given n and lda.",
87
+ );
88
+ if (x.length < (n - 1) * incx + 1)
89
+ throw new Error(
90
+ "x does not have enough elements for the given n and incx.",
91
+ );
92
+
93
+ const invertPipeline = await getPipeline(device, "strsv_invert_block");
94
+ const applyPipeline = await getPipeline(device, "strsv_apply_inverse");
95
+ const updatePipeline = await getPipeline(device, "strsv_update");
96
+
97
+ // Same forward/backward pairing the shaders use.
98
+ const forward = isNoTrans === isLower;
99
+ const blockStarts = [];
100
+ for (let s = 0; s < n; s += BLOCK_SIZE) blockStarts.push(s);
101
+ if (!forward) blockStarts.reverse();
102
+ const numBlocks = blockStarts.length;
103
+
104
+ const maxWg = device.limits.maxComputeWorkgroupsPerDimension;
105
+ const stride = device.limits.minUniformBufferOffsetAlignment;
106
+
107
+ let ABuffer = null;
108
+ let xBuffer = null;
109
+ let AinvBuffer = null;
110
+ let applyParamsBuffer = null;
111
+ let updateParamsBuffer = null;
112
+ let invertParams = null;
113
+
114
+ try {
115
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "strsv-A", false);
116
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "strsv-x", true);
117
+ // One BLOCK_SIZE x BLOCK_SIZE dense region per block (row-major), even
118
+ // though only a triangular half is ever nonzero — see strsv_invert_block.wgsl.
119
+ AinvBuffer = createStorageBuffer(
120
+ numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4,
121
+ "strsv-Ainv",
122
+ );
123
+
124
+ // Every block's {blockStart, blockEnd} is fixed by its natural index
125
+ // regardless of traversal direction, so both packed buffers are indexed
126
+ // by blockIndex (0..numBlocks-1), not by loop position.
127
+ const applyData = packBlockParams(numBlocks, stride, (blockIndex) => {
128
+ const blockStart = blockIndex * BLOCK_SIZE;
129
+ const blockEnd = Math.min(blockStart + BLOCK_SIZE, n);
130
+ return [incx, blockIndex, blockStart, blockEnd];
131
+ });
132
+ applyParamsBuffer = createSharedParamsBuffer(device, applyData, "strsv-apply-params");
133
+
134
+ const updateData = packBlockParams(numBlocks, stride, (blockIndex) => {
135
+ const blockStart = blockIndex * BLOCK_SIZE;
136
+ const blockEnd = Math.min(blockStart + BLOCK_SIZE, n);
137
+ return [n, incx, lda, isNoTrans ? 0 : 1, isLower ? 0 : 1, blockStart, blockEnd];
138
+ });
139
+ updateParamsBuffer = createSharedParamsBuffer(device, updateData, "strsv-update-params");
140
+
141
+ const { commandEncoder, querySet } = beginTimedEncoder();
142
+
143
+ // Pre-pass: every block's inverse, fully parallel, one dispatch.
144
+ invertParams = createParamsBuffer(
145
+ [
146
+ { value: n, type: "u32" },
147
+ { value: lda, type: "u32" },
148
+ { value: isNoTrans ? 0 : 1, type: "u32" },
149
+ { value: isLower ? 0 : 1, type: "u32" },
150
+ { value: isUnit ? 1 : 0, type: "u32" },
151
+ ],
152
+ "strsv-invert-params",
153
+ );
154
+ const invertBindGroup = createBindGroup(invertPipeline.getBindGroupLayout(0), [
155
+ ABuffer, AinvBuffer, invertParams,
156
+ ]);
157
+ const invertDesc = querySet
158
+ ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
159
+ : undefined;
160
+ encodePass(commandEncoder, invertPipeline, invertBindGroup, { x: BLOCK_SIZE, y: numBlocks }, invertDesc);
161
+
162
+ for (let bi = 0; bi < blockStarts.length; bi++) {
163
+ const blockStart = blockStarts[bi];
164
+ const blockEnd = Math.min(blockStart + BLOCK_SIZE, n);
165
+ const blockIndex = blockStart / BLOCK_SIZE;
166
+ const isLastPass = bi === blockStarts.length - 1;
167
+ const paramsOffset = blockIndex * stride;
168
+
169
+ const applyBindGroup = createBindGroup(applyPipeline.getBindGroupLayout(0), [
170
+ AinvBuffer, xBuffer, { buffer: applyParamsBuffer, offset: paramsOffset, size: 16 },
171
+ ]);
172
+
173
+ const applyDesc = isLastPass && querySet
174
+ ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
175
+ : undefined;
176
+ encodePass(commandEncoder, applyPipeline, applyBindGroup, 1, applyDesc);
177
+
178
+ const remaining = forward ? n - blockEnd : blockStart;
179
+ if (remaining === 0) continue;
180
+
181
+ const updateBindGroup = createBindGroup(updatePipeline.getBindGroupLayout(0), [
182
+ ABuffer, xBuffer, { buffer: updateParamsBuffer, offset: paramsOffset, size: 32 },
183
+ ]);
184
+
185
+ const wgCount = Math.min(remaining, maxWg);
186
+ encodePass(commandEncoder, updatePipeline, updateBindGroup, wgCount);
187
+ }
188
+
189
+ const ts = resolveTimestamp(commandEncoder, querySet);
190
+ const readBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
191
+
192
+ submit(commandEncoder);
193
+
194
+ const gpuTimeMs = await extractTimestamp(ts);
195
+
196
+ if (xIsGpu) {
197
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
198
+ return {};
199
+ }
200
+
201
+ const result = await extractResult(readBuffer, Float32Array);
202
+ if (gpuTimeMs !== undefined) return { x: result, gpuTimeMs };
203
+ return { x: result };
204
+ } finally {
205
+ if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
206
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
207
+ if (AinvBuffer) destroyBuffers(AinvBuffer);
208
+ if (applyParamsBuffer) destroyBuffers(applyParamsBuffer);
209
+ if (updateParamsBuffer) destroyBuffers(updateParamsBuffer);
210
+ if (invertParams) destroyBuffers(invertParams);
211
+ }
212
+ }
@@ -99,5 +99,5 @@ export async function extractTimestamp(ts) {
99
99
  tsReadBuffer.destroy();
100
100
  resolveBuffer.destroy(); // never mapped — no unmap() needed
101
101
  querySet.destroy();
102
- return Number(timestamps[1] - timestamps[0]) / 1e6;
102
+ return Math.max(0, Number(timestamps[1] - timestamps[0])) / 1e6;
103
103
  }
@@ -2,21 +2,25 @@
2
2
  import { getDevice } from "../init.mjs";
3
3
 
4
4
  /**
5
- * Creates a `GPUBindGroup` by mapping each buffer to sequential binding indices (0, 1, 2 …).
6
- * `resultBuffer` is appended last so its binding index follows all input buffers.
7
- * The order of `buffers` must match the `@binding` indices declared in the shader.
5
+ * Creates a `GPUBindGroup` by mapping each buffer to sequential binding indices
6
+ * (startBinding, startBinding + 1, …). The order of `buffers` must match the `@binding`s
7
+ * declared in the shader.
8
8
  * @param {GPUBindGroupLayout} layout - from `pipeline.getBindGroupLayout(0)`; the layout is derived automatically from the shader because `loadShader` uses `layout: "auto"`
9
- * @param {GPUBuffer[]} buffers - input and intermediate buffers in binding order
10
- * @param {GPUBuffer|null} [resultBuffer=null] - appended as the final binding if provided
9
+ * @param {(GPUBuffer|{buffer: GPUBuffer, offset: number, size: number})[]} buffers - input and
10
+ * intermediate buffers in binding order; a plain `GPUBuffer` binds the whole buffer, or pass
11
+ * `{buffer, offset, size}` to bind a sub-range — e.g. many small per-call param structs packed
12
+ * into one shared buffer instead of allocating a separate tiny buffer per call
13
+ * @param {number} [startBinding=0] - binding index of `buffers[0]`; nonzero only when a shader's
14
+ * own bindings don't start at 0 — e.g. a module built by concatenating two `.wgsl` files where
15
+ * the second file's bindings continue after the first's own bindings (if any)
11
16
  * @returns {GPUBindGroup}
12
17
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createBindGroup GPUDevice.createBindGroup()}
13
18
  */
14
- export function createBindGroup(layout, buffers, resultBuffer = null) {
19
+ export function createBindGroup(layout, buffers, startBinding = 0) {
15
20
  const device = getDevice();
16
- const allBuffers = resultBuffer ? [...buffers, resultBuffer] : [...buffers];
17
- const entries = allBuffers.map((buffer, i) => ({
18
- binding: i,
19
- resource: { buffer },
21
+ const entries = buffers.map((buffer, i) => ({
22
+ binding: startBinding + i,
23
+ resource: buffer instanceof GPUBuffer ? { buffer } : buffer,
20
24
  }));
21
25
  return device.createBindGroup({ layout, entries });
22
26
  }
@@ -12,7 +12,11 @@ export function destroyBuffers(...buffers) {
12
12
 
13
13
  /**
14
14
  * Creates a GPU storage buffer and uploads `data` into it via mapped-at-creation.
15
- * @param {Float32Array} data
15
+ * The mapped view is constructed from `data`'s own typed-array constructor, so
16
+ * bits are copied as-is regardless of element type (e.g. a Uint32Array's raw
17
+ * bit patterns are preserved — critical for dasum's aux half, which must
18
+ * never pass through a Float32Array view and risk NaN-bit-pattern canonicalization).
19
+ * @param {Float32Array|Uint32Array|Int32Array} data
16
20
  * @param {string} [label] - debug label visible in browser DevTools GPU inspection
17
21
  * @param {boolean} [readback=false] - add `COPY_SRC` so the buffer can be copied to a readback buffer
18
22
  * @throws {Error} if `data.byteLength` exceeds the device's `maxStorageBufferBindingSize`
@@ -45,7 +49,8 @@ export function uploadBuffer(data, label = "blas-input", readback = false) {
45
49
  mappedAtCreation: true,
46
50
  });
47
51
 
48
- const mappedArray = new Float32Array(buffer.getMappedRange());
52
+ const ViewCtor = data.constructor;
53
+ const mappedArray = new ViewCtor(buffer.getMappedRange());
49
54
  mappedArray.set(data);
50
55
  buffer.unmap();
51
56
 
@@ -2,6 +2,9 @@
2
2
  import { getDevice } from "../init.mjs";
3
3
  import { beginTimestamp, resolveTimestamp } from "./benchmark.mjs";
4
4
 
5
+ // Anchors the pass encoder to its command encoder to prevent premature GC.
6
+ const _passEncoders = new WeakMap();
7
+
5
8
  /**
6
9
  * Finalises `commandEncoder` into a command buffer and submits it to the GPU queue.
7
10
  * @param {GPUCommandEncoder} commandEncoder
@@ -13,24 +16,33 @@ export function submit(commandEncoder) {
13
16
  }
14
17
 
15
18
  /**
16
- * Encodes and submits a single compute pass: sets the pipeline and bind group,
17
- * dispatches workgroups, and optionally wraps the pass in GPU timestamp queries.
18
- * @param {GPUComputePipeline} pipeline
19
- * @param {GPUBindGroup} bindGroup
20
- * @param {number | { x: number, y: number }} workgroups - workgroup count; number for 1D dispatch, `{x, y}` for 2D
21
- * @returns {{ commandEncoder: GPUCommandEncoder, ts: any }} encoded commands and timestamp handle
19
+ * Starts a new command encoder together with its GPU timestamp query (if
20
+ * benchmarking is enabled) — the pairing every routine needs before encoding
21
+ * its first pass, whether that's `runComputePass`'s single pass or a
22
+ * multi-pass routine (e.g. strsv's blocked solve) encoding several by hand.
23
+ * @returns {{ commandEncoder: GPUCommandEncoder, querySet: GPUQuerySet|null, passDescriptor: object|undefined }}
22
24
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createCommandEncoder GPUDevice.createCommandEncoder()}
23
- * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/beginComputePass GPUCommandEncoder.beginComputePass()}
24
- * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUComputePassEncoder/setPipeline GPUComputePassEncoder.setPipeline()}
25
- * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUComputePassEncoder/setBindGroup GPUComputePassEncoder.setBindGroup()}
26
- * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUComputePassEncoder/dispatchWorkgroups GPUComputePassEncoder.dispatchWorkgroups()}
27
25
  */
28
- export function runComputePass(pipeline, bindGroup, workgroups) {
26
+ export function beginTimedEncoder() {
29
27
  const device = getDevice();
30
-
31
28
  const { querySet, passDescriptor } = beginTimestamp();
32
-
33
29
  const commandEncoder = device.createCommandEncoder();
30
+ return { commandEncoder, querySet, passDescriptor };
31
+ }
32
+
33
+ /**
34
+ * Encodes one compute pass (set pipeline, set bind group, dispatch, end) onto an
35
+ * existing command encoder — the shared building block behind `runComputePass`,
36
+ * also used directly by routines that need several passes on one encoder (e.g.
37
+ * strsv's blocked solve, where each pass depends on the previous one completing).
38
+ * @param {GPUCommandEncoder} commandEncoder
39
+ * @param {GPUComputePipeline} pipeline
40
+ * @param {GPUBindGroup} bindGroup
41
+ * @param {number | { x: number, y: number }} workgroups - workgroup count; number for 1D dispatch, `{x, y}` for 2D
42
+ * @param {object} [passDescriptor] - passed to `beginComputePass`, e.g. for timestamp writes
43
+ * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/beginComputePass GPUCommandEncoder.beginComputePass()}
44
+ */
45
+ export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, passDescriptor) {
34
46
  const passEncoder = commandEncoder.beginComputePass(passDescriptor);
35
47
 
36
48
  passEncoder.setPipeline(pipeline);
@@ -44,9 +56,23 @@ export function runComputePass(pipeline, bindGroup, workgroups) {
44
56
 
45
57
  passEncoder.end();
46
58
 
47
- const ts = resolveTimestamp(commandEncoder, querySet);
59
+ _passEncoders.set(commandEncoder, passEncoder);
60
+ }
48
61
 
49
- commandEncoder._passEncoder = passEncoder; // anchor — GC'd passEncoder may crash native encoder
62
+ /**
63
+ * Encodes and submits a single compute pass: sets the pipeline and bind group,
64
+ * dispatches workgroups, and optionally wraps the pass in GPU timestamp queries.
65
+ * @param {GPUComputePipeline} pipeline
66
+ * @param {GPUBindGroup} bindGroup
67
+ * @param {number | { x: number, y: number }} workgroups - workgroup count; number for 1D dispatch, `{x, y}` for 2D
68
+ * @returns {{ commandEncoder: GPUCommandEncoder, ts: any }} encoded commands and timestamp handle
69
+ * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createCommandEncoder GPUDevice.createCommandEncoder()}
70
+ */
71
+ export function runComputePass(pipeline, bindGroup, workgroups) {
72
+ const { commandEncoder, querySet, passDescriptor } = beginTimedEncoder();
73
+ encodePass(commandEncoder, pipeline, bindGroup, workgroups, passDescriptor);
74
+
75
+ const ts = resolveTimestamp(commandEncoder, querySet);
50
76
 
51
77
  return { commandEncoder, ts };
52
78
  }
@@ -0,0 +1,152 @@
1
+ /** @module devdocs/utility-functions/f64pack */
2
+
3
+ // Scratch buffers reused across calls for bit-pattern <-> float reinterpretation.
4
+ // dv8 is read/written exclusively through explicit big-endian DataView calls
5
+ // (never via Float64Array, whose byte order follows host endianness) so hi/lo
6
+ // word extraction is deterministic regardless of host architecture.
7
+ const buf8 = new ArrayBuffer(8);
8
+ const dv8 = new DataView(buf8);
9
+
10
+ const buf4 = new ArrayBuffer(4);
11
+ const u32View = new Uint32Array(buf4);
12
+ const f32View = new Float32Array(buf4);
13
+
14
+ function u32ToF32(bits) {
15
+ u32View[0] = bits >>> 0;
16
+ return f32View[0];
17
+ }
18
+
19
+ function f32ToU32(f) {
20
+ f32View[0] = f;
21
+ return u32View[0];
22
+ }
23
+
24
+ // Shared bit-scramble between an f64's raw fields (sign, 11-bit exponent,
25
+ // mantissa split as a 20-bit hi word + 32-bit lo word) and a [main, aux]
26
+ // pair. f64 has exactly 3 more exponent bits and 29 more mantissa bits than
27
+ // f32 — 3 + 29 = 32, i.e. exactly one f32's worth of bits, all of which fit
28
+ // in `aux`:
29
+ // aux.sign + aux.exponent[7:6] (3 bits) = f64's low 3 exponent bits
30
+ // aux.exponent[5:0] (6 bits) = top 6 of f64's low 29 mantissa bits
31
+ // aux.mantissa (23 bits) = bottom 23 of f64's low 29 mantissa bits
32
+ //
33
+ // `aux` is deliberately kept as a raw u32 integer, never converted to an
34
+ // actual float32 value (unlike `main`, which genuinely is a float32 for
35
+ // reuse in float-typed math like abs()). Any combination of the source bits
36
+ // above can land on aux.exponent === 0xff — i.e. aux's bit pattern reads as
37
+ // a NaN or Infinity if it's ever treated as a real float — and this can
38
+ // happen not just for unusual inputs but for the RESULT of an ordinary
39
+ // addition (see f64add.wgsl's computeSum), so it can't be filtered out at
40
+ // the input boundary. A JS engine or GPU canonicalizes a NaN bit pattern
41
+ // (sets its mantissa's quiet bit) the moment it passes through an actual
42
+ // Float32Array or f32-typed GPU buffer/register as a VALUE — silently
43
+ // corrupting the payload. Since aux never needs float semantics anywhere
44
+ // (WGSL only ever bitcasts it back to u32 for decode()), storing/uploading/
45
+ // reading it exclusively as Uint32Array bits (see GpuVector, buffer.mjs)
46
+ // sidesteps the whole class of corruption rather than working around it.
47
+ function fieldsToPacked(sign, rawExp, mantissaHi, lo) {
48
+ const expMain = rawExp >>> 3; // top 8 bits -> main's exponent
49
+ const expExtra = rawExp & 0x7; // bottom 3 bits -> stashed in aux
50
+
51
+ const mantTop3 = lo >>> 29; // top 3 bits of the low mantissa word
52
+ const mantMain = (mantissaHi << 3) | mantTop3; // 20 + 3 = 23 bits -> main's mantissa
53
+ const mantExtra29 = lo & 0x1fffffff; // bottom 29 bits -> stashed in aux
54
+
55
+ const mainBits = (sign << 31) | (expMain << 23) | mantMain;
56
+
57
+ const auxSign = (expExtra >>> 2) & 0x1;
58
+ const auxExpTop2 = expExtra & 0x3;
59
+ const auxExpBot6 = mantExtra29 >>> 23; // top 6 of the 29
60
+ const auxMant23 = mantExtra29 & 0x7fffff; // bottom 23 of the 29
61
+ const auxExp8 = (auxExpTop2 << 6) | auxExpBot6;
62
+
63
+ const auxBits = ((auxSign << 31) | (auxExp8 << 23) | auxMant23) >>> 0;
64
+
65
+ return [u32ToF32(mainBits), auxBits];
66
+ }
67
+
68
+ function packedToFields(main, auxBits) {
69
+ const mainBits = f32ToU32(main);
70
+ auxBits = auxBits >>> 0;
71
+
72
+ const sign = mainBits >>> 31;
73
+ const expMain = (mainBits >>> 23) & 0xff;
74
+ const mantMain = mainBits & 0x7fffff; // 23 bits
75
+
76
+ const auxSign = auxBits >>> 31;
77
+ const auxExp8 = (auxBits >>> 23) & 0xff;
78
+ const auxMant23 = auxBits & 0x7fffff;
79
+
80
+ const expExtra = (auxSign << 2) | (auxExp8 >>> 6); // 3 bits
81
+ const mantExtra29 = ((auxExp8 & 0x3f) << 23) | auxMant23; // 29 bits
82
+
83
+ const rawExp = (expMain << 3) | expExtra; // 11 bits
84
+ const mantissaHi = mantMain >>> 3; // top 20 bits
85
+ const mantTop3 = mantMain & 0x7; // bottom 3 bits of mantMain
86
+
87
+ const lo = ((mantTop3 << 29) | mantExtra29) >>> 0;
88
+
89
+ return { sign, rawExp, mantissaHi, lo };
90
+ }
91
+
92
+ // main (unlike aux) genuinely is a float32 value — it gets reinterpreted with
93
+ // real float ops (e.g. abs() in dasum.wgsl), so it can't be transported as
94
+ // raw u32 the way aux is. That leaves a narrow residual gap: main.exponent is
95
+ // `rawExp >>> 3`, which hits 0xff (main's own NaN/Infinity pattern) whenever
96
+ // the f64's raw exponent is >= 2040 — i.e. |value| >= ~1.4e306, plus actual
97
+ // ±Infinity/NaN (rawExp === 2047). Any real Float32Array/GPU round-trip of
98
+ // such a main would get silently NaN-canonicalized, the same corruption class
99
+ // aux was fixed for. Since main can't take that fix, reject the range instead
100
+ // of silently corrupting it.
101
+ const MIN_UNSAFE_RAW_EXP = 2040;
102
+
103
+ /**
104
+ * Packs a double into a [main, aux] pair for storage/transfer where only f32
105
+ * is available (e.g. WGSL, which has no f64 type). This is a raw bit
106
+ * repacking, not a value-preserving numeric split — neither half is a
107
+ * meaningful float on its own; only `unpackF64(main, aux)` reconstructs the
108
+ * original value.
109
+ * @param {number} value - finite double with |value| < ~1.4e306 (see
110
+ * MIN_UNSAFE_RAW_EXP above); larger magnitudes, ±Infinity, and NaN would
111
+ * produce a `main` whose bit pattern is itself NaN/Infinity-shaped, which
112
+ * silently corrupts on any real float32 round-trip, so they're rejected.
113
+ * @returns {[number, number]} `[main, aux]` — main is a float32 value, aux is
114
+ * a raw uint32 bit pattern (0 to 2^32-1). Store/transport aux exclusively
115
+ * via Uint32Array/`array<u32>` — never as an actual float value — see the
116
+ * comment above `fieldsToPacked`.
117
+ */
118
+ export function packF64(value) {
119
+ dv8.setFloat64(0, value, false);
120
+ const hi = dv8.getUint32(0, false); // sign(1) + exponent(11) + mantissa_hi(20)
121
+ const lo = dv8.getUint32(4, false); // mantissa_lo(32)
122
+
123
+ const sign = hi >>> 31;
124
+ const rawExp = (hi >>> 20) & 0x7ff; // 11 bits
125
+ const mantissaHi = hi & 0xfffff; // 20 bits
126
+
127
+ if (rawExp >= MIN_UNSAFE_RAW_EXP) {
128
+ throw new RangeError(
129
+ `packF64: |${value}| is too large to pack safely (must be finite with ` +
130
+ `magnitude below ~1.4e306); main's bit pattern would itself be NaN/` +
131
+ `Infinity-shaped and get silently corrupted by any real float32 round-trip`,
132
+ );
133
+ }
134
+
135
+ return fieldsToPacked(sign, rawExp, mantissaHi, lo);
136
+ }
137
+
138
+ /**
139
+ * Reconstructs the original double from the [main, aux] pair produced by
140
+ * `packF64`. Bit-exact inverse.
141
+ * @param {number} main - float32 value
142
+ * @param {number} aux - raw uint32 bit pattern (as read from a Uint32Array/`array<u32>`)
143
+ * @returns {number}
144
+ */
145
+ export function unpackF64(main, aux) {
146
+ const { sign, rawExp, mantissaHi, lo } = packedToFields(main, aux);
147
+ const hi = ((sign << 31) | (rawExp << 20) | mantissaHi) >>> 0;
148
+
149
+ dv8.setUint32(0, hi, false);
150
+ dv8.setUint32(4, lo, false);
151
+ return dv8.getFloat64(0, false);
152
+ }
@@ -8,18 +8,25 @@ const _pipelines = new WeakMap();
8
8
  * Returns a cached `GPUComputePipeline` for the given shader, compiling it on first use.
9
9
  * Pipelines are cached per device so reinitialization (new device) always recompiles.
10
10
  * @param {GPUDevice} device
11
- * @param {string} shaderName - filename without `.wgsl` extension (e.g. `"sscal"`)
11
+ * @param {string|string[]} shaderName - filename without `.wgsl` extension (e.g. `"sscal"`),
12
+ * or an array of filenames whose source gets concatenated into one module — WGSL has no
13
+ * `#include`, so this is how one shader's helper functions (e.g. f64add.wgsl's decode/
14
+ * encode/computeSum) get reused from another (e.g. dasum.wgsl) without duplicating them.
15
+ * @param {string} [entryPoint="main"] - which `@compute` function to run; only needed when
16
+ * concatenating shaders whose entry points aren't both named "main"
12
17
  * @returns {Promise<GPUComputePipeline>}
13
18
  */
14
- export async function getPipeline(device, shaderName) {
19
+ export async function getPipeline(device, shaderName, entryPoint = "main") {
15
20
  if (!_pipelines.has(device)) {
16
21
  _pipelines.set(device, new Map());
17
22
  }
18
23
  const byName = _pipelines.get(device);
19
- if (!byName.has(shaderName)) {
20
- byName.set(shaderName, await loadShader(shaderName));
24
+ const names = Array.isArray(shaderName) ? shaderName : [shaderName];
25
+ const key = `${names.join("+")}::${entryPoint}`;
26
+ if (!byName.has(key)) {
27
+ byName.set(key, await loadShader(names, entryPoint));
21
28
  }
22
- return byName.get(shaderName);
29
+ return byName.get(key);
23
30
  }
24
31
 
25
32
  /**
@@ -29,8 +36,8 @@ export async function getPipeline(device, shaderName) {
29
36
  * @returns {Promise<string>}
30
37
  */
31
38
  async function loadCode(shaderName) {
32
- // typeof guards against ReferenceError in Node.js — accessing window directly throws if it doesn't exist.
33
- if (typeof window !== "undefined") {
39
+ // Check for Node.js explicitly — `window` is undefined in Web Workers too, so it's not a reliable signal.
40
+ if (typeof process === "undefined" || !process.versions?.node) {
34
41
  const { shaderSources } = await import("../shaders/browser-shaders.mjs");
35
42
  const src = shaderSources[shaderName];
36
43
  if (!src) throw new Error(`Shader "${shaderName}" not found in browser bundle.`);
@@ -45,35 +52,43 @@ async function loadCode(shaderName) {
45
52
  }
46
53
 
47
54
  /**
48
- * Compiles a WGSL shader into a `GPUComputePipeline`. Throws with line-level detail if compilation fails, rather than surfacing a raw GPU error.
49
- * Uses `layout: "auto"` so the pipeline derives its bind group layout from the shader —
50
- * no manual layout definition needed.
51
- * @param {string} shaderName
55
+ * Compiles one or more WGSL shaders (concatenated, in order, into a single module) into a
56
+ * `GPUComputePipeline`. Throws with line-level detail if compilation fails, rather than
57
+ * surfacing a raw GPU error. Uses `layout: "auto"` so the pipeline derives its bind group
58
+ * layout from the shader — no manual layout definition needed.
59
+ * @param {string[]} shaderNames - filenames without `.wgsl`, concatenated in array order
60
+ * @param {string} [entryPoint="main"] - which `@compute` function in the combined module to run
52
61
  * @returns {Promise<GPUComputePipeline>}
53
62
  * @throws {Error} if the shader has compilation errors
54
63
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createShaderModule GPUDevice.createShaderModule()}
55
64
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUShaderModule/getCompilationInfo GPUShaderModule.getCompilationInfo()}
56
65
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createComputePipeline GPUDevice.createComputePipeline()}
57
66
  */
58
- export async function loadShader(shaderName) {
67
+ export async function loadShader(shaderNames, entryPoint = "main") {
59
68
  const device = getDevice();
60
- const code = await loadCode(shaderName);
69
+ const label = shaderNames.join("+");
70
+ const code = (await Promise.all(shaderNames.map(loadCode))).join("\n");
61
71
 
62
- const shaderModule = device.createShaderModule({ label: shaderName, code });
72
+ const shaderModule = device.createShaderModule({ label, code });
63
73
 
64
74
  const info = await shaderModule.getCompilationInfo();
65
75
  // GPUCompilationMessage: https://developer.mozilla.org/en-US/docs/Web/API/GPUCompilationMessage
66
76
  const errors = info.messages.filter((m) => m.type === "error");
67
77
  if (errors.length > 0) {
68
78
  throw new Error(
69
- `Shader "${shaderName}" compilation failed:\n${errors.map((m) => ` line ${m.lineNum}: ${m.message}`).join("\n")}`,
79
+ `Shader "${label}" compilation failed:\n${errors.map((m) => ` line ${m.lineNum}: ${m.message}`).join("\n")}`,
70
80
  );
71
81
  }
72
82
 
83
+ // Omit entryPoint when it's the default "main" rather than always passing it explicitly —
84
+ // this project's WebGPU backend is unstable (intermittent multi-minute hangs and wrong
85
+ // results, confirmed by bisection) when entryPoint is set explicitly, even to the shader's
86
+ // only/correct entry point. Auto-detecting the single entry point is the stable path.
87
+ const compute = entryPoint === "main" ? { module: shaderModule } : { module: shaderModule, entryPoint };
73
88
  const pipeline = device.createComputePipeline({
74
- label: shaderName,
89
+ label,
75
90
  layout: "auto",
76
- compute: { module: shaderModule },
91
+ compute,
77
92
  });
78
93
 
79
94
  pipeline._shaderModule = shaderModule; // anchor — GC'd shaderModule crashes native pipeline
@@ -12,8 +12,12 @@
12
12
  * @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUBuffer/unmap GPUBuffer.unmap()}
13
13
  */
14
14
  export async function extractResult(readBuffer, readbackType = Float32Array) {
15
- await readBuffer.mapAsync(GPUMapMode.READ);
16
- const result = new readbackType(readBuffer.getMappedRange().slice());
17
- readBuffer.unmap();
18
- return result;
15
+ try {
16
+ await readBuffer.mapAsync(GPUMapMode.READ);
17
+ const result = new readbackType(readBuffer.getMappedRange().slice());
18
+ readBuffer.unmap();
19
+ return result;
20
+ } finally {
21
+ readBuffer.destroy();
22
+ }
19
23
  }