wgblas 0.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 (43) hide show
  1. package/LICENSE +201 -0
  2. package/README.md +161 -0
  3. package/dist/wgblas.browser.js +422 -0
  4. package/index.d.mts +94 -0
  5. package/index.mjs +16 -0
  6. package/package.json +127 -0
  7. package/src/classes/GpuVector.d.mts +69 -0
  8. package/src/classes/GpuVector.mjs +33 -0
  9. package/src/devdocs.mjs +81 -0
  10. package/src/index.mjs +4 -0
  11. package/src/init.mjs +105 -0
  12. package/src/isamax/isamax.d.mts +46 -0
  13. package/src/isamax/isamax.mjs +115 -0
  14. package/src/random/random.d.mts +61 -0
  15. package/src/random/random.mjs +11 -0
  16. package/src/sasum/sasum.d.mts +44 -0
  17. package/src/sasum/sasum.mjs +98 -0
  18. package/src/saxpy/saxpy.d.mts +54 -0
  19. package/src/saxpy/saxpy.mjs +90 -0
  20. package/src/scopy/scopy.d.mts +50 -0
  21. package/src/scopy/scopy.mjs +87 -0
  22. package/src/sdot/sdot.d.mts +52 -0
  23. package/src/sdot/sdot.mjs +118 -0
  24. package/src/shaders/browser-shaders.mjs +27 -0
  25. package/src/shaders/index.mjs +27 -0
  26. package/src/snrm2/snrm2.d.mts +44 -0
  27. package/src/snrm2/snrm2.mjs +100 -0
  28. package/src/srot/srot.d.mts +62 -0
  29. package/src/srot/srot.mjs +95 -0
  30. package/src/srotm/srotm.d.mts +60 -0
  31. package/src/srotm/srotm.mjs +94 -0
  32. package/src/sscal/sscal.d.mts +46 -0
  33. package/src/sscal/sscal.mjs +71 -0
  34. package/src/sswap/sswap.d.mts +50 -0
  35. package/src/sswap/sswap.mjs +90 -0
  36. package/src/util/benchmark.mjs +103 -0
  37. package/src/util/bindgroup.mjs +22 -0
  38. package/src/util/buffer.mjs +160 -0
  39. package/src/util/compute.mjs +52 -0
  40. package/src/util/index.mjs +12 -0
  41. package/src/util/pipeline.mjs +82 -0
  42. package/src/util/result.mjs +19 -0
  43. package/src/util/workgroup.mjs +31 -0
@@ -0,0 +1,44 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+
3
+ /**
4
+ * Computes the sum of absolute values of a vector: result = sum(|x[i]|)
5
+ *
6
+ * {@includeCode ../../examples/sasum/sasum.js}
7
+ *
8
+ * **Browser (standalone HTML):**
9
+ * {@includeCode ../../examples/sasum/web/sasum.html}
10
+ *
11
+ * @param device - GPUDevice from `init()`
12
+ * @param n - number of elements (must be a positive integer)
13
+ * @param x - Float32Array input vector
14
+ * @param incx - stride for x (must be a positive integer)
15
+ * @returns absolute sum scalar — always a CPU readback, even for GpuVector inputs
16
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sasum/sasum.mjs#L18">Source code: sasum.mjs (L18)</a>
17
+ * @category BLAS Level 1
18
+ */
19
+ export declare function sasum(
20
+ device: GPUDevice,
21
+ n: number,
22
+ x: Float32Array,
23
+ incx: number,
24
+ ): Promise<{ asum: number } | { asum: number; gpuTimeMs: number }>;
25
+
26
+ /**
27
+ * Computes the sum of absolute values of a vector: result = sum(|x[i]|)
28
+ *
29
+ * {@includeCode ../../examples/sasum/gpuvec.sasum.js}
30
+ *
31
+ * @param device - GPUDevice from `init()`
32
+ * @param n - number of elements (must be a positive integer)
33
+ * @param x - GpuVector input vector
34
+ * @param incx - stride for x (must be a positive integer)
35
+ * @returns absolute sum scalar — always a CPU readback, even for GpuVector inputs
36
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sasum/sasum.mjs#L18">Source code: sasum.mjs (L18)</a>
37
+ * @category BLAS Level 1
38
+ */
39
+ export declare function sasum(
40
+ device: GPUDevice,
41
+ n: number,
42
+ x: GpuVector,
43
+ incx: number,
44
+ ): Promise<{ asum: number } | { asum: number; gpuTimeMs: number }>;
@@ -0,0 +1,98 @@
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
+
16
+ const WGS = 64; // workgroup size
17
+
18
+ export async function sasum(device, n, x, incx) {
19
+ const xIsGpu = x instanceof GpuVector;
20
+
21
+ if (!(device instanceof GPUDevice))
22
+ throw new Error("device must be a GPUDevice.");
23
+ if (!Number.isInteger(n) || !Number.isInteger(incx))
24
+ throw new Error("n and incx must be integers.");
25
+ if (incx <= 0) throw new Error("incx must be positive.");
26
+ if (!xIsGpu && !(x instanceof Float32Array))
27
+ throw new Error("x must be a Float32Array or GpuVector.");
28
+ if (n <= 0) return { asum: 0 };
29
+ if (x.length < (n - 1) * incx + 1)
30
+ throw new Error(
31
+ "x does not have enough elements for the given n and incx.",
32
+ );
33
+
34
+ const pipelineMain = await getPipeline(device, "sasum");
35
+ const pipelineReduce = await getPipeline(device, "reduction/sum");
36
+
37
+ const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sasum-x", false);
38
+ const partialsBuffer = createStorageBuffer(2 * WGS * 4, "sasum-partials"); // 2*WGS partial sums of f32
39
+ const resultBuffer = createResultBuffer(4, "sasum-result"); // final f32 scalar
40
+ const paramsBuffer = createParamsBuffer(
41
+ [
42
+ { value: n, type: "u32" },
43
+ { value: incx, type: "u32" },
44
+ ],
45
+ "sasum-params",
46
+ );
47
+
48
+ const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
49
+ xBuffer,
50
+ partialsBuffer,
51
+ paramsBuffer,
52
+ ]);
53
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
54
+ pipelineMain,
55
+ bgMain,
56
+ 2 * WGS,
57
+ ); // dispatch 2*WGS workgroups
58
+
59
+ submit(enc1);
60
+
61
+ const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
62
+ partialsBuffer,
63
+ resultBuffer,
64
+ ]);
65
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
66
+ pipelineReduce,
67
+ bgReduce,
68
+ 1,
69
+ ); // dispatch 1 workgroup to reduce the partial sums to a single result
70
+ const readBuffer = stageReadback(enc2, resultBuffer);
71
+
72
+ submit(enc2);
73
+
74
+ const [gpuTime1, gpuTime2, asumArr] = await Promise.all([
75
+ extractTimestamp(ts1),
76
+ extractTimestamp(ts2),
77
+ extractResult(readBuffer, Float32Array),
78
+ ]);
79
+
80
+ // asum is always a scalar readback — both paths return { asum }
81
+ if (xIsGpu) {
82
+ destroyBuffers(partialsBuffer, resultBuffer, paramsBuffer, readBuffer);
83
+ if (gpuTime1 !== undefined && gpuTime2 !== undefined)
84
+ return { asum: asumArr[0], gpuTimeMs: gpuTime1 + gpuTime2 };
85
+ return { asum: asumArr[0] };
86
+ }
87
+
88
+ destroyBuffers(
89
+ xBuffer,
90
+ partialsBuffer,
91
+ resultBuffer,
92
+ paramsBuffer,
93
+ readBuffer,
94
+ );
95
+ if (gpuTime1 !== undefined && gpuTime2 !== undefined)
96
+ return { asum: asumArr[0], gpuTimeMs: gpuTime1 + gpuTime2 };
97
+ return { asum: asumArr[0] };
98
+ }
@@ -0,0 +1,54 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+
3
+ /**
4
+ * Performs the operation y = alpha * x + y
5
+ *
6
+ * {@includeCode ../../examples/saxpy/saxpy.js}
7
+ *
8
+ * **Browser (standalone HTML):**
9
+ * {@includeCode ../../examples/saxpy/web/saxpy.html}
10
+ *
11
+ * @param device - GPUDevice from `init()`
12
+ * @param n - number of elements (must be a positive integer)
13
+ * @param alpha - scalar multiplier
14
+ * @param x - Float32Array input vector
15
+ * @param incx - stride for x (must be a positive integer)
16
+ * @param y - Float32Array input/output vector
17
+ * @param incy - stride for y (must be a positive integer)
18
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/saxpy/saxpy.mjs#L15">Source code: saxpy.mjs (L15)</a>
19
+ * @category BLAS Level 1
20
+ */
21
+ export declare function saxpy(
22
+ device: GPUDevice,
23
+ n: number,
24
+ alpha: number,
25
+ x: Float32Array,
26
+ incx: number,
27
+ y: Float32Array,
28
+ incy: number,
29
+ ): Promise<{ y: Float32Array } | { y: Float32Array; gpuTimeMs: number }>;
30
+
31
+ /**
32
+ * Performs the operation y = alpha * x + y
33
+ *
34
+ * {@includeCode ../../examples/saxpy/gpuvec.saxpy.js}
35
+ *
36
+ * @param device - GPUDevice from `init()`
37
+ * @param n - number of elements (must be a positive integer)
38
+ * @param alpha - scalar multiplier
39
+ * @param x - GpuVector input vector
40
+ * @param incx - stride for x (must be a positive integer)
41
+ * @param y - GpuVector input/output vector (mutated in place)
42
+ * @param incy - stride for y (must be a positive integer)
43
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/saxpy/saxpy.mjs#L15">Source code: saxpy.mjs (L15)</a>
44
+ * @category BLAS Level 1
45
+ */
46
+ export declare function saxpy(
47
+ device: GPUDevice,
48
+ n: number,
49
+ alpha: number,
50
+ x: GpuVector,
51
+ incx: number,
52
+ y: GpuVector,
53
+ incy: number,
54
+ ): Promise<{} | { gpuTimeMs: number }>;
@@ -0,0 +1,90 @@
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
+
15
+ export async function saxpy(device, n, alpha, x, incx, y, incy) {
16
+ const xIsGpu = x instanceof GpuVector;
17
+ const yIsGpu = y instanceof GpuVector;
18
+
19
+ if (!(device instanceof GPUDevice))
20
+ throw new Error("device must be a GPUDevice.");
21
+ if (
22
+ !Number.isInteger(n) ||
23
+ !Number.isInteger(incx) ||
24
+ !Number.isInteger(incy)
25
+ )
26
+ throw new Error("n, incx, and incy must be integers.");
27
+ if (isNaN(alpha)) throw new Error("alpha must not be NaN.");
28
+ if (!isFinite(alpha)) throw new Error("alpha must be finite.");
29
+ if (incx <= 0 || incy <= 0)
30
+ throw new Error("incx and incy must be positive.");
31
+ if (!xIsGpu && !(x instanceof Float32Array))
32
+ throw new Error("x must be a Float32Array or GpuVector.");
33
+ if (!yIsGpu && !(y instanceof Float32Array))
34
+ throw new Error("y must be a Float32Array or GpuVector.");
35
+ if (xIsGpu !== yIsGpu)
36
+ throw new Error(
37
+ "x and y must be the same type (both Float32Array or both GpuVector).",
38
+ );
39
+ if (n <= 0) return yIsGpu ? {} : { y };
40
+ if (x.length < (n - 1) * incx + 1)
41
+ throw new Error(
42
+ "x does not have enough elements for the given n and incx.",
43
+ );
44
+ if (y.length < (n - 1) * incy + 1)
45
+ throw new Error(
46
+ "y does not have enough elements for the given n and incy.",
47
+ );
48
+
49
+ const pipeline = await getPipeline(device, "saxpy");
50
+
51
+ const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "saxpy-x", false);
52
+ const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "saxpy-y", true);
53
+ const paramsBuffer = createParamsBuffer(
54
+ [
55
+ { value: n, type: "u32" },
56
+ { value: alpha, type: "f32" },
57
+ { value: incx, type: "u32" },
58
+ { value: incy, type: "u32" },
59
+ ],
60
+ "saxpy-params",
61
+ );
62
+
63
+ const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
64
+ xBuffer,
65
+ yBuffer,
66
+ paramsBuffer,
67
+ ]);
68
+ const { commandEncoder, ts } = runComputePass(
69
+ pipeline,
70
+ bindGroup,
71
+ calcWorkgroups(n),
72
+ );
73
+ const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
74
+
75
+ submit(commandEncoder);
76
+
77
+ const gpuTimeMs = await extractTimestamp(ts);
78
+
79
+ if (yIsGpu && xIsGpu) {
80
+ destroyBuffers(paramsBuffer);
81
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
82
+ return {};
83
+ }
84
+
85
+ const result = await extractResult(readBuffer, Float32Array);
86
+ destroyBuffers(xBuffer, yBuffer, paramsBuffer, readBuffer);
87
+
88
+ if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
89
+ return { y: result };
90
+ }
@@ -0,0 +1,50 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+
3
+ /**
4
+ * Performs the operation y = x
5
+ *
6
+ * {@includeCode ../../examples/scopy/scopy.js}
7
+ *
8
+ * **Browser (standalone HTML):**
9
+ * {@includeCode ../../examples/scopy/web/scopy.html}
10
+ *
11
+ * @param device - GPUDevice from `init()`
12
+ * @param n - number of elements (must be a positive integer)
13
+ * @param x - Float32Array input vector
14
+ * @param incx - stride for x (must be a positive integer)
15
+ * @param y - Float32Array output vector
16
+ * @param incy - stride for y (must be a positive integer)
17
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/scopy/scopy.mjs#L15">Source code: scopy.mjs (L15)</a>
18
+ * @category BLAS Level 1
19
+ */
20
+ export declare function scopy(
21
+ device: GPUDevice,
22
+ n: number,
23
+ x: Float32Array,
24
+ incx: number,
25
+ y: Float32Array,
26
+ incy: number,
27
+ ): Promise<{ y: Float32Array } | { y: Float32Array; gpuTimeMs: number }>;
28
+
29
+ /**
30
+ * Performs the operation y = x
31
+ *
32
+ * {@includeCode ../../examples/scopy/gpuvec.scopy.js}
33
+ *
34
+ * @param device - GPUDevice from `init()`
35
+ * @param n - number of elements (must be a positive integer)
36
+ * @param x - GpuVector input vector
37
+ * @param incx - stride for x (must be a positive integer)
38
+ * @param y - GpuVector output vector (mutated in place)
39
+ * @param incy - stride for y (must be a positive integer)
40
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/scopy/scopy.mjs#L15">Source code: scopy.mjs (L15)</a>
41
+ * @category BLAS Level 1
42
+ */
43
+ export declare function scopy(
44
+ device: GPUDevice,
45
+ n: number,
46
+ x: GpuVector,
47
+ incx: number,
48
+ y: GpuVector,
49
+ incy: number,
50
+ ): Promise<{} | { gpuTimeMs: number }>;
@@ -0,0 +1,87 @@
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
+
15
+ export async function scopy(device, n, x, incx, y, incy) {
16
+ const xIsGpu = x instanceof GpuVector;
17
+ const yIsGpu = y instanceof GpuVector;
18
+
19
+ if (!(device instanceof GPUDevice))
20
+ throw new Error("device must be a GPUDevice.");
21
+ if (
22
+ !Number.isInteger(n) ||
23
+ !Number.isInteger(incx) ||
24
+ !Number.isInteger(incy)
25
+ )
26
+ throw new Error("n, incx, and incy must be integers.");
27
+ if (incx <= 0 || incy <= 0)
28
+ throw new Error("incx and incy must be positive.");
29
+ if (!xIsGpu && !(x instanceof Float32Array))
30
+ throw new Error("x must be a Float32Array or GpuVector.");
31
+ if (!yIsGpu && !(y instanceof Float32Array))
32
+ throw new Error("y must be a Float32Array or GpuVector.");
33
+ if (xIsGpu !== yIsGpu)
34
+ throw new Error(
35
+ "x and y must be the same type (both Float32Array or both GpuVector).",
36
+ );
37
+ if (n <= 0) return yIsGpu ? {} : { y };
38
+ if (x.length < (n - 1) * incx + 1)
39
+ throw new Error(
40
+ "x does not have enough elements for the given n and incx.",
41
+ );
42
+ if (y.length < (n - 1) * incy + 1)
43
+ throw new Error(
44
+ "y does not have enough elements for the given n and incy.",
45
+ );
46
+
47
+ const pipeline = await getPipeline(device, "scopy");
48
+
49
+ const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "scopy-x", false);
50
+ const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "scopy-y", true);
51
+ const paramsBuffer = createParamsBuffer(
52
+ [
53
+ { value: n, type: "u32" },
54
+ { value: incx, type: "u32" },
55
+ { value: incy, type: "u32" },
56
+ ],
57
+ "scopy-params",
58
+ );
59
+
60
+ const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
61
+ xBuffer,
62
+ yBuffer,
63
+ paramsBuffer,
64
+ ]);
65
+ const { commandEncoder, ts } = runComputePass(
66
+ pipeline,
67
+ bindGroup,
68
+ calcWorkgroups(n),
69
+ );
70
+ const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
71
+
72
+ submit(commandEncoder);
73
+
74
+ const gpuTimeMs = await extractTimestamp(ts);
75
+
76
+ if (yIsGpu && xIsGpu) {
77
+ destroyBuffers(paramsBuffer);
78
+ if (gpuTimeMs !== undefined) return { gpuTimeMs };
79
+ return {};
80
+ }
81
+
82
+ const result = await extractResult(readBuffer, Float32Array);
83
+ destroyBuffers(xBuffer, yBuffer, paramsBuffer, readBuffer);
84
+
85
+ if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
86
+ return { y: result };
87
+ }
@@ -0,0 +1,52 @@
1
+ import { GpuVector } from "../classes/GpuVector.mjs";
2
+
3
+ /**
4
+ * Computes the dot product of two vectors: result = sum(x[i] * y[i])
5
+ *
6
+ * {@includeCode ../../examples/sdot/sdot.js}
7
+ *
8
+ * **Browser (standalone HTML):**
9
+ * {@includeCode ../../examples/sdot/web/sdot.html}
10
+ *
11
+ * @param device - GPUDevice from `init()`
12
+ * @param n - number of elements (must be a positive integer)
13
+ * @param x - Float32Array input vector
14
+ * @param incx - stride for x (must be a positive integer)
15
+ * @param y - Float32Array input vector
16
+ * @param incy - stride for y (must be a positive integer)
17
+ * @returns dot product scalar — always a CPU readback, even for GpuVector inputs
18
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sdot/sdot.mjs#L18">Source code: sdot.mjs (L18)</a>
19
+ * @category BLAS Level 1
20
+ */
21
+ export declare function sdot(
22
+ device: GPUDevice,
23
+ n: number,
24
+ x: Float32Array,
25
+ incx: number,
26
+ y: Float32Array,
27
+ incy: number,
28
+ ): Promise<{ dot: number } | { dot: number; gpuTimeMs: number }>;
29
+
30
+ /**
31
+ * Computes the dot product of two vectors: result = sum(x[i] * y[i])
32
+ *
33
+ * {@includeCode ../../examples/sdot/gpuvec.sdot.js}
34
+ *
35
+ * @param device - GPUDevice from `init()`
36
+ * @param n - number of elements (must be a positive integer)
37
+ * @param x - GpuVector input vector
38
+ * @param incx - stride for x (must be a positive integer)
39
+ * @param y - GpuVector input vector
40
+ * @param incy - stride for y (must be a positive integer)
41
+ * @returns dot product scalar — always a CPU readback, even for GpuVector inputs
42
+ * @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sdot/sdot.mjs#L18">Source code: sdot.mjs (L18)</a>
43
+ * @category BLAS Level 1
44
+ */
45
+ export declare function sdot(
46
+ device: GPUDevice,
47
+ n: number,
48
+ x: GpuVector,
49
+ incx: number,
50
+ y: GpuVector,
51
+ incy: number,
52
+ ): Promise<{ dot: number } | { dot: number; gpuTimeMs: number }>;
@@ -0,0 +1,118 @@
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
+
16
+ const WGS = 64; //workgroup size
17
+
18
+ export async function sdot(device, n, x, incx, y, incy) {
19
+ const xIsGpu = x instanceof GpuVector;
20
+ const yIsGpu = y instanceof GpuVector;
21
+
22
+ if (!(device instanceof GPUDevice))
23
+ throw new Error("device must be a GPUDevice.");
24
+ if (
25
+ !Number.isInteger(n) ||
26
+ !Number.isInteger(incx) ||
27
+ !Number.isInteger(incy)
28
+ )
29
+ throw new Error("n, incx, and incy must be integers.");
30
+ if (incx <= 0 || incy <= 0)
31
+ throw new Error("incx and incy must be positive.");
32
+ if (!xIsGpu && !(x instanceof Float32Array))
33
+ throw new Error("x must be a Float32Array or GpuVector.");
34
+ if (!yIsGpu && !(y instanceof Float32Array))
35
+ throw new Error("y must be a Float32Array or GpuVector.");
36
+ if (xIsGpu !== yIsGpu)
37
+ throw new Error(
38
+ "x and y must be the same type (both Float32Array or both GpuVector).",
39
+ );
40
+ if (n <= 0) return { dot: 0 };
41
+ if (x.length < (n - 1) * incx + 1)
42
+ throw new Error(
43
+ "x does not have enough elements for the given n and incx.",
44
+ );
45
+ if (y.length < (n - 1) * incy + 1)
46
+ throw new Error(
47
+ "y does not have enough elements for the given n and incy.",
48
+ );
49
+
50
+ const pipelineMain = await getPipeline(device, "sdot");
51
+ const pipelineReduce = await getPipeline(device, "reduction/sum");
52
+
53
+ const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sdot-x", false);
54
+ const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sdot-y", false);
55
+ const partialsBuffer = createStorageBuffer(2 * WGS * 4, "sdot-partials"); //to hold 2*WGS partial sums of f32
56
+ const resultBuffer = createResultBuffer(4, "sdot-result"); //to hold the final float32 dot product
57
+ const paramsBuffer = createParamsBuffer(
58
+ [
59
+ { value: n, type: "u32" },
60
+ { value: incx, type: "u32" },
61
+ { value: incy, type: "u32" },
62
+ ],
63
+ "sdot-params",
64
+ );
65
+
66
+ const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
67
+ xBuffer,
68
+ yBuffer,
69
+ partialsBuffer,
70
+ paramsBuffer,
71
+ ]);
72
+ const { commandEncoder: enc1, ts: ts1 } = runComputePass(
73
+ pipelineMain,
74
+ bgMain,
75
+ 2 * WGS,
76
+ ); //dispatch 2*WGS workgroups
77
+
78
+ submit(enc1);
79
+
80
+ const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
81
+ partialsBuffer,
82
+ resultBuffer,
83
+ ]);
84
+ const { commandEncoder: enc2, ts: ts2 } = runComputePass(
85
+ pipelineReduce,
86
+ bgReduce,
87
+ 1,
88
+ ); // dispatch 1 workgroup to reduce the partial sums to a single result
89
+ const readBuffer = stageReadback(enc2, resultBuffer);
90
+
91
+ submit(enc2);
92
+
93
+ const [gpuTime1, gpuTime2, dotArr] = await Promise.all([
94
+ extractTimestamp(ts1),
95
+ extractTimestamp(ts2),
96
+ extractResult(readBuffer, Float32Array),
97
+ ]);
98
+
99
+ // dot is always a scalar readback — both paths return { dot }
100
+ if (xIsGpu && yIsGpu) {
101
+ destroyBuffers(partialsBuffer, resultBuffer, paramsBuffer, readBuffer);
102
+ if (gpuTime1 !== undefined && gpuTime2 !== undefined)
103
+ return { dot: dotArr[0], gpuTimeMs: gpuTime1 + gpuTime2 };
104
+ return { dot: dotArr[0] };
105
+ }
106
+
107
+ destroyBuffers(
108
+ xBuffer,
109
+ yBuffer,
110
+ partialsBuffer,
111
+ resultBuffer,
112
+ paramsBuffer,
113
+ readBuffer,
114
+ );
115
+ if (gpuTime1 !== undefined && gpuTime2 !== undefined)
116
+ return { dot: dotArr[0], gpuTimeMs: gpuTime1 + gpuTime2 };
117
+ return { dot: dotArr[0] };
118
+ }
@@ -0,0 +1,27 @@
1
+ import argmax from "./reduction/argmax.wgsl";
2
+ import sum from "./reduction/sum.wgsl";
3
+ import sscal from "./sscal.wgsl";
4
+ import sswap from "./sswap.wgsl";
5
+ import saxpy from "./saxpy.wgsl";
6
+ import scopy from "./scopy.wgsl";
7
+ import sdot from "./sdot.wgsl";
8
+ import sasum from "./sasum.wgsl";
9
+ import snrm2 from "./snrm2.wgsl";
10
+ import srot from "./srot.wgsl";
11
+ import srotm from "./srotm.wgsl";
12
+ import isamax from "./isamax.wgsl";
13
+
14
+ export const shaderSources = {
15
+ "reduction/argmax": argmax,
16
+ "reduction/sum": sum,
17
+ sscal,
18
+ sswap,
19
+ saxpy,
20
+ scopy,
21
+ sdot,
22
+ sasum,
23
+ snrm2,
24
+ srot,
25
+ srotm,
26
+ isamax,
27
+ };
@@ -0,0 +1,27 @@
1
+ /**
2
+ * ## Structure
3
+ *
4
+ * `shaders/*.wgsl` — one WGSL compute shader per BLAS routine (sscal, saxpy, sdot, …).
5
+ *
6
+ * `shaders/browser-shaders.mjs` — the browser's runtime shader source. In Node.js, shaders are
7
+ * read directly from disk via `readFileSync`. In the browser there is no filesystem, so this file
8
+ * provides all shader strings inline. Vite bundles it by importing each `.wgsl` file as a string.
9
+ *
10
+ * ## Cross-shader patterns
11
+ *
12
+ * **Fixed workgroup size of 64.** Every shader declares `const WGS: u32 = 64` and
13
+ * `@workgroup_size(64)`. 64 is the minimum `maxComputeInvocationsPerWorkgroup` guaranteed across
14
+ * all WebGPU devices, so this works everywhere without querying device limits.
15
+ *
16
+ * **Single bind group.** All bindings use `@group(0)`. This means the JS side always calls
17
+ * `pipeline.getBindGroupLayout(0)` — no secondary groups to track.
18
+ *
19
+ * The `@binding` indices must match the position of each resource in the array passed to
20
+ * `createBindGroup` — it assigns `binding: 0, 1, 2 …` sequentially, with `resultBuffer` appended last.
21
+ *
22
+ * **All counts and strides are `u32`.** `n`, `x_inc`, `y_inc`, and any other index fields in the
23
+ * `Params` uniform struct are unsigned. This avoids implicit sign-extension when they appear in
24
+ * index expressions like `id * params.x_inc`.
25
+ *
26
+ * @module devdocs/shaders
27
+ */