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/sdot/sdot.mjs CHANGED
@@ -12,8 +12,9 @@ import { extractTimestamp } from "../util/benchmark.mjs";
12
12
  import { extractResult } from "../util/result.mjs";
13
13
  import { getPipeline } from "../util/pipeline.mjs";
14
14
  import { GpuVector } from "../classes/GpuVector.mjs";
15
+ import { WGS } from "../util/constants.mjs";
16
+ import { requireSameDevice } from "../util/device.mjs";
15
17
 
16
- const WGS = 64; //workgroup size
17
18
 
18
19
  export async function sdot(device, n, x, incx, y, incy) {
19
20
  const xIsGpu = x instanceof GpuVector;
@@ -21,6 +22,7 @@ export async function sdot(device, n, x, incx, y, incy) {
21
22
 
22
23
  if (!(device instanceof GPUDevice))
23
24
  throw new Error("device must be a GPUDevice.");
25
+ requireSameDevice(device, "sdot", { x, y });
24
26
  if (
25
27
  !Number.isInteger(n) ||
26
28
  !Number.isInteger(incx) ||
@@ -58,11 +60,11 @@ export async function sdot(device, n, x, incx, y, incy) {
58
60
  let readBuffer = null;
59
61
 
60
62
  try {
61
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sdot-x", false);
62
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sdot-y", false);
63
- partialsBuffer = createStorageBuffer(2 * WGS * 4, "sdot-partials"); //to hold 2*WGS partial sums of f32
64
- resultBuffer = createResultBuffer(4, "sdot-result"); //to hold the final float32 dot product
65
- paramsBuffer = createParamsBuffer(
63
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sdot-x", false);
64
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sdot-y", false);
65
+ partialsBuffer = createStorageBuffer(device, 2 * WGS * 4, "sdot-partials"); //to hold 2*WGS partial sums of f32
66
+ resultBuffer = createResultBuffer(device, 4, "sdot-result"); //to hold the final float32 dot product
67
+ paramsBuffer = createParamsBuffer(device,
66
68
  [
67
69
  { value: n, type: "u32" },
68
70
  { value: incx, type: "u32" },
@@ -71,32 +73,32 @@ export async function sdot(device, n, x, incx, y, incy) {
71
73
  "sdot-params",
72
74
  );
73
75
 
74
- const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
76
+ const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
75
77
  xBuffer,
76
78
  yBuffer,
77
79
  partialsBuffer,
78
80
  paramsBuffer,
79
81
  ]);
80
- const { commandEncoder: enc1, ts: ts1 } = runComputePass(
82
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(device,
81
83
  pipelineMain,
82
84
  bgMain,
83
85
  2 * WGS,
84
86
  ); //dispatch 2*WGS workgroups
85
87
 
86
- submit(enc1);
88
+ submit(device, enc1);
87
89
 
88
- const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
90
+ const bgReduce = createBindGroup(device, pipelineReduce.getBindGroupLayout(0), [
89
91
  partialsBuffer,
90
92
  resultBuffer,
91
93
  ]);
92
- const { commandEncoder: enc2, ts: ts2 } = runComputePass(
94
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(device,
93
95
  pipelineReduce,
94
96
  bgReduce,
95
97
  1,
96
98
  ); // dispatch 1 workgroup to reduce the partial sums to a single result
97
- readBuffer = stageReadback(enc2, resultBuffer);
99
+ readBuffer = stageReadback(device, enc2, resultBuffer);
98
100
 
99
- submit(enc2);
101
+ submit(device, enc2);
100
102
 
101
103
  const resultPromise = extractResult(readBuffer, Float32Array);
102
104
  readBuffer = null; // ownership transferred — extractResult's own finally destroys it
@@ -117,7 +119,7 @@ export async function sdot(device, n, x, incx, y, incy) {
117
119
  if (partialsBuffer) destroyBuffers(partialsBuffer);
118
120
  if (resultBuffer) destroyBuffers(resultBuffer);
119
121
  if (paramsBuffer) destroyBuffers(paramsBuffer);
120
- // Only reached if submit(enc2) threw before ownership was transferred above.
122
+ // Only reached if submit(device, enc2) threw before ownership was transferred above.
121
123
  if (readBuffer) destroyBuffers(readBuffer);
122
124
  }
123
125
  }
@@ -3,6 +3,8 @@ import {
3
3
  createParamsBuffer,
4
4
  stageReadback,
5
5
  destroyBuffers,
6
+ vec4ViewBinding,
7
+ vec4Usable,
6
8
  } from "../util/buffer.mjs";
7
9
  import { createBindGroup } from "../util/bindgroup.mjs";
8
10
  import { runComputePass, submit } from "../util/compute.mjs";
@@ -10,10 +12,10 @@ import { extractResult } from "../util/result.mjs";
10
12
  import { extractTimestamp } from "../util/benchmark.mjs";
11
13
  import { getPipeline } from "../util/pipeline.mjs";
12
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 { requireSameDevice } from "../util/device.mjs";
13
18
 
14
- const BM_SMALL = 32, BN_SMALL = 32; // sgemm_small.wgsl's block tile
15
- const BM_LARGE = 64, BN_LARGE = 64; // sgemm_large.wgsl's block tile
16
- const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // large tile needs >= a 6x6 grid of its own tiles to beat the small tile
17
19
 
18
20
  export async function sgemm(
19
21
  device, transA, transB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
@@ -24,6 +26,7 @@ export async function sgemm(
24
26
 
25
27
  if (!(device instanceof GPUDevice))
26
28
  throw new Error("device must be a GPUDevice.");
29
+ requireSameDevice(device, "sgemm", { A, B, C });
27
30
  if (transA !== "no-transpose" && transA !== "transpose")
28
31
  throw new Error("transA must be 'no-transpose' or 'transpose'.");
29
32
  if (transB !== "no-transpose" && transB !== "transpose")
@@ -135,10 +138,16 @@ export async function sgemm(
135
138
 
136
139
  const pipeline = await getPipeline(device, useLargeTile ? "sgemm_large" : "sgemm_small");
137
140
 
138
- const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sgemm-A", false);
139
- const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "sgemm-B", false);
140
- const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "sgemm-C", true);
141
- const paramsBuffer = createParamsBuffer(
141
+ const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemm-A", false);
142
+ const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "sgemm-B", false);
143
+ const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "sgemm-C", true);
144
+ // Vectorized-load enablement — kernel-side view after the column-major swap.
145
+ // op(A) is m×k (no-trans) or k×m (trans); op(B) is k×n or n×k.
146
+ const aNot = transA === "no-transpose";
147
+ const bNot = transB === "no-transpose";
148
+ const useVecA = aNot && vec4Usable(ABuffer, lda, m, k); // transposed-A vec stores bank-conflict smem — measured slower than scalar
149
+ const useVecB = vec4Usable(BBuffer, ldb, bNot ? k : n, bNot ? n : k);
150
+ const paramsBuffer = createParamsBuffer(device,
142
151
  [
143
152
  { value: m, type: "u32" },
144
153
  { value: n, type: "u32" },
@@ -150,31 +159,35 @@ export async function sgemm(
150
159
  { value: ldc, type: "u32" },
151
160
  { value: transA === "transpose" ? 1 : 0, type: "u32" },
152
161
  { value: transB === "transpose" ? 1 : 0, type: "u32" },
162
+ { value: useVecA ? 1 : 0, type: "u32" },
163
+ { value: useVecB ? 1 : 0, type: "u32" },
153
164
  ],
154
165
  "sgemm-params",
155
166
  );
156
167
 
157
168
  try {
158
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
169
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
159
170
  ABuffer,
171
+ vec4ViewBinding(device, ABuffer),
160
172
  BBuffer,
173
+ vec4ViewBinding(device, BBuffer),
161
174
  CBuffer,
162
175
  paramsBuffer,
163
176
  ]);
164
177
 
165
178
  const wgCount = useLargeTile
166
179
  ? {
167
- x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension),
168
- y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension),
180
+ x: requireWorkgroupCount(device, largeWgX, "sgemm", "x"),
181
+ y: requireWorkgroupCount(device, largeWgY, "sgemm", "y"),
169
182
  }
170
183
  : {
171
- x: Math.min(Math.ceil(n / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
172
- y: Math.min(Math.ceil(m / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
184
+ x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "sgemm", "x"),
185
+ y: requireWorkgroupCount(device, Math.ceil(m / BM_SMALL), "sgemm", "y"),
173
186
  };
174
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
175
- const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
187
+ const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
188
+ const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
176
189
 
177
- submit(commandEncoder);
190
+ submit(device, commandEncoder);
178
191
 
179
192
  const gpuTimeMs = await extractTimestamp(ts);
180
193
 
@@ -10,10 +10,10 @@ import { extractResult } from "../util/result.mjs";
10
10
  import { 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 uses — see sgemm.mjs
17
17
 
18
18
  export async function sgemmtr(
19
19
  device, uplo, transA, transB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
@@ -24,6 +24,7 @@ export async function sgemmtr(
24
24
 
25
25
  if (!(device instanceof GPUDevice))
26
26
  throw new Error("device must be a GPUDevice.");
27
+ requireSameDevice(device, "sgemmtr", { A, B, C });
27
28
  if (uplo !== "lower" && uplo !== "upper")
28
29
  throw new Error("uplo must be 'lower' or 'upper'.");
29
30
  if (transA !== "no-transpose" && transA !== "transpose")
@@ -142,10 +143,10 @@ export async function sgemmtr(
142
143
 
143
144
  const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
144
145
 
145
- const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sgemmtr-A", false);
146
- const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "sgemmtr-B", false);
147
- const CBuffer = CIsGpu ? C._buf : uploadBuffer(C, "sgemmtr-C", true);
148
- const paramsBuffer = createParamsBuffer(
146
+ const ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemmtr-A", false);
147
+ const BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "sgemmtr-B", false);
148
+ const CBuffer = CIsGpu ? C._buf : uploadBuffer(device, C, "sgemmtr-C", true);
149
+ const paramsBuffer = createParamsBuffer(device,
149
150
  [
150
151
  { value: m, type: "u32" },
151
152
  { value: n, type: "u32" },
@@ -163,7 +164,7 @@ export async function sgemmtr(
163
164
  );
164
165
 
165
166
  try {
166
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
167
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
167
168
  ABuffer,
168
169
  BBuffer,
169
170
  CBuffer,
@@ -172,17 +173,17 @@ export async function sgemmtr(
172
173
 
173
174
  const wgCount = useLargeTile
174
175
  ? {
175
- x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension),
176
- y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension),
176
+ x: requireWorkgroupCount(device, largeWgX, "sgemmtr", "x"),
177
+ y: requireWorkgroupCount(device, largeWgY, "sgemmtr", "y"),
177
178
  }
178
179
  : {
179
- x: Math.min(Math.ceil(n / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
180
- y: Math.min(Math.ceil(m / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension),
180
+ x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "sgemmtr", "x"),
181
+ y: requireWorkgroupCount(device, Math.ceil(m / BM_SMALL), "sgemmtr", "y"),
181
182
  };
182
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
183
- const readBuffer = CIsGpu ? null : stageReadback(commandEncoder, CBuffer);
183
+ const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
184
+ const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
184
185
 
185
- submit(commandEncoder);
186
+ submit(device, commandEncoder);
186
187
 
187
188
  const gpuTimeMs = await extractTimestamp(ts);
188
189
 
@@ -9,9 +9,10 @@ import { runComputePass, submit } from "../util/compute.mjs";
9
9
  import { extractResult } from "../util/result.mjs";
10
10
  import { extractTimestamp } from "../util/benchmark.mjs";
11
11
  import { getPipeline } from "../util/pipeline.mjs";
12
- import { calcWorkgroups } from "../util/workgroup.mjs";
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
16
 
16
17
  export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y, incy, layout = "row-major") {
17
18
  const AIsGpu = A instanceof GpuMatrix;
@@ -20,6 +21,7 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
20
21
 
21
22
  if (!(device instanceof GPUDevice))
22
23
  throw new Error("device must be a GPUDevice.");
24
+ requireSameDevice(device, "sgemv", { A, x, y });
23
25
  if (trans !== "no-transpose" && trans !== "transpose")
24
26
  throw new Error("trans must be 'no-transpose' or 'transpose'.");
25
27
  if (layout !== "row-major" && layout !== "column-major")
@@ -62,6 +64,8 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
62
64
  );
63
65
  if (xIsGpu && x._buf === y._buf)
64
66
  throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");
67
+ if (AIsGpu && yIsGpu && A._buf === y._buf)
68
+ throw new Error("A and y must not reference the same GPU buffer.");
65
69
  if (AIsGpu && lda !== A.lda)
66
70
  throw new Error("lda must match A.lda when A is a GpuMatrix.");
67
71
  if (AIsGpu && (A.rows < m || A.cols < n))
@@ -99,38 +103,46 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
99
103
  const shaderName = isNoTrans ? "sgemv_n" : "sgemv_t";
100
104
  const pipeline = await getPipeline(device, shaderName);
101
105
 
102
- const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sgemv-A", false);
103
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sgemv-x", false);
104
- const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sgemv-y", true);
105
- const paramsBuffer = createParamsBuffer(
106
- [
107
- { value: m, type: "u32" },
108
- { value: n, type: "u32" },
109
- { value: alpha, type: "f32" },
110
- { value: beta, type: "f32" },
111
- { value: incx, type: "u32" },
112
- { value: incy, type: "u32" },
113
- { value: lda, type: "u32" },
114
- ],
115
- "sgemv-params",
116
- );
106
+ let ABuffer = null;
107
+ let xBuffer = null;
108
+ let yBuffer = null;
109
+ let paramsBuffer = null;
117
110
 
118
111
  try {
119
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
112
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sgemv-A", false);
113
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sgemv-x", false);
114
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sgemv-y", true);
115
+ paramsBuffer = createParamsBuffer(device,
116
+ [
117
+ { value: m, type: "u32" },
118
+ { value: n, type: "u32" },
119
+ { value: alpha, type: "f32" },
120
+ { value: beta, type: "f32" },
121
+ { value: incx, type: "u32" },
122
+ { value: incy, type: "u32" },
123
+ { value: lda, type: "u32" },
124
+ ],
125
+ "sgemv-params",
126
+ );
127
+
128
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
120
129
  ABuffer,
121
130
  xBuffer,
122
131
  yBuffer,
123
132
  paramsBuffer,
124
133
  ]);
125
134
 
126
- // NoTrans: one workgroup per row (grid-stride handles overflow); Trans: one thread per output column.
135
+ // NoTrans: one workgroup per row — sgemv_n.wgsl is a grid-stride loop, so
136
+ // clamping here only costs parallelism. Trans: one thread per output
137
+ // column, and sgemv_t.wgsl indexes straight off global_invocation_id with
138
+ // no fallback, so an over-limit dispatch must be refused, not truncated.
127
139
  const wgCount = isNoTrans
128
140
  ? Math.min(m, device.limits.maxComputeWorkgroupsPerDimension)
129
- : calcWorkgroups(yLen);
130
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
131
- const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
141
+ : requireWorkgroups(device, "sgemv", yLen);
142
+ const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
143
+ const readBuffer = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
132
144
 
133
- submit(commandEncoder);
145
+ submit(device, commandEncoder);
134
146
 
135
147
  const gpuTimeMs = await extractTimestamp(ts);
136
148
 
@@ -143,10 +155,10 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
143
155
  if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
144
156
  return { y: result };
145
157
  } finally {
146
- if (!AIsGpu) destroyBuffers(ABuffer);
147
- if (!xIsGpu) destroyBuffers(xBuffer);
148
- if (!yIsGpu) destroyBuffers(yBuffer);
149
- destroyBuffers(paramsBuffer);
158
+ if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
159
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
160
+ if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
161
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
150
162
 
151
163
  }
152
164
  }
package/src/sger/sger.mjs CHANGED
@@ -11,12 +11,14 @@ 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 sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout = "row-major") {
16
17
  const AIsGpu = A instanceof GpuMatrix;
17
18
 
18
19
  if (!(device instanceof GPUDevice))
19
20
  throw new Error("device must be a GPUDevice.");
21
+ requireSameDevice(device, "sger", { A, x, y });
20
22
  if (layout !== "row-major" && layout !== "column-major")
21
23
  throw new Error("layout must be 'row-major' or 'column-major'.");
22
24
  if (typeof alpha !== "number")
@@ -89,10 +91,10 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
89
91
  let paramsBuffer = null;
90
92
 
91
93
  try {
92
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sger-x", false);
93
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sger-y", false);
94
- ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sger-A", true);
95
- paramsBuffer = createParamsBuffer(
94
+ xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sger-x", false);
95
+ yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sger-y", false);
96
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sger-A", true);
97
+ paramsBuffer = createParamsBuffer(device,
96
98
  [
97
99
  { value: m, type: "u32" },
98
100
  { value: n, type: "u32" },
@@ -104,7 +106,7 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
104
106
  "sger-params",
105
107
  );
106
108
 
107
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
109
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
108
110
  xBuffer,
109
111
  yBuffer,
110
112
  ABuffer,
@@ -114,10 +116,10 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
114
116
  // One workgroup per row of A; clamped to device limit — the shader's
115
117
  // grid-stride loop handles remaining rows when m > dispatch count.
116
118
  const wgCount = Math.min(m, device.limits.maxComputeWorkgroupsPerDimension);
117
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
118
- const readBuffer = AIsGpu ? null : stageReadback(commandEncoder, ABuffer);
119
+ const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
120
+ const readBuffer = AIsGpu ? null : stageReadback(device, commandEncoder, ABuffer);
119
121
 
120
- submit(commandEncoder);
122
+ submit(device, commandEncoder);
121
123
 
122
124
  const gpuTimeMs = await extractTimestamp(ts);
123
125
 
@@ -3,25 +3,175 @@
3
3
  *
4
4
  * `shaders/*.wgsl` — one WGSL compute shader per BLAS routine (sscal, saxpy, sdot, …).
5
5
  *
6
- * `shaders/browser-shaders.mjs` — the browser's runtime shader source. In Node.js, shaders are
7
- * read directly from disk via `readFileSync`. In the browser there is no filesystem, so this file
8
- * provides all shader strings inline. Vite bundles it by importing each `.wgsl` file as a string.
6
+ * `routineShaders` below is the single source of truth: routine name → the WGSL source(s)
7
+ * its `getPipeline()` calls actually reference, verified against every `src/<routine>/<routine>.mjs`
8
+ * rather than inferred from naming convention (see its doc comment for the exceptions). Each
9
+ * shader is imported right above the line that adds it — the import *is* the mapping entry, no
10
+ * separate block to cross-reference. `shaderSources`, the flat name → source registry the
11
+ * browser bundle's runtime lookup needs, is *derived* from `routineShaders` rather than
12
+ * hand-duplicated, so the two can never drift apart. In Node.js neither is read — shaders are
13
+ * `readFileSync` from disk directly; `scripts/build-browser.mjs` inlines this module into the
14
+ * browser's IIFE bundle via esbuild instead.
9
15
  *
10
16
  * ## Cross-shader patterns
11
17
  *
12
- * **Fixed workgroup size of 64.** Every shader declares `const WGS: u32 = 64` and
13
- * `@workgroup_size(64)`. 64 is the minimum `maxComputeInvocationsPerWorkgroup` guaranteed across
14
- * all WebGPU devices, so this works everywhere without querying device limits.
18
+ * **Single bind group.** Every shader with bindings uses `@group(0)` only — the JS side always
19
+ * calls `pipeline.getBindGroupLayout(0)`, no secondary groups to track. Binding order is
20
+ * consistent too: any read-only storage buffers come before read_write ones, with the
21
+ * `uniform Params` struct always last. `@binding` indices match the position of each resource in
22
+ * the array passed to `createBindGroup`, which appends `resultBuffer` last.
15
23
  *
16
- * **Single bind group.** All bindings use `@group(0)`. This means the JS side always calls
17
- * `pipeline.getBindGroupLayout(0)` — no secondary groups to track.
24
+ * **All counts and strides are `u32`.** `n`, `x_inc`, `y_inc`, and every other index/count field
25
+ * in a `Params` struct is unsigned, avoiding implicit sign-extension in index expressions like
26
+ * `id * params.x_inc`.
18
27
  *
19
- * The `@binding` indices must match the position of each resource in the array passed to
20
- * `createBindGroup` — it assigns `binding: 0, 1, 2 …` sequentially, with `resultBuffer` appended last.
21
- *
22
- * **All counts and strides are `u32`.** `n`, `x_inc`, `y_inc`, and any other index fields in the
23
- * `Params` uniform struct are unsigned. This avoids implicit sign-extension when they appear in
24
- * index expressions like `id * params.x_inc`.
28
+ * **Entry points don't have to be named `main`.** `loadShader` (`util/pipeline.mjs`)
29
+ * auto-detects the sole `@compute` function in a module instead of requiring a fixed name, so
30
+ * `dasum_main`, `strsv_invert_block_main`, etc. work without renaming.
25
31
  *
26
32
  * @module devdocs/shaders
27
33
  */
34
+
35
+ /**
36
+ * Routine name → the WGSL source(s) its `getPipeline()` calls reference. Keys are the exact
37
+ * shader names `getPipeline(device, name)` is called with — most routines have one, some pick
38
+ * one of several conditionally (e.g. sgemv's `sgemv_n`/`sgemv_t`, by `trans`), and some have no
39
+ * dedicated shader at all:
40
+ *
41
+ * - `sgemmtr`/`ssyrk`/`ssyr2k` all dispatch through `sgemmtr_small`/`sgemmtr_large`.
42
+ * - `strsm` reuses `strsv_invert_block` and `sscal`, plus its own `block_transfer` and the
43
+ * shared `sgemm_small`/`sgemm_large`.
44
+ * - `dasum`/`idamax` concatenate several f64 utility shaders with their own — see
45
+ * `getPipeline`'s `shaderName: string[]` behaviour.
46
+ * - `random` has no entry — CPU-only, no `getPipeline()` call.
47
+ *
48
+ * Built up entry by entry so each import sits next to the mapping entry that uses it.
49
+ * @public
50
+ */
51
+ export const routineShaders = {};
52
+
53
+ import sscal from "./sscal.wgsl";
54
+ routineShaders.sscal = { sscal };
55
+
56
+ import sswap from "./sswap.wgsl";
57
+ routineShaders.sswap = { sswap };
58
+
59
+ import saxpy from "./saxpy.wgsl";
60
+ routineShaders.saxpy = { saxpy };
61
+
62
+ import scopy from "./scopy.wgsl";
63
+ routineShaders.scopy = { scopy };
64
+
65
+ import sdot from "./sdot.wgsl";
66
+ import sum from "./reduction/sum.wgsl";
67
+ routineShaders.sdot = { sdot, "reduction/sum": sum };
68
+
69
+ import sasum from "./sasum.wgsl";
70
+ routineShaders.sasum = { sasum, "reduction/sum": sum };
71
+
72
+ import snrm2 from "./snrm2.wgsl";
73
+ import scaledSum from "./reduction/scaledSum.wgsl";
74
+ routineShaders.snrm2 = { snrm2, "reduction/scaledSum": scaledSum };
75
+
76
+ import isamax from "./isamax.wgsl";
77
+ import argmax from "./reduction/argmax.wgsl";
78
+ routineShaders.isamax = { isamax, "reduction/argmax": argmax };
79
+
80
+ import dekker from "./f64/dekker.wgsl";
81
+ import ddAbs from "./f64/utils/abs.wgsl";
82
+ import ddAddUtil from "./f64/utils/add.wgsl";
83
+ import dasum from "./dasum.wgsl";
84
+ import sumF64 from "./reduction/sumF64.wgsl";
85
+ routineShaders.dasum = {
86
+ "f64/dekker": dekker,
87
+ "f64/utils/abs": ddAbs,
88
+ "f64/utils/add": ddAddUtil,
89
+ dasum,
90
+ "reduction/sumF64": sumF64,
91
+ };
92
+
93
+ import ddGreater from "./f64/utils/greater.wgsl";
94
+ import ddEqual from "./f64/utils/equal.wgsl";
95
+ import idamax from "./idamax.wgsl";
96
+ import argmaxF64 from "./reduction/argmaxF64.wgsl";
97
+ routineShaders.idamax = {
98
+ "f64/dekker": dekker,
99
+ "f64/utils/abs": ddAbs,
100
+ "f64/utils/greater": ddGreater,
101
+ "f64/utils/equal": ddEqual,
102
+ idamax,
103
+ "reduction/argmaxF64": argmaxF64,
104
+ };
105
+
106
+ import srot from "./srot.wgsl";
107
+ routineShaders.srot = { srot };
108
+
109
+ import srotm from "./srotm.wgsl";
110
+ routineShaders.srotm = { srotm };
111
+
112
+ import sgemv_n from "./sgemv_n.wgsl";
113
+ import sgemv_t from "./sgemv_t.wgsl";
114
+ routineShaders.sgemv = { sgemv_n, sgemv_t }; // one or the other, picked by trans
115
+
116
+ import ssymv from "./ssymv.wgsl";
117
+ routineShaders.ssymv = { ssymv };
118
+
119
+ import strmv from "./strmv.wgsl";
120
+ routineShaders.strmv = { strmv };
121
+
122
+ import strsv_invert_block from "./strsv_invert_block.wgsl";
123
+ import strsv_apply_inverse from "./strsv_apply_inverse.wgsl";
124
+ import strsv_update from "./strsv_update.wgsl";
125
+ routineShaders.strsv = {
126
+ strsv_invert_block,
127
+ strsv_apply_inverse,
128
+ strsv_update,
129
+ };
130
+
131
+ import sger from "./sger.wgsl";
132
+ routineShaders.sger = { sger };
133
+
134
+ import ssyr from "./ssyr.wgsl";
135
+ routineShaders.ssyr = { ssyr };
136
+
137
+ import ssyr2 from "./ssyr2.wgsl";
138
+ routineShaders.ssyr2 = { ssyr2 };
139
+
140
+ import sgemm_small from "./sgemm_small.wgsl";
141
+ import sgemm_large from "./sgemm_large.wgsl";
142
+ routineShaders.sgemm = { sgemm_small, sgemm_large }; // one or the other, picked by a tile-size threshold
143
+
144
+ import sgemmtr_small from "./sgemmtr_small.wgsl";
145
+ import sgemmtr_large from "./sgemmtr_large.wgsl";
146
+ routineShaders.sgemmtr = { sgemmtr_small, sgemmtr_large };
147
+
148
+ routineShaders.ssyrk = { sgemmtr_small, sgemmtr_large }; // no shader of its own — rides on sgemmtr's
149
+ routineShaders.ssyr2k = { sgemmtr_small, sgemmtr_large }; // no shader of its own — rides on sgemmtr's
150
+
151
+ import symmetrize from "./symmetrize.wgsl";
152
+ routineShaders.ssymm = { sgemm_small, sgemm_large, symmetrize };
153
+
154
+ import triangularize from "./triangularize.wgsl";
155
+ routineShaders.strmm = { sgemm_small, sgemm_large, triangularize };
156
+
157
+ import blockTransfer from "./block_transfer.wgsl";
158
+ routineShaders.strsm = {
159
+ strsv_invert_block,
160
+ block_transfer: blockTransfer,
161
+ sscal,
162
+ sgemm_small,
163
+ sgemm_large,
164
+ };
165
+
166
+ /**
167
+ * Flat shader-name → WGSL source-string registry — what `getPipeline()`/`loadShader()` (see
168
+ * `util/pipeline.mjs`) actually look shaders up in, in the browser. Derived from
169
+ * `routineShaders` by merging every routine's shaders together; shared shaders (e.g.
170
+ * `"reduction/sum"`, used by two different routines above) collapse harmlessly here since
171
+ * every routine's copy is the same imported string, never independently authored text.
172
+ * @public
173
+ */
174
+ export const shaderSources = Object.assign(
175
+ {},
176
+ ...Object.values(routineShaders),
177
+ );