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
@@ -12,14 +12,14 @@ import { extractTimestamp } from "../util/benchmark.mjs";
12
12
  import { extractResult } from "../util/result.mjs";
13
13
  import { getPipeline } from "../util/pipeline.mjs";
14
14
  import { GpuVector } from "../classes/GpuVector.mjs";
15
-
16
- const WGS = 64;
15
+ import { WGS } from "../util/constants.mjs";
16
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
17
17
 
18
18
  export async function isamax(device, n, x, incx) {
19
19
  const xIsGpu = x instanceof GpuVector;
20
20
 
21
- if (!(device instanceof GPUDevice))
22
- throw new Error("device must be a GPUDevice.");
21
+ requireGpuDevice(device);
22
+ requireSameDevice(device, "isamax", { x });
23
23
  if (!Number.isInteger(n) || !Number.isInteger(incx))
24
24
  throw new Error("n and incx must be integers.");
25
25
  if (incx <= 0) throw new Error("incx must be positive.");
@@ -42,17 +42,20 @@ export async function isamax(device, n, x, incx) {
42
42
  let readBuffer = null;
43
43
 
44
44
  try {
45
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "isamax-x", false);
45
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "isamax-x", false);
46
46
  partialsValBuffer = createStorageBuffer(
47
+ device,
47
48
  2 * WGS * 4,
48
49
  "isamax-partials-val",
49
50
  ); //to hold 2*WGS partial max values of f32
50
51
  partialsIdxBuffer = createStorageBuffer(
52
+ device,
51
53
  2 * WGS * 4,
52
54
  "isamax-partials-idx",
53
55
  ); //to hold 2*WGS partial max indices of u32
54
- resultBuffer = createResultBuffer(4, "isamax-result"); // u32 index
56
+ resultBuffer = createResultBuffer(device, 4, "isamax-result"); // u32 index
55
57
  paramsBuffer = createParamsBuffer(
58
+ device,
56
59
  [
57
60
  { value: n, type: "u32" },
58
61
  { value: incx, type: "u32" },
@@ -60,33 +63,35 @@ export async function isamax(device, n, x, incx) {
60
63
  "isamax-params",
61
64
  );
62
65
 
63
- const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
66
+ const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
64
67
  xBuffer,
65
68
  partialsValBuffer,
66
69
  partialsIdxBuffer,
67
70
  paramsBuffer,
68
71
  ]);
69
72
  const { commandEncoder: enc1, ts: ts1 } = runComputePass(
73
+ device,
70
74
  pipelineMain,
71
75
  bgMain,
72
76
  2 * WGS,
73
77
  ); //dispatch 2*WGS workgroups
74
78
 
75
- submit(enc1);
79
+ submit(device, enc1);
76
80
 
77
- const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
78
- partialsValBuffer,
79
- partialsIdxBuffer,
80
- resultBuffer,
81
- ]);
81
+ const bgReduce = createBindGroup(
82
+ device,
83
+ pipelineReduce.getBindGroupLayout(0),
84
+ [partialsValBuffer, partialsIdxBuffer, resultBuffer],
85
+ );
82
86
  const { commandEncoder: enc2, ts: ts2 } = runComputePass(
87
+ device,
83
88
  pipelineReduce,
84
89
  bgReduce,
85
90
  1,
86
91
  ); // dispatch 1 workgroup to reduce the partial sums to a single result
87
- readBuffer = stageReadback(enc2, resultBuffer);
92
+ readBuffer = stageReadback(device, enc2, resultBuffer);
88
93
 
89
- submit(enc2);
94
+ submit(device, enc2);
90
95
 
91
96
  const resultPromise = extractResult(readBuffer, Uint32Array);
92
97
  readBuffer = null; // ownership transferred — extractResult's own finally destroys it
@@ -108,7 +113,7 @@ export async function isamax(device, n, x, incx) {
108
113
  if (partialsIdxBuffer) destroyBuffers(partialsIdxBuffer);
109
114
  if (resultBuffer) destroyBuffers(resultBuffer);
110
115
  if (paramsBuffer) destroyBuffers(paramsBuffer);
111
- // Only reached if submit(enc2) threw before ownership was transferred above.
116
+ // Only reached if submit(device, enc2) threw before ownership was transferred above.
112
117
  if (readBuffer) destroyBuffers(readBuffer);
113
118
  }
114
119
  }
@@ -4,22 +4,19 @@
4
4
  * @param n - number of elements
5
5
  * @param low - lower bound (default: -1)
6
6
  * @param high - upper bound (default: 1)
7
+ * @param seed - omit for genuine randomness (`Math.random`, the default);
8
+ * pass any number for a deterministic, reproducible sequence (mulberry32) —
9
+ * useful for regression tests that need "random-looking" data without
10
+ * flaking between runs
7
11
  *
8
- * @example Default range [-1, 1)
9
- * ```js
10
- * import { randomFloat32Array } from "wgblas";
12
+ * **Default range [-1, 1):**
13
+ * {@includeCode ../../examples/randomfloat32array/randomfloat32array.js}
11
14
  *
12
- * const x = randomFloat32Array(4);
13
- * console.log(x); // Float32Array [ -0.42, 0.81, -0.07, 0.55 ]
14
- * ```
15
+ * **Custom range [0, 10):**
16
+ * {@includeCode ../../examples/randomfloat32array-custom/randomfloat32array-custom.js}
15
17
  *
16
- * @example Custom range [0, 10)
17
- * ```js
18
- * import { randomFloat32Array } from "wgblas";
19
- *
20
- * const x = randomFloat32Array(4, 0, 10);
21
- * console.log(x);
22
- * ```
18
+ * **Seeded (deterministic):**
19
+ * {@includeCode ../../examples/randomfloat32array-seeded/randomfloat32array-seeded.js}
23
20
  * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/random/random.mjs#L1">Source code: random.mjs (L1)</a>
24
21
  * @category Utilities
25
22
  */
@@ -27,6 +24,7 @@ export declare function randomFloat32Array(
27
24
  n: number,
28
25
  low?: number,
29
26
  high?: number,
27
+ seed?: number,
30
28
  ): Float32Array;
31
29
 
32
30
  /**
@@ -35,22 +33,19 @@ export declare function randomFloat32Array(
35
33
  * @param n - number of elements
36
34
  * @param low - lower bound (default: -1)
37
35
  * @param high - upper bound (default: 1)
36
+ * @param seed - omit for genuine randomness (`Math.random`, the default);
37
+ * pass any number for a deterministic, reproducible sequence (mulberry32) —
38
+ * useful for regression tests that need "random-looking" data without
39
+ * flaking between runs
38
40
  *
39
- * @example Default range [-1, 1)
40
- * ```js
41
- * import { randomFloat64Array } from "wgblas";
42
- *
43
- * const x = randomFloat64Array(4);
44
- * console.log(x); // Float64Array [ -0.42, 0.81, -0.07, 0.55 ]
45
- * ```
41
+ * **Default range [-1, 1):**
42
+ * {@includeCode ../../examples/randomfloat64array/randomfloat64array.js}
46
43
  *
47
- * @example Custom range [0, 10)
48
- * ```js
49
- * import { randomFloat64Array } from "wgblas";
44
+ * **Custom range [0, 10):**
45
+ * {@includeCode ../../examples/randomfloat64array-custom/randomfloat64array-custom.js}
50
46
  *
51
- * const x = randomFloat64Array(4, 0, 10);
52
- * console.log(x);
53
- * ```
47
+ * **Seeded (deterministic):**
48
+ * {@includeCode ../../examples/randomfloat64array-seeded/randomfloat64array-seeded.js}
54
49
  * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/random/random.mjs#L7">Source code: random.mjs (L7)</a>
55
50
  * @category Utilities
56
51
  */
@@ -58,16 +53,20 @@ export declare function randomFloat64Array(
58
53
  n: number,
59
54
  low?: number,
60
55
  high?: number,
56
+ seed?: number,
61
57
  ): Float64Array;
62
58
 
63
59
  /**
64
- * Returns a Float32Array of n*lda elements, row-major with leading dimension
65
- * `lda`, representing an actual lower- or upper-triangular matrix for a
66
- * triangular routine (`strmv`/`strsv`) example: entries in the `uplo`
67
- * triangle are uniform in `[low, high)`, the n diagonal entries
68
- * (`A[i*lda+i]`) are uniform in `[diagLow, diagHigh)` — kept well away from 0
69
- * so a triangular solve doesn't divide by a near-zero pivot — and every entry
70
- * in the other triangle is 0.
60
+ * Returns a Float32Array of n*lda elements, with leading dimension `lda`
61
+ * under the given `layout`, representing an actual lower- or
62
+ * upper-triangular matrix for a triangular routine (`strmv`/`strsv`)
63
+ * example: entries in the `uplo` triangle are uniform in `[low, high)`, the
64
+ * n diagonal entries are uniform in `[diagLow, diagHigh)` — kept well away
65
+ * from 0 so a triangular solve doesn't divide by a near-zero pivot — and
66
+ * every entry in the other triangle is 0. `uplo` always describes the
67
+ * logical triangle, regardless of storage order: only the flat index each
68
+ * (row, col) maps to changes between layouts (`A[row*lda+col]` for
69
+ * row-major, `A[col*lda+row]` for column-major).
71
70
  *
72
71
  * @param n - matrix order (rows/cols read by the triangular routine)
73
72
  * @param lda - leading dimension; throws if `lda < n`
@@ -76,14 +75,12 @@ export declare function randomFloat64Array(
76
75
  * @param high - upper bound for off-diagonal entries (default: 1)
77
76
  * @param diagLow - lower bound for diagonal entries (default: 5)
78
77
  * @param diagHigh - upper bound for diagonal entries (default: 15)
78
+ * @param layout - `'row-major'` or `'column-major'` storage order (default: `'row-major'`)
79
79
  *
80
- * @example
81
- * ```js
82
- * import { randomTriangularFloat32Array } from "wgblas";
80
+ * {@includeCode ../../examples/randomtriangularfloat32array/randomtriangularfloat32array.js}
83
81
  *
84
- * const n = 4, lda = n;
85
- * const A = randomTriangularFloat32Array(n, lda, "lower");
86
- * ```
82
+ * **Column-major storage:**
83
+ * {@includeCode ../../examples/randomtriangularfloat32array-columnmajor/randomtriangularfloat32array-columnmajor.js}
87
84
  * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/random/random.mjs#L13">Source code: random.mjs (L13)</a>
88
85
  * @category Utilities
89
86
  */
@@ -95,4 +92,5 @@ export declare function randomTriangularFloat32Array(
95
92
  high?: number,
96
93
  diagLow?: number,
97
94
  diagHigh?: number,
95
+ layout?: 'row-major' | 'column-major',
98
96
  ): Float32Array;
@@ -1,28 +1,60 @@
1
- export function randomFloat32Array(n, low = -1, high = 1) {
1
+ // mulberry32: a small, fast, deterministic PRNG — not cryptographic, but
2
+ // good enough distribution for test/benchmark data. Returns a () => number
3
+ // generator producing values in [0, 1), same contract as Math.random.
4
+ function mulberry32(seed) {
5
+ let a = seed >>> 0;
6
+ return function () {
7
+ a = (a + 0x6d2b79f5) | 0;
8
+ let t = Math.imul(a ^ (a >>> 15), 1 | a);
9
+ t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t;
10
+ return ((t ^ (t >>> 14)) >>> 0) / 4294967296;
11
+ };
12
+ }
13
+
14
+ export function randomFloat32Array(n, low = -1, high = 1, seed) {
2
15
  const x = new Float32Array(n);
3
- for (let i = 0; i < n; i++) x[i] = low + Math.random() * (high - low);
16
+ const next = seed === undefined ? Math.random : mulberry32(seed);
17
+ for (let i = 0; i < n; i++) x[i] = low + next() * (high - low);
4
18
  return x;
5
19
  }
6
20
 
7
- export function randomFloat64Array(n, low = -1, high = 1) {
21
+ export function randomFloat64Array(n, low = -1, high = 1, seed) {
8
22
  const x = new Float64Array(n);
9
- for (let i = 0; i < n; i++) x[i] = low + Math.random() * (high - low);
23
+ const next = seed === undefined ? Math.random : mulberry32(seed);
24
+ for (let i = 0; i < n; i++) x[i] = low + next() * (high - low);
10
25
  return x;
11
26
  }
12
27
 
13
- export function randomTriangularFloat32Array(n, lda, uplo = "lower", low = -1, high = 1, diagLow = 5, diagHigh = 15) {
28
+ export function randomTriangularFloat32Array(
29
+ n,
30
+ lda,
31
+ uplo = "lower",
32
+ low = -1,
33
+ high = 1,
34
+ diagLow = 5,
35
+ diagHigh = 15,
36
+ layout = "row-major",
37
+ ) {
14
38
  if (uplo !== "lower" && uplo !== "upper")
15
39
  throw new Error("uplo must be 'lower' or 'upper'.");
40
+ if (layout !== "row-major" && layout !== "column-major")
41
+ throw new Error("layout must be 'row-major' or 'column-major'.");
16
42
  if (lda < n) throw new Error("lda must be >= n.");
17
43
 
44
+ // Row-major stores row i at A[i*lda+j]; column-major stores column j at
45
+ // A[j*lda+i] instead. `uplo` describes the logical triangle (unaffected by
46
+ // storage order) — only which flat index each (i, j) maps to changes.
47
+ const isColMajor = layout === "column-major";
48
+ const idx = (i, j) => (isColMajor ? j * lda + i : i * lda + j);
49
+
18
50
  const A = new Float32Array(n * lda);
19
51
  for (let i = 0; i < n; i++) {
20
52
  for (let j = 0; j < n; j++) {
21
53
  if (i === j) continue;
22
54
  const inTriangle = uplo === "lower" ? j < i : j > i;
23
- if (inTriangle) A[i * lda + j] = low + Math.random() * (high - low);
55
+ if (inTriangle) A[idx(i, j)] = low + Math.random() * (high - low);
24
56
  }
25
- A[i * lda + i] = diagLow + Math.random() * (diagHigh - diagLow);
57
+ A[idx(i, i)] = diagLow + Math.random() * (diagHigh - diagLow);
26
58
  }
27
59
  return A;
28
60
  }
@@ -1,7 +1,7 @@
1
1
  import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
- * Computes the sum of absolute values of a vector: result = sum(|x[i]|)
4
+ * Computes the sum of absolute values of a vector: $$\text{result} = \sum_{i} |x_i|$$
5
5
  *
6
6
  * {@includeCode ../../examples/sasum/sasum.js}
7
7
  *
@@ -24,7 +24,7 @@ export declare function sasum(
24
24
  ): Promise<{ asum: number } | { asum: number; gpuTimeMs: number }>;
25
25
 
26
26
  /**
27
- * Computes the sum of absolute values of a vector: result = sum(|x[i]|)
27
+ * Computes the sum of absolute values of a vector: $$\text{result} = \sum_{i} |x_i|$$
28
28
  *
29
29
  * {@includeCode ../../examples/sasum/gpu.sasum.js}
30
30
  *
@@ -12,14 +12,14 @@ import { extractTimestamp } from "../util/benchmark.mjs";
12
12
  import { extractResult } from "../util/result.mjs";
13
13
  import { getPipeline } from "../util/pipeline.mjs";
14
14
  import { GpuVector } from "../classes/GpuVector.mjs";
15
-
16
- const WGS = 64; // workgroup size
15
+ import { WGS } from "../util/constants.mjs";
16
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
17
17
 
18
18
  export async function sasum(device, n, x, incx) {
19
19
  const xIsGpu = x instanceof GpuVector;
20
20
 
21
- if (!(device instanceof GPUDevice))
22
- throw new Error("device must be a GPUDevice.");
21
+ requireGpuDevice(device);
22
+ requireSameDevice(device, "sasum", { x });
23
23
  if (!Number.isInteger(n) || !Number.isInteger(incx))
24
24
  throw new Error("n and incx must be integers.");
25
25
  if (incx <= 0) throw new Error("incx must be positive.");
@@ -41,10 +41,11 @@ export async function sasum(device, n, x, incx) {
41
41
  let readBuffer = null;
42
42
 
43
43
  try {
44
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sasum-x", false);
45
- partialsBuffer = createStorageBuffer(2 * WGS * 4, "sasum-partials"); // 2*WGS partial sums of f32
46
- resultBuffer = createResultBuffer(4, "sasum-result"); // final f32 scalar
44
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sasum-x", false);
45
+ partialsBuffer = createStorageBuffer(device, 2 * WGS * 4, "sasum-partials"); // 2*WGS partial sums of f32
46
+ resultBuffer = createResultBuffer(device, 4, "sasum-result"); // final f32 scalar
47
47
  paramsBuffer = createParamsBuffer(
48
+ device,
48
49
  [
49
50
  { value: n, type: "u32" },
50
51
  { value: incx, type: "u32" },
@@ -52,31 +53,34 @@ export async function sasum(device, n, x, incx) {
52
53
  "sasum-params",
53
54
  );
54
55
 
55
- const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
56
+ const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
56
57
  xBuffer,
57
58
  partialsBuffer,
58
59
  paramsBuffer,
59
60
  ]);
60
61
  const { commandEncoder: enc1, ts: ts1 } = runComputePass(
62
+ device,
61
63
  pipelineMain,
62
64
  bgMain,
63
65
  2 * WGS,
64
66
  ); // dispatch 2*WGS workgroups
65
67
 
66
- submit(enc1);
68
+ submit(device, enc1);
67
69
 
68
- const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
69
- partialsBuffer,
70
- resultBuffer,
71
- ]);
70
+ const bgReduce = createBindGroup(
71
+ device,
72
+ pipelineReduce.getBindGroupLayout(0),
73
+ [partialsBuffer, resultBuffer],
74
+ );
72
75
  const { commandEncoder: enc2, ts: ts2 } = runComputePass(
76
+ device,
73
77
  pipelineReduce,
74
78
  bgReduce,
75
79
  1,
76
80
  ); // dispatch 1 workgroup to reduce the partial sums to a single result
77
- readBuffer = stageReadback(enc2, resultBuffer);
81
+ readBuffer = stageReadback(device, enc2, resultBuffer);
78
82
 
79
- submit(enc2);
83
+ submit(device, enc2);
80
84
 
81
85
  const resultPromise = extractResult(readBuffer, Float32Array);
82
86
  readBuffer = null; // ownership transferred — extractResult's own finally destroys it
@@ -96,7 +100,7 @@ export async function sasum(device, n, x, incx) {
96
100
  if (partialsBuffer) destroyBuffers(partialsBuffer);
97
101
  if (resultBuffer) destroyBuffers(resultBuffer);
98
102
  if (paramsBuffer) destroyBuffers(paramsBuffer);
99
- // Only reached if submit(enc2) threw before ownership was transferred above.
103
+ // Only reached if submit(device, enc2) threw before ownership was transferred above.
100
104
  if (readBuffer) destroyBuffers(readBuffer);
101
105
  }
102
106
  }
@@ -1,7 +1,7 @@
1
1
  import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
- * Performs the operation y = alpha * x + y
4
+ * Performs the operation $$y \leftarrow \alpha x + y$$
5
5
  *
6
6
  * {@includeCode ../../examples/saxpy/saxpy.js}
7
7
  *
@@ -29,7 +29,7 @@ export declare function saxpy(
29
29
  ): Promise<{ y: Float32Array } | { y: Float32Array; gpuTimeMs: number }>;
30
30
 
31
31
  /**
32
- * Performs the operation y = alpha * x + y
32
+ * Performs the operation $$y \leftarrow \alpha x + y$$
33
33
  *
34
34
  * {@includeCode ../../examples/saxpy/gpu.saxpy.js}
35
35
  *
@@ -11,21 +11,21 @@ import { extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { calcWorkgroups } from "../util/workgroup.mjs";
13
13
  import { GpuVector } from "../classes/GpuVector.mjs";
14
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
14
15
 
15
16
  export async function saxpy(device, n, alpha, x, incx, y, incy) {
16
17
  const xIsGpu = x instanceof GpuVector;
17
18
  const yIsGpu = y instanceof GpuVector;
18
19
 
19
- if (!(device instanceof GPUDevice))
20
- throw new Error("device must be a GPUDevice.");
20
+ requireGpuDevice(device);
21
+ requireSameDevice(device, "saxpy", { x, y });
21
22
  if (
22
23
  !Number.isInteger(n) ||
23
24
  !Number.isInteger(incx) ||
24
25
  !Number.isInteger(incy)
25
26
  )
26
27
  throw new Error("n, incx, and incy must be integers.");
27
- if (typeof alpha !== "number")
28
- throw new Error("alpha must be a number.");
28
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
29
29
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
30
30
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
31
31
  if (incx <= 0 || incy <= 0)
@@ -56,9 +56,10 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
56
56
  let readBuffer = null;
57
57
 
58
58
  try {
59
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "saxpy-x", false);
60
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "saxpy-y", true);
59
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "saxpy-x", false);
60
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "saxpy-y", true);
61
61
  paramsBuffer = createParamsBuffer(
62
+ device,
62
63
  [
63
64
  { value: n, type: "u32" },
64
65
  { value: alpha, type: "f32" },
@@ -68,23 +69,25 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
68
69
  "saxpy-params",
69
70
  );
70
71
 
71
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
72
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
72
73
  xBuffer,
73
74
  yBuffer,
74
75
  paramsBuffer,
75
76
  ]);
76
77
  const { commandEncoder, ts } = runComputePass(
78
+ device,
77
79
  pipeline,
78
80
  bindGroup,
79
- calcWorkgroups(n),
81
+ calcWorkgroups(device, n),
80
82
  );
81
- readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
83
+ readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
82
84
 
83
- submit(commandEncoder);
85
+ submit(device, commandEncoder);
84
86
 
85
87
  const gpuTimeMs = await extractTimestamp(ts);
86
88
 
87
- if (yIsGpu && xIsGpu) {
89
+ if (yIsGpu) {
90
+ // xIsGpu === yIsGpu, enforced above
88
91
  if (gpuTimeMs !== undefined) return { gpuTimeMs };
89
92
  return {};
90
93
  }
@@ -1,7 +1,7 @@
1
1
  import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
- * Performs the operation y = x
4
+ * Performs the operation $$y \leftarrow x$$
5
5
  *
6
6
  * {@includeCode ../../examples/scopy/scopy.js}
7
7
  *
@@ -27,7 +27,7 @@ export declare function scopy(
27
27
  ): Promise<{ y: Float32Array } | { y: Float32Array; gpuTimeMs: number }>;
28
28
 
29
29
  /**
30
- * Performs the operation y = x
30
+ * Performs the operation $$y \leftarrow x$$
31
31
  *
32
32
  * {@includeCode ../../examples/scopy/gpu.scopy.js}
33
33
  *
@@ -11,13 +11,14 @@ import { extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { calcWorkgroups } from "../util/workgroup.mjs";
13
13
  import { GpuVector } from "../classes/GpuVector.mjs";
14
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
14
15
 
15
16
  export async function scopy(device, n, x, incx, y, incy) {
16
17
  const xIsGpu = x instanceof GpuVector;
17
18
  const yIsGpu = y instanceof GpuVector;
18
19
 
19
- if (!(device instanceof GPUDevice))
20
- throw new Error("device must be a GPUDevice.");
20
+ requireGpuDevice(device);
21
+ requireSameDevice(device, "scopy", { x, y });
21
22
  if (
22
23
  !Number.isInteger(n) ||
23
24
  !Number.isInteger(incx) ||
@@ -52,9 +53,10 @@ export async function scopy(device, n, x, incx, y, incy) {
52
53
  let readBuffer = null;
53
54
 
54
55
  try {
55
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "scopy-x", false);
56
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "scopy-y", true);
56
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "scopy-x", false);
57
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "scopy-y", true);
57
58
  paramsBuffer = createParamsBuffer(
59
+ device,
58
60
  [
59
61
  { value: n, type: "u32" },
60
62
  { value: incx, type: "u32" },
@@ -63,23 +65,25 @@ export async function scopy(device, n, x, incx, y, incy) {
63
65
  "scopy-params",
64
66
  );
65
67
 
66
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
68
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
67
69
  xBuffer,
68
70
  yBuffer,
69
71
  paramsBuffer,
70
72
  ]);
71
73
  const { commandEncoder, ts } = runComputePass(
74
+ device,
72
75
  pipeline,
73
76
  bindGroup,
74
- calcWorkgroups(n),
77
+ calcWorkgroups(device, n),
75
78
  );
76
- readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
79
+ readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
77
80
 
78
- submit(commandEncoder);
81
+ submit(device, commandEncoder);
79
82
 
80
83
  const gpuTimeMs = await extractTimestamp(ts);
81
84
 
82
- if (yIsGpu && xIsGpu) {
85
+ if (yIsGpu) {
86
+ // xIsGpu === yIsGpu, enforced above
83
87
  if (gpuTimeMs !== undefined) return { gpuTimeMs };
84
88
  return {};
85
89
  }
@@ -1,7 +1,7 @@
1
1
  import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
- * Computes the dot product of two vectors: result = sum(x[i] * y[i])
4
+ * Computes the dot product of two vectors: $$\text{result} = \sum_{i} x_i y_i$$
5
5
  *
6
6
  * {@includeCode ../../examples/sdot/sdot.js}
7
7
  *
@@ -28,7 +28,7 @@ export declare function sdot(
28
28
  ): Promise<{ dot: number } | { dot: number; gpuTimeMs: number }>;
29
29
 
30
30
  /**
31
- * Computes the dot product of two vectors: result = sum(x[i] * y[i])
31
+ * Computes the dot product of two vectors: $$\text{result} = \sum_{i} x_i y_i$$
32
32
  *
33
33
  * {@includeCode ../../examples/sdot/gpu.sdot.js}
34
34
  *