wgblas 2.0.0 → 2.1.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 (66) hide show
  1. package/README.md +18 -18
  2. package/dist/wgblas.browser.js +1273 -1239
  3. package/index.d.mts +38 -6
  4. package/package.json +2 -1
  5. package/src/classes/GpuMatrix.mjs +17 -10
  6. package/src/classes/GpuVector.mjs +28 -10
  7. package/src/dasum/dasum.d.mts +4 -4
  8. package/src/dasum/dasum.mjs +19 -17
  9. package/src/devdocs.mjs +13 -0
  10. package/src/idamax/idamax.d.mts +20 -2
  11. package/src/idamax/idamax.mjs +18 -16
  12. package/src/init.mjs +114 -56
  13. package/src/isamax/isamax.d.mts +20 -2
  14. package/src/isamax/isamax.mjs +16 -14
  15. package/src/random/random.d.mts +1 -0
  16. package/src/sasum/sasum.d.mts +2 -2
  17. package/src/sasum/sasum.mjs +15 -13
  18. package/src/saxpy/saxpy.d.mts +2 -2
  19. package/src/saxpy/saxpy.mjs +10 -8
  20. package/src/scopy/scopy.d.mts +2 -2
  21. package/src/scopy/scopy.mjs +10 -8
  22. package/src/sdot/sdot.d.mts +2 -2
  23. package/src/sdot/sdot.mjs +16 -14
  24. package/src/sgemm/sgemm.mjs +28 -15
  25. package/src/sgemmtr/sgemmtr.mjs +16 -15
  26. package/src/sgemv/sgemv.mjs +38 -26
  27. package/src/sger/sger.mjs +10 -8
  28. package/src/shaders/index.mjs +164 -14
  29. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  30. package/src/shaders/sgemm_large.wgsl +107 -18
  31. package/src/shaders/sgemm_small.wgsl +115 -15
  32. package/src/shaders/sgemmtr_large.wgsl +4 -1
  33. package/src/shaders/sgemmtr_small.wgsl +4 -1
  34. package/src/shaders/sgemv_n.wgsl +3 -1
  35. package/src/shaders/sgemv_t.wgsl +3 -1
  36. package/src/shaders/snrm2.wgsl +72 -23
  37. package/src/shaders/ssymv.wgsl +3 -1
  38. package/src/snrm2/snrm2.d.mts +2 -2
  39. package/src/snrm2/snrm2.mjs +33 -21
  40. package/src/srot/srot.d.mts +2 -4
  41. package/src/srot/srot.mjs +11 -9
  42. package/src/srotm/srotm.d.mts +2 -4
  43. package/src/srotm/srotm.mjs +19 -10
  44. package/src/sscal/sscal.d.mts +3 -3
  45. package/src/sscal/sscal.mjs +12 -10
  46. package/src/sswap/sswap.d.mts +2 -2
  47. package/src/sswap/sswap.mjs +11 -9
  48. package/src/ssymm/ssymm.mjs +31 -22
  49. package/src/ssymv/ssymv.mjs +10 -8
  50. package/src/ssyr/ssyr.mjs +9 -7
  51. package/src/ssyr2/ssyr2.mjs +10 -8
  52. package/src/ssyr2k/ssyr2k.mjs +18 -17
  53. package/src/ssyrk/ssyrk.mjs +18 -17
  54. package/src/strmm/strmm.mjs +47 -32
  55. package/src/strmv/strmv.mjs +10 -8
  56. package/src/strsm/strsm.mjs +54 -36
  57. package/src/strsv/strsv.mjs +16 -12
  58. package/src/util/benchmark.mjs +4 -6
  59. package/src/util/bindgroup.mjs +1 -3
  60. package/src/util/buffer.mjs +113 -19
  61. package/src/util/compute.mjs +6 -9
  62. package/src/util/constants.mjs +57 -0
  63. package/src/util/device.mjs +34 -0
  64. package/src/util/pipeline.mjs +5 -6
  65. package/src/util/workgroup.mjs +55 -7
  66. package/src/shaders/browser-shaders.mjs +0 -81
package/src/ssyr/ssyr.mjs CHANGED
@@ -11,6 +11,7 @@ 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
15
 
15
16
  export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "row-major") {
16
17
  const xIsGpu = x instanceof GpuVector;
@@ -18,6 +19,7 @@ export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "ro
18
19
 
19
20
  if (!(device instanceof GPUDevice))
20
21
  throw new Error("device must be a GPUDevice.");
22
+ requireSameDevice(device, "ssyr", { A, x });
21
23
  if (uplo !== "lower" && uplo !== "upper")
22
24
  throw new Error("uplo must be 'lower' or 'upper'.");
23
25
  if (layout !== "row-major" && layout !== "column-major")
@@ -63,9 +65,9 @@ export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "ro
63
65
  let paramsBuffer = null;
64
66
 
65
67
  try {
66
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "ssyr-x", false);
67
- ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyr-A", true);
68
- paramsBuffer = createParamsBuffer(
68
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "ssyr-x", false);
69
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyr-A", true);
70
+ paramsBuffer = createParamsBuffer(device,
69
71
  [
70
72
  { value: n, type: "u32" },
71
73
  { value: alpha, type: "f32" },
@@ -76,7 +78,7 @@ export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "ro
76
78
  "ssyr-params",
77
79
  );
78
80
 
79
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
81
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
80
82
  xBuffer,
81
83
  ABuffer,
82
84
  paramsBuffer,
@@ -85,10 +87,10 @@ export async function ssyr(device, uplo, n, alpha, x, incx, A, lda, layout = "ro
85
87
  // One workgroup per row of A; clamped to device limit — the shader's
86
88
  // grid-stride loop handles remaining rows when n > dispatch count.
87
89
  const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
88
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
89
- const readBuffer = AIsGpu ? null : stageReadback(commandEncoder, ABuffer);
90
+ const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
91
+ const readBuffer = AIsGpu ? null : stageReadback(device, commandEncoder, ABuffer);
90
92
 
91
- submit(commandEncoder);
93
+ submit(device, commandEncoder);
92
94
 
93
95
  const gpuTimeMs = await extractTimestamp(ts);
94
96
 
@@ -11,6 +11,7 @@ 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
15
 
15
16
  export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, layout = "row-major") {
16
17
  const xIsGpu = x instanceof GpuVector;
@@ -19,6 +20,7 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
19
20
 
20
21
  if (!(device instanceof GPUDevice))
21
22
  throw new Error("device must be a GPUDevice.");
23
+ requireSameDevice(device, "ssyr2", { A, x, y });
22
24
  if (uplo !== "lower" && uplo !== "upper")
23
25
  throw new Error("uplo must be 'lower' or 'upper'.");
24
26
  if (layout !== "row-major" && layout !== "column-major")
@@ -83,10 +85,10 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
83
85
  let paramsBuffer = null;
84
86
 
85
87
  try {
86
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "ssyr2-x", false);
87
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "ssyr2-y", false);
88
- ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyr2-A", true);
89
- paramsBuffer = createParamsBuffer(
88
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "ssyr2-x", false);
89
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "ssyr2-y", false);
90
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyr2-A", true);
91
+ paramsBuffer = createParamsBuffer(device,
90
92
  [
91
93
  { value: n, type: "u32" },
92
94
  { value: alpha, type: "f32" },
@@ -98,7 +100,7 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
98
100
  "ssyr2-params",
99
101
  );
100
102
 
101
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
103
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
102
104
  xBuffer,
103
105
  yBuffer,
104
106
  ABuffer,
@@ -108,10 +110,10 @@ export async function ssyr2(device, uplo, n, alpha, x, incx, y, incy, A, lda, la
108
110
  // One workgroup per row of A; clamped to device limit — the shader's
109
111
  // grid-stride loop handles remaining rows when n > dispatch count.
110
112
  const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
111
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
112
- const readBuffer = AIsGpu ? null : stageReadback(commandEncoder, ABuffer);
113
+ const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
114
+ const readBuffer = AIsGpu ? null : stageReadback(device, commandEncoder, ABuffer);
113
115
 
114
- submit(commandEncoder);
116
+ submit(device, commandEncoder);
115
117
 
116
118
  const gpuTimeMs = await extractTimestamp(ts);
117
119
 
@@ -10,10 +10,10 @@ 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
+ 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";
13
16
 
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
17
17
 
18
18
  // ssyr2k: C := uplo(alpha*op(A)*op(B)^T + alpha*op(B)*op(A)^T + beta*C). No
19
19
  // dedicated shader — two sgemmtr passes on one encoder, second with beta=1.
@@ -26,6 +26,7 @@ export async function ssyr2k(
26
26
 
27
27
  if (!(device instanceof GPUDevice))
28
28
  throw new Error("device must be a GPUDevice.");
29
+ requireSameDevice(device, "ssyr2k", { A, B, C });
29
30
  if (uplo !== "lower" && uplo !== "upper")
30
31
  throw new Error("uplo must be 'lower' or 'upper'.");
31
32
  if (trans !== "no-transpose" && trans !== "transpose")
@@ -131,24 +132,24 @@ export async function ssyr2k(
131
132
  const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
132
133
  const wgCount = useLargeTile
133
134
  ? {
134
- x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension),
135
- y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension),
135
+ x: requireWorkgroupCount(device, largeWgX, "ssyr2k", "x"),
136
+ y: requireWorkgroupCount(device, largeWgY, "ssyr2k", "y"),
136
137
  }
137
138
  : {
138
- x: Math.min(Math.ceil(n / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
139
- y: Math.min(Math.ceil(n / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
139
+ x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "ssyr2k", "x"),
140
+ y: requireWorkgroupCount(device, Math.ceil(n / BM_SMALL), "ssyr2k", "y"),
140
141
  };
141
142
 
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);
143
+ const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyr2k-A", false);
144
+ const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "ssyr2k-B", false);
145
+ const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssyr2k-C", true);
145
146
  let paramsBuffer1 = null, paramsBuffer2 = null;
146
147
 
147
148
  try {
148
149
  const pass1 = passShape(effTransA, ABuffer, lda, effTransB, BBuffer, ldb);
149
150
  const pass2 = passShape(effTransB, BBuffer, ldb, effTransA, ABuffer, lda);
150
151
 
151
- const makeParams = (p, betaVal) => createParamsBuffer(
152
+ const makeParams = (p, betaVal) => createParamsBuffer(device,
152
153
  [
153
154
  { value: n, type: "u32" },
154
155
  { value: n, type: "u32" },
@@ -167,19 +168,19 @@ export async function ssyr2k(
167
168
  paramsBuffer1 = makeParams(pass1, beta);
168
169
  paramsBuffer2 = makeParams(pass2, 1.0);
169
170
 
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]);
171
+ const bindGroup1 = createBindGroup(device, pipeline.getBindGroupLayout(0), [pass1.X, pass1.Y, CBuffer, paramsBuffer1]);
172
+ const bindGroup2 = createBindGroup(device, pipeline.getBindGroupLayout(0), [pass2.X, pass2.Y, CBuffer, paramsBuffer2]);
172
173
 
173
- const { commandEncoder, querySet } = beginTimedEncoder();
174
+ const { commandEncoder, querySet } = beginTimedEncoder(device);
174
175
  const desc1 = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
175
176
  const desc2 = querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
176
177
  encodePass(commandEncoder, pipeline, bindGroup1, wgCount, desc1);
177
178
  encodePass(commandEncoder, pipeline, bindGroup2, wgCount, desc2);
178
179
 
179
- const ts = resolveTimestamp(commandEncoder, querySet);
180
- const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
180
+ const ts = resolveTimestamp(device, commandEncoder, querySet);
181
+ const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
181
182
 
182
- submit(commandEncoder);
183
+ submit(device, commandEncoder);
183
184
 
184
185
  const gpuTimeMs = await extractTimestamp(ts);
185
186
 
@@ -11,10 +11,10 @@ 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
+ import { requireWorkgroupCount } from "../util/workgroup.mjs";
15
+ import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
16
+ import { requireSameDevice } from "../util/device.mjs";
14
17
 
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
18
18
 
19
19
  // ssyrk: C := uplo(alpha*op(A)*op(A)^T + beta*C). No dedicated shader —
20
20
  // sgemmtr's kernel with A duplicated into a separate B buffer (B := A).
@@ -26,6 +26,7 @@ export async function ssyrk(
26
26
 
27
27
  if (!(device instanceof GPUDevice))
28
28
  throw new Error("device must be a GPUDevice.");
29
+ requireSameDevice(device, "ssyrk", { A, C });
29
30
  if (uplo !== "lower" && uplo !== "upper")
30
31
  throw new Error("uplo must be 'lower' or 'upper'.");
31
32
  if (trans !== "no-transpose" && trans !== "transpose")
@@ -106,14 +107,14 @@ export async function ssyrk(
106
107
 
107
108
  const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
108
109
 
109
- const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "ssyrk-A", false);
110
- const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "ssyrk-C", true);
110
+ const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "ssyrk-A", false);
111
+ const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "ssyrk-C", true);
111
112
  // B := A, but as a genuinely separate buffer — re-uploaded for Float32Array,
112
113
  // GPU-copied (see below) for GpuMatrix, since the caller owns ABuffer.
113
114
  const BBuffer = AIsGpu
114
- ? createStorageBuffer(ABuffer.size, "ssyrk-B", GPUBufferUsage.COPY_DST)
115
- : uploadBuffer(A, "ssyrk-B", false);
116
- const paramsBuffer = createParamsBuffer(
115
+ ? createStorageBuffer(device, ABuffer.size, "ssyrk-B", GPUBufferUsage.COPY_DST)
116
+ : uploadBuffer(device, A, "ssyrk-B", false);
117
+ const paramsBuffer = createParamsBuffer(device,
117
118
  [
118
119
  { value: n, type: "u32" }, // gemmtr's m
119
120
  { value: n, type: "u32" }, // gemmtr's n
@@ -131,7 +132,7 @@ export async function ssyrk(
131
132
  );
132
133
 
133
134
  try {
134
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
135
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
135
136
  ABuffer,
136
137
  BBuffer,
137
138
  CBuffer,
@@ -140,22 +141,22 @@ export async function ssyrk(
140
141
 
141
142
  const wgCount = useLargeTile
142
143
  ? {
143
- x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension),
144
- y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension),
144
+ x: requireWorkgroupCount(device, largeWgX, "ssyrk", "x"),
145
+ y: requireWorkgroupCount(device, largeWgY, "ssyrk", "y"),
145
146
  }
146
147
  : {
147
- x: Math.min(Math.ceil(n / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
148
- y: Math.min(Math.ceil(n / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
148
+ x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "ssyrk", "x"),
149
+ y: requireWorkgroupCount(device, Math.ceil(n / BM_SMALL), "ssyrk", "y"),
149
150
  };
150
151
  // Manual encoder (not runComputePass) so the A->B duplicate copy lands
151
152
  // on the same command encoder, strictly before the compute pass reads B.
152
- const { commandEncoder, querySet, passDescriptor } = beginTimedEncoder();
153
+ const { commandEncoder, querySet, passDescriptor } = beginTimedEncoder(device);
153
154
  if (AIsGpu) commandEncoder.copyBufferToBuffer(ABuffer, 0, BBuffer, 0, ABuffer.size);
154
155
  encodePass(commandEncoder, pipeline, bindGroup, wgCount, passDescriptor);
155
- const ts = resolveTimestamp(commandEncoder, querySet);
156
- const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
156
+ const ts = resolveTimestamp(device, commandEncoder, querySet);
157
+ const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
157
158
 
158
- submit(commandEncoder);
159
+ submit(device, commandEncoder);
159
160
 
160
161
  const gpuTimeMs = await extractTimestamp(ts);
161
162
 
@@ -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,11 +12,11 @@ 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";
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 { TILE_WG_2D } from "../util/constants.mjs";
18
+ import { requireSameDevice } from "../util/device.mjs";
14
19
 
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)
19
20
 
20
21
  // strmm: B := alpha*op(A)*B (side='left') or alpha*B*op(A) (side='right'), A
21
22
  // triangular. Triangularize then sgemm, one command encoder. B is both
@@ -30,6 +31,7 @@ export async function strmm(
30
31
 
31
32
  if (!(device instanceof GPUDevice))
32
33
  throw new Error("device must be a GPUDevice.");
34
+ requireSameDevice(device, "strmm", { A, B });
33
35
  if (side !== "left" && side !== "right")
34
36
  throw new Error("side must be 'left' or 'right'.");
35
37
  if (uplo !== "lower" && uplo !== "upper")
@@ -110,29 +112,35 @@ export async function strmm(
110
112
  const triPipeline = await getPipeline(device, "triangularize");
111
113
  const gemmWgCount = useLargeTile
112
114
  ? {
113
- x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension),
114
- y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension),
115
+ x: requireWorkgroupCount(device, largeWgX, "strmm", "x"),
116
+ y: requireWorkgroupCount(device, largeWgY, "strmm", "y"),
115
117
  }
116
118
  : {
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
+ x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "strmm", "x"),
120
+ y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "strmm", "y"),
119
121
  };
120
122
 
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
- );
123
+ // Null-init here and allocate inside the try below, so a throw partway
124
+ // through the sequence still reaches finally with every handle visible
125
+ // (strsv.mjs is the reference for this pattern).
126
+ let ABuffer = null, BBuffer = null;
127
+ let AdenseBuffer = null, outBuffer = null;
131
128
  let triParams = null, gemmParams = null;
132
129
  let outBufferAdopted = false; // true once B._buf is repointed at outBuffer
133
130
 
134
131
  try {
135
- triParams = createParamsBuffer(
132
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmm-A", false);
133
+ // readback=true (COPY_SRC): BBuffer is also the source that seeds outBuffer.
134
+ BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strmm-B", true);
135
+ AdenseBuffer = createStorageBuffer(device, aOrder * ldDense * 4, "strmm-Adense");
136
+ // COPY_DST: seeded from B's own content before gemm runs, so stride-padding
137
+ // gaps (never written by gemm's tight m x n loop) keep B's original bytes
138
+ // 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,
141
+ );
142
+
143
+ triParams = createParamsBuffer(device,
136
144
  [
137
145
  { value: aOrder, type: "u32" },
138
146
  { value: lda, type: "u32" },
@@ -143,7 +151,7 @@ export async function strmm(
143
151
  ],
144
152
  "strmm-tri-params",
145
153
  );
146
- const triBindGroup = createBindGroup(triPipeline.getBindGroupLayout(0), [ABuffer, AdenseBuffer, triParams]);
154
+ const triBindGroup = createBindGroup(device, triPipeline.getBindGroupLayout(0), [ABuffer, AdenseBuffer, triParams]);
147
155
 
148
156
  // X/Y buffers and their own ld, matching swapXY above.
149
157
  const XBuffer = swapXY ? BBuffer : AdenseBuffer;
@@ -151,7 +159,7 @@ export async function strmm(
151
159
  const YBuffer = swapXY ? AdenseBuffer : BBuffer;
152
160
  const ldY = swapXY ? ldDense : ldb;
153
161
 
154
- gemmParams = createParamsBuffer(
162
+ gemmParams = createParamsBuffer(device,
155
163
  [
156
164
  { value: mg, type: "u32" },
157
165
  { value: ng, type: "u32" },
@@ -166,9 +174,16 @@ export async function strmm(
166
174
  ],
167
175
  "strmm-gemm-params",
168
176
  );
169
- const gemmBindGroup = createBindGroup(gemmPipeline.getBindGroupLayout(0), [XBuffer, YBuffer, outBuffer, gemmParams]);
170
-
171
- const { commandEncoder, querySet } = beginTimedEncoder();
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
+ ]);
185
+
186
+ const { commandEncoder, querySet } = beginTimedEncoder(device);
172
187
  // Seed outBuffer with B's own bytes first, so gemm's tight m x n write
173
188
  // leaves stride-padding gaps holding B's original content, not zero.
174
189
  // BBuffer may be larger than outBuffer (e.g. a validation-test baseline
@@ -176,13 +191,13 @@ export async function strmm(
176
191
  commandEncoder.copyBufferToBuffer(BBuffer, 0, outBuffer, 0, Math.min(BBuffer.size, outBuffer.size));
177
192
  const triDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
178
193
  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);
194
+ encodePass(commandEncoder, triPipeline, triBindGroup, { x: Math.ceil(aOrder / TILE_WG_2D), y: Math.ceil(aOrder / TILE_WG_2D) }, triDesc);
180
195
  encodePass(commandEncoder, gemmPipeline, gemmBindGroup, gemmWgCount, gemmDesc);
181
196
 
182
- const ts = resolveTimestamp(commandEncoder, querySet);
183
- const readBuffer = BIsGpu ? null : stageReadback(commandEncoder, outBuffer);
197
+ const ts = resolveTimestamp(device, commandEncoder, querySet);
198
+ const readBuffer = BIsGpu ? null : stageReadback(device, commandEncoder, outBuffer);
184
199
 
185
- submit(commandEncoder);
200
+ submit(device, commandEncoder);
186
201
 
187
202
  const gpuTimeMs = await extractTimestamp(ts);
188
203
 
@@ -201,10 +216,10 @@ export async function strmm(
201
216
  if (gpuTimeMs !== undefined) return { B: result, gpuTimeMs };
202
217
  return { B: result };
203
218
  } finally {
204
- if (!AIsGpu) destroyBuffers(ABuffer);
205
- if (!BIsGpu) destroyBuffers(BBuffer);
206
- destroyBuffers(AdenseBuffer);
207
- if (!outBufferAdopted) destroyBuffers(outBuffer);
219
+ if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
220
+ if (!BIsGpu && BBuffer) destroyBuffers(BBuffer);
221
+ if (AdenseBuffer) destroyBuffers(AdenseBuffer);
222
+ if (outBuffer && !outBufferAdopted) destroyBuffers(outBuffer);
208
223
  if (triParams) destroyBuffers(triParams);
209
224
  if (gemmParams) destroyBuffers(gemmParams);
210
225
  }
@@ -11,6 +11,7 @@ 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
15
 
15
16
  export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, incy, layout = "row-major") {
16
17
  const xIsGpu = x instanceof GpuVector;
@@ -20,6 +21,7 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
20
21
 
21
22
  if (!(device instanceof GPUDevice))
22
23
  throw new Error("device must be a GPUDevice.");
24
+ requireSameDevice(device, "strmv", { A, x, y });
23
25
  if (uplo !== "lower" && uplo !== "upper")
24
26
  throw new Error("uplo must be 'lower' or 'upper'.");
25
27
  if (trans !== "no-transpose" && trans !== "transpose")
@@ -92,10 +94,10 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
92
94
  let paramsBuffer = null;
93
95
 
94
96
  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);
98
- paramsBuffer = createParamsBuffer(
97
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strmv-A", false);
98
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "strmv-x", false);
99
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "strmv-y", true);
100
+ paramsBuffer = createParamsBuffer(device,
99
101
  [
100
102
  { value: n, type: "u32" },
101
103
  { value: incx, type: "u32" },
@@ -108,7 +110,7 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
108
110
  "strmv-params",
109
111
  );
110
112
 
111
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
113
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
112
114
  ABuffer,
113
115
  xBuffer,
114
116
  yBuffer,
@@ -116,10 +118,10 @@ export async function strmv(device, uplo, trans, diag, n, A, lda, x, incx, y, in
116
118
  ]);
117
119
 
118
120
  const wgCount = Math.min(n, device.limits.maxComputeWorkgroupsPerDimension);
119
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
120
- const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
121
+ const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
122
+ const readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
121
123
 
122
- submit(commandEncoder);
124
+ submit(device, commandEncoder);
123
125
 
124
126
  const gpuTimeMs = await extractTimestamp(ts);
125
127