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
@@ -10,39 +10,58 @@ import { extractResult } from "../util/result.mjs";
10
10
  import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
13
-
14
- const BM_SMALL = 32, BN_SMALL = 32; // sgemmtr_small.wgsl's block tile
15
- const BM_LARGE = 64, BN_LARGE = 64; // sgemmtr_large.wgsl's block tile
16
- const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // same threshold sgemm/sgemmtr/ssyrk use
13
+ import { requireWorkgroupCount } from "../util/workgroup.mjs";
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);
46
+ requireSameDevice(device, "ssyr2k", { A, B, C });
29
47
  if (uplo !== "lower" && uplo !== "upper")
30
48
  throw new Error("uplo must be 'lower' or 'upper'.");
31
49
  if (trans !== "no-transpose" && trans !== "transpose")
32
50
  throw new Error("trans must be 'no-transpose' or 'transpose'.");
33
51
  if (layout !== "row-major" && layout !== "column-major")
34
52
  throw new Error("layout must be 'row-major' or 'column-major'.");
35
- if (typeof alpha !== "number")
36
- throw new Error("alpha must be a number.");
53
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
37
54
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
38
55
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
39
- if (typeof beta !== "number")
40
- throw new Error("beta must be a number.");
56
+ if (typeof beta !== "number") throw new Error("beta must be a number.");
41
57
  if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
42
58
  if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
43
59
  if (
44
- !Number.isInteger(n) || !Number.isInteger(k) ||
45
- !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)
46
65
  )
47
66
  throw new Error("n, k, lda, ldb, and ldc must be integers.");
48
67
  if (!AIsGpu && !(A instanceof Float32Array))
@@ -56,6 +75,8 @@ export async function ssyr2k(
56
75
  if (CIsGpu && (!AIsGpu || !BIsGpu))
57
76
  throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");
58
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.");
59
80
  if (n === 0) return CIsGpu ? {} : { C };
60
81
 
61
82
  const effLayoutA = AIsGpu ? A.layout : layout;
@@ -68,14 +89,19 @@ export async function ssyr2k(
68
89
  const aOuter = trans === "no-transpose" ? aRows : aCols;
69
90
  const aInner = trans === "no-transpose" ? aCols : aRows;
70
91
  if (lda < aInner)
71
- 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
+ );
72
95
  if (AIsGpu) {
73
- 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.");
74
98
  const [aLogRows, aLogCols] = trans === "no-transpose" ? [n, k] : [k, n];
75
99
  if (A.rows < aLogRows || A.cols < aLogCols)
76
100
  throw new Error("A is too small for the given n, k, and trans.");
77
101
  } else if (A.length < (aOuter - 1) * lda + aInner) {
78
- 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
+ );
79
105
  }
80
106
 
81
107
  // B: same shape rule as A — netlib's syr2k shares one trans across both operands.
@@ -84,23 +110,32 @@ export async function ssyr2k(
84
110
  const bOuter = trans === "no-transpose" ? bRows : bCols;
85
111
  const bInner = trans === "no-transpose" ? bCols : bRows;
86
112
  if (ldb < bInner)
87
- 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
+ );
88
116
  if (BIsGpu) {
89
- 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.");
90
119
  const [bLogRows, bLogCols] = trans === "no-transpose" ? [n, k] : [k, n];
91
120
  if (B.rows < bLogRows || B.cols < bLogCols)
92
121
  throw new Error("B is too small for the given n, k, and trans.");
93
122
  } else if (B.length < (bOuter - 1) * ldb + bInner) {
94
- 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
+ );
95
126
  }
96
127
 
97
128
  // C: always n x n symmetric — layout only affects storage order, not size.
98
129
  if (ldc < n) throw new Error("ldc must be >= n.");
99
130
  if (CIsGpu) {
100
- if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
101
- 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.");
102
135
  } else if (C.length < (n - 1) * ldc + n) {
103
- 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
+ );
104
139
  }
105
140
 
106
141
  // Column-major A/B reinterpreted row-major is A^T/B^T — flip trans per operand.
@@ -111,75 +146,118 @@ export async function ssyr2k(
111
146
  if (effLayoutB === "column-major")
112
147
  effTransB = effTransB === "no-transpose" ? "transpose" : "no-transpose";
113
148
 
114
- const uploEff = effLayoutC === "column-major"
115
- ? (uplo === "lower" ? "upper" : "lower")
116
- : uplo;
149
+ const uploEff =
150
+ effLayoutC === "column-major"
151
+ ? uplo === "lower"
152
+ ? "upper"
153
+ : "lower"
154
+ : uplo;
117
155
  const flip = (t) => (t === "no-transpose" ? "transpose" : "no-transpose");
118
156
 
119
157
  // One pass: op(X) as-is, op(Y) transposed. transOwnY is explicit, not inferred.
120
158
  function passShape(transOwnX, X, ldX, transOwnY, Y, ldY) {
121
159
  const transX = transOwnX;
122
160
  const transY = flip(transOwnY);
123
- if (effLayoutC !== "column-major") return { transX, X, ldX, transY, Y, ldY };
124
- 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
+ };
125
171
  }
126
172
 
127
173
  // Shape-based auto-select — see sgemmtr_small.wgsl/sgemmtr_large.wgsl. m=n=n here (square C).
128
174
  const largeWgX = Math.ceil(n / BN_LARGE);
129
175
  const largeWgY = Math.ceil(n / BM_LARGE);
130
176
  const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
131
- const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
177
+ const pipeline = await getPipeline(
178
+ device,
179
+ useLargeTile ? "sgemmtr_large" : "sgemmtr_small",
180
+ );
132
181
  const wgCount = useLargeTile
133
182
  ? {
134
- x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension),
135
- y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension),
136
- }
183
+ x: requireWorkgroupCount(device, largeWgX, "ssyr2k", "x"),
184
+ y: requireWorkgroupCount(device, largeWgY, "ssyr2k", "y"),
185
+ }
137
186
  : {
138
- x: Math.min(Math.ceil(n / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
139
- y: Math.min(Math.ceil(n / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
140
- };
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
+ };
141
200
 
142
- const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyr2k-A", false);
143
- const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "ssyr2k-B", false);
144
- const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "ssyr2k-C", true);
145
- let paramsBuffer1 = null, paramsBuffer2 = null;
201
+ const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyr2k-A", false);
202
+ const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "ssyr2k-B", false);
203
+ const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssyr2k-C", true);
204
+ let paramsBuffer1 = null,
205
+ paramsBuffer2 = null;
146
206
 
147
207
  try {
148
208
  const pass1 = passShape(effTransA, ABuffer, lda, effTransB, BBuffer, ldb);
149
209
  const pass2 = passShape(effTransB, BBuffer, ldb, effTransA, ABuffer, lda);
150
210
 
151
- const makeParams = (p, betaVal) => createParamsBuffer(
152
- [
153
- { value: n, type: "u32" },
154
- { value: n, type: "u32" },
155
- { value: k, type: "u32" },
156
- { value: alpha, type: "f32" },
157
- { value: betaVal, type: "f32" },
158
- { value: p.ldX, type: "u32" },
159
- { value: p.ldY, type: "u32" },
160
- { value: ldc, type: "u32" },
161
- { value: p.transX === "transpose" ? 1 : 0, type: "u32" },
162
- { value: p.transY === "transpose" ? 1 : 0, type: "u32" },
163
- { value: uploEff === "upper" ? 1 : 0, type: "u32" },
164
- ],
165
- "ssyr2k-params",
166
- );
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
+ );
167
229
  paramsBuffer1 = makeParams(pass1, beta);
168
230
  paramsBuffer2 = makeParams(pass2, 1.0);
169
231
 
170
- const bindGroup1 = createBindGroup(pipeline.getBindGroupLayout(0), [pass1.X, pass1.Y, CBuffer, paramsBuffer1]);
171
- const bindGroup2 = createBindGroup(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
+ ]);
172
244
 
173
- const { commandEncoder, querySet } = beginTimedEncoder();
174
- const desc1 = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
175
- const desc2 = querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
245
+ const { commandEncoder, querySet } = beginTimedEncoder(device);
246
+ const desc1 = querySet
247
+ ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
248
+ : undefined;
249
+ const desc2 = querySet
250
+ ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
251
+ : undefined;
176
252
  encodePass(commandEncoder, pipeline, bindGroup1, wgCount, desc1);
177
253
  encodePass(commandEncoder, pipeline, bindGroup2, wgCount, desc2);
178
254
 
179
- const ts = resolveTimestamp(commandEncoder, querySet);
180
- const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
255
+ const ts = resolveTimestamp(device, commandEncoder, querySet);
256
+ const readBuffer = CIsGpu
257
+ ? null
258
+ : stageReadback(device, commandEncoder, CBuffer);
181
259
 
182
- submit(commandEncoder);
260
+ submit(device, commandEncoder);
183
261
 
184
262
  const gpuTimeMs = await extractTimestamp(ts);
185
263
 
@@ -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
@@ -11,36 +11,55 @@ import { extractResult } from "../util/result.mjs";
11
11
  import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
12
12
  import { getPipeline } from "../util/pipeline.mjs";
13
13
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
14
-
15
- const BM_SMALL = 32, BN_SMALL = 32; // sgemmtr_small.wgsl's block tile
16
- const BM_LARGE = 64, BN_LARGE = 64; // sgemmtr_large.wgsl's block tile
17
- const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // same threshold sgemm/sgemmtr use
14
+ import { requireWorkgroupCount } from "../util/workgroup.mjs";
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);
44
+ requireSameDevice(device, "ssyrk", { A, C });
29
45
  if (uplo !== "lower" && uplo !== "upper")
30
46
  throw new Error("uplo must be 'lower' or 'upper'.");
31
47
  if (trans !== "no-transpose" && trans !== "transpose")
32
48
  throw new Error("trans must be 'no-transpose' or 'transpose'.");
33
49
  if (layout !== "row-major" && layout !== "column-major")
34
50
  throw new Error("layout must be 'row-major' or 'column-major'.");
35
- if (typeof alpha !== "number")
36
- throw new Error("alpha must be a number.");
51
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
37
52
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
38
53
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
39
- if (typeof beta !== "number")
40
- throw new Error("beta must be a number.");
54
+ if (typeof beta !== "number") throw new Error("beta must be a number.");
41
55
  if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
42
56
  if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
43
- 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
+ )
44
63
  throw new Error("n, k, lda, and ldc must be integers.");
45
64
  if (!AIsGpu && !(A instanceof Float32Array))
46
65
  throw new Error("A must be a Float32Array or GpuMatrix.");
@@ -51,6 +70,7 @@ export async function ssyrk(
51
70
  if (CIsGpu && !AIsGpu)
52
71
  throw new Error("A must be a GpuMatrix when C is a GpuMatrix.");
53
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.");
54
74
  if (n === 0) return CIsGpu ? {} : { C };
55
75
 
56
76
  // GpuMatrix's own .layout wins over the shared `layout` argument.
@@ -63,23 +83,32 @@ export async function ssyrk(
63
83
  const aOuter = trans === "no-transpose" ? aRows : aCols;
64
84
  const aInner = trans === "no-transpose" ? aCols : aRows;
65
85
  if (lda < aInner)
66
- 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
+ );
67
89
  if (AIsGpu) {
68
- 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.");
69
92
  const [aLogRows, aLogCols] = trans === "no-transpose" ? [n, k] : [k, n];
70
93
  if (A.rows < aLogRows || A.cols < aLogCols)
71
94
  throw new Error("A is too small for the given n, k, and trans.");
72
95
  } else if (A.length < (aOuter - 1) * lda + aInner) {
73
- 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
+ );
74
99
  }
75
100
 
76
101
  // C: always n x n symmetric — layout only affects storage order, not size.
77
102
  if (ldc < n) throw new Error("ldc must be >= n.");
78
103
  if (CIsGpu) {
79
- if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
80
- 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.");
81
108
  } else if (C.length < (n - 1) * ldc + n) {
82
- 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
+ );
83
112
  }
84
113
 
85
114
  // Column-major A reinterpreted row-major is A^T — flip trans, same trick sgemm/sgemmtr use.
@@ -104,22 +133,31 @@ export async function ssyrk(
104
133
  const largeWgY = Math.ceil(n / BM_LARGE);
105
134
  const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
106
135
 
107
- const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
136
+ const pipeline = await getPipeline(
137
+ device,
138
+ useLargeTile ? "sgemmtr_large" : "sgemmtr_small",
139
+ );
108
140
 
109
- const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyrk-A", false);
110
- const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "ssyrk-C", true);
141
+ const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyrk-A", false);
142
+ const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssyrk-C", true);
111
143
  // B := A, but as a genuinely separate buffer — re-uploaded for Float32Array,
112
144
  // GPU-copied (see below) for GpuMatrix, since the caller owns ABuffer.
113
145
  const BBuffer = AIsGpu
114
- ? createStorageBuffer(ABuffer.size, "ssyrk-B", GPUBufferUsage.COPY_DST)
115
- : uploadBuffer(A, "ssyrk-B", false);
146
+ ? createStorageBuffer(
147
+ device,
148
+ ABuffer.size,
149
+ "ssyrk-B",
150
+ GPUBufferUsage.COPY_DST,
151
+ )
152
+ : uploadBuffer(device, A, "ssyrk-B", false);
116
153
  const paramsBuffer = createParamsBuffer(
154
+ device,
117
155
  [
118
- { value: n, type: "u32" }, // gemmtr's m
119
- { value: n, type: "u32" }, // gemmtr's n
120
- { 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" },
121
159
  { value: alpha, type: "f32" },
122
- { value: beta, type: "f32" },
160
+ { value: beta, type: "f32" },
123
161
  { value: lda, type: "u32" }, // gemmtr's lda
124
162
  { value: lda, type: "u32" }, // gemmtr's ldb — B := A, same lda
125
163
  { value: ldc, type: "u32" },
@@ -131,7 +169,7 @@ export async function ssyrk(
131
169
  );
132
170
 
133
171
  try {
134
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
172
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
135
173
  ABuffer,
136
174
  BBuffer,
137
175
  CBuffer,
@@ -140,22 +178,36 @@ export async function ssyrk(
140
178
 
141
179
  const wgCount = useLargeTile
142
180
  ? {
143
- x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension),
144
- y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension),
145
- }
181
+ x: requireWorkgroupCount(device, largeWgX, "ssyrk", "x"),
182
+ y: requireWorkgroupCount(device, largeWgY, "ssyrk", "y"),
183
+ }
146
184
  : {
147
- x: Math.min(Math.ceil(n / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
148
- y: Math.min(Math.ceil(n / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
149
- };
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
+ };
150
198
  // Manual encoder (not runComputePass) so the A->B duplicate copy lands
151
199
  // on the same command encoder, strictly before the compute pass reads B.
152
- const { commandEncoder, querySet, passDescriptor } = beginTimedEncoder();
153
- 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);
154
204
  encodePass(commandEncoder, pipeline, bindGroup, wgCount, passDescriptor);
155
- const ts = resolveTimestamp(commandEncoder, querySet);
156
- const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
205
+ const ts = resolveTimestamp(device, commandEncoder, querySet);
206
+ const readBuffer = CIsGpu
207
+ ? null
208
+ : stageReadback(device, commandEncoder, CBuffer);
157
209
 
158
- submit(commandEncoder);
210
+ submit(device, commandEncoder);
159
211
 
160
212
  const gpuTimeMs = await extractTimestamp(ts);
161
213
 
@@ -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 —