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.
- package/README.md +3 -0
- package/dist/wgblas.browser.js +1078 -37
- package/index.d.mts +6 -0
- package/index.mjs +6 -0
- package/package.json +32 -1
- package/src/classes/GpuMatrix.d.mts +85 -0
- package/src/classes/GpuMatrix.mjs +91 -0
- package/src/classes/GpuVector.d.mts +11 -7
- package/src/classes/GpuVector.mjs +31 -5
- package/src/dasum/dasum.d.mts +48 -0
- package/src/dasum/dasum.mjs +121 -0
- package/src/init.mjs +3 -1
- package/src/isamax/isamax.mjs +66 -67
- package/src/random/random.d.mts +37 -0
- package/src/random/random.mjs +17 -0
- package/src/sasum/sasum.mjs +55 -51
- package/src/saxpy/saxpy.mjs +48 -35
- package/src/scopy/scopy.mjs +43 -32
- package/src/sdot/sdot.mjs +60 -55
- package/src/sgemv/sgemv.d.mts +118 -0
- package/src/sgemv/sgemv.mjs +141 -0
- package/src/shaders/browser-shaders.mjs +20 -0
- package/src/shaders/dasum.wgsl +98 -0
- package/src/shaders/f64add.wgsl +281 -0
- package/src/shaders/isamax.wgsl +32 -9
- package/src/shaders/reduction/sumF64.wgsl +49 -0
- package/src/shaders/sasum.wgsl +18 -4
- package/src/shaders/sdot.wgsl +18 -4
- package/src/shaders/sgemv_n.wgsl +75 -0
- package/src/shaders/sgemv_t.wgsl +65 -0
- package/src/shaders/snrm2.wgsl +22 -4
- package/src/shaders/ssymv.wgsl +69 -0
- package/src/shaders/strmv.wgsl +103 -0
- package/src/shaders/strsv_apply_inverse.wgsl +47 -0
- package/src/shaders/strsv_invert_block.wgsl +109 -0
- package/src/shaders/strsv_update.wgsl +75 -0
- package/src/snrm2/snrm2.mjs +56 -52
- package/src/srot/srot.mjs +57 -41
- package/src/srotm/srotm.mjs +54 -38
- package/src/sscal/sscal.mjs +43 -32
- package/src/sswap/sswap.mjs +49 -34
- package/src/ssymv/ssymv.d.mts +109 -0
- package/src/ssymv/ssymv.mjs +130 -0
- package/src/strmv/strmv.d.mts +109 -0
- package/src/strmv/strmv.mjs +132 -0
- package/src/strsv/strsv.d.mts +98 -0
- package/src/strsv/strsv.mjs +212 -0
- package/src/util/benchmark.mjs +1 -1
- package/src/util/bindgroup.mjs +14 -10
- package/src/util/buffer.mjs +7 -2
- package/src/util/compute.mjs +41 -15
- package/src/util/f64pack.mjs +152 -0
- package/src/util/pipeline.mjs +32 -17
- package/src/util/result.mjs +8 -4
- package/src/util/workgroup.mjs +10 -10
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
// strsv_update: subtracts a solved block's contribution from every
|
|
2
|
+
// remaining row in parallel (one workgroup per row, like strmv.wgsl) —
|
|
3
|
+
// this is what turns strsv's O(n) sequential stages into O(n/blockSize).
|
|
4
|
+
// No diag/masking needed: this region never touches the diagonal.
|
|
5
|
+
|
|
6
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
7
|
+
@group(0) @binding(1) var<storage, read_write> x: array<f32>;
|
|
8
|
+
|
|
9
|
+
struct Params {
|
|
10
|
+
n: u32,
|
|
11
|
+
incx: u32,
|
|
12
|
+
lda: u32,
|
|
13
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
14
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
15
|
+
blockStart: u32,
|
|
16
|
+
blockEnd: u32, // exclusive
|
|
17
|
+
}
|
|
18
|
+
|
|
19
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
20
|
+
|
|
21
|
+
const WGS: u32 = 64u;
|
|
22
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
23
|
+
|
|
24
|
+
@compute @workgroup_size(64)
|
|
25
|
+
fn strsv_update_main(
|
|
26
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
27
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
28
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
29
|
+
) {
|
|
30
|
+
// forward: remaining rows are [blockEnd,n); backward: [0,blockStart).
|
|
31
|
+
let forward = (params.trans == 0u) == (params.uplo == 0u);
|
|
32
|
+
|
|
33
|
+
var rangeStart: u32;
|
|
34
|
+
var rangeEnd: u32;
|
|
35
|
+
|
|
36
|
+
if forward {
|
|
37
|
+
rangeStart = params.blockEnd;
|
|
38
|
+
rangeEnd = params.n;
|
|
39
|
+
} else {
|
|
40
|
+
rangeStart = 0u;
|
|
41
|
+
rangeEnd = params.blockStart;
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
if (rangeStart >= rangeEnd) { return; }
|
|
45
|
+
let count = rangeEnd - rangeStart;
|
|
46
|
+
|
|
47
|
+
for (var idx = wgid.x; idx < count; idx += nwg.x) {
|
|
48
|
+
let i = rangeStart + idx;
|
|
49
|
+
|
|
50
|
+
// No-trans reads A[i,j]; transpose reads A[j,i] — uplo only sets the range above.
|
|
51
|
+
var acc = 0.0f;
|
|
52
|
+
if params.trans == 0u {
|
|
53
|
+
for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
|
|
54
|
+
acc += A[i * params.lda + j] * x[j * params.incx];
|
|
55
|
+
}
|
|
56
|
+
} else {
|
|
57
|
+
for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
|
|
58
|
+
acc += A[j * params.lda + i] * x[j * params.incx];
|
|
59
|
+
}
|
|
60
|
+
}
|
|
61
|
+
|
|
62
|
+
// Parallel reduction: 64 → 1
|
|
63
|
+
scratch[lid.x] = acc;
|
|
64
|
+
workgroupBarrier();
|
|
65
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
66
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
67
|
+
workgroupBarrier();
|
|
68
|
+
}
|
|
69
|
+
|
|
70
|
+
if lid.x == 0u {
|
|
71
|
+
x[i * params.incx] -= scratch[0];
|
|
72
|
+
}
|
|
73
|
+
workgroupBarrier();
|
|
74
|
+
}
|
|
75
|
+
}
|
package/src/snrm2/snrm2.mjs
CHANGED
|
@@ -34,67 +34,71 @@ export async function snrm2(device, n, x, incx) {
|
|
|
34
34
|
const pipelineMain = await getPipeline(device, "snrm2");
|
|
35
35
|
const pipelineReduce = await getPipeline(device, "reduction/sum");
|
|
36
36
|
|
|
37
|
-
|
|
38
|
-
|
|
39
|
-
|
|
40
|
-
|
|
41
|
-
|
|
42
|
-
{ value: n, type: "u32" },
|
|
43
|
-
{ value: incx, type: "u32" },
|
|
44
|
-
],
|
|
45
|
-
"snrm2-params",
|
|
46
|
-
);
|
|
37
|
+
let xBuffer = null;
|
|
38
|
+
let partialsBuffer = null;
|
|
39
|
+
let resultBuffer = null;
|
|
40
|
+
let paramsBuffer = null;
|
|
41
|
+
let readBuffer = null;
|
|
47
42
|
|
|
48
|
-
|
|
49
|
-
xBuffer,
|
|
50
|
-
partialsBuffer,
|
|
51
|
-
|
|
52
|
-
|
|
53
|
-
|
|
54
|
-
|
|
55
|
-
|
|
56
|
-
|
|
57
|
-
|
|
43
|
+
try {
|
|
44
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "snrm2-x", false);
|
|
45
|
+
partialsBuffer = createStorageBuffer(2 * WGS * 4, "snrm2-partials"); // 2*WGS partial sums of f32
|
|
46
|
+
resultBuffer = createResultBuffer(4, "snrm2-result"); // final f32 scalar
|
|
47
|
+
paramsBuffer = createParamsBuffer(
|
|
48
|
+
[
|
|
49
|
+
{ value: n, type: "u32" },
|
|
50
|
+
{ value: incx, type: "u32" },
|
|
51
|
+
],
|
|
52
|
+
"snrm2-params",
|
|
53
|
+
);
|
|
54
|
+
|
|
55
|
+
const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
|
|
56
|
+
xBuffer,
|
|
57
|
+
partialsBuffer,
|
|
58
|
+
paramsBuffer,
|
|
59
|
+
]);
|
|
60
|
+
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
61
|
+
pipelineMain,
|
|
62
|
+
bgMain,
|
|
63
|
+
2 * WGS,
|
|
64
|
+
); // dispatch 2*WGS workgroups
|
|
58
65
|
|
|
59
|
-
|
|
66
|
+
submit(enc1);
|
|
60
67
|
|
|
61
|
-
|
|
62
|
-
|
|
63
|
-
|
|
64
|
-
|
|
65
|
-
|
|
66
|
-
|
|
67
|
-
|
|
68
|
-
|
|
69
|
-
|
|
70
|
-
|
|
68
|
+
const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
|
|
69
|
+
partialsBuffer,
|
|
70
|
+
resultBuffer,
|
|
71
|
+
]);
|
|
72
|
+
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
73
|
+
pipelineReduce,
|
|
74
|
+
bgReduce,
|
|
75
|
+
1,
|
|
76
|
+
); // reduce partials to single result
|
|
77
|
+
readBuffer = stageReadback(enc2, resultBuffer);
|
|
71
78
|
|
|
72
|
-
|
|
79
|
+
submit(enc2);
|
|
73
80
|
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
extractTimestamp(ts2),
|
|
77
|
-
extractResult(readBuffer, Float32Array),
|
|
78
|
-
]);
|
|
81
|
+
const resultPromise = extractResult(readBuffer, Float32Array);
|
|
82
|
+
readBuffer = null; // ownership transferred — extractResult's own finally destroys it
|
|
79
83
|
|
|
80
|
-
|
|
81
|
-
|
|
84
|
+
const [gpuTime1, gpuTime2, sqsumArr] = await Promise.all([
|
|
85
|
+
extractTimestamp(ts1),
|
|
86
|
+
extractTimestamp(ts2),
|
|
87
|
+
resultPromise,
|
|
88
|
+
]);
|
|
89
|
+
|
|
90
|
+
// sqrt is taken on CPU after the GPU sum-of-squares reduction
|
|
91
|
+
const nrm2 = Math.sqrt(sqsumArr[0]);
|
|
82
92
|
|
|
83
|
-
if (xIsGpu) {
|
|
84
|
-
destroyBuffers(partialsBuffer, resultBuffer, paramsBuffer, readBuffer);
|
|
85
93
|
if (gpuTime1 !== undefined && gpuTime2 !== undefined)
|
|
86
94
|
return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
|
|
87
95
|
return { nrm2 };
|
|
96
|
+
} finally {
|
|
97
|
+
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
98
|
+
if (partialsBuffer) destroyBuffers(partialsBuffer);
|
|
99
|
+
if (resultBuffer) destroyBuffers(resultBuffer);
|
|
100
|
+
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
101
|
+
// Only reached if submit(enc2) threw before ownership was transferred above.
|
|
102
|
+
if (readBuffer) destroyBuffers(readBuffer);
|
|
88
103
|
}
|
|
89
|
-
|
|
90
|
-
destroyBuffers(
|
|
91
|
-
xBuffer,
|
|
92
|
-
partialsBuffer,
|
|
93
|
-
resultBuffer,
|
|
94
|
-
paramsBuffer,
|
|
95
|
-
readBuffer,
|
|
96
|
-
);
|
|
97
|
-
if (gpuTime1 !== undefined && gpuTime2 !== undefined)
|
|
98
|
-
return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
|
|
99
|
-
return { nrm2 };
|
|
100
104
|
}
|
package/src/srot/srot.mjs
CHANGED
|
@@ -24,9 +24,11 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
24
24
|
!Number.isInteger(incy)
|
|
25
25
|
)
|
|
26
26
|
throw new Error("n, incx, and incy must be integers.");
|
|
27
|
-
if (
|
|
28
|
-
if (
|
|
29
|
-
if (
|
|
27
|
+
if (typeof c !== "number") throw new Error("c must be a number.");
|
|
28
|
+
if (typeof s !== "number") throw new Error("s must be a number.");
|
|
29
|
+
if (Number.isNaN(c) || Number.isNaN(s)) throw new Error("c and s must not be NaN.");
|
|
30
|
+
if (!Number.isFinite(c)) throw new Error("c must be finite.");
|
|
31
|
+
if (!Number.isFinite(s)) throw new Error("s must be finite.");
|
|
30
32
|
if (incx <= 0 || incy <= 0)
|
|
31
33
|
throw new Error("incx and incy must be positive.");
|
|
32
34
|
if (!xIsGpu && !(x instanceof Float32Array))
|
|
@@ -49,47 +51,61 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
49
51
|
|
|
50
52
|
const pipeline = await getPipeline(device, "srot");
|
|
51
53
|
|
|
52
|
-
|
|
53
|
-
|
|
54
|
-
|
|
55
|
-
|
|
56
|
-
|
|
57
|
-
{ value: c, type: "f32" },
|
|
58
|
-
{ value: s, type: "f32" },
|
|
59
|
-
{ value: incx, type: "u32" },
|
|
60
|
-
{ value: incy, type: "u32" },
|
|
61
|
-
],
|
|
62
|
-
"srot-params",
|
|
63
|
-
);
|
|
54
|
+
let xBuffer = null;
|
|
55
|
+
let yBuffer = null;
|
|
56
|
+
let paramsBuffer = null;
|
|
57
|
+
let readX = null;
|
|
58
|
+
let readY = null;
|
|
64
59
|
|
|
65
|
-
|
|
66
|
-
xBuffer,
|
|
67
|
-
yBuffer,
|
|
68
|
-
paramsBuffer
|
|
69
|
-
|
|
70
|
-
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
|
|
60
|
+
try {
|
|
61
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srot-x", true);
|
|
62
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srot-y", true);
|
|
63
|
+
paramsBuffer = createParamsBuffer(
|
|
64
|
+
[
|
|
65
|
+
{ value: n, type: "u32" },
|
|
66
|
+
{ value: c, type: "f32" },
|
|
67
|
+
{ value: s, type: "f32" },
|
|
68
|
+
{ value: incx, type: "u32" },
|
|
69
|
+
{ value: incy, type: "u32" },
|
|
70
|
+
],
|
|
71
|
+
"srot-params",
|
|
72
|
+
);
|
|
78
73
|
|
|
79
|
-
|
|
74
|
+
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
75
|
+
xBuffer,
|
|
76
|
+
yBuffer,
|
|
77
|
+
paramsBuffer,
|
|
78
|
+
]);
|
|
79
|
+
const { commandEncoder, ts } = runComputePass(
|
|
80
|
+
pipeline,
|
|
81
|
+
bindGroup,
|
|
82
|
+
calcWorkgroups(n),
|
|
83
|
+
);
|
|
84
|
+
readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
85
|
+
readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
86
|
+
submit(commandEncoder);
|
|
80
87
|
|
|
81
|
-
|
|
82
|
-
destroyBuffers(paramsBuffer);
|
|
83
|
-
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
84
|
-
return {};
|
|
85
|
-
}
|
|
88
|
+
const gpuTimeMs = await extractTimestamp(ts);
|
|
86
89
|
|
|
87
|
-
|
|
88
|
-
|
|
89
|
-
|
|
90
|
-
|
|
91
|
-
destroyBuffers(xBuffer, yBuffer, paramsBuffer, readX, readY);
|
|
90
|
+
if (xIsGpu && yIsGpu) {
|
|
91
|
+
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
92
|
+
return {};
|
|
93
|
+
}
|
|
92
94
|
|
|
93
|
-
|
|
94
|
-
|
|
95
|
+
const xPromise = extractResult(readX, Float32Array);
|
|
96
|
+
const yPromise = extractResult(readY, Float32Array);
|
|
97
|
+
readX = null; // ownership transferred — extractResult's own finally destroys it
|
|
98
|
+
readY = null;
|
|
99
|
+
const [xResult, yResult] = await Promise.all([xPromise, yPromise]);
|
|
100
|
+
|
|
101
|
+
if (gpuTimeMs !== undefined) return { x: xResult, y: yResult, gpuTimeMs };
|
|
102
|
+
return { x: xResult, y: yResult };
|
|
103
|
+
} finally {
|
|
104
|
+
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
105
|
+
if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
|
|
106
|
+
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
107
|
+
// Only reached if extractTimestamp threw before ownership was transferred above.
|
|
108
|
+
if (readX) destroyBuffers(readX);
|
|
109
|
+
if (readY) destroyBuffers(readY);
|
|
110
|
+
}
|
|
95
111
|
}
|
package/src/srotm/srotm.mjs
CHANGED
|
@@ -48,47 +48,63 @@ export async function srotm(device, n, x, incx, y, incy, param) {
|
|
|
48
48
|
|
|
49
49
|
const pipeline = await getPipeline(device, "srotm");
|
|
50
50
|
|
|
51
|
-
|
|
52
|
-
|
|
53
|
-
|
|
54
|
-
|
|
55
|
-
|
|
56
|
-
|
|
57
|
-
{ value: incx, type: "u32" },
|
|
58
|
-
{ value: incy, type: "u32" },
|
|
59
|
-
],
|
|
60
|
-
"srotm-params",
|
|
61
|
-
);
|
|
51
|
+
let xBuffer = null;
|
|
52
|
+
let yBuffer = null;
|
|
53
|
+
let paramBuffer = null;
|
|
54
|
+
let paramsBuffer = null;
|
|
55
|
+
let readX = null;
|
|
56
|
+
let readY = null;
|
|
62
57
|
|
|
63
|
-
|
|
64
|
-
xBuffer,
|
|
65
|
-
yBuffer,
|
|
66
|
-
paramBuffer,
|
|
67
|
-
paramsBuffer
|
|
68
|
-
|
|
69
|
-
|
|
70
|
-
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
|
|
74
|
-
|
|
75
|
-
const readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
76
|
-
submit(commandEncoder);
|
|
58
|
+
try {
|
|
59
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srotm-x", true);
|
|
60
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srotm-y", true);
|
|
61
|
+
paramBuffer = uploadBuffer(param, "srotm-param", false);
|
|
62
|
+
paramsBuffer = createParamsBuffer(
|
|
63
|
+
[
|
|
64
|
+
{ value: n, type: "u32" },
|
|
65
|
+
{ value: incx, type: "u32" },
|
|
66
|
+
{ value: incy, type: "u32" },
|
|
67
|
+
],
|
|
68
|
+
"srotm-params",
|
|
69
|
+
);
|
|
77
70
|
|
|
78
|
-
|
|
71
|
+
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
72
|
+
xBuffer,
|
|
73
|
+
yBuffer,
|
|
74
|
+
paramBuffer,
|
|
75
|
+
paramsBuffer,
|
|
76
|
+
]);
|
|
77
|
+
const { commandEncoder, ts } = runComputePass(
|
|
78
|
+
pipeline,
|
|
79
|
+
bindGroup,
|
|
80
|
+
calcWorkgroups(n),
|
|
81
|
+
);
|
|
82
|
+
readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
83
|
+
readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
84
|
+
submit(commandEncoder);
|
|
79
85
|
|
|
80
|
-
|
|
81
|
-
destroyBuffers(paramBuffer, paramsBuffer);
|
|
82
|
-
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
83
|
-
return {};
|
|
84
|
-
}
|
|
86
|
+
const gpuTimeMs = await extractTimestamp(ts);
|
|
85
87
|
|
|
86
|
-
|
|
87
|
-
|
|
88
|
-
|
|
89
|
-
|
|
90
|
-
destroyBuffers(xBuffer, yBuffer, paramBuffer, paramsBuffer, readX, readY);
|
|
88
|
+
if (xIsGpu && yIsGpu) {
|
|
89
|
+
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
90
|
+
return {};
|
|
91
|
+
}
|
|
91
92
|
|
|
92
|
-
|
|
93
|
-
|
|
93
|
+
const xPromise = extractResult(readX, Float32Array);
|
|
94
|
+
const yPromise = extractResult(readY, Float32Array);
|
|
95
|
+
readX = null; // ownership transferred — extractResult's own finally destroys it
|
|
96
|
+
readY = null;
|
|
97
|
+
const [xResult, yResult] = await Promise.all([xPromise, yPromise]);
|
|
98
|
+
|
|
99
|
+
if (gpuTimeMs !== undefined) return { x: xResult, y: yResult, gpuTimeMs };
|
|
100
|
+
return { x: xResult, y: yResult };
|
|
101
|
+
} finally {
|
|
102
|
+
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
103
|
+
if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
|
|
104
|
+
if (paramBuffer) destroyBuffers(paramBuffer);
|
|
105
|
+
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
106
|
+
// Only reached if extractTimestamp threw before ownership was transferred above.
|
|
107
|
+
if (readX) destroyBuffers(readX);
|
|
108
|
+
if (readY) destroyBuffers(readY);
|
|
109
|
+
}
|
|
94
110
|
}
|
package/src/sscal/sscal.mjs
CHANGED
|
@@ -19,8 +19,10 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
19
19
|
throw new Error("device must be a GPUDevice.");
|
|
20
20
|
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
21
21
|
throw new Error("n and incx must be integers.");
|
|
22
|
-
if (
|
|
23
|
-
|
|
22
|
+
if (typeof alpha !== "number")
|
|
23
|
+
throw new Error("alpha must be a number.");
|
|
24
|
+
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
25
|
+
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
24
26
|
if (incx <= 0) throw new Error("incx must be positive.");
|
|
25
27
|
if (!(x instanceof Float32Array) && !(x instanceof GpuVector))
|
|
26
28
|
throw new Error("x must be a Float32Array or GpuVector.");
|
|
@@ -32,40 +34,49 @@ export async function sscal(device, n, alpha, x, incx) {
|
|
|
32
34
|
|
|
33
35
|
const pipeline = await getPipeline(device, "sscal");
|
|
34
36
|
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
|
|
38
|
-
{ value: n, type: "u32" },
|
|
39
|
-
{ value: alpha, type: "f32" },
|
|
40
|
-
{ value: incx, type: "u32" },
|
|
41
|
-
],
|
|
42
|
-
"sscal-params",
|
|
43
|
-
);
|
|
37
|
+
let xBuffer = null;
|
|
38
|
+
let paramsBuffer = null;
|
|
39
|
+
let readBuffer = null;
|
|
44
40
|
|
|
45
|
-
|
|
46
|
-
xBuffer,
|
|
47
|
-
paramsBuffer
|
|
48
|
-
|
|
49
|
-
|
|
50
|
-
|
|
51
|
-
|
|
52
|
-
|
|
53
|
-
|
|
54
|
-
|
|
41
|
+
try {
|
|
42
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sscal-x", true);
|
|
43
|
+
paramsBuffer = createParamsBuffer(
|
|
44
|
+
[
|
|
45
|
+
{ value: n, type: "u32" },
|
|
46
|
+
{ value: alpha, type: "f32" },
|
|
47
|
+
{ value: incx, type: "u32" },
|
|
48
|
+
],
|
|
49
|
+
"sscal-params",
|
|
50
|
+
);
|
|
55
51
|
|
|
56
|
-
|
|
52
|
+
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
53
|
+
xBuffer,
|
|
54
|
+
paramsBuffer,
|
|
55
|
+
]);
|
|
56
|
+
const { commandEncoder, ts } = runComputePass(
|
|
57
|
+
pipeline,
|
|
58
|
+
bindGroup,
|
|
59
|
+
calcWorkgroups(n),
|
|
60
|
+
);
|
|
61
|
+
readBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
57
62
|
|
|
58
|
-
|
|
63
|
+
submit(commandEncoder);
|
|
59
64
|
|
|
60
|
-
|
|
61
|
-
destroyBuffers(paramsBuffer);
|
|
62
|
-
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
63
|
-
return {};
|
|
64
|
-
}
|
|
65
|
+
const gpuTimeMs = await extractTimestamp(ts);
|
|
65
66
|
|
|
66
|
-
|
|
67
|
-
|
|
67
|
+
if (xIsGpu) {
|
|
68
|
+
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
69
|
+
return {};
|
|
70
|
+
}
|
|
68
71
|
|
|
69
|
-
|
|
70
|
-
|
|
72
|
+
const result = await extractResult(readBuffer, Float32Array);
|
|
73
|
+
readBuffer = null; // extractResult already destroyed it
|
|
74
|
+
if (gpuTimeMs !== undefined) return { x: result, gpuTimeMs };
|
|
75
|
+
return result;
|
|
76
|
+
} finally {
|
|
77
|
+
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
78
|
+
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
79
|
+
// Only reached if extractTimestamp threw before extractResult ran.
|
|
80
|
+
if (readBuffer) destroyBuffers(readBuffer);
|
|
81
|
+
}
|
|
71
82
|
}
|
package/src/sswap/sswap.mjs
CHANGED
|
@@ -46,45 +46,60 @@ export async function sswap(device, n, x, incx, y, incy) {
|
|
|
46
46
|
|
|
47
47
|
const pipeline = await getPipeline(device, "sswap");
|
|
48
48
|
|
|
49
|
-
|
|
50
|
-
|
|
51
|
-
|
|
52
|
-
|
|
53
|
-
|
|
54
|
-
{ value: incx, type: "u32" },
|
|
55
|
-
{ value: incy, type: "u32" },
|
|
56
|
-
],
|
|
57
|
-
"sswap-params",
|
|
58
|
-
);
|
|
49
|
+
let xBuffer = null;
|
|
50
|
+
let yBuffer = null;
|
|
51
|
+
let paramsBuffer = null;
|
|
52
|
+
let xReadBuffer = null;
|
|
53
|
+
let yReadBuffer = null;
|
|
59
54
|
|
|
60
|
-
|
|
61
|
-
xBuffer,
|
|
62
|
-
yBuffer,
|
|
63
|
-
paramsBuffer
|
|
64
|
-
|
|
65
|
-
|
|
66
|
-
|
|
67
|
-
|
|
68
|
-
|
|
69
|
-
|
|
70
|
-
|
|
71
|
-
const yReadBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
55
|
+
try {
|
|
56
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sswap-x", true);
|
|
57
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sswap-y", true);
|
|
58
|
+
paramsBuffer = createParamsBuffer(
|
|
59
|
+
[
|
|
60
|
+
{ value: n, type: "u32" },
|
|
61
|
+
{ value: incx, type: "u32" },
|
|
62
|
+
{ value: incy, type: "u32" },
|
|
63
|
+
],
|
|
64
|
+
"sswap-params",
|
|
65
|
+
);
|
|
72
66
|
|
|
73
|
-
|
|
67
|
+
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
68
|
+
xBuffer,
|
|
69
|
+
yBuffer,
|
|
70
|
+
paramsBuffer,
|
|
71
|
+
]);
|
|
72
|
+
const { commandEncoder, ts } = runComputePass(
|
|
73
|
+
pipeline,
|
|
74
|
+
bindGroup,
|
|
75
|
+
calcWorkgroups(n),
|
|
76
|
+
);
|
|
77
|
+
xReadBuffer = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
78
|
+
yReadBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
74
79
|
|
|
75
|
-
|
|
80
|
+
submit(commandEncoder);
|
|
76
81
|
|
|
77
|
-
|
|
78
|
-
destroyBuffers(paramsBuffer);
|
|
79
|
-
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
80
|
-
return {};
|
|
81
|
-
}
|
|
82
|
+
const gpuTimeMs = await extractTimestamp(ts);
|
|
82
83
|
|
|
83
|
-
|
|
84
|
-
|
|
84
|
+
if (xIsGpu && yIsGpu) {
|
|
85
|
+
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
86
|
+
return {};
|
|
87
|
+
}
|
|
85
88
|
|
|
86
|
-
|
|
89
|
+
const resultX = await extractResult(xReadBuffer, Float32Array);
|
|
90
|
+
xReadBuffer = null; // extractResult already destroyed it
|
|
91
|
+
const resultY = await extractResult(yReadBuffer, Float32Array);
|
|
92
|
+
yReadBuffer = null; // extractResult already destroyed it
|
|
87
93
|
|
|
88
|
-
|
|
89
|
-
|
|
94
|
+
if (gpuTimeMs !== undefined) return { x: resultX, y: resultY, gpuTimeMs };
|
|
95
|
+
return { x: resultX, y: resultY };
|
|
96
|
+
} finally {
|
|
97
|
+
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
98
|
+
if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
|
|
99
|
+
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
100
|
+
// Only reached if extractTimestamp or extractResult threw before
|
|
101
|
+
// clearing these — on the success path they're already null.
|
|
102
|
+
if (xReadBuffer) destroyBuffers(xReadBuffer);
|
|
103
|
+
if (yReadBuffer) destroyBuffers(yReadBuffer);
|
|
104
|
+
}
|
|
90
105
|
}
|