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
@@ -13,24 +13,40 @@ import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
13
13
  import { getPipeline } from "../util/pipeline.mjs";
14
14
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
15
15
  import { requireWorkgroupCount } from "../util/workgroup.mjs";
16
- import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
16
+ import {
17
+ BM_SMALL,
18
+ BN_SMALL,
19
+ BM_LARGE,
20
+ BN_LARGE,
21
+ LARGE_TILE_WORKGROUP_THRESHOLD,
22
+ } from "../util/constants.mjs";
17
23
  import { TILE_WG_2D } from "../util/constants.mjs";
18
- import { requireSameDevice } from "../util/device.mjs";
19
-
24
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
20
25
 
21
26
  // strmm: B := alpha*op(A)*B (side='left') or alpha*B*op(A) (side='right'), A
22
27
  // triangular. Triangularize then sgemm, one command encoder. B is both
23
28
  // input and output, so gemm writes to a fresh buffer (no aliasing race),
24
29
  // copied back into B (GpuMatrix) or read back directly (Float32Array).
25
30
  export async function strmm(
26
- device, side, uplo, transA, diag, m, n, alpha, A, lda, B, ldb, layout = "row-major",
31
+ device,
32
+ side,
33
+ uplo,
34
+ transA,
35
+ diag,
36
+ m,
37
+ n,
38
+ alpha,
39
+ A,
40
+ lda,
41
+ B,
42
+ ldb,
43
+ layout = "row-major",
27
44
  ) {
28
45
  const AIsGpu = A instanceof GpuMatrix;
29
46
  const BIsGpu = B instanceof GpuMatrix;
30
47
  const isUnit = diag === "unit";
31
48
 
32
- if (!(device instanceof GPUDevice))
33
- throw new Error("device must be a GPUDevice.");
49
+ requireGpuDevice(device);
34
50
  requireSameDevice(device, "strmm", { A, B });
35
51
  if (side !== "left" && side !== "right")
36
52
  throw new Error("side must be 'left' or 'right'.");
@@ -42,11 +58,15 @@ export async function strmm(
42
58
  throw new Error("diag must be 'unit' or 'non-unit'.");
43
59
  if (layout !== "row-major" && layout !== "column-major")
44
60
  throw new Error("layout must be 'row-major' or 'column-major'.");
45
- if (typeof alpha !== "number")
46
- throw new Error("alpha must be a number.");
61
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
47
62
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
48
63
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
49
- if (!Number.isInteger(m) || !Number.isInteger(n) || !Number.isInteger(lda) || !Number.isInteger(ldb))
64
+ if (
65
+ !Number.isInteger(m) ||
66
+ !Number.isInteger(n) ||
67
+ !Number.isInteger(lda) ||
68
+ !Number.isInteger(ldb)
69
+ )
50
70
  throw new Error("m, n, lda, and ldb must be integers.");
51
71
  if (!AIsGpu && !(A instanceof Float32Array))
52
72
  throw new Error("A must be a Float32Array or GpuMatrix.");
@@ -62,37 +82,59 @@ export async function strmm(
62
82
 
63
83
  // A: triangular, order = m (side='left') or n (side='right').
64
84
  const aOrder = side === "left" ? m : n;
65
- if (lda < aOrder) throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
85
+ if (lda < aOrder)
86
+ throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
66
87
  if (AIsGpu) {
67
- if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
68
- if (A.rows < aOrder || A.cols < aOrder) throw new Error("A is too small for the given m/n and side.");
88
+ if (lda !== A.lda)
89
+ throw new Error("lda must match A.lda when A is a GpuMatrix.");
90
+ if (A.rows < aOrder || A.cols < aOrder)
91
+ throw new Error("A is too small for the given m/n and side.");
69
92
  } else if (A.length < (aOrder - 1) * lda + aOrder) {
70
- throw new Error("A does not have enough elements for the given dimensions and lda.");
93
+ throw new Error(
94
+ "A does not have enough elements for the given dimensions and lda.",
95
+ );
71
96
  }
72
97
 
73
98
  // B: always m x n, overwritten in place with the same ldb.
74
99
  const bOuter = effLayoutB === "column-major" ? n : m;
75
100
  const bInner = effLayoutB === "column-major" ? m : n;
76
101
  if (ldb < bInner)
77
- throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
102
+ throw new Error(
103
+ `ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
104
+ );
78
105
  if (BIsGpu) {
79
- if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
80
- if (B.rows < m || B.cols < n) throw new Error("B is too small for the given m and n.");
106
+ if (ldb !== B.lda)
107
+ throw new Error("ldb must match B.lda when B is a GpuMatrix.");
108
+ if (B.rows < m || B.cols < n)
109
+ throw new Error("B is too small for the given m and n.");
81
110
  } else if (B.length < (bOuter - 1) * ldb + bInner) {
82
- throw new Error("B does not have enough elements for the given dimensions and ldb.");
111
+ throw new Error(
112
+ "B does not have enough elements for the given dimensions and ldb.",
113
+ );
83
114
  }
84
115
 
85
116
  // A isn't symmetric: column-major = genuine transpose, so flip transA;
86
117
  // transposing also swaps which triangle looks stored, so flip uplo too.
87
- const uploEffA = effLayoutA === "column-major" ? (uplo === "lower" ? "upper" : "lower") : uplo;
88
- const transEffA = effLayoutA === "column-major" ? (transA === "no-transpose" ? "transpose" : "no-transpose") : transA;
118
+ const uploEffA =
119
+ effLayoutA === "column-major"
120
+ ? uplo === "lower"
121
+ ? "upper"
122
+ : "lower"
123
+ : uplo;
124
+ const transEffA =
125
+ effLayoutA === "column-major"
126
+ ? transA === "no-transpose"
127
+ ? "transpose"
128
+ : "no-transpose"
129
+ : transA;
89
130
 
90
131
  const transB = effLayoutB === "column-major" ? "transpose" : "no-transpose";
91
132
  const transDense = "no-transpose"; // Adense already embodies op(A)
92
133
 
93
134
  // X*Y (X=Adense,Y=B for side='left', swapped for 'right'). Column-major
94
135
  // output: compute (B_out)^T instead — sgemm's own trick, same as ssymm's.
95
- let mg = m, ng = n;
136
+ let mg = m,
137
+ ng = n;
96
138
  const kg = aOrder;
97
139
  let transX = side === "left" ? transDense : transB;
98
140
  let transY = side === "left" ? transB : transDense;
@@ -108,39 +150,63 @@ export async function strmm(
108
150
  const largeWgX = Math.ceil(ng / BN_LARGE);
109
151
  const largeWgY = Math.ceil(mg / BM_LARGE);
110
152
  const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
111
- const gemmPipeline = await getPipeline(device, useLargeTile ? "sgemm_large" : "sgemm_small");
153
+ const gemmPipeline = await getPipeline(
154
+ device,
155
+ useLargeTile ? "sgemm_large" : "sgemm_small",
156
+ );
112
157
  const triPipeline = await getPipeline(device, "triangularize");
113
158
  const gemmWgCount = useLargeTile
114
159
  ? {
115
- x: requireWorkgroupCount(device, largeWgX, "strmm", "x"),
116
- y: requireWorkgroupCount(device, largeWgY, "strmm", "y"),
117
- }
160
+ x: requireWorkgroupCount(device, largeWgX, "strmm", "x"),
161
+ y: requireWorkgroupCount(device, largeWgY, "strmm", "y"),
162
+ }
118
163
  : {
119
- x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "strmm", "x"),
120
- y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "strmm", "y"),
121
- };
164
+ x: requireWorkgroupCount(
165
+ device,
166
+ Math.ceil(ng / BN_SMALL),
167
+ "strmm",
168
+ "x",
169
+ ),
170
+ y: requireWorkgroupCount(
171
+ device,
172
+ Math.ceil(mg / BM_SMALL),
173
+ "strmm",
174
+ "y",
175
+ ),
176
+ };
122
177
 
123
178
  // Null-init here and allocate inside the try below, so a throw partway
124
179
  // through the sequence still reaches finally with every handle visible
125
180
  // (strsv.mjs is the reference for this pattern).
126
- let ABuffer = null, BBuffer = null;
127
- let AdenseBuffer = null, outBuffer = null;
128
- let triParams = null, gemmParams = null;
181
+ let ABuffer = null,
182
+ BBuffer = null;
183
+ let AdenseBuffer = null,
184
+ outBuffer = null;
185
+ let triParams = null,
186
+ gemmParams = null;
129
187
  let outBufferAdopted = false; // true once B._buf is repointed at outBuffer
130
188
 
131
189
  try {
132
190
  ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmm-A", false);
133
191
  // readback=true (COPY_SRC): BBuffer is also the source that seeds outBuffer.
134
192
  BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strmm-B", true);
135
- AdenseBuffer = createStorageBuffer(device, aOrder * ldDense * 4, "strmm-Adense");
193
+ AdenseBuffer = createStorageBuffer(
194
+ device,
195
+ aOrder * ldDense * 4,
196
+ "strmm-Adense",
197
+ );
136
198
  // COPY_DST: seeded from B's own content before gemm runs, so stride-padding
137
199
  // gaps (never written by gemm's tight m x n loop) keep B's original bytes
138
200
  // instead of reading back as zero. COPY_SRC: read back / adopted by B after.
139
- outBuffer = createStorageBuffer(device,
140
- bOuter * ldb * 4, "strmm-out", GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST,
201
+ outBuffer = createStorageBuffer(
202
+ device,
203
+ bOuter * ldb * 4,
204
+ "strmm-out",
205
+ GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST,
141
206
  );
142
207
 
143
- triParams = createParamsBuffer(device,
208
+ triParams = createParamsBuffer(
209
+ device,
144
210
  [
145
211
  { value: aOrder, type: "u32" },
146
212
  { value: lda, type: "u32" },
@@ -151,7 +217,11 @@ export async function strmm(
151
217
  ],
152
218
  "strmm-tri-params",
153
219
  );
154
- const triBindGroup = createBindGroup(device, triPipeline.getBindGroupLayout(0), [ABuffer, AdenseBuffer, triParams]);
220
+ const triBindGroup = createBindGroup(
221
+ device,
222
+ triPipeline.getBindGroupLayout(0),
223
+ [ABuffer, AdenseBuffer, triParams],
224
+ );
155
225
 
156
226
  // X/Y buffers and their own ld, matching swapXY above.
157
227
  const XBuffer = swapXY ? BBuffer : AdenseBuffer;
@@ -159,13 +229,14 @@ export async function strmm(
159
229
  const YBuffer = swapXY ? AdenseBuffer : BBuffer;
160
230
  const ldY = swapXY ? ldDense : ldb;
161
231
 
162
- gemmParams = createParamsBuffer(device,
232
+ gemmParams = createParamsBuffer(
233
+ device,
163
234
  [
164
- { value: mg, type: "u32" },
165
- { value: ng, type: "u32" },
166
- { value: kg, type: "u32" },
235
+ { value: mg, type: "u32" },
236
+ { value: ng, type: "u32" },
237
+ { value: kg, type: "u32" },
167
238
  { value: alpha, type: "f32" },
168
- { value: 0.0, type: "f32" }, // beta — strmm has no C accumulation term
239
+ { value: 0.0, type: "f32" }, // beta — strmm has no C accumulation term
169
240
  { value: ldX, type: "u32" },
170
241
  { value: ldY, type: "u32" },
171
242
  { value: ldb, type: "u32" },
@@ -174,28 +245,56 @@ export async function strmm(
174
245
  ],
175
246
  "strmm-gemm-params",
176
247
  );
177
- const gemmBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
178
- XBuffer,
179
- vec4ViewBinding(device, XBuffer),
180
- YBuffer,
181
- vec4ViewBinding(device, YBuffer),
182
- outBuffer,
183
- gemmParams,
184
- ]);
248
+ const gemmBindGroup = createBindGroup(
249
+ device,
250
+ gemmPipeline.getBindGroupLayout(0),
251
+ [
252
+ XBuffer,
253
+ vec4ViewBinding(device, XBuffer),
254
+ YBuffer,
255
+ vec4ViewBinding(device, YBuffer),
256
+ outBuffer,
257
+ gemmParams,
258
+ ],
259
+ );
185
260
 
186
261
  const { commandEncoder, querySet } = beginTimedEncoder(device);
187
262
  // Seed outBuffer with B's own bytes first, so gemm's tight m x n write
188
263
  // leaves stride-padding gaps holding B's original content, not zero.
189
264
  // BBuffer may be larger than outBuffer (e.g. a validation-test baseline
190
265
  // over-provisioned for a bigger ldb it might later be substituted with).
191
- commandEncoder.copyBufferToBuffer(BBuffer, 0, outBuffer, 0, Math.min(BBuffer.size, outBuffer.size));
192
- const triDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
193
- const gemmDesc = querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
194
- encodePass(commandEncoder, triPipeline, triBindGroup, { x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) }, triDesc);
195
- encodePass(commandEncoder, gemmPipeline, gemmBindGroup, gemmWgCount, gemmDesc);
266
+ commandEncoder.copyBufferToBuffer(
267
+ BBuffer,
268
+ 0,
269
+ outBuffer,
270
+ 0,
271
+ Math.min(BBuffer.size, outBuffer.size),
272
+ );
273
+ const triDesc = querySet
274
+ ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
275
+ : undefined;
276
+ const gemmDesc = querySet
277
+ ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
278
+ : undefined;
279
+ encodePass(
280
+ commandEncoder,
281
+ triPipeline,
282
+ triBindGroup,
283
+ { x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) },
284
+ triDesc,
285
+ );
286
+ encodePass(
287
+ commandEncoder,
288
+ gemmPipeline,
289
+ gemmBindGroup,
290
+ gemmWgCount,
291
+ gemmDesc,
292
+ );
196
293
 
197
294
  const ts = resolveTimestamp(device, commandEncoder, querySet);
198
- const readBuffer = BIsGpu ? null : stageReadback(device, commandEncoder, outBuffer);
295
+ const readBuffer = BIsGpu
296
+ ? null
297
+ : stageReadback(device, commandEncoder, outBuffer);
199
298
 
200
299
  submit(device, commandEncoder);
201
300
 
@@ -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 triangular matrix-vector operation y = op(A) * x
5
+ * Performs the triangular matrix-vector operation $$y \leftarrow \mathrm{op}(A) x$$
6
6
  *
7
7
  * A is an n×n triangular matrix stored in row-major order. Only the triangle
8
8
  * specified by `uplo` is referenced; the other triangle is not accessed.
@@ -45,7 +45,7 @@ export declare function strmv(
45
45
  ): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
46
46
 
47
47
  /**
48
- * Performs the triangular matrix-vector operation y = op(A) * x
48
+ * Performs the triangular matrix-vector operation $$y \leftarrow \mathrm{op}(A) x$$
49
49
  *
50
50
  * x and y are kept resident on the GPU. A must be a GpuMatrix; its own
51
51
  * `layout` (set at `GpuMatrix.from` time) determines the operation — there is
@@ -11,16 +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 { requireSameDevice } from "../util/device.mjs";
14
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
15
15
 
16
- export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, incy, layout = "row-major") {
16
+ export async function strmv(
17
+ device,
18
+ uplo,
19
+ trans,
20
+ diag,
21
+ n,
22
+ A,
23
+ lda,
24
+ x,
25
+ incx,
26
+ y,
27
+ incy,
28
+ layout = "row-major",
29
+ ) {
17
30
  const xIsGpu = x instanceof GpuVector;
18
31
  const yIsGpu = y instanceof GpuVector;
19
32
  const AIsGpu = A instanceof GpuMatrix;
20
33
  const isUnit = diag === "unit";
21
34
 
22
- if (!(device instanceof GPUDevice))
23
- throw new Error("device must be a GPUDevice.");
35
+ requireGpuDevice(device);
24
36
  requireSameDevice(device, "strmv", { A, x, y });
25
37
  if (uplo !== "lower" && uplo !== "upper")
26
38
  throw new Error("uplo must be 'lower' or 'upper'.");
@@ -68,9 +80,7 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
68
80
  if (n === 0) return yIsGpu ? {} : { y };
69
81
 
70
82
  if (!AIsGpu && A.length < (n - 1) * lda + n)
71
- throw new Error(
72
- "A does not have enough elements for the given n and lda.",
73
- );
83
+ throw new Error("A does not have enough elements for the given n and lda.");
74
84
  if (x.length < (n - 1) * incx + 1)
75
85
  throw new Error(
76
86
  "x does not have enough elements for the given n and incx.",
@@ -84,7 +94,9 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
84
94
  const effLayout = AIsGpu ? A.layout : layout;
85
95
  const isColMajor = effLayout === "column-major";
86
96
  const isLower = isColMajor ? uplo === "upper" : uplo === "lower";
87
- const isNoTrans = isColMajor ? trans === "transpose" : trans === "no-transpose";
97
+ const isNoTrans = isColMajor
98
+ ? trans === "transpose"
99
+ : trans === "no-transpose";
88
100
 
89
101
  const pipeline = await getPipeline(device, "strmv");
90
102
 
@@ -97,15 +109,16 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
97
109
  ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmv-A", false);
98
110
  xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "strmv-x", false);
99
111
  yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "strmv-y", true);
100
- paramsBuffer = createParamsBuffer(device,
112
+ paramsBuffer = createParamsBuffer(
113
+ device,
101
114
  [
102
- { value: n, type: "u32" },
103
- { value: incx, type: "u32" },
104
- { value: incy, type: "u32" },
105
- { value: lda, type: "u32" },
115
+ { value: n, type: "u32" },
116
+ { value: incx, type: "u32" },
117
+ { value: incy, type: "u32" },
118
+ { value: lda, type: "u32" },
106
119
  { value: isNoTrans ? 0 : 1, type: "u32" },
107
- { value: isLower ? 0 : 1, type: "u32" },
108
- { value: isUnit ? 1 : 0, type: "u32" },
120
+ { value: isLower ? 0 : 1, type: "u32" },
121
+ { value: isUnit ? 1 : 0, type: "u32" },
109
122
  ],
110
123
  "strmv-params",
111
124
  );
@@ -118,8 +131,15 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
118
131
  ]);
119
132
 
120
133
  const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
121
- const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
122
- const readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
134
+ const { commandEncoder, ts } = runComputePass(
135
+ device,
136
+ pipeline,
137
+ bindGroup,
138
+ wgCount,
139
+ );
140
+ const readBuffer = yIsGpu
141
+ ? null
142
+ : stageReadback(device, commandEncoder, yBuffer);
123
143
 
124
144
  submit(device, commandEncoder);
125
145
 
@@ -2,8 +2,9 @@ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
2
 
3
3
  /**
4
4
  * Solves the triangular matrix equation
5
- * op(A) * X = alpha * B (`side='left'`) or
6
- * X * op(A) = alpha * B (`side='right'`), overwriting `B` with the
5
+ * $$\mathrm{op}(A) X = \alpha B \quad (\texttt{side='left'})$$
6
+ * $$X \mathrm{op}(A) = \alpha B \quad (\texttt{side='right'})$$
7
+ * overwriting `B` with the
7
8
  * solution `X` — `A` is triangular, only its `uplo` triangle stored; `B` is
8
9
  * a general m×n matrix.
9
10
  *
@@ -57,8 +58,9 @@ export declare function strsm(
57
58
 
58
59
  /**
59
60
  * Solves the triangular matrix equation
60
- * op(A) * X = alpha * B (`side='left'`) or
61
- * X * op(A) = alpha * B (`side='right'`), overwriting `B` in place with `X`.
61
+ * $$\mathrm{op}(A) X = \alpha B \quad (\texttt{side='left'})$$
62
+ * $$X \mathrm{op}(A) = \alpha B \quad (\texttt{side='right'})$$
63
+ * overwriting `B` in place with `X`.
62
64
  *
63
65
  * A and B are both kept GPU-resident. Each matrix's own `layout` (set at
64
66
  * `GpuMatrix.from` time) determines the operation — there is no separate