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.
- package/LICENSE +1 -1
- package/README.md +36 -54
- package/dist/wgblas.browser.js +779 -34
- package/index.d.mts +13 -0
- package/index.mjs +8 -0
- package/package.json +56 -3
- package/src/dasum/dasum.d.mts +2 -2
- package/src/dasum/dasum.mjs +6 -4
- package/src/idamax/idamax.d.mts +51 -0
- package/src/idamax/idamax.mjs +128 -0
- package/src/init.mjs +9 -1
- package/src/isamax/isamax.d.mts +1 -1
- package/src/sasum/sasum.d.mts +1 -1
- package/src/saxpy/saxpy.d.mts +1 -1
- package/src/scopy/scopy.d.mts +1 -1
- package/src/sdot/sdot.d.mts +1 -1
- package/src/sgemm/sgemm.d.mts +102 -0
- package/src/sgemm/sgemm.mjs +195 -0
- package/src/sgemmtr/sgemmtr.d.mts +104 -0
- package/src/sgemmtr/sgemmtr.mjs +203 -0
- package/src/sgemv/sgemv.d.mts +1 -38
- package/src/sgemv/sgemv.mjs +4 -0
- package/src/sger/sger.d.mts +1 -34
- package/src/sger/sger.mjs +2 -0
- package/src/shaders/block_transfer.wgsl +42 -0
- package/src/shaders/browser-shaders.mjs +26 -0
- package/src/shaders/dasum.wgsl +3 -2
- package/src/shaders/f64/dekker.wgsl +4 -85
- package/src/shaders/f64/utils/abs.wgsl +10 -0
- package/src/shaders/f64/utils/add.wgsl +77 -0
- package/src/shaders/f64/utils/equal.wgsl +7 -0
- package/src/shaders/f64/utils/greater.wgsl +12 -0
- package/src/shaders/f64/utils/multiply.wgsl +81 -0
- package/src/shaders/idamax.wgsl +96 -0
- package/src/shaders/reduction/argmaxF64.wgsl +50 -0
- package/src/shaders/reduction/sumF64.wgsl +2 -2
- package/src/shaders/sgemm_large.wgsl +117 -0
- package/src/shaders/sgemm_small.wgsl +112 -0
- package/src/shaders/sgemmtr_large.wgsl +117 -0
- package/src/shaders/sgemmtr_small.wgsl +110 -0
- package/src/shaders/symmetrize.wgsl +31 -0
- package/src/shaders/triangularize.wgsl +44 -0
- package/src/snrm2/snrm2.d.mts +1 -1
- package/src/srot/srot.d.mts +1 -1
- package/src/srotm/srotm.d.mts +1 -1
- package/src/sscal/sscal.d.mts +1 -1
- package/src/sswap/sswap.d.mts +1 -1
- package/src/ssymm/ssymm.d.mts +103 -0
- package/src/ssymm/ssymm.mjs +209 -0
- package/src/ssymv/ssymv.d.mts +1 -36
- package/src/ssymv/ssymv.mjs +2 -0
- package/src/ssyr/ssyr.d.mts +1 -30
- package/src/ssyr/ssyr.mjs +2 -0
- package/src/ssyr2/ssyr2.d.mts +1 -34
- package/src/ssyr2/ssyr2.mjs +2 -0
- package/src/ssyr2k/ssyr2k.d.mts +100 -0
- package/src/ssyr2k/ssyr2k.mjs +201 -0
- package/src/ssyrk/ssyrk.d.mts +90 -0
- package/src/ssyrk/ssyrk.mjs +176 -0
- package/src/strmm/strmm.d.mts +100 -0
- package/src/strmm/strmm.mjs +211 -0
- package/src/strmv/strmv.d.mts +1 -36
- package/src/strmv/strmv.mjs +2 -0
- package/src/strsm/strsm.d.mts +99 -0
- package/src/strsm/strsm.mjs +342 -0
- package/src/strsv/strsv.d.mts +1 -32
- package/src/strsv/strsv.mjs +2 -0
- package/src/util/buffer.mjs +4 -2
- package/src/util/compute.mjs +6 -3
- 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": "
|
|
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.
|
|
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",
|
package/src/dasum/dasum.d.mts
CHANGED
|
@@ -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
|
|
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/
|
|
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)
|
package/src/dasum/dasum.mjs
CHANGED
|
@@ -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
|
|
38
|
-
//
|
|
39
|
-
|
|
40
|
-
const
|
|
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
|
-
|
|
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
|
|
package/src/isamax/isamax.d.mts
CHANGED
|
@@ -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/
|
|
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)
|
package/src/sasum/sasum.d.mts
CHANGED
|
@@ -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/
|
|
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)
|
package/src/saxpy/saxpy.d.mts
CHANGED
|
@@ -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/
|
|
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)
|
package/src/scopy/scopy.d.mts
CHANGED
|
@@ -29,7 +29,7 @@ export declare function scopy(
|
|
|
29
29
|
/**
|
|
30
30
|
* Performs the operation y = x
|
|
31
31
|
*
|
|
32
|
-
* {@includeCode ../../examples/scopy/
|
|
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)
|
package/src/sdot/sdot.d.mts
CHANGED
|
@@ -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/
|
|
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 }>;
|