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
@@ -13,19 +13,37 @@ import { 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";
17
- import { requireSameDevice } from "../util/device.mjs";
18
-
16
+ import {
17
+ BM_SMALL,
18
+ BN_SMALL,
19
+ BM_LARGE,
20
+ BN_LARGE,
21
+ LARGE_TILE_WORKGROUP_THRESHOLD,
22
+ } from "../util/constants.mjs";
23
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
19
24
 
20
25
  export async function sgemm(
21
- device, transA, transB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
26
+ device,
27
+ transA,
28
+ transB,
29
+ m,
30
+ n,
31
+ k,
32
+ alpha,
33
+ A,
34
+ lda,
35
+ B,
36
+ ldb,
37
+ beta,
38
+ C,
39
+ ldc,
40
+ layout = "row-major",
22
41
  ) {
23
42
  let AIsGpu = A instanceof GpuMatrix;
24
43
  let BIsGpu = B instanceof GpuMatrix;
25
44
  const CIsGpu = C instanceof GpuMatrix;
26
45
 
27
- if (!(device instanceof GPUDevice))
28
- throw new Error("device must be a GPUDevice.");
46
+ requireGpuDevice(device);
29
47
  requireSameDevice(device, "sgemm", { A, B, C });
30
48
  if (transA !== "no-transpose" && transA !== "transpose")
31
49
  throw new Error("transA must be 'no-transpose' or 'transpose'.");
@@ -33,12 +51,10 @@ export async function sgemm(
33
51
  throw new Error("transB must be 'no-transpose' or 'transpose'.");
34
52
  if (layout !== "row-major" && layout !== "column-major")
35
53
  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.");
54
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
38
55
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
39
56
  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.");
57
+ if (typeof beta !== "number") throw new Error("beta must be a number.");
42
58
  if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
43
59
  if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
44
60
  if (
@@ -60,7 +76,10 @@ export async function sgemm(
60
76
  throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");
61
77
  if (CIsGpu && (!AIsGpu || !BIsGpu))
62
78
  throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");
63
- if (m < 0 || n < 0 || k < 0) throw new Error("m, n, and k must be non-negative.");
79
+ if (m < 0 || n < 0 || k < 0)
80
+ throw new Error("m, n, and k must be non-negative.");
81
+ if (lda <= 0 || ldb <= 0 || ldc <= 0)
82
+ throw new Error("lda, ldb, and ldc must be positive.");
64
83
  if (m === 0 || n === 0) return CIsGpu ? {} : { C };
65
84
 
66
85
  // GpuMatrix's own .layout wins over the shared `layout` argument.
@@ -75,14 +94,19 @@ export async function sgemm(
75
94
  const aOuter = transA === "no-transpose" ? aRows : aCols; // # of stored chunks
76
95
  const aInner = transA === "no-transpose" ? aCols : aRows; // required chunk length
77
96
  if (lda < aInner)
78
- throw new Error(`lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`);
97
+ throw new Error(
98
+ `lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`,
99
+ );
79
100
  if (AIsGpu) {
80
- if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
101
+ if (lda !== A.lda)
102
+ throw new Error("lda must match A.lda when A is a GpuMatrix.");
81
103
  const [aLogRows, aLogCols] = transA === "no-transpose" ? [m, k] : [k, m];
82
104
  if (A.rows < aLogRows || A.cols < aLogCols)
83
105
  throw new Error("A is too small for the given m, k, and transA.");
84
106
  } else if (A.length < (aOuter - 1) * lda + aInner) {
85
- throw new Error("A does not have enough elements for the given dimensions and lda.");
107
+ throw new Error(
108
+ "A does not have enough elements for the given dimensions and lda.",
109
+ );
86
110
  }
87
111
 
88
112
  // B: same reasoning as A, with op(B) = k x n.
@@ -91,26 +115,37 @@ export async function sgemm(
91
115
  const bOuter = transB === "no-transpose" ? bRows : bCols;
92
116
  const bInner = transB === "no-transpose" ? bCols : bRows;
93
117
  if (ldb < bInner)
94
- throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
118
+ throw new Error(
119
+ `ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
120
+ );
95
121
  if (BIsGpu) {
96
- if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
122
+ if (ldb !== B.lda)
123
+ throw new Error("ldb must match B.lda when B is a GpuMatrix.");
97
124
  const [bLogRows, bLogCols] = transB === "no-transpose" ? [k, n] : [n, k];
98
125
  if (B.rows < bLogRows || B.cols < bLogCols)
99
126
  throw new Error("B is too small for the given n, k, and transB.");
100
127
  } else if (B.length < (bOuter - 1) * ldb + bInner) {
101
- throw new Error("B does not have enough elements for the given dimensions and ldb.");
128
+ throw new Error(
129
+ "B does not have enough elements for the given dimensions and ldb.",
130
+ );
102
131
  }
103
132
 
104
133
  // C: always m x n (no trans flag) — layout only affects lda/storage order.
105
134
  const cOuter = effLayoutC === "column-major" ? n : m;
106
135
  const cInner = effLayoutC === "column-major" ? m : n;
107
136
  if (ldc < cInner)
108
- throw new Error(`ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`);
137
+ throw new Error(
138
+ `ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`,
139
+ );
109
140
  if (CIsGpu) {
110
- if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
111
- if (C.rows < m || C.cols < n) throw new Error("C is too small for the given m and n.");
141
+ if (ldc !== C.lda)
142
+ throw new Error("ldc must match C.lda when C is a GpuMatrix.");
143
+ if (C.rows < m || C.cols < n)
144
+ throw new Error("C is too small for the given m and n.");
112
145
  } else if (C.length < (cOuter - 1) * ldc + cInner) {
113
- throw new Error("C does not have enough elements for the given dimensions and ldc.");
146
+ throw new Error(
147
+ "C does not have enough elements for the given dimensions and ldc.",
148
+ );
114
149
  }
115
150
 
116
151
  // Column-major A/B reinterpreted row-major is A^T/B^T — flip the trans flag.
@@ -136,7 +171,10 @@ export async function sgemm(
136
171
  const largeWgY = Math.ceil(m / BM_LARGE);
137
172
  const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
138
173
 
139
- const pipeline = await getPipeline(device, useLargeTile ? "sgemm_large" : "sgemm_small");
174
+ const pipeline = await getPipeline(
175
+ device,
176
+ useLargeTile ? "sgemm_large" : "sgemm_small",
177
+ );
140
178
 
141
179
  const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemm-A", false);
142
180
  const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "sgemm-B", false);
@@ -147,13 +185,14 @@ export async function sgemm(
147
185
  const bNot = transB === "no-transpose";
148
186
  const useVecA = aNot && vec4Usable(ABuffer, lda, m, k); // transposed-A vec stores bank-conflict smem — measured slower than scalar
149
187
  const useVecB = vec4Usable(BBuffer, ldb, bNot ? k : n, bNot ? n : k);
150
- const paramsBuffer = createParamsBuffer(device,
188
+ const paramsBuffer = createParamsBuffer(
189
+ device,
151
190
  [
152
- { value: m, type: "u32" },
153
- { value: n, type: "u32" },
154
- { value: k, type: "u32" },
191
+ { value: m, type: "u32" },
192
+ { value: n, type: "u32" },
193
+ { value: k, type: "u32" },
155
194
  { value: alpha, type: "f32" },
156
- { value: beta, type: "f32" },
195
+ { value: beta, type: "f32" },
157
196
  { value: lda, type: "u32" },
158
197
  { value: ldb, type: "u32" },
159
198
  { value: ldc, type: "u32" },
@@ -177,15 +216,32 @@ export async function sgemm(
177
216
 
178
217
  const wgCount = useLargeTile
179
218
  ? {
180
- x: requireWorkgroupCount(device, largeWgX, "sgemm", "x"),
181
- y: requireWorkgroupCount(device, largeWgY, "sgemm", "y"),
182
- }
219
+ x: requireWorkgroupCount(device, largeWgX, "sgemm", "x"),
220
+ y: requireWorkgroupCount(device, largeWgY, "sgemm", "y"),
221
+ }
183
222
  : {
184
- x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "sgemm", "x"),
185
- y: requireWorkgroupCount(device, Math.ceil(m / BM_SMALL), "sgemm", "y"),
186
- };
187
- const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
188
- const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
223
+ x: requireWorkgroupCount(
224
+ device,
225
+ Math.ceil(n / BN_SMALL),
226
+ "sgemm",
227
+ "x",
228
+ ),
229
+ y: requireWorkgroupCount(
230
+ device,
231
+ Math.ceil(m / BM_SMALL),
232
+ "sgemm",
233
+ "y",
234
+ ),
235
+ };
236
+ const { commandEncoder, ts } = runComputePass(
237
+ device,
238
+ pipeline,
239
+ bindGroup,
240
+ wgCount,
241
+ );
242
+ const readBuffer = CIsGpu
243
+ ? null
244
+ : stageReadback(device, commandEncoder, CBuffer);
189
245
 
190
246
  submit(device, commandEncoder);
191
247
 
@@ -1,7 +1,8 @@
1
1
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
2
 
3
3
  /**
4
- * Performs the matrix-matrix operation C := uplo(alpha * op(A) * op(B) + beta * C) —
4
+ * Performs the matrix-matrix operation $$C \leftarrow \mathrm{uplo}(\alpha \mathrm{op}(A) \mathrm{op}(B) + \beta C)$$
5
+ *
5
6
  * `sgemm`'s operation, but only the triangle of C named by `uplo` is read or
6
7
  * written (`'lower'`: `col <= row`, `'upper'`: `col >= row`; C need not be
7
8
  * square — the test applies over the full m×n grid).
@@ -57,7 +58,7 @@ export declare function sgemmtr(
57
58
  ): Promise<{ C: Float32Array; gpuTimeMs?: number }>;
58
59
 
59
60
  /**
60
- * Performs the matrix-matrix operation C := uplo(alpha * op(A) * op(B) + beta * C)
61
+ * Performs the matrix-matrix operation $$C \leftarrow \mathrm{uplo}(\alpha \mathrm{op}(A) \mathrm{op}(B) + \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
@@ -11,19 +11,38 @@ import { 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
  export async function sgemmtr(
19
- device, uplo, transA, transB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
24
+ device,
25
+ uplo,
26
+ transA,
27
+ transB,
28
+ m,
29
+ n,
30
+ k,
31
+ alpha,
32
+ A,
33
+ lda,
34
+ B,
35
+ ldb,
36
+ beta,
37
+ C,
38
+ ldc,
39
+ layout = "row-major",
20
40
  ) {
21
41
  let AIsGpu = A instanceof GpuMatrix;
22
42
  let BIsGpu = B instanceof GpuMatrix;
23
43
  const CIsGpu = C instanceof GpuMatrix;
24
44
 
25
- if (!(device instanceof GPUDevice))
26
- throw new Error("device must be a GPUDevice.");
45
+ requireGpuDevice(device);
27
46
  requireSameDevice(device, "sgemmtr", { A, B, C });
28
47
  if (uplo !== "lower" && uplo !== "upper")
29
48
  throw new Error("uplo must be 'lower' or 'upper'.");
@@ -33,12 +52,10 @@ export async function sgemmtr(
33
52
  throw new Error("transB must be 'no-transpose' or 'transpose'.");
34
53
  if (layout !== "row-major" && layout !== "column-major")
35
54
  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.");
55
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
38
56
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
39
57
  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.");
58
+ if (typeof beta !== "number") throw new Error("beta must be a number.");
42
59
  if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
43
60
  if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
44
61
  if (
@@ -60,7 +77,10 @@ export async function sgemmtr(
60
77
  throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");
61
78
  if (CIsGpu && (!AIsGpu || !BIsGpu))
62
79
  throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");
63
- if (m < 0 || n < 0 || k < 0) throw new Error("m, n, and k must be non-negative.");
80
+ if (m < 0 || n < 0 || k < 0)
81
+ throw new Error("m, n, and k must be non-negative.");
82
+ if (lda <= 0 || ldb <= 0 || ldc <= 0)
83
+ throw new Error("lda, ldb, and ldc must be positive.");
64
84
  if (m === 0 || n === 0) return CIsGpu ? {} : { C };
65
85
 
66
86
  // GpuMatrix's own .layout wins over the shared `layout` argument.
@@ -75,14 +95,19 @@ export async function sgemmtr(
75
95
  const aOuter = transA === "no-transpose" ? aRows : aCols; // # of stored chunks
76
96
  const aInner = transA === "no-transpose" ? aCols : aRows; // required chunk length
77
97
  if (lda < aInner)
78
- throw new Error(`lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`);
98
+ throw new Error(
99
+ `lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`,
100
+ );
79
101
  if (AIsGpu) {
80
- if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
102
+ if (lda !== A.lda)
103
+ throw new Error("lda must match A.lda when A is a GpuMatrix.");
81
104
  const [aLogRows, aLogCols] = transA === "no-transpose" ? [m, k] : [k, m];
82
105
  if (A.rows < aLogRows || A.cols < aLogCols)
83
106
  throw new Error("A is too small for the given m, k, and transA.");
84
107
  } else if (A.length < (aOuter - 1) * lda + aInner) {
85
- throw new Error("A does not have enough elements for the given dimensions and lda.");
108
+ throw new Error(
109
+ "A does not have enough elements for the given dimensions and lda.",
110
+ );
86
111
  }
87
112
 
88
113
  // B: same reasoning as A, with op(B) = k x n.
@@ -91,26 +116,37 @@ export async function sgemmtr(
91
116
  const bOuter = transB === "no-transpose" ? bRows : bCols;
92
117
  const bInner = transB === "no-transpose" ? bCols : bRows;
93
118
  if (ldb < bInner)
94
- throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
119
+ throw new Error(
120
+ `ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
121
+ );
95
122
  if (BIsGpu) {
96
- if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
123
+ if (ldb !== B.lda)
124
+ throw new Error("ldb must match B.lda when B is a GpuMatrix.");
97
125
  const [bLogRows, bLogCols] = transB === "no-transpose" ? [k, n] : [n, k];
98
126
  if (B.rows < bLogRows || B.cols < bLogCols)
99
127
  throw new Error("B is too small for the given n, k, and transB.");
100
128
  } else if (B.length < (bOuter - 1) * ldb + bInner) {
101
- throw new Error("B does not have enough elements for the given dimensions and ldb.");
129
+ throw new Error(
130
+ "B does not have enough elements for the given dimensions and ldb.",
131
+ );
102
132
  }
103
133
 
104
134
  // C: always m x n (no trans flag) — layout only affects lda/storage order.
105
135
  const cOuter = effLayoutC === "column-major" ? n : m;
106
136
  const cInner = effLayoutC === "column-major" ? m : n;
107
137
  if (ldc < cInner)
108
- throw new Error(`ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`);
138
+ throw new Error(
139
+ `ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`,
140
+ );
109
141
  if (CIsGpu) {
110
- if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
111
- if (C.rows < m || C.cols < n) throw new Error("C is too small for the given m and n.");
142
+ if (ldc !== C.lda)
143
+ throw new Error("ldc must match C.lda when C is a GpuMatrix.");
144
+ if (C.rows < m || C.cols < n)
145
+ throw new Error("C is too small for the given m and n.");
112
146
  } else if (C.length < (cOuter - 1) * ldc + cInner) {
113
- throw new Error("C does not have enough elements for the given dimensions and ldc.");
147
+ throw new Error(
148
+ "C does not have enough elements for the given dimensions and ldc.",
149
+ );
114
150
  }
115
151
 
116
152
  // Column-major A/B reinterpreted row-major is A^T/B^T — flip the trans flag.
@@ -141,18 +177,22 @@ export async function sgemmtr(
141
177
  const largeWgY = Math.ceil(m / BM_LARGE);
142
178
  const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
143
179
 
144
- const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
180
+ const pipeline = await getPipeline(
181
+ device,
182
+ useLargeTile ? "sgemmtr_large" : "sgemmtr_small",
183
+ );
145
184
 
146
185
  const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemmtr-A", false);
147
186
  const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "sgemmtr-B", false);
148
187
  const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "sgemmtr-C", true);
149
- const paramsBuffer = createParamsBuffer(device,
188
+ const paramsBuffer = createParamsBuffer(
189
+ device,
150
190
  [
151
- { value: m, type: "u32" },
152
- { value: n, type: "u32" },
153
- { value: k, type: "u32" },
191
+ { value: m, type: "u32" },
192
+ { value: n, type: "u32" },
193
+ { value: k, type: "u32" },
154
194
  { value: alpha, type: "f32" },
155
- { value: beta, type: "f32" },
195
+ { value: beta, type: "f32" },
156
196
  { value: lda, type: "u32" },
157
197
  { value: ldb, type: "u32" },
158
198
  { value: ldc, type: "u32" },
@@ -173,15 +213,32 @@ export async function sgemmtr(
173
213
 
174
214
  const wgCount = useLargeTile
175
215
  ? {
176
- x: requireWorkgroupCount(device, largeWgX, "sgemmtr", "x"),
177
- y: requireWorkgroupCount(device, largeWgY, "sgemmtr", "y"),
178
- }
216
+ x: requireWorkgroupCount(device, largeWgX, "sgemmtr", "x"),
217
+ y: requireWorkgroupCount(device, largeWgY, "sgemmtr", "y"),
218
+ }
179
219
  : {
180
- x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "sgemmtr", "x"),
181
- y: requireWorkgroupCount(device, Math.ceil(m / BM_SMALL), "sgemmtr", "y"),
182
- };
183
- const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
184
- const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
220
+ x: requireWorkgroupCount(
221
+ device,
222
+ Math.ceil(n / BN_SMALL),
223
+ "sgemmtr",
224
+ "x",
225
+ ),
226
+ y: requireWorkgroupCount(
227
+ device,
228
+ Math.ceil(m / BM_SMALL),
229
+ "sgemmtr",
230
+ "y",
231
+ ),
232
+ };
233
+ const { commandEncoder, ts } = runComputePass(
234
+ device,
235
+ pipeline,
236
+ bindGroup,
237
+ wgCount,
238
+ );
239
+ const readBuffer = CIsGpu
240
+ ? null
241
+ : stageReadback(device, commandEncoder, CBuffer);
185
242
 
186
243
  submit(device, commandEncoder);
187
244
 
@@ -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 matrix-vector operation y = alpha * op(A) * x + beta * y
5
+ * Performs the matrix-vector operation $$y \leftarrow \alpha \mathrm{op}(A) x + \beta y$$
6
6
  *
7
7
  * - `trans='no-transpose'`: op(A) = A, x is length n, y is length m
8
8
  * - `trans='transpose'`: op(A) = A^T, x is length m, y is length n
@@ -51,7 +51,7 @@ export declare function sgemv(
51
51
  ): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
52
52
 
53
53
  /**
54
- * Performs the matrix-vector operation y = alpha * op(A) * x + beta * y
54
+ * Performs the matrix-vector operation $$y \leftarrow \alpha \mathrm{op}(A) x + \beta y$$
55
55
  *
56
56
  * x and y are kept resident on the GPU. A must be a GpuMatrix; its own
57
57
  * `layout` (set at `GpuMatrix.from` time) determines the operation — there is
@@ -12,26 +12,37 @@ import { getPipeline } from "../util/pipeline.mjs";
12
12
  import { requireWorkgroups } from "../util/workgroup.mjs";
13
13
  import { GpuVector } from "../classes/GpuVector.mjs";
14
14
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
15
- import { requireSameDevice } from "../util/device.mjs";
15
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
16
16
 
17
- export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y, incy, layout = "row-major") {
17
+ export async function sgemv(
18
+ device,
19
+ trans,
20
+ m,
21
+ n,
22
+ alpha,
23
+ A,
24
+ lda,
25
+ x,
26
+ incx,
27
+ beta,
28
+ y,
29
+ incy,
30
+ layout = "row-major",
31
+ ) {
18
32
  const AIsGpu = A instanceof GpuMatrix;
19
33
  const xIsGpu = x instanceof GpuVector;
20
34
  const yIsGpu = y instanceof GpuVector;
21
35
 
22
- if (!(device instanceof GPUDevice))
23
- throw new Error("device must be a GPUDevice.");
36
+ requireGpuDevice(device);
24
37
  requireSameDevice(device, "sgemv", { A, x, y });
25
38
  if (trans !== "no-transpose" && trans !== "transpose")
26
39
  throw new Error("trans must be 'no-transpose' or 'transpose'.");
27
40
  if (layout !== "row-major" && layout !== "column-major")
28
41
  throw new Error("layout must be 'row-major' or 'column-major'.");
29
- if (typeof alpha !== "number")
30
- throw new Error("alpha must be a number.");
42
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
31
43
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
32
44
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
33
- if (typeof beta !== "number")
34
- throw new Error("beta must be a number.");
45
+ if (typeof beta !== "number") throw new Error("beta must be a number.");
35
46
  if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
36
47
  if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
37
48
  if (
@@ -55,15 +66,13 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
55
66
  "x and y must be the same type (both Float32Array or both GpuVector).",
56
67
  );
57
68
  if (xIsGpu && !AIsGpu)
58
- throw new Error(
59
- "A must be a GpuMatrix when x and y are GpuVectors.",
60
- );
69
+ throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");
61
70
  if (AIsGpu && !xIsGpu)
71
+ throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");
72
+ if (xIsGpu && x._buf === y._buf)
62
73
  throw new Error(
63
- "x and y must be GpuVectors when A is a GpuMatrix.",
74
+ "x and y must not reference the same GPU buffer when both are GpuVectors.",
64
75
  );
65
- if (xIsGpu && x._buf === y._buf)
66
- throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");
67
76
  if (AIsGpu && yIsGpu && A._buf === y._buf)
68
77
  throw new Error("A and y must not reference the same GPU buffer.");
69
78
  if (AIsGpu && lda !== A.lda)
@@ -101,7 +110,7 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
101
110
  );
102
111
 
103
112
  const shaderName = isNoTrans ? "sgemv_n" : "sgemv_t";
104
- const pipeline = await getPipeline(device, shaderName);
113
+ const pipeline = await getPipeline(device, shaderName);
105
114
 
106
115
  let ABuffer = null;
107
116
  let xBuffer = null;
@@ -112,15 +121,16 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
112
121
  ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemv-A", false);
113
122
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sgemv-x", false);
114
123
  yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sgemv-y", true);
115
- paramsBuffer = createParamsBuffer(device,
124
+ paramsBuffer = createParamsBuffer(
125
+ device,
116
126
  [
117
- { value: m, type: "u32" },
118
- { value: n, type: "u32" },
127
+ { value: m, type: "u32" },
128
+ { value: n, type: "u32" },
119
129
  { value: alpha, type: "f32" },
120
- { value: beta, type: "f32" },
121
- { value: incx, type: "u32" },
122
- { value: incy, type: "u32" },
123
- { value: lda, type: "u32" },
130
+ { value: beta, type: "f32" },
131
+ { value: incx, type: "u32" },
132
+ { value: incy, type: "u32" },
133
+ { value: lda, type: "u32" },
124
134
  ],
125
135
  "sgemv-params",
126
136
  );
@@ -139,8 +149,15 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
139
149
  const wgCount = isNoTrans
140
150
  ? Math.min(m, device.limits.maxComputeWorkgroupsPerDimension)
141
151
  : requireWorkgroups(device, "sgemv", yLen);
142
- const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
143
- const readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
152
+ const { commandEncoder, ts } = runComputePass(
153
+ device,
154
+ pipeline,
155
+ bindGroup,
156
+ wgCount,
157
+ );
158
+ const readBuffer = yIsGpu
159
+ ? null
160
+ : stageReadback(device, commandEncoder, yBuffer);
144
161
 
145
162
  submit(device, commandEncoder);
146
163
 
@@ -159,6 +176,5 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
159
176
  if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
160
177
  if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
161
178
  if (paramsBuffer) destroyBuffers(paramsBuffer);
162
-
163
179
  }
164
180
  }
@@ -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 rank-1 update A = alpha * x * y^T + A
5
+ * Performs the rank-1 update $$A \leftarrow \alpha x y^{T} + A$$
6
6
  *
7
7
  * A is an m×n matrix stored in row-major order, updated in place. `lda` is
8
8
  * the leading dimension (number of floats between the start of consecutive
@@ -43,7 +43,7 @@ export declare function sger(
43
43
  ): Promise<{ A: Float32Array; gpuTimeMs?: number }>;
44
44
 
45
45
  /**
46
- * Performs the rank-1 update A = alpha * x * y^T + A
46
+ * Performs the rank-1 update $$A \leftarrow \alpha x y^{T} + A$$
47
47
  *
48
48
  * x, y, and A are all kept resident on the GPU. `A`'s own `layout` (set at
49
49
  * `GpuMatrix.from` time) determines the operation — there is no separate