wgblas 2.1.0 → 2.2.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 +2 -0
- package/dist/wgblas.browser.js +1007 -43
- package/index.d.mts +26 -53
- package/index.mjs +11 -0
- package/package.json +132 -63
- package/src/classes/Complex32.d.mts +43 -0
- package/src/classes/Complex32.mjs +82 -0
- package/src/classes/Complex64.d.mts +44 -0
- package/src/classes/Complex64.mjs +76 -0
- package/src/classes/GpuMatrix.d.mts +41 -41
- package/src/classes/GpuMatrix.mjs +112 -10
- package/src/classes/GpuVector.d.mts +36 -40
- package/src/classes/GpuVector.mjs +39 -2
- package/src/cscal/cscal.d.mts +47 -0
- package/src/cscal/cscal.mjs +98 -0
- package/src/dasum/dasum.mjs +31 -15
- package/src/daxpy/daxpy.d.mts +56 -0
- package/src/daxpy/daxpy.mjs +150 -0
- package/src/dcopy/dcopy.d.mts +52 -0
- package/src/dcopy/dcopy.mjs +140 -0
- package/src/ddot/ddot.d.mts +62 -0
- package/src/ddot/ddot.mjs +184 -0
- package/src/dnrm2/dnrm2.d.mts +50 -0
- package/src/dnrm2/dnrm2.mjs +189 -0
- package/src/drot/drot.d.mts +67 -0
- package/src/drot/drot.mjs +170 -0
- package/src/drotm/drotm.d.mts +67 -0
- package/src/drotm/drotm.mjs +171 -0
- package/src/dscal/dscal.d.mts +52 -0
- package/src/dscal/dscal.mjs +119 -0
- package/src/dswap/dswap.d.mts +57 -0
- package/src/dswap/dswap.mjs +155 -0
- package/src/idamax/idamax.mjs +49 -19
- package/src/init.mjs +6 -3
- package/src/isamax/isamax.mjs +17 -14
- package/src/random/random.d.mts +37 -40
- package/src/random/random.mjs +39 -7
- package/src/sasum/sasum.mjs +13 -11
- package/src/saxpy/saxpy.mjs +9 -8
- package/src/scopy/scopy.mjs +8 -6
- package/src/sdot/sdot.mjs +13 -11
- package/src/sgemm/sgemm.d.mts +2 -2
- package/src/sgemm/sgemm.mjs +91 -35
- package/src/sgemmtr/sgemmtr.d.mts +3 -2
- package/src/sgemmtr/sgemmtr.mjs +92 -35
- package/src/sgemv/sgemv.d.mts +2 -2
- package/src/sgemv/sgemv.mjs +41 -25
- package/src/sger/sger.d.mts +2 -2
- package/src/sger/sger.mjs +38 -16
- package/src/shaders/__test_pipeline_a.wgsl +3 -0
- package/src/shaders/__test_pipeline_b.wgsl +2 -0
- package/src/shaders/cscal.wgsl +33 -0
- package/src/shaders/daxpy.wgsl +66 -0
- package/src/shaders/dcopy.wgsl +34 -0
- package/src/shaders/ddot.wgsl +106 -0
- package/src/shaders/dnrm2.wgsl +167 -0
- package/src/shaders/drot.wgsl +81 -0
- package/src/shaders/drotm.wgsl +99 -0
- package/src/shaders/dscal.wgsl +60 -0
- package/src/shaders/dswap.wgsl +38 -0
- package/src/shaders/f64/utils/add.wgsl +6 -0
- package/src/shaders/f64/utils/divide.wgsl +45 -0
- package/src/shaders/f64/utils/multiply.wgsl +19 -10
- package/src/shaders/f64/utils/sqrt.wgsl +45 -0
- package/src/shaders/index.mjs +69 -0
- package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
- package/src/snrm2/snrm2.mjs +20 -14
- package/src/srot/srot.mjs +10 -7
- package/src/srotm/srotm.mjs +9 -12
- package/src/sscal/sscal.mjs +7 -7
- package/src/sswap/sswap.mjs +14 -8
- package/src/ssymm/ssymm.d.mts +5 -4
- package/src/ssymm/ssymm.mjs +142 -55
- package/src/ssymv/ssymv.d.mts +2 -2
- package/src/ssymv/ssymv.mjs +42 -23
- package/src/ssyr/ssyr.d.mts +2 -2
- package/src/ssyr/ssyr.mjs +34 -15
- package/src/ssyr2/ssyr2.d.mts +2 -2
- package/src/ssyr2/ssyr2.mjs +43 -18
- package/src/ssyr2k/ssyr2k.d.mts +3 -2
- package/src/ssyr2k/ssyr2k.mjs +132 -55
- package/src/ssyrk/ssyrk.d.mts +3 -2
- package/src/ssyrk/ssyrk.mjs +84 -33
- package/src/strmm/strmm.d.mts +5 -4
- package/src/strmm/strmm.mjs +153 -54
- package/src/strmv/strmv.d.mts +2 -2
- package/src/strmv/strmv.mjs +37 -17
- package/src/strsm/strsm.d.mts +6 -4
- package/src/strsm/strsm.mjs +418 -172
- package/src/strsv/strsv.d.mts +5 -3
- package/src/strsv/strsv.mjs +82 -31
- package/src/util/benchmark.mjs +5 -3
- package/src/util/buffer.mjs +33 -12
- package/src/util/complex.mjs +87 -0
- package/src/util/compute.mjs +14 -8
- package/src/util/device.mjs +18 -3
- package/src/util/pipeline.mjs +40 -5
- package/src/util/workgroup.mjs +23 -6
- package/src/shaders/f64add.wgsl +0 -281
package/src/strsv/strsv.d.mts
CHANGED
|
@@ -2,8 +2,9 @@ import { GpuVector } from "../classes/GpuVector.mjs";
|
|
|
2
2
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
3
3
|
|
|
4
4
|
/**
|
|
5
|
-
* Solves the triangular system
|
|
6
|
-
*
|
|
5
|
+
* Solves the triangular system for $x$, in place — x holds b on input, the
|
|
6
|
+
* solution on output:
|
|
7
|
+
* $$\mathrm{op}(A) x = b$$
|
|
7
8
|
*
|
|
8
9
|
* A is an n×n triangular matrix stored in row-major order. Only the triangle
|
|
9
10
|
* specified by `uplo` is referenced; the other triangle is not accessed.
|
|
@@ -42,7 +43,8 @@ export declare function strsv(
|
|
|
42
43
|
): Promise<{ x: Float32Array; gpuTimeMs?: number }>;
|
|
43
44
|
|
|
44
45
|
/**
|
|
45
|
-
* Solves the triangular system
|
|
46
|
+
* Solves the triangular system for $x$, in place:
|
|
47
|
+
* $$\mathrm{op}(A) x = b$$
|
|
46
48
|
*
|
|
47
49
|
* x is kept resident on the GPU (mutated in place). A must be a GpuMatrix;
|
|
48
50
|
* its own `layout` (set at `GpuMatrix.from` time) determines the operation —
|
package/src/strsv/strsv.mjs
CHANGED
|
@@ -13,7 +13,7 @@ import { getPipeline } from "../util/pipeline.mjs";
|
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
14
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
15
15
|
import { BLOCK_SIZE } from "../util/constants.mjs";
|
|
16
|
-
import { requireSameDevice } from "../util/device.mjs";
|
|
16
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
17
17
|
|
|
18
18
|
// Blocked triangular solve via explicit block inversion (invert/apply/update passes) instead of barrier-per-row substitution.
|
|
19
19
|
|
|
@@ -39,13 +39,23 @@ function createSharedParamsBuffer(device, data, label) {
|
|
|
39
39
|
return buffer;
|
|
40
40
|
}
|
|
41
41
|
|
|
42
|
-
export async function strsv(
|
|
42
|
+
export async function strsv(
|
|
43
|
+
device,
|
|
44
|
+
uplo,
|
|
45
|
+
trans,
|
|
46
|
+
diag,
|
|
47
|
+
n,
|
|
48
|
+
A,
|
|
49
|
+
lda,
|
|
50
|
+
x,
|
|
51
|
+
incx,
|
|
52
|
+
layout = "row-major",
|
|
53
|
+
) {
|
|
43
54
|
const xIsGpu = x instanceof GpuVector;
|
|
44
55
|
const AIsGpu = A instanceof GpuMatrix;
|
|
45
56
|
const isUnit = diag === "unit";
|
|
46
57
|
|
|
47
|
-
|
|
48
|
-
throw new Error("device must be a GPUDevice.");
|
|
58
|
+
requireGpuDevice(device);
|
|
49
59
|
requireSameDevice(device, "strsv", { A, x });
|
|
50
60
|
if (uplo !== "lower" && uplo !== "upper")
|
|
51
61
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
@@ -77,9 +87,7 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
77
87
|
if (n === 0) return xIsGpu ? {} : { x };
|
|
78
88
|
|
|
79
89
|
if (!AIsGpu && A.length < (n - 1) * lda + n)
|
|
80
|
-
throw new Error(
|
|
81
|
-
"A does not have enough elements for the given n and lda.",
|
|
82
|
-
);
|
|
90
|
+
throw new Error("A does not have enough elements for the given n and lda.");
|
|
83
91
|
if (x.length < (n - 1) * incx + 1)
|
|
84
92
|
throw new Error(
|
|
85
93
|
"x does not have enough elements for the given n and incx.",
|
|
@@ -89,7 +97,9 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
89
97
|
const effLayout = AIsGpu ? A.layout : layout;
|
|
90
98
|
const isColMajor = effLayout === "column-major";
|
|
91
99
|
const isLower = isColMajor ? uplo === "upper" : uplo === "lower";
|
|
92
|
-
const isNoTrans = isColMajor
|
|
100
|
+
const isNoTrans = isColMajor
|
|
101
|
+
? trans === "transpose"
|
|
102
|
+
: trans === "no-transpose";
|
|
93
103
|
|
|
94
104
|
const invertPipeline = await getPipeline(device, "strsv_invert_block");
|
|
95
105
|
const applyPipeline = await getPipeline(device, "strsv_apply_inverse");
|
|
@@ -117,7 +127,8 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
117
127
|
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "strsv-x", true);
|
|
118
128
|
// One BLOCK_SIZE x BLOCK_SIZE dense region per block (row-major), even
|
|
119
129
|
// though only a triangular half is ever nonzero — see strsv_invert_block.wgsl.
|
|
120
|
-
AinvBuffer = createStorageBuffer(
|
|
130
|
+
AinvBuffer = createStorageBuffer(
|
|
131
|
+
device,
|
|
121
132
|
numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4,
|
|
122
133
|
"strsv-Ainv",
|
|
123
134
|
);
|
|
@@ -130,35 +141,60 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
130
141
|
const blockEnd = Math.min(blockStart + BLOCK_SIZE, n);
|
|
131
142
|
return [incx, blockIndex, blockStart, blockEnd];
|
|
132
143
|
});
|
|
133
|
-
applyParamsBuffer = createSharedParamsBuffer(
|
|
144
|
+
applyParamsBuffer = createSharedParamsBuffer(
|
|
145
|
+
device,
|
|
146
|
+
applyData,
|
|
147
|
+
"strsv-apply-params",
|
|
148
|
+
);
|
|
134
149
|
|
|
135
150
|
const updateData = packBlockParams(numBlocks, stride, (blockIndex) => {
|
|
136
151
|
const blockStart = blockIndex * BLOCK_SIZE;
|
|
137
152
|
const blockEnd = Math.min(blockStart + BLOCK_SIZE, n);
|
|
138
|
-
return [
|
|
153
|
+
return [
|
|
154
|
+
n,
|
|
155
|
+
incx,
|
|
156
|
+
lda,
|
|
157
|
+
isNoTrans ? 0 : 1,
|
|
158
|
+
isLower ? 0 : 1,
|
|
159
|
+
blockStart,
|
|
160
|
+
blockEnd,
|
|
161
|
+
];
|
|
139
162
|
});
|
|
140
|
-
updateParamsBuffer = createSharedParamsBuffer(
|
|
163
|
+
updateParamsBuffer = createSharedParamsBuffer(
|
|
164
|
+
device,
|
|
165
|
+
updateData,
|
|
166
|
+
"strsv-update-params",
|
|
167
|
+
);
|
|
141
168
|
|
|
142
169
|
const { commandEncoder, querySet } = beginTimedEncoder(device);
|
|
143
170
|
|
|
144
171
|
// Pre-pass: every block's inverse, fully parallel, one dispatch.
|
|
145
|
-
invertParams = createParamsBuffer(
|
|
172
|
+
invertParams = createParamsBuffer(
|
|
173
|
+
device,
|
|
146
174
|
[
|
|
147
|
-
{ value: n,
|
|
148
|
-
{ value: lda,
|
|
175
|
+
{ value: n, type: "u32" },
|
|
176
|
+
{ value: lda, type: "u32" },
|
|
149
177
|
{ value: isNoTrans ? 0 : 1, type: "u32" },
|
|
150
|
-
{ value: isLower ? 0 : 1,
|
|
151
|
-
{ value: isUnit ? 1 : 0,
|
|
178
|
+
{ value: isLower ? 0 : 1, type: "u32" },
|
|
179
|
+
{ value: isUnit ? 1 : 0, type: "u32" },
|
|
152
180
|
],
|
|
153
181
|
"strsv-invert-params",
|
|
154
182
|
);
|
|
155
|
-
const invertBindGroup = createBindGroup(
|
|
156
|
-
|
|
157
|
-
|
|
183
|
+
const invertBindGroup = createBindGroup(
|
|
184
|
+
device,
|
|
185
|
+
invertPipeline.getBindGroupLayout(0),
|
|
186
|
+
[ABuffer, AinvBuffer, invertParams],
|
|
187
|
+
);
|
|
158
188
|
const invertDesc = querySet
|
|
159
189
|
? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
|
|
160
190
|
: undefined;
|
|
161
|
-
encodePass(
|
|
191
|
+
encodePass(
|
|
192
|
+
commandEncoder,
|
|
193
|
+
invertPipeline,
|
|
194
|
+
invertBindGroup,
|
|
195
|
+
{ x: BLOCK_SIZE, y: numBlocks },
|
|
196
|
+
invertDesc,
|
|
197
|
+
);
|
|
162
198
|
|
|
163
199
|
for (let bi = 0; bi < blockStarts.length; bi++) {
|
|
164
200
|
const blockStart = blockStarts[bi];
|
|
@@ -167,28 +203,43 @@ export async function strsv(device, uplo, trans, diag, n, A, lda, x, incx, layou
|
|
|
167
203
|
const isLastPass = bi === blockStarts.length - 1;
|
|
168
204
|
const paramsOffset = blockIndex * stride;
|
|
169
205
|
|
|
170
|
-
const applyBindGroup = createBindGroup(
|
|
171
|
-
|
|
172
|
-
|
|
206
|
+
const applyBindGroup = createBindGroup(
|
|
207
|
+
device,
|
|
208
|
+
applyPipeline.getBindGroupLayout(0),
|
|
209
|
+
[
|
|
210
|
+
AinvBuffer,
|
|
211
|
+
xBuffer,
|
|
212
|
+
{ buffer: applyParamsBuffer, offset: paramsOffset, size: 16 },
|
|
213
|
+
],
|
|
214
|
+
);
|
|
173
215
|
|
|
174
|
-
const applyDesc =
|
|
175
|
-
|
|
176
|
-
|
|
216
|
+
const applyDesc =
|
|
217
|
+
isLastPass && querySet
|
|
218
|
+
? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
|
|
219
|
+
: undefined;
|
|
177
220
|
encodePass(commandEncoder, applyPipeline, applyBindGroup, 1, applyDesc);
|
|
178
221
|
|
|
179
222
|
const remaining = forward ? n - blockEnd : blockStart;
|
|
180
223
|
if (remaining === 0) continue;
|
|
181
224
|
|
|
182
|
-
const updateBindGroup = createBindGroup(
|
|
183
|
-
|
|
184
|
-
|
|
225
|
+
const updateBindGroup = createBindGroup(
|
|
226
|
+
device,
|
|
227
|
+
updatePipeline.getBindGroupLayout(0),
|
|
228
|
+
[
|
|
229
|
+
ABuffer,
|
|
230
|
+
xBuffer,
|
|
231
|
+
{ buffer: updateParamsBuffer, offset: paramsOffset, size: 32 },
|
|
232
|
+
],
|
|
233
|
+
);
|
|
185
234
|
|
|
186
235
|
const wgCount = Math.min(remaining, maxWg);
|
|
187
236
|
encodePass(commandEncoder, updatePipeline, updateBindGroup, wgCount);
|
|
188
237
|
}
|
|
189
238
|
|
|
190
239
|
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
|
191
|
-
const readBuffer = xIsGpu
|
|
240
|
+
const readBuffer = xIsGpu
|
|
241
|
+
? null
|
|
242
|
+
: stageReadback(device, commandEncoder, xBuffer);
|
|
192
243
|
|
|
193
244
|
submit(device, commandEncoder);
|
|
194
245
|
|
package/src/util/benchmark.mjs
CHANGED
|
@@ -71,9 +71,11 @@ export function resolveTimestamp(device, commandEncoder, querySet) {
|
|
|
71
71
|
usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.MAP_READ,
|
|
72
72
|
});
|
|
73
73
|
commandEncoder.copyBufferToBuffer(
|
|
74
|
-
resolveBuffer,
|
|
75
|
-
|
|
76
|
-
|
|
74
|
+
resolveBuffer,
|
|
75
|
+
0, // src, srcOffset
|
|
76
|
+
tsReadBuffer,
|
|
77
|
+
0, // dst, dstOffset
|
|
78
|
+
16, // full 16 bytes (both timestamps)
|
|
77
79
|
);
|
|
78
80
|
// resolveBuffer is returned to prevent GC — the copy command is only encoded here, not yet executed.
|
|
79
81
|
return { tsReadBuffer, resolveBuffer, querySet };
|
package/src/util/buffer.mjs
CHANGED
|
@@ -29,7 +29,7 @@ function requireStorageSize(device, byteSize, label) {
|
|
|
29
29
|
if (byteSize > maxSize) {
|
|
30
30
|
throw new Error(
|
|
31
31
|
`Buffer "${label}" needs ${byteSize} bytes, exceeding this device's ` +
|
|
32
|
-
|
|
32
|
+
`maxStorageBufferBindingSize (${maxSize} bytes). The operands are too large for this device.`,
|
|
33
33
|
);
|
|
34
34
|
}
|
|
35
35
|
}
|
|
@@ -51,7 +51,12 @@ function requireStorageSize(device, byteSize, label) {
|
|
|
51
51
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUBuffer/unmap GPUBuffer.unmap()}
|
|
52
52
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUSupportedLimits GPUSupportedLimits} (`maxStorageBufferBindingSize`)
|
|
53
53
|
*/
|
|
54
|
-
export function uploadBuffer(
|
|
54
|
+
export function uploadBuffer(
|
|
55
|
+
device,
|
|
56
|
+
data,
|
|
57
|
+
label = "blas-input",
|
|
58
|
+
readback = false,
|
|
59
|
+
) {
|
|
55
60
|
const byteSize = data.byteLength;
|
|
56
61
|
requireStorageSize(device, byteSize, label);
|
|
57
62
|
|
|
@@ -86,7 +91,12 @@ export function uploadBuffer(device, data, label = "blas-input", readback = fals
|
|
|
86
91
|
* @throws {Error} if `size` exceeds the device's `maxStorageBufferBindingSize`
|
|
87
92
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createBuffer GPUDevice.createBuffer()}
|
|
88
93
|
*/
|
|
89
|
-
export function createStorageBuffer(
|
|
94
|
+
export function createStorageBuffer(
|
|
95
|
+
device,
|
|
96
|
+
size,
|
|
97
|
+
label = "blas-storage",
|
|
98
|
+
extraUsage = 0,
|
|
99
|
+
) {
|
|
90
100
|
requireStorageSize(device, size, label);
|
|
91
101
|
return device.createBuffer({
|
|
92
102
|
label,
|
|
@@ -125,7 +135,6 @@ export function createResultBuffer(device, size, label = "blas-result") {
|
|
|
125
135
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/copyBufferToBuffer GPUCommandEncoder.copyBufferToBuffer()}
|
|
126
136
|
*/
|
|
127
137
|
export function stageReadback(device, commandEncoder, sourceBuffer) {
|
|
128
|
-
|
|
129
138
|
// COPY_DST: receives the copyBufferToBuffer transfer; MAP_READ: lets the CPU map and read it back.
|
|
130
139
|
const readBuffer = device.createBuffer({
|
|
131
140
|
label: "blas-readback",
|
|
@@ -134,9 +143,11 @@ export function stageReadback(device, commandEncoder, sourceBuffer) {
|
|
|
134
143
|
});
|
|
135
144
|
|
|
136
145
|
commandEncoder.copyBufferToBuffer(
|
|
137
|
-
sourceBuffer,
|
|
138
|
-
|
|
139
|
-
|
|
146
|
+
sourceBuffer,
|
|
147
|
+
0, // src, srcOffset
|
|
148
|
+
readBuffer,
|
|
149
|
+
0, // dst, dstOffset
|
|
150
|
+
sourceBuffer.size, // full copy, no partial reads
|
|
140
151
|
);
|
|
141
152
|
|
|
142
153
|
return readBuffer;
|
|
@@ -179,10 +190,17 @@ function vec4FallbackBuffer(device) {
|
|
|
179
190
|
export function vec4ViewBinding(device, entry) {
|
|
180
191
|
const buffer = entry instanceof GPUBuffer ? entry : entry.buffer;
|
|
181
192
|
const offset = entry instanceof GPUBuffer ? 0 : (entry.offset ?? 0);
|
|
182
|
-
const avail =
|
|
193
|
+
const avail =
|
|
194
|
+
entry instanceof GPUBuffer
|
|
195
|
+
? entry.size
|
|
196
|
+
: (entry.size ?? buffer.size - offset);
|
|
183
197
|
const size = Math.floor(avail / VEC4_ELEM_BYTES) * VEC4_ELEM_BYTES;
|
|
184
198
|
if (size < VEC4_ELEM_BYTES) {
|
|
185
|
-
return {
|
|
199
|
+
return {
|
|
200
|
+
buffer: vec4FallbackBuffer(device),
|
|
201
|
+
offset: 0,
|
|
202
|
+
size: VEC4_ELEM_BYTES,
|
|
203
|
+
};
|
|
186
204
|
}
|
|
187
205
|
return { buffer, offset, size };
|
|
188
206
|
}
|
|
@@ -206,12 +224,16 @@ export function vec4Usable(entry, stride, outerCount, innerCount) {
|
|
|
206
224
|
if (stride % 4 !== 0) return false;
|
|
207
225
|
const buffer = entry instanceof GPUBuffer ? entry : entry.buffer;
|
|
208
226
|
const offset = entry instanceof GPUBuffer ? 0 : (entry.offset ?? 0);
|
|
209
|
-
const avail =
|
|
227
|
+
const avail =
|
|
228
|
+
entry instanceof GPUBuffer
|
|
229
|
+
? buffer.size
|
|
230
|
+
: (entry.size ?? buffer.size - offset);
|
|
210
231
|
const viewFloats = Math.floor(avail / VEC4_ELEM_BYTES) * 4;
|
|
211
232
|
if (viewFloats <= 0) return false;
|
|
212
233
|
// Highest flat index any masked-in component can touch; usable iff its
|
|
213
234
|
// containing vec4 ends within the view.
|
|
214
|
-
const maxFlat =
|
|
235
|
+
const maxFlat =
|
|
236
|
+
(Math.max(outerCount, 1) - 1) * stride + (Math.max(innerCount, 1) - 1);
|
|
215
237
|
return Math.floor(maxFlat / 4) * 4 + 4 <= viewFloats;
|
|
216
238
|
}
|
|
217
239
|
|
|
@@ -226,7 +248,6 @@ export function vec4Usable(entry, stride, outerCount, innerCount) {
|
|
|
226
248
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUQueue/writeBuffer GPUQueue.writeBuffer()}
|
|
227
249
|
*/
|
|
228
250
|
export function createParamsBuffer(device, params, label = "blas-params") {
|
|
229
|
-
|
|
230
251
|
const rawSize = params.length * 4;
|
|
231
252
|
const size = Math.ceil(rawSize / 16) * 16;
|
|
232
253
|
|
|
@@ -0,0 +1,87 @@
|
|
|
1
|
+
/** @module devdocs/utility-functions/complex */
|
|
2
|
+
import { splitDoubleDouble, mergeDoubleDouble } from "./f64.mjs";
|
|
3
|
+
import { Complex64, Complex64Array } from "../classes/Complex64.mjs";
|
|
4
|
+
|
|
5
|
+
// Complex32Array <-> interleaved [re0, im0, re1, im1, ...] f32 buffer — one
|
|
6
|
+
// buffer, not a (hi, lo) pair like Float64Array (see f64.mjs), since re/im
|
|
7
|
+
// need no error compensation. Matches cuBLAS/stdlib's own complex layout.
|
|
8
|
+
|
|
9
|
+
/**
|
|
10
|
+
* Interleaves the first `n` elements of a Complex32Array into a flat f32
|
|
11
|
+
* buffer ready for `uploadBuffer`.
|
|
12
|
+
* @param {import("../classes/Complex32.mjs").Complex32Array} data
|
|
13
|
+
* @param {number} [n] - element count to interleave (default: data.length)
|
|
14
|
+
* @returns {Float32Array} length `2*n`, [re0, im0, re1, im1, ...]
|
|
15
|
+
* @public
|
|
16
|
+
*/
|
|
17
|
+
export function interleaveComplex32(data, n = data.length) {
|
|
18
|
+
const flat = new Float32Array(n * 2);
|
|
19
|
+
for (let i = 0; i < n; i++) {
|
|
20
|
+
flat[i * 2] = data[i].re;
|
|
21
|
+
flat[i * 2 + 1] = data[i].im;
|
|
22
|
+
}
|
|
23
|
+
return flat;
|
|
24
|
+
}
|
|
25
|
+
|
|
26
|
+
// Complex64Array <-> a double-double (hi, lo) pair of interleaved f32
|
|
27
|
+
// buffers. re and im each need their own (hi, lo) split (see f64.mjs) —
|
|
28
|
+
// zipped together per channel, [reHi0,imHi0,...] / [reLo0,imLo0,...],
|
|
29
|
+
// rather than four separate buffers, so GpuVector/GpuMatrix's existing
|
|
30
|
+
// two-buffer shape covers this dtype with no restructuring.
|
|
31
|
+
|
|
32
|
+
/**
|
|
33
|
+
* Splits the first `n` elements of a Complex64Array into an interleaved
|
|
34
|
+
* double-double (hi, lo) pair of f32 buffers.
|
|
35
|
+
* @param {Complex64Array} data
|
|
36
|
+
* @param {number} [n] - element count to split (default: data.length)
|
|
37
|
+
* @returns {{hi: Float32Array, lo: Float32Array}} each length `2*n`,
|
|
38
|
+
* [reHi0, imHi0, reHi1, imHi1, ...] / [reLo0, imLo0, reLo1, imLo1, ...]
|
|
39
|
+
* @public
|
|
40
|
+
*/
|
|
41
|
+
export function splitComplex64(data, n = data.length) {
|
|
42
|
+
const re = new Float64Array(n);
|
|
43
|
+
const im = new Float64Array(n);
|
|
44
|
+
for (let i = 0; i < n; i++) {
|
|
45
|
+
re[i] = data[i].re;
|
|
46
|
+
im[i] = data[i].im;
|
|
47
|
+
}
|
|
48
|
+
const { hi: reHi, lo: reLo } = splitDoubleDouble(re);
|
|
49
|
+
const { hi: imHi, lo: imLo } = splitDoubleDouble(im);
|
|
50
|
+
|
|
51
|
+
const hi = new Float32Array(n * 2);
|
|
52
|
+
const lo = new Float32Array(n * 2);
|
|
53
|
+
for (let i = 0; i < n; i++) {
|
|
54
|
+
hi[i * 2] = reHi[i];
|
|
55
|
+
hi[i * 2 + 1] = imHi[i];
|
|
56
|
+
lo[i * 2] = reLo[i];
|
|
57
|
+
lo[i * 2 + 1] = imLo[i];
|
|
58
|
+
}
|
|
59
|
+
return { hi, lo };
|
|
60
|
+
}
|
|
61
|
+
|
|
62
|
+
/**
|
|
63
|
+
* Reassembles a Complex64Array from an interleaved double-double (hi, lo)
|
|
64
|
+
* pair of f32 buffers — the inverse of splitComplex64.
|
|
65
|
+
* @param {Float32Array} hi - [reHi0, imHi0, reHi1, imHi1, ...]
|
|
66
|
+
* @param {Float32Array} lo - [reLo0, imLo0, reLo1, imLo1, ...]
|
|
67
|
+
* @returns {Complex64Array}
|
|
68
|
+
* @public
|
|
69
|
+
*/
|
|
70
|
+
export function mergeComplex64(hi, lo) {
|
|
71
|
+
const n = hi.length / 2;
|
|
72
|
+
const reHi = new Float32Array(n),
|
|
73
|
+
reLo = new Float32Array(n);
|
|
74
|
+
const imHi = new Float32Array(n),
|
|
75
|
+
imLo = new Float32Array(n);
|
|
76
|
+
for (let i = 0; i < n; i++) {
|
|
77
|
+
reHi[i] = hi[i * 2];
|
|
78
|
+
imHi[i] = hi[i * 2 + 1];
|
|
79
|
+
reLo[i] = lo[i * 2];
|
|
80
|
+
imLo[i] = lo[i * 2 + 1];
|
|
81
|
+
}
|
|
82
|
+
const re = mergeDoubleDouble(reHi, reLo);
|
|
83
|
+
const im = mergeDoubleDouble(imHi, imLo);
|
|
84
|
+
const out = new Complex64Array(n);
|
|
85
|
+
for (let i = 0; i < n; i++) out[i] = new Complex64(re[i], im[i]);
|
|
86
|
+
return out;
|
|
87
|
+
}
|
package/src/util/compute.mjs
CHANGED
|
@@ -1,9 +1,6 @@
|
|
|
1
1
|
/** @module devdocs/utility-functions/compute */
|
|
2
2
|
import { beginTimestamp, resolveTimestamp } from "./benchmark.mjs";
|
|
3
3
|
|
|
4
|
-
// Anchors the pass encoder to its command encoder to prevent premature GC.
|
|
5
|
-
const _passEncoders = new WeakMap();
|
|
6
|
-
|
|
7
4
|
/**
|
|
8
5
|
* Finalises `commandEncoder` into a command buffer and submits it to the GPU queue.
|
|
9
6
|
* @param {GPUCommandEncoder} commandEncoder
|
|
@@ -40,7 +37,13 @@ export function beginTimedEncoder(device) {
|
|
|
40
37
|
* @param {object} [passDescriptor] - passed to `beginComputePass`, e.g. for timestamp writes
|
|
41
38
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUCommandEncoder/beginComputePass GPUCommandEncoder.beginComputePass()}
|
|
42
39
|
*/
|
|
43
|
-
export function encodePass(
|
|
40
|
+
export function encodePass(
|
|
41
|
+
commandEncoder,
|
|
42
|
+
pipeline,
|
|
43
|
+
bindGroup,
|
|
44
|
+
workgroups,
|
|
45
|
+
passDescriptor,
|
|
46
|
+
) {
|
|
44
47
|
const passEncoder = commandEncoder.beginComputePass(passDescriptor);
|
|
45
48
|
|
|
46
49
|
passEncoder.setPipeline(pipeline);
|
|
@@ -50,12 +53,14 @@ export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, pass
|
|
|
50
53
|
passEncoder.dispatchWorkgroups(workgroups);
|
|
51
54
|
} else {
|
|
52
55
|
// `?? 1` is load-bearing — confirmed passing undefined (from an {x,y}-only caller) crashes the process, not just no-ops.
|
|
53
|
-
passEncoder.dispatchWorkgroups(
|
|
56
|
+
passEncoder.dispatchWorkgroups(
|
|
57
|
+
workgroups.x,
|
|
58
|
+
workgroups.y,
|
|
59
|
+
workgroups.z ?? 1,
|
|
60
|
+
);
|
|
54
61
|
}
|
|
55
62
|
|
|
56
63
|
passEncoder.end();
|
|
57
|
-
|
|
58
|
-
_passEncoders.set(commandEncoder, passEncoder);
|
|
59
64
|
}
|
|
60
65
|
|
|
61
66
|
/**
|
|
@@ -69,7 +74,8 @@ export function encodePass(commandEncoder, pipeline, bindGroup, workgroups, pass
|
|
|
69
74
|
* @see {@link https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/createCommandEncoder GPUDevice.createCommandEncoder()}
|
|
70
75
|
*/
|
|
71
76
|
export function runComputePass(device, pipeline, bindGroup, workgroups) {
|
|
72
|
-
const { commandEncoder, querySet, passDescriptor } =
|
|
77
|
+
const { commandEncoder, querySet, passDescriptor } =
|
|
78
|
+
beginTimedEncoder(device);
|
|
73
79
|
encodePass(commandEncoder, pipeline, bindGroup, workgroups, passDescriptor);
|
|
74
80
|
|
|
75
81
|
const ts = resolveTimestamp(device, commandEncoder, querySet);
|
package/src/util/device.mjs
CHANGED
|
@@ -2,6 +2,20 @@
|
|
|
2
2
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
3
3
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
4
4
|
|
|
5
|
+
/**
|
|
6
|
+
* Throws if `device` is not a `GPUDevice`.
|
|
7
|
+
*
|
|
8
|
+
* Every routine's first guard, extracted since the check and message are
|
|
9
|
+
* identical across all of them.
|
|
10
|
+
*
|
|
11
|
+
* @param {GPUDevice} device - the value to check
|
|
12
|
+
* @throws {Error} if `device` is not a `GPUDevice`
|
|
13
|
+
*/
|
|
14
|
+
export function requireGpuDevice(device) {
|
|
15
|
+
if (!(device instanceof GPUDevice))
|
|
16
|
+
throw new Error("device must be a GPUDevice.");
|
|
17
|
+
}
|
|
18
|
+
|
|
5
19
|
/**
|
|
6
20
|
* Throws if any GPU-resident operand belongs to a device other than the one
|
|
7
21
|
* the routine was called with.
|
|
@@ -22,12 +36,13 @@ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
|
22
36
|
*/
|
|
23
37
|
export function requireSameDevice(device, routine, operands) {
|
|
24
38
|
for (const [name, value] of Object.entries(operands)) {
|
|
25
|
-
if (!(value instanceof GpuVector) && !(value instanceof GpuMatrix))
|
|
39
|
+
if (!(value instanceof GpuVector) && !(value instanceof GpuMatrix))
|
|
40
|
+
continue;
|
|
26
41
|
if (value.device !== device) {
|
|
27
42
|
throw new Error(
|
|
28
43
|
`${routine}: ${name} belongs to a different GPUDevice than the one passed in. ` +
|
|
29
|
-
|
|
30
|
-
|
|
44
|
+
"GPU buffers cannot be shared across devices — recreate the operand on this " +
|
|
45
|
+
"device, or call the routine with the device that owns it.",
|
|
31
46
|
);
|
|
32
47
|
}
|
|
33
48
|
}
|
package/src/util/pipeline.mjs
CHANGED
|
@@ -23,7 +23,14 @@ export async function getPipeline(device, shaderName, entryPoint = "main") {
|
|
|
23
23
|
const names = Array.isArray(shaderName) ? shaderName : [shaderName];
|
|
24
24
|
const key = `${names.join("+")}::${entryPoint}`;
|
|
25
25
|
if (!byName.has(key)) {
|
|
26
|
-
|
|
26
|
+
// Cache the in-flight promise (not its resolved value) so a concurrent
|
|
27
|
+
// call awaits the same compile instead of starting a duplicate one;
|
|
28
|
+
// drop the entry on failure so a later call can retry.
|
|
29
|
+
const pending = loadShader(device, names, entryPoint).catch((err) => {
|
|
30
|
+
byName.delete(key);
|
|
31
|
+
throw err;
|
|
32
|
+
});
|
|
33
|
+
byName.set(key, pending);
|
|
27
34
|
}
|
|
28
35
|
return byName.get(key);
|
|
29
36
|
}
|
|
@@ -39,7 +46,8 @@ async function loadCode(shaderName) {
|
|
|
39
46
|
if (typeof process === "undefined" || !process.versions?.node) {
|
|
40
47
|
const { shaderSources } = await import("../shaders/index.mjs");
|
|
41
48
|
const src = shaderSources[shaderName];
|
|
42
|
-
if (!src)
|
|
49
|
+
if (!src)
|
|
50
|
+
throw new Error(`Shader "${shaderName}" not found in browser bundle.`);
|
|
43
51
|
return src;
|
|
44
52
|
} else {
|
|
45
53
|
const { readFileSync } = await import("fs");
|
|
@@ -66,8 +74,32 @@ async function loadCode(shaderName) {
|
|
|
66
74
|
*/
|
|
67
75
|
export async function loadShader(device, shaderNames, entryPoint = "main") {
|
|
68
76
|
const label = shaderNames.join("+");
|
|
69
|
-
const
|
|
77
|
+
const codes = await Promise.all(shaderNames.map(loadCode));
|
|
70
78
|
|
|
79
|
+
// Per-file line ranges within the concatenated module, so a compile error
|
|
80
|
+
// reports as "<file>.wgsl:<local line>" instead of a whole-module line
|
|
81
|
+
// number — confusing for multi-file pipelines like dasum's.
|
|
82
|
+
let offset = 0;
|
|
83
|
+
const ranges = codes.map((c, i) => {
|
|
84
|
+
const lineCount = c.split("\n").length;
|
|
85
|
+
const range = {
|
|
86
|
+
name: shaderNames[i],
|
|
87
|
+
startLine: offset + 1,
|
|
88
|
+
endLine: offset + lineCount,
|
|
89
|
+
};
|
|
90
|
+
offset += lineCount;
|
|
91
|
+
return range;
|
|
92
|
+
});
|
|
93
|
+
const locate = (lineNum) => {
|
|
94
|
+
const range =
|
|
95
|
+
lineNum &&
|
|
96
|
+
ranges.find((r) => lineNum >= r.startLine && lineNum <= r.endLine);
|
|
97
|
+
return range
|
|
98
|
+
? `${range.name}.wgsl:${lineNum - range.startLine + 1}`
|
|
99
|
+
: `line ${lineNum}`;
|
|
100
|
+
};
|
|
101
|
+
|
|
102
|
+
const code = codes.join("\n");
|
|
71
103
|
const shaderModule = device.createShaderModule({ label, code });
|
|
72
104
|
|
|
73
105
|
const info = await shaderModule.getCompilationInfo();
|
|
@@ -75,7 +107,7 @@ export async function loadShader(device, shaderNames, entryPoint = "main") {
|
|
|
75
107
|
const errors = info.messages.filter((m) => m.type === "error");
|
|
76
108
|
if (errors.length > 0) {
|
|
77
109
|
throw new Error(
|
|
78
|
-
`Shader "${label}" compilation failed:\n${errors.map((m) => `
|
|
110
|
+
`Shader "${label}" compilation failed:\n${errors.map((m) => ` ${locate(m.lineNum)}: ${m.message}`).join("\n")}`,
|
|
79
111
|
);
|
|
80
112
|
}
|
|
81
113
|
|
|
@@ -83,7 +115,10 @@ export async function loadShader(device, shaderNames, entryPoint = "main") {
|
|
|
83
115
|
// this project's WebGPU backend is unstable (intermittent multi-minute hangs and wrong
|
|
84
116
|
// results, confirmed by bisection) when entryPoint is set explicitly, even to the shader's
|
|
85
117
|
// only/correct entry point. Auto-detecting the single entry point is the stable path.
|
|
86
|
-
const compute =
|
|
118
|
+
const compute =
|
|
119
|
+
entryPoint === "main"
|
|
120
|
+
? { module: shaderModule }
|
|
121
|
+
: { module: shaderModule, entryPoint };
|
|
87
122
|
const pipeline = device.createComputePipeline({
|
|
88
123
|
label,
|
|
89
124
|
layout: "auto",
|
package/src/util/workgroup.mjs
CHANGED
|
@@ -2,7 +2,10 @@
|
|
|
2
2
|
// Fixed sizes match the shader declarations (WGS = 64 for 1D, 8×8 = 64 threads
|
|
3
3
|
// for 2D) — see constants.mjs, which is where both values are defined and
|
|
4
4
|
// where the WGSL cross-check hangs off.
|
|
5
|
-
import {
|
|
5
|
+
import {
|
|
6
|
+
WGS as WORKGROUP_SIZE_1D,
|
|
7
|
+
TILE_WG_2D as WORKGROUP_SIZE_2D,
|
|
8
|
+
} from "./constants.mjs";
|
|
6
9
|
|
|
7
10
|
/**
|
|
8
11
|
* Calculates the number of workgroups to dispatch, clamped to the device's
|
|
@@ -53,8 +56,8 @@ export function requireWorkgroupCount(device, count, routine, dim = "x") {
|
|
|
53
56
|
if (count > max)
|
|
54
57
|
throw new Error(
|
|
55
58
|
`${routine}: this problem needs ${count} workgroups in ${dim}, but the device allows ` +
|
|
56
|
-
|
|
57
|
-
|
|
59
|
+
`${max} (maxComputeWorkgroupsPerDimension). The operands are too large for this device — ` +
|
|
60
|
+
`split the operation into smaller blocks.`,
|
|
58
61
|
);
|
|
59
62
|
return count;
|
|
60
63
|
}
|
|
@@ -71,9 +74,23 @@ export function requireWorkgroupCount(device, count, routine, dim = "x") {
|
|
|
71
74
|
*/
|
|
72
75
|
export function requireWorkgroups(device, routine, rows, cols) {
|
|
73
76
|
if (cols === undefined)
|
|
74
|
-
return requireWorkgroupCount(
|
|
77
|
+
return requireWorkgroupCount(
|
|
78
|
+
device,
|
|
79
|
+
Math.ceil(rows / WORKGROUP_SIZE_1D),
|
|
80
|
+
routine,
|
|
81
|
+
);
|
|
75
82
|
return {
|
|
76
|
-
x: requireWorkgroupCount(
|
|
77
|
-
|
|
83
|
+
x: requireWorkgroupCount(
|
|
84
|
+
device,
|
|
85
|
+
Math.ceil(cols / WORKGROUP_SIZE_2D),
|
|
86
|
+
routine,
|
|
87
|
+
"x",
|
|
88
|
+
),
|
|
89
|
+
y: requireWorkgroupCount(
|
|
90
|
+
device,
|
|
91
|
+
Math.ceil(rows / WORKGROUP_SIZE_2D),
|
|
92
|
+
routine,
|
|
93
|
+
"y",
|
|
94
|
+
),
|
|
78
95
|
};
|
|
79
96
|
}
|