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