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
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
/**
|
|
2
|
+
* A single 64-bit (f64) complex number: a real and an imaginary component.
|
|
3
|
+
* Same shape as Complex32, with one deliberate difference: no f32 rounding
|
|
4
|
+
* — a plain JS number already is an f64, so re/im are kept at their full
|
|
5
|
+
* native double precision.
|
|
6
|
+
*
|
|
7
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/classes/Complex64.mjs#L11">Source code: Complex64.mjs (L11)</a>
|
|
8
|
+
* @category Classes
|
|
9
|
+
*/
|
|
10
|
+
export declare class Complex64 {
|
|
11
|
+
/**
|
|
12
|
+
* @param re - real component, full f64 precision
|
|
13
|
+
* @param im - imaginary component, full f64 precision
|
|
14
|
+
*
|
|
15
|
+
* {@includeCode ../../examples/complex64/complex64.js}
|
|
16
|
+
*/
|
|
17
|
+
constructor(re: number, im: number);
|
|
18
|
+
|
|
19
|
+
/** Real component (full f64 precision). */
|
|
20
|
+
re: number;
|
|
21
|
+
|
|
22
|
+
/** Imaginary component (full f64 precision). */
|
|
23
|
+
im: number;
|
|
24
|
+
}
|
|
25
|
+
|
|
26
|
+
/**
|
|
27
|
+
* An array of Complex64 values — array-of-structs, the f64 sibling of
|
|
28
|
+
* Complex32Array (same overloads, same interleaved-pairs convention),
|
|
29
|
+
* backed by full-precision Complex64 elements instead of f32-rounded
|
|
30
|
+
* Complex32 ones.
|
|
31
|
+
*
|
|
32
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/classes/Complex64.mjs#L34">Source code: Complex64.mjs (L34)</a>
|
|
33
|
+
* @category Classes
|
|
34
|
+
*/
|
|
35
|
+
export declare class Complex64Array extends Array<Complex64> {
|
|
36
|
+
/**
|
|
37
|
+
* @param arg - a length (fills with that many zero-valued Complex64
|
|
38
|
+
* entries), a flat interleaved `[re, im, re, im, ...]` list of numbers,
|
|
39
|
+
* or an iterable of existing Complex64 instances to copy
|
|
40
|
+
*
|
|
41
|
+
* {@includeCode ../../examples/complex64array/complex64array.js}
|
|
42
|
+
*/
|
|
43
|
+
constructor(arg?: number | Iterable<number> | Iterable<Complex64>);
|
|
44
|
+
}
|
|
@@ -0,0 +1,76 @@
|
|
|
1
|
+
/** @module devdocs/classes/Complex64 */
|
|
2
|
+
|
|
3
|
+
/**
|
|
4
|
+
* A single 64-bit (f64) complex number: a real and an imaginary component.
|
|
5
|
+
* Same shape as Complex32 (see Complex32.mjs) with one deliberate
|
|
6
|
+
* difference: no `Math.fround` rounding here. A plain JS number already is
|
|
7
|
+
* an f64, so re/im are kept at their full native double precision instead
|
|
8
|
+
* of being truncated down to f32.
|
|
9
|
+
*/
|
|
10
|
+
export class Complex64 {
|
|
11
|
+
/**
|
|
12
|
+
* @param {number} re - real component, full f64 precision
|
|
13
|
+
* @param {number} im - imaginary component, full f64 precision
|
|
14
|
+
*/
|
|
15
|
+
constructor(re, im) {
|
|
16
|
+
this.re = re;
|
|
17
|
+
this.im = im;
|
|
18
|
+
}
|
|
19
|
+
}
|
|
20
|
+
|
|
21
|
+
/**
|
|
22
|
+
* An array of Complex64 values — array-of-structs, the f64 sibling of
|
|
23
|
+
* Complex32Array (same overloads, same interleaved-pairs convention),
|
|
24
|
+
* backed by full-precision Complex64 elements instead of f32-rounded
|
|
25
|
+
* Complex32 ones.
|
|
26
|
+
*
|
|
27
|
+
* new Complex64Array() // empty
|
|
28
|
+
* new Complex64Array(length) // length zero-valued entries
|
|
29
|
+
* new Complex64Array([1, 5, 3, 8]) // flat [re, im, re, im, ...] pairs
|
|
30
|
+
* new Complex64Array([z1, z2]) // copies existing Complex64 instances
|
|
31
|
+
*/
|
|
32
|
+
export class Complex64Array extends Array {
|
|
33
|
+
/**
|
|
34
|
+
* @param {number|Iterable<number>|Iterable<Complex64>} [arg] - a length,
|
|
35
|
+
* a flat interleaved [re, im, ...] list of numbers, or an iterable of
|
|
36
|
+
* existing Complex64 instances
|
|
37
|
+
*/
|
|
38
|
+
constructor(arg) {
|
|
39
|
+
if (arg === undefined) {
|
|
40
|
+
super();
|
|
41
|
+
return;
|
|
42
|
+
}
|
|
43
|
+
if (typeof arg === "number") {
|
|
44
|
+
super(arg);
|
|
45
|
+
for (let i = 0; i < arg; i++) this[i] = new Complex64(0, 0);
|
|
46
|
+
return;
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
const items = Array.from(arg);
|
|
50
|
+
super();
|
|
51
|
+
if (items.length === 0) return;
|
|
52
|
+
|
|
53
|
+
if (items[0] instanceof Complex64) {
|
|
54
|
+
for (const z of items) {
|
|
55
|
+
if (!(z instanceof Complex64))
|
|
56
|
+
throw new Error(
|
|
57
|
+
"Complex64Array expects every element to be a Complex64.",
|
|
58
|
+
);
|
|
59
|
+
this.push(z);
|
|
60
|
+
}
|
|
61
|
+
return;
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
if (items.length % 2 !== 0)
|
|
65
|
+
throw new Error(
|
|
66
|
+
"Complex64Array expects an even number of interleaved [re, im, ...] values.",
|
|
67
|
+
);
|
|
68
|
+
for (let i = 0; i < items.length; i += 2) {
|
|
69
|
+
if (typeof items[i] !== "number" || typeof items[i + 1] !== "number")
|
|
70
|
+
throw new Error(
|
|
71
|
+
"Complex64Array expects interleaved [re, im, ...] values to be numbers.",
|
|
72
|
+
);
|
|
73
|
+
this.push(new Complex64(items[i], items[i + 1]));
|
|
74
|
+
}
|
|
75
|
+
}
|
|
76
|
+
}
|
|
@@ -1,3 +1,6 @@
|
|
|
1
|
+
import { Complex32Array } from "./Complex32.mjs";
|
|
2
|
+
import { Complex64Array } from "./Complex64.mjs";
|
|
3
|
+
|
|
1
4
|
/**
|
|
2
5
|
* Represents a Float32Array (or Float64Array) matrix stored in GPU memory,
|
|
3
6
|
* row-major or column-major.
|
|
@@ -30,72 +33,69 @@ export declare class GpuMatrix {
|
|
|
30
33
|
/** Storage layout this matrix was created with — every routine that accepts a GpuMatrix reads this automatically. */
|
|
31
34
|
readonly layout: 'row-major' | 'column-major';
|
|
32
35
|
|
|
36
|
+
/** Typed array (or complex array) constructor used when reading data back from the GPU. */
|
|
37
|
+
readonly dtype: Float32ArrayConstructor | Float64ArrayConstructor | typeof Complex32Array | typeof Complex64Array;
|
|
38
|
+
|
|
33
39
|
/**
|
|
34
|
-
* Uploads a Float32Array
|
|
35
|
-
* or column-major. A Float64Array is split
|
|
36
|
-
* f32 pair per element (WGSL has no f64
|
|
37
|
-
* buffers internally; `read()` reassembles
|
|
38
|
-
* gives ~48 bits of mantissa (vs. 24 for a
|
|
39
|
-
* f64 precision (52 bits), so
|
|
40
|
-
* bit-exact with the original input.
|
|
40
|
+
* Uploads a Float32Array, Float64Array, Complex32Array, or Complex64Array
|
|
41
|
+
* matrix to GPU memory, row-major or column-major. A Float64Array is split
|
|
42
|
+
* into a double-double (hi, lo) f32 pair per element (WGSL has no f64
|
|
43
|
+
* type) and stored across two GPU buffers internally; `read()` reassembles
|
|
44
|
+
* doubles from these pairs. This gives ~48 bits of mantissa (vs. 24 for a
|
|
45
|
+
* single f32) but less than true f64 precision (52 bits), so
|
|
46
|
+
* round-tripped values are not always bit-exact with the original input.
|
|
47
|
+
* A Complex32Array is stored interleaved (`[re0, im0, re1, im1, ...]`) in
|
|
48
|
+
* one buffer; a Complex64Array gets the same double-double split applied
|
|
49
|
+
* independently to its real and imaginary components.
|
|
41
50
|
*
|
|
42
51
|
* `rows`/`cols` always describe the logical shape regardless of layout.
|
|
43
52
|
* `lda` defaults to `cols` (row-major) or `rows` (column-major) — dense, no
|
|
44
53
|
* padding. `data` must have at least `rows * lda` (row-major) or
|
|
45
54
|
* `cols * lda` (column-major) elements.
|
|
46
55
|
*
|
|
56
|
+
* Omitting the device falls back to the one from the last {@link init} call
|
|
57
|
+
* — the historical form, and fine for a single-GPU program. Pass a device
|
|
58
|
+
* explicitly (matching every routine's own `(device, ...)` convention)
|
|
59
|
+
* when driving more than one GPU at once, since a GpuMatrix is bound for
|
|
60
|
+
* life to whichever device created it.
|
|
61
|
+
*
|
|
47
62
|
* @param data - matrix data, in the order matching `layout`
|
|
48
63
|
* @param rows - number of rows
|
|
49
64
|
* @param cols - number of columns
|
|
50
65
|
* @param lda - leading dimension (default: `cols` for row-major, `rows` for column-major)
|
|
51
66
|
* @param layout - storage layout (default: `'row-major'`)
|
|
52
67
|
*
|
|
53
|
-
* @
|
|
54
|
-
* ```js
|
|
55
|
-
* import { init, GpuMatrix } from "wgblas";
|
|
56
|
-
*
|
|
57
|
-
* await init();
|
|
58
|
-
* // 2×3 matrix: [[1,2,3],[4,5,6]]
|
|
59
|
-
* const mat = GpuMatrix.from(new Float32Array([1,2,3,4,5,6]), 2, 3);
|
|
60
|
-
* console.log(mat.rows, mat.cols, mat.lda); // 2 3 3
|
|
68
|
+
* {@includeCode ../../examples/gpumatrix-from/gpumatrix-from.js}
|
|
61
69
|
*
|
|
62
|
-
*
|
|
63
|
-
*
|
|
64
|
-
* ```
|
|
70
|
+
* **Explicit device (multi-GPU):**
|
|
71
|
+
* {@includeCode ../../examples/gpumatrix-from-device/gpumatrix-from-device.js}
|
|
65
72
|
*/
|
|
66
|
-
static from(data: Float32Array | Float64Array, rows: number, cols: number, lda?: number, layout?: 'row-major' | 'column-major'): GpuMatrix;
|
|
73
|
+
static from(data: Float32Array | Float64Array | Complex32Array | Complex64Array, rows: number, cols: number, lda?: number, layout?: 'row-major' | 'column-major'): GpuMatrix;
|
|
74
|
+
/**
|
|
75
|
+
* @param device - GPUDevice from `init()` — the matrix is bound to this device for life
|
|
76
|
+
* @param data - matrix data, in the order matching `layout`
|
|
77
|
+
* @param rows - number of rows
|
|
78
|
+
* @param cols - number of columns
|
|
79
|
+
* @param lda - leading dimension (default: `cols` for row-major, `rows` for column-major)
|
|
80
|
+
* @param layout - storage layout (default: `'row-major'`)
|
|
81
|
+
*/
|
|
82
|
+
static from(device: GPUDevice, data: Float32Array | Float64Array | Complex32Array | Complex64Array, rows: number, cols: number, lda?: number, layout?: 'row-major' | 'column-major'): GpuMatrix;
|
|
67
83
|
|
|
68
84
|
/**
|
|
69
85
|
* Downloads the matrix from GPU memory and returns a dense array of shape
|
|
70
|
-
* `rows × cols`, in the same layout it was created with
|
|
71
|
-
*
|
|
72
|
-
* the
|
|
73
|
-
* returned array is always tightly packed.
|
|
74
|
-
*
|
|
75
|
-
* @example
|
|
76
|
-
* ```js
|
|
77
|
-
* import { init, GpuMatrix } from "wgblas";
|
|
86
|
+
* `rows × cols`, in the same layout and type it was created with. If `lda`
|
|
87
|
+
* exceeds the dense minimum, the leading-dimension padding is stripped so
|
|
88
|
+
* the returned array is always tightly packed.
|
|
78
89
|
*
|
|
79
|
-
*
|
|
80
|
-
* const mat = GpuMatrix.from(new Float32Array([1,2,3,4,5,6]), 2, 3);
|
|
81
|
-
* const data = await mat.read();
|
|
82
|
-
* console.log(data); // Float32Array [1, 2, 3, 4, 5, 6]
|
|
83
|
-
* ```
|
|
90
|
+
* {@includeCode ../../examples/gpumatrix-read/gpumatrix-read.js}
|
|
84
91
|
*/
|
|
85
|
-
read(): Promise<Float32Array | Float64Array>;
|
|
92
|
+
read(): Promise<Float32Array | Float64Array | Complex32Array | Complex64Array>;
|
|
86
93
|
|
|
87
94
|
/**
|
|
88
95
|
* Destroys the underlying GPU buffer. Call when the matrix is no longer
|
|
89
96
|
* needed to free GPU memory.
|
|
90
97
|
*
|
|
91
|
-
* @
|
|
92
|
-
* ```js
|
|
93
|
-
* import { init, GpuMatrix } from "wgblas";
|
|
94
|
-
*
|
|
95
|
-
* await init();
|
|
96
|
-
* const mat = GpuMatrix.from(new Float32Array([1,2,3,4,5,6]), 2, 3);
|
|
97
|
-
* mat.destroy();
|
|
98
|
-
* ```
|
|
98
|
+
* {@includeCode ../../examples/gpumatrix-destroy/gpumatrix-destroy.js}
|
|
99
99
|
*/
|
|
100
100
|
destroy(): void;
|
|
101
101
|
}
|
|
@@ -2,15 +2,34 @@ import { getDevice } from "../init.mjs";
|
|
|
2
2
|
import { uploadBuffer, stageReadback } from "../util/buffer.mjs";
|
|
3
3
|
import { extractResult } from "../util/result.mjs";
|
|
4
4
|
import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
|
|
5
|
+
import {
|
|
6
|
+
interleaveComplex32,
|
|
7
|
+
splitComplex64,
|
|
8
|
+
mergeComplex64,
|
|
9
|
+
} from "../util/complex.mjs";
|
|
10
|
+
import { Complex32Array } from "./Complex32.mjs";
|
|
11
|
+
import { Complex64Array } from "./Complex64.mjs";
|
|
5
12
|
|
|
6
13
|
export class GpuMatrix {
|
|
7
|
-
constructor(
|
|
14
|
+
constructor(
|
|
15
|
+
buffer,
|
|
16
|
+
rows,
|
|
17
|
+
cols,
|
|
18
|
+
lda,
|
|
19
|
+
loBuffer = null,
|
|
20
|
+
layout = "row-major",
|
|
21
|
+
device = null,
|
|
22
|
+
dtype = Float32Array,
|
|
23
|
+
) {
|
|
8
24
|
this._buf = buffer;
|
|
9
25
|
this._loBuf = loBuffer; // Non-null only for Float64Array-backed matrices
|
|
10
26
|
this.rows = rows;
|
|
11
27
|
this.cols = cols;
|
|
12
|
-
this.lda
|
|
28
|
+
this.lda = lda;
|
|
13
29
|
this.layout = layout;
|
|
30
|
+
this.dtype = dtype; // disambiguates Complex32Array from Float32Array — both have _loBuf === null
|
|
31
|
+
// See GpuVector: a GPUBuffer is bound to one device for life.
|
|
32
|
+
this.device = device ?? getDevice();
|
|
14
33
|
}
|
|
15
34
|
|
|
16
35
|
/**
|
|
@@ -21,21 +40,35 @@ export class GpuMatrix {
|
|
|
21
40
|
* no padding). `data` must have at least `rows * lda` (row-major) or
|
|
22
41
|
* `cols * lda` (column-major) elements.
|
|
23
42
|
*/
|
|
24
|
-
static from(
|
|
43
|
+
static from(deviceOrData, ...rest) {
|
|
44
|
+
const explicit = deviceOrData instanceof GPUDevice;
|
|
45
|
+
const device = explicit ? deviceOrData : getDevice();
|
|
46
|
+
const data = explicit ? rest.shift() : deviceOrData;
|
|
47
|
+
let [rows, cols, lda, layout = "row-major"] = rest;
|
|
48
|
+
|
|
25
49
|
if (layout !== "row-major" && layout !== "column-major")
|
|
26
50
|
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
27
51
|
const isRowMajor = layout === "row-major";
|
|
28
52
|
if (lda === undefined) lda = isRowMajor ? cols : rows;
|
|
29
53
|
|
|
30
|
-
if (
|
|
31
|
-
|
|
54
|
+
if (
|
|
55
|
+
!(data instanceof Float32Array) &&
|
|
56
|
+
!(data instanceof Float64Array) &&
|
|
57
|
+
!(data instanceof Complex32Array) &&
|
|
58
|
+
!(data instanceof Complex64Array)
|
|
59
|
+
)
|
|
60
|
+
throw new Error(
|
|
61
|
+
"GpuMatrix.from expects a Float32Array, Float64Array, Complex32Array, or Complex64Array.",
|
|
62
|
+
);
|
|
32
63
|
if (!Number.isInteger(rows) || rows <= 0)
|
|
33
64
|
throw new Error("rows must be a positive integer.");
|
|
34
65
|
if (!Number.isInteger(cols) || cols <= 0)
|
|
35
66
|
throw new Error("cols must be a positive integer.");
|
|
36
67
|
const minLda = isRowMajor ? cols : rows;
|
|
37
68
|
if (!Number.isInteger(lda) || lda < minLda)
|
|
38
|
-
throw new Error(
|
|
69
|
+
throw new Error(
|
|
70
|
+
`lda must be an integer >= ${isRowMajor ? "cols" : "rows"}.`,
|
|
71
|
+
);
|
|
39
72
|
|
|
40
73
|
// Row-major: `rows` chunks of length `lda` (only the first `cols` of each used).
|
|
41
74
|
// Column-major: `cols` chunks of length `lda` (only the first `rows` of each used).
|
|
@@ -48,39 +81,112 @@ export class GpuMatrix {
|
|
|
48
81
|
if (data instanceof Float64Array) {
|
|
49
82
|
const n = outerCount * lda;
|
|
50
83
|
const { hi, lo } = splitDoubleDouble(data.subarray(0, n));
|
|
51
|
-
const hiBuf = uploadBuffer(hi, "gpu-matrix-f64-hi", true);
|
|
52
|
-
const loBuf = uploadBuffer(lo, "gpu-matrix-f64-lo", true);
|
|
53
|
-
return new GpuMatrix(
|
|
84
|
+
const hiBuf = uploadBuffer(device, hi, "gpu-matrix-f64-hi", true);
|
|
85
|
+
const loBuf = uploadBuffer(device, lo, "gpu-matrix-f64-lo", true);
|
|
86
|
+
return new GpuMatrix(
|
|
87
|
+
hiBuf,
|
|
88
|
+
rows,
|
|
89
|
+
cols,
|
|
90
|
+
lda,
|
|
91
|
+
loBuf,
|
|
92
|
+
layout,
|
|
93
|
+
device,
|
|
94
|
+
Float64Array,
|
|
95
|
+
);
|
|
54
96
|
}
|
|
55
97
|
|
|
56
|
-
|
|
57
|
-
|
|
98
|
+
if (data instanceof Complex32Array) {
|
|
99
|
+
const buf = uploadBuffer(
|
|
100
|
+
device,
|
|
101
|
+
interleaveComplex32(data, outerCount * lda),
|
|
102
|
+
"gpu-matrix-complex32",
|
|
103
|
+
true,
|
|
104
|
+
);
|
|
105
|
+
return new GpuMatrix(
|
|
106
|
+
buf,
|
|
107
|
+
rows,
|
|
108
|
+
cols,
|
|
109
|
+
lda,
|
|
110
|
+
null,
|
|
111
|
+
layout,
|
|
112
|
+
device,
|
|
113
|
+
Complex32Array,
|
|
114
|
+
);
|
|
115
|
+
}
|
|
116
|
+
|
|
117
|
+
if (data instanceof Complex64Array) {
|
|
118
|
+
const { hi, lo } = splitComplex64(data, outerCount * lda);
|
|
119
|
+
const hiBuf = uploadBuffer(device, hi, "gpu-matrix-complex64-hi", true);
|
|
120
|
+
const loBuf = uploadBuffer(device, lo, "gpu-matrix-complex64-lo", true);
|
|
121
|
+
return new GpuMatrix(
|
|
122
|
+
hiBuf,
|
|
123
|
+
rows,
|
|
124
|
+
cols,
|
|
125
|
+
lda,
|
|
126
|
+
loBuf,
|
|
127
|
+
layout,
|
|
128
|
+
device,
|
|
129
|
+
Complex64Array,
|
|
130
|
+
);
|
|
131
|
+
}
|
|
132
|
+
|
|
133
|
+
const buf = uploadBuffer(
|
|
134
|
+
device,
|
|
135
|
+
data.subarray(0, outerCount * lda),
|
|
136
|
+
"gpu-matrix",
|
|
137
|
+
true,
|
|
138
|
+
);
|
|
139
|
+
return new GpuMatrix(buf, rows, cols, lda, null, layout, device);
|
|
58
140
|
}
|
|
59
141
|
|
|
60
142
|
async read() {
|
|
61
|
-
const device =
|
|
143
|
+
const device = this.device;
|
|
62
144
|
const enc = device.createCommandEncoder();
|
|
63
|
-
const rb = stageReadback(enc, this._buf);
|
|
145
|
+
const rb = stageReadback(device, enc, this._buf);
|
|
64
146
|
device.queue.submit([enc.finish()]);
|
|
65
147
|
|
|
66
148
|
const isRowMajor = this.layout !== "column-major";
|
|
67
149
|
const outerCount = isRowMajor ? this.rows : this.cols;
|
|
68
|
-
const innerLen
|
|
150
|
+
const innerLen = isRowMajor ? this.cols : this.rows;
|
|
151
|
+
|
|
152
|
+
if (this.dtype === Complex32Array) {
|
|
153
|
+
const raw = new Complex32Array(await extractResult(rb, Float32Array));
|
|
154
|
+
if (this.lda === innerLen) return raw;
|
|
155
|
+
const out = new Complex32Array(outerCount * innerLen);
|
|
156
|
+
for (let r = 0; r < outerCount; r++)
|
|
157
|
+
for (let c = 0; c < innerLen; c++)
|
|
158
|
+
out[r * innerLen + c] = raw[r * this.lda + c];
|
|
159
|
+
return out;
|
|
160
|
+
}
|
|
69
161
|
|
|
70
162
|
if (this._loBuf) {
|
|
71
163
|
const encLo = device.createCommandEncoder();
|
|
72
|
-
const rbLo = stageReadback(encLo, this._loBuf);
|
|
164
|
+
const rbLo = stageReadback(device, encLo, this._loBuf);
|
|
73
165
|
device.queue.submit([encLo.finish()]);
|
|
74
166
|
|
|
75
167
|
const [hi, lo] = await Promise.all([
|
|
76
168
|
extractResult(rb, Float32Array),
|
|
77
169
|
extractResult(rbLo, Float32Array),
|
|
78
170
|
]);
|
|
171
|
+
|
|
172
|
+
if (this.dtype === Complex64Array) {
|
|
173
|
+
const raw = mergeComplex64(hi, lo);
|
|
174
|
+
if (this.lda === innerLen) return raw;
|
|
175
|
+
const out = new Complex64Array(outerCount * innerLen);
|
|
176
|
+
for (let r = 0; r < outerCount; r++)
|
|
177
|
+
for (let c = 0; c < innerLen; c++)
|
|
178
|
+
out[r * innerLen + c] = raw[r * this.lda + c];
|
|
179
|
+
return out;
|
|
180
|
+
}
|
|
181
|
+
|
|
79
182
|
const raw = mergeDoubleDouble(hi, lo);
|
|
80
183
|
if (this.lda === innerLen) return raw;
|
|
81
184
|
const out = new Float64Array(outerCount * innerLen);
|
|
82
185
|
for (let r = 0; r < outerCount; r++)
|
|
83
|
-
out.set(
|
|
186
|
+
out.set(
|
|
187
|
+
raw.subarray(r * this.lda, r * this.lda + innerLen),
|
|
188
|
+
r * innerLen,
|
|
189
|
+
);
|
|
84
190
|
return out;
|
|
85
191
|
}
|
|
86
192
|
|
|
@@ -88,7 +194,10 @@ export class GpuMatrix {
|
|
|
88
194
|
if (this.lda === innerLen) return raw;
|
|
89
195
|
const out = new Float32Array(outerCount * innerLen);
|
|
90
196
|
for (let r = 0; r < outerCount; r++)
|
|
91
|
-
out.set(
|
|
197
|
+
out.set(
|
|
198
|
+
raw.subarray(r * this.lda, r * this.lda + innerLen),
|
|
199
|
+
r * innerLen,
|
|
200
|
+
);
|
|
92
201
|
return out;
|
|
93
202
|
}
|
|
94
203
|
|
|
@@ -1,3 +1,6 @@
|
|
|
1
|
+
import { Complex32Array } from "./Complex32.mjs";
|
|
2
|
+
import { Complex64Array } from "./Complex64.mjs";
|
|
3
|
+
|
|
1
4
|
/**
|
|
2
5
|
* Represents a Float32Array stored in GPU memory.
|
|
3
6
|
*
|
|
@@ -12,65 +15,58 @@ export declare class GpuVector {
|
|
|
12
15
|
/** Number of elements in the vector. */
|
|
13
16
|
readonly length: number;
|
|
14
17
|
|
|
15
|
-
/** Typed array constructor used when reading data back from the GPU. */
|
|
16
|
-
readonly dtype: Float32ArrayConstructor | Float64ArrayConstructor;
|
|
18
|
+
/** Typed array (or complex array) constructor used when reading data back from the GPU. */
|
|
19
|
+
readonly dtype: Float32ArrayConstructor | Float64ArrayConstructor | typeof Complex32Array | typeof Complex64Array;
|
|
17
20
|
|
|
18
21
|
/**
|
|
19
|
-
* Uploads a Float32Array
|
|
20
|
-
* split into a double-double (hi, lo) f32
|
|
21
|
-
* f64 type) and stored across two GPU
|
|
22
|
-
* reassembles doubles from these pairs. This
|
|
23
|
-
* (vs. 24 for a single f32) but less than true
|
|
24
|
-
* round-tripped values are not always
|
|
22
|
+
* Uploads a Float32Array, Float64Array, Complex32Array, or Complex64Array
|
|
23
|
+
* to GPU memory. A Float64Array is split into a double-double (hi, lo) f32
|
|
24
|
+
* pair per element (WGSL has no f64 type) and stored across two GPU
|
|
25
|
+
* buffers internally; `read()` reassembles doubles from these pairs. This
|
|
26
|
+
* gives ~48 bits of mantissa (vs. 24 for a single f32) but less than true
|
|
27
|
+
* f64 precision (52 bits), so round-tripped values are not always
|
|
28
|
+
* bit-exact with the original input. A Complex32Array is stored
|
|
29
|
+
* interleaved (`[re0, im0, re1, im1, ...]`) in one buffer; a
|
|
30
|
+
* Complex64Array gets the same double-double split applied independently
|
|
31
|
+
* to its real and imaginary components, interleaved per (hi, lo) channel.
|
|
32
|
+
*
|
|
33
|
+
* Omitting the device falls back to the one from the last {@link init} call
|
|
34
|
+
* — the historical form, and fine for a single-GPU program. Pass a device
|
|
35
|
+
* explicitly (matching every routine's own `(device, ...)` convention)
|
|
36
|
+
* when driving more than one GPU at once, since a GpuVector is bound for
|
|
37
|
+
* life to whichever device created it.
|
|
25
38
|
*
|
|
26
39
|
* @param data - input vector data
|
|
27
40
|
* @returns GpuVector backed by a GPU buffer
|
|
28
41
|
*
|
|
29
|
-
* @
|
|
30
|
-
* ```js
|
|
31
|
-
* import { init, GpuVector } from "wgblas";
|
|
32
|
-
*
|
|
33
|
-
* await init();
|
|
34
|
-
* const vec = GpuVector.from(new Float32Array([1, 2, 3, 4]));
|
|
35
|
-
* console.log("length:", vec.length, "dtype:", vec.dtype.name);
|
|
42
|
+
* {@includeCode ../../examples/gpuvector-from/gpuvector-from.js}
|
|
36
43
|
*
|
|
37
|
-
*
|
|
38
|
-
*
|
|
39
|
-
* ```
|
|
44
|
+
* **Explicit device (multi-GPU):**
|
|
45
|
+
* {@includeCode ../../examples/gpuvector-from-device/gpuvector-from-device.js}
|
|
40
46
|
*/
|
|
41
|
-
static from(data: Float32Array | Float64Array): GpuVector;
|
|
47
|
+
static from(data: Float32Array | Float64Array | Complex32Array | Complex64Array): GpuVector;
|
|
48
|
+
/**
|
|
49
|
+
* @param device - GPUDevice from `init()` — the vector is bound to this device for life
|
|
50
|
+
* @param data - input vector data
|
|
51
|
+
* @returns GpuVector backed by a GPU buffer
|
|
52
|
+
*/
|
|
53
|
+
static from(device: GPUDevice, data: Float32Array | Float64Array | Complex32Array | Complex64Array): GpuVector;
|
|
42
54
|
|
|
43
55
|
/**
|
|
44
56
|
* Reads the vector data back from GPU memory.
|
|
45
57
|
*
|
|
46
|
-
* @returns vector data
|
|
47
|
-
*
|
|
58
|
+
* @returns vector data in the same shape it was created from — a
|
|
59
|
+
* Float32Array, Float64Array, Complex32Array, or Complex64Array
|
|
48
60
|
*
|
|
49
|
-
* @
|
|
50
|
-
* ```js
|
|
51
|
-
* import { init, GpuVector } from "wgblas";
|
|
52
|
-
*
|
|
53
|
-
* await init();
|
|
54
|
-
* const vec = GpuVector.from(new Float32Array([1, 2, 3, 4]));
|
|
55
|
-
* const data = await vec.read();
|
|
56
|
-
* console.log(data);
|
|
57
|
-
* ```
|
|
61
|
+
* {@includeCode ../../examples/gpuvector-read/gpuvector-read.js}
|
|
58
62
|
*/
|
|
59
|
-
read(): Promise<Float32Array | Float64Array>;
|
|
63
|
+
read(): Promise<Float32Array | Float64Array | Complex32Array | Complex64Array>;
|
|
60
64
|
|
|
61
65
|
/**
|
|
62
66
|
* Destroys the underlying GPU buffer. Call when the vector is no longer needed
|
|
63
67
|
* to free GPU memory — especially important in long-running programs.
|
|
64
68
|
*
|
|
65
|
-
* @
|
|
66
|
-
* ```js
|
|
67
|
-
* import { init, GpuVector } from "wgblas";
|
|
68
|
-
*
|
|
69
|
-
* await init();
|
|
70
|
-
* const vec = GpuVector.from(new Float32Array([1, 2, 3, 4]));
|
|
71
|
-
* vec.destroy();
|
|
72
|
-
* console.log("GPU buffer released");
|
|
73
|
-
* ```
|
|
69
|
+
* {@includeCode ../../examples/gpuvector-destroy/gpuvector-destroy.js}
|
|
74
70
|
*/
|
|
75
71
|
destroy(): void;
|
|
76
72
|
}
|
|
@@ -2,45 +2,100 @@ import { getDevice } from "../init.mjs";
|
|
|
2
2
|
import { uploadBuffer, stageReadback } from "../util/buffer.mjs";
|
|
3
3
|
import { extractResult } from "../util/result.mjs";
|
|
4
4
|
import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
|
|
5
|
+
import {
|
|
6
|
+
interleaveComplex32,
|
|
7
|
+
splitComplex64,
|
|
8
|
+
mergeComplex64,
|
|
9
|
+
} from "../util/complex.mjs";
|
|
10
|
+
import { Complex32Array } from "./Complex32.mjs";
|
|
11
|
+
import { Complex64Array } from "./Complex64.mjs";
|
|
5
12
|
|
|
6
13
|
export class GpuVector {
|
|
7
|
-
constructor(
|
|
14
|
+
constructor(
|
|
15
|
+
buffer,
|
|
16
|
+
length,
|
|
17
|
+
dtype = Float32Array,
|
|
18
|
+
loBuffer = null,
|
|
19
|
+
device = null,
|
|
20
|
+
) {
|
|
8
21
|
this._buf = buffer;
|
|
9
22
|
this._loBuf = loBuffer; // Non-null only for Float64Array-backed vectors
|
|
10
23
|
this.length = length;
|
|
11
24
|
this.dtype = dtype;
|
|
25
|
+
// A GPUBuffer belongs to exactly one device and WebGPU rejects any attempt
|
|
26
|
+
// to use it with another, so every handle remembers where it lives. Routines
|
|
27
|
+
// check this to reject mixed-device operands with a clear message instead of
|
|
28
|
+
// a raw GPUValidationError.
|
|
29
|
+
this.device = device ?? getDevice();
|
|
12
30
|
}
|
|
13
31
|
|
|
14
|
-
|
|
32
|
+
/**
|
|
33
|
+
* Uploads a vector to GPU memory.
|
|
34
|
+
*
|
|
35
|
+
* Pass the target `GPUDevice` first — matching every routine's own
|
|
36
|
+
* `(device, ...)` convention. Omitting it falls back to the device from the
|
|
37
|
+
* last `init()`, which is the historical form and only works single-device.
|
|
38
|
+
*
|
|
39
|
+
* @param {GPUDevice|Float32Array|Float64Array} deviceOrData
|
|
40
|
+
*/
|
|
41
|
+
static from(deviceOrData, maybeData) {
|
|
42
|
+
const explicit = deviceOrData instanceof GPUDevice;
|
|
43
|
+
const device = explicit ? deviceOrData : getDevice();
|
|
44
|
+
const data = explicit ? maybeData : deviceOrData;
|
|
45
|
+
|
|
15
46
|
if (data instanceof Float64Array) {
|
|
16
47
|
const { hi, lo } = splitDoubleDouble(data);
|
|
17
|
-
const hiBuf = uploadBuffer(hi, "gpu-vector-f64-hi", true);
|
|
18
|
-
const loBuf = uploadBuffer(lo, "gpu-vector-f64-lo", true);
|
|
19
|
-
return new GpuVector(hiBuf, data.length, Float64Array, loBuf);
|
|
48
|
+
const hiBuf = uploadBuffer(device, hi, "gpu-vector-f64-hi", true);
|
|
49
|
+
const loBuf = uploadBuffer(device, lo, "gpu-vector-f64-lo", true);
|
|
50
|
+
return new GpuVector(hiBuf, data.length, Float64Array, loBuf, device);
|
|
51
|
+
}
|
|
52
|
+
if (data instanceof Complex32Array) {
|
|
53
|
+
const buf = uploadBuffer(
|
|
54
|
+
device,
|
|
55
|
+
interleaveComplex32(data),
|
|
56
|
+
"gpu-vector-complex32",
|
|
57
|
+
true,
|
|
58
|
+
);
|
|
59
|
+
return new GpuVector(buf, data.length, Complex32Array, null, device);
|
|
60
|
+
}
|
|
61
|
+
if (data instanceof Complex64Array) {
|
|
62
|
+
const { hi, lo } = splitComplex64(data);
|
|
63
|
+
const hiBuf = uploadBuffer(device, hi, "gpu-vector-complex64-hi", true);
|
|
64
|
+
const loBuf = uploadBuffer(device, lo, "gpu-vector-complex64-lo", true);
|
|
65
|
+
return new GpuVector(hiBuf, data.length, Complex64Array, loBuf, device);
|
|
20
66
|
}
|
|
21
67
|
if (!(data instanceof Float32Array)) {
|
|
22
|
-
throw new Error(
|
|
68
|
+
throw new Error(
|
|
69
|
+
"GpuVector.from expects a Float32Array, Float64Array, Complex32Array, or Complex64Array.",
|
|
70
|
+
);
|
|
23
71
|
}
|
|
24
|
-
const buf = uploadBuffer(data, "gpu-vector", true);
|
|
25
|
-
return new GpuVector(buf, data.length, data.constructor);
|
|
72
|
+
const buf = uploadBuffer(device, data, "gpu-vector", true);
|
|
73
|
+
return new GpuVector(buf, data.length, data.constructor, null, device);
|
|
26
74
|
}
|
|
27
75
|
|
|
28
76
|
async read() {
|
|
29
|
-
const device =
|
|
77
|
+
const device = this.device;
|
|
30
78
|
const enc = device.createCommandEncoder();
|
|
31
|
-
const rb = stageReadback(enc, this._buf);
|
|
79
|
+
const rb = stageReadback(device, enc, this._buf);
|
|
32
80
|
device.queue.submit([enc.finish()]);
|
|
33
81
|
|
|
82
|
+
if (this.dtype === Complex32Array) {
|
|
83
|
+
// Complex32Array's own interleaved-numbers constructor overload does
|
|
84
|
+
// the de-interleaving — see complex.mjs.
|
|
85
|
+
return new Complex32Array(await extractResult(rb, Float32Array));
|
|
86
|
+
}
|
|
87
|
+
|
|
34
88
|
if (!this._loBuf) return extractResult(rb, this.dtype);
|
|
35
89
|
|
|
36
90
|
const encLo = device.createCommandEncoder();
|
|
37
|
-
const rbLo = stageReadback(encLo, this._loBuf);
|
|
91
|
+
const rbLo = stageReadback(device, encLo, this._loBuf);
|
|
38
92
|
device.queue.submit([encLo.finish()]);
|
|
39
93
|
|
|
40
94
|
const [hi, lo] = await Promise.all([
|
|
41
95
|
extractResult(rb, Float32Array),
|
|
42
96
|
extractResult(rbLo, Float32Array),
|
|
43
97
|
]);
|
|
98
|
+
if (this.dtype === Complex64Array) return mergeComplex64(hi, lo);
|
|
44
99
|
return mergeDoubleDouble(hi, lo);
|
|
45
100
|
}
|
|
46
101
|
|