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.
Files changed (124) hide show
  1. package/README.md +20 -18
  2. package/dist/wgblas.browser.js +2172 -1174
  3. package/index.d.mts +49 -44
  4. package/index.mjs +11 -0
  5. package/package.json +133 -63
  6. package/src/classes/Complex32.d.mts +43 -0
  7. package/src/classes/Complex32.mjs +82 -0
  8. package/src/classes/Complex64.d.mts +44 -0
  9. package/src/classes/Complex64.mjs +76 -0
  10. package/src/classes/GpuMatrix.d.mts +41 -41
  11. package/src/classes/GpuMatrix.mjs +126 -17
  12. package/src/classes/GpuVector.d.mts +36 -40
  13. package/src/classes/GpuVector.mjs +66 -11
  14. package/src/cscal/cscal.d.mts +47 -0
  15. package/src/cscal/cscal.mjs +98 -0
  16. package/src/dasum/dasum.d.mts +4 -4
  17. package/src/dasum/dasum.mjs +38 -20
  18. package/src/daxpy/daxpy.d.mts +56 -0
  19. package/src/daxpy/daxpy.mjs +150 -0
  20. package/src/dcopy/dcopy.d.mts +52 -0
  21. package/src/dcopy/dcopy.mjs +140 -0
  22. package/src/ddot/ddot.d.mts +62 -0
  23. package/src/ddot/ddot.mjs +184 -0
  24. package/src/devdocs.mjs +13 -0
  25. package/src/dnrm2/dnrm2.d.mts +50 -0
  26. package/src/dnrm2/dnrm2.mjs +189 -0
  27. package/src/drot/drot.d.mts +67 -0
  28. package/src/drot/drot.mjs +170 -0
  29. package/src/drotm/drotm.d.mts +67 -0
  30. package/src/drotm/drotm.mjs +171 -0
  31. package/src/dscal/dscal.d.mts +52 -0
  32. package/src/dscal/dscal.mjs +119 -0
  33. package/src/dswap/dswap.d.mts +57 -0
  34. package/src/dswap/dswap.mjs +155 -0
  35. package/src/idamax/idamax.d.mts +20 -2
  36. package/src/idamax/idamax.mjs +56 -24
  37. package/src/init.mjs +117 -56
  38. package/src/isamax/isamax.d.mts +20 -2
  39. package/src/isamax/isamax.mjs +21 -16
  40. package/src/random/random.d.mts +37 -39
  41. package/src/random/random.mjs +39 -7
  42. package/src/sasum/sasum.d.mts +2 -2
  43. package/src/sasum/sasum.mjs +20 -16
  44. package/src/saxpy/saxpy.d.mts +2 -2
  45. package/src/saxpy/saxpy.mjs +14 -11
  46. package/src/scopy/scopy.d.mts +2 -2
  47. package/src/scopy/scopy.mjs +13 -9
  48. package/src/sdot/sdot.d.mts +2 -2
  49. package/src/sdot/sdot.mjs +21 -17
  50. package/src/sgemm/sgemm.d.mts +2 -2
  51. package/src/sgemm/sgemm.mjs +109 -40
  52. package/src/sgemmtr/sgemmtr.d.mts +3 -2
  53. package/src/sgemmtr/sgemmtr.mjs +98 -40
  54. package/src/sgemv/sgemv.d.mts +2 -2
  55. package/src/sgemv/sgemv.mjs +69 -41
  56. package/src/sger/sger.d.mts +2 -2
  57. package/src/sger/sger.mjs +43 -19
  58. package/src/shaders/__test_pipeline_a.wgsl +3 -0
  59. package/src/shaders/__test_pipeline_b.wgsl +2 -0
  60. package/src/shaders/cscal.wgsl +33 -0
  61. package/src/shaders/daxpy.wgsl +66 -0
  62. package/src/shaders/dcopy.wgsl +34 -0
  63. package/src/shaders/ddot.wgsl +106 -0
  64. package/src/shaders/dnrm2.wgsl +167 -0
  65. package/src/shaders/drot.wgsl +81 -0
  66. package/src/shaders/drotm.wgsl +99 -0
  67. package/src/shaders/dscal.wgsl +60 -0
  68. package/src/shaders/dswap.wgsl +38 -0
  69. package/src/shaders/f64/utils/add.wgsl +6 -0
  70. package/src/shaders/f64/utils/divide.wgsl +45 -0
  71. package/src/shaders/f64/utils/multiply.wgsl +19 -10
  72. package/src/shaders/f64/utils/sqrt.wgsl +45 -0
  73. package/src/shaders/index.mjs +233 -14
  74. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  75. package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
  76. package/src/shaders/sgemm_large.wgsl +107 -18
  77. package/src/shaders/sgemm_small.wgsl +115 -15
  78. package/src/shaders/sgemmtr_large.wgsl +4 -1
  79. package/src/shaders/sgemmtr_small.wgsl +4 -1
  80. package/src/shaders/sgemv_n.wgsl +3 -1
  81. package/src/shaders/sgemv_t.wgsl +3 -1
  82. package/src/shaders/snrm2.wgsl +72 -23
  83. package/src/shaders/ssymv.wgsl +3 -1
  84. package/src/snrm2/snrm2.d.mts +2 -2
  85. package/src/snrm2/snrm2.mjs +41 -23
  86. package/src/srot/srot.d.mts +2 -4
  87. package/src/srot/srot.mjs +16 -11
  88. package/src/srotm/srotm.d.mts +2 -4
  89. package/src/srotm/srotm.mjs +17 -11
  90. package/src/sscal/sscal.d.mts +3 -3
  91. package/src/sscal/sscal.mjs +14 -12
  92. package/src/sswap/sswap.d.mts +2 -2
  93. package/src/sswap/sswap.mjs +18 -10
  94. package/src/ssymm/ssymm.d.mts +5 -4
  95. package/src/ssymm/ssymm.mjs +150 -54
  96. package/src/ssymv/ssymv.d.mts +2 -2
  97. package/src/ssymv/ssymv.mjs +47 -26
  98. package/src/ssyr/ssyr.d.mts +2 -2
  99. package/src/ssyr/ssyr.mjs +38 -17
  100. package/src/ssyr2/ssyr2.d.mts +2 -2
  101. package/src/ssyr2/ssyr2.mjs +48 -21
  102. package/src/ssyr2k/ssyr2k.d.mts +3 -2
  103. package/src/ssyr2k/ssyr2k.mjs +140 -62
  104. package/src/ssyrk/ssyrk.d.mts +3 -2
  105. package/src/ssyrk/ssyrk.mjs +91 -39
  106. package/src/strmm/strmm.d.mts +5 -4
  107. package/src/strmm/strmm.mjs +174 -60
  108. package/src/strmv/strmv.d.mts +2 -2
  109. package/src/strmv/strmv.mjs +42 -20
  110. package/src/strsm/strsm.d.mts +6 -4
  111. package/src/strsm/strsm.mjs +438 -174
  112. package/src/strsv/strsv.d.mts +5 -3
  113. package/src/strsv/strsv.mjs +89 -34
  114. package/src/util/benchmark.mjs +9 -9
  115. package/src/util/bindgroup.mjs +1 -3
  116. package/src/util/buffer.mjs +139 -24
  117. package/src/util/complex.mjs +87 -0
  118. package/src/util/compute.mjs +19 -16
  119. package/src/util/constants.mjs +57 -0
  120. package/src/util/device.mjs +49 -0
  121. package/src/util/pipeline.mjs +44 -10
  122. package/src/util/workgroup.mjs +72 -7
  123. package/src/shaders/browser-shaders.mjs +0 -81
  124. 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 or Float64Array matrix to GPU memory, row-major
35
- * or column-major. A Float64Array is split into a double-double (hi, lo)
36
- * f32 pair per element (WGSL has no f64 type) and stored across two GPU
37
- * buffers internally; `read()` reassembles doubles from these pairs. This
38
- * gives ~48 bits of mantissa (vs. 24 for a single f32) but less than true
39
- * f64 precision (52 bits), so round-tripped values are not always
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
- * @example
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
- * // Same logical matrix, column-major storage
63
- * const matCol = GpuMatrix.from(new Float32Array([1,4,2,5,3,6]), 2, 3, undefined, "column-major");
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 — a `Float32Array`,
71
- * or a `Float64Array` if this matrix was created from one. If `lda` exceeds
72
- * the dense minimum, the leading-dimension padding is stripped so 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
- * await init();
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
- * @example
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(buffer, rows, cols, lda, loBuffer = null, layout = "row-major") {
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 = 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(data, rows, cols, lda, layout = "row-major") {
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 (!(data instanceof Float32Array) && !(data instanceof Float64Array))
31
- throw new Error("GpuMatrix.from expects a Float32Array or Float64Array.");
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(`lda must be an integer >= ${isRowMajor ? "cols" : "rows"}.`);
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(hiBuf, rows, cols, lda, loBuf, layout);
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
- const buf = uploadBuffer(data.subarray(0, outerCount * lda), "gpu-matrix", true);
57
- return new GpuMatrix(buf, rows, cols, lda, null, layout);
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 = getDevice();
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 = isRowMajor ? this.cols : this.rows;
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(raw.subarray(r * this.lda, r * this.lda + innerLen), r * innerLen);
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(raw.subarray(r * this.lda, r * this.lda + innerLen), r * innerLen);
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 or Float64Array to GPU memory. A Float64Array is
20
- * split into a double-double (hi, lo) f32 pair per element (WGSL has no
21
- * f64 type) and stored across two GPU buffers internally; `read()`
22
- * reassembles doubles from these pairs. This gives ~48 bits of mantissa
23
- * (vs. 24 for a single f32) but less than true f64 precision (52 bits), so
24
- * round-tripped values are not always bit-exact with the original input.
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
- * @example
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
- * const dvec = GpuVector.from(new Float64Array([1.1, 2.2, 3.3]));
38
- * console.log("dtype:", dvec.dtype.name); // Float64Array
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 as a Float32Array, or a Float64Array if this vector
47
- * was created from one
58
+ * @returns vector data in the same shape it was created from — a
59
+ * Float32Array, Float64Array, Complex32Array, or Complex64Array
48
60
  *
49
- * @example
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
- * @example
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(buffer, length, dtype = Float32Array, loBuffer = null) {
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
- static from(data) {
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("GpuVector.from expects a Float32Array or Float64Array.");
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 = getDevice();
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