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.
Files changed (124) hide show
  1. package/README.md +20 -18
  2. package/dist/wgblas.browser.js +2172 -1174
  3. package/index.d.mts +49 -44
  4. package/index.mjs +11 -0
  5. package/package.json +133 -63
  6. package/src/classes/Complex32.d.mts +43 -0
  7. package/src/classes/Complex32.mjs +82 -0
  8. package/src/classes/Complex64.d.mts +44 -0
  9. package/src/classes/Complex64.mjs +76 -0
  10. package/src/classes/GpuMatrix.d.mts +41 -41
  11. package/src/classes/GpuMatrix.mjs +126 -17
  12. package/src/classes/GpuVector.d.mts +36 -40
  13. package/src/classes/GpuVector.mjs +66 -11
  14. package/src/cscal/cscal.d.mts +47 -0
  15. package/src/cscal/cscal.mjs +98 -0
  16. package/src/dasum/dasum.d.mts +4 -4
  17. package/src/dasum/dasum.mjs +38 -20
  18. package/src/daxpy/daxpy.d.mts +56 -0
  19. package/src/daxpy/daxpy.mjs +150 -0
  20. package/src/dcopy/dcopy.d.mts +52 -0
  21. package/src/dcopy/dcopy.mjs +140 -0
  22. package/src/ddot/ddot.d.mts +62 -0
  23. package/src/ddot/ddot.mjs +184 -0
  24. package/src/devdocs.mjs +13 -0
  25. package/src/dnrm2/dnrm2.d.mts +50 -0
  26. package/src/dnrm2/dnrm2.mjs +189 -0
  27. package/src/drot/drot.d.mts +67 -0
  28. package/src/drot/drot.mjs +170 -0
  29. package/src/drotm/drotm.d.mts +67 -0
  30. package/src/drotm/drotm.mjs +171 -0
  31. package/src/dscal/dscal.d.mts +52 -0
  32. package/src/dscal/dscal.mjs +119 -0
  33. package/src/dswap/dswap.d.mts +57 -0
  34. package/src/dswap/dswap.mjs +155 -0
  35. package/src/idamax/idamax.d.mts +20 -2
  36. package/src/idamax/idamax.mjs +56 -24
  37. package/src/init.mjs +117 -56
  38. package/src/isamax/isamax.d.mts +20 -2
  39. package/src/isamax/isamax.mjs +21 -16
  40. package/src/random/random.d.mts +37 -39
  41. package/src/random/random.mjs +39 -7
  42. package/src/sasum/sasum.d.mts +2 -2
  43. package/src/sasum/sasum.mjs +20 -16
  44. package/src/saxpy/saxpy.d.mts +2 -2
  45. package/src/saxpy/saxpy.mjs +14 -11
  46. package/src/scopy/scopy.d.mts +2 -2
  47. package/src/scopy/scopy.mjs +13 -9
  48. package/src/sdot/sdot.d.mts +2 -2
  49. package/src/sdot/sdot.mjs +21 -17
  50. package/src/sgemm/sgemm.d.mts +2 -2
  51. package/src/sgemm/sgemm.mjs +109 -40
  52. package/src/sgemmtr/sgemmtr.d.mts +3 -2
  53. package/src/sgemmtr/sgemmtr.mjs +98 -40
  54. package/src/sgemv/sgemv.d.mts +2 -2
  55. package/src/sgemv/sgemv.mjs +69 -41
  56. package/src/sger/sger.d.mts +2 -2
  57. package/src/sger/sger.mjs +43 -19
  58. package/src/shaders/__test_pipeline_a.wgsl +3 -0
  59. package/src/shaders/__test_pipeline_b.wgsl +2 -0
  60. package/src/shaders/cscal.wgsl +33 -0
  61. package/src/shaders/daxpy.wgsl +66 -0
  62. package/src/shaders/dcopy.wgsl +34 -0
  63. package/src/shaders/ddot.wgsl +106 -0
  64. package/src/shaders/dnrm2.wgsl +167 -0
  65. package/src/shaders/drot.wgsl +81 -0
  66. package/src/shaders/drotm.wgsl +99 -0
  67. package/src/shaders/dscal.wgsl +60 -0
  68. package/src/shaders/dswap.wgsl +38 -0
  69. package/src/shaders/f64/utils/add.wgsl +6 -0
  70. package/src/shaders/f64/utils/divide.wgsl +45 -0
  71. package/src/shaders/f64/utils/multiply.wgsl +19 -10
  72. package/src/shaders/f64/utils/sqrt.wgsl +45 -0
  73. package/src/shaders/index.mjs +233 -14
  74. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  75. package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
  76. package/src/shaders/sgemm_large.wgsl +107 -18
  77. package/src/shaders/sgemm_small.wgsl +115 -15
  78. package/src/shaders/sgemmtr_large.wgsl +4 -1
  79. package/src/shaders/sgemmtr_small.wgsl +4 -1
  80. package/src/shaders/sgemv_n.wgsl +3 -1
  81. package/src/shaders/sgemv_t.wgsl +3 -1
  82. package/src/shaders/snrm2.wgsl +72 -23
  83. package/src/shaders/ssymv.wgsl +3 -1
  84. package/src/snrm2/snrm2.d.mts +2 -2
  85. package/src/snrm2/snrm2.mjs +41 -23
  86. package/src/srot/srot.d.mts +2 -4
  87. package/src/srot/srot.mjs +16 -11
  88. package/src/srotm/srotm.d.mts +2 -4
  89. package/src/srotm/srotm.mjs +17 -11
  90. package/src/sscal/sscal.d.mts +3 -3
  91. package/src/sscal/sscal.mjs +14 -12
  92. package/src/sswap/sswap.d.mts +2 -2
  93. package/src/sswap/sswap.mjs +18 -10
  94. package/src/ssymm/ssymm.d.mts +5 -4
  95. package/src/ssymm/ssymm.mjs +150 -54
  96. package/src/ssymv/ssymv.d.mts +2 -2
  97. package/src/ssymv/ssymv.mjs +47 -26
  98. package/src/ssyr/ssyr.d.mts +2 -2
  99. package/src/ssyr/ssyr.mjs +38 -17
  100. package/src/ssyr2/ssyr2.d.mts +2 -2
  101. package/src/ssyr2/ssyr2.mjs +48 -21
  102. package/src/ssyr2k/ssyr2k.d.mts +3 -2
  103. package/src/ssyr2k/ssyr2k.mjs +140 -62
  104. package/src/ssyrk/ssyrk.d.mts +3 -2
  105. package/src/ssyrk/ssyrk.mjs +91 -39
  106. package/src/strmm/strmm.d.mts +5 -4
  107. package/src/strmm/strmm.mjs +174 -60
  108. package/src/strmv/strmv.d.mts +2 -2
  109. package/src/strmv/strmv.mjs +42 -20
  110. package/src/strsm/strsm.d.mts +6 -4
  111. package/src/strsm/strsm.mjs +438 -174
  112. package/src/strsv/strsv.d.mts +5 -3
  113. package/src/strsv/strsv.mjs +89 -34
  114. package/src/util/benchmark.mjs +9 -9
  115. package/src/util/bindgroup.mjs +1 -3
  116. package/src/util/buffer.mjs +139 -24
  117. package/src/util/complex.mjs +87 -0
  118. package/src/util/compute.mjs +19 -16
  119. package/src/util/constants.mjs +57 -0
  120. package/src/util/device.mjs +49 -0
  121. package/src/util/pipeline.mjs +44 -10
  122. package/src/util/workgroup.mjs +72 -7
  123. package/src/shaders/browser-shaders.mjs +0 -81
  124. package/src/shaders/f64add.wgsl +0 -281
@@ -3,25 +3,244 @@
3
3
  *
4
4
  * `shaders/*.wgsl` — one WGSL compute shader per BLAS routine (sscal, saxpy, sdot, …).
5
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.
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
- * **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.
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
- * **Single bind group.** All bindings use `@group(0)`. This means the JS side always calls
17
- * `pipeline.getBindGroupLayout(0)` — no secondary groups to track.
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
- * 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`.
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
- // Neither vectorized loads (kernel 6) nor warp-tiling (kernel 10) beat this
10
- // at the sizes tried, including warp-tiled variants in the same sweep at
11
- // BM=64/128.
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: array<f32>;
25
- @group(0) @binding(1) var<storage, read> B: array<f32>;
26
- @group(0) @binding(2) var<storage, read_write> C: array<f32>;
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(3) var<uniform> params: Params;
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
- for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
74
- let gRowA = blockRow + innerRowA + loadOffset;
75
- let gColA = t * BK + innerColA;
76
- let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
77
- As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
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
- for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
80
- let gRowB = t * BK + innerRowB + loadOffset;
81
- let gColB = blockCol + innerColB;
82
- let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
83
- Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
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
- C[cIdx] = params.alpha * threadResults[resIdxM * TN + resIdxN] + params.beta * C[cIdx];
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
  }