wgblas 1.2.1 → 2.1.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 (96) hide show
  1. package/LICENSE +1 -1
  2. package/README.md +54 -72
  3. package/dist/wgblas.browser.js +1637 -858
  4. package/index.d.mts +51 -6
  5. package/index.mjs +8 -0
  6. package/package.json +56 -2
  7. package/src/classes/GpuMatrix.mjs +17 -10
  8. package/src/classes/GpuVector.mjs +28 -10
  9. package/src/dasum/dasum.d.mts +6 -6
  10. package/src/dasum/dasum.mjs +25 -21
  11. package/src/devdocs.mjs +13 -0
  12. package/src/idamax/idamax.d.mts +69 -0
  13. package/src/idamax/idamax.mjs +130 -0
  14. package/src/init.mjs +115 -49
  15. package/src/isamax/isamax.d.mts +21 -3
  16. package/src/isamax/isamax.mjs +16 -14
  17. package/src/random/random.d.mts +1 -0
  18. package/src/sasum/sasum.d.mts +3 -3
  19. package/src/sasum/sasum.mjs +15 -13
  20. package/src/saxpy/saxpy.d.mts +3 -3
  21. package/src/saxpy/saxpy.mjs +10 -8
  22. package/src/scopy/scopy.d.mts +3 -3
  23. package/src/scopy/scopy.mjs +10 -8
  24. package/src/sdot/sdot.d.mts +3 -3
  25. package/src/sdot/sdot.mjs +16 -14
  26. package/src/sgemm/sgemm.d.mts +102 -0
  27. package/src/sgemm/sgemm.mjs +208 -0
  28. package/src/sgemmtr/sgemmtr.d.mts +104 -0
  29. package/src/sgemmtr/sgemmtr.mjs +204 -0
  30. package/src/sgemv/sgemv.d.mts +1 -38
  31. package/src/sgemv/sgemv.mjs +42 -26
  32. package/src/sger/sger.d.mts +1 -34
  33. package/src/sger/sger.mjs +12 -8
  34. package/src/shaders/block_transfer.wgsl +42 -0
  35. package/src/shaders/dasum.wgsl +3 -2
  36. package/src/shaders/f64/dekker.wgsl +4 -85
  37. package/src/shaders/f64/utils/abs.wgsl +10 -0
  38. package/src/shaders/f64/utils/add.wgsl +77 -0
  39. package/src/shaders/f64/utils/equal.wgsl +7 -0
  40. package/src/shaders/f64/utils/greater.wgsl +12 -0
  41. package/src/shaders/f64/utils/multiply.wgsl +81 -0
  42. package/src/shaders/idamax.wgsl +96 -0
  43. package/src/shaders/index.mjs +164 -14
  44. package/src/shaders/reduction/argmaxF64.wgsl +50 -0
  45. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  46. package/src/shaders/reduction/sumF64.wgsl +2 -2
  47. package/src/shaders/sgemm_large.wgsl +206 -0
  48. package/src/shaders/sgemm_small.wgsl +212 -0
  49. package/src/shaders/sgemmtr_large.wgsl +120 -0
  50. package/src/shaders/sgemmtr_small.wgsl +113 -0
  51. package/src/shaders/sgemv_n.wgsl +3 -1
  52. package/src/shaders/sgemv_t.wgsl +3 -1
  53. package/src/shaders/snrm2.wgsl +72 -23
  54. package/src/shaders/ssymv.wgsl +3 -1
  55. package/src/shaders/symmetrize.wgsl +31 -0
  56. package/src/shaders/triangularize.wgsl +44 -0
  57. package/src/snrm2/snrm2.d.mts +3 -3
  58. package/src/snrm2/snrm2.mjs +33 -21
  59. package/src/srot/srot.d.mts +3 -5
  60. package/src/srot/srot.mjs +11 -9
  61. package/src/srotm/srotm.d.mts +3 -5
  62. package/src/srotm/srotm.mjs +19 -10
  63. package/src/sscal/sscal.d.mts +4 -4
  64. package/src/sscal/sscal.mjs +12 -10
  65. package/src/sswap/sswap.d.mts +3 -3
  66. package/src/sswap/sswap.mjs +11 -9
  67. package/src/ssymm/ssymm.d.mts +103 -0
  68. package/src/ssymm/ssymm.mjs +218 -0
  69. package/src/ssymv/ssymv.d.mts +1 -36
  70. package/src/ssymv/ssymv.mjs +12 -8
  71. package/src/ssyr/ssyr.d.mts +1 -30
  72. package/src/ssyr/ssyr.mjs +11 -7
  73. package/src/ssyr2/ssyr2.d.mts +1 -34
  74. package/src/ssyr2/ssyr2.mjs +12 -8
  75. package/src/ssyr2k/ssyr2k.d.mts +100 -0
  76. package/src/ssyr2k/ssyr2k.mjs +202 -0
  77. package/src/ssyrk/ssyrk.d.mts +90 -0
  78. package/src/ssyrk/ssyrk.mjs +177 -0
  79. package/src/strmm/strmm.d.mts +100 -0
  80. package/src/strmm/strmm.mjs +226 -0
  81. package/src/strmv/strmv.d.mts +1 -36
  82. package/src/strmv/strmv.mjs +12 -8
  83. package/src/strsm/strsm.d.mts +99 -0
  84. package/src/strsm/strsm.mjs +360 -0
  85. package/src/strsv/strsv.d.mts +1 -32
  86. package/src/strsv/strsv.mjs +18 -12
  87. package/src/util/benchmark.mjs +4 -6
  88. package/src/util/bindgroup.mjs +1 -3
  89. package/src/util/buffer.mjs +116 -20
  90. package/src/util/compute.mjs +12 -12
  91. package/src/util/constants.mjs +57 -0
  92. package/src/util/device.mjs +34 -0
  93. package/src/util/f64.mjs +3 -3
  94. package/src/util/pipeline.mjs +5 -6
  95. package/src/util/workgroup.mjs +55 -7
  96. package/src/shaders/browser-shaders.mjs +0 -55
package/index.d.mts CHANGED
@@ -17,6 +17,7 @@ export { sasum } from "./src/sasum/sasum.mjs";
17
17
  export { dasum } from "./src/dasum/dasum.mjs";
18
18
  export { snrm2 } from "./src/snrm2/snrm2.mjs";
19
19
  export { isamax } from "./src/isamax/isamax.mjs";
20
+ export { idamax } from "./src/idamax/idamax.mjs";
20
21
  export { srot } from "./src/srot/srot.mjs";
21
22
  export { srotm } from "./src/srotm/srotm.mjs";
22
23
  export { sgemv } from "./src/sgemv/sgemv.mjs";
@@ -26,15 +27,35 @@ export { strsv } from "./src/strsv/strsv.mjs";
26
27
  export { sger } from "./src/sger/sger.mjs";
27
28
  export { ssyr } from "./src/ssyr/ssyr.mjs";
28
29
  export { ssyr2 } from "./src/ssyr2/ssyr2.mjs";
30
+ export { sgemm } from "./src/sgemm/sgemm.mjs";
31
+ export { sgemmtr } from "./src/sgemmtr/sgemmtr.mjs";
32
+ export { ssyrk } from "./src/ssyrk/ssyrk.mjs";
33
+ export { ssyr2k } from "./src/ssyr2k/ssyr2k.mjs";
34
+ export { ssymm } from "./src/ssymm/ssymm.mjs";
35
+ export { strmm } from "./src/strmm/strmm.mjs";
36
+ export { strsm } from "./src/strsm/strsm.mjs";
29
37
 
30
38
  /**
31
39
  * Initializes the WebGPU device.
32
40
  *
41
+ * Devices are cached per option set: the same options return the same device,
42
+ * different options a separate one — so one process can drive several GPUs
43
+ * (`"high-performance"` and `"low-power"` typically resolve to the discrete and
44
+ * integrated adapters). Each routine dispatches to the device you pass it; GPU
45
+ * buffers cannot cross devices, so a `GpuVector`/`GpuMatrix` from one is
46
+ * rejected by another. {@link cleanup} releases them all.
47
+ *
33
48
  * @param options.powerPreference - GPU power preference (default: `"high-performance"`).
34
49
  * This is a hint to the browser: on dual-GPU systems, `"high-performance"` typically favors the discrete GPU
35
50
  * and `"low-power"` favors the integrated one.
36
51
  * See [MDN: GPU.requestAdapter()](https://developer.mozilla.org/en-US/docs/Web/API/GPU/requestAdapter).
37
- * @param options.benchmark - enable GPU timestamp queries; BLAS functions return `{ result, gpuTimeMs }` (default: `false`)
52
+ * @param options.benchmark - enable GPU timestamp queries; BLAS functions then also return `gpuTimeMs`
53
+ * alongside their normal result (e.g. sscal returns `{ x, gpuTimeMs }`, saxpy returns `{ y, gpuTimeMs }`) —
54
+ * see each routine's own docs for its exact return shape (default: `false`)
55
+ * @param options.dumpShaders - Node-only. Forwards Dawn's `dump_shaders` debug toggle, printing
56
+ * each pipeline's WGSL and compiled backend IR (SPIR-V/Vulkan, MSL/Metal, or HLSL/D3D12,
57
+ * whichever Dawn picked) to stderr as it compiles. A Dawn passthrough, not a wgblas format —
58
+ * no effect in the browser, which gives pages no API to request compiled shader IR (default: `false`)
38
59
  *
39
60
  * @example Default (high-performance GPU)
40
61
  * ```js
@@ -59,21 +80,42 @@ export { ssyr2 } from "./src/ssyr2/ssyr2.mjs";
59
80
  * const n = 5;
60
81
  * const alpha = 2.0;
61
82
  * const x = new Float32Array([1, 2, 3, 4, 5]);
62
- * const { result, gpuTimeMs } = await sscal(device, n, alpha, x, 1);
83
+ * const { x: result, gpuTimeMs } = await sscal(device, n, alpha, x, 1);
63
84
  * console.log(`Result: [${Array.from(result).join(", ")}]`);
64
85
  * console.log(`GPU time: ${gpuTimeMs.toFixed(3)} ms`);
65
86
  * ```
87
+ * @example Two GPUs at once
88
+ * ```js
89
+ * import { init, cleanup, gpuName, sscal } from "wgblas";
90
+ * const dGpu = await init({ powerPreference: "high-performance" });
91
+ * const iGpu = await init({ powerPreference: "low-power" });
92
+ * console.log(gpuName(dGpu).description, "and", gpuName(iGpu).description);
93
+ * const [a, b] = await Promise.all([
94
+ * sscal(dGpu, 4, 2, new Float32Array([1, 2, 3, 4]), 1),
95
+ * sscal(iGpu, 4, 5, new Float32Array([1, 2, 3, 4]), 1),
96
+ * ]);
97
+ * cleanup(); // releases both
98
+ * ```
99
+ *
66
100
  * @see [Source code: init.mjs](https://github.com/manit2004/wgblas/blob/main/src/init.mjs#L18-L54)
67
101
  * @category Core
68
102
  */
69
103
  export declare function init(options?: {
70
104
  powerPreference?: GPUPowerPreference;
71
105
  benchmark?: boolean;
106
+ dumpShaders?: boolean;
72
107
  }): Promise<GPUDevice>;
73
108
 
74
109
  /**
75
- * Destroys the WebGPU device, releases the adapter, resets benchmark state, and fires all internal
76
- * cleanup callbacks (e.g. releasing cached GPU pipelines and buffers). Call when done (required in Node.js to prevent crash on exit).
110
+ * Destroys devices created by {@link init} and releases their cached pipelines and buffers.
111
+ * Call when done (required in Node.js to prevent crash on exit).
112
+ *
113
+ * With no argument, releases every device at once. Pass a device to release
114
+ * just that one and leave the others usable — handy when driving several GPUs.
115
+ * Unknown or already-released devices are ignored, so this is safe to call
116
+ * more than once.
117
+ *
118
+ * @param device - the device to release; omit to release all of them.
77
119
  *
78
120
  * @example
79
121
  * ```js
@@ -86,11 +128,14 @@ export declare function init(options?: {
86
128
  * @see [Source code: init.mjs](https://github.com/manit2004/wgblas/blob/main/src/init.mjs#L56-L65)
87
129
  * @category Core
88
130
  */
89
- export declare function cleanup(): void;
131
+ export declare function cleanup(device?: GPUDevice): void;
90
132
 
91
133
  /**
92
134
  * Returns the GPU device name from the WebGPU adapter info. Must be called after `init()`.
93
135
  *
136
+ * @param device - which device to report on; defaults to the one from the first
137
+ * `init()` call. Pass it explicitly when driving more than one GPU.
138
+ *
94
139
  * @example
95
140
  * ```js
96
141
  * import { init, gpuName } from "wgblas";
@@ -101,4 +146,4 @@ export declare function cleanup(): void;
101
146
  * @see [Source code: init.mjs](https://github.com/manit2004/wgblas/blob/main/src/init.mjs#L81-L87)
102
147
  * @category Core
103
148
  */
104
- export declare function gpuName(): { description: string; device: string };
149
+ export declare function gpuName(device?: GPUDevice): { description: string; device: string };
package/index.mjs CHANGED
@@ -15,6 +15,7 @@ export { sasum } from "./src/sasum/sasum.mjs";
15
15
  export { dasum } from "./src/dasum/dasum.mjs";
16
16
  export { snrm2 } from "./src/snrm2/snrm2.mjs";
17
17
  export { isamax } from "./src/isamax/isamax.mjs";
18
+ export { idamax } from "./src/idamax/idamax.mjs";
18
19
  export { srot } from "./src/srot/srot.mjs";
19
20
  export { srotm } from "./src/srotm/srotm.mjs";
20
21
  export { sgemv } from "./src/sgemv/sgemv.mjs";
@@ -24,3 +25,10 @@ export { strsv } from "./src/strsv/strsv.mjs";
24
25
  export { sger } from "./src/sger/sger.mjs";
25
26
  export { ssyr } from "./src/ssyr/ssyr.mjs";
26
27
  export { ssyr2 } from "./src/ssyr2/ssyr2.mjs";
28
+ export { sgemm } from "./src/sgemm/sgemm.mjs";
29
+ export { sgemmtr } from "./src/sgemmtr/sgemmtr.mjs";
30
+ export { ssyrk } from "./src/ssyrk/ssyrk.mjs";
31
+ export { ssyr2k } from "./src/ssyr2k/ssyr2k.mjs";
32
+ export { ssymm } from "./src/ssymm/ssymm.mjs";
33
+ export { strmm } from "./src/strmm/strmm.mjs";
34
+ export { strsm } from "./src/strsm/strsm.mjs";
package/package.json CHANGED
@@ -1,6 +1,6 @@
1
1
  {
2
2
  "name": "wgblas",
3
- "version": "1.2.1",
3
+ "version": "2.1.0",
4
4
  "description": "BLAS on WebGPU",
5
5
  "type": "module",
6
6
  "main": "index.mjs",
@@ -59,6 +59,10 @@
59
59
  "import": "./src/isamax/isamax.mjs",
60
60
  "types": "./src/isamax/isamax.d.mts"
61
61
  },
62
+ "./idamax": {
63
+ "import": "./src/idamax/idamax.mjs",
64
+ "types": "./src/idamax/idamax.d.mts"
65
+ },
62
66
  "./srot": {
63
67
  "import": "./src/srot/srot.mjs",
64
68
  "types": "./src/srot/srot.d.mts"
@@ -94,6 +98,34 @@
94
98
  "./ssyr2": {
95
99
  "import": "./src/ssyr2/ssyr2.mjs",
96
100
  "types": "./src/ssyr2/ssyr2.d.mts"
101
+ },
102
+ "./sgemm": {
103
+ "import": "./src/sgemm/sgemm.mjs",
104
+ "types": "./src/sgemm/sgemm.d.mts"
105
+ },
106
+ "./sgemmtr": {
107
+ "import": "./src/sgemmtr/sgemmtr.mjs",
108
+ "types": "./src/sgemmtr/sgemmtr.d.mts"
109
+ },
110
+ "./ssyrk": {
111
+ "import": "./src/ssyrk/ssyrk.mjs",
112
+ "types": "./src/ssyrk/ssyrk.d.mts"
113
+ },
114
+ "./ssyr2k": {
115
+ "import": "./src/ssyr2k/ssyr2k.mjs",
116
+ "types": "./src/ssyr2k/ssyr2k.d.mts"
117
+ },
118
+ "./ssymm": {
119
+ "import": "./src/ssymm/ssymm.mjs",
120
+ "types": "./src/ssymm/ssymm.d.mts"
121
+ },
122
+ "./strmm": {
123
+ "import": "./src/strmm/strmm.mjs",
124
+ "types": "./src/strmm/strmm.d.mts"
125
+ },
126
+ "./strsm": {
127
+ "import": "./src/strsm/strsm.mjs",
128
+ "types": "./src/strsm/strsm.d.mts"
97
129
  }
98
130
  },
99
131
  "files": [
@@ -110,7 +142,26 @@
110
142
  },
111
143
  "keywords": [
112
144
  "BLAS",
113
- "WebGPU"
145
+ "WebGPU",
146
+ "linear-algebra",
147
+ "linear",
148
+ "algebra",
149
+ "subroutines",
150
+ "level 1",
151
+ "level 2",
152
+ "level 3",
153
+ "gpu",
154
+ "gpgpu",
155
+ "wgsl",
156
+ "compute-shader",
157
+ "matrix",
158
+ "vector",
159
+ "matmul",
160
+ "dot product",
161
+ "float32",
162
+ "float64",
163
+ "numerical-computing",
164
+ "scientific-computing"
114
165
  ],
115
166
  "author": "Manit Roy",
116
167
  "license": "Apache-2.0",
@@ -129,11 +180,13 @@
129
180
  "@commitlint/config-conventional": "^21.2.0",
130
181
  "@eslint/js": "^10.0.1",
131
182
  "@stdlib/blas-base-dasum": "^0.4.1",
183
+ "@stdlib/blas-base-idamax": "^0.1.1",
132
184
  "@stdlib/blas-base-isamax": "^0.1.1",
133
185
  "@stdlib/blas-base-sasum": "^0.3.1",
134
186
  "@stdlib/blas-base-saxpy": "^0.3.1",
135
187
  "@stdlib/blas-base-scopy": "^0.3.1",
136
188
  "@stdlib/blas-base-sdot": "^0.3.1",
189
+ "@stdlib/blas-base-sgemm": "^0.1.1",
137
190
  "@stdlib/blas-base-sgemv": "^0.1.1",
138
191
  "@stdlib/blas-base-sger": "^0.1.1",
139
192
  "@stdlib/blas-base-snrm2": "^0.3.1",
@@ -156,6 +209,7 @@
156
209
  "lint-staged": "^17.0.8",
157
210
  "prettier": "^3.9.4",
158
211
  "typedoc": "^0.28.19",
212
+ "typedoc-plugin-katex": "^0.1.2",
159
213
  "typescript": "^6.0.3",
160
214
  "vite": "^8.0.16"
161
215
  },
@@ -4,13 +4,15 @@ import { extractResult } from "../util/result.mjs";
4
4
  import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
5
5
 
6
6
  export class GpuMatrix {
7
- constructor(buffer, rows, cols, lda, loBuffer = null, layout = "row-major") {
7
+ constructor(buffer, rows, cols, lda, loBuffer = null, layout = "row-major", device = null) {
8
8
  this._buf = buffer;
9
9
  this._loBuf = loBuffer; // Non-null only for Float64Array-backed matrices
10
10
  this.rows = rows;
11
11
  this.cols = cols;
12
12
  this.lda = lda;
13
13
  this.layout = layout;
14
+ // See GpuVector: a GPUBuffer is bound to one device for life.
15
+ this.device = device ?? getDevice();
14
16
  }
15
17
 
16
18
  /**
@@ -21,7 +23,12 @@ export class GpuMatrix {
21
23
  * no padding). `data` must have at least `rows * lda` (row-major) or
22
24
  * `cols * lda` (column-major) elements.
23
25
  */
24
- static from(data, rows, cols, lda, layout = "row-major") {
26
+ static from(deviceOrData, ...rest) {
27
+ const explicit = deviceOrData instanceof GPUDevice;
28
+ const device = explicit ? deviceOrData : getDevice();
29
+ const data = explicit ? rest.shift() : deviceOrData;
30
+ let [rows, cols, lda, layout = "row-major"] = rest;
31
+
25
32
  if (layout !== "row-major" && layout !== "column-major")
26
33
  throw new Error("layout must be 'row-major' or 'column-major'.");
27
34
  const isRowMajor = layout === "row-major";
@@ -48,19 +55,19 @@ export class GpuMatrix {
48
55
  if (data instanceof Float64Array) {
49
56
  const n = outerCount * lda;
50
57
  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);
58
+ const hiBuf = uploadBuffer(device, hi, "gpu-matrix-f64-hi", true);
59
+ const loBuf = uploadBuffer(device, lo, "gpu-matrix-f64-lo", true);
60
+ return new GpuMatrix(hiBuf, rows, cols, lda, loBuf, layout, device);
54
61
  }
55
62
 
56
- const buf = uploadBuffer(data.subarray(0, outerCount * lda), "gpu-matrix", true);
57
- return new GpuMatrix(buf, rows, cols, lda, null, layout);
63
+ const buf = uploadBuffer(device, data.subarray(0, outerCount * lda), "gpu-matrix", true);
64
+ return new GpuMatrix(buf, rows, cols, lda, null, layout, device);
58
65
  }
59
66
 
60
67
  async read() {
61
- const device = getDevice();
68
+ const device = this.device;
62
69
  const enc = device.createCommandEncoder();
63
- const rb = stageReadback(enc, this._buf);
70
+ const rb = stageReadback(device, enc, this._buf);
64
71
  device.queue.submit([enc.finish()]);
65
72
 
66
73
  const isRowMajor = this.layout !== "column-major";
@@ -69,7 +76,7 @@ export class GpuMatrix {
69
76
 
70
77
  if (this._loBuf) {
71
78
  const encLo = device.createCommandEncoder();
72
- const rbLo = stageReadback(encLo, this._loBuf);
79
+ const rbLo = stageReadback(device, encLo, this._loBuf);
73
80
  device.queue.submit([encLo.finish()]);
74
81
 
75
82
  const [hi, lo] = await Promise.all([
@@ -4,37 +4,55 @@ import { extractResult } from "../util/result.mjs";
4
4
  import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
5
5
 
6
6
  export class GpuVector {
7
- constructor(buffer, length, dtype = Float32Array, loBuffer = null) {
7
+ constructor(buffer, length, dtype = Float32Array, loBuffer = null, device = null) {
8
8
  this._buf = buffer;
9
9
  this._loBuf = loBuffer; // Non-null only for Float64Array-backed vectors
10
10
  this.length = length;
11
11
  this.dtype = dtype;
12
+ // A GPUBuffer belongs to exactly one device and WebGPU rejects any attempt
13
+ // to use it with another, so every handle remembers where it lives. Routines
14
+ // check this to reject mixed-device operands with a clear message instead of
15
+ // a raw GPUValidationError.
16
+ this.device = device ?? getDevice();
12
17
  }
13
18
 
14
- static from(data) {
19
+ /**
20
+ * Uploads a vector to GPU memory.
21
+ *
22
+ * Pass the target `GPUDevice` first — matching every routine's own
23
+ * `(device, ...)` convention. Omitting it falls back to the device from the
24
+ * last `init()`, which is the historical form and only works single-device.
25
+ *
26
+ * @param {GPUDevice|Float32Array|Float64Array} deviceOrData
27
+ */
28
+ static from(deviceOrData, maybeData) {
29
+ const explicit = deviceOrData instanceof GPUDevice;
30
+ const device = explicit ? deviceOrData : getDevice();
31
+ const data = explicit ? maybeData : deviceOrData;
32
+
15
33
  if (data instanceof Float64Array) {
16
34
  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);
35
+ const hiBuf = uploadBuffer(device, hi, "gpu-vector-f64-hi", true);
36
+ const loBuf = uploadBuffer(device, lo, "gpu-vector-f64-lo", true);
37
+ return new GpuVector(hiBuf, data.length, Float64Array, loBuf, device);
20
38
  }
21
39
  if (!(data instanceof Float32Array)) {
22
40
  throw new Error("GpuVector.from expects a Float32Array or Float64Array.");
23
41
  }
24
- const buf = uploadBuffer(data, "gpu-vector", true);
25
- return new GpuVector(buf, data.length, data.constructor);
42
+ const buf = uploadBuffer(device, data, "gpu-vector", true);
43
+ return new GpuVector(buf, data.length, data.constructor, null, device);
26
44
  }
27
45
 
28
46
  async read() {
29
- const device = getDevice();
47
+ const device = this.device;
30
48
  const enc = device.createCommandEncoder();
31
- const rb = stageReadback(enc, this._buf);
49
+ const rb = stageReadback(device, enc, this._buf);
32
50
  device.queue.submit([enc.finish()]);
33
51
 
34
52
  if (!this._loBuf) return extractResult(rb, this.dtype);
35
53
 
36
54
  const encLo = device.createCommandEncoder();
37
- const rbLo = stageReadback(encLo, this._loBuf);
55
+ const rbLo = stageReadback(device, encLo, this._loBuf);
38
56
  device.queue.submit([encLo.finish()]);
39
57
 
40
58
  const [hi, lo] = await Promise.all([
@@ -2,10 +2,10 @@ import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
4
  * Computes the sum of absolute values of a vector of doubles in extended
5
- * precision: result = sum(|x[i]|). Each element of `x` has abs() applied,
6
- * then is split into a (hi, lo) double-double f32 pair (see
7
- * `splitDoubleDouble`/`f64.mjs`) since WGSL has no f64 type; accumulation
8
- * uses Dekker's double-double algorithm (see `shaders/f64/dekker.wgsl`), giving ~48 bits
5
+ * precision: $$\text{result} = \sum_{i} |x_i|$$
6
+ * Each element of `x` has abs() applied, then is split into a (hi, lo) double-double f32 pair
7
+ * (see `splitDoubleDouble`/`f64.mjs`) since WGSL has no f64 type; accumulation
8
+ * uses Dekker's double-double algorithm (see `shaders/f64/`), giving ~48 bits
9
9
  * of mantissa — more than a single f32 (24 bits) but less than true f64
10
10
  * (52 bits), so results are not bit-exact with a CPU double.
11
11
  *
@@ -31,9 +31,9 @@ export declare function dasum(
31
31
 
32
32
  /**
33
33
  * Computes the sum of absolute values of a vector of doubles in double
34
- * precision: result = sum(|x[i]|).
34
+ * precision: $$\text{result} = \sum_{i} |x_i|$$.
35
35
  *
36
- * {@includeCode ../../examples/dasum/gpuvec.dasum.js}
36
+ * {@includeCode ../../examples/dasum/gpu.dasum.js}
37
37
  *
38
38
  * @param device - GPUDevice from `init()`
39
39
  * @param n - number of elements (must be a positive integer)
@@ -13,14 +13,16 @@ import { extractResult } from "../util/result.mjs";
13
13
  import { getPipeline } from "../util/pipeline.mjs";
14
14
  import { GpuVector } from "../classes/GpuVector.mjs";
15
15
  import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
16
+ import { WGS } from "../util/constants.mjs";
17
+ import { requireSameDevice } from "../util/device.mjs";
16
18
 
17
- const WGS = 64; // workgroup size
18
19
 
19
20
  export async function dasum(device, n, x, incx) {
20
21
  const xIsGpu = x instanceof GpuVector;
21
22
 
22
23
  if (!(device instanceof GPUDevice))
23
24
  throw new Error("device must be a GPUDevice.");
25
+ requireSameDevice(device, "dasum", { x });
24
26
  if (!Number.isInteger(n) || !Number.isInteger(incx))
25
27
  throw new Error("n and incx must be integers.");
26
28
  if (incx <= 0) throw new Error("incx must be positive.");
@@ -34,10 +36,12 @@ export async function dasum(device, n, x, incx) {
34
36
  "x does not have enough elements for the given n and incx.",
35
37
  );
36
38
 
37
- // Concatenated with f64/dekker.wgsl for its DD struct/ddAdd helpers (WGSL
38
- // has no #include); entryPoint omitted since each module has only one @compute.
39
- const pipelineMain = await getPipeline(device, ["f64/dekker", "dasum"]);
40
- const pipelineReduce = await getPipeline(device, ["f64/dekker", "reduction/sumF64"]);
39
+ // Concatenated with f64/dekker.wgsl (DD struct), f64/utils/abs.wgsl
40
+ // (ddAbs), and f64/utils/add.wgsl (ddAddProtected) — WGSL has no
41
+ // #include; entryPoint omitted since each module has only one @compute.
42
+ const f64Deps = ["f64/dekker", "f64/utils/abs", "f64/utils/add"];
43
+ const pipelineMain = await getPipeline(device, [...f64Deps, "dasum"]);
44
+ const pipelineReduce = await getPipeline(device, [...f64Deps, "reduction/sumF64"]);
41
45
 
42
46
  let xHiBuffer = null;
43
47
  let xLoBuffer = null;
@@ -55,14 +59,14 @@ export async function dasum(device, n, x, incx) {
55
59
  xLoBuffer = x._loBuf;
56
60
  } else {
57
61
  const { hi, lo } = splitDoubleDouble(x.map(Math.abs));
58
- xHiBuffer = uploadBuffer(hi, "dasum-xHi", false);
59
- xLoBuffer = uploadBuffer(lo, "dasum-xLo", false);
62
+ xHiBuffer = uploadBuffer(device, hi, "dasum-xHi", false);
63
+ xLoBuffer = uploadBuffer(device, lo, "dasum-xLo", false);
60
64
  }
61
- partialsHiBuffer = createStorageBuffer(2 * WGS * 4, "dasum-partialsHi");
62
- partialsLoBuffer = createStorageBuffer(2 * WGS * 4, "dasum-partialsLo");
63
- resultHiBuffer = createResultBuffer(4, "dasum-result-hi");
64
- resultLoBuffer = createResultBuffer(4, "dasum-result-lo");
65
- paramsBuffer = createParamsBuffer(
65
+ partialsHiBuffer = createStorageBuffer(device, 2 * WGS * 4, "dasum-partialsHi");
66
+ partialsLoBuffer = createStorageBuffer(device, 2 * WGS * 4, "dasum-partialsLo");
67
+ resultHiBuffer = createResultBuffer(device, 4, "dasum-result-hi");
68
+ resultLoBuffer = createResultBuffer(device, 4, "dasum-result-lo");
69
+ paramsBuffer = createParamsBuffer(device,
66
70
  [
67
71
  { value: n, type: "u32" },
68
72
  { value: incx, type: "u32" },
@@ -70,31 +74,31 @@ export async function dasum(device, n, x, incx) {
70
74
  "dasum-params",
71
75
  );
72
76
 
73
- const bgMain = createBindGroup(
77
+ const bgMain = createBindGroup(device,
74
78
  pipelineMain.getBindGroupLayout(0),
75
79
  [xHiBuffer, xLoBuffer, partialsHiBuffer, partialsLoBuffer, paramsBuffer],
76
80
  );
77
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(
81
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
78
82
  pipelineMain,
79
83
  bgMain,
80
84
  2 * WGS,
81
85
  ); // dispatch 2*WGS workgroups
82
86
 
83
- submit(enc1);
87
+ submit(device, enc1);
84
88
 
85
- const bgReduce = createBindGroup(
89
+ const bgReduce = createBindGroup(device,
86
90
  pipelineReduce.getBindGroupLayout(0),
87
91
  [partialsHiBuffer, partialsLoBuffer, resultHiBuffer, resultLoBuffer],
88
92
  );
89
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(
93
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
90
94
  pipelineReduce,
91
95
  bgReduce,
92
96
  1,
93
97
  ); // dispatch 1 workgroup to reduce the partial sums to a single result
94
- readHiBuffer = stageReadback(enc2, resultHiBuffer);
95
- readLoBuffer = stageReadback(enc2, resultLoBuffer);
98
+ readHiBuffer = stageReadback(device, enc2, resultHiBuffer);
99
+ readLoBuffer = stageReadback(device, enc2, resultLoBuffer);
96
100
 
97
- submit(enc2);
101
+ submit(device, enc2);
98
102
 
99
103
  const hiPromise = extractResult(readHiBuffer, Float32Array);
100
104
  const loPromise = extractResult(readLoBuffer, Float32Array);
@@ -121,7 +125,7 @@ export async function dasum(device, n, x, incx) {
121
125
  if (resultHiBuffer) destroyBuffers(resultHiBuffer);
122
126
  if (resultLoBuffer) destroyBuffers(resultLoBuffer);
123
127
  if (paramsBuffer) destroyBuffers(paramsBuffer);
124
- // Only reached if submit(enc2) threw before ownership was transferred above.
128
+ // Only reached if submit(device, enc2) threw before ownership was transferred above.
125
129
  if (readHiBuffer) destroyBuffers(readHiBuffer);
126
130
  if (readLoBuffer) destroyBuffers(readLoBuffer);
127
131
  }
package/src/devdocs.mjs CHANGED
@@ -4,6 +4,19 @@
4
4
  * repository and explains the reasoning behind it — not just what the code
5
5
  * does, but why it is shaped the way it is.
6
6
  *
7
+ * ## How These Docs Are Organized
8
+ *
9
+ * The sidebar's top-level modules mirror the repository's top-level folders.
10
+ * `devdocs` stands in for `src/` — the module you're reading right now is a
11
+ * narrated walkthrough of it, with `devdocs/blas-routines` and
12
+ * `devdocs/shaders` covering `src/index.mjs` and `src/shaders/`
13
+ * respectively. Its siblings — `assets`, `benchmarks`, `examples`, `scripts`,
14
+ * `tests` — each document the identically-named top-level folder.
15
+ *
16
+ * `docs` is the one exception: it's the public API reference generated from
17
+ * `index.d.mts`, not a tour of the top-level `docs/` folder. That folder is
18
+ * this site's own generated output — the two just happen to share a name.
19
+ *
7
20
  * ## What is BLAS?
8
21
  *
9
22
  * BLAS (Basic Linear Algebra Subprograms) is a standard API for vector and
@@ -0,0 +1,69 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+
3
+ /**
4
+ * Returns the 0-based index of the element with the largest absolute value,
5
+ * for a vector of doubles: $$\text{index} = \arg\max_{i} |x_i|$$
6
+ * Each element of `x` is split into a (hi, lo)
7
+ * double-double f32 pair (see `splitDoubleDouble`/`f64.mjs`) since WGSL has
8
+ * no f64 type; comparisons use the double-double pair directly (hi, falling
9
+ * back to lo on an exact tie), giving ~48 bits of discriminating precision —
10
+ * more than a single f32 (24 bits) but less than true f64 (52 bits). Ties
11
+ * are broken in favour of the lower index, matching CBLAS behaviour.
12
+ *
13
+ * **NaN handling.** The search compares with `>`, which is false for NaN, so
14
+ * NaN elements are skipped rather than selected: a vector of all NaN returns
15
+ * `0`. This matches CBLAS except when `x[0]` itself is NaN — CBLAS seeds its
16
+ * running maximum from `x[0]`, and since no comparison against NaN succeeds it
17
+ * returns `0` however large the later elements are, whereas this returns the
18
+ * index of the largest non-NaN element. `+-Infinity` compares normally and is
19
+ * selected as the maximum.
20
+ *
21
+ * {@includeCode ../../examples/idamax/idamax.js}
22
+ *
23
+ * **Browser (standalone HTML):**
24
+ * {@includeCode ../../examples/idamax/web/idamax.html}
25
+ *
26
+ * @param device - GPUDevice from `init()`
27
+ * @param n - number of elements (must be a positive integer)
28
+ * @param x - Float64Array input vector
29
+ * @param incx - stride for x (must be a positive integer)
30
+ * @returns 0-based index of max |x[i]|
31
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/idamax/idamax.mjs#L19">Source code: idamax.mjs (L19)</a>
32
+ * @category BLAS Level 1
33
+ */
34
+ export declare function idamax(
35
+ device: GPUDevice,
36
+ n: number,
37
+ x: Float64Array,
38
+ incx: number,
39
+ ): Promise<{ index: number } | { index: number; gpuTimeMs: number }>;
40
+
41
+ /**
42
+ * Returns the 0-based index of the element with the largest absolute value:
43
+ * $$\text{index} = \arg\max_{i} |x_i|$$
44
+ * Ties are broken in favour of the lower index, matching CBLAS behaviour.
45
+ *
46
+ * **NaN handling.** The search compares with `>`, which is false for NaN, so
47
+ * NaN elements are skipped rather than selected: a vector of all NaN returns
48
+ * `0`. This matches CBLAS except when `x[0]` itself is NaN — CBLAS seeds its
49
+ * running maximum from `x[0]`, and since no comparison against NaN succeeds it
50
+ * returns `0` however large the later elements are, whereas this returns the
51
+ * index of the largest non-NaN element. `+-Infinity` compares normally and is
52
+ * selected as the maximum.
53
+ *
54
+ * {@includeCode ../../examples/idamax/gpu.idamax.js}
55
+ *
56
+ * @param device - GPUDevice from `init()`
57
+ * @param n - number of elements (must be a positive integer)
58
+ * @param x - Float64Array-backed GpuVector input vector
59
+ * @param incx - stride for x (must be a positive integer)
60
+ * @returns 0-based index of max |x[i]|
61
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/idamax/idamax.mjs#L19">Source code: idamax.mjs (L19)</a>
62
+ * @category BLAS Level 1
63
+ */
64
+ export declare function idamax(
65
+ device: GPUDevice,
66
+ n: number,
67
+ x: GpuVector,
68
+ incx: number,
69
+ ): Promise<{ index: number } | { index: number; gpuTimeMs: number }>;