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
package/src/srot/srot.mjs CHANGED
@@ -11,14 +11,13 @@ import { extractResult } from "../util/result.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 srot(device, n, x, incx, y, incy, c, s) {
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, "srot", { x, y });
23
22
  if (
24
23
  !Number.isInteger(n) ||
@@ -28,7 +27,8 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
28
27
  throw new Error("n, incx, and incy must be integers.");
29
28
  if (typeof c !== "number") throw new Error("c must be a number.");
30
29
  if (typeof s !== "number") throw new Error("s must be a number.");
31
- if (Number.isNaN(c) || Number.isNaN(s)) throw new Error("c and s must not be NaN.");
30
+ if (Number.isNaN(c) || Number.isNaN(s))
31
+ throw new Error("c and s must not be NaN.");
32
32
  if (!Number.isFinite(c)) throw new Error("c must be finite.");
33
33
  if (!Number.isFinite(s)) throw new Error("s must be finite.");
34
34
  if (incx <= 0 || incy <= 0)
@@ -62,7 +62,8 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
62
62
  try {
63
63
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "srot-x", true);
64
64
  yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "srot-y", true);
65
- paramsBuffer = createParamsBuffer(device,
65
+ paramsBuffer = createParamsBuffer(
66
+ device,
66
67
  [
67
68
  { value: n, type: "u32" },
68
69
  { value: c, type: "f32" },
@@ -78,7 +79,8 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
78
79
  yBuffer,
79
80
  paramsBuffer,
80
81
  ]);
81
- const { commandEncoder, ts } = runComputePass(device,
82
+ const { commandEncoder, ts } = runComputePass(
83
+ device,
82
84
  pipeline,
83
85
  bindGroup,
84
86
  calcWorkgroups(device, n),
@@ -89,7 +91,8 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
89
91
 
90
92
  const gpuTimeMs = await extractTimestamp(ts);
91
93
 
92
- if (xIsGpu && yIsGpu) {
94
+ if (xIsGpu) {
95
+ // xIsGpu === yIsGpu, enforced above
93
96
  if (gpuTimeMs !== undefined) return { gpuTimeMs };
94
97
  return {};
95
98
  }
@@ -11,14 +11,13 @@ import { extractResult } from "../util/result.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 srotm(device, n, x, incx, y, incy, param) {
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, "srotm", { x, y });
23
22
  if (
24
23
  !Number.isInteger(n) ||
@@ -28,12 +27,7 @@ export async function srotm(device, n, x, incx, y, incy, param) {
28
27
  throw new Error("n, incx, and incy must be integers.");
29
28
  if (!(param instanceof Float32Array) || param.length !== 5)
30
29
  throw new Error("param must be a Float32Array of length 5.");
31
- if (
32
- param[0] !== -2 &&
33
- param[0] !== -1 &&
34
- param[0] !== 0 &&
35
- param[0] !== 1
36
- )
30
+ if (param[0] !== -2 && param[0] !== -1 && param[0] !== 0 && param[0] !== 1)
37
31
  throw new Error("param[0] (flag) must be one of -2, -1, 0, or 1.");
38
32
  if (incx <= 0 || incy <= 0)
39
33
  throw new Error("incx and incy must be positive.");
@@ -68,7 +62,8 @@ export async function srotm(device, n, x, incx, y, incy, param) {
68
62
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "srotm-x", true);
69
63
  yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "srotm-y", true);
70
64
  paramBuffer = uploadBuffer(device, param, "srotm-param", false);
71
- paramsBuffer = createParamsBuffer(device,
65
+ paramsBuffer = createParamsBuffer(
66
+ device,
72
67
  [
73
68
  { value: n, type: "u32" },
74
69
  { value: incx, type: "u32" },
@@ -83,7 +78,8 @@ export async function srotm(device, n, x, incx, y, incy, param) {
83
78
  paramBuffer,
84
79
  paramsBuffer,
85
80
  ]);
86
- const { commandEncoder, ts } = runComputePass(device,
81
+ const { commandEncoder, ts } = runComputePass(
82
+ device,
87
83
  pipeline,
88
84
  bindGroup,
89
85
  calcWorkgroups(device, n),
@@ -94,7 +90,8 @@ export async function srotm(device, n, x, incx, y, incy, param) {
94
90
 
95
91
  const gpuTimeMs = await extractTimestamp(ts);
96
92
 
97
- if (xIsGpu && yIsGpu) {
93
+ if (xIsGpu) {
94
+ // xIsGpu === yIsGpu, enforced above
98
95
  if (gpuTimeMs !== undefined) return { gpuTimeMs };
99
96
  return {};
100
97
  }
@@ -11,18 +11,16 @@ 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 sscal(device, n, alpha, x, incx) {
17
17
  const xIsGpu = x instanceof GpuVector;
18
18
 
19
- if (!(device instanceof GPUDevice))
20
- throw new Error("device must be a GPUDevice.");
19
+ requireGpuDevice(device);
21
20
  requireSameDevice(device, "sscal", { x });
22
21
  if (!Number.isInteger(n) || !Number.isInteger(incx))
23
22
  throw new Error("n and incx must be integers.");
24
- if (typeof alpha !== "number")
25
- throw new Error("alpha must be a number.");
23
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
26
24
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
27
25
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
28
26
  if (incx <= 0) throw new Error("incx must be positive.");
@@ -42,7 +40,8 @@ export async function sscal(device, n, alpha, x, incx) {
42
40
 
43
41
  try {
44
42
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sscal-x", true);
45
- paramsBuffer = createParamsBuffer(device,
43
+ paramsBuffer = createParamsBuffer(
44
+ device,
46
45
  [
47
46
  { value: n, type: "u32" },
48
47
  { value: alpha, type: "f32" },
@@ -55,7 +54,8 @@ export async function sscal(device, n, alpha, x, incx) {
55
54
  xBuffer,
56
55
  paramsBuffer,
57
56
  ]);
58
- const { commandEncoder, ts } = runComputePass(device,
57
+ const { commandEncoder, ts } = runComputePass(
58
+ device,
59
59
  pipeline,
60
60
  bindGroup,
61
61
  calcWorkgroups(device, n),
@@ -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 sswap(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, "sswap", { x, y });
23
22
  if (
24
23
  !Number.isInteger(n) ||
@@ -57,7 +56,8 @@ export async function sswap(device, n, x, incx, y, incy) {
57
56
  try {
58
57
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sswap-x", true);
59
58
  yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sswap-y", true);
60
- paramsBuffer = createParamsBuffer(device,
59
+ paramsBuffer = createParamsBuffer(
60
+ device,
61
61
  [
62
62
  { value: n, type: "u32" },
63
63
  { value: incx, type: "u32" },
@@ -71,19 +71,25 @@ export async function sswap(device, n, x, incx, y, incy) {
71
71
  yBuffer,
72
72
  paramsBuffer,
73
73
  ]);
74
- const { commandEncoder, ts } = runComputePass(device,
74
+ const { commandEncoder, ts } = runComputePass(
75
+ device,
75
76
  pipeline,
76
77
  bindGroup,
77
78
  calcWorkgroups(device, n),
78
79
  );
79
- xReadBuffer = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
80
- yReadBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
80
+ xReadBuffer = xIsGpu
81
+ ? null
82
+ : stageReadback(device, commandEncoder, xBuffer);
83
+ yReadBuffer = yIsGpu
84
+ ? null
85
+ : stageReadback(device, commandEncoder, yBuffer);
81
86
 
82
87
  submit(device, commandEncoder);
83
88
 
84
89
  const gpuTimeMs = await extractTimestamp(ts);
85
90
 
86
- if (xIsGpu && yIsGpu) {
91
+ if (xIsGpu) {
92
+ // xIsGpu === yIsGpu, enforced above (x.constructor !== y.constructor throws)
87
93
  if (gpuTimeMs !== undefined) return { gpuTimeMs };
88
94
  return {};
89
95
  }
@@ -2,8 +2,9 @@ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
2
 
3
3
  /**
4
4
  * Performs the symmetric matrix-matrix operation
5
- * C := alpha * A * B + beta * C (`side='left'`) or
6
- * C := alpha * B * A + beta * C (`side='right'`) — `A` is symmetric, only
5
+ * $$C \leftarrow \alpha A B + \beta C \quad (\texttt{side='left'})$$
6
+ * $$C \leftarrow \alpha B A + \beta C \quad (\texttt{side='right'})$$
7
+ * `A` is symmetric, only
7
8
  * its `uplo` triangle stored; `B` and `C` are general m×n matrices.
8
9
  *
9
10
  * - `side='left'`: `A` is m×m — `A` premultiplies `B`
@@ -59,8 +60,8 @@ export declare function ssymm(
59
60
 
60
61
  /**
61
62
  * Performs the symmetric matrix-matrix operation
62
- * C := alpha * A * B + beta * C (`side='left'`) or
63
- * C := alpha * B * A + beta * C (`side='right'`)
63
+ * $$C \leftarrow \alpha A B + \beta C \quad (\texttt{side='left'})$$
64
+ * $$C \leftarrow \alpha B A + \beta C \quad (\texttt{side='right'})$$
64
65
  *
65
66
  * A, B, and C are all kept GPU-resident. Each matrix's own `layout` (set at
66
67
  * `GpuMatrix.from` time) determines the operation — there is no separate
@@ -13,23 +13,40 @@ import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
13
13
  import { getPipeline } from "../util/pipeline.mjs";
14
14
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
15
15
  import { requireWorkgroupCount } from "../util/workgroup.mjs";
16
- import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
16
+ import {
17
+ BM_SMALL,
18
+ BN_SMALL,
19
+ BM_LARGE,
20
+ BN_LARGE,
21
+ LARGE_TILE_WORKGROUP_THRESHOLD,
22
+ } from "../util/constants.mjs";
17
23
  import { TILE_WG_2D } from "../util/constants.mjs";
18
- import { requireSameDevice } from "../util/device.mjs";
19
-
24
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
20
25
 
21
26
  // ssymm: C := alpha*A*B + beta*C (side='left') or alpha*B*A + beta*C
22
27
  // (side='right'), A symmetric. No fused kernel — symmetrize then sgemm,
23
28
  // both on one command encoder. See symmetrize.wgsl.
24
29
  export async function ssymm(
25
- device, side, uplo, m, n, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
30
+ device,
31
+ side,
32
+ uplo,
33
+ m,
34
+ n,
35
+ alpha,
36
+ A,
37
+ lda,
38
+ B,
39
+ ldb,
40
+ beta,
41
+ C,
42
+ ldc,
43
+ layout = "row-major",
26
44
  ) {
27
45
  const AIsGpu = A instanceof GpuMatrix;
28
46
  const BIsGpu = B instanceof GpuMatrix;
29
47
  const CIsGpu = C instanceof GpuMatrix;
30
48
 
31
- if (!(device instanceof GPUDevice))
32
- throw new Error("device must be a GPUDevice.");
49
+ requireGpuDevice(device);
33
50
  requireSameDevice(device, "ssymm", { A, B, C });
34
51
  if (side !== "left" && side !== "right")
35
52
  throw new Error("side must be 'left' or 'right'.");
@@ -37,17 +54,18 @@ export async function ssymm(
37
54
  throw new Error("uplo must be 'lower' or 'upper'.");
38
55
  if (layout !== "row-major" && layout !== "column-major")
39
56
  throw new Error("layout must be 'row-major' or 'column-major'.");
40
- if (typeof alpha !== "number")
41
- throw new Error("alpha must be a number.");
57
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
42
58
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
43
59
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
44
- if (typeof beta !== "number")
45
- throw new Error("beta must be a number.");
60
+ if (typeof beta !== "number") throw new Error("beta must be a number.");
46
61
  if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
47
62
  if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
48
63
  if (
49
- !Number.isInteger(m) || !Number.isInteger(n) ||
50
- !Number.isInteger(lda) || !Number.isInteger(ldb) || !Number.isInteger(ldc)
64
+ !Number.isInteger(m) ||
65
+ !Number.isInteger(n) ||
66
+ !Number.isInteger(lda) ||
67
+ !Number.isInteger(ldb) ||
68
+ !Number.isInteger(ldc)
51
69
  )
52
70
  throw new Error("m, n, lda, ldb, and ldc must be integers.");
53
71
  if (!AIsGpu && !(A instanceof Float32Array))
@@ -69,48 +87,71 @@ export async function ssymm(
69
87
 
70
88
  // A: symmetric, order = m (side='left') or n (side='right').
71
89
  const aOrder = side === "left" ? m : n;
72
- if (lda < aOrder) throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
90
+ if (lda < aOrder)
91
+ throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
73
92
  if (AIsGpu) {
74
- if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
75
- if (A.rows < aOrder || A.cols < aOrder) throw new Error("A is too small for the given m/n and side.");
93
+ if (lda !== A.lda)
94
+ throw new Error("lda must match A.lda when A is a GpuMatrix.");
95
+ if (A.rows < aOrder || A.cols < aOrder)
96
+ throw new Error("A is too small for the given m/n and side.");
76
97
  } else if (A.length < (aOrder - 1) * lda + aOrder) {
77
- throw new Error("A does not have enough elements for the given dimensions and lda.");
98
+ throw new Error(
99
+ "A does not have enough elements for the given dimensions and lda.",
100
+ );
78
101
  }
79
102
 
80
103
  // B: always m x n, no trans flag — same shape rule as sgemm's C.
81
104
  const bOuter = effLayoutB === "column-major" ? n : m;
82
105
  const bInner = effLayoutB === "column-major" ? m : n;
83
106
  if (ldb < bInner)
84
- throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
107
+ throw new Error(
108
+ `ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
109
+ );
85
110
  if (BIsGpu) {
86
- if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
87
- if (B.rows < m || B.cols < n) throw new Error("B is too small for the given m and n.");
111
+ if (ldb !== B.lda)
112
+ throw new Error("ldb must match B.lda when B is a GpuMatrix.");
113
+ if (B.rows < m || B.cols < n)
114
+ throw new Error("B is too small for the given m and n.");
88
115
  } else if (B.length < (bOuter - 1) * ldb + bInner) {
89
- throw new Error("B does not have enough elements for the given dimensions and ldb.");
116
+ throw new Error(
117
+ "B does not have enough elements for the given dimensions and ldb.",
118
+ );
90
119
  }
91
120
 
92
121
  // C: always m x n.
93
122
  const cOuter = effLayoutC === "column-major" ? n : m;
94
123
  const cInner = effLayoutC === "column-major" ? m : n;
95
124
  if (ldc < cInner)
96
- throw new Error(`ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`);
125
+ throw new Error(
126
+ `ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`,
127
+ );
97
128
  if (CIsGpu) {
98
- if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
99
- if (C.rows < m || C.cols < n) throw new Error("C is too small for the given m and n.");
129
+ if (ldc !== C.lda)
130
+ throw new Error("ldc must match C.lda when C is a GpuMatrix.");
131
+ if (C.rows < m || C.cols < n)
132
+ throw new Error("C is too small for the given m and n.");
100
133
  } else if (C.length < (cOuter - 1) * ldc + cInner) {
101
- throw new Error("C does not have enough elements for the given dimensions and ldc.");
134
+ throw new Error(
135
+ "C does not have enough elements for the given dimensions and ldc.",
136
+ );
102
137
  }
103
138
 
104
139
  // A = A^T, so column-major storage is still A, but the populated
105
140
  // triangle swaps — uplo flips (same reasoning ssyr/ssyrk use).
106
- const uploEffA = effLayoutA === "column-major" ? (uplo === "lower" ? "upper" : "lower") : uplo;
141
+ const uploEffA =
142
+ effLayoutA === "column-major"
143
+ ? uplo === "lower"
144
+ ? "upper"
145
+ : "lower"
146
+ : uplo;
107
147
 
108
148
  const transB = effLayoutB === "column-major" ? "transpose" : "no-transpose";
109
149
  const transDense = "no-transpose"; // Adense is always row-major
110
150
 
111
151
  // X*Y (X=Adense,Y=B for side='left', swapped for 'right'). Column-major
112
152
  // C: compute C^T instead (swap X/Y, flip trans, swap m_g/n_g) — sgemm's own trick.
113
- let mg = m, ng = n;
153
+ let mg = m,
154
+ ng = n;
114
155
  const kg = aOrder;
115
156
  let transX = side === "left" ? transDense : transB;
116
157
  let transY = side === "left" ? transB : transDense;
@@ -126,26 +167,45 @@ export async function ssymm(
126
167
  const largeWgX = Math.ceil(ng / BN_LARGE);
127
168
  const largeWgY = Math.ceil(mg / BM_LARGE);
128
169
  const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
129
- const gemmPipeline = await getPipeline(device, useLargeTile ? "sgemm_large" : "sgemm_small");
170
+ const gemmPipeline = await getPipeline(
171
+ device,
172
+ useLargeTile ? "sgemm_large" : "sgemm_small",
173
+ );
130
174
  const symPipeline = await getPipeline(device, "symmetrize");
131
175
  const gemmWgCount = useLargeTile
132
176
  ? {
133
- x: requireWorkgroupCount(device, largeWgX, "ssymm", "x"),
134
- y: requireWorkgroupCount(device, largeWgY, "ssymm", "y"),
135
- }
177
+ x: requireWorkgroupCount(device, largeWgX, "ssymm", "x"),
178
+ y: requireWorkgroupCount(device, largeWgY, "ssymm", "y"),
179
+ }
136
180
  : {
137
- x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "ssymm", "x"),
138
- y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "ssymm", "y"),
139
- };
181
+ x: requireWorkgroupCount(
182
+ device,
183
+ Math.ceil(ng / BN_SMALL),
184
+ "ssymm",
185
+ "x",
186
+ ),
187
+ y: requireWorkgroupCount(
188
+ device,
189
+ Math.ceil(mg / BM_SMALL),
190
+ "ssymm",
191
+ "y",
192
+ ),
193
+ };
140
194
 
141
195
  const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssymm-A", false);
142
196
  const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "ssymm-B", false);
143
197
  const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssymm-C", true);
144
- const AdenseBuffer = createStorageBuffer(device, aOrder * ldDense * 4, "ssymm-Adense");
145
- let symParams = null, gemmParams = null;
198
+ const AdenseBuffer = createStorageBuffer(
199
+ device,
200
+ aOrder * ldDense * 4,
201
+ "ssymm-Adense",
202
+ );
203
+ let symParams = null,
204
+ gemmParams = null;
146
205
 
147
206
  try {
148
- symParams = createParamsBuffer(device,
207
+ symParams = createParamsBuffer(
208
+ device,
149
209
  [
150
210
  { value: aOrder, type: "u32" },
151
211
  { value: lda, type: "u32" },
@@ -154,7 +214,11 @@ export async function ssymm(
154
214
  ],
155
215
  "ssymm-sym-params",
156
216
  );
157
- const symBindGroup = createBindGroup(device, symPipeline.getBindGroupLayout(0), [ABuffer, AdenseBuffer, symParams]);
217
+ const symBindGroup = createBindGroup(
218
+ device,
219
+ symPipeline.getBindGroupLayout(0),
220
+ [ABuffer, AdenseBuffer, symParams],
221
+ );
158
222
 
159
223
  // X/Y buffers and their own ld, matching swapXY above.
160
224
  const XBuffer = swapXY ? BBuffer : AdenseBuffer;
@@ -162,13 +226,14 @@ export async function ssymm(
162
226
  const YBuffer = swapXY ? AdenseBuffer : BBuffer;
163
227
  const ldY = swapXY ? ldDense : ldb;
164
228
 
165
- gemmParams = createParamsBuffer(device,
229
+ gemmParams = createParamsBuffer(
230
+ device,
166
231
  [
167
- { value: mg, type: "u32" },
168
- { value: ng, type: "u32" },
169
- { value: kg, type: "u32" },
232
+ { value: mg, type: "u32" },
233
+ { value: ng, type: "u32" },
234
+ { value: kg, type: "u32" },
170
235
  { value: alpha, type: "f32" },
171
- { value: beta, type: "f32" },
236
+ { value: beta, type: "f32" },
172
237
  { value: ldX, type: "u32" },
173
238
  { value: ldY, type: "u32" },
174
239
  { value: ldc, type: "u32" },
@@ -177,23 +242,45 @@ export async function ssymm(
177
242
  ],
178
243
  "ssymm-gemm-params",
179
244
  );
180
- const gemmBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
181
- XBuffer,
182
- vec4ViewBinding(device, XBuffer),
183
- YBuffer,
184
- vec4ViewBinding(device, YBuffer),
185
- CBuffer,
186
- gemmParams,
187
- ]);
245
+ const gemmBindGroup = createBindGroup(
246
+ device,
247
+ gemmPipeline.getBindGroupLayout(0),
248
+ [
249
+ XBuffer,
250
+ vec4ViewBinding(device, XBuffer),
251
+ YBuffer,
252
+ vec4ViewBinding(device, YBuffer),
253
+ CBuffer,
254
+ gemmParams,
255
+ ],
256
+ );
188
257
 
189
258
  const { commandEncoder, querySet } = beginTimedEncoder(device);
190
- const symDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
191
- const gemmDesc = querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
192
- encodePass(commandEncoder, symPipeline, symBindGroup, { x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) }, symDesc);
193
- encodePass(commandEncoder, gemmPipeline, gemmBindGroup, gemmWgCount, gemmDesc);
259
+ const symDesc = querySet
260
+ ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
261
+ : undefined;
262
+ const gemmDesc = querySet
263
+ ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
264
+ : undefined;
265
+ encodePass(
266
+ commandEncoder,
267
+ symPipeline,
268
+ symBindGroup,
269
+ { x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) },
270
+ symDesc,
271
+ );
272
+ encodePass(
273
+ commandEncoder,
274
+ gemmPipeline,
275
+ gemmBindGroup,
276
+ gemmWgCount,
277
+ gemmDesc,
278
+ );
194
279
 
195
280
  const ts = resolveTimestamp(device, commandEncoder, querySet);
196
- const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
281
+ const readBuffer = CIsGpu
282
+ ? null
283
+ : stageReadback(device, commandEncoder, CBuffer);
197
284
 
198
285
  submit(device, commandEncoder);
199
286
 
@@ -2,7 +2,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
2
2
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
3
3
 
4
4
  /**
5
- * Performs the symmetric matrix-vector operation y = alpha * A * x + beta * y
5
+ * Performs the symmetric matrix-vector operation $$y \leftarrow \alpha A x + \beta y$$
6
6
  *
7
7
  * A is an n×n symmetric matrix stored in row-major order. Only the triangle
8
8
  * specified by `uplo` is referenced; the other triangle is inferred by symmetry.
@@ -45,7 +45,7 @@ export declare function ssymv(
45
45
  ): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
46
46
 
47
47
  /**
48
- * Performs the symmetric matrix-vector operation y = alpha * A * x + beta * y
48
+ * Performs the symmetric matrix-vector operation $$y \leftarrow \alpha A x + \beta y$$
49
49
  *
50
50
  * x and y are kept resident on the GPU. A must be a GpuMatrix; its own
51
51
  * `layout` (set at `GpuMatrix.from` time) determines the operation — there is