wgblas 2.0.0 → 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 (66) hide show
  1. package/README.md +18 -18
  2. package/dist/wgblas.browser.js +1273 -1239
  3. package/index.d.mts +38 -6
  4. package/package.json +2 -1
  5. package/src/classes/GpuMatrix.mjs +17 -10
  6. package/src/classes/GpuVector.mjs +28 -10
  7. package/src/dasum/dasum.d.mts +4 -4
  8. package/src/dasum/dasum.mjs +19 -17
  9. package/src/devdocs.mjs +13 -0
  10. package/src/idamax/idamax.d.mts +20 -2
  11. package/src/idamax/idamax.mjs +18 -16
  12. package/src/init.mjs +114 -56
  13. package/src/isamax/isamax.d.mts +20 -2
  14. package/src/isamax/isamax.mjs +16 -14
  15. package/src/random/random.d.mts +1 -0
  16. package/src/sasum/sasum.d.mts +2 -2
  17. package/src/sasum/sasum.mjs +15 -13
  18. package/src/saxpy/saxpy.d.mts +2 -2
  19. package/src/saxpy/saxpy.mjs +10 -8
  20. package/src/scopy/scopy.d.mts +2 -2
  21. package/src/scopy/scopy.mjs +10 -8
  22. package/src/sdot/sdot.d.mts +2 -2
  23. package/src/sdot/sdot.mjs +16 -14
  24. package/src/sgemm/sgemm.mjs +28 -15
  25. package/src/sgemmtr/sgemmtr.mjs +16 -15
  26. package/src/sgemv/sgemv.mjs +38 -26
  27. package/src/sger/sger.mjs +10 -8
  28. package/src/shaders/index.mjs +164 -14
  29. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  30. package/src/shaders/sgemm_large.wgsl +107 -18
  31. package/src/shaders/sgemm_small.wgsl +115 -15
  32. package/src/shaders/sgemmtr_large.wgsl +4 -1
  33. package/src/shaders/sgemmtr_small.wgsl +4 -1
  34. package/src/shaders/sgemv_n.wgsl +3 -1
  35. package/src/shaders/sgemv_t.wgsl +3 -1
  36. package/src/shaders/snrm2.wgsl +72 -23
  37. package/src/shaders/ssymv.wgsl +3 -1
  38. package/src/snrm2/snrm2.d.mts +2 -2
  39. package/src/snrm2/snrm2.mjs +33 -21
  40. package/src/srot/srot.d.mts +2 -4
  41. package/src/srot/srot.mjs +11 -9
  42. package/src/srotm/srotm.d.mts +2 -4
  43. package/src/srotm/srotm.mjs +19 -10
  44. package/src/sscal/sscal.d.mts +3 -3
  45. package/src/sscal/sscal.mjs +12 -10
  46. package/src/sswap/sswap.d.mts +2 -2
  47. package/src/sswap/sswap.mjs +11 -9
  48. package/src/ssymm/ssymm.mjs +31 -22
  49. package/src/ssymv/ssymv.mjs +10 -8
  50. package/src/ssyr/ssyr.mjs +9 -7
  51. package/src/ssyr2/ssyr2.mjs +10 -8
  52. package/src/ssyr2k/ssyr2k.mjs +18 -17
  53. package/src/ssyrk/ssyrk.mjs +18 -17
  54. package/src/strmm/strmm.mjs +47 -32
  55. package/src/strmv/strmv.mjs +10 -8
  56. package/src/strsm/strsm.mjs +54 -36
  57. package/src/strsv/strsv.mjs +16 -12
  58. package/src/util/benchmark.mjs +4 -6
  59. package/src/util/bindgroup.mjs +1 -3
  60. package/src/util/buffer.mjs +113 -19
  61. package/src/util/compute.mjs +6 -9
  62. package/src/util/constants.mjs +57 -0
  63. package/src/util/device.mjs +34 -0
  64. package/src/util/pipeline.mjs +5 -6
  65. package/src/util/workgroup.mjs +55 -7
  66. package/src/shaders/browser-shaders.mjs +0 -81
package/index.d.mts CHANGED
@@ -38,11 +38,20 @@ export { strsm } from "./src/strsm/strsm.mjs";
38
38
  /**
39
39
  * Initializes the WebGPU device.
40
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
+ *
41
48
  * @param options.powerPreference - GPU power preference (default: `"high-performance"`).
42
49
  * This is a hint to the browser: on dual-GPU systems, `"high-performance"` typically favors the discrete GPU
43
50
  * and `"low-power"` favors the integrated one.
44
51
  * See [MDN: GPU.requestAdapter()](https://developer.mozilla.org/en-US/docs/Web/API/GPU/requestAdapter).
45
- * @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`)
46
55
  * @param options.dumpShaders - Node-only. Forwards Dawn's `dump_shaders` debug toggle, printing
47
56
  * each pipeline's WGSL and compiled backend IR (SPIR-V/Vulkan, MSL/Metal, or HLSL/D3D12,
48
57
  * whichever Dawn picked) to stderr as it compiles. A Dawn passthrough, not a wgblas format —
@@ -71,10 +80,23 @@ export { strsm } from "./src/strsm/strsm.mjs";
71
80
  * const n = 5;
72
81
  * const alpha = 2.0;
73
82
  * const x = new Float32Array([1, 2, 3, 4, 5]);
74
- * const { result, gpuTimeMs } = await sscal(device, n, alpha, x, 1);
83
+ * const { x: result, gpuTimeMs } = await sscal(device, n, alpha, x, 1);
75
84
  * console.log(`Result: [${Array.from(result).join(", ")}]`);
76
85
  * console.log(`GPU time: ${gpuTimeMs.toFixed(3)} ms`);
77
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
+ *
78
100
  * @see [Source code: init.mjs](https://github.com/manit2004/wgblas/blob/main/src/init.mjs#L18-L54)
79
101
  * @category Core
80
102
  */
@@ -85,8 +107,15 @@ export declare function init(options?: {
85
107
  }): Promise<GPUDevice>;
86
108
 
87
109
  /**
88
- * Destroys the WebGPU device, releases the adapter, resets benchmark state, and fires all internal
89
- * 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.
90
119
  *
91
120
  * @example
92
121
  * ```js
@@ -99,11 +128,14 @@ export declare function init(options?: {
99
128
  * @see [Source code: init.mjs](https://github.com/manit2004/wgblas/blob/main/src/init.mjs#L56-L65)
100
129
  * @category Core
101
130
  */
102
- export declare function cleanup(): void;
131
+ export declare function cleanup(device?: GPUDevice): void;
103
132
 
104
133
  /**
105
134
  * Returns the GPU device name from the WebGPU adapter info. Must be called after `init()`.
106
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
+ *
107
139
  * @example
108
140
  * ```js
109
141
  * import { init, gpuName } from "wgblas";
@@ -114,4 +146,4 @@ export declare function cleanup(): void;
114
146
  * @see [Source code: init.mjs](https://github.com/manit2004/wgblas/blob/main/src/init.mjs#L81-L87)
115
147
  * @category Core
116
148
  */
117
- export declare function gpuName(): { description: string; device: string };
149
+ export declare function gpuName(device?: GPUDevice): { description: string; device: string };
package/package.json CHANGED
@@ -1,6 +1,6 @@
1
1
  {
2
2
  "name": "wgblas",
3
- "version": "2.0.0",
3
+ "version": "2.1.0",
4
4
  "description": "BLAS on WebGPU",
5
5
  "type": "module",
6
6
  "main": "index.mjs",
@@ -209,6 +209,7 @@
209
209
  "lint-staged": "^17.0.8",
210
210
  "prettier": "^3.9.4",
211
211
  "typedoc": "^0.28.19",
212
+ "typedoc-plugin-katex": "^0.1.2",
212
213
  "typescript": "^6.0.3",
213
214
  "vite": "^8.0.16"
214
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,9 +2,9 @@ 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
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
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.
@@ -31,7 +31,7 @@ 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
36
  * {@includeCode ../../examples/dasum/gpu.dasum.js}
37
37
  *
@@ -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.");
@@ -57,14 +59,14 @@ export async function dasum(device, n, x, incx) {
57
59
  xLoBuffer = x._loBuf;
58
60
  } else {
59
61
  const { hi, lo } = splitDoubleDouble(x.map(Math.abs));
60
- xHiBuffer = uploadBuffer(hi, "dasum-xHi", false);
61
- xLoBuffer = uploadBuffer(lo, "dasum-xLo", false);
62
+ xHiBuffer = uploadBuffer(device, hi, "dasum-xHi", false);
63
+ xLoBuffer = uploadBuffer(device, lo, "dasum-xLo", false);
62
64
  }
63
- partialsHiBuffer = createStorageBuffer(2 * WGS * 4, "dasum-partialsHi");
64
- partialsLoBuffer = createStorageBuffer(2 * WGS * 4, "dasum-partialsLo");
65
- resultHiBuffer = createResultBuffer(4, "dasum-result-hi");
66
- resultLoBuffer = createResultBuffer(4, "dasum-result-lo");
67
- 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,
68
70
  [
69
71
  { value: n, type: "u32" },
70
72
  { value: incx, type: "u32" },
@@ -72,31 +74,31 @@ export async function dasum(device, n, x, incx) {
72
74
  "dasum-params",
73
75
  );
74
76
 
75
- const bgMain = createBindGroup(
77
+ const bgMain = createBindGroup(device,
76
78
  pipelineMain.getBindGroupLayout(0),
77
79
  [xHiBuffer, xLoBuffer, partialsHiBuffer, partialsLoBuffer, paramsBuffer],
78
80
  );
79
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(
81
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
80
82
  pipelineMain,
81
83
  bgMain,
82
84
  2 * WGS,
83
85
  ); // dispatch 2*WGS workgroups
84
86
 
85
- submit(enc1);
87
+ submit(device, enc1);
86
88
 
87
- const bgReduce = createBindGroup(
89
+ const bgReduce = createBindGroup(device,
88
90
  pipelineReduce.getBindGroupLayout(0),
89
91
  [partialsHiBuffer, partialsLoBuffer, resultHiBuffer, resultLoBuffer],
90
92
  );
91
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(
93
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
92
94
  pipelineReduce,
93
95
  bgReduce,
94
96
  1,
95
97
  ); // dispatch 1 workgroup to reduce the partial sums to a single result
96
- readHiBuffer = stageReadback(enc2, resultHiBuffer);
97
- readLoBuffer = stageReadback(enc2, resultLoBuffer);
98
+ readHiBuffer = stageReadback(device, enc2, resultHiBuffer);
99
+ readLoBuffer = stageReadback(device, enc2, resultLoBuffer);
98
100
 
99
- submit(enc2);
101
+ submit(device, enc2);
100
102
 
101
103
  const hiPromise = extractResult(readHiBuffer, Float32Array);
102
104
  const loPromise = extractResult(readLoBuffer, Float32Array);
@@ -123,7 +125,7 @@ export async function dasum(device, n, x, incx) {
123
125
  if (resultHiBuffer) destroyBuffers(resultHiBuffer);
124
126
  if (resultLoBuffer) destroyBuffers(resultLoBuffer);
125
127
  if (paramsBuffer) destroyBuffers(paramsBuffer);
126
- // Only reached if submit(enc2) threw before ownership was transferred above.
128
+ // Only reached if submit(device, enc2) threw before ownership was transferred above.
127
129
  if (readHiBuffer) destroyBuffers(readHiBuffer);
128
130
  if (readLoBuffer) destroyBuffers(readLoBuffer);
129
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
@@ -2,13 +2,22 @@ import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
4
  * Returns the 0-based index of the element with the largest absolute value,
5
- * for a vector of doubles. Each element of `x` is split into a (hi, lo)
5
+ * for a vector of doubles: $$\text{index} = \arg\max_{i} |x_i|$$
6
+ * Each element of `x` is split into a (hi, lo)
6
7
  * double-double f32 pair (see `splitDoubleDouble`/`f64.mjs`) since WGSL has
7
8
  * no f64 type; comparisons use the double-double pair directly (hi, falling
8
9
  * back to lo on an exact tie), giving ~48 bits of discriminating precision —
9
10
  * more than a single f32 (24 bits) but less than true f64 (52 bits). Ties
10
11
  * are broken in favour of the lower index, matching CBLAS behaviour.
11
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
+ *
12
21
  * {@includeCode ../../examples/idamax/idamax.js}
13
22
  *
14
23
  * **Browser (standalone HTML):**
@@ -30,9 +39,18 @@ export declare function idamax(
30
39
  ): Promise<{ index: number } | { index: number; gpuTimeMs: number }>;
31
40
 
32
41
  /**
33
- * Returns the 0-based index of the element with the largest absolute value.
42
+ * Returns the 0-based index of the element with the largest absolute value:
43
+ * $$\text{index} = \arg\max_{i} |x_i|$$
34
44
  * Ties are broken in favour of the lower index, matching CBLAS behaviour.
35
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
+ *
36
54
  * {@includeCode ../../examples/idamax/gpu.idamax.js}
37
55
  *
38
56
  * @param device - GPUDevice from `init()`
@@ -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 } from "../util/f64.mjs";
16
+ import { WGS } from "../util/constants.mjs";
17
+ import { requireSameDevice } from "../util/device.mjs";
16
18
 
17
- const WGS = 64;
18
19
 
19
20
  export async function idamax(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, "idamax", { 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.");
@@ -54,14 +56,14 @@ export async function idamax(device, n, x, incx) {
54
56
  xLoBuffer = x._loBuf;
55
57
  } else {
56
58
  const { hi, lo } = splitDoubleDouble(x);
57
- xHiBuffer = uploadBuffer(hi, "idamax-xHi", false);
58
- xLoBuffer = uploadBuffer(lo, "idamax-xLo", false);
59
+ xHiBuffer = uploadBuffer(device, hi, "idamax-xHi", false);
60
+ xLoBuffer = uploadBuffer(device, lo, "idamax-xLo", false);
59
61
  }
60
- partialsValHiBuffer = createStorageBuffer(2 * WGS * 4, "idamax-partials-val-hi");
61
- partialsValLoBuffer = createStorageBuffer(2 * WGS * 4, "idamax-partials-val-lo");
62
- partialsIdxBuffer = createStorageBuffer(2 * WGS * 4, "idamax-partials-idx");
63
- resultBuffer = createResultBuffer(4, "idamax-result"); // u32 index
64
- paramsBuffer = createParamsBuffer(
62
+ partialsValHiBuffer = createStorageBuffer(device, 2 * WGS * 4, "idamax-partials-val-hi");
63
+ partialsValLoBuffer = createStorageBuffer(device, 2 * WGS * 4, "idamax-partials-val-lo");
64
+ partialsIdxBuffer = createStorageBuffer(device, 2 * WGS * 4, "idamax-partials-idx");
65
+ resultBuffer = createResultBuffer(device, 4, "idamax-result"); // u32 index
66
+ paramsBuffer = createParamsBuffer(device,
65
67
  [
66
68
  { value: n, type: "u32" },
67
69
  { value: incx, type: "u32" },
@@ -69,7 +71,7 @@ export async function idamax(device, n, x, incx) {
69
71
  "idamax-params",
70
72
  );
71
73
 
72
- const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
74
+ const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
73
75
  xHiBuffer,
74
76
  xLoBuffer,
75
77
  partialsValHiBuffer,
@@ -77,28 +79,28 @@ export async function idamax(device, n, x, incx) {
77
79
  partialsIdxBuffer,
78
80
  paramsBuffer,
79
81
  ]);
80
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(
82
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
81
83
  pipelineMain,
82
84
  bgMain,
83
85
  2 * WGS,
84
86
  ); // dispatch 2*WGS workgroups
85
87
 
86
- submit(enc1);
88
+ submit(device, enc1);
87
89
 
88
- const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
90
+ const bgReduce = createBindGroup(device, pipelineReduce.getBindGroupLayout(0), [
89
91
  partialsValHiBuffer,
90
92
  partialsValLoBuffer,
91
93
  partialsIdxBuffer,
92
94
  resultBuffer,
93
95
  ]);
94
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(
96
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
95
97
  pipelineReduce,
96
98
  bgReduce,
97
99
  1,
98
100
  ); // dispatch 1 workgroup to reduce the partials to a single index
99
- readBuffer = stageReadback(enc2, resultBuffer);
101
+ readBuffer = stageReadback(device, enc2, resultBuffer);
100
102
 
101
- submit(enc2);
103
+ submit(device, enc2);
102
104
 
103
105
  const resultPromise = extractResult(readBuffer, Uint32Array);
104
106
  readBuffer = null; // ownership transferred — extractResult's own finally destroys it
@@ -122,7 +124,7 @@ export async function idamax(device, n, x, incx) {
122
124
  if (partialsIdxBuffer) destroyBuffers(partialsIdxBuffer);
123
125
  if (resultBuffer) destroyBuffers(resultBuffer);
124
126
  if (paramsBuffer) destroyBuffers(paramsBuffer);
125
- // Only reached if submit(enc2) threw before ownership was transferred above.
127
+ // Only reached if submit(device, enc2) threw before ownership was transferred above.
126
128
  if (readBuffer) destroyBuffers(readBuffer);
127
129
  }
128
130
  }