wgblas 2.1.0 → 2.2.1

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 (99) hide show
  1. package/README.md +2 -0
  2. package/dist/wgblas.browser.js +1016 -43
  3. package/index.d.mts +26 -53
  4. package/index.mjs +11 -0
  5. package/package.json +132 -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 +112 -10
  12. package/src/classes/GpuVector.d.mts +36 -40
  13. package/src/classes/GpuVector.mjs +39 -2
  14. package/src/cscal/cscal.d.mts +47 -0
  15. package/src/cscal/cscal.mjs +98 -0
  16. package/src/dasum/dasum.mjs +31 -15
  17. package/src/daxpy/daxpy.d.mts +56 -0
  18. package/src/daxpy/daxpy.mjs +150 -0
  19. package/src/dcopy/dcopy.d.mts +52 -0
  20. package/src/dcopy/dcopy.mjs +140 -0
  21. package/src/ddot/ddot.d.mts +62 -0
  22. package/src/ddot/ddot.mjs +184 -0
  23. package/src/dnrm2/dnrm2.d.mts +50 -0
  24. package/src/dnrm2/dnrm2.mjs +189 -0
  25. package/src/drot/drot.d.mts +67 -0
  26. package/src/drot/drot.mjs +170 -0
  27. package/src/drotm/drotm.d.mts +67 -0
  28. package/src/drotm/drotm.mjs +171 -0
  29. package/src/dscal/dscal.d.mts +52 -0
  30. package/src/dscal/dscal.mjs +119 -0
  31. package/src/dswap/dswap.d.mts +57 -0
  32. package/src/dswap/dswap.mjs +155 -0
  33. package/src/idamax/idamax.mjs +49 -19
  34. package/src/init.mjs +6 -3
  35. package/src/isamax/isamax.mjs +17 -14
  36. package/src/random/random.d.mts +37 -40
  37. package/src/random/random.mjs +39 -7
  38. package/src/sasum/sasum.mjs +13 -11
  39. package/src/saxpy/saxpy.mjs +9 -8
  40. package/src/scopy/scopy.mjs +8 -6
  41. package/src/sdot/sdot.mjs +13 -11
  42. package/src/sgemm/sgemm.d.mts +2 -2
  43. package/src/sgemm/sgemm.mjs +91 -35
  44. package/src/sgemmtr/sgemmtr.d.mts +3 -2
  45. package/src/sgemmtr/sgemmtr.mjs +92 -35
  46. package/src/sgemv/sgemv.d.mts +2 -2
  47. package/src/sgemv/sgemv.mjs +41 -25
  48. package/src/sger/sger.d.mts +2 -2
  49. package/src/sger/sger.mjs +38 -16
  50. package/src/shaders/__test_pipeline_a.wgsl +3 -0
  51. package/src/shaders/__test_pipeline_b.wgsl +2 -0
  52. package/src/shaders/cscal.wgsl +33 -0
  53. package/src/shaders/daxpy.wgsl +66 -0
  54. package/src/shaders/dcopy.wgsl +34 -0
  55. package/src/shaders/ddot.wgsl +106 -0
  56. package/src/shaders/dnrm2.wgsl +167 -0
  57. package/src/shaders/drot.wgsl +81 -0
  58. package/src/shaders/drotm.wgsl +99 -0
  59. package/src/shaders/dscal.wgsl +60 -0
  60. package/src/shaders/dswap.wgsl +38 -0
  61. package/src/shaders/f64/utils/add.wgsl +6 -0
  62. package/src/shaders/f64/utils/divide.wgsl +45 -0
  63. package/src/shaders/f64/utils/multiply.wgsl +29 -11
  64. package/src/shaders/f64/utils/sqrt.wgsl +45 -0
  65. package/src/shaders/index.mjs +69 -0
  66. package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
  67. package/src/snrm2/snrm2.mjs +20 -14
  68. package/src/srot/srot.mjs +10 -7
  69. package/src/srotm/srotm.mjs +9 -12
  70. package/src/sscal/sscal.mjs +7 -7
  71. package/src/sswap/sswap.mjs +14 -8
  72. package/src/ssymm/ssymm.d.mts +5 -4
  73. package/src/ssymm/ssymm.mjs +142 -55
  74. package/src/ssymv/ssymv.d.mts +2 -2
  75. package/src/ssymv/ssymv.mjs +42 -23
  76. package/src/ssyr/ssyr.d.mts +2 -2
  77. package/src/ssyr/ssyr.mjs +34 -15
  78. package/src/ssyr2/ssyr2.d.mts +2 -2
  79. package/src/ssyr2/ssyr2.mjs +43 -18
  80. package/src/ssyr2k/ssyr2k.d.mts +3 -2
  81. package/src/ssyr2k/ssyr2k.mjs +132 -55
  82. package/src/ssyrk/ssyrk.d.mts +3 -2
  83. package/src/ssyrk/ssyrk.mjs +84 -33
  84. package/src/strmm/strmm.d.mts +5 -4
  85. package/src/strmm/strmm.mjs +153 -54
  86. package/src/strmv/strmv.d.mts +2 -2
  87. package/src/strmv/strmv.mjs +37 -17
  88. package/src/strsm/strsm.d.mts +6 -4
  89. package/src/strsm/strsm.mjs +418 -172
  90. package/src/strsv/strsv.d.mts +5 -3
  91. package/src/strsv/strsv.mjs +82 -31
  92. package/src/util/benchmark.mjs +5 -3
  93. package/src/util/buffer.mjs +33 -12
  94. package/src/util/complex.mjs +87 -0
  95. package/src/util/compute.mjs +14 -8
  96. package/src/util/device.mjs +18 -3
  97. package/src/util/pipeline.mjs +40 -5
  98. package/src/util/workgroup.mjs +23 -6
  99. package/src/shaders/f64add.wgsl +0 -281
@@ -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,32 @@ 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", device = null) {
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
14
31
  // See GpuVector: a GPUBuffer is bound to one device for life.
15
32
  this.device = device ?? getDevice();
16
33
  }
@@ -34,15 +51,24 @@ export class GpuMatrix {
34
51
  const isRowMajor = layout === "row-major";
35
52
  if (lda === undefined) lda = isRowMajor ? cols : rows;
36
53
 
37
- if (!(data instanceof Float32Array) && !(data instanceof Float64Array))
38
- 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
+ );
39
63
  if (!Number.isInteger(rows) || rows <= 0)
40
64
  throw new Error("rows must be a positive integer.");
41
65
  if (!Number.isInteger(cols) || cols <= 0)
42
66
  throw new Error("cols must be a positive integer.");
43
67
  const minLda = isRowMajor ? cols : rows;
44
68
  if (!Number.isInteger(lda) || lda < minLda)
45
- 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
+ );
46
72
 
47
73
  // Row-major: `rows` chunks of length `lda` (only the first `cols` of each used).
48
74
  // Column-major: `cols` chunks of length `lda` (only the first `rows` of each used).
@@ -57,10 +83,59 @@ export class GpuMatrix {
57
83
  const { hi, lo } = splitDoubleDouble(data.subarray(0, n));
58
84
  const hiBuf = uploadBuffer(device, hi, "gpu-matrix-f64-hi", true);
59
85
  const loBuf = uploadBuffer(device, lo, "gpu-matrix-f64-lo", true);
60
- return new GpuMatrix(hiBuf, rows, cols, lda, loBuf, layout, device);
86
+ return new GpuMatrix(
87
+ hiBuf,
88
+ rows,
89
+ cols,
90
+ lda,
91
+ loBuf,
92
+ layout,
93
+ device,
94
+ Float64Array,
95
+ );
96
+ }
97
+
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
+ );
61
131
  }
62
132
 
63
- const buf = uploadBuffer(device, data.subarray(0, outerCount * lda), "gpu-matrix", true);
133
+ const buf = uploadBuffer(
134
+ device,
135
+ data.subarray(0, outerCount * lda),
136
+ "gpu-matrix",
137
+ true,
138
+ );
64
139
  return new GpuMatrix(buf, rows, cols, lda, null, layout, device);
65
140
  }
66
141
 
@@ -72,7 +147,17 @@ export class GpuMatrix {
72
147
 
73
148
  const isRowMajor = this.layout !== "column-major";
74
149
  const outerCount = isRowMajor ? this.rows : this.cols;
75
- 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
+ }
76
161
 
77
162
  if (this._loBuf) {
78
163
  const encLo = device.createCommandEncoder();
@@ -83,11 +168,25 @@ export class GpuMatrix {
83
168
  extractResult(rb, Float32Array),
84
169
  extractResult(rbLo, Float32Array),
85
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
+
86
182
  const raw = mergeDoubleDouble(hi, lo);
87
183
  if (this.lda === innerLen) return raw;
88
184
  const out = new Float64Array(outerCount * innerLen);
89
185
  for (let r = 0; r < outerCount; r++)
90
- 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
+ );
91
190
  return out;
92
191
  }
93
192
 
@@ -95,7 +194,10 @@ export class GpuMatrix {
95
194
  if (this.lda === innerLen) return raw;
96
195
  const out = new Float32Array(outerCount * innerLen);
97
196
  for (let r = 0; r < outerCount; r++)
98
- 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
+ );
99
201
  return out;
100
202
  }
101
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,9 +2,22 @@ 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, device = 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;
@@ -36,8 +49,25 @@ export class GpuVector {
36
49
  const loBuf = uploadBuffer(device, lo, "gpu-vector-f64-lo", true);
37
50
  return new GpuVector(hiBuf, data.length, Float64Array, loBuf, device);
38
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);
66
+ }
39
67
  if (!(data instanceof Float32Array)) {
40
- 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
+ );
41
71
  }
42
72
  const buf = uploadBuffer(device, data, "gpu-vector", true);
43
73
  return new GpuVector(buf, data.length, data.constructor, null, device);
@@ -49,6 +79,12 @@ export class GpuVector {
49
79
  const rb = stageReadback(device, enc, this._buf);
50
80
  device.queue.submit([enc.finish()]);
51
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
+
52
88
  if (!this._loBuf) return extractResult(rb, this.dtype);
53
89
 
54
90
  const encLo = device.createCommandEncoder();
@@ -59,6 +95,7 @@ export class GpuVector {
59
95
  extractResult(rb, Float32Array),
60
96
  extractResult(rbLo, Float32Array),
61
97
  ]);
98
+ if (this.dtype === Complex64Array) return mergeComplex64(hi, lo);
62
99
  return mergeDoubleDouble(hi, lo);
63
100
  }
64
101
 
@@ -0,0 +1,47 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+ import { Complex32, Complex32Array } from "../classes/Complex32.mjs";
3
+
4
+ /**
5
+ * Scales a complex vector by a complex constant: $$x \leftarrow \alpha x$$
6
+ *
7
+ * {@includeCode ../../examples/cscal/cscal.js}
8
+ *
9
+ * **Browser (standalone HTML):**
10
+ * {@includeCode ../../examples/cscal/web/cscal.html}
11
+ *
12
+ * @param device - GPUDevice from `init()`
13
+ * @param n - number of elements to scale (must be a positive integer)
14
+ * @param alpha - complex scalar multiplier
15
+ * @param x - Complex32Array input/output vector
16
+ * @param incx - stride for x (must be a positive integer)
17
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/cscal/cscal.mjs">Source code: cscal.mjs</a>
18
+ * @category BLAS Level 1
19
+ */
20
+ export declare function cscal(
21
+ device: GPUDevice,
22
+ n: number,
23
+ alpha: Complex32,
24
+ x: Complex32Array,
25
+ incx: number,
26
+ ): Promise<{ x: Complex32Array } | { x: Complex32Array; gpuTimeMs: number }>;
27
+
28
+ /**
29
+ * Scales a complex vector by a complex constant: $$x \leftarrow \alpha x$$
30
+ *
31
+ * {@includeCode ../../examples/cscal/gpu.cscal.js}
32
+ *
33
+ * @param device - GPUDevice from `init()`
34
+ * @param n - number of elements to scale (must be a positive integer)
35
+ * @param alpha - complex scalar multiplier
36
+ * @param x - Complex32Array-backed GpuVector input/output vector (mutated in place)
37
+ * @param incx - stride for x (must be a positive integer)
38
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/cscal/cscal.mjs">Source code: cscal.mjs</a>
39
+ * @category BLAS Level 1
40
+ */
41
+ export declare function cscal(
42
+ device: GPUDevice,
43
+ n: number,
44
+ alpha: Complex32,
45
+ x: GpuVector,
46
+ incx: number,
47
+ ): Promise<{} | { gpuTimeMs: number }>;