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