wgblas 2.1.0 → 2.2.1

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (99) hide show
  1. package/README.md +2 -0
  2. package/dist/wgblas.browser.js +1016 -43
  3. package/index.d.mts +26 -53
  4. package/index.mjs +11 -0
  5. package/package.json +132 -63
  6. package/src/classes/Complex32.d.mts +43 -0
  7. package/src/classes/Complex32.mjs +82 -0
  8. package/src/classes/Complex64.d.mts +44 -0
  9. package/src/classes/Complex64.mjs +76 -0
  10. package/src/classes/GpuMatrix.d.mts +41 -41
  11. package/src/classes/GpuMatrix.mjs +112 -10
  12. package/src/classes/GpuVector.d.mts +36 -40
  13. package/src/classes/GpuVector.mjs +39 -2
  14. package/src/cscal/cscal.d.mts +47 -0
  15. package/src/cscal/cscal.mjs +98 -0
  16. package/src/dasum/dasum.mjs +31 -15
  17. package/src/daxpy/daxpy.d.mts +56 -0
  18. package/src/daxpy/daxpy.mjs +150 -0
  19. package/src/dcopy/dcopy.d.mts +52 -0
  20. package/src/dcopy/dcopy.mjs +140 -0
  21. package/src/ddot/ddot.d.mts +62 -0
  22. package/src/ddot/ddot.mjs +184 -0
  23. package/src/dnrm2/dnrm2.d.mts +50 -0
  24. package/src/dnrm2/dnrm2.mjs +189 -0
  25. package/src/drot/drot.d.mts +67 -0
  26. package/src/drot/drot.mjs +170 -0
  27. package/src/drotm/drotm.d.mts +67 -0
  28. package/src/drotm/drotm.mjs +171 -0
  29. package/src/dscal/dscal.d.mts +52 -0
  30. package/src/dscal/dscal.mjs +119 -0
  31. package/src/dswap/dswap.d.mts +57 -0
  32. package/src/dswap/dswap.mjs +155 -0
  33. package/src/idamax/idamax.mjs +49 -19
  34. package/src/init.mjs +6 -3
  35. package/src/isamax/isamax.mjs +17 -14
  36. package/src/random/random.d.mts +37 -40
  37. package/src/random/random.mjs +39 -7
  38. package/src/sasum/sasum.mjs +13 -11
  39. package/src/saxpy/saxpy.mjs +9 -8
  40. package/src/scopy/scopy.mjs +8 -6
  41. package/src/sdot/sdot.mjs +13 -11
  42. package/src/sgemm/sgemm.d.mts +2 -2
  43. package/src/sgemm/sgemm.mjs +91 -35
  44. package/src/sgemmtr/sgemmtr.d.mts +3 -2
  45. package/src/sgemmtr/sgemmtr.mjs +92 -35
  46. package/src/sgemv/sgemv.d.mts +2 -2
  47. package/src/sgemv/sgemv.mjs +41 -25
  48. package/src/sger/sger.d.mts +2 -2
  49. package/src/sger/sger.mjs +38 -16
  50. package/src/shaders/__test_pipeline_a.wgsl +3 -0
  51. package/src/shaders/__test_pipeline_b.wgsl +2 -0
  52. package/src/shaders/cscal.wgsl +33 -0
  53. package/src/shaders/daxpy.wgsl +66 -0
  54. package/src/shaders/dcopy.wgsl +34 -0
  55. package/src/shaders/ddot.wgsl +106 -0
  56. package/src/shaders/dnrm2.wgsl +167 -0
  57. package/src/shaders/drot.wgsl +81 -0
  58. package/src/shaders/drotm.wgsl +99 -0
  59. package/src/shaders/dscal.wgsl +60 -0
  60. package/src/shaders/dswap.wgsl +38 -0
  61. package/src/shaders/f64/utils/add.wgsl +6 -0
  62. package/src/shaders/f64/utils/divide.wgsl +45 -0
  63. package/src/shaders/f64/utils/multiply.wgsl +29 -11
  64. package/src/shaders/f64/utils/sqrt.wgsl +45 -0
  65. package/src/shaders/index.mjs +69 -0
  66. package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
  67. package/src/snrm2/snrm2.mjs +20 -14
  68. package/src/srot/srot.mjs +10 -7
  69. package/src/srotm/srotm.mjs +9 -12
  70. package/src/sscal/sscal.mjs +7 -7
  71. package/src/sswap/sswap.mjs +14 -8
  72. package/src/ssymm/ssymm.d.mts +5 -4
  73. package/src/ssymm/ssymm.mjs +142 -55
  74. package/src/ssymv/ssymv.d.mts +2 -2
  75. package/src/ssymv/ssymv.mjs +42 -23
  76. package/src/ssyr/ssyr.d.mts +2 -2
  77. package/src/ssyr/ssyr.mjs +34 -15
  78. package/src/ssyr2/ssyr2.d.mts +2 -2
  79. package/src/ssyr2/ssyr2.mjs +43 -18
  80. package/src/ssyr2k/ssyr2k.d.mts +3 -2
  81. package/src/ssyr2k/ssyr2k.mjs +132 -55
  82. package/src/ssyrk/ssyrk.d.mts +3 -2
  83. package/src/ssyrk/ssyrk.mjs +84 -33
  84. package/src/strmm/strmm.d.mts +5 -4
  85. package/src/strmm/strmm.mjs +153 -54
  86. package/src/strmv/strmv.d.mts +2 -2
  87. package/src/strmv/strmv.mjs +37 -17
  88. package/src/strsm/strsm.d.mts +6 -4
  89. package/src/strsm/strsm.mjs +418 -172
  90. package/src/strsv/strsv.d.mts +5 -3
  91. package/src/strsv/strsv.mjs +82 -31
  92. package/src/util/benchmark.mjs +5 -3
  93. package/src/util/buffer.mjs +33 -12
  94. package/src/util/complex.mjs +87 -0
  95. package/src/util/compute.mjs +14 -8
  96. package/src/util/device.mjs +18 -3
  97. package/src/util/pipeline.mjs +40 -5
  98. package/src/util/workgroup.mjs +23 -6
  99. package/src/shaders/f64add.wgsl +0 -281
@@ -14,14 +14,12 @@ import { getPipeline } from "../util/pipeline.mjs";
14
14
  import { GpuVector } from "../classes/GpuVector.mjs";
15
15
  import { splitDoubleDouble } from "../util/f64.mjs";
16
16
  import { WGS } from "../util/constants.mjs";
17
- import { requireSameDevice } from "../util/device.mjs";
18
-
17
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
19
18
 
20
19
  export async function idamax(device, n, x, incx) {
21
20
  const xIsGpu = x instanceof GpuVector;
22
21
 
23
- if (!(device instanceof GPUDevice))
24
- throw new Error("device must be a GPUDevice.");
22
+ requireGpuDevice(device);
25
23
  requireSameDevice(device, "idamax", { x });
26
24
  if (!Number.isInteger(n) || !Number.isInteger(incx))
27
25
  throw new Error("n and incx must be integers.");
@@ -37,9 +35,22 @@ export async function idamax(device, n, x, incx) {
37
35
  );
38
36
 
39
37
  // Concatenated f64 helpers (WGSL has no #include); ddAbs is unconditional in idamax.wgsl, so x is split as-is.
40
- const f64Deps = ["f64/dekker", "f64/utils/abs", "f64/utils/greater", "f64/utils/equal"];
41
- const pipelineMain = await getPipeline(device, [...f64Deps, "idamax"], "idamax_main");
42
- 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
+ );
43
54
 
44
55
  let xHiBuffer = null;
45
56
  let xLoBuffer = null;
@@ -59,11 +70,24 @@ export async function idamax(device, n, x, incx) {
59
70
  xHiBuffer = uploadBuffer(device, hi, "idamax-xHi", false);
60
71
  xLoBuffer = uploadBuffer(device, lo, "idamax-xLo", false);
61
72
  }
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");
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
+ );
65
88
  resultBuffer = createResultBuffer(device, 4, "idamax-result"); // u32 index
66
- paramsBuffer = createParamsBuffer(device,
89
+ paramsBuffer = createParamsBuffer(
90
+ device,
67
91
  [
68
92
  { value: n, type: "u32" },
69
93
  { value: incx, type: "u32" },
@@ -79,7 +103,8 @@ export async function idamax(device, n, x, incx) {
79
103
  partialsIdxBuffer,
80
104
  paramsBuffer,
81
105
  ]);
82
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
106
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
107
+ device,
83
108
  pipelineMain,
84
109
  bgMain,
85
110
  2 * WGS,
@@ -87,13 +112,18 @@ export async function idamax(device, n, x, incx) {
87
112
 
88
113
  submit(device, enc1);
89
114
 
90
- const bgReduce = createBindGroup(device, pipelineReduce.getBindGroupLayout(0), [
91
- partialsValHiBuffer,
92
- partialsValLoBuffer,
93
- partialsIdxBuffer,
94
- resultBuffer,
95
- ]);
96
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
115
+ const bgReduce = createBindGroup(
116
+ device,
117
+ pipelineReduce.getBindGroupLayout(0),
118
+ [
119
+ partialsValHiBuffer,
120
+ partialsValLoBuffer,
121
+ partialsIdxBuffer,
122
+ resultBuffer,
123
+ ],
124
+ );
125
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
126
+ device,
97
127
  pipelineReduce,
98
128
  bgReduce,
99
129
  1,
package/src/init.mjs CHANGED
@@ -21,7 +21,8 @@ const _meta = new WeakMap();
21
21
  // never mention one (GpuVector.from(data), GpuMatrix.from(data, ...)).
22
22
  let _primary = null;
23
23
 
24
- const optionsKey = ({ powerPreference, benchmark }) => `${powerPreference}::${benchmark}`;
24
+ const optionsKey = ({ powerPreference, benchmark }) =>
25
+ `${powerPreference}::${benchmark}`;
25
26
 
26
27
  // ── Public API ───────────────────────────────────────────────────────────────
27
28
 
@@ -53,7 +54,9 @@ export async function init({
53
54
  _dumpShaders = dumpShaders;
54
55
  } else {
55
56
  if (dumpShaders)
56
- console.warn("dumpShaders has no effect in the browser — see init()'s docs.");
57
+ console.warn(
58
+ "dumpShaders has no effect in the browser — see init()'s docs.",
59
+ );
57
60
  _gpu = navigator.gpu;
58
61
  }
59
62
  } else if (dumpShaders !== _dumpShaders && typeof window === "undefined") {
@@ -62,7 +65,7 @@ export async function init({
62
65
  // init() cannot change it.
63
66
  console.warn(
64
67
  `dumpShaders: ${dumpShaders} was requested, but the WebGPU instance was already created with ` +
65
- `dumpShaders: ${_dumpShaders}. The first init() call fixes this for the process.`,
68
+ `dumpShaders: ${_dumpShaders}. The first init() call fixes this for the process.`,
66
69
  );
67
70
  }
68
71
 
@@ -13,14 +13,12 @@ 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 { WGS } from "../util/constants.mjs";
16
- import { requireSameDevice } from "../util/device.mjs";
17
-
16
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
18
17
 
19
18
  export async function isamax(device, n, x, incx) {
20
19
  const xIsGpu = x instanceof GpuVector;
21
20
 
22
- if (!(device instanceof GPUDevice))
23
- throw new Error("device must be a GPUDevice.");
21
+ requireGpuDevice(device);
24
22
  requireSameDevice(device, "isamax", { x });
25
23
  if (!Number.isInteger(n) || !Number.isInteger(incx))
26
24
  throw new Error("n and incx must be integers.");
@@ -45,16 +43,19 @@ export async function isamax(device, n, x, incx) {
45
43
 
46
44
  try {
47
45
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "isamax-x", false);
48
- partialsValBuffer = createStorageBuffer(device,
46
+ partialsValBuffer = createStorageBuffer(
47
+ device,
49
48
  2 * WGS * 4,
50
49
  "isamax-partials-val",
51
50
  ); //to hold 2*WGS partial max values of f32
52
- partialsIdxBuffer = createStorageBuffer(device,
51
+ partialsIdxBuffer = createStorageBuffer(
52
+ device,
53
53
  2 * WGS * 4,
54
54
  "isamax-partials-idx",
55
55
  ); //to hold 2*WGS partial max indices of u32
56
56
  resultBuffer = createResultBuffer(device, 4, "isamax-result"); // u32 index
57
- paramsBuffer = createParamsBuffer(device,
57
+ paramsBuffer = createParamsBuffer(
58
+ device,
58
59
  [
59
60
  { value: n, type: "u32" },
60
61
  { value: incx, type: "u32" },
@@ -68,7 +69,8 @@ export async function isamax(device, n, x, incx) {
68
69
  partialsIdxBuffer,
69
70
  paramsBuffer,
70
71
  ]);
71
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
72
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
73
+ device,
72
74
  pipelineMain,
73
75
  bgMain,
74
76
  2 * WGS,
@@ -76,12 +78,13 @@ export async function isamax(device, n, x, incx) {
76
78
 
77
79
  submit(device, enc1);
78
80
 
79
- const bgReduce = createBindGroup(device, pipelineReduce.getBindGroupLayout(0), [
80
- partialsValBuffer,
81
- partialsIdxBuffer,
82
- resultBuffer,
83
- ]);
84
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
81
+ const bgReduce = createBindGroup(
82
+ device,
83
+ pipelineReduce.getBindGroupLayout(0),
84
+ [partialsValBuffer, partialsIdxBuffer, resultBuffer],
85
+ );
86
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
87
+ device,
85
88
  pipelineReduce,
86
89
  bgReduce,
87
90
  1,
@@ -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,15 +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
- * console.log(A);
87
- * ```
82
+ * **Column-major storage:**
83
+ * {@includeCode ../../examples/randomtriangularfloat32array-columnmajor/randomtriangularfloat32array-columnmajor.js}
88
84
  * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/random/random.mjs#L13">Source code: random.mjs (L13)</a>
89
85
  * @category Utilities
90
86
  */
@@ -96,4 +92,5 @@ export declare function randomTriangularFloat32Array(
96
92
  high?: number,
97
93
  diagLow?: number,
98
94
  diagHigh?: number,
95
+ layout?: 'row-major' | 'column-major',
99
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
  }
@@ -13,14 +13,12 @@ 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 { WGS } from "../util/constants.mjs";
16
- import { requireSameDevice } from "../util/device.mjs";
17
-
16
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
18
17
 
19
18
  export async function sasum(device, n, x, incx) {
20
19
  const xIsGpu = x instanceof GpuVector;
21
20
 
22
- if (!(device instanceof GPUDevice))
23
- throw new Error("device must be a GPUDevice.");
21
+ requireGpuDevice(device);
24
22
  requireSameDevice(device, "sasum", { x });
25
23
  if (!Number.isInteger(n) || !Number.isInteger(incx))
26
24
  throw new Error("n and incx must be integers.");
@@ -46,7 +44,8 @@ export async function sasum(device, n, x, incx) {
46
44
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sasum-x", false);
47
45
  partialsBuffer = createStorageBuffer(device, 2 * WGS * 4, "sasum-partials"); // 2*WGS partial sums of f32
48
46
  resultBuffer = createResultBuffer(device, 4, "sasum-result"); // final f32 scalar
49
- paramsBuffer = createParamsBuffer(device,
47
+ paramsBuffer = createParamsBuffer(
48
+ device,
50
49
  [
51
50
  { value: n, type: "u32" },
52
51
  { value: incx, type: "u32" },
@@ -59,7 +58,8 @@ export async function sasum(device, n, x, incx) {
59
58
  partialsBuffer,
60
59
  paramsBuffer,
61
60
  ]);
62
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
61
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
62
+ device,
63
63
  pipelineMain,
64
64
  bgMain,
65
65
  2 * WGS,
@@ -67,11 +67,13 @@ export async function sasum(device, n, x, incx) {
67
67
 
68
68
  submit(device, enc1);
69
69
 
70
- const bgReduce = createBindGroup(device, pipelineReduce.getBindGroupLayout(0), [
71
- partialsBuffer,
72
- resultBuffer,
73
- ]);
74
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
70
+ const bgReduce = createBindGroup(
71
+ device,
72
+ pipelineReduce.getBindGroupLayout(0),
73
+ [partialsBuffer, resultBuffer],
74
+ );
75
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
76
+ device,
75
77
  pipelineReduce,
76
78
  bgReduce,
77
79
  1,
@@ -11,14 +11,13 @@ 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 { requireSameDevice } from "../util/device.mjs";
14
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
15
15
 
16
16
  export async function saxpy(device, n, alpha, x, incx, y, incy) {
17
17
  const xIsGpu = x instanceof GpuVector;
18
18
  const yIsGpu = y instanceof GpuVector;
19
19
 
20
- if (!(device instanceof GPUDevice))
21
- throw new Error("device must be a GPUDevice.");
20
+ requireGpuDevice(device);
22
21
  requireSameDevice(device, "saxpy", { x, y });
23
22
  if (
24
23
  !Number.isInteger(n) ||
@@ -26,8 +25,7 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
26
25
  !Number.isInteger(incy)
27
26
  )
28
27
  throw new Error("n, incx, and incy must be integers.");
29
- if (typeof alpha !== "number")
30
- throw new Error("alpha must be a number.");
28
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
31
29
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
32
30
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
33
31
  if (incx <= 0 || incy <= 0)
@@ -60,7 +58,8 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
60
58
  try {
61
59
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "saxpy-x", false);
62
60
  yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "saxpy-y", true);
63
- paramsBuffer = createParamsBuffer(device,
61
+ paramsBuffer = createParamsBuffer(
62
+ device,
64
63
  [
65
64
  { value: n, type: "u32" },
66
65
  { value: alpha, type: "f32" },
@@ -75,7 +74,8 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
75
74
  yBuffer,
76
75
  paramsBuffer,
77
76
  ]);
78
- const { commandEncoder, ts } = runComputePass(device,
77
+ const { commandEncoder, ts } = runComputePass(
78
+ device,
79
79
  pipeline,
80
80
  bindGroup,
81
81
  calcWorkgroups(device, n),
@@ -86,7 +86,8 @@ export async function saxpy(device, n, alpha, x, incx, y, incy) {
86
86
 
87
87
  const gpuTimeMs = await extractTimestamp(ts);
88
88
 
89
- if (yIsGpu && xIsGpu) {
89
+ if (yIsGpu) {
90
+ // xIsGpu === yIsGpu, enforced above
90
91
  if (gpuTimeMs !== undefined) return { gpuTimeMs };
91
92
  return {};
92
93
  }
@@ -11,14 +11,13 @@ 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 { requireSameDevice } from "../util/device.mjs";
14
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
15
15
 
16
16
  export async function scopy(device, n, x, incx, y, incy) {
17
17
  const xIsGpu = x instanceof GpuVector;
18
18
  const yIsGpu = y instanceof GpuVector;
19
19
 
20
- if (!(device instanceof GPUDevice))
21
- throw new Error("device must be a GPUDevice.");
20
+ requireGpuDevice(device);
22
21
  requireSameDevice(device, "scopy", { x, y });
23
22
  if (
24
23
  !Number.isInteger(n) ||
@@ -56,7 +55,8 @@ export async function scopy(device, n, x, incx, y, incy) {
56
55
  try {
57
56
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "scopy-x", false);
58
57
  yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "scopy-y", true);
59
- paramsBuffer = createParamsBuffer(device,
58
+ paramsBuffer = createParamsBuffer(
59
+ device,
60
60
  [
61
61
  { value: n, type: "u32" },
62
62
  { value: incx, type: "u32" },
@@ -70,7 +70,8 @@ export async function scopy(device, n, x, incx, y, incy) {
70
70
  yBuffer,
71
71
  paramsBuffer,
72
72
  ]);
73
- const { commandEncoder, ts } = runComputePass(device,
73
+ const { commandEncoder, ts } = runComputePass(
74
+ device,
74
75
  pipeline,
75
76
  bindGroup,
76
77
  calcWorkgroups(device, n),
@@ -81,7 +82,8 @@ export async function scopy(device, n, x, incx, y, incy) {
81
82
 
82
83
  const gpuTimeMs = await extractTimestamp(ts);
83
84
 
84
- if (yIsGpu && xIsGpu) {
85
+ if (yIsGpu) {
86
+ // xIsGpu === yIsGpu, enforced above
85
87
  if (gpuTimeMs !== undefined) return { gpuTimeMs };
86
88
  return {};
87
89
  }
package/src/sdot/sdot.mjs CHANGED
@@ -13,15 +13,13 @@ 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 { WGS } from "../util/constants.mjs";
16
- import { requireSameDevice } from "../util/device.mjs";
17
-
16
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
18
17
 
19
18
  export async function sdot(device, n, x, incx, y, incy) {
20
19
  const xIsGpu = x instanceof GpuVector;
21
20
  const yIsGpu = y instanceof GpuVector;
22
21
 
23
- if (!(device instanceof GPUDevice))
24
- throw new Error("device must be a GPUDevice.");
22
+ requireGpuDevice(device);
25
23
  requireSameDevice(device, "sdot", { x, y });
26
24
  if (
27
25
  !Number.isInteger(n) ||
@@ -64,7 +62,8 @@ export async function sdot(device, n, x, incx, y, incy) {
64
62
  yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sdot-y", false);
65
63
  partialsBuffer = createStorageBuffer(device, 2 * WGS * 4, "sdot-partials"); //to hold 2*WGS partial sums of f32
66
64
  resultBuffer = createResultBuffer(device, 4, "sdot-result"); //to hold the final float32 dot product
67
- paramsBuffer = createParamsBuffer(device,
65
+ paramsBuffer = createParamsBuffer(
66
+ device,
68
67
  [
69
68
  { value: n, type: "u32" },
70
69
  { value: incx, type: "u32" },
@@ -79,7 +78,8 @@ export async function sdot(device, n, x, incx, y, incy) {
79
78
  partialsBuffer,
80
79
  paramsBuffer,
81
80
  ]);
82
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
81
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
82
+ device,
83
83
  pipelineMain,
84
84
  bgMain,
85
85
  2 * WGS,
@@ -87,11 +87,13 @@ export async function sdot(device, n, x, incx, y, incy) {
87
87
 
88
88
  submit(device, enc1);
89
89
 
90
- const bgReduce = createBindGroup(device, pipelineReduce.getBindGroupLayout(0), [
91
- partialsBuffer,
92
- resultBuffer,
93
- ]);
94
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
90
+ const bgReduce = createBindGroup(
91
+ device,
92
+ pipelineReduce.getBindGroupLayout(0),
93
+ [partialsBuffer, resultBuffer],
94
+ );
95
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
96
+ device,
95
97
  pipelineReduce,
96
98
  bgReduce,
97
99
  1,
@@ -1,7 +1,7 @@
1
1
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
2
 
3
3
  /**
4
- * Performs the matrix-matrix operation C = alpha * op(A) * op(B) + beta * C
4
+ * Performs the matrix-matrix operation $$C \leftarrow \alpha \mathrm{op}(A) \mathrm{op}(B) + \beta C$$
5
5
  *
6
6
  * - `transA/transB='no-transpose'`: op(A) = A (m×k), op(B) = B (k×n)
7
7
  * - `transA/transB='transpose'`: op(A) = A^T, op(B) = B^T
@@ -57,7 +57,7 @@ export declare function sgemm(
57
57
  ): Promise<{ C: Float32Array; gpuTimeMs?: number }>;
58
58
 
59
59
  /**
60
- * Performs the matrix-matrix operation C = alpha * op(A) * op(B) + beta * C
60
+ * Performs the matrix-matrix operation $$C \leftarrow \alpha \mathrm{op}(A) \mathrm{op}(B) + \beta C$$
61
61
  *
62
62
  * A, B, and C are all kept GPU-resident. Each matrix's own `layout` (set at
63
63
  * `GpuMatrix.from` time) determines the operation — there is no separate