wgblas 2.0.0 → 2.2.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/README.md +20 -18
- package/dist/wgblas.browser.js +2172 -1174
- package/index.d.mts +49 -44
- package/index.mjs +11 -0
- package/package.json +133 -63
- package/src/classes/Complex32.d.mts +43 -0
- package/src/classes/Complex32.mjs +82 -0
- package/src/classes/Complex64.d.mts +44 -0
- package/src/classes/Complex64.mjs +76 -0
- package/src/classes/GpuMatrix.d.mts +41 -41
- package/src/classes/GpuMatrix.mjs +126 -17
- package/src/classes/GpuVector.d.mts +36 -40
- package/src/classes/GpuVector.mjs +66 -11
- package/src/cscal/cscal.d.mts +47 -0
- package/src/cscal/cscal.mjs +98 -0
- package/src/dasum/dasum.d.mts +4 -4
- package/src/dasum/dasum.mjs +38 -20
- package/src/daxpy/daxpy.d.mts +56 -0
- package/src/daxpy/daxpy.mjs +150 -0
- package/src/dcopy/dcopy.d.mts +52 -0
- package/src/dcopy/dcopy.mjs +140 -0
- package/src/ddot/ddot.d.mts +62 -0
- package/src/ddot/ddot.mjs +184 -0
- package/src/devdocs.mjs +13 -0
- package/src/dnrm2/dnrm2.d.mts +50 -0
- package/src/dnrm2/dnrm2.mjs +189 -0
- package/src/drot/drot.d.mts +67 -0
- package/src/drot/drot.mjs +170 -0
- package/src/drotm/drotm.d.mts +67 -0
- package/src/drotm/drotm.mjs +171 -0
- package/src/dscal/dscal.d.mts +52 -0
- package/src/dscal/dscal.mjs +119 -0
- package/src/dswap/dswap.d.mts +57 -0
- package/src/dswap/dswap.mjs +155 -0
- package/src/idamax/idamax.d.mts +20 -2
- package/src/idamax/idamax.mjs +56 -24
- package/src/init.mjs +117 -56
- package/src/isamax/isamax.d.mts +20 -2
- package/src/isamax/isamax.mjs +21 -16
- package/src/random/random.d.mts +37 -39
- package/src/random/random.mjs +39 -7
- package/src/sasum/sasum.d.mts +2 -2
- package/src/sasum/sasum.mjs +20 -16
- package/src/saxpy/saxpy.d.mts +2 -2
- package/src/saxpy/saxpy.mjs +14 -11
- package/src/scopy/scopy.d.mts +2 -2
- package/src/scopy/scopy.mjs +13 -9
- package/src/sdot/sdot.d.mts +2 -2
- package/src/sdot/sdot.mjs +21 -17
- package/src/sgemm/sgemm.d.mts +2 -2
- package/src/sgemm/sgemm.mjs +109 -40
- package/src/sgemmtr/sgemmtr.d.mts +3 -2
- package/src/sgemmtr/sgemmtr.mjs +98 -40
- package/src/sgemv/sgemv.d.mts +2 -2
- package/src/sgemv/sgemv.mjs +69 -41
- package/src/sger/sger.d.mts +2 -2
- package/src/sger/sger.mjs +43 -19
- package/src/shaders/__test_pipeline_a.wgsl +3 -0
- package/src/shaders/__test_pipeline_b.wgsl +2 -0
- package/src/shaders/cscal.wgsl +33 -0
- package/src/shaders/daxpy.wgsl +66 -0
- package/src/shaders/dcopy.wgsl +34 -0
- package/src/shaders/ddot.wgsl +106 -0
- package/src/shaders/dnrm2.wgsl +167 -0
- package/src/shaders/drot.wgsl +81 -0
- package/src/shaders/drotm.wgsl +99 -0
- package/src/shaders/dscal.wgsl +60 -0
- package/src/shaders/dswap.wgsl +38 -0
- package/src/shaders/f64/utils/add.wgsl +6 -0
- package/src/shaders/f64/utils/divide.wgsl +45 -0
- package/src/shaders/f64/utils/multiply.wgsl +19 -10
- package/src/shaders/f64/utils/sqrt.wgsl +45 -0
- package/src/shaders/index.mjs +233 -14
- package/src/shaders/reduction/scaledSum.wgsl +65 -0
- package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
- package/src/shaders/sgemm_large.wgsl +107 -18
- package/src/shaders/sgemm_small.wgsl +115 -15
- package/src/shaders/sgemmtr_large.wgsl +4 -1
- package/src/shaders/sgemmtr_small.wgsl +4 -1
- package/src/shaders/sgemv_n.wgsl +3 -1
- package/src/shaders/sgemv_t.wgsl +3 -1
- package/src/shaders/snrm2.wgsl +72 -23
- package/src/shaders/ssymv.wgsl +3 -1
- package/src/snrm2/snrm2.d.mts +2 -2
- package/src/snrm2/snrm2.mjs +41 -23
- package/src/srot/srot.d.mts +2 -4
- package/src/srot/srot.mjs +16 -11
- package/src/srotm/srotm.d.mts +2 -4
- package/src/srotm/srotm.mjs +17 -11
- package/src/sscal/sscal.d.mts +3 -3
- package/src/sscal/sscal.mjs +14 -12
- package/src/sswap/sswap.d.mts +2 -2
- package/src/sswap/sswap.mjs +18 -10
- package/src/ssymm/ssymm.d.mts +5 -4
- package/src/ssymm/ssymm.mjs +150 -54
- package/src/ssymv/ssymv.d.mts +2 -2
- package/src/ssymv/ssymv.mjs +47 -26
- package/src/ssyr/ssyr.d.mts +2 -2
- package/src/ssyr/ssyr.mjs +38 -17
- package/src/ssyr2/ssyr2.d.mts +2 -2
- package/src/ssyr2/ssyr2.mjs +48 -21
- package/src/ssyr2k/ssyr2k.d.mts +3 -2
- package/src/ssyr2k/ssyr2k.mjs +140 -62
- package/src/ssyrk/ssyrk.d.mts +3 -2
- package/src/ssyrk/ssyrk.mjs +91 -39
- package/src/strmm/strmm.d.mts +5 -4
- package/src/strmm/strmm.mjs +174 -60
- package/src/strmv/strmv.d.mts +2 -2
- package/src/strmv/strmv.mjs +42 -20
- package/src/strsm/strsm.d.mts +6 -4
- package/src/strsm/strsm.mjs +438 -174
- package/src/strsv/strsv.d.mts +5 -3
- package/src/strsv/strsv.mjs +89 -34
- package/src/util/benchmark.mjs +9 -9
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +139 -24
- package/src/util/complex.mjs +87 -0
- package/src/util/compute.mjs +19 -16
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +49 -0
- package/src/util/pipeline.mjs +44 -10
- package/src/util/workgroup.mjs +72 -7
- package/src/shaders/browser-shaders.mjs +0 -81
- package/src/shaders/f64add.wgsl +0 -281
package/src/shaders/index.mjs
CHANGED
|
@@ -3,25 +3,244 @@
|
|
|
3
3
|
*
|
|
4
4
|
* `shaders/*.wgsl` — one WGSL compute shader per BLAS routine (sscal, saxpy, sdot, …).
|
|
5
5
|
*
|
|
6
|
-
* `
|
|
7
|
-
*
|
|
8
|
-
*
|
|
6
|
+
* `routineShaders` below is the single source of truth: routine name → the WGSL source(s)
|
|
7
|
+
* its `getPipeline()` calls actually reference, verified against every `src/<routine>/<routine>.mjs`
|
|
8
|
+
* rather than inferred from naming convention (see its doc comment for the exceptions). Each
|
|
9
|
+
* shader is imported right above the line that adds it — the import *is* the mapping entry, no
|
|
10
|
+
* separate block to cross-reference. `shaderSources`, the flat name → source registry the
|
|
11
|
+
* browser bundle's runtime lookup needs, is *derived* from `routineShaders` rather than
|
|
12
|
+
* hand-duplicated, so the two can never drift apart. In Node.js neither is read — shaders are
|
|
13
|
+
* `readFileSync` from disk directly; `scripts/build-browser.mjs` inlines this module into the
|
|
14
|
+
* browser's IIFE bundle via esbuild instead.
|
|
9
15
|
*
|
|
10
16
|
* ## Cross-shader patterns
|
|
11
17
|
*
|
|
12
|
-
* **
|
|
13
|
-
*
|
|
14
|
-
*
|
|
18
|
+
* **Single bind group.** Every shader with bindings uses `@group(0)` only — the JS side always
|
|
19
|
+
* calls `pipeline.getBindGroupLayout(0)`, no secondary groups to track. Binding order is
|
|
20
|
+
* consistent too: any read-only storage buffers come before read_write ones, with the
|
|
21
|
+
* `uniform Params` struct always last. `@binding` indices match the position of each resource in
|
|
22
|
+
* the array passed to `createBindGroup`, which appends `resultBuffer` last.
|
|
15
23
|
*
|
|
16
|
-
* **
|
|
17
|
-
* `
|
|
24
|
+
* **All counts and strides are `u32`.** `n`, `x_inc`, `y_inc`, and every other index/count field
|
|
25
|
+
* in a `Params` struct is unsigned, avoiding implicit sign-extension in index expressions like
|
|
26
|
+
* `id * params.x_inc`.
|
|
18
27
|
*
|
|
19
|
-
*
|
|
20
|
-
*
|
|
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`.
|
|
28
|
+
* **Entry points don't have to be named `main`.** `loadShader` (`util/pipeline.mjs`)
|
|
29
|
+
* auto-detects the sole `@compute` function in a module instead of requiring a fixed name, so
|
|
30
|
+
* `dasum_main`, `strsv_invert_block_main`, etc. work without renaming.
|
|
25
31
|
*
|
|
26
32
|
* @module devdocs/shaders
|
|
27
33
|
*/
|
|
34
|
+
|
|
35
|
+
/**
|
|
36
|
+
* Routine name → the WGSL source(s) its `getPipeline()` calls reference. Keys are the exact
|
|
37
|
+
* shader names `getPipeline(device, name)` is called with — most routines have one, some pick
|
|
38
|
+
* one of several conditionally (e.g. sgemv's `sgemv_n`/`sgemv_t`, by `trans`), and some have no
|
|
39
|
+
* dedicated shader at all:
|
|
40
|
+
*
|
|
41
|
+
* - `sgemmtr`/`ssyrk`/`ssyr2k` all dispatch through `sgemmtr_small`/`sgemmtr_large`.
|
|
42
|
+
* - `strsm` reuses `strsv_invert_block` and `sscal`, plus its own `block_transfer` and the
|
|
43
|
+
* shared `sgemm_small`/`sgemm_large`.
|
|
44
|
+
* - `dasum`/`idamax` concatenate several f64 utility shaders with their own — see
|
|
45
|
+
* `getPipeline`'s `shaderName: string[]` behaviour.
|
|
46
|
+
* - `random` has no entry — CPU-only, no `getPipeline()` call.
|
|
47
|
+
*
|
|
48
|
+
* Built up entry by entry so each import sits next to the mapping entry that uses it.
|
|
49
|
+
* @public
|
|
50
|
+
*/
|
|
51
|
+
export const routineShaders = {};
|
|
52
|
+
|
|
53
|
+
import sscal from "./sscal.wgsl";
|
|
54
|
+
routineShaders.sscal = { sscal };
|
|
55
|
+
|
|
56
|
+
import cscal from "./cscal.wgsl";
|
|
57
|
+
routineShaders.cscal = { cscal };
|
|
58
|
+
|
|
59
|
+
import sswap from "./sswap.wgsl";
|
|
60
|
+
routineShaders.sswap = { sswap };
|
|
61
|
+
|
|
62
|
+
import dswap from "./dswap.wgsl"; // f64 sibling of sswap — pure data movement, no arithmetic, no reduction/barrier shader needed
|
|
63
|
+
routineShaders.dswap = { dswap };
|
|
64
|
+
|
|
65
|
+
import saxpy from "./saxpy.wgsl";
|
|
66
|
+
routineShaders.saxpy = { saxpy };
|
|
67
|
+
|
|
68
|
+
import scopy from "./scopy.wgsl";
|
|
69
|
+
routineShaders.scopy = { scopy };
|
|
70
|
+
|
|
71
|
+
import dcopy from "./dcopy.wgsl"; // f64 sibling of scopy — pure data movement, no arithmetic, no reduction/barrier shader needed
|
|
72
|
+
routineShaders.dcopy = { dcopy };
|
|
73
|
+
|
|
74
|
+
import sdot from "./sdot.wgsl";
|
|
75
|
+
import sum from "./reduction/sum.wgsl";
|
|
76
|
+
routineShaders.sdot = { sdot, "reduction/sum": sum };
|
|
77
|
+
|
|
78
|
+
import sasum from "./sasum.wgsl";
|
|
79
|
+
routineShaders.sasum = { sasum, "reduction/sum": sum };
|
|
80
|
+
|
|
81
|
+
import snrm2 from "./snrm2.wgsl";
|
|
82
|
+
import scaledSum from "./reduction/scaledSum.wgsl";
|
|
83
|
+
routineShaders.snrm2 = { snrm2, "reduction/scaledSum": scaledSum };
|
|
84
|
+
|
|
85
|
+
import isamax from "./isamax.wgsl";
|
|
86
|
+
import argmax from "./reduction/argmax.wgsl";
|
|
87
|
+
routineShaders.isamax = { isamax, "reduction/argmax": argmax };
|
|
88
|
+
|
|
89
|
+
import dekker from "./f64/dekker.wgsl";
|
|
90
|
+
import ddAbs from "./f64/utils/abs.wgsl";
|
|
91
|
+
import ddAddUtil from "./f64/utils/add.wgsl";
|
|
92
|
+
import dasum from "./dasum.wgsl";
|
|
93
|
+
import sumF64 from "./reduction/sumF64.wgsl";
|
|
94
|
+
routineShaders.dasum = {
|
|
95
|
+
"f64/dekker": dekker,
|
|
96
|
+
"f64/utils/abs": ddAbs,
|
|
97
|
+
"f64/utils/add": ddAddUtil,
|
|
98
|
+
dasum,
|
|
99
|
+
"reduction/sumF64": sumF64,
|
|
100
|
+
};
|
|
101
|
+
|
|
102
|
+
import ddMulUtil from "./f64/utils/multiply.wgsl";
|
|
103
|
+
import ddot from "./ddot.wgsl";
|
|
104
|
+
// multiply.wgsl needs dekker's DD struct and add.wgsl's fsub/negf and
|
|
105
|
+
// fastTwoSumProtected, so those two precede it here.
|
|
106
|
+
routineShaders.ddot = {
|
|
107
|
+
"f64/dekker": dekker,
|
|
108
|
+
"f64/utils/add": ddAddUtil,
|
|
109
|
+
"f64/utils/multiply": ddMulUtil,
|
|
110
|
+
ddot,
|
|
111
|
+
"reduction/sumF64": sumF64,
|
|
112
|
+
};
|
|
113
|
+
|
|
114
|
+
import dscal from "./dscal.wgsl"; // f64 sibling of sscal — no reduction shader needed, unlike dasum/ddot
|
|
115
|
+
routineShaders.dscal = {
|
|
116
|
+
"f64/dekker": dekker,
|
|
117
|
+
"f64/utils/add": ddAddUtil,
|
|
118
|
+
"f64/utils/multiply": ddMulUtil,
|
|
119
|
+
dscal,
|
|
120
|
+
};
|
|
121
|
+
|
|
122
|
+
import daxpy from "./daxpy.wgsl"; // f64 sibling of saxpy — one ddMulProtected + one ddAddProtected per element, no reduction shader needed
|
|
123
|
+
routineShaders.daxpy = {
|
|
124
|
+
"f64/dekker": dekker,
|
|
125
|
+
"f64/utils/add": ddAddUtil,
|
|
126
|
+
"f64/utils/multiply": ddMulUtil,
|
|
127
|
+
daxpy,
|
|
128
|
+
};
|
|
129
|
+
|
|
130
|
+
import ddGreater from "./f64/utils/greater.wgsl";
|
|
131
|
+
import ddEqual from "./f64/utils/equal.wgsl";
|
|
132
|
+
import idamax from "./idamax.wgsl";
|
|
133
|
+
import argmaxF64 from "./reduction/argmaxF64.wgsl";
|
|
134
|
+
routineShaders.idamax = {
|
|
135
|
+
"f64/dekker": dekker,
|
|
136
|
+
"f64/utils/abs": ddAbs,
|
|
137
|
+
"f64/utils/greater": ddGreater,
|
|
138
|
+
"f64/utils/equal": ddEqual,
|
|
139
|
+
idamax,
|
|
140
|
+
"reduction/argmaxF64": argmaxF64,
|
|
141
|
+
};
|
|
142
|
+
|
|
143
|
+
import srot from "./srot.wgsl";
|
|
144
|
+
routineShaders.srot = { srot };
|
|
145
|
+
|
|
146
|
+
import drot from "./drot.wgsl"; // f64 sibling of srot — four ddMulProtected + two ddAddProtected per element, no reduction shader needed
|
|
147
|
+
routineShaders.drot = {
|
|
148
|
+
"f64/dekker": dekker,
|
|
149
|
+
"f64/utils/add": ddAddUtil,
|
|
150
|
+
"f64/utils/multiply": ddMulUtil,
|
|
151
|
+
drot,
|
|
152
|
+
};
|
|
153
|
+
|
|
154
|
+
import srotm from "./srotm.wgsl";
|
|
155
|
+
routineShaders.srotm = { srotm };
|
|
156
|
+
|
|
157
|
+
import drotm from "./drotm.wgsl"; // f64 sibling of srotm — four ddMulProtected + two ddAddProtected per element, no reduction shader needed
|
|
158
|
+
routineShaders.drotm = {
|
|
159
|
+
"f64/dekker": dekker,
|
|
160
|
+
"f64/utils/add": ddAddUtil,
|
|
161
|
+
"f64/utils/multiply": ddMulUtil,
|
|
162
|
+
drotm,
|
|
163
|
+
};
|
|
164
|
+
|
|
165
|
+
import ddDivUtil from "./f64/utils/divide.wgsl";
|
|
166
|
+
import ddSqrtUtil from "./f64/utils/sqrt.wgsl";
|
|
167
|
+
import dnrm2 from "./dnrm2.wgsl"; // f64 sibling of snrm2 — scaled accumulation (Blue's algorithm) ported to double-double, branch-free (select()) since ddDivProtected/ddMulProtected/ddAddProtected's barriers need every thread to take the same path
|
|
168
|
+
import scaledSumF64 from "./reduction/scaledSumF64.wgsl";
|
|
169
|
+
routineShaders.dnrm2 = {
|
|
170
|
+
"f64/dekker": dekker,
|
|
171
|
+
"f64/utils/abs": ddAbs,
|
|
172
|
+
"f64/utils/greater": ddGreater,
|
|
173
|
+
"f64/utils/add": ddAddUtil,
|
|
174
|
+
"f64/utils/multiply": ddMulUtil,
|
|
175
|
+
"f64/utils/divide": ddDivUtil,
|
|
176
|
+
"f64/utils/sqrt": ddSqrtUtil,
|
|
177
|
+
dnrm2,
|
|
178
|
+
"reduction/scaledSumF64": scaledSumF64,
|
|
179
|
+
};
|
|
180
|
+
|
|
181
|
+
import sgemv_n from "./sgemv_n.wgsl";
|
|
182
|
+
import sgemv_t from "./sgemv_t.wgsl";
|
|
183
|
+
routineShaders.sgemv = { sgemv_n, sgemv_t }; // one or the other, picked by trans
|
|
184
|
+
|
|
185
|
+
import ssymv from "./ssymv.wgsl";
|
|
186
|
+
routineShaders.ssymv = { ssymv };
|
|
187
|
+
|
|
188
|
+
import strmv from "./strmv.wgsl";
|
|
189
|
+
routineShaders.strmv = { strmv };
|
|
190
|
+
|
|
191
|
+
import strsv_invert_block from "./strsv_invert_block.wgsl";
|
|
192
|
+
import strsv_apply_inverse from "./strsv_apply_inverse.wgsl";
|
|
193
|
+
import strsv_update from "./strsv_update.wgsl";
|
|
194
|
+
routineShaders.strsv = {
|
|
195
|
+
strsv_invert_block,
|
|
196
|
+
strsv_apply_inverse,
|
|
197
|
+
strsv_update,
|
|
198
|
+
};
|
|
199
|
+
|
|
200
|
+
import sger from "./sger.wgsl";
|
|
201
|
+
routineShaders.sger = { sger };
|
|
202
|
+
|
|
203
|
+
import ssyr from "./ssyr.wgsl";
|
|
204
|
+
routineShaders.ssyr = { ssyr };
|
|
205
|
+
|
|
206
|
+
import ssyr2 from "./ssyr2.wgsl";
|
|
207
|
+
routineShaders.ssyr2 = { ssyr2 };
|
|
208
|
+
|
|
209
|
+
import sgemm_small from "./sgemm_small.wgsl";
|
|
210
|
+
import sgemm_large from "./sgemm_large.wgsl";
|
|
211
|
+
routineShaders.sgemm = { sgemm_small, sgemm_large }; // one or the other, picked by a tile-size threshold
|
|
212
|
+
|
|
213
|
+
import sgemmtr_small from "./sgemmtr_small.wgsl";
|
|
214
|
+
import sgemmtr_large from "./sgemmtr_large.wgsl";
|
|
215
|
+
routineShaders.sgemmtr = { sgemmtr_small, sgemmtr_large };
|
|
216
|
+
|
|
217
|
+
routineShaders.ssyrk = { sgemmtr_small, sgemmtr_large }; // no shader of its own — rides on sgemmtr's
|
|
218
|
+
routineShaders.ssyr2k = { sgemmtr_small, sgemmtr_large }; // no shader of its own — rides on sgemmtr's
|
|
219
|
+
|
|
220
|
+
import symmetrize from "./symmetrize.wgsl";
|
|
221
|
+
routineShaders.ssymm = { sgemm_small, sgemm_large, symmetrize };
|
|
222
|
+
|
|
223
|
+
import triangularize from "./triangularize.wgsl";
|
|
224
|
+
routineShaders.strmm = { sgemm_small, sgemm_large, triangularize };
|
|
225
|
+
|
|
226
|
+
import blockTransfer from "./block_transfer.wgsl";
|
|
227
|
+
routineShaders.strsm = {
|
|
228
|
+
strsv_invert_block,
|
|
229
|
+
block_transfer: blockTransfer,
|
|
230
|
+
sscal,
|
|
231
|
+
sgemm_small,
|
|
232
|
+
sgemm_large,
|
|
233
|
+
};
|
|
234
|
+
|
|
235
|
+
/**
|
|
236
|
+
* Flat shader-name → WGSL source-string registry — what `getPipeline()`/`loadShader()` (see
|
|
237
|
+
* `util/pipeline.mjs`) actually look shaders up in, in the browser. Derived from
|
|
238
|
+
* `routineShaders` by merging every routine's shaders together; shared shaders (e.g.
|
|
239
|
+
* `"reduction/sum"`, used by two different routines above) collapse harmlessly here since
|
|
240
|
+
* every routine's copy is the same imported string, never independently authored text.
|
|
241
|
+
* @public
|
|
242
|
+
*/
|
|
243
|
+
export const shaderSources = Object.assign(
|
|
244
|
+
{},
|
|
245
|
+
...Object.values(routineShaders),
|
|
246
|
+
);
|
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
// scaledSum reduction: collapses 2*WGS (scale, ssq) partials from
|
|
2
|
+
// snrm2.wgsl into the final norm — sqrt(scale² · ssq) == scale · sqrt(ssq).
|
|
3
|
+
// Mirrors reduction/sum.wgsl's shape exactly, merging via ssqMerge (see
|
|
4
|
+
// snrm2.wgsl for the derivation) instead of plain `+`, and taking the final
|
|
5
|
+
// sqrt here rather than on the CPU — unlike sasum/sdot's plain sum, "sum of
|
|
6
|
+
// squares" isn't a meaningful standalone value to hand back, only
|
|
7
|
+
// scale·sqrt(ssq) is.
|
|
8
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
9
|
+
// partialsScale/partialsSsq must have exactly 2*WGS entries each.
|
|
10
|
+
|
|
11
|
+
@group(0) @binding(0) var<storage, read> partialsScale: array<f32>;
|
|
12
|
+
@group(0) @binding(1) var<storage, read> partialsSsq: array<f32>;
|
|
13
|
+
@group(0) @binding(2) var<storage, read_write> result: array<f32>;
|
|
14
|
+
|
|
15
|
+
const WGS: u32 = 64;
|
|
16
|
+
|
|
17
|
+
// True sum-of-squares represented so far == scale² · ssq — see snrm2.wgsl.
|
|
18
|
+
struct ScaleSsq {
|
|
19
|
+
scale: f32,
|
|
20
|
+
ssq: f32,
|
|
21
|
+
}
|
|
22
|
+
|
|
23
|
+
// Associative merge of two independent (scale, ssq) partials.
|
|
24
|
+
fn ssqMerge(a: ScaleSsq, b: ScaleSsq) -> ScaleSsq {
|
|
25
|
+
if (a.scale == 0.0 && b.scale == 0.0) { return ScaleSsq(0.0, 1.0); }
|
|
26
|
+
if (a.scale >= b.scale) {
|
|
27
|
+
let r = b.scale / a.scale;
|
|
28
|
+
return ScaleSsq(a.scale, a.ssq + b.ssq * r * r);
|
|
29
|
+
}
|
|
30
|
+
let r = a.scale / b.scale;
|
|
31
|
+
return ScaleSsq(b.scale, b.ssq + a.ssq * r * r);
|
|
32
|
+
}
|
|
33
|
+
|
|
34
|
+
var<workgroup> tileScale: array<f32, 64>;
|
|
35
|
+
var<workgroup> tileSsq: array<f32, 64>;
|
|
36
|
+
|
|
37
|
+
@compute @workgroup_size(64)
|
|
38
|
+
fn reduce_scaled(
|
|
39
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
40
|
+
) {
|
|
41
|
+
let i = lid.x;
|
|
42
|
+
let merged0 = ssqMerge(
|
|
43
|
+
ScaleSsq(partialsScale[i], partialsSsq[i]),
|
|
44
|
+
ScaleSsq(partialsScale[i + WGS], partialsSsq[i + WGS]),
|
|
45
|
+
);
|
|
46
|
+
tileScale[i] = merged0.scale;
|
|
47
|
+
tileSsq[i] = merged0.ssq;
|
|
48
|
+
workgroupBarrier();
|
|
49
|
+
|
|
50
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
51
|
+
if (i < s) {
|
|
52
|
+
let merged = ssqMerge(
|
|
53
|
+
ScaleSsq(tileScale[i], tileSsq[i]),
|
|
54
|
+
ScaleSsq(tileScale[i + s], tileSsq[i + s]),
|
|
55
|
+
);
|
|
56
|
+
tileScale[i] = merged.scale;
|
|
57
|
+
tileSsq[i] = merged.ssq;
|
|
58
|
+
}
|
|
59
|
+
workgroupBarrier();
|
|
60
|
+
}
|
|
61
|
+
|
|
62
|
+
if (i == 0u) {
|
|
63
|
+
result[0] = tileScale[0] * sqrt(tileSsq[0]);
|
|
64
|
+
}
|
|
65
|
+
}
|
|
@@ -0,0 +1,93 @@
|
|
|
1
|
+
// scaledSum reduction (f64, double-double): collapses 2*WGS (scale, ssq) DD
|
|
2
|
+
// partials from dnrm2.wgsl into the final norm — sqrt(scale² · ssq) ==
|
|
3
|
+
// scale · sqrt(ssq), via ddMulProtected/ddSqrtProtected. Mirrors
|
|
4
|
+
// reduction/scaledSum.wgsl's shape exactly; ssqMergeProtected is duplicated
|
|
5
|
+
// from dnrm2.wgsl rather than shared via f64/utils/ — see that file's own
|
|
6
|
+
// header for why (same convention the f32 pair already uses).
|
|
7
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
8
|
+
// partialsScale*/partialsSsq* must have exactly 2*WGS entries each.
|
|
9
|
+
|
|
10
|
+
@group(0) @binding(0) var<storage, read> partialsScaleHi: array<f32>;
|
|
11
|
+
@group(0) @binding(1) var<storage, read> partialsScaleLo: array<f32>;
|
|
12
|
+
@group(0) @binding(2) var<storage, read> partialsSsqHi: array<f32>;
|
|
13
|
+
@group(0) @binding(3) var<storage, read> partialsSsqLo: array<f32>;
|
|
14
|
+
@group(0) @binding(4) var<storage, read_write> resultHi: array<f32, 1>;
|
|
15
|
+
@group(0) @binding(5) var<storage, read_write> resultLo: array<f32, 1>;
|
|
16
|
+
|
|
17
|
+
const WGS: u32 = 64;
|
|
18
|
+
|
|
19
|
+
struct ScaleSsq {
|
|
20
|
+
scale: DD,
|
|
21
|
+
ssq: DD,
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
fn ddSelect(a: DD, b: DD, cond: bool) -> DD {
|
|
25
|
+
return DD(select(a.hi, b.hi, cond), select(a.lo, b.lo, cond));
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
// Associative merge of two independent (scale, ssq) partials — see
|
|
29
|
+
// dnrm2.wgsl for the derivation and why this is branch-free.
|
|
30
|
+
fn ssqMergeProtected(a: ScaleSsq, b: ScaleSsq, threadSlot: u32) -> ScaleSsq {
|
|
31
|
+
let isBigger = !ddGreater(b.scale, a.scale); // a.scale >= b.scale
|
|
32
|
+
let bigger = ddSelect(b.scale, a.scale, isBigger);
|
|
33
|
+
let smaller = ddSelect(a.scale, b.scale, isBigger);
|
|
34
|
+
let biggerSsq = ddSelect(b.ssq, a.ssq, isBigger);
|
|
35
|
+
let smallerSsq = ddSelect(a.ssq, b.ssq, isBigger);
|
|
36
|
+
let biggerIsZero = bigger.hi == 0.0;
|
|
37
|
+
let safeBigger = ddSelect(bigger, DD(1.0, 0.0), biggerIsZero);
|
|
38
|
+
let r = ddDivProtected(smaller, safeBigger, threadSlot);
|
|
39
|
+
let rsq = ddMulProtected(r, r, threadSlot);
|
|
40
|
+
let smallerSsqTimesRsq = ddMulProtected(smallerSsq, rsq, threadSlot);
|
|
41
|
+
let newSsq = ddAddProtected(biggerSsq, smallerSsqTimesRsq, threadSlot);
|
|
42
|
+
return ScaleSsq(bigger, newSsq);
|
|
43
|
+
}
|
|
44
|
+
|
|
45
|
+
var<workgroup> tileScaleHi: array<f32, 64>;
|
|
46
|
+
var<workgroup> tileScaleLo: array<f32, 64>;
|
|
47
|
+
var<workgroup> tileSsqHi: array<f32, 64>;
|
|
48
|
+
var<workgroup> tileSsqLo: array<f32, 64>;
|
|
49
|
+
|
|
50
|
+
@compute @workgroup_size(64)
|
|
51
|
+
fn reduce_scaled_f64(
|
|
52
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
53
|
+
) {
|
|
54
|
+
let i = lid.x;
|
|
55
|
+
let a = ScaleSsq(DD(partialsScaleHi[i], partialsScaleLo[i]), DD(partialsSsqHi[i], partialsSsqLo[i]));
|
|
56
|
+
let b = ScaleSsq(DD(partialsScaleHi[i + WGS], partialsScaleLo[i + WGS]), DD(partialsSsqHi[i + WGS], partialsSsqLo[i + WGS]));
|
|
57
|
+
let merged0 = ssqMergeProtected(a, b, i);
|
|
58
|
+
tileScaleHi[i] = merged0.scale.hi;
|
|
59
|
+
tileScaleLo[i] = merged0.scale.lo;
|
|
60
|
+
tileSsqHi[i] = merged0.ssq.hi;
|
|
61
|
+
tileSsqLo[i] = merged0.ssq.lo;
|
|
62
|
+
workgroupBarrier();
|
|
63
|
+
|
|
64
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
65
|
+
let partner = select(i, i + s, i < s);
|
|
66
|
+
let ai = ScaleSsq(DD(tileScaleHi[i], tileScaleLo[i]), DD(tileSsqHi[i], tileSsqLo[i]));
|
|
67
|
+
let bi = ScaleSsq(DD(tileScaleHi[partner], tileScaleLo[partner]), DD(tileSsqHi[partner], tileSsqLo[partner]));
|
|
68
|
+
let merged = ssqMergeProtected(ai, bi, i);
|
|
69
|
+
workgroupBarrier();
|
|
70
|
+
if (i < s) {
|
|
71
|
+
tileScaleHi[i] = merged.scale.hi;
|
|
72
|
+
tileScaleLo[i] = merged.scale.lo;
|
|
73
|
+
tileSsqHi[i] = merged.ssq.hi;
|
|
74
|
+
tileSsqLo[i] = merged.ssq.lo;
|
|
75
|
+
}
|
|
76
|
+
workgroupBarrier();
|
|
77
|
+
}
|
|
78
|
+
|
|
79
|
+
// ddSqrtProtected/ddMulProtected's own workgroupBarrier()s need every
|
|
80
|
+
// thread to call them — every thread redundantly computes the same final
|
|
81
|
+
// scale·sqrt(ssq) from tile[0] (still visible to all after the reduction
|
|
82
|
+
// above), and only the write-back is conditional. Guarding the calls
|
|
83
|
+
// themselves behind `if (i == 0u)` (as the plain-f32 original safely
|
|
84
|
+
// does with its unprotected `sqrt()`) would leave 63 threads never
|
|
85
|
+
// reaching a barrier the one remaining thread still needs.
|
|
86
|
+
let scale = DD(tileScaleHi[0], tileScaleLo[0]);
|
|
87
|
+
let ssq = DD(tileSsqHi[0], tileSsqLo[0]);
|
|
88
|
+
let result = ddMulProtected(scale, ddSqrtProtected(ssq, i), i);
|
|
89
|
+
if (i == 0u) {
|
|
90
|
+
resultHi[0] = result.hi;
|
|
91
|
+
resultLo[0] = result.lo;
|
|
92
|
+
}
|
|
93
|
+
}
|
|
@@ -6,9 +6,13 @@
|
|
|
6
6
|
// single-tier baseline at n=512, +84% at n=1024. But BM=64 loses to BM=32
|
|
7
7
|
// below a 6x6=36 workgroup grid (not enough workgroups to fill the GPU at
|
|
8
8
|
// that tile size), hence the two-tier split rather than one global config.
|
|
9
|
-
//
|
|
10
|
-
//
|
|
11
|
-
//
|
|
9
|
+
//
|
|
10
|
+
// A and B are bound twice — scalar array<f32> and array<vec4<f32>> views of
|
|
11
|
+
// the same GPUBuffer (see vec4ViewBinding) — so each tile load can issue
|
|
12
|
+
// 16-byte vector reads along op(A)/op(B)'s contiguous dimension when the
|
|
13
|
+
// stride allows it (stride % 4 == 0 keeps every row base 16-byte aligned).
|
|
14
|
+
// Transposed or odd-stride operands take the scalar path; both paths
|
|
15
|
+
// zero-fill out-of-bounds components identically.
|
|
12
16
|
|
|
13
17
|
const BM: u32 = 64u;
|
|
14
18
|
const BN: u32 = 64u;
|
|
@@ -21,9 +25,11 @@ const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 128
|
|
|
21
25
|
const STRIDE_A: u32 = NUM_THREADS / BK; // rows of As covered per load-loop step
|
|
22
26
|
const STRIDE_B: u32 = NUM_THREADS / BN; // rows of Bs covered per load-loop step
|
|
23
27
|
|
|
24
|
-
@group(0) @binding(0) var<storage, read> A:
|
|
25
|
-
@group(0) @binding(1) var<storage, read>
|
|
26
|
-
@group(0) @binding(2) var<storage,
|
|
28
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
29
|
+
@group(0) @binding(1) var<storage, read> A4: array<vec4<f32>>;
|
|
30
|
+
@group(0) @binding(2) var<storage, read> B: array<f32>;
|
|
31
|
+
@group(0) @binding(3) var<storage, read> B4: array<vec4<f32>>;
|
|
32
|
+
@group(0) @binding(4) var<storage, read_write> C: array<f32>;
|
|
27
33
|
|
|
28
34
|
struct Params {
|
|
29
35
|
m: u32,
|
|
@@ -36,9 +42,11 @@ struct Params {
|
|
|
36
42
|
ldc: u32,
|
|
37
43
|
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
38
44
|
transB: u32,
|
|
45
|
+
useVecA: u32, // 1 = A's vec4 view covers every in-bounds element (see vec4Usable)
|
|
46
|
+
useVecB: u32, // 1 = B's vec4 view covers every in-bounds element
|
|
39
47
|
}
|
|
40
48
|
|
|
41
|
-
@group(0) @binding(
|
|
49
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
42
50
|
|
|
43
51
|
var<workgroup> As: array<f32, BM * BK>;
|
|
44
52
|
var<workgroup> Bs: array<f32, BK * BN>;
|
|
@@ -70,17 +78,95 @@ fn main(
|
|
|
70
78
|
|
|
71
79
|
let numTiles = (params.k + BK - 1u) / BK;
|
|
72
80
|
for (var t = 0u; t < numTiles; t++) {
|
|
73
|
-
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
|
|
81
|
+
// ── Load the BM×BK A tile into As (vectorized along op(A)'s fast dim
|
|
82
|
+
// when lda allows; every branch here is dispatch-uniform) ──
|
|
83
|
+
if (params.useVecA == 1u && params.transA == 0u) {
|
|
84
|
+
// No-transpose: columns contiguous. Each thread loads one vec4 of 4
|
|
85
|
+
// columns; 64 rows × 2 column-lanes = NUM_THREADS exactly, single pass.
|
|
86
|
+
let r4 = tid / (BK / 4u);
|
|
87
|
+
let c4 = tid % (BK / 4u);
|
|
88
|
+
let gRow = blockRow + r4;
|
|
89
|
+
let gCol = t * BK + c4 * 4u;
|
|
90
|
+
var v = A4[(gRow * params.lda + gCol) / 4u];
|
|
91
|
+
let rowOK = gRow < params.m;
|
|
92
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.k);
|
|
93
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.k);
|
|
94
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.k);
|
|
95
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.k);
|
|
96
|
+
As[r4 * BK + c4 * 4u] = v.x;
|
|
97
|
+
As[r4 * BK + c4 * 4u + 1u] = v.y;
|
|
98
|
+
As[r4 * BK + c4 * 4u + 2u] = v.z;
|
|
99
|
+
As[r4 * BK + c4 * 4u + 3u] = v.w;
|
|
100
|
+
} else if (params.useVecA == 1u && params.transA != 0u) {
|
|
101
|
+
// Transpose: rows contiguous within a column. Each thread loads one
|
|
102
|
+
// vec4 of 4 rows; 16 row-lanes × 8 columns = NUM_THREADS, single pass.
|
|
103
|
+
let r4 = tid % (BM / 4u);
|
|
104
|
+
let c = tid / (BM / 4u);
|
|
105
|
+
let gRow = blockRow + r4 * 4u;
|
|
106
|
+
let gCol = t * BK + c;
|
|
107
|
+
var v = A4[(gCol * params.lda + gRow) / 4u];
|
|
108
|
+
let colOK = gCol < params.k;
|
|
109
|
+
v.x = select(0.0, v.x, colOK && gRow < params.m);
|
|
110
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.m);
|
|
111
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.m);
|
|
112
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.m);
|
|
113
|
+
As[(r4 * 4u) * BK + c] = v.x;
|
|
114
|
+
As[(r4 * 4u + 1u) * BK + c] = v.y;
|
|
115
|
+
As[(r4 * 4u + 2u) * BK + c] = v.z;
|
|
116
|
+
As[(r4 * 4u + 3u) * BK + c] = v.w;
|
|
117
|
+
} else {
|
|
118
|
+
// Scalar fallback: odd stride or unhandled orientation.
|
|
119
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
120
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
121
|
+
let gColA = t * BK + innerColA;
|
|
122
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
123
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
124
|
+
}
|
|
78
125
|
}
|
|
79
|
-
|
|
80
|
-
|
|
81
|
-
|
|
82
|
-
|
|
83
|
-
|
|
126
|
+
|
|
127
|
+
// ── Load the BK×BN B tile into Bs ──
|
|
128
|
+
if (params.useVecB == 1u && params.transB == 0u) {
|
|
129
|
+
// No-transpose: columns contiguous. 8 rows × 16 column-lanes cover the
|
|
130
|
+
// tile in one pass (BK = NUM_THREADS / (BN/4)).
|
|
131
|
+
let r = tid / (BN / 4u);
|
|
132
|
+
let c4 = tid % (BN / 4u);
|
|
133
|
+
let gRow = t * BK + r;
|
|
134
|
+
let gCol = blockCol + c4 * 4u;
|
|
135
|
+
var v = B4[(gRow * params.ldb + gCol) / 4u];
|
|
136
|
+
let rowOK = gRow < params.k;
|
|
137
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.n);
|
|
138
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.n);
|
|
139
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.n);
|
|
140
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.n);
|
|
141
|
+
Bs[r * BN + c4 * 4u] = v.x;
|
|
142
|
+
Bs[r * BN + c4 * 4u + 1u] = v.y;
|
|
143
|
+
Bs[r * BN + c4 * 4u + 2u] = v.z;
|
|
144
|
+
Bs[r * BN + c4 * 4u + 3u] = v.w;
|
|
145
|
+
} else if (params.useVecB == 1u && params.transB != 0u) {
|
|
146
|
+
// Transpose: rows contiguous within a column. 2 row-lanes × 64 columns
|
|
147
|
+
// cover the tile in one pass (BN = NUM_THREADS / (BK/4)).
|
|
148
|
+
let r4 = tid % (BK / 4u);
|
|
149
|
+
let c = tid / (BK / 4u);
|
|
150
|
+
let gRow = t * BK + r4 * 4u;
|
|
151
|
+
let gCol = blockCol + c;
|
|
152
|
+
var v = B4[(gCol * params.ldb + gRow) / 4u];
|
|
153
|
+
let colOK = gCol < params.n;
|
|
154
|
+
v.x = select(0.0, v.x, colOK && gRow < params.k);
|
|
155
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.k);
|
|
156
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.k);
|
|
157
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.k);
|
|
158
|
+
Bs[(r4 * 4u) * BN + c] = v.x;
|
|
159
|
+
Bs[(r4 * 4u + 1u) * BN + c] = v.y;
|
|
160
|
+
Bs[(r4 * 4u + 2u) * BN + c] = v.z;
|
|
161
|
+
Bs[(r4 * 4u + 3u) * BN + c] = v.w;
|
|
162
|
+
} else {
|
|
163
|
+
// Scalar fallback.
|
|
164
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
165
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
166
|
+
let gColB = blockCol + innerColB;
|
|
167
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
168
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
169
|
+
}
|
|
84
170
|
}
|
|
85
171
|
|
|
86
172
|
workgroupBarrier();
|
|
@@ -109,7 +195,10 @@ fn main(
|
|
|
109
195
|
let col = blockCol + threadCol * TN + resIdxN;
|
|
110
196
|
if (col < params.n) {
|
|
111
197
|
let cIdx = row * params.ldc + col;
|
|
112
|
-
|
|
198
|
+
// BLAS beta==0 semantics: C is written, not accumulated — must not
|
|
199
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
200
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
201
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
113
202
|
}
|
|
114
203
|
}
|
|
115
204
|
}
|