wgblas 0.1.2 → 1.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 (64) hide show
  1. package/README.md +3 -0
  2. package/dist/wgblas.browser.js +1249 -37
  3. package/index.d.mts +9 -0
  4. package/index.mjs +9 -0
  5. package/package.json +47 -1
  6. package/src/classes/GpuMatrix.d.mts +98 -0
  7. package/src/classes/GpuMatrix.mjs +109 -0
  8. package/src/classes/GpuVector.d.mts +11 -7
  9. package/src/classes/GpuVector.mjs +31 -5
  10. package/src/dasum/dasum.d.mts +48 -0
  11. package/src/dasum/dasum.mjs +121 -0
  12. package/src/init.mjs +3 -1
  13. package/src/isamax/isamax.mjs +66 -67
  14. package/src/random/random.d.mts +37 -0
  15. package/src/random/random.mjs +17 -0
  16. package/src/sasum/sasum.mjs +55 -51
  17. package/src/saxpy/saxpy.mjs +48 -35
  18. package/src/scopy/scopy.mjs +43 -32
  19. package/src/sdot/sdot.mjs +60 -55
  20. package/src/sgemv/sgemv.d.mts +127 -0
  21. package/src/sgemv/sgemv.mjs +148 -0
  22. package/src/sger/sger.d.mts +111 -0
  23. package/src/sger/sger.mjs +136 -0
  24. package/src/shaders/browser-shaders.mjs +26 -0
  25. package/src/shaders/dasum.wgsl +98 -0
  26. package/src/shaders/f64add.wgsl +281 -0
  27. package/src/shaders/isamax.wgsl +32 -9
  28. package/src/shaders/reduction/sumF64.wgsl +49 -0
  29. package/src/shaders/sasum.wgsl +18 -4
  30. package/src/shaders/sdot.wgsl +18 -4
  31. package/src/shaders/sgemv_n.wgsl +75 -0
  32. package/src/shaders/sgemv_t.wgsl +65 -0
  33. package/src/shaders/sger.wgsl +48 -0
  34. package/src/shaders/snrm2.wgsl +22 -4
  35. package/src/shaders/ssymv.wgsl +69 -0
  36. package/src/shaders/ssyr.wgsl +60 -0
  37. package/src/shaders/ssyr2.wgsl +63 -0
  38. package/src/shaders/strmv.wgsl +103 -0
  39. package/src/shaders/strsv_apply_inverse.wgsl +47 -0
  40. package/src/shaders/strsv_invert_block.wgsl +109 -0
  41. package/src/shaders/strsv_update.wgsl +75 -0
  42. package/src/snrm2/snrm2.mjs +56 -52
  43. package/src/srot/srot.mjs +57 -41
  44. package/src/srotm/srotm.mjs +54 -38
  45. package/src/sscal/sscal.mjs +43 -32
  46. package/src/sswap/sswap.mjs +49 -34
  47. package/src/ssymv/ssymv.d.mts +117 -0
  48. package/src/ssymv/ssymv.mjs +135 -0
  49. package/src/ssyr/ssyr.d.mts +100 -0
  50. package/src/ssyr/ssyr.mjs +106 -0
  51. package/src/ssyr2/ssyr2.d.mts +112 -0
  52. package/src/ssyr2/ssyr2.mjs +130 -0
  53. package/src/strmv/strmv.d.mts +117 -0
  54. package/src/strmv/strmv.mjs +138 -0
  55. package/src/strsv/strsv.d.mts +106 -0
  56. package/src/strsv/strsv.mjs +207 -0
  57. package/src/util/benchmark.mjs +1 -1
  58. package/src/util/bindgroup.mjs +14 -10
  59. package/src/util/buffer.mjs +7 -2
  60. package/src/util/compute.mjs +41 -15
  61. package/src/util/f64pack.mjs +152 -0
  62. package/src/util/pipeline.mjs +32 -17
  63. package/src/util/result.mjs +8 -4
  64. package/src/util/workgroup.mjs +10 -10
@@ -0,0 +1,127 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
3
+
4
+ /**
5
+ * Performs the matrix-vector operation y = alpha * op(A) * x + beta * y
6
+ *
7
+ * - `trans='no-transpose'`: op(A) = A, x is length n, y is length m
8
+ * - `trans='transpose'`: op(A) = A^T, x is length m, y is length n
9
+ *
10
+ * A is an m×n matrix stored in row-major order. `lda` is the leading dimension
11
+ * (number of floats between the start of consecutive rows — must be >= n).
12
+ *
13
+ * {@includeCode ../../examples/sgemv/sgemv.js}
14
+ *
15
+ * **Browser (standalone HTML):**
16
+ * {@includeCode ../../examples/sgemv/web/sgemv.html}
17
+ *
18
+ * @param device - GPUDevice from `init()`
19
+ * @param trans - `'no-transpose'` for A, `'transpose'` for A^T
20
+ * @param m - number of rows in A
21
+ * @param n - number of columns in A
22
+ * @param alpha - scalar multiplier for op(A)*x
23
+ * @param A - Float32Array, row-major or column-major (see `layout`), at least
24
+ * (m-1)*lda+n elements for row-major or (n-1)*lda+m elements for column-major
25
+ * @param lda - leading dimension of A (>= n for row-major, >= m for column-major)
26
+ * @param x - Float32Array input vector
27
+ * @param incx - stride for x (must be a positive integer)
28
+ * @param beta - scalar multiplier for y
29
+ * @param y - Float32Array input/output vector
30
+ * @param incy - stride for y (must be a positive integer)
31
+ * @param layout - storage layout of `A` (default: `'row-major'`); column-major
32
+ * swaps the effective `m`/`n` and flips `trans` internally (op(A) stays
33
+ * what you asked for either way — x/y keep their original lengths)
34
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemv/sgemv.mjs#L16">Source code: sgemv.mjs (L16)</a>
35
+ * @category BLAS Level 2
36
+ */
37
+ export declare function sgemv(
38
+ device: GPUDevice,
39
+ trans: 'no-transpose' | 'transpose',
40
+ m: number,
41
+ n: number,
42
+ alpha: number,
43
+ A: Float32Array,
44
+ lda: number,
45
+ x: Float32Array,
46
+ incx: number,
47
+ beta: number,
48
+ y: Float32Array,
49
+ incy: number,
50
+ layout?: 'row-major' | 'column-major',
51
+ ): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
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
+ /**
91
+ * Performs the matrix-vector operation y = alpha * op(A) * x + beta * y
92
+ *
93
+ * x and y are kept resident on the GPU. A must be a GpuMatrix; its own
94
+ * `layout` (set at `GpuMatrix.from` time) determines the operation — there is
95
+ * no separate `layout` argument here.
96
+ *
97
+ * {@includeCode ../../examples/sgemv/gpuvec.sgemv.js}
98
+ *
99
+ * @param device - GPUDevice from `init()`
100
+ * @param trans - `'no-transpose'` for A, `'transpose'` for A^T
101
+ * @param m - number of rows in A
102
+ * @param n - number of columns in A
103
+ * @param alpha - scalar multiplier for op(A)*x
104
+ * @param A - GpuMatrix
105
+ * @param lda - leading dimension of A (must equal A.lda)
106
+ * @param x - GpuVector input vector (not mutated)
107
+ * @param incx - stride for x (must be a positive integer)
108
+ * @param beta - scalar multiplier for y
109
+ * @param y - GpuVector input/output vector (mutated in place)
110
+ * @param incy - stride for y (must be a positive integer)
111
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemv/sgemv.mjs#L16">Source code: sgemv.mjs (L16)</a>
112
+ * @category BLAS Level 2
113
+ */
114
+ export declare function sgemv(
115
+ device: GPUDevice,
116
+ trans: 'no-transpose' | 'transpose',
117
+ m: number,
118
+ n: number,
119
+ alpha: number,
120
+ A: GpuMatrix,
121
+ lda: number,
122
+ x: GpuVector,
123
+ incx: number,
124
+ beta: number,
125
+ y: GpuVector,
126
+ incy: number,
127
+ ): Promise<{ gpuTimeMs?: number }>;
@@ -0,0 +1,148 @@
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 { calcWorkgroups } from "../util/workgroup.mjs";
13
+ import { GpuVector } from "../classes/GpuVector.mjs";
14
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
15
+
16
+ export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y, incy, layout = "row-major") {
17
+ const AIsGpu = A instanceof GpuMatrix;
18
+ const xIsGpu = x instanceof GpuVector;
19
+ const yIsGpu = y instanceof GpuVector;
20
+
21
+ if (!(device instanceof GPUDevice))
22
+ throw new Error("device must be a GPUDevice.");
23
+ if (trans !== "no-transpose" && trans !== "transpose")
24
+ throw new Error("trans must be 'no-transpose' or 'transpose'.");
25
+ if (layout !== "row-major" && layout !== "column-major")
26
+ throw new Error("layout must be 'row-major' or 'column-major'.");
27
+ if (typeof alpha !== "number")
28
+ throw new Error("alpha must be a number.");
29
+ if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
30
+ if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
31
+ if (typeof beta !== "number")
32
+ throw new Error("beta must be a number.");
33
+ if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
34
+ if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
35
+ if (
36
+ !Number.isInteger(m) ||
37
+ !Number.isInteger(n) ||
38
+ !Number.isInteger(incx) ||
39
+ !Number.isInteger(incy) ||
40
+ !Number.isInteger(lda)
41
+ )
42
+ throw new Error("m, n, incx, incy, and lda must be integers.");
43
+ if (incx <= 0 || incy <= 0)
44
+ throw new Error("incx and incy must be positive.");
45
+ if (!AIsGpu && !(A instanceof Float32Array))
46
+ throw new Error("A must be a Float32Array or GpuMatrix.");
47
+ if (!xIsGpu && !(x instanceof Float32Array))
48
+ throw new Error("x must be a Float32Array or GpuVector.");
49
+ if (!yIsGpu && !(y instanceof Float32Array))
50
+ throw new Error("y must be a Float32Array or GpuVector.");
51
+ if (xIsGpu !== yIsGpu)
52
+ throw new Error(
53
+ "x and y must be the same type (both Float32Array or both GpuVector).",
54
+ );
55
+ if (xIsGpu && !AIsGpu)
56
+ throw new Error(
57
+ "A must be a GpuMatrix when x and y are GpuVectors.",
58
+ );
59
+ if (xIsGpu && x._buf === y._buf)
60
+ throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");
61
+ if (AIsGpu && lda !== A.lda)
62
+ throw new Error("lda must match A.lda when A is a GpuMatrix.");
63
+ if (AIsGpu && (A.rows < m || A.cols < n))
64
+ throw new Error("A is too small for the given m and n.");
65
+ if (m < 0 || n < 0) throw new Error("m and n must be non-negative.");
66
+ if (m === 0 || n === 0) return yIsGpu ? {} : { y };
67
+
68
+ // Column-major A reinterpreted row-major is A^T — swap to the kernel's
69
+ // view of m/n/trans now that the shape checks above are done.
70
+ const effLayout = AIsGpu ? A.layout : layout;
71
+ if (effLayout === "column-major") {
72
+ [m, n] = [n, m];
73
+ trans = trans === "no-transpose" ? "transpose" : "no-transpose";
74
+ }
75
+ const isNoTrans = trans === "no-transpose";
76
+
77
+ // NoTrans: x has n elements, y has m elements; Trans: x has m elements, y has n elements.
78
+ const xLen = isNoTrans ? n : m;
79
+ const yLen = isNoTrans ? m : n;
80
+
81
+ if (lda < n) throw new Error("lda must be >= n.");
82
+ if (!AIsGpu && A.length < (m - 1) * lda + n)
83
+ throw new Error(
84
+ "A does not have enough elements for the given m, n, and lda.",
85
+ );
86
+ if (x.length < (xLen - 1) * incx + 1)
87
+ throw new Error(
88
+ "x does not have enough elements for the given dimensions and incx.",
89
+ );
90
+ if (y.length < (yLen - 1) * incy + 1)
91
+ throw new Error(
92
+ "y does not have enough elements for the given dimensions and incy.",
93
+ );
94
+
95
+ const shaderName = isNoTrans ? "sgemv_n" : "sgemv_t";
96
+ const pipeline = await getPipeline(device, shaderName);
97
+
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
+ );
113
+
114
+ try {
115
+ const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
116
+ ABuffer,
117
+ xBuffer,
118
+ yBuffer,
119
+ paramsBuffer,
120
+ ]);
121
+
122
+ // NoTrans: one workgroup per row (grid-stride handles overflow); Trans: one thread per output column.
123
+ const wgCount = isNoTrans
124
+ ? 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);
128
+
129
+ submit(commandEncoder);
130
+
131
+ const gpuTimeMs = await extractTimestamp(ts);
132
+
133
+ if (yIsGpu) {
134
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
135
+ return {};
136
+ }
137
+
138
+ const result = await extractResult(readBuffer, Float32Array);
139
+ if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
140
+ return { y: result };
141
+ } finally {
142
+ if (!AIsGpu) destroyBuffers(ABuffer);
143
+ if (!xIsGpu) destroyBuffers(xBuffer);
144
+ if (!yIsGpu) destroyBuffers(yBuffer);
145
+ destroyBuffers(paramsBuffer);
146
+
147
+ }
148
+ }
@@ -0,0 +1,111 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
3
+
4
+ /**
5
+ * Performs the rank-1 update A = alpha * x * y^T + A
6
+ *
7
+ * A is an m×n matrix stored in row-major order, updated in place. `lda` is
8
+ * the leading dimension (number of floats between the start of consecutive
9
+ * rows — must be >= n).
10
+ *
11
+ * {@includeCode ../../examples/sger/sger.js}
12
+ *
13
+ * **Browser (standalone HTML):**
14
+ * {@includeCode ../../examples/sger/web/sger.html}
15
+ *
16
+ * @param device - GPUDevice from `init()`
17
+ * @param m - number of rows in A (length of x)
18
+ * @param n - number of columns in A (length of y)
19
+ * @param alpha - scalar multiplier for x*y^T
20
+ * @param x - Float32Array input vector, length at least (m-1)*incx+1
21
+ * @param incx - stride for x (must be a positive integer)
22
+ * @param y - Float32Array input vector, length at least (n-1)*incy+1
23
+ * @param incy - stride for y (must be a positive integer)
24
+ * @param A - Float32Array, row-major or column-major (see `layout`), at least
25
+ * (m-1)*lda+n elements for row-major or (n-1)*lda+m elements for column-major
26
+ * @param lda - leading dimension of A (>= n for row-major, >= m for column-major)
27
+ * @param layout - storage layout of `A` (default: `'row-major'`)
28
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sger/sger.mjs#L15">Source code: sger.mjs (L15)</a>
29
+ * @category BLAS Level 2
30
+ */
31
+ export declare function sger(
32
+ device: GPUDevice,
33
+ m: number,
34
+ n: number,
35
+ alpha: number,
36
+ x: Float32Array,
37
+ incx: number,
38
+ y: Float32Array,
39
+ incy: number,
40
+ A: Float32Array,
41
+ lda: number,
42
+ layout?: 'row-major' | 'column-major',
43
+ ): Promise<{ A: Float32Array; gpuTimeMs?: number }>;
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
+ /**
79
+ * Performs the rank-1 update A = alpha * x * y^T + A
80
+ *
81
+ * x, y, and A are all kept resident on the GPU. `A`'s own `layout` (set at
82
+ * `GpuMatrix.from` time) determines the operation — there is no separate
83
+ * `layout` argument here.
84
+ *
85
+ * {@includeCode ../../examples/sger/gpuvec.sger.js}
86
+ *
87
+ * @param device - GPUDevice from `init()`
88
+ * @param m - number of rows in A
89
+ * @param n - number of columns in A
90
+ * @param alpha - scalar multiplier for x*y^T
91
+ * @param x - GpuVector input vector (not mutated)
92
+ * @param incx - stride for x (must be a positive integer)
93
+ * @param y - GpuVector input vector (not mutated)
94
+ * @param incy - stride for y (must be a positive integer)
95
+ * @param A - GpuMatrix, mutated in place
96
+ * @param lda - leading dimension of A (must equal A.lda)
97
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sger/sger.mjs#L15">Source code: sger.mjs (L15)</a>
98
+ * @category BLAS Level 2
99
+ */
100
+ export declare function sger(
101
+ device: GPUDevice,
102
+ m: number,
103
+ n: number,
104
+ alpha: number,
105
+ x: GpuVector,
106
+ incx: number,
107
+ y: GpuVector,
108
+ incy: number,
109
+ A: GpuMatrix,
110
+ lda: number,
111
+ ): Promise<{ gpuTimeMs?: number }>;
@@ -0,0 +1,136 @@
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 { GpuVector } from "../classes/GpuVector.mjs";
13
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
14
+
15
+ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout = "row-major") {
16
+ const AIsGpu = A instanceof GpuMatrix;
17
+
18
+ if (!(device instanceof GPUDevice))
19
+ throw new Error("device must be a GPUDevice.");
20
+ if (layout !== "row-major" && layout !== "column-major")
21
+ throw new Error("layout must be 'row-major' or 'column-major'.");
22
+ if (typeof alpha !== "number")
23
+ throw new Error("alpha must be a number.");
24
+ if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
25
+ if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
26
+ if (
27
+ !Number.isInteger(m) ||
28
+ !Number.isInteger(n) ||
29
+ !Number.isInteger(incx) ||
30
+ !Number.isInteger(incy) ||
31
+ !Number.isInteger(lda)
32
+ )
33
+ throw new Error("m, n, incx, incy, and lda must be integers.");
34
+ if (incx <= 0 || incy <= 0)
35
+ throw new Error("incx and incy must be positive.");
36
+ if (!AIsGpu && !(A instanceof Float32Array))
37
+ throw new Error("A must be a Float32Array or GpuMatrix.");
38
+ if (AIsGpu && lda !== A.lda)
39
+ throw new Error("lda must match A.lda when A is a GpuMatrix.");
40
+ // A.rows/A.cols are fixed regardless of layout — check against original m/n before the swap below.
41
+ if (AIsGpu && (A.rows < m || A.cols < n))
42
+ throw new Error("A is too small for the given m and n.");
43
+
44
+ // GpuMatrix's own layout wins over the argument; column-major A reinterpreted row-major is A^T, so swap x/y and m/n to reproduce it (transpose of a rank-1 update swaps its two vectors).
45
+ const effLayout = AIsGpu ? A.layout : layout;
46
+ if (effLayout === "column-major") {
47
+ [m, n] = [n, m];
48
+ [x, y] = [y, x];
49
+ [incx, incy] = [incy, incx];
50
+ }
51
+
52
+ const xIsGpu = x instanceof GpuVector;
53
+ const yIsGpu = y instanceof GpuVector;
54
+
55
+ if (lda < n) throw new Error("lda must be >= n.");
56
+ if (!xIsGpu && !(x instanceof Float32Array))
57
+ throw new Error("x must be a Float32Array or GpuVector.");
58
+ if (!yIsGpu && !(y instanceof Float32Array))
59
+ throw new Error("y must be a Float32Array or GpuVector.");
60
+ if (xIsGpu !== yIsGpu)
61
+ throw new Error(
62
+ "x and y must be the same type (both Float32Array or both GpuVector).",
63
+ );
64
+ if (xIsGpu && !AIsGpu)
65
+ throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");
66
+ if (AIsGpu && xIsGpu && A._buf === x._buf)
67
+ throw new Error("A and x must not reference the same GPU buffer.");
68
+ if (AIsGpu && yIsGpu && A._buf === y._buf)
69
+ throw new Error("A and y must not reference the same GPU buffer.");
70
+ if (m < 0 || n < 0) throw new Error("m and n must be non-negative.");
71
+ if (m === 0 || n === 0) return AIsGpu ? {} : { A };
72
+
73
+ if (!AIsGpu && A.length < (m - 1) * lda + n)
74
+ throw new Error(
75
+ "A does not have enough elements for the given m, n, and lda.",
76
+ );
77
+ if (x.length < (m - 1) * incx + 1)
78
+ throw new Error("x does not have enough elements for the given m and incx.");
79
+ if (y.length < (n - 1) * incy + 1)
80
+ throw new Error("y does not have enough elements for the given n and incy.");
81
+
82
+ const pipeline = await getPipeline(device, "sger");
83
+
84
+ let xBuffer = null;
85
+ let yBuffer = null;
86
+ let ABuffer = null;
87
+ let paramsBuffer = null;
88
+
89
+ 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
+ [
95
+ { value: m, type: "u32" },
96
+ { value: n, type: "u32" },
97
+ { value: alpha, type: "f32" },
98
+ { value: incx, type: "u32" },
99
+ { value: incy, type: "u32" },
100
+ { value: lda, type: "u32" },
101
+ ],
102
+ "sger-params",
103
+ );
104
+
105
+ const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
106
+ xBuffer,
107
+ yBuffer,
108
+ ABuffer,
109
+ paramsBuffer,
110
+ ]);
111
+
112
+ // One workgroup per row of A; clamped to device limit — the shader's
113
+ // grid-stride loop handles remaining rows when m > dispatch count.
114
+ const wgCount = Math.min(m, device.limits.maxComputeWorkgroupsPerDimension);
115
+ const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
116
+ const readBuffer = AIsGpu ? null : stageReadback(commandEncoder, ABuffer);
117
+
118
+ submit(commandEncoder);
119
+
120
+ const gpuTimeMs = await extractTimestamp(ts);
121
+
122
+ if (AIsGpu) {
123
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
124
+ return {};
125
+ }
126
+
127
+ const result = await extractResult(readBuffer, Float32Array);
128
+ if (gpuTimeMs !== undefined) return { A: result, gpuTimeMs };
129
+ return { A: result };
130
+ } finally {
131
+ if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
132
+ if (!yIsGpu && yBuffer) destroyBuffers(yBuffer);
133
+ if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
134
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
135
+ }
136
+ }
@@ -1,5 +1,6 @@
1
1
  import argmax from "./reduction/argmax.wgsl";
2
2
  import sum from "./reduction/sum.wgsl";
3
+ import sumF64 from "./reduction/sumF64.wgsl";
3
4
  import sscal from "./sscal.wgsl";
4
5
  import sswap from "./sswap.wgsl";
5
6
  import saxpy from "./saxpy.wgsl";
@@ -10,10 +11,23 @@ import snrm2 from "./snrm2.wgsl";
10
11
  import srot from "./srot.wgsl";
11
12
  import srotm from "./srotm.wgsl";
12
13
  import isamax from "./isamax.wgsl";
14
+ import sgemv_n from "./sgemv_n.wgsl";
15
+ import sgemv_t from "./sgemv_t.wgsl";
16
+ import ssymv from "./ssymv.wgsl";
17
+ import strmv from "./strmv.wgsl";
18
+ import sger from "./sger.wgsl";
19
+ import ssyr from "./ssyr.wgsl";
20
+ import ssyr2 from "./ssyr2.wgsl";
21
+ import f64add from "./f64add.wgsl";
22
+ import dasum from "./dasum.wgsl";
23
+ import strsv_invert_block from "./strsv_invert_block.wgsl";
24
+ import strsv_apply_inverse from "./strsv_apply_inverse.wgsl";
25
+ import strsv_update from "./strsv_update.wgsl";
13
26
 
14
27
  export const shaderSources = {
15
28
  "reduction/argmax": argmax,
16
29
  "reduction/sum": sum,
30
+ "reduction/sumF64": sumF64,
17
31
  sscal,
18
32
  sswap,
19
33
  saxpy,
@@ -24,4 +38,16 @@ export const shaderSources = {
24
38
  srot,
25
39
  srotm,
26
40
  isamax,
41
+ sgemv_n,
42
+ sgemv_t,
43
+ ssymv,
44
+ strmv,
45
+ sger,
46
+ ssyr,
47
+ ssyr2,
48
+ f64add,
49
+ dasum,
50
+ strsv_invert_block,
51
+ strsv_apply_inverse,
52
+ strsv_update,
27
53
  };
@@ -0,0 +1,98 @@
1
+ // dasum: result = sum(|x[i]|)
2
+ // pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sumF64.wgsl.
3
+ // Same structure as sasum.wgsl — every value is now a [main, aux] pair
4
+ // (see src/util/f64pack.mjs) and every `+`/`+=` is computeSum via addPair
5
+ // instead of plain f32 addition. Concatenated after f64add.wgsl by
6
+ // getPipeline (WGSL has no #include), reusing its decode/encode/computeSum/
7
+ // addFields and Packed struct — f64add.wgsl declares no bindings and no entry
8
+ // point of its own (just helper functions), so bindings here start at 0 and
9
+ // the entry point is simply `dasum_main`.
10
+ //
11
+ // xAux/partialsAux are array<u32>, not array<f32> — aux's bits must never
12
+ // pass through an f32-typed storage slot (NaN-bit-pattern corruption risk,
13
+ // see f64pack.mjs and the Packed struct comment above decode()/encode() in
14
+ // f64add.wgsl); Packed (from f64add.wgsl) keeps aux as u32 in registers/
15
+ // workgroup memory too.
16
+ //
17
+ // Per-thread accumulation (acc0..acc3) stays in DECODED Fields form for the
18
+ // entire strided loop below, via addFields — not re-encoded to Packed and
19
+ // re-decoded on every single element like a naive version would. Only the
20
+ // freshly-loaded x[idx] needs decoding each iteration (unavoidable, it's new
21
+ // data every time); the running total never leaves Fields form until the
22
+ // four accumulators are combined and encoded exactly once, right before
23
+ // writing into workgroup-shared `tile`. The cross-thread reduction tree
24
+ // after that still goes through Packed per level (unavoidable — each level
25
+ // combines values that live in different threads' registers via shared
26
+ // memory), but that's a fixed 6 levels regardless of n, unlike the strided
27
+ // loop above whose iteration count scales with n.
28
+
29
+ @group(0) @binding(0) var<storage, read> xMain: array<f32>;
30
+ @group(0) @binding(1) var<storage, read> xAux: array<u32>;
31
+ @group(0) @binding(2) var<storage, read_write> partialsMain: array<f32>;
32
+ @group(0) @binding(3) var<storage, read_write> partialsAux: array<u32>;
33
+ @group(0) @binding(4) var<uniform> params: Params;
34
+
35
+ struct Params {
36
+ n: u32,
37
+ x_inc: u32,
38
+ }
39
+
40
+ const WGS: u32 = 64;
41
+
42
+ var<workgroup> tile: array<Packed, 64>;
43
+
44
+ // a + b, where a/b are [main, aux] pairs — computeSum takes decoded Fields.
45
+ // Only used for the cross-thread reduction tree below; the per-thread
46
+ // strided loop uses addFields directly instead (see module comment).
47
+ fn addPair(a: Packed, b: Packed) -> Packed {
48
+ return computeSum(decode(bitcast<u32>(a.main), a.aux), decode(bitcast<u32>(b.main), b.aux));
49
+ }
50
+
51
+ // |x| for a packed double is abs(main) with aux untouched — only main's
52
+ // sign bit carries the double's sign (see fieldsToPacked() in f64pack.mjs).
53
+ // Returns decoded Fields directly (not Packed) for the per-thread loop.
54
+ fn absFields(idx: u32) -> Fields {
55
+ return decode(bitcast<u32>(abs(xMain[idx])), xAux[idx]);
56
+ }
57
+
58
+ @compute @workgroup_size(64)
59
+ fn dasum_main(
60
+ @builtin(global_invocation_id) gid: vec3u,
61
+ @builtin(local_invocation_id) lid: vec3u,
62
+ @builtin(workgroup_id) wgid: vec3u,
63
+ @builtin(num_workgroups) num_wg: vec3u,
64
+ ) {
65
+ var acc0: Fields = Fields(0u, 0u, 0u, 0u);
66
+ var acc1: Fields = Fields(0u, 0u, 0u, 0u);
67
+ var acc2: Fields = Fields(0u, 0u, 0u, 0u);
68
+ var acc3: Fields = Fields(0u, 0u, 0u, 0u);
69
+
70
+ let stride = num_wg.x * WGS;
71
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
72
+
73
+ for (var id = gid.x; id < n4_floor; id += 4u * stride) {
74
+ acc0 = addFields(acc0, absFields( id * params.x_inc));
75
+ acc1 = addFields(acc1, absFields((id + stride) * params.x_inc));
76
+ acc2 = addFields(acc2, absFields((id + 2u * stride) * params.x_inc));
77
+ acc3 = addFields(acc3, absFields((id + 3u * stride) * params.x_inc));
78
+ }
79
+ for (var id = n4_floor + gid.x; id < params.n; id += stride) {
80
+ acc0 = addFields(acc0, absFields(id * params.x_inc));
81
+ }
82
+
83
+ // Combine the 4 per-thread accumulators in Fields form too — still no
84
+ // encode/decode needed, since none of them have touched Packed yet.
85
+ let combined = addFields(addFields(acc0, acc1), addFields(acc2, acc3));
86
+ tile[lid.x] = encode(combined.sign, combined.rawExp, combined.mantissaHi, combined.lo);
87
+ workgroupBarrier();
88
+
89
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
90
+ if (lid.x < s) { tile[lid.x] = addPair(tile[lid.x], tile[lid.x + s]); }
91
+ workgroupBarrier();
92
+ }
93
+
94
+ if (lid.x == 0u) {
95
+ partialsMain[wgid.x] = tile[0].main;
96
+ partialsAux[wgid.x] = tile[0].aux;
97
+ }
98
+ }