wgblas 1.2.0 → 2.0.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 (70) hide show
  1. package/LICENSE +1 -1
  2. package/README.md +36 -54
  3. package/dist/wgblas.browser.js +779 -34
  4. package/index.d.mts +13 -0
  5. package/index.mjs +8 -0
  6. package/package.json +56 -3
  7. package/src/dasum/dasum.d.mts +2 -2
  8. package/src/dasum/dasum.mjs +6 -4
  9. package/src/idamax/idamax.d.mts +51 -0
  10. package/src/idamax/idamax.mjs +128 -0
  11. package/src/init.mjs +9 -1
  12. package/src/isamax/isamax.d.mts +1 -1
  13. package/src/sasum/sasum.d.mts +1 -1
  14. package/src/saxpy/saxpy.d.mts +1 -1
  15. package/src/scopy/scopy.d.mts +1 -1
  16. package/src/sdot/sdot.d.mts +1 -1
  17. package/src/sgemm/sgemm.d.mts +102 -0
  18. package/src/sgemm/sgemm.mjs +195 -0
  19. package/src/sgemmtr/sgemmtr.d.mts +104 -0
  20. package/src/sgemmtr/sgemmtr.mjs +203 -0
  21. package/src/sgemv/sgemv.d.mts +1 -38
  22. package/src/sgemv/sgemv.mjs +4 -0
  23. package/src/sger/sger.d.mts +1 -34
  24. package/src/sger/sger.mjs +2 -0
  25. package/src/shaders/block_transfer.wgsl +42 -0
  26. package/src/shaders/browser-shaders.mjs +26 -0
  27. package/src/shaders/dasum.wgsl +3 -2
  28. package/src/shaders/f64/dekker.wgsl +4 -85
  29. package/src/shaders/f64/utils/abs.wgsl +10 -0
  30. package/src/shaders/f64/utils/add.wgsl +77 -0
  31. package/src/shaders/f64/utils/equal.wgsl +7 -0
  32. package/src/shaders/f64/utils/greater.wgsl +12 -0
  33. package/src/shaders/f64/utils/multiply.wgsl +81 -0
  34. package/src/shaders/idamax.wgsl +96 -0
  35. package/src/shaders/reduction/argmaxF64.wgsl +50 -0
  36. package/src/shaders/reduction/sumF64.wgsl +2 -2
  37. package/src/shaders/sgemm_large.wgsl +117 -0
  38. package/src/shaders/sgemm_small.wgsl +112 -0
  39. package/src/shaders/sgemmtr_large.wgsl +117 -0
  40. package/src/shaders/sgemmtr_small.wgsl +110 -0
  41. package/src/shaders/symmetrize.wgsl +31 -0
  42. package/src/shaders/triangularize.wgsl +44 -0
  43. package/src/snrm2/snrm2.d.mts +1 -1
  44. package/src/srot/srot.d.mts +1 -1
  45. package/src/srotm/srotm.d.mts +1 -1
  46. package/src/sscal/sscal.d.mts +1 -1
  47. package/src/sswap/sswap.d.mts +1 -1
  48. package/src/ssymm/ssymm.d.mts +103 -0
  49. package/src/ssymm/ssymm.mjs +209 -0
  50. package/src/ssymv/ssymv.d.mts +1 -36
  51. package/src/ssymv/ssymv.mjs +2 -0
  52. package/src/ssyr/ssyr.d.mts +1 -30
  53. package/src/ssyr/ssyr.mjs +2 -0
  54. package/src/ssyr2/ssyr2.d.mts +1 -34
  55. package/src/ssyr2/ssyr2.mjs +2 -0
  56. package/src/ssyr2k/ssyr2k.d.mts +100 -0
  57. package/src/ssyr2k/ssyr2k.mjs +201 -0
  58. package/src/ssyrk/ssyrk.d.mts +90 -0
  59. package/src/ssyrk/ssyrk.mjs +176 -0
  60. package/src/strmm/strmm.d.mts +100 -0
  61. package/src/strmm/strmm.mjs +211 -0
  62. package/src/strmv/strmv.d.mts +1 -36
  63. package/src/strmv/strmv.mjs +2 -0
  64. package/src/strsm/strsm.d.mts +99 -0
  65. package/src/strsm/strsm.mjs +342 -0
  66. package/src/strsv/strsv.d.mts +1 -32
  67. package/src/strsv/strsv.mjs +2 -0
  68. package/src/util/buffer.mjs +4 -2
  69. package/src/util/compute.mjs +6 -3
  70. package/src/util/f64.mjs +3 -3
package/index.d.mts CHANGED
@@ -17,6 +17,7 @@ export { sasum } from "./src/sasum/sasum.mjs";
17
17
  export { dasum } from "./src/dasum/dasum.mjs";
18
18
  export { snrm2 } from "./src/snrm2/snrm2.mjs";
19
19
  export { isamax } from "./src/isamax/isamax.mjs";
20
+ export { idamax } from "./src/idamax/idamax.mjs";
20
21
  export { srot } from "./src/srot/srot.mjs";
21
22
  export { srotm } from "./src/srotm/srotm.mjs";
22
23
  export { sgemv } from "./src/sgemv/sgemv.mjs";
@@ -26,6 +27,13 @@ export { strsv } from "./src/strsv/strsv.mjs";
26
27
  export { sger } from "./src/sger/sger.mjs";
27
28
  export { ssyr } from "./src/ssyr/ssyr.mjs";
28
29
  export { ssyr2 } from "./src/ssyr2/ssyr2.mjs";
30
+ export { sgemm } from "./src/sgemm/sgemm.mjs";
31
+ export { sgemmtr } from "./src/sgemmtr/sgemmtr.mjs";
32
+ export { ssyrk } from "./src/ssyrk/ssyrk.mjs";
33
+ export { ssyr2k } from "./src/ssyr2k/ssyr2k.mjs";
34
+ export { ssymm } from "./src/ssymm/ssymm.mjs";
35
+ export { strmm } from "./src/strmm/strmm.mjs";
36
+ export { strsm } from "./src/strsm/strsm.mjs";
29
37
 
30
38
  /**
31
39
  * Initializes the WebGPU device.
@@ -35,6 +43,10 @@ export { ssyr2 } from "./src/ssyr2/ssyr2.mjs";
35
43
  * and `"low-power"` favors the integrated one.
36
44
  * See [MDN: GPU.requestAdapter()](https://developer.mozilla.org/en-US/docs/Web/API/GPU/requestAdapter).
37
45
  * @param options.benchmark - enable GPU timestamp queries; BLAS functions return `{ result, gpuTimeMs }` (default: `false`)
46
+ * @param options.dumpShaders - Node-only. Forwards Dawn's `dump_shaders` debug toggle, printing
47
+ * each pipeline's WGSL and compiled backend IR (SPIR-V/Vulkan, MSL/Metal, or HLSL/D3D12,
48
+ * whichever Dawn picked) to stderr as it compiles. A Dawn passthrough, not a wgblas format —
49
+ * no effect in the browser, which gives pages no API to request compiled shader IR (default: `false`)
38
50
  *
39
51
  * @example Default (high-performance GPU)
40
52
  * ```js
@@ -69,6 +81,7 @@ export { ssyr2 } from "./src/ssyr2/ssyr2.mjs";
69
81
  export declare function init(options?: {
70
82
  powerPreference?: GPUPowerPreference;
71
83
  benchmark?: boolean;
84
+ dumpShaders?: boolean;
72
85
  }): Promise<GPUDevice>;
73
86
 
74
87
  /**
package/index.mjs CHANGED
@@ -15,6 +15,7 @@ export { sasum } from "./src/sasum/sasum.mjs";
15
15
  export { dasum } from "./src/dasum/dasum.mjs";
16
16
  export { snrm2 } from "./src/snrm2/snrm2.mjs";
17
17
  export { isamax } from "./src/isamax/isamax.mjs";
18
+ export { idamax } from "./src/idamax/idamax.mjs";
18
19
  export { srot } from "./src/srot/srot.mjs";
19
20
  export { srotm } from "./src/srotm/srotm.mjs";
20
21
  export { sgemv } from "./src/sgemv/sgemv.mjs";
@@ -24,3 +25,10 @@ export { strsv } from "./src/strsv/strsv.mjs";
24
25
  export { sger } from "./src/sger/sger.mjs";
25
26
  export { ssyr } from "./src/ssyr/ssyr.mjs";
26
27
  export { ssyr2 } from "./src/ssyr2/ssyr2.mjs";
28
+ export { sgemm } from "./src/sgemm/sgemm.mjs";
29
+ export { sgemmtr } from "./src/sgemmtr/sgemmtr.mjs";
30
+ export { ssyrk } from "./src/ssyrk/ssyrk.mjs";
31
+ export { ssyr2k } from "./src/ssyr2k/ssyr2k.mjs";
32
+ export { ssymm } from "./src/ssymm/ssymm.mjs";
33
+ export { strmm } from "./src/strmm/strmm.mjs";
34
+ export { strsm } from "./src/strsm/strsm.mjs";
package/package.json CHANGED
@@ -1,6 +1,6 @@
1
1
  {
2
2
  "name": "wgblas",
3
- "version": "1.2.0",
3
+ "version": "2.0.0",
4
4
  "description": "BLAS on WebGPU",
5
5
  "type": "module",
6
6
  "main": "index.mjs",
@@ -59,6 +59,10 @@
59
59
  "import": "./src/isamax/isamax.mjs",
60
60
  "types": "./src/isamax/isamax.d.mts"
61
61
  },
62
+ "./idamax": {
63
+ "import": "./src/idamax/idamax.mjs",
64
+ "types": "./src/idamax/idamax.d.mts"
65
+ },
62
66
  "./srot": {
63
67
  "import": "./src/srot/srot.mjs",
64
68
  "types": "./src/srot/srot.d.mts"
@@ -94,6 +98,34 @@
94
98
  "./ssyr2": {
95
99
  "import": "./src/ssyr2/ssyr2.mjs",
96
100
  "types": "./src/ssyr2/ssyr2.d.mts"
101
+ },
102
+ "./sgemm": {
103
+ "import": "./src/sgemm/sgemm.mjs",
104
+ "types": "./src/sgemm/sgemm.d.mts"
105
+ },
106
+ "./sgemmtr": {
107
+ "import": "./src/sgemmtr/sgemmtr.mjs",
108
+ "types": "./src/sgemmtr/sgemmtr.d.mts"
109
+ },
110
+ "./ssyrk": {
111
+ "import": "./src/ssyrk/ssyrk.mjs",
112
+ "types": "./src/ssyrk/ssyrk.d.mts"
113
+ },
114
+ "./ssyr2k": {
115
+ "import": "./src/ssyr2k/ssyr2k.mjs",
116
+ "types": "./src/ssyr2k/ssyr2k.d.mts"
117
+ },
118
+ "./ssymm": {
119
+ "import": "./src/ssymm/ssymm.mjs",
120
+ "types": "./src/ssymm/ssymm.d.mts"
121
+ },
122
+ "./strmm": {
123
+ "import": "./src/strmm/strmm.mjs",
124
+ "types": "./src/strmm/strmm.d.mts"
125
+ },
126
+ "./strsm": {
127
+ "import": "./src/strsm/strsm.mjs",
128
+ "types": "./src/strsm/strsm.d.mts"
97
129
  }
98
130
  },
99
131
  "files": [
@@ -110,7 +142,26 @@
110
142
  },
111
143
  "keywords": [
112
144
  "BLAS",
113
- "WebGPU"
145
+ "WebGPU",
146
+ "linear-algebra",
147
+ "linear",
148
+ "algebra",
149
+ "subroutines",
150
+ "level 1",
151
+ "level 2",
152
+ "level 3",
153
+ "gpu",
154
+ "gpgpu",
155
+ "wgsl",
156
+ "compute-shader",
157
+ "matrix",
158
+ "vector",
159
+ "matmul",
160
+ "dot product",
161
+ "float32",
162
+ "float64",
163
+ "numerical-computing",
164
+ "scientific-computing"
114
165
  ],
115
166
  "author": "Manit Roy",
116
167
  "license": "Apache-2.0",
@@ -120,7 +171,7 @@
120
171
  "bugs": {
121
172
  "url": "https://github.com/manit2004/wgblas/issues"
122
173
  },
123
- "homepage": "https://github.com/manit2004/wgblas#readme",
174
+ "homepage": "https://manit2004.github.io/wgblas/",
124
175
  "dependencies": {
125
176
  "webgpu": "^0.4.0"
126
177
  },
@@ -129,11 +180,13 @@
129
180
  "@commitlint/config-conventional": "^21.2.0",
130
181
  "@eslint/js": "^10.0.1",
131
182
  "@stdlib/blas-base-dasum": "^0.4.1",
183
+ "@stdlib/blas-base-idamax": "^0.1.1",
132
184
  "@stdlib/blas-base-isamax": "^0.1.1",
133
185
  "@stdlib/blas-base-sasum": "^0.3.1",
134
186
  "@stdlib/blas-base-saxpy": "^0.3.1",
135
187
  "@stdlib/blas-base-scopy": "^0.3.1",
136
188
  "@stdlib/blas-base-sdot": "^0.3.1",
189
+ "@stdlib/blas-base-sgemm": "^0.1.1",
137
190
  "@stdlib/blas-base-sgemv": "^0.1.1",
138
191
  "@stdlib/blas-base-sger": "^0.1.1",
139
192
  "@stdlib/blas-base-snrm2": "^0.3.1",
@@ -5,7 +5,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
5
5
  * precision: result = sum(|x[i]|). Each element of `x` has abs() applied,
6
6
  * then is split into a (hi, lo) double-double f32 pair (see
7
7
  * `splitDoubleDouble`/`f64.mjs`) since WGSL has no f64 type; accumulation
8
- * uses Dekker's double-double algorithm (see `shaders/f64/dekker.wgsl`), giving ~48 bits
8
+ * uses Dekker's double-double algorithm (see `shaders/f64/`), giving ~48 bits
9
9
  * of mantissa — more than a single f32 (24 bits) but less than true f64
10
10
  * (52 bits), so results are not bit-exact with a CPU double.
11
11
  *
@@ -33,7 +33,7 @@ export declare function dasum(
33
33
  * Computes the sum of absolute values of a vector of doubles in double
34
34
  * precision: result = sum(|x[i]|).
35
35
  *
36
- * {@includeCode ../../examples/dasum/gpuvec.dasum.js}
36
+ * {@includeCode ../../examples/dasum/gpu.dasum.js}
37
37
  *
38
38
  * @param device - GPUDevice from `init()`
39
39
  * @param n - number of elements (must be a positive integer)
@@ -34,10 +34,12 @@ export async function dasum(device, n, x, incx) {
34
34
  "x does not have enough elements for the given n and incx.",
35
35
  );
36
36
 
37
- // Concatenated with f64/dekker.wgsl for its DD struct/ddAdd helpers (WGSL
38
- // has no #include); entryPoint omitted since each module has only one @compute.
39
- const pipelineMain = await getPipeline(device, ["f64/dekker", "dasum"]);
40
- const pipelineReduce = await getPipeline(device, ["f64/dekker", "reduction/sumF64"]);
37
+ // Concatenated with f64/dekker.wgsl (DD struct), f64/utils/abs.wgsl
38
+ // (ddAbs), and f64/utils/add.wgsl (ddAddProtected) — WGSL has no
39
+ // #include; entryPoint omitted since each module has only one @compute.
40
+ const f64Deps = ["f64/dekker", "f64/utils/abs", "f64/utils/add"];
41
+ const pipelineMain = await getPipeline(device, [...f64Deps, "dasum"]);
42
+ const pipelineReduce = await getPipeline(device, [...f64Deps, "reduction/sumF64"]);
41
43
 
42
44
  let xHiBuffer = null;
43
45
  let xLoBuffer = null;
@@ -0,0 +1,51 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+
3
+ /**
4
+ * Returns the 0-based index of the element with the largest absolute value,
5
+ * for a vector of doubles. Each element of `x` is split into a (hi, lo)
6
+ * double-double f32 pair (see `splitDoubleDouble`/`f64.mjs`) since WGSL has
7
+ * no f64 type; comparisons use the double-double pair directly (hi, falling
8
+ * back to lo on an exact tie), giving ~48 bits of discriminating precision —
9
+ * more than a single f32 (24 bits) but less than true f64 (52 bits). Ties
10
+ * are broken in favour of the lower index, matching CBLAS behaviour.
11
+ *
12
+ * {@includeCode ../../examples/idamax/idamax.js}
13
+ *
14
+ * **Browser (standalone HTML):**
15
+ * {@includeCode ../../examples/idamax/web/idamax.html}
16
+ *
17
+ * @param device - GPUDevice from `init()`
18
+ * @param n - number of elements (must be a positive integer)
19
+ * @param x - Float64Array input vector
20
+ * @param incx - stride for x (must be a positive integer)
21
+ * @returns 0-based index of max |x[i]|
22
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/idamax/idamax.mjs#L19">Source code: idamax.mjs (L19)</a>
23
+ * @category BLAS Level 1
24
+ */
25
+ export declare function idamax(
26
+ device: GPUDevice,
27
+ n: number,
28
+ x: Float64Array,
29
+ incx: number,
30
+ ): Promise<{ index: number } | { index: number; gpuTimeMs: number }>;
31
+
32
+ /**
33
+ * Returns the 0-based index of the element with the largest absolute value.
34
+ * Ties are broken in favour of the lower index, matching CBLAS behaviour.
35
+ *
36
+ * {@includeCode ../../examples/idamax/gpu.idamax.js}
37
+ *
38
+ * @param device - GPUDevice from `init()`
39
+ * @param n - number of elements (must be a positive integer)
40
+ * @param x - Float64Array-backed GpuVector input vector
41
+ * @param incx - stride for x (must be a positive integer)
42
+ * @returns 0-based index of max |x[i]|
43
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/idamax/idamax.mjs#L19">Source code: idamax.mjs (L19)</a>
44
+ * @category BLAS Level 1
45
+ */
46
+ export declare function idamax(
47
+ device: GPUDevice,
48
+ n: number,
49
+ x: GpuVector,
50
+ incx: number,
51
+ ): Promise<{ index: number } | { index: number; gpuTimeMs: number }>;
@@ -0,0 +1,128 @@
1
+ import {
2
+ uploadBuffer,
3
+ createStorageBuffer,
4
+ createParamsBuffer,
5
+ createResultBuffer,
6
+ stageReadback,
7
+ destroyBuffers,
8
+ } from "../util/buffer.mjs";
9
+ import { createBindGroup } from "../util/bindgroup.mjs";
10
+ import { runComputePass, submit } from "../util/compute.mjs";
11
+ import { extractTimestamp } from "../util/benchmark.mjs";
12
+ import { extractResult } from "../util/result.mjs";
13
+ import { getPipeline } from "../util/pipeline.mjs";
14
+ import { GpuVector } from "../classes/GpuVector.mjs";
15
+ import { splitDoubleDouble } from "../util/f64.mjs";
16
+
17
+ const WGS = 64;
18
+
19
+ export async function idamax(device, n, x, incx) {
20
+ const xIsGpu = x instanceof GpuVector;
21
+
22
+ if (!(device instanceof GPUDevice))
23
+ throw new Error("device must be a GPUDevice.");
24
+ if (!Number.isInteger(n) || !Number.isInteger(incx))
25
+ throw new Error("n and incx must be integers.");
26
+ if (incx <= 0) throw new Error("incx must be positive.");
27
+ if (!xIsGpu && !(x instanceof Float64Array))
28
+ throw new Error("x must be a Float64Array or GpuVector.");
29
+ if (xIsGpu && x.dtype !== Float64Array)
30
+ throw new Error("x must be a Float64Array-backed GpuVector.");
31
+ if (n <= 0) return { index: 0 };
32
+ if (x.length < (n - 1) * incx + 1)
33
+ throw new Error(
34
+ "x does not have enough elements for the given n and incx.",
35
+ );
36
+
37
+ // Concatenated f64 helpers (WGSL has no #include); ddAbs is unconditional in idamax.wgsl, so x is split as-is.
38
+ const f64Deps = ["f64/dekker", "f64/utils/abs", "f64/utils/greater", "f64/utils/equal"];
39
+ const pipelineMain = await getPipeline(device, [...f64Deps, "idamax"], "idamax_main");
40
+ const pipelineReduce = await getPipeline(device, [...f64Deps, "reduction/argmaxF64"], "reduce_f64");
41
+
42
+ let xHiBuffer = null;
43
+ let xLoBuffer = null;
44
+ let partialsValHiBuffer = null;
45
+ let partialsValLoBuffer = null;
46
+ let partialsIdxBuffer = null;
47
+ let resultBuffer = null;
48
+ let paramsBuffer = null;
49
+ let readBuffer = null;
50
+
51
+ try {
52
+ if (xIsGpu) {
53
+ xHiBuffer = x._buf;
54
+ xLoBuffer = x._loBuf;
55
+ } else {
56
+ const { hi, lo } = splitDoubleDouble(x);
57
+ xHiBuffer = uploadBuffer(hi, "idamax-xHi", false);
58
+ xLoBuffer = uploadBuffer(lo, "idamax-xLo", false);
59
+ }
60
+ partialsValHiBuffer = createStorageBuffer(2 * WGS * 4, "idamax-partials-val-hi");
61
+ partialsValLoBuffer = createStorageBuffer(2 * WGS * 4, "idamax-partials-val-lo");
62
+ partialsIdxBuffer = createStorageBuffer(2 * WGS * 4, "idamax-partials-idx");
63
+ resultBuffer = createResultBuffer(4, "idamax-result"); // u32 index
64
+ paramsBuffer = createParamsBuffer(
65
+ [
66
+ { value: n, type: "u32" },
67
+ { value: incx, type: "u32" },
68
+ ],
69
+ "idamax-params",
70
+ );
71
+
72
+ const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
73
+ xHiBuffer,
74
+ xLoBuffer,
75
+ partialsValHiBuffer,
76
+ partialsValLoBuffer,
77
+ partialsIdxBuffer,
78
+ paramsBuffer,
79
+ ]);
80
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
81
+ pipelineMain,
82
+ bgMain,
83
+ 2 * WGS,
84
+ ); // dispatch 2*WGS workgroups
85
+
86
+ submit(enc1);
87
+
88
+ const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
89
+ partialsValHiBuffer,
90
+ partialsValLoBuffer,
91
+ partialsIdxBuffer,
92
+ resultBuffer,
93
+ ]);
94
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
95
+ pipelineReduce,
96
+ bgReduce,
97
+ 1,
98
+ ); // dispatch 1 workgroup to reduce the partials to a single index
99
+ readBuffer = stageReadback(enc2, resultBuffer);
100
+
101
+ submit(enc2);
102
+
103
+ const resultPromise = extractResult(readBuffer, Uint32Array);
104
+ readBuffer = null; // ownership transferred — extractResult's own finally destroys it
105
+
106
+ const [gpuTime1, gpuTime2, idxArr] = await Promise.all([
107
+ extractTimestamp(ts1),
108
+ extractTimestamp(ts2),
109
+ resultPromise,
110
+ ]);
111
+
112
+ const index = idxArr[0];
113
+
114
+ if (gpuTime1 !== undefined && gpuTime2 !== undefined)
115
+ return { index, gpuTimeMs: gpuTime1 + gpuTime2 };
116
+ return { index };
117
+ } finally {
118
+ if (!xIsGpu && xHiBuffer) destroyBuffers(xHiBuffer);
119
+ if (!xIsGpu && xLoBuffer) destroyBuffers(xLoBuffer);
120
+ if (partialsValHiBuffer) destroyBuffers(partialsValHiBuffer);
121
+ if (partialsValLoBuffer) destroyBuffers(partialsValLoBuffer);
122
+ if (partialsIdxBuffer) destroyBuffers(partialsIdxBuffer);
123
+ if (resultBuffer) destroyBuffers(resultBuffer);
124
+ if (paramsBuffer) destroyBuffers(paramsBuffer);
125
+ // Only reached if submit(enc2) threw before ownership was transferred above.
126
+ if (readBuffer) destroyBuffers(readBuffer);
127
+ }
128
+ }
package/src/init.mjs CHANGED
@@ -11,6 +11,7 @@ let _benchmarkEnabled = false;
11
11
  export async function init({
12
12
  powerPreference = "high-performance",
13
13
  benchmark = false,
14
+ dumpShaders = false,
14
15
  } = {}) {
15
16
  if (_device) {
16
17
  return _device;
@@ -23,9 +24,16 @@ export async function init({
23
24
  if (typeof window === "undefined") {
24
25
  const { create, globals } = await import("webgpu");
25
26
  Object.assign(globalThis, globals);
26
- gpu = create([]);
27
+ // dumpShaders forwards Dawn's own debug toggle — prints each pipeline's
28
+ // WGSL and compiled backend IR to stderr. Node-only; see index.d.mts.
29
+ const toggles = dumpShaders
30
+ ? ["enable-dawn-features=dump_shaders,disable_symbol_renaming"]
31
+ : [];
32
+ gpu = create(toggles);
27
33
  _gpu = gpu;
28
34
  } else {
35
+ if (dumpShaders)
36
+ console.warn("dumpShaders has no effect in the browser — see init()'s docs.");
29
37
  gpu = navigator.gpu;
30
38
  }
31
39
 
@@ -28,7 +28,7 @@ export declare function isamax(
28
28
  * Returns the 0-based index of the element with the largest absolute value.
29
29
  * Ties are broken in favour of the lower index, matching CBLAS behaviour.
30
30
  *
31
- * {@includeCode ../../examples/isamax/gpuvec.isamax.js}
31
+ * {@includeCode ../../examples/isamax/gpu.isamax.js}
32
32
  *
33
33
  * @param device - GPUDevice from `init()`
34
34
  * @param n - number of elements (must be a positive integer)
@@ -26,7 +26,7 @@ export declare function sasum(
26
26
  /**
27
27
  * Computes the sum of absolute values of a vector: result = sum(|x[i]|)
28
28
  *
29
- * {@includeCode ../../examples/sasum/gpuvec.sasum.js}
29
+ * {@includeCode ../../examples/sasum/gpu.sasum.js}
30
30
  *
31
31
  * @param device - GPUDevice from `init()`
32
32
  * @param n - number of elements (must be a positive integer)
@@ -31,7 +31,7 @@ export declare function saxpy(
31
31
  /**
32
32
  * Performs the operation y = alpha * x + y
33
33
  *
34
- * {@includeCode ../../examples/saxpy/gpuvec.saxpy.js}
34
+ * {@includeCode ../../examples/saxpy/gpu.saxpy.js}
35
35
  *
36
36
  * @param device - GPUDevice from `init()`
37
37
  * @param n - number of elements (must be a positive integer)
@@ -29,7 +29,7 @@ export declare function scopy(
29
29
  /**
30
30
  * Performs the operation y = x
31
31
  *
32
- * {@includeCode ../../examples/scopy/gpuvec.scopy.js}
32
+ * {@includeCode ../../examples/scopy/gpu.scopy.js}
33
33
  *
34
34
  * @param device - GPUDevice from `init()`
35
35
  * @param n - number of elements (must be a positive integer)
@@ -30,7 +30,7 @@ export declare function sdot(
30
30
  /**
31
31
  * Computes the dot product of two vectors: result = sum(x[i] * y[i])
32
32
  *
33
- * {@includeCode ../../examples/sdot/gpuvec.sdot.js}
33
+ * {@includeCode ../../examples/sdot/gpu.sdot.js}
34
34
  *
35
35
  * @param device - GPUDevice from `init()`
36
36
  * @param n - number of elements (must be a positive integer)
@@ -0,0 +1,102 @@
1
+ import { GpuMatrix } from "../classes/GpuMatrix.mjs";
2
+
3
+ /**
4
+ * Performs the matrix-matrix operation C = alpha * op(A) * op(B) + beta * C
5
+ *
6
+ * - `transA/transB='no-transpose'`: op(A) = A (m×k), op(B) = B (k×n)
7
+ * - `transA/transB='transpose'`: op(A) = A^T, op(B) = B^T
8
+ *
9
+ * A, B, C are row-major or column-major (see `layout`) — backed by one of
10
+ * two shared-memory-tiled, register-blocked kernels chosen by shape:
11
+ * `sgemm_small.wgsl` (BM=BN=32) below a 6x6 workgroup grid, `sgemm_large.wgsl`
12
+ * (BM=BN=64) above it.
13
+ *
14
+ * {@includeCode ../../examples/sgemm/sgemm.js}
15
+ *
16
+ * **Browser (standalone HTML):**
17
+ * {@includeCode ../../examples/sgemm/web/sgemm.html}
18
+ *
19
+ * @param device - GPUDevice from `init()`
20
+ * @param transA - `'no-transpose'` for A, `'transpose'` for A^T
21
+ * @param transB - `'no-transpose'` for B, `'transpose'` for B^T
22
+ * @param m - rows of op(A) and C
23
+ * @param n - columns of op(B) and C
24
+ * @param k - columns of op(A), rows of op(B)
25
+ * @param alpha - scalar multiplier for op(A)*op(B)
26
+ * @param A - Float32Array, row-major or column-major (see `layout`)
27
+ * @param lda - leading dimension of A as stored
28
+ * @param B - Float32Array, row-major or column-major (see `layout`)
29
+ * @param ldb - leading dimension of B as stored
30
+ * @param beta - scalar multiplier for C
31
+ * @param C - Float32Array input/output matrix, row-major or column-major
32
+ * @param ldc - leading dimension of C as stored
33
+ * @param layout - storage layout shared by A/B/C when they're Float32Array
34
+ * (default: `'row-major'`); column-major A/B flips the respective trans
35
+ * flag internally, column-major C computes C^T = op(B)^T*op(A)^T instead
36
+ * (same underlying bytes) — op(A)*op(B) stays what you asked for either way
37
+ * @returns updated C as a Float32Array
38
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemm/sgemm.mjs#L17">Source code: sgemm.mjs (L17)</a>
39
+ * @category BLAS Level 3
40
+ */
41
+ export declare function sgemm(
42
+ device: GPUDevice,
43
+ transA: 'no-transpose' | 'transpose',
44
+ transB: 'no-transpose' | 'transpose',
45
+ m: number,
46
+ n: number,
47
+ k: number,
48
+ alpha: number,
49
+ A: Float32Array,
50
+ lda: number,
51
+ B: Float32Array,
52
+ ldb: number,
53
+ beta: number,
54
+ C: Float32Array,
55
+ ldc: number,
56
+ layout?: 'row-major' | 'column-major',
57
+ ): Promise<{ C: Float32Array; gpuTimeMs?: number }>;
58
+
59
+ /**
60
+ * Performs the matrix-matrix operation C = alpha * op(A) * op(B) + beta * C
61
+ *
62
+ * A, B, and C are all kept GPU-resident. Each matrix's own `layout` (set at
63
+ * `GpuMatrix.from` time) determines the operation — there is no separate
64
+ * `layout` argument here. A and B must be GpuMatrix whenever C is, and vice
65
+ * versa — mixing a GpuMatrix with a plain Float32Array is not supported.
66
+ *
67
+ * {@includeCode ../../examples/sgemm/gpu.sgemm.js}
68
+ *
69
+ * @param device - GPUDevice from `init()`
70
+ * @param transA - `'no-transpose'` for A, `'transpose'` for A^T
71
+ * @param transB - `'no-transpose'` for B, `'transpose'` for B^T
72
+ * @param m - rows of op(A) and C
73
+ * @param n - columns of op(B) and C
74
+ * @param k - columns of op(A), rows of op(B)
75
+ * @param alpha - scalar multiplier for op(A)*op(B)
76
+ * @param A - GpuMatrix
77
+ * @param lda - leading dimension of A (must equal A.lda)
78
+ * @param B - GpuMatrix
79
+ * @param ldb - leading dimension of B (must equal B.lda)
80
+ * @param beta - scalar multiplier for C
81
+ * @param C - GpuMatrix (mutated in place)
82
+ * @param ldc - leading dimension of C (must equal C.lda)
83
+ * @returns no C — it stays GPU-resident; call `C.read()` yourself for a CPU readback (see the example)
84
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemm/sgemm.mjs#L17">Source code: sgemm.mjs (L17)</a>
85
+ * @category BLAS Level 3
86
+ */
87
+ export declare function sgemm(
88
+ device: GPUDevice,
89
+ transA: 'no-transpose' | 'transpose',
90
+ transB: 'no-transpose' | 'transpose',
91
+ m: number,
92
+ n: number,
93
+ k: number,
94
+ alpha: number,
95
+ A: GpuMatrix,
96
+ lda: number,
97
+ B: GpuMatrix,
98
+ ldb: number,
99
+ beta: number,
100
+ C: GpuMatrix,
101
+ ldc: number,
102
+ ): Promise<{ gpuTimeMs?: number }>;