wgblas 2.1.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 (99) hide show
  1. package/README.md +2 -0
  2. package/dist/wgblas.browser.js +1007 -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 +19 -10
  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
@@ -11,21 +11,38 @@ import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
13
13
  import { requireWorkgroupCount } from "../util/workgroup.mjs";
14
- import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
15
- import { requireSameDevice } from "../util/device.mjs";
16
-
14
+ import {
15
+ BM_SMALL,
16
+ BN_SMALL,
17
+ BM_LARGE,
18
+ BN_LARGE,
19
+ LARGE_TILE_WORKGROUP_THRESHOLD,
20
+ } from "../util/constants.mjs";
21
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
17
22
 
18
23
  // ssyr2k: C := uplo(alpha*op(A)*op(B)^T + alpha*op(B)*op(A)^T + beta*C). No
19
24
  // dedicated shader — two sgemmtr passes on one encoder, second with beta=1.
20
25
  export async function ssyr2k(
21
- device, uplo, trans, n, k, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
26
+ device,
27
+ uplo,
28
+ trans,
29
+ n,
30
+ k,
31
+ alpha,
32
+ A,
33
+ lda,
34
+ B,
35
+ ldb,
36
+ beta,
37
+ C,
38
+ ldc,
39
+ layout = "row-major",
22
40
  ) {
23
41
  const AIsGpu = A instanceof GpuMatrix;
24
42
  const BIsGpu = B instanceof GpuMatrix;
25
43
  const CIsGpu = C instanceof GpuMatrix;
26
44
 
27
- if (!(device instanceof GPUDevice))
28
- throw new Error("device must be a GPUDevice.");
45
+ requireGpuDevice(device);
29
46
  requireSameDevice(device, "ssyr2k", { A, B, C });
30
47
  if (uplo !== "lower" && uplo !== "upper")
31
48
  throw new Error("uplo must be 'lower' or 'upper'.");
@@ -33,17 +50,18 @@ export async function ssyr2k(
33
50
  throw new Error("trans must be 'no-transpose' or 'transpose'.");
34
51
  if (layout !== "row-major" && layout !== "column-major")
35
52
  throw new Error("layout must be 'row-major' or 'column-major'.");
36
- if (typeof alpha !== "number")
37
- throw new Error("alpha must be a number.");
53
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
38
54
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
39
55
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
40
- if (typeof beta !== "number")
41
- throw new Error("beta must be a number.");
56
+ if (typeof beta !== "number") throw new Error("beta must be a number.");
42
57
  if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
43
58
  if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
44
59
  if (
45
- !Number.isInteger(n) || !Number.isInteger(k) ||
46
- !Number.isInteger(lda) || !Number.isInteger(ldb) || !Number.isInteger(ldc)
60
+ !Number.isInteger(n) ||
61
+ !Number.isInteger(k) ||
62
+ !Number.isInteger(lda) ||
63
+ !Number.isInteger(ldb) ||
64
+ !Number.isInteger(ldc)
47
65
  )
48
66
  throw new Error("n, k, lda, ldb, and ldc must be integers.");
49
67
  if (!AIsGpu && !(A instanceof Float32Array))
@@ -57,6 +75,8 @@ export async function ssyr2k(
57
75
  if (CIsGpu && (!AIsGpu || !BIsGpu))
58
76
  throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");
59
77
  if (n < 0 || k < 0) throw new Error("n and k must be non-negative.");
78
+ if (lda <= 0 || ldb <= 0 || ldc <= 0)
79
+ throw new Error("lda, ldb, and ldc must be positive.");
60
80
  if (n === 0) return CIsGpu ? {} : { C };
61
81
 
62
82
  const effLayoutA = AIsGpu ? A.layout : layout;
@@ -69,14 +89,19 @@ export async function ssyr2k(
69
89
  const aOuter = trans === "no-transpose" ? aRows : aCols;
70
90
  const aInner = trans === "no-transpose" ? aCols : aRows;
71
91
  if (lda < aInner)
72
- throw new Error(`lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`);
92
+ throw new Error(
93
+ `lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`,
94
+ );
73
95
  if (AIsGpu) {
74
- if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
96
+ if (lda !== A.lda)
97
+ throw new Error("lda must match A.lda when A is a GpuMatrix.");
75
98
  const [aLogRows, aLogCols] = trans === "no-transpose" ? [n, k] : [k, n];
76
99
  if (A.rows < aLogRows || A.cols < aLogCols)
77
100
  throw new Error("A is too small for the given n, k, and trans.");
78
101
  } else if (A.length < (aOuter - 1) * lda + aInner) {
79
- throw new Error("A does not have enough elements for the given dimensions and lda.");
102
+ throw new Error(
103
+ "A does not have enough elements for the given dimensions and lda.",
104
+ );
80
105
  }
81
106
 
82
107
  // B: same shape rule as A — netlib's syr2k shares one trans across both operands.
@@ -85,23 +110,32 @@ export async function ssyr2k(
85
110
  const bOuter = trans === "no-transpose" ? bRows : bCols;
86
111
  const bInner = trans === "no-transpose" ? bCols : bRows;
87
112
  if (ldb < bInner)
88
- throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
113
+ throw new Error(
114
+ `ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
115
+ );
89
116
  if (BIsGpu) {
90
- if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
117
+ if (ldb !== B.lda)
118
+ throw new Error("ldb must match B.lda when B is a GpuMatrix.");
91
119
  const [bLogRows, bLogCols] = trans === "no-transpose" ? [n, k] : [k, n];
92
120
  if (B.rows < bLogRows || B.cols < bLogCols)
93
121
  throw new Error("B is too small for the given n, k, and trans.");
94
122
  } else if (B.length < (bOuter - 1) * ldb + bInner) {
95
- throw new Error("B does not have enough elements for the given dimensions and ldb.");
123
+ throw new Error(
124
+ "B does not have enough elements for the given dimensions and ldb.",
125
+ );
96
126
  }
97
127
 
98
128
  // C: always n x n symmetric — layout only affects storage order, not size.
99
129
  if (ldc < n) throw new Error("ldc must be >= n.");
100
130
  if (CIsGpu) {
101
- if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
102
- if (C.rows < n || C.cols < n) throw new Error("C is too small for the given n.");
131
+ if (ldc !== C.lda)
132
+ throw new Error("ldc must match C.lda when C is a GpuMatrix.");
133
+ if (C.rows < n || C.cols < n)
134
+ throw new Error("C is too small for the given n.");
103
135
  } else if (C.length < (n - 1) * ldc + n) {
104
- throw new Error("C does not have enough elements for the given dimensions and ldc.");
136
+ throw new Error(
137
+ "C does not have enough elements for the given dimensions and ldc.",
138
+ );
105
139
  }
106
140
 
107
141
  // Column-major A/B reinterpreted row-major is A^T/B^T — flip trans per operand.
@@ -112,73 +146,116 @@ export async function ssyr2k(
112
146
  if (effLayoutB === "column-major")
113
147
  effTransB = effTransB === "no-transpose" ? "transpose" : "no-transpose";
114
148
 
115
- const uploEff = effLayoutC === "column-major"
116
- ? (uplo === "lower" ? "upper" : "lower")
117
- : uplo;
149
+ const uploEff =
150
+ effLayoutC === "column-major"
151
+ ? uplo === "lower"
152
+ ? "upper"
153
+ : "lower"
154
+ : uplo;
118
155
  const flip = (t) => (t === "no-transpose" ? "transpose" : "no-transpose");
119
156
 
120
157
  // One pass: op(X) as-is, op(Y) transposed. transOwnY is explicit, not inferred.
121
158
  function passShape(transOwnX, X, ldX, transOwnY, Y, ldY) {
122
159
  const transX = transOwnX;
123
160
  const transY = flip(transOwnY);
124
- if (effLayoutC !== "column-major") return { transX, X, ldX, transY, Y, ldY };
125
- return { transX: flip(transY), X: Y, ldX: ldY, transY: flip(transX), Y: X, ldY: ldX };
161
+ if (effLayoutC !== "column-major")
162
+ return { transX, X, ldX, transY, Y, ldY };
163
+ return {
164
+ transX: flip(transY),
165
+ X: Y,
166
+ ldX: ldY,
167
+ transY: flip(transX),
168
+ Y: X,
169
+ ldY: ldX,
170
+ };
126
171
  }
127
172
 
128
173
  // Shape-based auto-select — see sgemmtr_small.wgsl/sgemmtr_large.wgsl. m=n=n here (square C).
129
174
  const largeWgX = Math.ceil(n / BN_LARGE);
130
175
  const largeWgY = Math.ceil(n / BM_LARGE);
131
176
  const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
132
- const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
177
+ const pipeline = await getPipeline(
178
+ device,
179
+ useLargeTile ? "sgemmtr_large" : "sgemmtr_small",
180
+ );
133
181
  const wgCount = useLargeTile
134
182
  ? {
135
- x: requireWorkgroupCount(device, largeWgX, "ssyr2k", "x"),
136
- y: requireWorkgroupCount(device, largeWgY, "ssyr2k", "y"),
137
- }
183
+ x: requireWorkgroupCount(device, largeWgX, "ssyr2k", "x"),
184
+ y: requireWorkgroupCount(device, largeWgY, "ssyr2k", "y"),
185
+ }
138
186
  : {
139
- x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "ssyr2k", "x"),
140
- y: requireWorkgroupCount(device, Math.ceil(n / BM_SMALL), "ssyr2k", "y"),
141
- };
187
+ x: requireWorkgroupCount(
188
+ device,
189
+ Math.ceil(n / BN_SMALL),
190
+ "ssyr2k",
191
+ "x",
192
+ ),
193
+ y: requireWorkgroupCount(
194
+ device,
195
+ Math.ceil(n / BM_SMALL),
196
+ "ssyr2k",
197
+ "y",
198
+ ),
199
+ };
142
200
 
143
201
  const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyr2k-A", false);
144
202
  const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "ssyr2k-B", false);
145
203
  const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssyr2k-C", true);
146
- let paramsBuffer1 = null, paramsBuffer2 = null;
204
+ let paramsBuffer1 = null,
205
+ paramsBuffer2 = null;
147
206
 
148
207
  try {
149
208
  const pass1 = passShape(effTransA, ABuffer, lda, effTransB, BBuffer, ldb);
150
209
  const pass2 = passShape(effTransB, BBuffer, ldb, effTransA, ABuffer, lda);
151
210
 
152
- const makeParams = (p, betaVal) => createParamsBuffer(device,
153
- [
154
- { value: n, type: "u32" },
155
- { value: n, type: "u32" },
156
- { value: k, type: "u32" },
157
- { value: alpha, type: "f32" },
158
- { value: betaVal, type: "f32" },
159
- { value: p.ldX, type: "u32" },
160
- { value: p.ldY, type: "u32" },
161
- { value: ldc, type: "u32" },
162
- { value: p.transX === "transpose" ? 1 : 0, type: "u32" },
163
- { value: p.transY === "transpose" ? 1 : 0, type: "u32" },
164
- { value: uploEff === "upper" ? 1 : 0, type: "u32" },
165
- ],
166
- "ssyr2k-params",
167
- );
211
+ const makeParams = (p, betaVal) =>
212
+ createParamsBuffer(
213
+ device,
214
+ [
215
+ { value: n, type: "u32" },
216
+ { value: n, type: "u32" },
217
+ { value: k, type: "u32" },
218
+ { value: alpha, type: "f32" },
219
+ { value: betaVal, type: "f32" },
220
+ { value: p.ldX, type: "u32" },
221
+ { value: p.ldY, type: "u32" },
222
+ { value: ldc, type: "u32" },
223
+ { value: p.transX === "transpose" ? 1 : 0, type: "u32" },
224
+ { value: p.transY === "transpose" ? 1 : 0, type: "u32" },
225
+ { value: uploEff === "upper" ? 1 : 0, type: "u32" },
226
+ ],
227
+ "ssyr2k-params",
228
+ );
168
229
  paramsBuffer1 = makeParams(pass1, beta);
169
230
  paramsBuffer2 = makeParams(pass2, 1.0);
170
231
 
171
- const bindGroup1 = createBindGroup(device, pipeline.getBindGroupLayout(0), [pass1.X, pass1.Y, CBuffer, paramsBuffer1]);
172
- const bindGroup2 = createBindGroup(device, pipeline.getBindGroupLayout(0), [pass2.X, pass2.Y, CBuffer, paramsBuffer2]);
232
+ const bindGroup1 = createBindGroup(device, pipeline.getBindGroupLayout(0), [
233
+ pass1.X,
234
+ pass1.Y,
235
+ CBuffer,
236
+ paramsBuffer1,
237
+ ]);
238
+ const bindGroup2 = createBindGroup(device, pipeline.getBindGroupLayout(0), [
239
+ pass2.X,
240
+ pass2.Y,
241
+ CBuffer,
242
+ paramsBuffer2,
243
+ ]);
173
244
 
174
245
  const { commandEncoder, querySet } = beginTimedEncoder(device);
175
- const desc1 = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
176
- const desc2 = querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
246
+ const desc1 = querySet
247
+ ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
248
+ : undefined;
249
+ const desc2 = querySet
250
+ ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
251
+ : undefined;
177
252
  encodePass(commandEncoder, pipeline, bindGroup1, wgCount, desc1);
178
253
  encodePass(commandEncoder, pipeline, bindGroup2, wgCount, desc2);
179
254
 
180
255
  const ts = resolveTimestamp(device, commandEncoder, querySet);
181
- const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
256
+ const readBuffer = CIsGpu
257
+ ? null
258
+ : stageReadback(device, commandEncoder, CBuffer);
182
259
 
183
260
  submit(device, commandEncoder);
184
261
 
@@ -1,7 +1,8 @@
1
1
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
2
 
3
3
  /**
4
- * Performs the symmetric rank-k update C := uplo(alpha * op(A) * op(A)^T + beta * C) —
4
+ * Performs the symmetric rank-k update $$C \leftarrow \mathrm{uplo}(\alpha \mathrm{op}(A) \mathrm{op}(A)^{T} + \beta C)$$
5
+ *
5
6
  * only the triangle of C named by `uplo` is read or written (`'lower'`: `col <= row`,
6
7
  * `'upper'`: `col >= row`). C is always n×n.
7
8
  *
@@ -51,7 +52,7 @@ export declare function ssyrk(
51
52
  ): Promise<{ C: Float32Array; gpuTimeMs?: number }>;
52
53
 
53
54
  /**
54
- * Performs the symmetric rank-k update C := uplo(alpha * op(A) * op(A)^T + beta * C)
55
+ * Performs the symmetric rank-k update $$C \leftarrow \mathrm{uplo}(\alpha \mathrm{op}(A) \mathrm{op}(A)^{T} + \beta C)$$
55
56
  *
56
57
  * A and C are both kept GPU-resident. Each matrix's own `layout` (set at
57
58
  * `GpuMatrix.from` time) determines the operation — there is no separate
@@ -12,20 +12,35 @@ import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
12
12
  import { getPipeline } from "../util/pipeline.mjs";
13
13
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
14
14
  import { requireWorkgroupCount } from "../util/workgroup.mjs";
15
- import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
16
- import { requireSameDevice } from "../util/device.mjs";
17
-
15
+ import {
16
+ BM_SMALL,
17
+ BN_SMALL,
18
+ BM_LARGE,
19
+ BN_LARGE,
20
+ LARGE_TILE_WORKGROUP_THRESHOLD,
21
+ } from "../util/constants.mjs";
22
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
18
23
 
19
24
  // ssyrk: C := uplo(alpha*op(A)*op(A)^T + beta*C). No dedicated shader —
20
25
  // sgemmtr's kernel with A duplicated into a separate B buffer (B := A).
21
26
  export async function ssyrk(
22
- device, uplo, trans, n, k, alpha, A, lda, beta, C, ldc, layout = "row-major",
27
+ device,
28
+ uplo,
29
+ trans,
30
+ n,
31
+ k,
32
+ alpha,
33
+ A,
34
+ lda,
35
+ beta,
36
+ C,
37
+ ldc,
38
+ layout = "row-major",
23
39
  ) {
24
40
  const AIsGpu = A instanceof GpuMatrix;
25
41
  const CIsGpu = C instanceof GpuMatrix;
26
42
 
27
- if (!(device instanceof GPUDevice))
28
- throw new Error("device must be a GPUDevice.");
43
+ requireGpuDevice(device);
29
44
  requireSameDevice(device, "ssyrk", { A, C });
30
45
  if (uplo !== "lower" && uplo !== "upper")
31
46
  throw new Error("uplo must be 'lower' or 'upper'.");
@@ -33,15 +48,18 @@ export async function ssyrk(
33
48
  throw new Error("trans must be 'no-transpose' or 'transpose'.");
34
49
  if (layout !== "row-major" && layout !== "column-major")
35
50
  throw new Error("layout must be 'row-major' or 'column-major'.");
36
- if (typeof alpha !== "number")
37
- throw new Error("alpha must be a number.");
51
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
38
52
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
39
53
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
40
- if (typeof beta !== "number")
41
- throw new Error("beta must be a number.");
54
+ if (typeof beta !== "number") throw new Error("beta must be a number.");
42
55
  if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
43
56
  if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
44
- if (!Number.isInteger(n) || !Number.isInteger(k) || !Number.isInteger(lda) || !Number.isInteger(ldc))
57
+ if (
58
+ !Number.isInteger(n) ||
59
+ !Number.isInteger(k) ||
60
+ !Number.isInteger(lda) ||
61
+ !Number.isInteger(ldc)
62
+ )
45
63
  throw new Error("n, k, lda, and ldc must be integers.");
46
64
  if (!AIsGpu && !(A instanceof Float32Array))
47
65
  throw new Error("A must be a Float32Array or GpuMatrix.");
@@ -52,6 +70,7 @@ export async function ssyrk(
52
70
  if (CIsGpu && !AIsGpu)
53
71
  throw new Error("A must be a GpuMatrix when C is a GpuMatrix.");
54
72
  if (n < 0 || k < 0) throw new Error("n and k must be non-negative.");
73
+ if (lda <= 0 || ldc <= 0) throw new Error("lda and ldc must be positive.");
55
74
  if (n === 0) return CIsGpu ? {} : { C };
56
75
 
57
76
  // GpuMatrix's own .layout wins over the shared `layout` argument.
@@ -64,23 +83,32 @@ export async function ssyrk(
64
83
  const aOuter = trans === "no-transpose" ? aRows : aCols;
65
84
  const aInner = trans === "no-transpose" ? aCols : aRows;
66
85
  if (lda < aInner)
67
- throw new Error(`lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`);
86
+ throw new Error(
87
+ `lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`,
88
+ );
68
89
  if (AIsGpu) {
69
- if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
90
+ if (lda !== A.lda)
91
+ throw new Error("lda must match A.lda when A is a GpuMatrix.");
70
92
  const [aLogRows, aLogCols] = trans === "no-transpose" ? [n, k] : [k, n];
71
93
  if (A.rows < aLogRows || A.cols < aLogCols)
72
94
  throw new Error("A is too small for the given n, k, and trans.");
73
95
  } else if (A.length < (aOuter - 1) * lda + aInner) {
74
- throw new Error("A does not have enough elements for the given dimensions and lda.");
96
+ throw new Error(
97
+ "A does not have enough elements for the given dimensions and lda.",
98
+ );
75
99
  }
76
100
 
77
101
  // C: always n x n symmetric — layout only affects storage order, not size.
78
102
  if (ldc < n) throw new Error("ldc must be >= n.");
79
103
  if (CIsGpu) {
80
- if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
81
- if (C.rows < n || C.cols < n) throw new Error("C is too small for the given n.");
104
+ if (ldc !== C.lda)
105
+ throw new Error("ldc must match C.lda when C is a GpuMatrix.");
106
+ if (C.rows < n || C.cols < n)
107
+ throw new Error("C is too small for the given n.");
82
108
  } else if (C.length < (n - 1) * ldc + n) {
83
- throw new Error("C does not have enough elements for the given dimensions and ldc.");
109
+ throw new Error(
110
+ "C does not have enough elements for the given dimensions and ldc.",
111
+ );
84
112
  }
85
113
 
86
114
  // Column-major A reinterpreted row-major is A^T — flip trans, same trick sgemm/sgemmtr use.
@@ -105,22 +133,31 @@ export async function ssyrk(
105
133
  const largeWgY = Math.ceil(n / BM_LARGE);
106
134
  const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
107
135
 
108
- const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
136
+ const pipeline = await getPipeline(
137
+ device,
138
+ useLargeTile ? "sgemmtr_large" : "sgemmtr_small",
139
+ );
109
140
 
110
141
  const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyrk-A", false);
111
142
  const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssyrk-C", true);
112
143
  // B := A, but as a genuinely separate buffer — re-uploaded for Float32Array,
113
144
  // GPU-copied (see below) for GpuMatrix, since the caller owns ABuffer.
114
145
  const BBuffer = AIsGpu
115
- ? createStorageBuffer(device, ABuffer.size, "ssyrk-B", GPUBufferUsage.COPY_DST)
146
+ ? createStorageBuffer(
147
+ device,
148
+ ABuffer.size,
149
+ "ssyrk-B",
150
+ GPUBufferUsage.COPY_DST,
151
+ )
116
152
  : uploadBuffer(device, A, "ssyrk-B", false);
117
- const paramsBuffer = createParamsBuffer(device,
153
+ const paramsBuffer = createParamsBuffer(
154
+ device,
118
155
  [
119
- { value: n, type: "u32" }, // gemmtr's m
120
- { value: n, type: "u32" }, // gemmtr's n
121
- { value: k, type: "u32" },
156
+ { value: n, type: "u32" }, // gemmtr's m
157
+ { value: n, type: "u32" }, // gemmtr's n
158
+ { value: k, type: "u32" },
122
159
  { value: alpha, type: "f32" },
123
- { value: beta, type: "f32" },
160
+ { value: beta, type: "f32" },
124
161
  { value: lda, type: "u32" }, // gemmtr's lda
125
162
  { value: lda, type: "u32" }, // gemmtr's ldb — B := A, same lda
126
163
  { value: ldc, type: "u32" },
@@ -141,20 +178,34 @@ export async function ssyrk(
141
178
 
142
179
  const wgCount = useLargeTile
143
180
  ? {
144
- x: requireWorkgroupCount(device, largeWgX, "ssyrk", "x"),
145
- y: requireWorkgroupCount(device, largeWgY, "ssyrk", "y"),
146
- }
181
+ x: requireWorkgroupCount(device, largeWgX, "ssyrk", "x"),
182
+ y: requireWorkgroupCount(device, largeWgY, "ssyrk", "y"),
183
+ }
147
184
  : {
148
- x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "ssyrk", "x"),
149
- y: requireWorkgroupCount(device, Math.ceil(n / BM_SMALL), "ssyrk", "y"),
150
- };
185
+ x: requireWorkgroupCount(
186
+ device,
187
+ Math.ceil(n / BN_SMALL),
188
+ "ssyrk",
189
+ "x",
190
+ ),
191
+ y: requireWorkgroupCount(
192
+ device,
193
+ Math.ceil(n / BM_SMALL),
194
+ "ssyrk",
195
+ "y",
196
+ ),
197
+ };
151
198
  // Manual encoder (not runComputePass) so the A->B duplicate copy lands
152
199
  // on the same command encoder, strictly before the compute pass reads B.
153
- const { commandEncoder, querySet, passDescriptor } = beginTimedEncoder(device);
154
- if (AIsGpu) commandEncoder.copyBufferToBuffer(ABuffer, 0, BBuffer, 0, ABuffer.size);
200
+ const { commandEncoder, querySet, passDescriptor } =
201
+ beginTimedEncoder(device);
202
+ if (AIsGpu)
203
+ commandEncoder.copyBufferToBuffer(ABuffer, 0, BBuffer, 0, ABuffer.size);
155
204
  encodePass(commandEncoder, pipeline, bindGroup, wgCount, passDescriptor);
156
205
  const ts = resolveTimestamp(device, commandEncoder, querySet);
157
- const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
206
+ const readBuffer = CIsGpu
207
+ ? null
208
+ : stageReadback(device, commandEncoder, CBuffer);
158
209
 
159
210
  submit(device, commandEncoder);
160
211
 
@@ -2,8 +2,9 @@ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
2
 
3
3
  /**
4
4
  * Performs the triangular matrix-matrix operation
5
- * B := alpha * op(A) * B (`side='left'`) or
6
- * B := alpha * B * op(A) (`side='right'`) — `A` is triangular, only its
5
+ * $$B \leftarrow \alpha \mathrm{op}(A) B \quad (\texttt{side='left'})$$
6
+ * $$B \leftarrow \alpha B \mathrm{op}(A) \quad (\texttt{side='right'})$$
7
+ * `A` is triangular, only its
7
8
  * `uplo` triangle stored; `B` is a general m×n matrix, overwritten in place.
8
9
  *
9
10
  * - `side='left'`: `A` is m×m — `A` premultiplies `B`
@@ -58,8 +59,8 @@ export declare function strmm(
58
59
 
59
60
  /**
60
61
  * Performs the triangular matrix-matrix operation
61
- * B := alpha * op(A) * B (`side='left'`) or
62
- * B := alpha * B * op(A) (`side='right'`)
62
+ * $$B \leftarrow \alpha \mathrm{op}(A) B \quad (\texttt{side='left'})$$
63
+ * $$B \leftarrow \alpha B \mathrm{op}(A) \quad (\texttt{side='right'})$$
63
64
  *
64
65
  * A and B are both kept GPU-resident; B is mutated in place. Each matrix's
65
66
  * own `layout` (set at `GpuMatrix.from` time) determines the operation —