wgblas 1.2.1 → 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 (96) hide show
  1. package/LICENSE +1 -1
  2. package/README.md +54 -72
  3. package/dist/wgblas.browser.js +1637 -858
  4. package/index.d.mts +51 -6
  5. package/index.mjs +8 -0
  6. package/package.json +56 -2
  7. package/src/classes/GpuMatrix.mjs +17 -10
  8. package/src/classes/GpuVector.mjs +28 -10
  9. package/src/dasum/dasum.d.mts +6 -6
  10. package/src/dasum/dasum.mjs +25 -21
  11. package/src/devdocs.mjs +13 -0
  12. package/src/idamax/idamax.d.mts +69 -0
  13. package/src/idamax/idamax.mjs +130 -0
  14. package/src/init.mjs +115 -49
  15. package/src/isamax/isamax.d.mts +21 -3
  16. package/src/isamax/isamax.mjs +16 -14
  17. package/src/random/random.d.mts +1 -0
  18. package/src/sasum/sasum.d.mts +3 -3
  19. package/src/sasum/sasum.mjs +15 -13
  20. package/src/saxpy/saxpy.d.mts +3 -3
  21. package/src/saxpy/saxpy.mjs +10 -8
  22. package/src/scopy/scopy.d.mts +3 -3
  23. package/src/scopy/scopy.mjs +10 -8
  24. package/src/sdot/sdot.d.mts +3 -3
  25. package/src/sdot/sdot.mjs +16 -14
  26. package/src/sgemm/sgemm.d.mts +102 -0
  27. package/src/sgemm/sgemm.mjs +208 -0
  28. package/src/sgemmtr/sgemmtr.d.mts +104 -0
  29. package/src/sgemmtr/sgemmtr.mjs +204 -0
  30. package/src/sgemv/sgemv.d.mts +1 -38
  31. package/src/sgemv/sgemv.mjs +42 -26
  32. package/src/sger/sger.d.mts +1 -34
  33. package/src/sger/sger.mjs +12 -8
  34. package/src/shaders/block_transfer.wgsl +42 -0
  35. package/src/shaders/dasum.wgsl +3 -2
  36. package/src/shaders/f64/dekker.wgsl +4 -85
  37. package/src/shaders/f64/utils/abs.wgsl +10 -0
  38. package/src/shaders/f64/utils/add.wgsl +77 -0
  39. package/src/shaders/f64/utils/equal.wgsl +7 -0
  40. package/src/shaders/f64/utils/greater.wgsl +12 -0
  41. package/src/shaders/f64/utils/multiply.wgsl +81 -0
  42. package/src/shaders/idamax.wgsl +96 -0
  43. package/src/shaders/index.mjs +164 -14
  44. package/src/shaders/reduction/argmaxF64.wgsl +50 -0
  45. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  46. package/src/shaders/reduction/sumF64.wgsl +2 -2
  47. package/src/shaders/sgemm_large.wgsl +206 -0
  48. package/src/shaders/sgemm_small.wgsl +212 -0
  49. package/src/shaders/sgemmtr_large.wgsl +120 -0
  50. package/src/shaders/sgemmtr_small.wgsl +113 -0
  51. package/src/shaders/sgemv_n.wgsl +3 -1
  52. package/src/shaders/sgemv_t.wgsl +3 -1
  53. package/src/shaders/snrm2.wgsl +72 -23
  54. package/src/shaders/ssymv.wgsl +3 -1
  55. package/src/shaders/symmetrize.wgsl +31 -0
  56. package/src/shaders/triangularize.wgsl +44 -0
  57. package/src/snrm2/snrm2.d.mts +3 -3
  58. package/src/snrm2/snrm2.mjs +33 -21
  59. package/src/srot/srot.d.mts +3 -5
  60. package/src/srot/srot.mjs +11 -9
  61. package/src/srotm/srotm.d.mts +3 -5
  62. package/src/srotm/srotm.mjs +19 -10
  63. package/src/sscal/sscal.d.mts +4 -4
  64. package/src/sscal/sscal.mjs +12 -10
  65. package/src/sswap/sswap.d.mts +3 -3
  66. package/src/sswap/sswap.mjs +11 -9
  67. package/src/ssymm/ssymm.d.mts +103 -0
  68. package/src/ssymm/ssymm.mjs +218 -0
  69. package/src/ssymv/ssymv.d.mts +1 -36
  70. package/src/ssymv/ssymv.mjs +12 -8
  71. package/src/ssyr/ssyr.d.mts +1 -30
  72. package/src/ssyr/ssyr.mjs +11 -7
  73. package/src/ssyr2/ssyr2.d.mts +1 -34
  74. package/src/ssyr2/ssyr2.mjs +12 -8
  75. package/src/ssyr2k/ssyr2k.d.mts +100 -0
  76. package/src/ssyr2k/ssyr2k.mjs +202 -0
  77. package/src/ssyrk/ssyrk.d.mts +90 -0
  78. package/src/ssyrk/ssyrk.mjs +177 -0
  79. package/src/strmm/strmm.d.mts +100 -0
  80. package/src/strmm/strmm.mjs +226 -0
  81. package/src/strmv/strmv.d.mts +1 -36
  82. package/src/strmv/strmv.mjs +12 -8
  83. package/src/strsm/strsm.d.mts +99 -0
  84. package/src/strsm/strsm.mjs +360 -0
  85. package/src/strsv/strsv.d.mts +1 -32
  86. package/src/strsv/strsv.mjs +18 -12
  87. package/src/util/benchmark.mjs +4 -6
  88. package/src/util/bindgroup.mjs +1 -3
  89. package/src/util/buffer.mjs +116 -20
  90. package/src/util/compute.mjs +12 -12
  91. package/src/util/constants.mjs +57 -0
  92. package/src/util/device.mjs +34 -0
  93. package/src/util/f64.mjs +3 -3
  94. package/src/util/pipeline.mjs +5 -6
  95. package/src/util/workgroup.mjs +55 -7
  96. package/src/shaders/browser-shaders.mjs +0 -55
@@ -0,0 +1,204 @@
1
+ import {
2
+ uploadBuffer,
3
+ createParamsBuffer,
4
+ stageReadback,
5
+ destroyBuffers,
6
+ } from "../util/buffer.mjs";
7
+ import { createBindGroup } from "../util/bindgroup.mjs";
8
+ import { runComputePass, submit } from "../util/compute.mjs";
9
+ import { extractResult } from "../util/result.mjs";
10
+ import { extractTimestamp } from "../util/benchmark.mjs";
11
+ import { getPipeline } from "../util/pipeline.mjs";
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";
16
+
17
+
18
+ export async function sgemmtr(
19
+ device, uplo, transA, transB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, layout = "row-major",
20
+ ) {
21
+ let AIsGpu = A instanceof GpuMatrix;
22
+ let BIsGpu = B instanceof GpuMatrix;
23
+ const CIsGpu = C instanceof GpuMatrix;
24
+
25
+ if (!(device instanceof GPUDevice))
26
+ throw new Error("device must be a GPUDevice.");
27
+ requireSameDevice(device, "sgemmtr", { A, B, C });
28
+ if (uplo !== "lower" && uplo !== "upper")
29
+ throw new Error("uplo must be 'lower' or 'upper'.");
30
+ if (transA !== "no-transpose" && transA !== "transpose")
31
+ throw new Error("transA must be 'no-transpose' or 'transpose'.");
32
+ if (transB !== "no-transpose" && transB !== "transpose")
33
+ throw new Error("transB must be 'no-transpose' or 'transpose'.");
34
+ if (layout !== "row-major" && layout !== "column-major")
35
+ throw new Error("layout must be 'row-major' or 'column-major'.");
36
+ if (typeof alpha !== "number")
37
+ throw new Error("alpha must be a number.");
38
+ if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
39
+ if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
40
+ if (typeof beta !== "number")
41
+ throw new Error("beta must be a number.");
42
+ if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
43
+ if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
44
+ if (
45
+ !Number.isInteger(m) ||
46
+ !Number.isInteger(n) ||
47
+ !Number.isInteger(k) ||
48
+ !Number.isInteger(lda) ||
49
+ !Number.isInteger(ldb) ||
50
+ !Number.isInteger(ldc)
51
+ )
52
+ throw new Error("m, n, k, lda, ldb, and ldc must be integers.");
53
+ if (!AIsGpu && !(A instanceof Float32Array))
54
+ throw new Error("A must be a Float32Array or GpuMatrix.");
55
+ if (!BIsGpu && !(B instanceof Float32Array))
56
+ throw new Error("B must be a Float32Array or GpuMatrix.");
57
+ if (!CIsGpu && !(C instanceof Float32Array))
58
+ throw new Error("C must be a Float32Array or GpuMatrix.");
59
+ if ((AIsGpu || BIsGpu) && !CIsGpu)
60
+ throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");
61
+ if (CIsGpu && (!AIsGpu || !BIsGpu))
62
+ throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");
63
+ if (m < 0 || n < 0 || k < 0) throw new Error("m, n, and k must be non-negative.");
64
+ if (m === 0 || n === 0) return CIsGpu ? {} : { C };
65
+
66
+ // GpuMatrix's own .layout wins over the shared `layout` argument.
67
+ const effLayoutA = AIsGpu ? A.layout : layout;
68
+ const effLayoutB = BIsGpu ? B.layout : layout;
69
+ const effLayoutC = CIsGpu ? C.layout : layout;
70
+
71
+ // Shape validation, before any layout-driven swapping below.
72
+ // A: op(A) is m x k; A itself is m x k or k x m depending on transA.
73
+ const aRows = effLayoutA === "column-major" ? k : m;
74
+ const aCols = effLayoutA === "column-major" ? m : k;
75
+ const aOuter = transA === "no-transpose" ? aRows : aCols; // # of stored chunks
76
+ const aInner = transA === "no-transpose" ? aCols : aRows; // required chunk length
77
+ if (lda < aInner)
78
+ throw new Error(`lda must be >= ${effLayoutA === "column-major" ? "rows" : "cols"} of A as stored.`);
79
+ if (AIsGpu) {
80
+ if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
81
+ const [aLogRows, aLogCols] = transA === "no-transpose" ? [m, k] : [k, m];
82
+ if (A.rows < aLogRows || A.cols < aLogCols)
83
+ throw new Error("A is too small for the given m, k, and transA.");
84
+ } else if (A.length < (aOuter - 1) * lda + aInner) {
85
+ throw new Error("A does not have enough elements for the given dimensions and lda.");
86
+ }
87
+
88
+ // B: same reasoning as A, with op(B) = k x n.
89
+ const bRows = effLayoutB === "column-major" ? n : k;
90
+ const bCols = effLayoutB === "column-major" ? k : n;
91
+ const bOuter = transB === "no-transpose" ? bRows : bCols;
92
+ const bInner = transB === "no-transpose" ? bCols : bRows;
93
+ if (ldb < bInner)
94
+ throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
95
+ if (BIsGpu) {
96
+ if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
97
+ const [bLogRows, bLogCols] = transB === "no-transpose" ? [k, n] : [n, k];
98
+ if (B.rows < bLogRows || B.cols < bLogCols)
99
+ throw new Error("B is too small for the given n, k, and transB.");
100
+ } else if (B.length < (bOuter - 1) * ldb + bInner) {
101
+ throw new Error("B does not have enough elements for the given dimensions and ldb.");
102
+ }
103
+
104
+ // C: always m x n (no trans flag) — layout only affects lda/storage order.
105
+ const cOuter = effLayoutC === "column-major" ? n : m;
106
+ const cInner = effLayoutC === "column-major" ? m : n;
107
+ if (ldc < cInner)
108
+ throw new Error(`ldc must be >= ${effLayoutC === "column-major" ? "rows" : "cols"} of C as stored.`);
109
+ if (CIsGpu) {
110
+ if (ldc !== C.lda) throw new Error("ldc must match C.lda when C is a GpuMatrix.");
111
+ if (C.rows < m || C.cols < n) throw new Error("C is too small for the given m and n.");
112
+ } else if (C.length < (cOuter - 1) * ldc + cInner) {
113
+ throw new Error("C does not have enough elements for the given dimensions and ldc.");
114
+ }
115
+
116
+ // Column-major A/B reinterpreted row-major is A^T/B^T — flip the trans flag.
117
+ if (effLayoutA === "column-major")
118
+ transA = transA === "no-transpose" ? "transpose" : "no-transpose";
119
+ if (effLayoutB === "column-major")
120
+ transB = transB === "no-transpose" ? "transpose" : "no-transpose";
121
+
122
+ // Column-major C: compute C^T = op(B)^T * op(A)^T instead (swap A/B, flip trans, swap m<->n)
123
+ // — same trick sgemm uses. uplo(C) in row/col terms becomes uplo(C^T) with row/col swapped,
124
+ // i.e. the opposite triangle, so uplo must flip here too (same reasoning ssyr's isLower flip
125
+ // uses for column-major A) — nothing else that's uplo-specific needs to change, since the
126
+ // shader's row/col test is applied to whatever (m, n) grid it's actually given.
127
+ if (effLayoutC === "column-major") {
128
+ [A, B] = [B, A];
129
+ [AIsGpu, BIsGpu] = [BIsGpu, AIsGpu];
130
+ [lda, ldb] = [ldb, lda];
131
+ [transA, transB] = [
132
+ transB === "no-transpose" ? "transpose" : "no-transpose",
133
+ transA === "no-transpose" ? "transpose" : "no-transpose",
134
+ ];
135
+ [m, n] = [n, m];
136
+ uplo = uplo === "lower" ? "upper" : "lower";
137
+ }
138
+
139
+ // Shape-based auto-select — see sgemmtr_small.wgsl/sgemmtr_large.wgsl.
140
+ const largeWgX = Math.ceil(n / BN_LARGE);
141
+ const largeWgY = Math.ceil(m / BM_LARGE);
142
+ const useLargeTile = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
143
+
144
+ const pipeline = await getPipeline(device, useLargeTile ? "sgemmtr_large" : "sgemmtr_small");
145
+
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,
150
+ [
151
+ { value: m, type: "u32" },
152
+ { value: n, type: "u32" },
153
+ { value: k, type: "u32" },
154
+ { value: alpha, type: "f32" },
155
+ { value: beta, type: "f32" },
156
+ { value: lda, type: "u32" },
157
+ { value: ldb, type: "u32" },
158
+ { value: ldc, type: "u32" },
159
+ { value: transA === "transpose" ? 1 : 0, type: "u32" },
160
+ { value: transB === "transpose" ? 1 : 0, type: "u32" },
161
+ { value: uplo === "upper" ? 1 : 0, type: "u32" },
162
+ ],
163
+ "sgemmtr-params",
164
+ );
165
+
166
+ try {
167
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
168
+ ABuffer,
169
+ BBuffer,
170
+ CBuffer,
171
+ paramsBuffer,
172
+ ]);
173
+
174
+ const wgCount = useLargeTile
175
+ ? {
176
+ x: requireWorkgroupCount(device, largeWgX, "sgemmtr", "x"),
177
+ y: requireWorkgroupCount(device, largeWgY, "sgemmtr", "y"),
178
+ }
179
+ : {
180
+ x: requireWorkgroupCount(device, Math.ceil(n / BN_SMALL), "sgemmtr", "x"),
181
+ y: requireWorkgroupCount(device, Math.ceil(m / BM_SMALL), "sgemmtr", "y"),
182
+ };
183
+ const { commandEncoder, ts } = runComputePass(device, pipeline, bindGroup, wgCount);
184
+ const readBuffer = CIsGpu ? null : stageReadback(device, commandEncoder, CBuffer);
185
+
186
+ submit(device, commandEncoder);
187
+
188
+ const gpuTimeMs = await extractTimestamp(ts);
189
+
190
+ if (CIsGpu) {
191
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
192
+ return {};
193
+ }
194
+
195
+ const result = await extractResult(readBuffer, Float32Array);
196
+ if (gpuTimeMs !== undefined) return { C: result, gpuTimeMs };
197
+ return { C: result };
198
+ } finally {
199
+ if (!AIsGpu) destroyBuffers(ABuffer);
200
+ if (!BIsGpu) destroyBuffers(BBuffer);
201
+ if (!CIsGpu) destroyBuffers(CBuffer);
202
+ destroyBuffers(paramsBuffer);
203
+ }
204
+ }
@@ -50,43 +50,6 @@ export declare function sgemv(
50
50
  layout?: 'row-major' | 'column-major',
51
51
  ): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
52
52
 
53
- /**
54
- * Performs the matrix-vector operation y = alpha * op(A) * x + beta * y
55
- *
56
- * A is kept GPU-resident; x and y are CPU Float32Arrays. `A`'s own `layout`
57
- * (set at `GpuMatrix.from` time) determines the operation — there is no
58
- * separate `layout` argument here.
59
- *
60
- * @param device - GPUDevice from `init()`
61
- * @param trans - `'no-transpose'` for A, `'transpose'` for A^T
62
- * @param m - number of rows in A
63
- * @param n - number of columns in A
64
- * @param alpha - scalar multiplier for op(A)*x
65
- * @param A - GpuMatrix, GPU-resident
66
- * @param lda - leading dimension of A (must equal A.lda)
67
- * @param x - Float32Array input vector
68
- * @param incx - stride for x (must be a positive integer)
69
- * @param beta - scalar multiplier for y
70
- * @param y - Float32Array input/output vector
71
- * @param incy - stride for y (must be a positive integer)
72
- * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemv/sgemv.mjs#L16">Source code: sgemv.mjs (L16)</a>
73
- * @category BLAS Level 2
74
- */
75
- export declare function sgemv(
76
- device: GPUDevice,
77
- trans: 'no-transpose' | 'transpose',
78
- m: number,
79
- n: number,
80
- alpha: number,
81
- A: GpuMatrix,
82
- lda: number,
83
- x: Float32Array,
84
- incx: number,
85
- beta: number,
86
- y: Float32Array,
87
- incy: number,
88
- ): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
89
-
90
53
  /**
91
54
  * Performs the matrix-vector operation y = alpha * op(A) * x + beta * y
92
55
  *
@@ -94,7 +57,7 @@ export declare function sgemv(
94
57
  * `layout` (set at `GpuMatrix.from` time) determines the operation — there is
95
58
  * no separate `layout` argument here.
96
59
  *
97
- * {@includeCode ../../examples/sgemv/gpuvec.sgemv.js}
60
+ * {@includeCode ../../examples/sgemv/gpu.sgemv.js}
98
61
  *
99
62
  * @param device - GPUDevice from `init()`
100
63
  * @param trans - `'no-transpose'` for A, `'transpose'` for A^T
@@ -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")
@@ -56,8 +58,14 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
56
58
  throw new Error(
57
59
  "A must be a GpuMatrix when x and y are GpuVectors.",
58
60
  );
61
+ if (AIsGpu && !xIsGpu)
62
+ throw new Error(
63
+ "x and y must be GpuVectors when A is a GpuMatrix.",
64
+ );
59
65
  if (xIsGpu && x._buf === y._buf)
60
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.");
61
69
  if (AIsGpu && lda !== A.lda)
62
70
  throw new Error("lda must match A.lda when A is a GpuMatrix.");
63
71
  if (AIsGpu && (A.rows < m || A.cols < n))
@@ -95,38 +103,46 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
95
103
  const shaderName = isNoTrans ? "sgemv_n" : "sgemv_t";
96
104
  const pipeline = await getPipeline(device, shaderName);
97
105
 
98
- const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sgemv-A", false);
99
- const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sgemv-x", false);
100
- const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sgemv-y", true);
101
- const paramsBuffer = createParamsBuffer(
102
- [
103
- { value: m, type: "u32" },
104
- { value: n, type: "u32" },
105
- { value: alpha, type: "f32" },
106
- { value: beta, type: "f32" },
107
- { value: incx, type: "u32" },
108
- { value: incy, type: "u32" },
109
- { value: lda, type: "u32" },
110
- ],
111
- "sgemv-params",
112
- );
106
+ let ABuffer = null;
107
+ let xBuffer = null;
108
+ let yBuffer = null;
109
+ let paramsBuffer = null;
113
110
 
114
111
  try {
115
- 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), [
116
129
  ABuffer,
117
130
  xBuffer,
118
131
  yBuffer,
119
132
  paramsBuffer,
120
133
  ]);
121
134
 
122
- // 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.
123
139
  const wgCount = isNoTrans
124
140
  ? Math.min(m, device.limits.maxComputeWorkgroupsPerDimension)
125
- : calcWorkgroups(yLen);
126
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
127
- 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);
128
144
 
129
- submit(commandEncoder);
145
+ submit(device, commandEncoder);
130
146
 
131
147
  const gpuTimeMs = await extractTimestamp(ts);
132
148
 
@@ -139,10 +155,10 @@ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y
139
155
  if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
140
156
  return { y: result };
141
157
  } finally {
142
- if (!AIsGpu) destroyBuffers(ABuffer);
143
- if (!xIsGpu) destroyBuffers(xBuffer);
144
- if (!yIsGpu) destroyBuffers(yBuffer);
145
- 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);
146
162
 
147
163
  }
148
164
  }
@@ -42,39 +42,6 @@ export declare function sger(
42
42
  layout?: 'row-major' | 'column-major',
43
43
  ): Promise<{ A: Float32Array; gpuTimeMs?: number }>;
44
44
 
45
- /**
46
- * Performs the rank-1 update A = alpha * x * y^T + A
47
- *
48
- * A is kept GPU-resident; x and y are CPU Float32Arrays. `A`'s own `layout`
49
- * (set at `GpuMatrix.from` time) determines the operation — there is no
50
- * separate `layout` argument here.
51
- *
52
- * @param device - GPUDevice from `init()`
53
- * @param m - number of rows in A
54
- * @param n - number of columns in A
55
- * @param alpha - scalar multiplier for x*y^T
56
- * @param x - Float32Array input vector
57
- * @param incx - stride for x (must be a positive integer)
58
- * @param y - Float32Array input vector
59
- * @param incy - stride for y (must be a positive integer)
60
- * @param A - GpuMatrix, GPU-resident
61
- * @param lda - leading dimension of A (must equal A.lda)
62
- * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sger/sger.mjs#L15">Source code: sger.mjs (L15)</a>
63
- * @category BLAS Level 2
64
- */
65
- export declare function sger(
66
- device: GPUDevice,
67
- m: number,
68
- n: number,
69
- alpha: number,
70
- x: Float32Array,
71
- incx: number,
72
- y: Float32Array,
73
- incy: number,
74
- A: GpuMatrix,
75
- lda: number,
76
- ): Promise<{ gpuTimeMs?: number }>;
77
-
78
45
  /**
79
46
  * Performs the rank-1 update A = alpha * x * y^T + A
80
47
  *
@@ -82,7 +49,7 @@ export declare function sger(
82
49
  * `GpuMatrix.from` time) determines the operation — there is no separate
83
50
  * `layout` argument here.
84
51
  *
85
- * {@includeCode ../../examples/sger/gpuvec.sger.js}
52
+ * {@includeCode ../../examples/sger/gpu.sger.js}
86
53
  *
87
54
  * @param device - GPUDevice from `init()`
88
55
  * @param m - number of rows in A
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")
@@ -63,6 +65,8 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
63
65
  );
64
66
  if (xIsGpu && !AIsGpu)
65
67
  throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");
68
+ if (AIsGpu && !xIsGpu)
69
+ throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");
66
70
  if (AIsGpu && xIsGpu && A._buf === x._buf)
67
71
  throw new Error("A and x must not reference the same GPU buffer.");
68
72
  if (AIsGpu && yIsGpu && A._buf === y._buf)
@@ -87,10 +91,10 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
87
91
  let paramsBuffer = null;
88
92
 
89
93
  try {
90
- xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sger-x", false);
91
- yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sger-y", false);
92
- ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sger-A", true);
93
- 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,
94
98
  [
95
99
  { value: m, type: "u32" },
96
100
  { value: n, type: "u32" },
@@ -102,7 +106,7 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
102
106
  "sger-params",
103
107
  );
104
108
 
105
- const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
109
+ const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
106
110
  xBuffer,
107
111
  yBuffer,
108
112
  ABuffer,
@@ -112,10 +116,10 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
112
116
  // One workgroup per row of A; clamped to device limit — the shader's
113
117
  // grid-stride loop handles remaining rows when m > dispatch count.
114
118
  const wgCount = Math.min(m, device.limits.maxComputeWorkgroupsPerDimension);
115
- const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
116
- 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);
117
121
 
118
- submit(commandEncoder);
122
+ submit(device, commandEncoder);
119
123
 
120
124
  const gpuTimeMs = await extractTimestamp(ts);
121
125
 
@@ -0,0 +1,42 @@
1
+ // block_transfer: gather/scatter/scatter-subtract between a tight (blockLen
2
+ // x otherLen) block and a sub-range of a strided (any ld, row/col-major)
3
+ // buffer — needed since block offsets aren't 256-byte-aligned and block
4
+ // rows/cols aren't always one contiguous range for copyBufferToBuffer.
5
+
6
+ @group(0) @binding(0) var<storage, read_write> block: array<f32>; // blockLen x otherLen, block[i*otherLen+j]
7
+ @group(0) @binding(1) var<storage, read_write> strided: array<f32>; // B's or A's own buffer
8
+
9
+ struct Params {
10
+ blockStart: u32,
11
+ blockLen: u32,
12
+ otherStart: u32,
13
+ otherLen: u32,
14
+ ld: u32,
15
+ isColMajor: u32, // 0 = row-major addressing, 1 = column-major (row/col swapped)
16
+ blockIsRow: u32, // 1 = blockStart indexes strided's rows, 0 = its columns
17
+ mode: u32, // 0 = scatter (strided := block), 1 = scatter_sub (strided -= block), 2 = gather (block := strided)
18
+ }
19
+
20
+ @group(0) @binding(2) var<uniform> params: Params;
21
+
22
+ @compute @workgroup_size(8, 8)
23
+ fn main(@builtin(global_invocation_id) gid: vec3u) {
24
+ let i = gid.y; // index along the blocked axis, within the block
25
+ let j = gid.x; // index along the other axis, within the block
26
+ if (i >= params.blockLen || j >= params.otherLen) {
27
+ return;
28
+ }
29
+
30
+ let row = select(params.otherStart + j, params.blockStart + i, params.blockIsRow == 1u);
31
+ let col = select(params.blockStart + i, params.otherStart + j, params.blockIsRow == 1u);
32
+ let stridedIdx = select(row * params.ld + col, col * params.ld + row, params.isColMajor == 1u);
33
+ let blockIdx = i * params.otherLen + j;
34
+
35
+ if (params.mode == 2u) {
36
+ block[blockIdx] = strided[stridedIdx];
37
+ } else if (params.mode == 1u) {
38
+ strided[stridedIdx] -= block[blockIdx];
39
+ } else {
40
+ strided[stridedIdx] = block[blockIdx];
41
+ }
42
+ }
@@ -1,6 +1,7 @@
1
1
  // dasum: sum(|x[i]|), double-double (Dekker). Same ILP=4 shape as sasum.wgsl;
2
- // see f64/dekker.wgsl for ddAddProtected and why plain ddAdd isn't safe.
3
- // GpuVector input isn't pre-abs'd, so ddAbs() applies unconditionally below.
2
+ // see f64/utils/add.wgsl for ddAddProtected and why plain ddAdd isn't safe.
3
+ // GpuVector input isn't pre-abs'd, so ddAbs() (f64/utils/abs.wgsl) applies
4
+ // unconditionally below.
4
5
 
5
6
  @group(0) @binding(0) var<storage, read> xHi: array<f32>;
6
7
  @group(0) @binding(1) var<storage, read> xLo: array<f32>;
@@ -7,93 +7,12 @@
7
7
  //
8
8
  // No bindings, no entry point — a helper library, concatenated with a
9
9
  // consumer's own bindings/entry point by getPipeline (WGSL has no #include).
10
+ // The DD struct lives here — abs.wgsl/add.wgsl/greater.wgsl/equal.wgsl all
11
+ // use it but don't redefine it (WGSL errors on duplicate struct definitions
12
+ // once concatenated), so any consumer using those must concatenate this
13
+ // file too, first.
10
14
 
11
15
  struct DD {
12
16
  hi: f32,
13
17
  lo: f32,
14
18
  }
15
-
16
- // |a| for a double-double pair. Negation is exact (no rounding), so this is
17
- // just a sign flip on both components — hi alone determines the pair's sign.
18
- fn ddAbs(a: DD) -> DD {
19
- if (a.hi < 0.0) {
20
- return DD(-a.hi, -a.lo);
21
- }
22
- return a;
23
- }
24
-
25
- // ── A real compiler bug — read before touching anything below ──────────────
26
- //
27
- // twoSum/fastTwoSum's error term `e` should be nonzero (that's the point —
28
- // `s` is rounded). A front-end-optimizer bug zeros it anyway, on both NVIDIA
29
- // and Mesa (ANV + llvmpipe), via two different mechanisms needing two fixes:
30
- // bitcast-based subtraction (`fsub`/`negf`, fixes NVIDIA) and materializing
31
- // the sum through workgroup memory + workgroupBarrier() (fixes Mesa). Only
32
- // both together (ddAddProtected) is verified correct everywhere — the plain
33
- // twoSum/fastTwoSum/ddAdd below are reference-only, not safe to use.
34
- fn negf(x: f32) -> f32 {
35
- return bitcast<f32>(bitcast<u32>(x) ^ 0x80000000u);
36
- }
37
- fn fsub(a: f32, b: f32) -> f32 {
38
- return a + negf(b);
39
- }
40
-
41
- // Knuth/Møller's TwoSum: s = fl(a+b), e = exact rounding error, a+b == s+e.
42
- // Works for any a, b. UNPROTECTED — see header above.
43
- fn twoSum(a: f32, b: f32) -> DD {
44
- let s = a + b;
45
- let v = s - a;
46
- let e = (a - (s - v)) + (b - v);
47
- return DD(s, e);
48
- }
49
-
50
- // Dekker's Fast-Two-Sum: same contract, but only correct when |a| >= |b|.
51
- // UNPROTECTED — see header above.
52
- fn fastTwoSum(a: f32, b: f32) -> DD {
53
- let s = a + b;
54
- let e = b - (s - a);
55
- return DD(s, e);
56
- }
57
-
58
- // Double-double addition (Dekker's Add2). UNPROTECTED — see header above.
59
- fn ddAdd(a: DD, b: DD) -> DD {
60
- let s = twoSum(a.hi, b.hi);
61
- let loSum = a.lo + b.lo;
62
- return fastTwoSum(s.hi, s.lo + loSum);
63
- }
64
-
65
- // ── Protected variants — use these ──────────────────────────────────────────
66
- //
67
- // Bitcast subtraction + workgroup-barrier materialization, verified correct
68
- // on all three backends tested. Costs a real barrier: fine for O(1)-per-
69
- // thread or O(log n) reduction use, not a long per-element loop. A
70
- // workgroupBarrier() requires uniform control flow, so:
71
- // - `threadSlot` must be unique per concurrent caller (e.g. local_invocation_index).
72
- // - Every thread in the workgroup must call this the same number of times
73
- // — including ones whose result gets discarded. Compute unconditionally;
74
- // only the write-back should be conditional.
75
- var<workgroup> dekkerScratch: array<f32, 64>;
76
-
77
- fn twoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
78
- dekkerScratch[threadSlot] = a + b;
79
- workgroupBarrier();
80
- let s = dekkerScratch[threadSlot];
81
- let v = fsub(s, a);
82
- let e = fsub(a, fsub(s, v)) + fsub(b, v);
83
- return DD(s, e);
84
- }
85
-
86
- fn fastTwoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
87
- dekkerScratch[threadSlot] = a + b;
88
- workgroupBarrier();
89
- let s = dekkerScratch[threadSlot];
90
- let e = fsub(b, fsub(s, a));
91
- return DD(s, e);
92
- }
93
-
94
- // Protected double-double addition — same contract as ddAdd, but exact.
95
- fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
96
- let s = twoSumProtected(a.hi, b.hi, threadSlot);
97
- let loSum = a.lo + b.lo;
98
- return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
99
- }
@@ -0,0 +1,10 @@
1
+ // Requires f64/dekker.wgsl concatenated first for the DD struct.
2
+
3
+ // |a| for a double-double pair. Negation is exact (no rounding), so this is
4
+ // just a sign flip on both components — hi alone determines the pair's sign.
5
+ fn ddAbs(a: DD) -> DD {
6
+ if (a.hi < 0.0) {
7
+ return DD(-a.hi, -a.lo);
8
+ }
9
+ return a;
10
+ }