wgblas 2.0.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 +20 -18
- package/dist/wgblas.browser.js +2172 -1174
- package/index.d.mts +49 -44
- package/index.mjs +11 -0
- package/package.json +133 -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 +126 -17
- package/src/classes/GpuVector.d.mts +36 -40
- package/src/classes/GpuVector.mjs +66 -11
- package/src/cscal/cscal.d.mts +47 -0
- package/src/cscal/cscal.mjs +98 -0
- package/src/dasum/dasum.d.mts +4 -4
- package/src/dasum/dasum.mjs +38 -20
- 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/devdocs.mjs +13 -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.d.mts +20 -2
- package/src/idamax/idamax.mjs +56 -24
- package/src/init.mjs +117 -56
- package/src/isamax/isamax.d.mts +20 -2
- package/src/isamax/isamax.mjs +21 -16
- package/src/random/random.d.mts +37 -39
- package/src/random/random.mjs +39 -7
- package/src/sasum/sasum.d.mts +2 -2
- package/src/sasum/sasum.mjs +20 -16
- package/src/saxpy/saxpy.d.mts +2 -2
- package/src/saxpy/saxpy.mjs +14 -11
- package/src/scopy/scopy.d.mts +2 -2
- package/src/scopy/scopy.mjs +13 -9
- package/src/sdot/sdot.d.mts +2 -2
- package/src/sdot/sdot.mjs +21 -17
- package/src/sgemm/sgemm.d.mts +2 -2
- package/src/sgemm/sgemm.mjs +109 -40
- package/src/sgemmtr/sgemmtr.d.mts +3 -2
- package/src/sgemmtr/sgemmtr.mjs +98 -40
- package/src/sgemv/sgemv.d.mts +2 -2
- package/src/sgemv/sgemv.mjs +69 -41
- package/src/sger/sger.d.mts +2 -2
- package/src/sger/sger.mjs +43 -19
- 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 +233 -14
- package/src/shaders/reduction/scaledSum.wgsl +65 -0
- package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
- package/src/shaders/sgemm_large.wgsl +107 -18
- package/src/shaders/sgemm_small.wgsl +115 -15
- package/src/shaders/sgemmtr_large.wgsl +4 -1
- package/src/shaders/sgemmtr_small.wgsl +4 -1
- package/src/shaders/sgemv_n.wgsl +3 -1
- package/src/shaders/sgemv_t.wgsl +3 -1
- package/src/shaders/snrm2.wgsl +72 -23
- package/src/shaders/ssymv.wgsl +3 -1
- package/src/snrm2/snrm2.d.mts +2 -2
- package/src/snrm2/snrm2.mjs +41 -23
- package/src/srot/srot.d.mts +2 -4
- package/src/srot/srot.mjs +16 -11
- package/src/srotm/srotm.d.mts +2 -4
- package/src/srotm/srotm.mjs +17 -11
- package/src/sscal/sscal.d.mts +3 -3
- package/src/sscal/sscal.mjs +14 -12
- package/src/sswap/sswap.d.mts +2 -2
- package/src/sswap/sswap.mjs +18 -10
- package/src/ssymm/ssymm.d.mts +5 -4
- package/src/ssymm/ssymm.mjs +150 -54
- package/src/ssymv/ssymv.d.mts +2 -2
- package/src/ssymv/ssymv.mjs +47 -26
- package/src/ssyr/ssyr.d.mts +2 -2
- package/src/ssyr/ssyr.mjs +38 -17
- package/src/ssyr2/ssyr2.d.mts +2 -2
- package/src/ssyr2/ssyr2.mjs +48 -21
- package/src/ssyr2k/ssyr2k.d.mts +3 -2
- package/src/ssyr2k/ssyr2k.mjs +140 -62
- package/src/ssyrk/ssyrk.d.mts +3 -2
- package/src/ssyrk/ssyrk.mjs +91 -39
- package/src/strmm/strmm.d.mts +5 -4
- package/src/strmm/strmm.mjs +174 -60
- package/src/strmv/strmv.d.mts +2 -2
- package/src/strmv/strmv.mjs +42 -20
- package/src/strsm/strsm.d.mts +6 -4
- package/src/strsm/strsm.mjs +438 -174
- package/src/strsv/strsv.d.mts +5 -3
- package/src/strsv/strsv.mjs +89 -34
- package/src/util/benchmark.mjs +9 -9
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +139 -24
- package/src/util/complex.mjs +87 -0
- package/src/util/compute.mjs +19 -16
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +49 -0
- package/src/util/pipeline.mjs +44 -10
- package/src/util/workgroup.mjs +72 -7
- package/src/shaders/browser-shaders.mjs +0 -81
- package/src/shaders/f64add.wgsl +0 -281
package/src/isamax/isamax.mjs
CHANGED
|
@@ -12,14 +12,14 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
12
12
|
import { extractResult } from "../util/result.mjs";
|
|
13
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
14
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
15
|
-
|
|
16
|
-
|
|
15
|
+
import { WGS } from "../util/constants.mjs";
|
|
16
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
17
17
|
|
|
18
18
|
export async function isamax(device, n, x, incx) {
|
|
19
19
|
const xIsGpu = x instanceof GpuVector;
|
|
20
20
|
|
|
21
|
-
|
|
22
|
-
|
|
21
|
+
requireGpuDevice(device);
|
|
22
|
+
requireSameDevice(device, "isamax", { x });
|
|
23
23
|
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
24
24
|
throw new Error("n and incx must be integers.");
|
|
25
25
|
if (incx <= 0) throw new Error("incx must be positive.");
|
|
@@ -42,17 +42,20 @@ export async function isamax(device, n, x, incx) {
|
|
|
42
42
|
let readBuffer = null;
|
|
43
43
|
|
|
44
44
|
try {
|
|
45
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "isamax-x", false);
|
|
45
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "isamax-x", false);
|
|
46
46
|
partialsValBuffer = createStorageBuffer(
|
|
47
|
+
device,
|
|
47
48
|
2 * WGS * 4,
|
|
48
49
|
"isamax-partials-val",
|
|
49
50
|
); //to hold 2*WGS partial max values of f32
|
|
50
51
|
partialsIdxBuffer = createStorageBuffer(
|
|
52
|
+
device,
|
|
51
53
|
2 * WGS * 4,
|
|
52
54
|
"isamax-partials-idx",
|
|
53
55
|
); //to hold 2*WGS partial max indices of u32
|
|
54
|
-
resultBuffer = createResultBuffer(4, "isamax-result"); // u32 index
|
|
56
|
+
resultBuffer = createResultBuffer(device, 4, "isamax-result"); // u32 index
|
|
55
57
|
paramsBuffer = createParamsBuffer(
|
|
58
|
+
device,
|
|
56
59
|
[
|
|
57
60
|
{ value: n, type: "u32" },
|
|
58
61
|
{ value: incx, type: "u32" },
|
|
@@ -60,33 +63,35 @@ export async function isamax(device, n, x, incx) {
|
|
|
60
63
|
"isamax-params",
|
|
61
64
|
);
|
|
62
65
|
|
|
63
|
-
const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
|
|
66
|
+
const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
|
|
64
67
|
xBuffer,
|
|
65
68
|
partialsValBuffer,
|
|
66
69
|
partialsIdxBuffer,
|
|
67
70
|
paramsBuffer,
|
|
68
71
|
]);
|
|
69
72
|
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
73
|
+
device,
|
|
70
74
|
pipelineMain,
|
|
71
75
|
bgMain,
|
|
72
76
|
2 * WGS,
|
|
73
77
|
); //dispatch 2*WGS workgroups
|
|
74
78
|
|
|
75
|
-
submit(enc1);
|
|
79
|
+
submit(device, enc1);
|
|
76
80
|
|
|
77
|
-
const bgReduce = createBindGroup(
|
|
78
|
-
|
|
79
|
-
|
|
80
|
-
resultBuffer,
|
|
81
|
-
|
|
81
|
+
const bgReduce = createBindGroup(
|
|
82
|
+
device,
|
|
83
|
+
pipelineReduce.getBindGroupLayout(0),
|
|
84
|
+
[partialsValBuffer, partialsIdxBuffer, resultBuffer],
|
|
85
|
+
);
|
|
82
86
|
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
87
|
+
device,
|
|
83
88
|
pipelineReduce,
|
|
84
89
|
bgReduce,
|
|
85
90
|
1,
|
|
86
91
|
); // dispatch 1 workgroup to reduce the partial sums to a single result
|
|
87
|
-
readBuffer = stageReadback(enc2, resultBuffer);
|
|
92
|
+
readBuffer = stageReadback(device, enc2, resultBuffer);
|
|
88
93
|
|
|
89
|
-
submit(enc2);
|
|
94
|
+
submit(device, enc2);
|
|
90
95
|
|
|
91
96
|
const resultPromise = extractResult(readBuffer, Uint32Array);
|
|
92
97
|
readBuffer = null; // ownership transferred — extractResult's own finally destroys it
|
|
@@ -108,7 +113,7 @@ export async function isamax(device, n, x, incx) {
|
|
|
108
113
|
if (partialsIdxBuffer) destroyBuffers(partialsIdxBuffer);
|
|
109
114
|
if (resultBuffer) destroyBuffers(resultBuffer);
|
|
110
115
|
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
111
|
-
// Only reached if submit(enc2) threw before ownership was transferred above.
|
|
116
|
+
// Only reached if submit(device, enc2) threw before ownership was transferred above.
|
|
112
117
|
if (readBuffer) destroyBuffers(readBuffer);
|
|
113
118
|
}
|
|
114
119
|
}
|
package/src/random/random.d.mts
CHANGED
|
@@ -4,22 +4,19 @@
|
|
|
4
4
|
* @param n - number of elements
|
|
5
5
|
* @param low - lower bound (default: -1)
|
|
6
6
|
* @param high - upper bound (default: 1)
|
|
7
|
+
* @param seed - omit for genuine randomness (`Math.random`, the default);
|
|
8
|
+
* pass any number for a deterministic, reproducible sequence (mulberry32) —
|
|
9
|
+
* useful for regression tests that need "random-looking" data without
|
|
10
|
+
* flaking between runs
|
|
7
11
|
*
|
|
8
|
-
*
|
|
9
|
-
*
|
|
10
|
-
* import { randomFloat32Array } from "wgblas";
|
|
12
|
+
* **Default range [-1, 1):**
|
|
13
|
+
* {@includeCode ../../examples/randomfloat32array/randomfloat32array.js}
|
|
11
14
|
*
|
|
12
|
-
*
|
|
13
|
-
*
|
|
14
|
-
* ```
|
|
15
|
+
* **Custom range [0, 10):**
|
|
16
|
+
* {@includeCode ../../examples/randomfloat32array-custom/randomfloat32array-custom.js}
|
|
15
17
|
*
|
|
16
|
-
*
|
|
17
|
-
*
|
|
18
|
-
* import { randomFloat32Array } from "wgblas";
|
|
19
|
-
*
|
|
20
|
-
* const x = randomFloat32Array(4, 0, 10);
|
|
21
|
-
* console.log(x);
|
|
22
|
-
* ```
|
|
18
|
+
* **Seeded (deterministic):**
|
|
19
|
+
* {@includeCode ../../examples/randomfloat32array-seeded/randomfloat32array-seeded.js}
|
|
23
20
|
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/random/random.mjs#L1">Source code: random.mjs (L1)</a>
|
|
24
21
|
* @category Utilities
|
|
25
22
|
*/
|
|
@@ -27,6 +24,7 @@ export declare function randomFloat32Array(
|
|
|
27
24
|
n: number,
|
|
28
25
|
low?: number,
|
|
29
26
|
high?: number,
|
|
27
|
+
seed?: number,
|
|
30
28
|
): Float32Array;
|
|
31
29
|
|
|
32
30
|
/**
|
|
@@ -35,22 +33,19 @@ export declare function randomFloat32Array(
|
|
|
35
33
|
* @param n - number of elements
|
|
36
34
|
* @param low - lower bound (default: -1)
|
|
37
35
|
* @param high - upper bound (default: 1)
|
|
36
|
+
* @param seed - omit for genuine randomness (`Math.random`, the default);
|
|
37
|
+
* pass any number for a deterministic, reproducible sequence (mulberry32) —
|
|
38
|
+
* useful for regression tests that need "random-looking" data without
|
|
39
|
+
* flaking between runs
|
|
38
40
|
*
|
|
39
|
-
*
|
|
40
|
-
*
|
|
41
|
-
* import { randomFloat64Array } from "wgblas";
|
|
42
|
-
*
|
|
43
|
-
* const x = randomFloat64Array(4);
|
|
44
|
-
* console.log(x); // Float64Array [ -0.42, 0.81, -0.07, 0.55 ]
|
|
45
|
-
* ```
|
|
41
|
+
* **Default range [-1, 1):**
|
|
42
|
+
* {@includeCode ../../examples/randomfloat64array/randomfloat64array.js}
|
|
46
43
|
*
|
|
47
|
-
*
|
|
48
|
-
*
|
|
49
|
-
* import { randomFloat64Array } from "wgblas";
|
|
44
|
+
* **Custom range [0, 10):**
|
|
45
|
+
* {@includeCode ../../examples/randomfloat64array-custom/randomfloat64array-custom.js}
|
|
50
46
|
*
|
|
51
|
-
*
|
|
52
|
-
*
|
|
53
|
-
* ```
|
|
47
|
+
* **Seeded (deterministic):**
|
|
48
|
+
* {@includeCode ../../examples/randomfloat64array-seeded/randomfloat64array-seeded.js}
|
|
54
49
|
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/random/random.mjs#L7">Source code: random.mjs (L7)</a>
|
|
55
50
|
* @category Utilities
|
|
56
51
|
*/
|
|
@@ -58,16 +53,20 @@ export declare function randomFloat64Array(
|
|
|
58
53
|
n: number,
|
|
59
54
|
low?: number,
|
|
60
55
|
high?: number,
|
|
56
|
+
seed?: number,
|
|
61
57
|
): Float64Array;
|
|
62
58
|
|
|
63
59
|
/**
|
|
64
|
-
* Returns a Float32Array of n*lda elements,
|
|
65
|
-
* `
|
|
66
|
-
* triangular routine (`strmv`/`strsv`)
|
|
67
|
-
* triangle are uniform in `[low, high)`, the
|
|
68
|
-
*
|
|
69
|
-
* so a triangular solve doesn't divide by a near-zero pivot — and
|
|
70
|
-
* in the other triangle is 0.
|
|
60
|
+
* Returns a Float32Array of n*lda elements, with leading dimension `lda`
|
|
61
|
+
* under the given `layout`, representing an actual lower- or
|
|
62
|
+
* upper-triangular matrix for a triangular routine (`strmv`/`strsv`)
|
|
63
|
+
* example: entries in the `uplo` triangle are uniform in `[low, high)`, the
|
|
64
|
+
* n diagonal entries are uniform in `[diagLow, diagHigh)` — kept well away
|
|
65
|
+
* from 0 so a triangular solve doesn't divide by a near-zero pivot — and
|
|
66
|
+
* every entry in the other triangle is 0. `uplo` always describes the
|
|
67
|
+
* logical triangle, regardless of storage order: only the flat index each
|
|
68
|
+
* (row, col) maps to changes between layouts (`A[row*lda+col]` for
|
|
69
|
+
* row-major, `A[col*lda+row]` for column-major).
|
|
71
70
|
*
|
|
72
71
|
* @param n - matrix order (rows/cols read by the triangular routine)
|
|
73
72
|
* @param lda - leading dimension; throws if `lda < n`
|
|
@@ -76,14 +75,12 @@ export declare function randomFloat64Array(
|
|
|
76
75
|
* @param high - upper bound for off-diagonal entries (default: 1)
|
|
77
76
|
* @param diagLow - lower bound for diagonal entries (default: 5)
|
|
78
77
|
* @param diagHigh - upper bound for diagonal entries (default: 15)
|
|
78
|
+
* @param layout - `'row-major'` or `'column-major'` storage order (default: `'row-major'`)
|
|
79
79
|
*
|
|
80
|
-
* @
|
|
81
|
-
* ```js
|
|
82
|
-
* import { randomTriangularFloat32Array } from "wgblas";
|
|
80
|
+
* {@includeCode ../../examples/randomtriangularfloat32array/randomtriangularfloat32array.js}
|
|
83
81
|
*
|
|
84
|
-
*
|
|
85
|
-
*
|
|
86
|
-
* ```
|
|
82
|
+
* **Column-major storage:**
|
|
83
|
+
* {@includeCode ../../examples/randomtriangularfloat32array-columnmajor/randomtriangularfloat32array-columnmajor.js}
|
|
87
84
|
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/random/random.mjs#L13">Source code: random.mjs (L13)</a>
|
|
88
85
|
* @category Utilities
|
|
89
86
|
*/
|
|
@@ -95,4 +92,5 @@ export declare function randomTriangularFloat32Array(
|
|
|
95
92
|
high?: number,
|
|
96
93
|
diagLow?: number,
|
|
97
94
|
diagHigh?: number,
|
|
95
|
+
layout?: 'row-major' | 'column-major',
|
|
98
96
|
): Float32Array;
|
package/src/random/random.mjs
CHANGED
|
@@ -1,28 +1,60 @@
|
|
|
1
|
-
|
|
1
|
+
// mulberry32: a small, fast, deterministic PRNG — not cryptographic, but
|
|
2
|
+
// good enough distribution for test/benchmark data. Returns a () => number
|
|
3
|
+
// generator producing values in [0, 1), same contract as Math.random.
|
|
4
|
+
function mulberry32(seed) {
|
|
5
|
+
let a = seed >>> 0;
|
|
6
|
+
return function () {
|
|
7
|
+
a = (a + 0x6d2b79f5) | 0;
|
|
8
|
+
let t = Math.imul(a ^ (a >>> 15), 1 | a);
|
|
9
|
+
t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t;
|
|
10
|
+
return ((t ^ (t >>> 14)) >>> 0) / 4294967296;
|
|
11
|
+
};
|
|
12
|
+
}
|
|
13
|
+
|
|
14
|
+
export function randomFloat32Array(n, low = -1, high = 1, seed) {
|
|
2
15
|
const x = new Float32Array(n);
|
|
3
|
-
|
|
16
|
+
const next = seed === undefined ? Math.random : mulberry32(seed);
|
|
17
|
+
for (let i = 0; i < n; i++) x[i] = low + next() * (high - low);
|
|
4
18
|
return x;
|
|
5
19
|
}
|
|
6
20
|
|
|
7
|
-
export function randomFloat64Array(n, low = -1, high = 1) {
|
|
21
|
+
export function randomFloat64Array(n, low = -1, high = 1, seed) {
|
|
8
22
|
const x = new Float64Array(n);
|
|
9
|
-
|
|
23
|
+
const next = seed === undefined ? Math.random : mulberry32(seed);
|
|
24
|
+
for (let i = 0; i < n; i++) x[i] = low + next() * (high - low);
|
|
10
25
|
return x;
|
|
11
26
|
}
|
|
12
27
|
|
|
13
|
-
export function randomTriangularFloat32Array(
|
|
28
|
+
export function randomTriangularFloat32Array(
|
|
29
|
+
n,
|
|
30
|
+
lda,
|
|
31
|
+
uplo = "lower",
|
|
32
|
+
low = -1,
|
|
33
|
+
high = 1,
|
|
34
|
+
diagLow = 5,
|
|
35
|
+
diagHigh = 15,
|
|
36
|
+
layout = "row-major",
|
|
37
|
+
) {
|
|
14
38
|
if (uplo !== "lower" && uplo !== "upper")
|
|
15
39
|
throw new Error("uplo must be 'lower' or 'upper'.");
|
|
40
|
+
if (layout !== "row-major" && layout !== "column-major")
|
|
41
|
+
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
16
42
|
if (lda < n) throw new Error("lda must be >= n.");
|
|
17
43
|
|
|
44
|
+
// Row-major stores row i at A[i*lda+j]; column-major stores column j at
|
|
45
|
+
// A[j*lda+i] instead. `uplo` describes the logical triangle (unaffected by
|
|
46
|
+
// storage order) — only which flat index each (i, j) maps to changes.
|
|
47
|
+
const isColMajor = layout === "column-major";
|
|
48
|
+
const idx = (i, j) => (isColMajor ? j * lda + i : i * lda + j);
|
|
49
|
+
|
|
18
50
|
const A = new Float32Array(n * lda);
|
|
19
51
|
for (let i = 0; i < n; i++) {
|
|
20
52
|
for (let j = 0; j < n; j++) {
|
|
21
53
|
if (i === j) continue;
|
|
22
54
|
const inTriangle = uplo === "lower" ? j < i : j > i;
|
|
23
|
-
if (inTriangle) A[i
|
|
55
|
+
if (inTriangle) A[idx(i, j)] = low + Math.random() * (high - low);
|
|
24
56
|
}
|
|
25
|
-
A[i
|
|
57
|
+
A[idx(i, i)] = diagLow + Math.random() * (diagHigh - diagLow);
|
|
26
58
|
}
|
|
27
59
|
return A;
|
|
28
60
|
}
|
package/src/sasum/sasum.d.mts
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
|
-
* Computes the sum of absolute values of a vector: result =
|
|
4
|
+
* Computes the sum of absolute values of a vector: $$\text{result} = \sum_{i} |x_i|$$
|
|
5
5
|
*
|
|
6
6
|
* {@includeCode ../../examples/sasum/sasum.js}
|
|
7
7
|
*
|
|
@@ -24,7 +24,7 @@ export declare function sasum(
|
|
|
24
24
|
): Promise<{ asum: number } | { asum: number; gpuTimeMs: number }>;
|
|
25
25
|
|
|
26
26
|
/**
|
|
27
|
-
* Computes the sum of absolute values of a vector: result =
|
|
27
|
+
* Computes the sum of absolute values of a vector: $$\text{result} = \sum_{i} |x_i|$$
|
|
28
28
|
*
|
|
29
29
|
* {@includeCode ../../examples/sasum/gpu.sasum.js}
|
|
30
30
|
*
|
package/src/sasum/sasum.mjs
CHANGED
|
@@ -12,14 +12,14 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
12
12
|
import { extractResult } from "../util/result.mjs";
|
|
13
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
14
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
15
|
-
|
|
16
|
-
|
|
15
|
+
import { WGS } from "../util/constants.mjs";
|
|
16
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
17
17
|
|
|
18
18
|
export async function sasum(device, n, x, incx) {
|
|
19
19
|
const xIsGpu = x instanceof GpuVector;
|
|
20
20
|
|
|
21
|
-
|
|
22
|
-
|
|
21
|
+
requireGpuDevice(device);
|
|
22
|
+
requireSameDevice(device, "sasum", { x });
|
|
23
23
|
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
24
24
|
throw new Error("n and incx must be integers.");
|
|
25
25
|
if (incx <= 0) throw new Error("incx must be positive.");
|
|
@@ -41,10 +41,11 @@ export async function sasum(device, n, x, incx) {
|
|
|
41
41
|
let readBuffer = null;
|
|
42
42
|
|
|
43
43
|
try {
|
|
44
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sasum-x", false);
|
|
45
|
-
partialsBuffer = createStorageBuffer(2 * WGS * 4, "sasum-partials"); // 2*WGS partial sums of f32
|
|
46
|
-
resultBuffer = createResultBuffer(4, "sasum-result"); // final f32 scalar
|
|
44
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sasum-x", false);
|
|
45
|
+
partialsBuffer = createStorageBuffer(device, 2 * WGS * 4, "sasum-partials"); // 2*WGS partial sums of f32
|
|
46
|
+
resultBuffer = createResultBuffer(device, 4, "sasum-result"); // final f32 scalar
|
|
47
47
|
paramsBuffer = createParamsBuffer(
|
|
48
|
+
device,
|
|
48
49
|
[
|
|
49
50
|
{ value: n, type: "u32" },
|
|
50
51
|
{ value: incx, type: "u32" },
|
|
@@ -52,31 +53,34 @@ export async function sasum(device, n, x, incx) {
|
|
|
52
53
|
"sasum-params",
|
|
53
54
|
);
|
|
54
55
|
|
|
55
|
-
const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
|
|
56
|
+
const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
|
|
56
57
|
xBuffer,
|
|
57
58
|
partialsBuffer,
|
|
58
59
|
paramsBuffer,
|
|
59
60
|
]);
|
|
60
61
|
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
62
|
+
device,
|
|
61
63
|
pipelineMain,
|
|
62
64
|
bgMain,
|
|
63
65
|
2 * WGS,
|
|
64
66
|
); // dispatch 2*WGS workgroups
|
|
65
67
|
|
|
66
|
-
submit(enc1);
|
|
68
|
+
submit(device, enc1);
|
|
67
69
|
|
|
68
|
-
const bgReduce = createBindGroup(
|
|
69
|
-
|
|
70
|
-
|
|
71
|
-
|
|
70
|
+
const bgReduce = createBindGroup(
|
|
71
|
+
device,
|
|
72
|
+
pipelineReduce.getBindGroupLayout(0),
|
|
73
|
+
[partialsBuffer, resultBuffer],
|
|
74
|
+
);
|
|
72
75
|
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
76
|
+
device,
|
|
73
77
|
pipelineReduce,
|
|
74
78
|
bgReduce,
|
|
75
79
|
1,
|
|
76
80
|
); // dispatch 1 workgroup to reduce the partial sums to a single result
|
|
77
|
-
readBuffer = stageReadback(enc2, resultBuffer);
|
|
81
|
+
readBuffer = stageReadback(device, enc2, resultBuffer);
|
|
78
82
|
|
|
79
|
-
submit(enc2);
|
|
83
|
+
submit(device, enc2);
|
|
80
84
|
|
|
81
85
|
const resultPromise = extractResult(readBuffer, Float32Array);
|
|
82
86
|
readBuffer = null; // ownership transferred — extractResult's own finally destroys it
|
|
@@ -96,7 +100,7 @@ export async function sasum(device, n, x, incx) {
|
|
|
96
100
|
if (partialsBuffer) destroyBuffers(partialsBuffer);
|
|
97
101
|
if (resultBuffer) destroyBuffers(resultBuffer);
|
|
98
102
|
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
99
|
-
// Only reached if submit(enc2) threw before ownership was transferred above.
|
|
103
|
+
// Only reached if submit(device, enc2) threw before ownership was transferred above.
|
|
100
104
|
if (readBuffer) destroyBuffers(readBuffer);
|
|
101
105
|
}
|
|
102
106
|
}
|
package/src/saxpy/saxpy.d.mts
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
|
-
* Performs the operation y
|
|
4
|
+
* Performs the operation $$y \leftarrow \alpha x + y$$
|
|
5
5
|
*
|
|
6
6
|
* {@includeCode ../../examples/saxpy/saxpy.js}
|
|
7
7
|
*
|
|
@@ -29,7 +29,7 @@ export declare function saxpy(
|
|
|
29
29
|
): Promise<{ y: Float32Array } | { y: Float32Array; gpuTimeMs: number }>;
|
|
30
30
|
|
|
31
31
|
/**
|
|
32
|
-
* Performs the operation y
|
|
32
|
+
* Performs the operation $$y \leftarrow \alpha x + y$$
|
|
33
33
|
*
|
|
34
34
|
* {@includeCode ../../examples/saxpy/gpu.saxpy.js}
|
|
35
35
|
*
|
package/src/saxpy/saxpy.mjs
CHANGED
|
@@ -11,21 +11,21 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function saxpy(device, n, alpha, x, incx, y, incy) {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
17
18
|
const yIsGpu = y instanceof GpuVector;
|
|
18
19
|
|
|
19
|
-
|
|
20
|
-
|
|
20
|
+
requireGpuDevice(device);
|
|
21
|
+
requireSameDevice(device, "saxpy", { x, y });
|
|
21
22
|
if (
|
|
22
23
|
!Number.isInteger(n) ||
|
|
23
24
|
!Number.isInteger(incx) ||
|
|
24
25
|
!Number.isInteger(incy)
|
|
25
26
|
)
|
|
26
27
|
throw new Error("n, incx, and incy must be integers.");
|
|
27
|
-
if (typeof alpha !== "number")
|
|
28
|
-
throw new Error("alpha must be a number.");
|
|
28
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
29
29
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
30
30
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
31
31
|
if (incx <= 0 || incy <= 0)
|
|
@@ -56,9 +56,10 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
|
|
|
56
56
|
let readBuffer = null;
|
|
57
57
|
|
|
58
58
|
try {
|
|
59
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "saxpy-x", false);
|
|
60
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "saxpy-y", true);
|
|
59
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "saxpy-x", false);
|
|
60
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "saxpy-y", true);
|
|
61
61
|
paramsBuffer = createParamsBuffer(
|
|
62
|
+
device,
|
|
62
63
|
[
|
|
63
64
|
{ value: n, type: "u32" },
|
|
64
65
|
{ value: alpha, type: "f32" },
|
|
@@ -68,23 +69,25 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
|
|
|
68
69
|
"saxpy-params",
|
|
69
70
|
);
|
|
70
71
|
|
|
71
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
72
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
72
73
|
xBuffer,
|
|
73
74
|
yBuffer,
|
|
74
75
|
paramsBuffer,
|
|
75
76
|
]);
|
|
76
77
|
const { commandEncoder, ts } = runComputePass(
|
|
78
|
+
device,
|
|
77
79
|
pipeline,
|
|
78
80
|
bindGroup,
|
|
79
|
-
calcWorkgroups(n),
|
|
81
|
+
calcWorkgroups(device, n),
|
|
80
82
|
);
|
|
81
|
-
readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
83
|
+
readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
82
84
|
|
|
83
|
-
submit(commandEncoder);
|
|
85
|
+
submit(device, commandEncoder);
|
|
84
86
|
|
|
85
87
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
86
88
|
|
|
87
|
-
if (yIsGpu
|
|
89
|
+
if (yIsGpu) {
|
|
90
|
+
// xIsGpu === yIsGpu, enforced above
|
|
88
91
|
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
89
92
|
return {};
|
|
90
93
|
}
|
package/src/scopy/scopy.d.mts
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
|
-
* Performs the operation y
|
|
4
|
+
* Performs the operation $$y \leftarrow x$$
|
|
5
5
|
*
|
|
6
6
|
* {@includeCode ../../examples/scopy/scopy.js}
|
|
7
7
|
*
|
|
@@ -27,7 +27,7 @@ export declare function scopy(
|
|
|
27
27
|
): Promise<{ y: Float32Array } | { y: Float32Array; gpuTimeMs: number }>;
|
|
28
28
|
|
|
29
29
|
/**
|
|
30
|
-
* Performs the operation y
|
|
30
|
+
* Performs the operation $$y \leftarrow x$$
|
|
31
31
|
*
|
|
32
32
|
* {@includeCode ../../examples/scopy/gpu.scopy.js}
|
|
33
33
|
*
|
package/src/scopy/scopy.mjs
CHANGED
|
@@ -11,13 +11,14 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function scopy(device, n, x, incx, y, incy) {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
17
18
|
const yIsGpu = y instanceof GpuVector;
|
|
18
19
|
|
|
19
|
-
|
|
20
|
-
|
|
20
|
+
requireGpuDevice(device);
|
|
21
|
+
requireSameDevice(device, "scopy", { x, y });
|
|
21
22
|
if (
|
|
22
23
|
!Number.isInteger(n) ||
|
|
23
24
|
!Number.isInteger(incx) ||
|
|
@@ -52,9 +53,10 @@ export async function scopy(device, n, x, incx, y, incy) {
|
|
|
52
53
|
let readBuffer = null;
|
|
53
54
|
|
|
54
55
|
try {
|
|
55
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "scopy-x", false);
|
|
56
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "scopy-y", true);
|
|
56
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "scopy-x", false);
|
|
57
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "scopy-y", true);
|
|
57
58
|
paramsBuffer = createParamsBuffer(
|
|
59
|
+
device,
|
|
58
60
|
[
|
|
59
61
|
{ value: n, type: "u32" },
|
|
60
62
|
{ value: incx, type: "u32" },
|
|
@@ -63,23 +65,25 @@ export async function scopy(device, n, x, incx, y, incy) {
|
|
|
63
65
|
"scopy-params",
|
|
64
66
|
);
|
|
65
67
|
|
|
66
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
68
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
67
69
|
xBuffer,
|
|
68
70
|
yBuffer,
|
|
69
71
|
paramsBuffer,
|
|
70
72
|
]);
|
|
71
73
|
const { commandEncoder, ts } = runComputePass(
|
|
74
|
+
device,
|
|
72
75
|
pipeline,
|
|
73
76
|
bindGroup,
|
|
74
|
-
calcWorkgroups(n),
|
|
77
|
+
calcWorkgroups(device, n),
|
|
75
78
|
);
|
|
76
|
-
readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
79
|
+
readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
77
80
|
|
|
78
|
-
submit(commandEncoder);
|
|
81
|
+
submit(device, commandEncoder);
|
|
79
82
|
|
|
80
83
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
81
84
|
|
|
82
|
-
if (yIsGpu
|
|
85
|
+
if (yIsGpu) {
|
|
86
|
+
// xIsGpu === yIsGpu, enforced above
|
|
83
87
|
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
84
88
|
return {};
|
|
85
89
|
}
|
package/src/sdot/sdot.d.mts
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
|
-
* Computes the dot product of two vectors: result =
|
|
4
|
+
* Computes the dot product of two vectors: $$\text{result} = \sum_{i} x_i y_i$$
|
|
5
5
|
*
|
|
6
6
|
* {@includeCode ../../examples/sdot/sdot.js}
|
|
7
7
|
*
|
|
@@ -28,7 +28,7 @@ export declare function sdot(
|
|
|
28
28
|
): Promise<{ dot: number } | { dot: number; gpuTimeMs: number }>;
|
|
29
29
|
|
|
30
30
|
/**
|
|
31
|
-
* Computes the dot product of two vectors: result =
|
|
31
|
+
* Computes the dot product of two vectors: $$\text{result} = \sum_{i} x_i y_i$$
|
|
32
32
|
*
|
|
33
33
|
* {@includeCode ../../examples/sdot/gpu.sdot.js}
|
|
34
34
|
*
|