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,155 @@
1
+ import {
2
+ uploadBuffer,
3
+ createParamsBuffer,
4
+ stageReadback,
5
+ destroyBuffers,
6
+ } from "../util/buffer.mjs";
7
+ import { createBindGroup } from "../util/bindgroup.mjs";
8
+ import { runComputePass, submit } from "../util/compute.mjs";
9
+ import { extractResult } from "../util/result.mjs";
10
+ import { extractTimestamp } from "../util/benchmark.mjs";
11
+ import { getPipeline } from "../util/pipeline.mjs";
12
+ import { calcWorkgroups } from "../util/workgroup.mjs";
13
+ import { GpuVector } from "../classes/GpuVector.mjs";
14
+ import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
15
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
16
+
17
+ // dswap: x <-> y, double-double (Dekker) f64 emulation of sswap — x and y
18
+ // are each split into an f32 (hi, lo) pair; WGSL has no f64 type. Unlike
19
+ // dscal/daxpy/ddot, a swap has no arithmetic at all, so it needs no Dekker
20
+ // shader dependencies (dekker.wgsl/add.wgsl/multiply.wgsl) — dswap.wgsl is
21
+ // entirely self-contained.
22
+ export async function dswap(device, n, x, incx, y, incy) {
23
+ const xIsGpu = x instanceof GpuVector;
24
+ const yIsGpu = y instanceof GpuVector;
25
+
26
+ requireGpuDevice(device);
27
+ requireSameDevice(device, "dswap", { x, y });
28
+ if (
29
+ !Number.isInteger(n) ||
30
+ !Number.isInteger(incx) ||
31
+ !Number.isInteger(incy)
32
+ )
33
+ throw new Error("n, incx, and incy must be integers.");
34
+ if (incx <= 0 || incy <= 0)
35
+ throw new Error("incx and incy must be positive.");
36
+ if (!(x instanceof Float64Array) && !xIsGpu)
37
+ throw new Error("x must be a Float64Array or GpuVector.");
38
+ if (!(y instanceof Float64Array) && !yIsGpu)
39
+ throw new Error("y must be a Float64Array or GpuVector.");
40
+ if (xIsGpu && x.dtype !== Float64Array)
41
+ throw new Error("x must be a Float64Array-backed GpuVector.");
42
+ if (yIsGpu && y.dtype !== Float64Array)
43
+ throw new Error("y must be a Float64Array-backed GpuVector.");
44
+ if (xIsGpu !== yIsGpu)
45
+ throw new Error(
46
+ "x and y must be the same type (both Float64Array or both GpuVector).",
47
+ );
48
+ if (n <= 0) return xIsGpu ? {} : { x, y };
49
+ if (x.length < (n - 1) * incx + 1)
50
+ throw new Error(
51
+ "x does not have enough elements for the given n and incx.",
52
+ );
53
+ if (y.length < (n - 1) * incy + 1)
54
+ throw new Error(
55
+ "y does not have enough elements for the given n and incy.",
56
+ );
57
+
58
+ const pipeline = await getPipeline(device, "dswap");
59
+
60
+ let xHiBuffer = null;
61
+ let xLoBuffer = null;
62
+ let yHiBuffer = null;
63
+ let yLoBuffer = null;
64
+ let paramsBuffer = null;
65
+ let xReadHiBuffer = null;
66
+ let xReadLoBuffer = null;
67
+ let yReadHiBuffer = null;
68
+ let yReadLoBuffer = null;
69
+
70
+ try {
71
+ if (xIsGpu) {
72
+ xHiBuffer = x._buf;
73
+ xLoBuffer = x._loBuf;
74
+ yHiBuffer = y._buf;
75
+ yLoBuffer = y._loBuf;
76
+ } else {
77
+ const xSplit = splitDoubleDouble(x);
78
+ const ySplit = splitDoubleDouble(y);
79
+ xHiBuffer = uploadBuffer(device, xSplit.hi, "dswap-xHi", true);
80
+ xLoBuffer = uploadBuffer(device, xSplit.lo, "dswap-xLo", true);
81
+ yHiBuffer = uploadBuffer(device, ySplit.hi, "dswap-yHi", true);
82
+ yLoBuffer = uploadBuffer(device, ySplit.lo, "dswap-yLo", true);
83
+ }
84
+ paramsBuffer = createParamsBuffer(
85
+ device,
86
+ [
87
+ { value: n, type: "u32" },
88
+ { value: incx, type: "u32" },
89
+ { value: incy, type: "u32" },
90
+ ],
91
+ "dswap-params",
92
+ );
93
+
94
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
95
+ xHiBuffer,
96
+ xLoBuffer,
97
+ yHiBuffer,
98
+ yLoBuffer,
99
+ paramsBuffer,
100
+ ]);
101
+ const { commandEncoder, ts } = runComputePass(
102
+ device,
103
+ pipeline,
104
+ bindGroup,
105
+ calcWorkgroups(device, n),
106
+ );
107
+ xReadHiBuffer = xIsGpu
108
+ ? null
109
+ : stageReadback(device, commandEncoder, xHiBuffer);
110
+ xReadLoBuffer = xIsGpu
111
+ ? null
112
+ : stageReadback(device, commandEncoder, xLoBuffer);
113
+ yReadHiBuffer = yIsGpu
114
+ ? null
115
+ : stageReadback(device, commandEncoder, yHiBuffer);
116
+ yReadLoBuffer = yIsGpu
117
+ ? null
118
+ : stageReadback(device, commandEncoder, yLoBuffer);
119
+
120
+ submit(device, commandEncoder);
121
+
122
+ const gpuTimeMs = await extractTimestamp(ts);
123
+
124
+ if (xIsGpu) {
125
+ // xIsGpu === yIsGpu, enforced above
126
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
127
+ return {};
128
+ }
129
+
130
+ const xHi = await extractResult(xReadHiBuffer, Float32Array);
131
+ xReadHiBuffer = null; // extractResult already destroyed it
132
+ const xLo = await extractResult(xReadLoBuffer, Float32Array);
133
+ xReadLoBuffer = null;
134
+ const yHi = await extractResult(yReadHiBuffer, Float32Array);
135
+ yReadHiBuffer = null;
136
+ const yLo = await extractResult(yReadLoBuffer, Float32Array);
137
+ yReadLoBuffer = null;
138
+ const resultX = mergeDoubleDouble(xHi, xLo);
139
+ const resultY = mergeDoubleDouble(yHi, yLo);
140
+ if (gpuTimeMs !== undefined) return { x: resultX, y: resultY, gpuTimeMs };
141
+ return { x: resultX, y: resultY };
142
+ } finally {
143
+ if (!xIsGpu && xHiBuffer) destroyBuffers(xHiBuffer);
144
+ if (!xIsGpu && xLoBuffer) destroyBuffers(xLoBuffer);
145
+ if (!yIsGpu && yHiBuffer) destroyBuffers(yHiBuffer);
146
+ if (!yIsGpu && yLoBuffer) destroyBuffers(yLoBuffer);
147
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
148
+ // Only reached if extractTimestamp or extractResult threw before
149
+ // clearing these — on the success path they're already null.
150
+ if (xReadHiBuffer) destroyBuffers(xReadHiBuffer);
151
+ if (xReadLoBuffer) destroyBuffers(xReadLoBuffer);
152
+ if (yReadHiBuffer) destroyBuffers(yReadHiBuffer);
153
+ if (yReadLoBuffer) destroyBuffers(yReadLoBuffer);
154
+ }
155
+ }
@@ -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,14 @@ 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
-
17
- const WGS = 64;
16
+ import { WGS } from "../util/constants.mjs";
17
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
18
18
 
19
19
  export async function idamax(device, n, x, incx) {
20
20
  const xIsGpu = x instanceof GpuVector;
21
21
 
22
- if (!(device instanceof GPUDevice))
23
- throw new Error("device must be a GPUDevice.");
22
+ requireGpuDevice(device);
23
+ requireSameDevice(device, "idamax", { x });
24
24
  if (!Number.isInteger(n) || !Number.isInteger(incx))
25
25
  throw new Error("n and incx must be integers.");
26
26
  if (incx <= 0) throw new Error("incx must be positive.");
@@ -35,9 +35,22 @@ export async function idamax(device, n, x, incx) {
35
35
  );
36
36
 
37
37
  // Concatenated f64 helpers (WGSL has no #include); ddAbs is unconditional in idamax.wgsl, so x is split as-is.
38
- const f64Deps = ["f64/dekker", "f64/utils/abs", "f64/utils/greater", "f64/utils/equal"];
39
- const pipelineMain = await getPipeline(device, [...f64Deps, "idamax"], "idamax_main");
40
- const pipelineReduce = await getPipeline(device, [...f64Deps, "reduction/argmaxF64"], "reduce_f64");
38
+ const f64Deps = [
39
+ "f64/dekker",
40
+ "f64/utils/abs",
41
+ "f64/utils/greater",
42
+ "f64/utils/equal",
43
+ ];
44
+ const pipelineMain = await getPipeline(
45
+ device,
46
+ [...f64Deps, "idamax"],
47
+ "idamax_main",
48
+ );
49
+ const pipelineReduce = await getPipeline(
50
+ device,
51
+ [...f64Deps, "reduction/argmaxF64"],
52
+ "reduce_f64",
53
+ );
41
54
 
42
55
  let xHiBuffer = null;
43
56
  let xLoBuffer = null;
@@ -54,14 +67,27 @@ export async function idamax(device, n, x, incx) {
54
67
  xLoBuffer = x._loBuf;
55
68
  } else {
56
69
  const { hi, lo } = splitDoubleDouble(x);
57
- xHiBuffer = uploadBuffer(hi, "idamax-xHi", false);
58
- xLoBuffer = uploadBuffer(lo, "idamax-xLo", false);
70
+ xHiBuffer = uploadBuffer(device, hi, "idamax-xHi", false);
71
+ xLoBuffer = uploadBuffer(device, lo, "idamax-xLo", false);
59
72
  }
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
73
+ partialsValHiBuffer = createStorageBuffer(
74
+ device,
75
+ 2 * WGS * 4,
76
+ "idamax-partials-val-hi",
77
+ );
78
+ partialsValLoBuffer = createStorageBuffer(
79
+ device,
80
+ 2 * WGS * 4,
81
+ "idamax-partials-val-lo",
82
+ );
83
+ partialsIdxBuffer = createStorageBuffer(
84
+ device,
85
+ 2 * WGS * 4,
86
+ "idamax-partials-idx",
87
+ );
88
+ resultBuffer = createResultBuffer(device, 4, "idamax-result"); // u32 index
64
89
  paramsBuffer = createParamsBuffer(
90
+ device,
65
91
  [
66
92
  { value: n, type: "u32" },
67
93
  { value: incx, type: "u32" },
@@ -69,7 +95,7 @@ export async function idamax(device, n, x, incx) {
69
95
  "idamax-params",
70
96
  );
71
97
 
72
- const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
98
+ const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
73
99
  xHiBuffer,
74
100
  xLoBuffer,
75
101
  partialsValHiBuffer,
@@ -78,27 +104,33 @@ export async function idamax(device, n, x, incx) {
78
104
  paramsBuffer,
79
105
  ]);
80
106
  const { commandEncoder: enc1, ts: ts1 } = runComputePass(
107
+ device,
81
108
  pipelineMain,
82
109
  bgMain,
83
110
  2 * WGS,
84
111
  ); // dispatch 2*WGS workgroups
85
112
 
86
- submit(enc1);
113
+ submit(device, enc1);
87
114
 
88
- const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
89
- partialsValHiBuffer,
90
- partialsValLoBuffer,
91
- partialsIdxBuffer,
92
- resultBuffer,
93
- ]);
115
+ const bgReduce = createBindGroup(
116
+ device,
117
+ pipelineReduce.getBindGroupLayout(0),
118
+ [
119
+ partialsValHiBuffer,
120
+ partialsValLoBuffer,
121
+ partialsIdxBuffer,
122
+ resultBuffer,
123
+ ],
124
+ );
94
125
  const { commandEncoder: enc2, ts: ts2 } = runComputePass(
126
+ device,
95
127
  pipelineReduce,
96
128
  bgReduce,
97
129
  1,
98
130
  ); // dispatch 1 workgroup to reduce the partials to a single index
99
- readBuffer = stageReadback(enc2, resultBuffer);
131
+ readBuffer = stageReadback(device, enc2, resultBuffer);
100
132
 
101
- submit(enc2);
133
+ submit(device, enc2);
102
134
 
103
135
  const resultPromise = extractResult(readBuffer, Uint32Array);
104
136
  readBuffer = null; // ownership transferred — extractResult's own finally destroys it
@@ -122,7 +154,7 @@ export async function idamax(device, n, x, incx) {
122
154
  if (partialsIdxBuffer) destroyBuffers(partialsIdxBuffer);
123
155
  if (resultBuffer) destroyBuffers(resultBuffer);
124
156
  if (paramsBuffer) destroyBuffers(paramsBuffer);
125
- // Only reached if submit(enc2) threw before ownership was transferred above.
157
+ // Only reached if submit(device, enc2) threw before ownership was transferred above.
126
158
  if (readBuffer) destroyBuffers(readBuffer);
127
159
  }
128
160
  }
package/src/init.mjs CHANGED
@@ -1,10 +1,28 @@
1
1
  import { benchmarkMode } from "./util/benchmark.mjs";
2
2
 
3
- let _device = null;
4
- let _adapter = null;
5
- let _gpu = null; // eslint-disable-line no-unused-vars
6
- let _benchmarkEnabled = false;
7
-
3
+ // One WebGPU instance for the whole process, never released. A GPUAdapter
4
+ // yields at most one working device, so a second device needs a second
5
+ // adapter — but creating a second *instance* aborts Dawn during native
6
+ // teardown (std::system_error), whether the instances overlap or are made one
7
+ // after another. So the instance is built once and every adapter comes from
8
+ // it; cleanup() releases devices, not this.
9
+ let _gpu = null;
10
+ let _dumpShaders = false; // instance-level Dawn toggle, fixed when _gpu is made
11
+
12
+ // Resolved-options key -> GPUDevice. init() returns the cached device for a
13
+ // given option set and creates one per distinct set, so a process can drive
14
+ // several GPUs at once (e.g. discrete via "high-performance", integrated via
15
+ // "low-power").
16
+ const _devices = new Map();
17
+ // GPUDevice -> { adapter, benchmark, options }. Benchmark support is a
18
+ // property of the device (its requiredFeatures), not of the library.
19
+ const _meta = new WeakMap();
20
+ // The device from the first init(); what getDevice() returns for callers that
21
+ // never mention one (GpuVector.from(data), GpuMatrix.from(data, ...)).
22
+ let _primary = null;
23
+
24
+ const optionsKey = ({ powerPreference, benchmark }) =>
25
+ `${powerPreference}::${benchmark}`;
8
26
 
9
27
  // ── Public API ───────────────────────────────────────────────────────────────
10
28
 
@@ -13,103 +31,146 @@ export async function init({
13
31
  benchmark = false,
14
32
  dumpShaders = false,
15
33
  } = {}) {
16
- if (_device) {
17
- return _device;
18
- }
34
+ const options = { powerPreference, benchmark, dumpShaders };
35
+ const key = optionsKey(options);
36
+
37
+ // Same options: idempotent, hand back the device already built for them.
38
+ const cached = _devices.get(key);
39
+ if (cached) return cached;
19
40
 
20
- let gpu;
21
41
  // Browser exposes WebGPU natively via navigator.gpu.
22
42
  // Node.js has no navigator, so we polyfill using the "webgpu" npm package which also
23
43
  // injects WebGPU globals (GPUBufferUsage, GPUShaderStage, etc.) into globalThis.
24
- if (typeof window === "undefined") {
25
- const { create, globals } = await import("webgpu");
26
- Object.assign(globalThis, globals);
27
- // dumpShaders forwards Dawn's own debug toggle — prints each pipeline's
28
- // WGSL and compiled backend IR to stderr. Node-only; see index.d.mts.
29
- const toggles = dumpShaders
30
- ? ["enable-dawn-features=dump_shaders,disable_symbol_renaming"]
31
- : [];
32
- gpu = create(toggles);
33
- _gpu = gpu;
34
- } else {
35
- if (dumpShaders)
36
- console.warn("dumpShaders has no effect in the browser — see init()'s docs.");
37
- gpu = navigator.gpu;
44
+ if (!_gpu) {
45
+ if (typeof window === "undefined") {
46
+ const { create, globals } = await import("webgpu");
47
+ Object.assign(globalThis, globals);
48
+ // dumpShaders forwards Dawn's own debug toggle — prints each pipeline's
49
+ // WGSL and compiled backend IR to stderr. Node-only; see index.d.mts.
50
+ const toggles = dumpShaders
51
+ ? ["enable-dawn-features=dump_shaders,disable_symbol_renaming"]
52
+ : [];
53
+ _gpu = create(toggles);
54
+ _dumpShaders = dumpShaders;
55
+ } else {
56
+ if (dumpShaders)
57
+ console.warn(
58
+ "dumpShaders has no effect in the browser — see init()'s docs.",
59
+ );
60
+ _gpu = navigator.gpu;
61
+ }
62
+ } else if (dumpShaders !== _dumpShaders && typeof window === "undefined") {
63
+ // Unlike powerPreference and benchmark, dumpShaders is a toggle on the Dawn
64
+ // instance rather than the device, and the instance is shared, so a later
65
+ // init() cannot change it.
66
+ console.warn(
67
+ `dumpShaders: ${dumpShaders} was requested, but the WebGPU instance was already created with ` +
68
+ `dumpShaders: ${_dumpShaders}. The first init() call fixes this for the process.`,
69
+ );
38
70
  }
39
71
 
40
- if (!gpu) {
72
+ if (!_gpu) {
41
73
  throw new Error("WebGPU not supported in this environment.");
42
74
  }
43
75
 
44
- _adapter =
45
- (await gpu.requestAdapter({ powerPreference })) ??
46
- (await gpu.requestAdapter());
47
- if (!_adapter) {
76
+ // A fresh adapter per device: requesting a device consumes its adapter, so
77
+ // reusing one would hand back an already-lost device.
78
+ const adapter =
79
+ (await _gpu.requestAdapter({ powerPreference })) ??
80
+ (await _gpu.requestAdapter());
81
+ if (!adapter) {
48
82
  throw new Error("No WebGPU adapter found.");
49
83
  }
50
84
 
51
- _benchmarkEnabled = benchmark;
52
- const bmConfig = benchmarkMode(_adapter, benchmark);
85
+ const bmConfig = benchmarkMode(adapter, benchmark);
53
86
  const features = [...(bmConfig.requiredFeatures ?? [])];
54
- _device = await _adapter.requestDevice({ requiredFeatures: features });
87
+ const device = await adapter.requestDevice({ requiredFeatures: features });
55
88
  // Fires for any GPU error not caught by a pushErrorScope/popErrorScope pair — surfaces silent GPU failures to the console.
56
89
  // See: https://developer.mozilla.org/en-US/docs/Web/API/GPUDevice/uncapturederror_event
57
- _device.addEventListener("uncapturederror", (e) => {
90
+ device.addEventListener("uncapturederror", (e) => {
58
91
  console.error("Uncaptured GPU error:", e.error.message);
59
92
  });
60
93
 
61
- return _device;
94
+ // benchmarkMode() drops the feature when the adapter can't do timestamp
95
+ // queries, so record what was actually granted rather than what was asked
96
+ // for — otherwise beginTimestamp() would build a query set on a device that
97
+ // never requested the feature.
98
+ const benchmarkGranted = features.includes("timestamp-query");
99
+ _meta.set(device, { adapter, benchmark: benchmarkGranted, options });
100
+ _devices.set(key, device);
101
+ if (!_primary) _primary = device;
102
+
103
+ return device;
62
104
  }
63
105
 
64
- export function cleanup() {
65
- if (_device) {
66
- _device.destroy();
67
- _device = null;
106
+ export function cleanup(device) {
107
+ if (device === undefined) {
108
+ for (const d of _devices.values()) d.destroy();
109
+ _devices.clear();
110
+ _primary = null;
111
+ return;
68
112
  }
69
- _adapter = null;
70
- _gpu = null;
71
- _benchmarkEnabled = false;
113
+
114
+ // Releasing one device of several. Unknown or already-released devices are a
115
+ // no-op so teardown paths can call this unguarded.
116
+ const meta = _meta.get(device);
117
+ if (!meta) return;
118
+ _devices.delete(optionsKey(meta.options));
119
+ _meta.delete(device);
120
+ device.destroy();
121
+
122
+ // getDevice() must keep answering while any device is left, so promote a
123
+ // survivor when the primary is the one being released.
124
+ if (_primary === device) _primary = _devices.values().next().value ?? null;
72
125
  }
73
126
 
74
- export function gpuName() {
75
- if (!_adapter) {
127
+ export function gpuName(device = _primary) {
128
+ const meta = device && _meta.get(device);
129
+ if (!meta) {
76
130
  throw new Error("WebGPU adapter not initialized — call init() first.");
77
131
  }
78
- const { device, description } = _adapter.info;
132
+ const { device: deviceName, description } = meta.adapter.info;
79
133
  return {
80
134
  description: description || "unknown",
81
- device: device || "unknown",
135
+ device: deviceName || "unknown",
82
136
  };
83
137
  }
84
138
 
85
139
  // ── Library internals (not part of the public API) ───────────────────────────
86
140
 
87
- /** @returns {boolean} whether benchmark mode was enabled in the last `init()` call */
88
- export function isBenchmarkEnabled() {
89
- return _benchmarkEnabled;
141
+ /**
142
+ * Whether benchmark mode is active for `device` — i.e. it was created with
143
+ * `benchmark: true` *and* its adapter actually supports timestamp queries.
144
+ * @param {GPUDevice} [device] - defaults to the first-initialized device
145
+ * @returns {boolean}
146
+ */
147
+ export function isBenchmarkEnabled(device = _primary) {
148
+ return _meta.get(device)?.benchmark ?? false;
90
149
  }
91
150
 
92
151
  /**
93
- * Returns the active `GPUDevice`. Throws if `init()` has not been called.
152
+ * Returns the device from the first `init()` call — the default for callers
153
+ * that don't name one. Throws if `init()` has not been called.
94
154
  * @returns {GPUDevice}
95
- * @throws {Error} if the device is not initialized
155
+ * @throws {Error} if no device is initialized
96
156
  */
97
157
  export function getDevice() {
98
- if (!_device) {
158
+ if (!_primary) {
99
159
  throw new Error("WebGPU device not initialized — call init() first.");
100
160
  }
101
- return _device;
161
+ return _primary;
102
162
  }
103
163
 
104
164
  /**
105
- * Returns the active `GPUAdapter`. Throws if `init()` has not been called.
165
+ * Returns the `GPUAdapter` backing `device`. Throws if it isn't initialized.
166
+ * @param {GPUDevice} [device] - defaults to the first-initialized device
106
167
  * @returns {GPUAdapter}
107
168
  * @throws {Error} if the adapter is not initialized
108
169
  */
109
- export function getAdapter() {
110
- if (!_adapter) {
170
+ export function getAdapter(device = _primary) {
171
+ const meta = device && _meta.get(device);
172
+ if (!meta) {
111
173
  throw new Error("WebGPU adapter not initialized — call init() first.");
112
174
  }
113
- return _adapter;
175
+ return meta.adapter;
114
176
  }
115
-
@@ -1,9 +1,18 @@
1
1
  import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
- * Returns the 0-based index of the element with the largest absolute value.
4
+ * Returns the 0-based index of the element with the largest absolute value:
5
+ * $$\text{index} = \arg\max_{i} |x_i|$$
5
6
  * Ties are broken in favour of the lower index, matching CBLAS behaviour.
6
7
  *
8
+ * **NaN handling.** The search compares with `>`, which is false for NaN, so
9
+ * NaN elements are skipped rather than selected: a vector of all NaN returns
10
+ * `0`. This matches CBLAS except when `x[0]` itself is NaN — CBLAS seeds its
11
+ * running maximum from `x[0]`, and since no comparison against NaN succeeds it
12
+ * returns `0` however large the later elements are, whereas this returns the
13
+ * index of the largest non-NaN element. `+-Infinity` compares normally and is
14
+ * selected as the maximum.
15
+ *
7
16
  * {@includeCode ../../examples/isamax/isamax.js}
8
17
  *
9
18
  * **Browser (standalone HTML):**
@@ -25,9 +34,18 @@ export declare function isamax(
25
34
  ): Promise<{ index: number } | { index: number; gpuTimeMs: number }>;
26
35
 
27
36
  /**
28
- * Returns the 0-based index of the element with the largest absolute value.
37
+ * Returns the 0-based index of the element with the largest absolute value:
38
+ * $$\text{index} = \arg\max_{i} |x_i|$$
29
39
  * Ties are broken in favour of the lower index, matching CBLAS behaviour.
30
40
  *
41
+ * **NaN handling.** The search compares with `>`, which is false for NaN, so
42
+ * NaN elements are skipped rather than selected: a vector of all NaN returns
43
+ * `0`. This matches CBLAS except when `x[0]` itself is NaN — CBLAS seeds its
44
+ * running maximum from `x[0]`, and since no comparison against NaN succeeds it
45
+ * returns `0` however large the later elements are, whereas this returns the
46
+ * index of the largest non-NaN element. `+-Infinity` compares normally and is
47
+ * selected as the maximum.
48
+ *
31
49
  * {@includeCode ../../examples/isamax/gpu.isamax.js}
32
50
  *
33
51
  * @param device - GPUDevice from `init()`