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
@@ -11,14 +11,28 @@ import { extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { GpuVector } from "../classes/GpuVector.mjs";
13
13
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
14
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
14
15
 
15
- export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, incy, layout = "row-major") {
16
+ export async function ssymv(
17
+ device,
18
+ uplo,
19
+ n,
20
+ alpha,
21
+ A,
22
+ lda,
23
+ x,
24
+ incx,
25
+ beta,
26
+ y,
27
+ incy,
28
+ layout = "row-major",
29
+ ) {
16
30
  const xIsGpu = x instanceof GpuVector;
17
31
  const yIsGpu = y instanceof GpuVector;
18
32
  const AIsGpu = A instanceof GpuMatrix;
19
33
 
20
- if (!(device instanceof GPUDevice))
21
- throw new Error("device must be a GPUDevice.");
34
+ requireGpuDevice(device);
35
+ requireSameDevice(device, "ssymv", { A, x, y });
22
36
  if (uplo !== "lower" && uplo !== "upper")
23
37
  throw new Error("uplo must be 'lower' or 'upper'.");
24
38
  if (layout !== "row-major" && layout !== "column-major")
@@ -30,12 +44,10 @@ export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, in
30
44
  !Number.isInteger(lda)
31
45
  )
32
46
  throw new Error("n, incx, incy, and lda must be integers.");
33
- if (typeof alpha !== "number")
34
- throw new Error("alpha must be a number.");
47
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
35
48
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
36
49
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
37
- if (typeof beta !== "number")
38
- throw new Error("beta must be a number.");
50
+ if (typeof beta !== "number") throw new Error("beta must be a number.");
39
51
  if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
40
52
  if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
41
53
  if (incx <= 0 || incy <= 0)
@@ -56,7 +68,9 @@ export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, in
56
68
  if (AIsGpu && !xIsGpu)
57
69
  throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");
58
70
  if (xIsGpu && x._buf === y._buf)
59
- throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");
71
+ throw new Error(
72
+ "x and y must not reference the same GPU buffer when both are GpuVectors.",
73
+ );
60
74
  if (AIsGpu && lda !== A.lda)
61
75
  throw new Error("lda must match A.lda when A is a GpuMatrix.");
62
76
  if (AIsGpu && (A.rows < n || A.cols < n))
@@ -65,9 +79,7 @@ export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, in
65
79
  if (n === 0) return yIsGpu ? {} : { y };
66
80
 
67
81
  if (!AIsGpu && A.length < (n - 1) * lda + n)
68
- throw new Error(
69
- "A does not have enough elements for the given n and lda.",
70
- );
82
+ throw new Error("A does not have enough elements for the given n and lda.");
71
83
  if (x.length < (n - 1) * incx + 1)
72
84
  throw new Error(
73
85
  "x does not have enough elements for the given n and incx.",
@@ -79,7 +91,8 @@ export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, in
79
91
 
80
92
  // GpuMatrix's own layout wins over the argument; A is symmetric, so column-major A reinterpreted row-major just flips which triangle is stored — flip uplo to match.
81
93
  const effLayout = AIsGpu ? A.layout : layout;
82
- const isLower = effLayout === "column-major" ? uplo === "upper" : uplo === "lower";
94
+ const isLower =
95
+ effLayout === "column-major" ? uplo === "upper" : uplo === "lower";
83
96
 
84
97
  const pipeline = await getPipeline(device, "ssymv");
85
98
 
@@ -89,23 +102,24 @@ export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, in
89
102
  let paramsBuffer = null;
90
103
 
91
104
  try {
92
- ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssymv-A", false);
93
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "ssymv-x", false);
94
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "ssymv-y", true);
105
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssymv-A", false);
106
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "ssymv-x", false);
107
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "ssymv-y", true);
95
108
  paramsBuffer = createParamsBuffer(
109
+ device,
96
110
  [
97
- { value: n, type: "u32" },
98
- { value: alpha, type: "f32" },
99
- { value: beta, type: "f32" },
100
- { value: incx, type: "u32" },
101
- { value: incy, type: "u32" },
102
- { value: lda, type: "u32" },
111
+ { value: n, type: "u32" },
112
+ { value: alpha, type: "f32" },
113
+ { value: beta, type: "f32" },
114
+ { value: incx, type: "u32" },
115
+ { value: incy, type: "u32" },
116
+ { value: lda, type: "u32" },
103
117
  { value: isLower ? 0 : 1, type: "u32" },
104
118
  ],
105
119
  "ssymv-params",
106
120
  );
107
121
 
108
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
122
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
109
123
  ABuffer,
110
124
  xBuffer,
111
125
  yBuffer,
@@ -113,10 +127,17 @@ export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, in
113
127
  ]);
114
128
 
115
129
  const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
116
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
117
- const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
130
+ const { commandEncoder, ts } = runComputePass(
131
+ device,
132
+ pipeline,
133
+ bindGroup,
134
+ wgCount,
135
+ );
136
+ const readBuffer = yIsGpu
137
+ ? null
138
+ : stageReadback(device, commandEncoder, yBuffer);
118
139
 
119
- submit(commandEncoder);
140
+ submit(device, commandEncoder);
120
141
 
121
142
  const gpuTimeMs = await extractTimestamp(ts);
122
143
 
@@ -134,4 +155,4 @@ export async function ssymv(device, uplo, n, alpha, A, lda, x, incx, beta, y, in
134
155
  if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
135
156
  if (paramsBuffer) destroyBuffers(paramsBuffer);
136
157
  }
137
- }
158
+ }
@@ -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 rank-1 update A = alpha * x * x^T + A
5
+ * Performs the symmetric rank-1 update $$A \leftarrow \alpha x x^{T} + A$$
6
6
  *
7
7
  * A is an n×n symmetric matrix stored in row-major order, updated in place.
8
8
  * Only the triangle specified by `uplo` is referenced and updated; the other
@@ -40,7 +40,7 @@ export declare function ssyr(
40
40
  ): Promise<{ A: Float32Array; gpuTimeMs?: number }>;
41
41
 
42
42
  /**
43
- * Performs the symmetric rank-1 update A = alpha * x * x^T + A
43
+ * Performs the symmetric rank-1 update $$A \leftarrow \alpha x x^{T} + A$$
44
44
  *
45
45
  * x and A are both kept resident on the GPU. `A`'s own `layout` (set at
46
46
  * `GpuMatrix.from` time) determines the operation — there is no separate
package/src/ssyr/ssyr.mjs CHANGED
@@ -11,21 +11,31 @@ import { extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { GpuVector } from "../classes/GpuVector.mjs";
13
13
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
14
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
14
15
 
15
- export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "row-major") {
16
+ export async function ssyr(
17
+ device,
18
+ uplo,
19
+ n,
20
+ alpha,
21
+ x,
22
+ incx,
23
+ A,
24
+ lda,
25
+ layout = "row-major",
26
+ ) {
16
27
  const xIsGpu = x instanceof GpuVector;
17
28
  const AIsGpu = A instanceof GpuMatrix;
18
29
 
19
- if (!(device instanceof GPUDevice))
20
- throw new Error("device must be a GPUDevice.");
30
+ requireGpuDevice(device);
31
+ requireSameDevice(device, "ssyr", { A, x });
21
32
  if (uplo !== "lower" && uplo !== "upper")
22
33
  throw new Error("uplo must be 'lower' or 'upper'.");
23
34
  if (layout !== "row-major" && layout !== "column-major")
24
35
  throw new Error("layout must be 'row-major' or 'column-major'.");
25
36
  if (!Number.isInteger(n) || !Number.isInteger(incx) || !Number.isInteger(lda))
26
37
  throw new Error("n, incx, and lda must be integers.");
27
- if (typeof alpha !== "number")
28
- throw new Error("alpha must be a number.");
38
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
29
39
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
30
40
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
31
41
  if (incx <= 0) throw new Error("incx must be positive.");
@@ -50,11 +60,14 @@ export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "ro
50
60
  if (!AIsGpu && A.length < (n - 1) * lda + n)
51
61
  throw new Error("A does not have enough elements for the given n and lda.");
52
62
  if (x.length < (n - 1) * incx + 1)
53
- throw new Error("x does not have enough elements for the given n and incx.");
63
+ throw new Error(
64
+ "x does not have enough elements for the given n and incx.",
65
+ );
54
66
 
55
67
  // GpuMatrix's own layout wins over the argument; A is symmetric, so column-major A reinterpreted row-major just flips which triangle is stored — flip uplo to match.
56
68
  const effLayout = AIsGpu ? A.layout : layout;
57
- const isLower = effLayout === "column-major" ? uplo === "upper" : uplo === "lower";
69
+ const isLower =
70
+ effLayout === "column-major" ? uplo === "upper" : uplo === "lower";
58
71
 
59
72
  const pipeline = await getPipeline(device, "ssyr");
60
73
 
@@ -63,20 +76,21 @@ export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "ro
63
76
  let paramsBuffer = null;
64
77
 
65
78
  try {
66
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "ssyr-x", false);
67
- ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyr-A", true);
79
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "ssyr-x", false);
80
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyr-A", true);
68
81
  paramsBuffer = createParamsBuffer(
82
+ device,
69
83
  [
70
- { value: n, type: "u32" },
71
- { value: alpha, type: "f32" },
72
- { value: incx, type: "u32" },
73
- { value: lda, type: "u32" },
84
+ { value: n, type: "u32" },
85
+ { value: alpha, type: "f32" },
86
+ { value: incx, type: "u32" },
87
+ { value: lda, type: "u32" },
74
88
  { value: isLower ? 0 : 1, type: "u32" },
75
89
  ],
76
90
  "ssyr-params",
77
91
  );
78
92
 
79
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
93
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
80
94
  xBuffer,
81
95
  ABuffer,
82
96
  paramsBuffer,
@@ -85,10 +99,17 @@ export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "ro
85
99
  // One workgroup per row of A; clamped to device limit — the shader's
86
100
  // grid-stride loop handles remaining rows when n > dispatch count.
87
101
  const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
88
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
89
- const readBuffer = AIsGpu ? null : stageReadback(commandEncoder, ABuffer);
102
+ const { commandEncoder, ts } = runComputePass(
103
+ device,
104
+ pipeline,
105
+ bindGroup,
106
+ wgCount,
107
+ );
108
+ const readBuffer = AIsGpu
109
+ ? null
110
+ : stageReadback(device, commandEncoder, ABuffer);
90
111
 
91
- submit(commandEncoder);
112
+ submit(device, commandEncoder);
92
113
 
93
114
  const gpuTimeMs = await extractTimestamp(ts);
94
115
 
@@ -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 rank-2 update A = alpha * x * y^T + alpha * y * x^T + A
5
+ * Performs the symmetric rank-2 update $$A \leftarrow \alpha x y^{T} + \alpha y x^{T} + A$$
6
6
  *
7
7
  * A is an n×n symmetric matrix stored in row-major order, updated in place.
8
8
  * Only the triangle specified by `uplo` is referenced and updated; the other
@@ -44,7 +44,7 @@ export declare function ssyr2(
44
44
  ): Promise<{ A: Float32Array; gpuTimeMs?: number }>;
45
45
 
46
46
  /**
47
- * Performs the symmetric rank-2 update A = alpha * x * y^T + alpha * y * x^T + A
47
+ * Performs the symmetric rank-2 update $$A \leftarrow \alpha x y^{T} + \alpha y x^{T} + A$$
48
48
  *
49
49
  * x, y, and A are all kept resident on the GPU. `A`'s own `layout` (set at
50
50
  * `GpuMatrix.from` time) determines the operation — there is no separate
@@ -11,14 +11,27 @@ import { extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { GpuVector } from "../classes/GpuVector.mjs";
13
13
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
14
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
14
15
 
15
- export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, layout = "row-major") {
16
+ export async function ssyr2(
17
+ device,
18
+ uplo,
19
+ n,
20
+ alpha,
21
+ x,
22
+ incx,
23
+ y,
24
+ incy,
25
+ A,
26
+ lda,
27
+ layout = "row-major",
28
+ ) {
16
29
  const xIsGpu = x instanceof GpuVector;
17
30
  const yIsGpu = y instanceof GpuVector;
18
31
  const AIsGpu = A instanceof GpuMatrix;
19
32
 
20
- if (!(device instanceof GPUDevice))
21
- throw new Error("device must be a GPUDevice.");
33
+ requireGpuDevice(device);
34
+ requireSameDevice(device, "ssyr2", { A, x, y });
22
35
  if (uplo !== "lower" && uplo !== "upper")
23
36
  throw new Error("uplo must be 'lower' or 'upper'.");
24
37
  if (layout !== "row-major" && layout !== "column-major")
@@ -30,8 +43,7 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
30
43
  !Number.isInteger(lda)
31
44
  )
32
45
  throw new Error("n, incx, incy, and lda must be integers.");
33
- if (typeof alpha !== "number")
34
- throw new Error("alpha must be a number.");
46
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
35
47
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
36
48
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
37
49
  if (incx <= 0 || incy <= 0)
@@ -56,7 +68,9 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
56
68
  if (AIsGpu && yIsGpu && A._buf === y._buf)
57
69
  throw new Error("A and y must not reference the same GPU buffer.");
58
70
  if (xIsGpu && x._buf === y._buf)
59
- throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");
71
+ throw new Error(
72
+ "x and y must not reference the same GPU buffer when both are GpuVectors.",
73
+ );
60
74
  if (AIsGpu && lda !== A.lda)
61
75
  throw new Error("lda must match A.lda when A is a GpuMatrix.");
62
76
  if (AIsGpu && (A.rows < n || A.cols < n))
@@ -67,13 +81,18 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
67
81
  if (!AIsGpu && A.length < (n - 1) * lda + n)
68
82
  throw new Error("A does not have enough elements for the given n and lda.");
69
83
  if (x.length < (n - 1) * incx + 1)
70
- throw new Error("x does not have enough elements for the given n and incx.");
84
+ throw new Error(
85
+ "x does not have enough elements for the given n and incx.",
86
+ );
71
87
  if (y.length < (n - 1) * incy + 1)
72
- throw new Error("y does not have enough elements for the given n and incy.");
88
+ throw new Error(
89
+ "y does not have enough elements for the given n and incy.",
90
+ );
73
91
 
74
92
  // GpuMatrix's own layout wins over the argument; A is symmetric, so column-major A reinterpreted row-major just flips which triangle is stored (no x/y swap needed — x*y^T+y*x^T is already symmetric under swapping them).
75
93
  const effLayout = AIsGpu ? A.layout : layout;
76
- const isLower = effLayout === "column-major" ? uplo === "upper" : uplo === "lower";
94
+ const isLower =
95
+ effLayout === "column-major" ? uplo === "upper" : uplo === "lower";
77
96
 
78
97
  const pipeline = await getPipeline(device, "ssyr2");
79
98
 
@@ -83,22 +102,23 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
83
102
  let paramsBuffer = null;
84
103
 
85
104
  try {
86
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "ssyr2-x", false);
87
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "ssyr2-y", false);
88
- ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyr2-A", true);
105
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "ssyr2-x", false);
106
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "ssyr2-y", false);
107
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyr2-A", true);
89
108
  paramsBuffer = createParamsBuffer(
109
+ device,
90
110
  [
91
- { value: n, type: "u32" },
92
- { value: alpha, type: "f32" },
93
- { value: incx, type: "u32" },
94
- { value: incy, type: "u32" },
95
- { value: lda, type: "u32" },
111
+ { value: n, type: "u32" },
112
+ { value: alpha, type: "f32" },
113
+ { value: incx, type: "u32" },
114
+ { value: incy, type: "u32" },
115
+ { value: lda, type: "u32" },
96
116
  { value: isLower ? 0 : 1, type: "u32" },
97
117
  ],
98
118
  "ssyr2-params",
99
119
  );
100
120
 
101
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
121
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
102
122
  xBuffer,
103
123
  yBuffer,
104
124
  ABuffer,
@@ -108,10 +128,17 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
108
128
  // One workgroup per row of A; clamped to device limit — the shader's
109
129
  // grid-stride loop handles remaining rows when n > dispatch count.
110
130
  const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
111
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
112
- const readBuffer = AIsGpu ? null : stageReadback(commandEncoder, ABuffer);
131
+ const { commandEncoder, ts } = runComputePass(
132
+ device,
133
+ pipeline,
134
+ bindGroup,
135
+ wgCount,
136
+ );
137
+ const readBuffer = AIsGpu
138
+ ? null
139
+ : stageReadback(device, commandEncoder, ABuffer);
113
140
 
114
- submit(commandEncoder);
141
+ submit(device, commandEncoder);
115
142
 
116
143
  const gpuTimeMs = await extractTimestamp(ts);
117
144
 
@@ -2,7 +2,8 @@ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
2
 
3
3
  /**
4
4
  * Performs the symmetric rank-2k update
5
- * C := uplo(alpha * op(A) * op(B)^T + alpha * op(B) * op(A)^T + beta * C) —
5
+ * $$C \leftarrow \mathrm{uplo}(\alpha \mathrm{op}(A) \mathrm{op}(B)^{T} + \alpha \mathrm{op}(B) \mathrm{op}(A)^{T} + \beta C)$$
6
+ *
6
7
  * only the triangle of C named by `uplo` is read or written (`'lower'`:
7
8
  * `col <= row`, `'upper'`: `col >= row`). C is always n×n.
8
9
  *
@@ -57,7 +58,7 @@ export declare function ssyr2k(
57
58
 
58
59
  /**
59
60
  * Performs the symmetric rank-2k update
60
- * C := uplo(alpha * op(A) * op(B)^T + alpha * op(B) * op(A)^T + beta * C)
61
+ * $$C \leftarrow \mathrm{uplo}(\alpha \mathrm{op}(A) \mathrm{op}(B)^{T} + \alpha \mathrm{op}(B) \mathrm{op}(A)^{T} + \beta C)$$
61
62
  *
62
63
  * A, B, and C are all kept GPU-resident. Each matrix's own `layout` (set at
63
64
  * `GpuMatrix.from` time) determines the operation — there is no separate