wgblas 0.1.2 → 1.1.0

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (64) hide show
  1. package/README.md +3 -0
  2. package/dist/wgblas.browser.js +1249 -37
  3. package/index.d.mts +9 -0
  4. package/index.mjs +9 -0
  5. package/package.json +47 -1
  6. package/src/classes/GpuMatrix.d.mts +98 -0
  7. package/src/classes/GpuMatrix.mjs +109 -0
  8. package/src/classes/GpuVector.d.mts +11 -7
  9. package/src/classes/GpuVector.mjs +31 -5
  10. package/src/dasum/dasum.d.mts +48 -0
  11. package/src/dasum/dasum.mjs +121 -0
  12. package/src/init.mjs +3 -1
  13. package/src/isamax/isamax.mjs +66 -67
  14. package/src/random/random.d.mts +37 -0
  15. package/src/random/random.mjs +17 -0
  16. package/src/sasum/sasum.mjs +55 -51
  17. package/src/saxpy/saxpy.mjs +48 -35
  18. package/src/scopy/scopy.mjs +43 -32
  19. package/src/sdot/sdot.mjs +60 -55
  20. package/src/sgemv/sgemv.d.mts +127 -0
  21. package/src/sgemv/sgemv.mjs +148 -0
  22. package/src/sger/sger.d.mts +111 -0
  23. package/src/sger/sger.mjs +136 -0
  24. package/src/shaders/browser-shaders.mjs +26 -0
  25. package/src/shaders/dasum.wgsl +98 -0
  26. package/src/shaders/f64add.wgsl +281 -0
  27. package/src/shaders/isamax.wgsl +32 -9
  28. package/src/shaders/reduction/sumF64.wgsl +49 -0
  29. package/src/shaders/sasum.wgsl +18 -4
  30. package/src/shaders/sdot.wgsl +18 -4
  31. package/src/shaders/sgemv_n.wgsl +75 -0
  32. package/src/shaders/sgemv_t.wgsl +65 -0
  33. package/src/shaders/sger.wgsl +48 -0
  34. package/src/shaders/snrm2.wgsl +22 -4
  35. package/src/shaders/ssymv.wgsl +69 -0
  36. package/src/shaders/ssyr.wgsl +60 -0
  37. package/src/shaders/ssyr2.wgsl +63 -0
  38. package/src/shaders/strmv.wgsl +103 -0
  39. package/src/shaders/strsv_apply_inverse.wgsl +47 -0
  40. package/src/shaders/strsv_invert_block.wgsl +109 -0
  41. package/src/shaders/strsv_update.wgsl +75 -0
  42. package/src/snrm2/snrm2.mjs +56 -52
  43. package/src/srot/srot.mjs +57 -41
  44. package/src/srotm/srotm.mjs +54 -38
  45. package/src/sscal/sscal.mjs +43 -32
  46. package/src/sswap/sswap.mjs +49 -34
  47. package/src/ssymv/ssymv.d.mts +117 -0
  48. package/src/ssymv/ssymv.mjs +135 -0
  49. package/src/ssyr/ssyr.d.mts +100 -0
  50. package/src/ssyr/ssyr.mjs +106 -0
  51. package/src/ssyr2/ssyr2.d.mts +112 -0
  52. package/src/ssyr2/ssyr2.mjs +130 -0
  53. package/src/strmv/strmv.d.mts +117 -0
  54. package/src/strmv/strmv.mjs +138 -0
  55. package/src/strsv/strsv.d.mts +106 -0
  56. package/src/strsv/strsv.mjs +207 -0
  57. package/src/util/benchmark.mjs +1 -1
  58. package/src/util/bindgroup.mjs +14 -10
  59. package/src/util/buffer.mjs +7 -2
  60. package/src/util/compute.mjs +41 -15
  61. package/src/util/f64pack.mjs +152 -0
  62. package/src/util/pipeline.mjs +32 -17
  63. package/src/util/result.mjs +8 -4
  64. package/src/util/workgroup.mjs +10 -10
@@ -0,0 +1,281 @@
1
+ // f64add: adds two doubles, each packed as a [main, aux] pair (main: f32
2
+ // value, aux: raw u32 bits — see src/util/f64pack.mjs; decode()/encode()
3
+ // below are the WGSL mirror of that file's packedToFields()/fieldsToPacked()),
4
+ // producing the sum as another [main, aux] pair.
5
+ //
6
+ // Implements IEEE-754 binary64 addition (align, add/subtract significands,
7
+ // normalize, round-to-nearest-even) using only u32 bitwise/integer
8
+ // arithmetic — WGSL has no 64-bit integer type or arbitrary-precision
9
+ // integers, so each operand's 53-bit significand is carried as a two-word
10
+ // (hi, lo) pair, widened by 3 bits at the bottom to hold guard/round/sticky
11
+ // information while aligning exponents.
12
+
13
+ const EXP_ALL_ONES: u32 = 0x7ffu;
14
+ const BIAS: i32 = 1023;
15
+ const QUIET_NAN_MANTISSA_HI: u32 = 1u << 19u; // bit51 of the 52-bit mantissa -> canonical quiet NaN
16
+
17
+ struct Fields {
18
+ sign: u32,
19
+ rawExp: u32,
20
+ mantissaHi: u32, // 20 bits
21
+ lo: u32, // 32 bits
22
+ }
23
+
24
+ // A packed [main, aux] result — aux stays a raw u32; it must never be stored
25
+ // as an array<f32>/treated as a real float (bit pattern can land on a NaN/
26
+ // Infinity exponent for perfectly ordinary doubles — an f32-typed storage
27
+ // slot canonicalizes/corrupts that on any round-trip). See f64pack.mjs's
28
+ // comment above fieldsToPacked, and dasum.wgsl's xAux/partialsAux bindings.
29
+ struct Packed {
30
+ main: f32,
31
+ aux: u32,
32
+ }
33
+
34
+ // Mirrors packedToFields() in f64pack.mjs.
35
+ fn decode(mainBits: u32, auxBits: u32) -> Fields {
36
+ let sign = mainBits >> 31u;
37
+ let expMain = (mainBits >> 23u) & 0xffu;
38
+ let mantMain = mainBits & 0x7fffffu;
39
+
40
+ let auxSign = auxBits >> 31u;
41
+ let auxExp8 = (auxBits >> 23u) & 0xffu;
42
+ let auxMant23 = auxBits & 0x7fffffu;
43
+
44
+ let expExtra = (auxSign << 2u) | (auxExp8 >> 6u);
45
+ let mantExtra29 = ((auxExp8 & 0x3fu) << 23u) | auxMant23;
46
+
47
+ let rawExp = (expMain << 3u) | expExtra;
48
+ let mantissaHi = mantMain >> 3u;
49
+ let mantTop3 = mantMain & 0x7u;
50
+ let lo = (mantTop3 << 29u) | mantExtra29;
51
+
52
+ return Fields(sign, rawExp, mantissaHi, lo);
53
+ }
54
+
55
+ // Mirrors fieldsToPacked() in f64pack.mjs.
56
+ fn encode(sign: u32, rawExp: u32, mantissaHi: u32, lo: u32) -> Packed {
57
+ let expMain = rawExp >> 3u;
58
+ let expExtra = rawExp & 0x7u;
59
+
60
+ let mantTop3 = lo >> 29u;
61
+ let mantMain = (mantissaHi << 3u) | mantTop3;
62
+ let mantExtra29 = lo & 0x1fffffffu;
63
+
64
+ let mainBits = (sign << 31u) | (expMain << 23u) | mantMain;
65
+
66
+ let auxSign = (expExtra >> 2u) & 0x1u;
67
+ let auxExpTop2 = expExtra & 0x3u;
68
+ let auxExpBot6 = mantExtra29 >> 23u;
69
+ let auxMant23 = mantExtra29 & 0x7fffffu;
70
+ let auxExp8 = (auxExpTop2 << 6u) | auxExpBot6;
71
+
72
+ let auxBits = (auxSign << 31u) | (auxExp8 << 23u) | auxMant23;
73
+
74
+ return Packed(bitcast<f32>(mainBits), auxBits);
75
+ }
76
+
77
+ struct Pair { hi: u32, lo: u32 }
78
+ struct Shifted { hi: u32, lo: u32, sticky: u32 }
79
+
80
+ // Two-word right shift by 0..64+ bits, folding every shifted-out 1 bit into
81
+ // a returned sticky flag — used only for the (potentially huge) exponent
82
+ // alignment shift, where exact bits can't all be kept.
83
+ fn shr_sticky(hi: u32, lo: u32, n: u32) -> Shifted {
84
+ if (n == 0u) {
85
+ return Shifted(hi, lo, 0u);
86
+ }
87
+ if (n >= 64u) {
88
+ return Shifted(0u, 0u, select(0u, 1u, hi != 0u || lo != 0u));
89
+ }
90
+ if (n < 32u) {
91
+ let stickyBits = lo & ((1u << n) - 1u);
92
+ let newLo = (lo >> n) | (hi << (32u - n));
93
+ let newHi = hi >> n;
94
+ return Shifted(newHi, newLo, select(0u, 1u, stickyBits != 0u));
95
+ }
96
+ if (n == 32u) {
97
+ return Shifted(0u, hi, select(0u, 1u, lo != 0u));
98
+ }
99
+ let m = n - 32u;
100
+ let stickyBits = lo | (hi & ((1u << m) - 1u));
101
+ let newLo = hi >> m;
102
+ return Shifted(0u, newLo, select(0u, 1u, stickyBits != 0u));
103
+ }
104
+
105
+ // Two-word left shift by 0..63 bits — used only to renormalize after
106
+ // cancellation, by an amount that exactly matches the leading-zero count,
107
+ // so nothing meaningful is ever lost off the top.
108
+ fn shl(hi: u32, lo: u32, n: u32) -> Pair {
109
+ if (n == 0u) {
110
+ return Pair(hi, lo);
111
+ }
112
+ if (n < 32u) {
113
+ let newHi = (hi << n) | (lo >> (32u - n));
114
+ let newLo = lo << n;
115
+ return Pair(newHi, newLo);
116
+ }
117
+ if (n == 32u) {
118
+ return Pair(lo, 0u);
119
+ }
120
+ let m = n - 32u;
121
+ return Pair(lo << m, 0u);
122
+ }
123
+
124
+ fn add64(aHi: u32, aLo: u32, bHi: u32, bLo: u32) -> Pair {
125
+ let sumLo = aLo + bLo;
126
+ let carry = select(0u, 1u, sumLo < aLo); // wrapped around -> there was a carry
127
+ let sumHi = aHi + bHi + carry;
128
+ return Pair(sumHi, sumLo);
129
+ }
130
+
131
+ // Assumes (aHi:aLo) >= (bHi:bLo) — callers guarantee this so no sign handling is needed.
132
+ fn sub64(aHi: u32, aLo: u32, bHi: u32, bLo: u32) -> Pair {
133
+ let borrow = select(0u, 1u, aLo < bLo);
134
+ let diffLo = aLo - bLo;
135
+ let diffHi = aHi - bHi - borrow;
136
+ return Pair(diffHi, diffLo);
137
+ }
138
+
139
+ fn ge64(aHi: u32, aLo: u32, bHi: u32, bLo: u32) -> bool {
140
+ return aHi > bHi || (aHi == bHi && aLo >= bLo);
141
+ }
142
+
143
+ // The actual IEEE-754 addition, returning decoded Fields rather than an
144
+ // encoded Packed pair — lets a caller that's accumulating many values in a
145
+ // row (e.g. dasum.wgsl's per-thread reduction loop) keep the running total
146
+ // in Fields form the whole time, only encoding once at the very end, instead
147
+ // of paying a decode+encode round-trip on every single addition. computeSum
148
+ // (below) is the Packed-in/Packed-out convenience wrapper around this.
149
+ fn addFields(a: Fields, b: Fields) -> Fields {
150
+ let aIsNaN = a.rawExp == EXP_ALL_ONES && (a.mantissaHi != 0u || a.lo != 0u);
151
+ let bIsNaN = b.rawExp == EXP_ALL_ONES && (b.mantissaHi != 0u || b.lo != 0u);
152
+ if (aIsNaN || bIsNaN) {
153
+ return Fields(0u, EXP_ALL_ONES, QUIET_NAN_MANTISSA_HI, 0u);
154
+ }
155
+
156
+ let aIsInf = a.rawExp == EXP_ALL_ONES; // mantissa==0 here since NaN is excluded above
157
+ let bIsInf = b.rawExp == EXP_ALL_ONES;
158
+ if (aIsInf && bIsInf) {
159
+ if (a.sign != b.sign) {
160
+ return Fields(0u, EXP_ALL_ONES, QUIET_NAN_MANTISSA_HI, 0u);
161
+ }
162
+ return Fields(a.sign, EXP_ALL_ONES, 0u, 0u);
163
+ }
164
+ if (aIsInf) { return Fields(a.sign, EXP_ALL_ONES, 0u, 0u); }
165
+ if (bIsInf) { return Fields(b.sign, EXP_ALL_ONES, 0u, 0u); }
166
+
167
+ let aIsZero = a.rawExp == 0u && a.mantissaHi == 0u && a.lo == 0u;
168
+ let bIsZero = b.rawExp == 0u && b.mantissaHi == 0u && b.lo == 0u;
169
+ if (aIsZero && bIsZero) {
170
+ return Fields(a.sign & b.sign, 0u, 0u, 0u);
171
+ }
172
+ if (aIsZero) { return Fields(b.sign, b.rawExp, b.mantissaHi, b.lo); }
173
+ if (bIsZero) { return Fields(a.sign, a.rawExp, a.mantissaHi, a.lo); }
174
+
175
+ // Effective (unbiased) exponent — subnormals share the smallest normal
176
+ // exponent for alignment purposes and have no implicit leading 1.
177
+ var expA = i32(a.rawExp) - BIAS;
178
+ if (a.rawExp == 0u) { expA = 1 - BIAS; }
179
+ var expB = i32(b.rawExp) - BIAS;
180
+ if (b.rawExp == 0u) { expB = 1 - BIAS; }
181
+
182
+ let implicitA = select(0u, 1u, a.rawExp != 0u);
183
+ let implicitB = select(0u, 1u, b.rawExp != 0u);
184
+
185
+ // Widen each 53-bit significand (implicit + 52-bit mantissa) by 3 zero
186
+ // bits at the bottom — room for guard/round/sticky once alignment shifts happen.
187
+ let sigHiA = (implicitA << 23u) | (a.mantissaHi << 3u) | (a.lo >> 29u);
188
+ let sigLoA = a.lo << 3u;
189
+ let sigHiB = (implicitB << 23u) | (b.mantissaHi << 3u) | (b.lo >> 29u);
190
+ let sigLoB = b.lo << 3u;
191
+
192
+ // P = the operand with the larger exponent (Q = the other); on a tie, P =
193
+ // whichever has the larger significand — keeps subtraction below always
194
+ // non-negative without needing signed magnitudes.
195
+ var signP: u32; var expP: i32; var sigHiP: u32; var sigLoP: u32;
196
+ var signQ: u32; var expQ: i32; var sigHiQ: u32; var sigLoQ: u32;
197
+ if (expA > expB || (expA == expB && ge64(sigHiA, sigLoA, sigHiB, sigLoB))) {
198
+ signP = a.sign; expP = expA; sigHiP = sigHiA; sigLoP = sigLoA;
199
+ signQ = b.sign; expQ = expB; sigHiQ = sigHiB; sigLoQ = sigLoB;
200
+ } else {
201
+ signP = b.sign; expP = expB; sigHiP = sigHiB; sigLoP = sigLoB;
202
+ signQ = a.sign; expQ = expA; sigHiQ = sigHiA; sigLoQ = sigLoA;
203
+ }
204
+
205
+ let diff = u32(expP - expQ);
206
+ let shiftedQ = shr_sticky(sigHiQ, sigLoQ, diff);
207
+ let alignedHiQ = shiftedQ.hi;
208
+ let alignedLoQ = shiftedQ.lo | shiftedQ.sticky; // fold sticky into bit0
209
+
210
+ var sumHi: u32; var sumLo: u32;
211
+ if (signP == signQ) {
212
+ let s = add64(sigHiP, sigLoP, alignedHiQ, alignedLoQ);
213
+ sumHi = s.hi; sumLo = s.lo;
214
+ } else {
215
+ let s = sub64(sigHiP, sigLoP, alignedHiQ, alignedLoQ); // P >= Q(aligned) by construction
216
+ sumHi = s.hi; sumLo = s.lo;
217
+ }
218
+
219
+ if (sumHi == 0u && sumLo == 0u) {
220
+ return Fields(0u, 0u, 0u, 0u); // exact cancellation -> +0
221
+ }
222
+
223
+ // commonExp2: per-bit scale of sumLo's bit0 in the widened representation.
224
+ let commonExp2 = expP - 55;
225
+
226
+ var leadPos: i32;
227
+ if (sumHi != 0u) {
228
+ leadPos = 32 + i32(31u - countLeadingZeros(sumHi));
229
+ } else {
230
+ leadPos = i32(31u - countLeadingZeros(sumLo));
231
+ }
232
+ let tentativeExp = leadPos + commonExp2;
233
+ var targetLSBScale = tentativeExp - 52;
234
+ if (tentativeExp < -1022) { targetLSBScale = -1074; }
235
+ let shiftAmt = targetLSBScale - commonExp2;
236
+
237
+ var keepHi: u32; var keepLo: u32;
238
+ if (shiftAmt <= 0) {
239
+ let sh = shl(sumHi, sumLo, u32(-shiftAmt)); // exact — cancellation only, never loses bits
240
+ keepHi = sh.hi; keepLo = sh.lo;
241
+ } else {
242
+ // Only reached without cancellation (same-sign add, or a tied-exponent
243
+ // subtract with no shrinkage) — shiftAmt here is always exactly 3 or 4,
244
+ // so the dropped bits are fully known from sumLo directly (no sticky
245
+ // approximation needed, unlike the Q-alignment shift above).
246
+ let n = u32(shiftAmt);
247
+ let remainder = sumLo & ((1u << n) - 1u);
248
+ let halfway = 1u << (n - 1u);
249
+ let sh = shr_sticky(sumHi, sumLo, n);
250
+ keepHi = sh.hi; keepLo = sh.lo;
251
+ if (remainder > halfway || (remainder == halfway && (keepLo & 1u) != 0u)) {
252
+ let inc = add64(keepHi, keepLo, 0u, 1u);
253
+ keepHi = inc.hi; keepLo = inc.lo;
254
+ }
255
+ }
256
+
257
+ var resultExpBase = targetLSBScale;
258
+ if ((keepHi & (1u << 21u)) != 0u) { // rounding carried past the 53-bit budget
259
+ let sh = shr_sticky(keepHi, keepLo, 1u); // dropped bit is guaranteed 0 here
260
+ keepHi = sh.hi; keepLo = sh.lo;
261
+ resultExpBase = resultExpBase + 1;
262
+ }
263
+
264
+ let resultSign = signP;
265
+ if ((keepHi & (1u << 20u)) != 0u) { // normal-shaped result
266
+ let unbiasedExp = 52 + resultExpBase;
267
+ let rawExpFinal = unbiasedExp + BIAS;
268
+ if (rawExpFinal >= 2047) {
269
+ return Fields(resultSign, EXP_ALL_ONES, 0u, 0u); // overflow -> Infinity
270
+ }
271
+ return Fields(resultSign, u32(rawExpFinal), keepHi & 0xfffffu, keepLo);
272
+ }
273
+ return Fields(resultSign, 0u, keepHi & 0xfffffu, keepLo); // subnormal result
274
+ }
275
+
276
+ // Packed-in/Packed-out convenience wrapper around addFields — encodes once,
277
+ // after the math, rather than addFields itself needing to know about Packed.
278
+ fn computeSum(a: Fields, b: Fields) -> Packed {
279
+ let f = addFields(a, b);
280
+ return encode(f.sign, f.rawExp, f.mantissaHi, f.lo);
281
+ }
@@ -25,19 +25,42 @@ fn main(
25
25
  ) {
26
26
  // -1.0 is a safe sentinel: any |x[i]| >= 0 beats it,
27
27
  // so workgroups with no elements lose gracefully in the epilogue.
28
- var best_val: f32 = -1.0;
29
- var best_idx: u32 = 0u;
28
+ var best_val0: f32 = -1.0; var best_idx0: u32 = 0u;
29
+ var best_val1: f32 = -1.0; var best_idx1: u32 = 0u;
30
+ var best_val2: f32 = -1.0; var best_idx2: u32 = 0u;
31
+ var best_val3: f32 = -1.0; var best_idx3: u32 = 0u;
30
32
 
31
- for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
33
+ let stride = num_wg.x * WGS;
34
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
35
+
36
+ for (var id = gid.x; id < n4_floor; id += 4u * stride) {
37
+ let v0 = abs(x[ id * params.x_inc]);
38
+ let v1 = abs(x[(id + stride) * params.x_inc]);
39
+ let v2 = abs(x[(id + 2u * stride) * params.x_inc]);
40
+ let v3 = abs(x[(id + 3u * stride) * params.x_inc]);
41
+ if (v0 > best_val0) { best_val0 = v0; best_idx0 = id; }
42
+ if (v1 > best_val1) { best_val1 = v1; best_idx1 = id + stride; }
43
+ if (v2 > best_val2) { best_val2 = v2; best_idx2 = id + 2u * stride; }
44
+ if (v3 > best_val3) { best_val3 = v3; best_idx3 = id + 3u * stride; }
45
+ }
46
+ for (var id = n4_floor + gid.x; id < params.n; id += stride) {
32
47
  let v = abs(x[id * params.x_inc]);
33
- if (v > best_val) {
34
- best_val = v;
35
- best_idx = id;
36
- }
48
+ if (v > best_val0) { best_val0 = v; best_idx0 = id; }
49
+ }
50
+
51
+ // merge 4 independent lanes; prefer lower index on tie (first occurrence wins)
52
+ if (best_val1 > best_val0 || (best_val1 == best_val0 && best_idx1 < best_idx0)) {
53
+ best_val0 = best_val1; best_idx0 = best_idx1;
54
+ }
55
+ if (best_val2 > best_val0 || (best_val2 == best_val0 && best_idx2 < best_idx0)) {
56
+ best_val0 = best_val2; best_idx0 = best_idx2;
57
+ }
58
+ if (best_val3 > best_val0 || (best_val3 == best_val0 && best_idx3 < best_idx0)) {
59
+ best_val0 = best_val3; best_idx0 = best_idx3;
37
60
  }
38
61
 
39
- tile_val[lid.x] = best_val;
40
- tile_idx[lid.x] = best_idx;
62
+ tile_val[lid.x] = best_val0;
63
+ tile_idx[lid.x] = best_idx0;
41
64
  workgroupBarrier();
42
65
 
43
66
  for (var s = WGS / 2u; s > 0u; s >>= 1u) {
@@ -0,0 +1,49 @@
1
+ // sum reduction (f64): collapses 2*WGS partial [main, aux] pairs into one,
2
+ // using computeSum instead of plain f32 `+` (see reduction/sum.wgsl for the
3
+ // f32 original this mirrors).
4
+ // dispatch: 1 workgroup of WGS threads. partialsMain/partialsAux must have
5
+ // exactly 2*WGS entries each.
6
+ //
7
+ // Concatenated after f64add.wgsl by getPipeline (WGSL has no #include),
8
+ // reusing its decode/encode/computeSum and Packed struct — f64add.wgsl
9
+ // declares no bindings and no entry point of its own (just helper functions),
10
+ // so bindings here start at 0 and the entry point is simply `reduce_f64`.
11
+ //
12
+ // partialsAux/result's aux slot are array<u32>, not array<f32> — aux's bits
13
+ // must never pass through an f32-typed storage slot (NaN-bit-pattern
14
+ // corruption risk, see f64pack.mjs and the Packed struct comment above
15
+ // decode()/encode() in f64add.wgsl).
16
+
17
+ @group(0) @binding(0) var<storage, read> partialsMain: array<f32>;
18
+ @group(0) @binding(1) var<storage, read> partialsAux: array<u32>;
19
+ @group(0) @binding(2) var<storage, read_write> resultMain: array<f32, 1>;
20
+ @group(0) @binding(3) var<storage, read_write> resultAux: array<u32, 1>;
21
+
22
+ const WGS: u32 = 64;
23
+
24
+ var<workgroup> tile: array<Packed, 64>;
25
+
26
+ fn addPair(a: Packed, b: Packed) -> Packed {
27
+ return computeSum(decode(bitcast<u32>(a.main), a.aux), decode(bitcast<u32>(b.main), b.aux));
28
+ }
29
+
30
+ @compute @workgroup_size(64)
31
+ fn reduce_f64(
32
+ @builtin(local_invocation_id) lid: vec3u,
33
+ ) {
34
+ let i = lid.x;
35
+ let a = Packed(partialsMain[i], partialsAux[i]);
36
+ let b = Packed(partialsMain[i + WGS], partialsAux[i + WGS]);
37
+ tile[i] = addPair(a, b);
38
+ workgroupBarrier();
39
+
40
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
41
+ if (i < s) { tile[i] = addPair(tile[i], tile[i + s]); }
42
+ workgroupBarrier();
43
+ }
44
+
45
+ if (i == 0u) {
46
+ resultMain[0] = tile[0].main;
47
+ resultAux[0] = tile[0].aux;
48
+ }
49
+ }
@@ -21,11 +21,25 @@ fn main(
21
21
  @builtin(workgroup_id) wgid: vec3u,
22
22
  @builtin(num_workgroups) num_wg: vec3u,
23
23
  ) {
24
- var acc: f32 = 0.0;
25
- for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
26
- acc += abs(x[id * params.x_inc]);
24
+ var acc0: f32 = 0.0;
25
+ var acc1: f32 = 0.0;
26
+ var acc2: f32 = 0.0;
27
+ var acc3: f32 = 0.0;
28
+
29
+ let stride = num_wg.x * WGS;
30
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
31
+
32
+ for (var id = gid.x; id < n4_floor; id += 4u * stride) {
33
+ acc0 += abs(x[ id * params.x_inc]);
34
+ acc1 += abs(x[(id + stride) * params.x_inc]);
35
+ acc2 += abs(x[(id + 2u * stride) * params.x_inc]);
36
+ acc3 += abs(x[(id + 3u * stride) * params.x_inc]);
27
37
  }
28
- tile[lid.x] = acc;
38
+ for (var id = n4_floor + gid.x; id < params.n; id += stride) {
39
+ acc0 += abs(x[id * params.x_inc]);
40
+ }
41
+
42
+ tile[lid.x] = acc0 + acc1 + acc2 + acc3;
29
43
  workgroupBarrier();
30
44
 
31
45
  for (var s = WGS / 2u; s > 0u; s >>= 1u) {
@@ -23,11 +23,25 @@ fn main(
23
23
  @builtin(workgroup_id) wgid: vec3u,
24
24
  @builtin(num_workgroups) num_wg: vec3u,
25
25
  ) {
26
- var acc: f32 = 0.0;
27
- for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
28
- acc += x[id * params.x_inc] * y[id * params.y_inc];
26
+ var acc0: f32 = 0.0;
27
+ var acc1: f32 = 0.0;
28
+ var acc2: f32 = 0.0;
29
+ var acc3: f32 = 0.0;
30
+
31
+ let stride = num_wg.x * WGS;
32
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
33
+
34
+ for (var id = gid.x; id < n4_floor; id += 4u * stride) {
35
+ acc0 += x[ id * params.x_inc] * y[ id * params.y_inc];
36
+ acc1 += x[(id + stride) * params.x_inc] * y[(id + stride) * params.y_inc];
37
+ acc2 += x[(id + 2u * stride) * params.x_inc] * y[(id + 2u * stride) * params.y_inc];
38
+ acc3 += x[(id + 3u * stride) * params.x_inc] * y[(id + 3u * stride) * params.y_inc];
29
39
  }
30
- tile[lid.x] = acc;
40
+ for (var id = n4_floor + gid.x; id < params.n; id += stride) {
41
+ acc0 += x[id * params.x_inc] * y[id * params.y_inc];
42
+ }
43
+
44
+ tile[lid.x] = acc0 + acc1 + acc2 + acc3;
31
45
  workgroupBarrier();
32
46
 
33
47
  for (var s = WGS / 2u; s > 0u; s >>= 1u) {
@@ -0,0 +1,75 @@
1
+ // sgemv_n: y = alpha * A * x + beta * y (A is m×n row-major, no-transpose)
2
+ //
3
+ // One workgroup per output row, with a grid-stride outer loop so the shader
4
+ // still covers all rows when m exceeds maxComputeWorkgroupsPerDimension.
5
+ // Threads stride through A[row, :] and x with coalesced reads (consecutive
6
+ // threads → consecutive addresses). Four independent accumulators let the GPU
7
+ // pipeline memory requests across iterations (ILP=4), hiding the
8
+ // global-memory latency.
9
+
10
+ @group(0) @binding(0) var<storage, read> A: array<f32>;
11
+ @group(0) @binding(1) var<storage, read> x: array<f32>;
12
+ @group(0) @binding(2) var<storage, read_write> y: array<f32>;
13
+
14
+ struct Params {
15
+ m: u32,
16
+ n: u32,
17
+ alpha: f32,
18
+ beta: f32,
19
+ incx: u32,
20
+ incy: u32,
21
+ lda: u32,
22
+ }
23
+
24
+ @group(0) @binding(3) var<uniform> params: Params;
25
+
26
+ const WGS: u32 = 64u;
27
+ var<workgroup> scratch: array<f32, 64>;
28
+
29
+ @compute @workgroup_size(64)
30
+ fn main(
31
+ @builtin(workgroup_id) wgid: vec3u,
32
+ @builtin(local_invocation_id) lid: vec3u,
33
+ @builtin(num_workgroups) nwg: vec3u,
34
+ ) {
35
+ // Grid-stride loop: each workgroup handles ceil(m / nwg.x) rows.
36
+ for (var row = wgid.x; row < params.m; row += nwg.x) {
37
+ let row_base = row * params.lda;
38
+ var acc0: f32 = 0.0;
39
+ var acc1: f32 = 0.0;
40
+ var acc2: f32 = 0.0;
41
+ var acc3: f32 = 0.0;
42
+
43
+ // 4-unrolled loop: each iteration issues 4 independent loads for A and x.
44
+ // The accumulators are independent so the GPU can overlap the memory
45
+ // requests rather than serialising them behind a dependency chain.
46
+ let n4_floor = (params.n / (4u * WGS)) * (4u * WGS);
47
+ for (var j: u32 = lid.x; j < n4_floor; j += 4u * WGS) {
48
+ acc0 += A[row_base + j ] * x[ j * params.incx];
49
+ acc1 += A[row_base + j + WGS ] * x[(j + WGS) * params.incx];
50
+ acc2 += A[row_base + j + 2u * WGS ] * x[(j + 2u * WGS) * params.incx];
51
+ acc3 += A[row_base + j + 3u * WGS ] * x[(j + 3u * WGS) * params.incx];
52
+ }
53
+ // Scalar tail: at most 3*WGS elements left after the unrolled block.
54
+ for (var j: u32 = n4_floor + lid.x; j < params.n; j += WGS) {
55
+ acc0 += A[row_base + j] * x[j * params.incx];
56
+ }
57
+
58
+ // Parallel reduction: 64 → 32 → 16 → 8 → 4 → 2 → 1
59
+ scratch[lid.x] = acc0 + acc1 + acc2 + acc3;
60
+ workgroupBarrier();
61
+ for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
62
+ if lid.x < stride {
63
+ scratch[lid.x] += scratch[lid.x + stride];
64
+ }
65
+ workgroupBarrier();
66
+ }
67
+
68
+ if lid.x == 0u {
69
+ let yi = row * params.incy;
70
+ y[yi] = params.alpha * scratch[0] + params.beta * y[yi];
71
+ }
72
+ // All 64 threads must agree before the next row reuses scratch[].
73
+ workgroupBarrier();
74
+ }
75
+ }
@@ -0,0 +1,65 @@
1
+ // sgemv_t: y = alpha * A^T * x + beta * y (A is m×n row-major, transposed)
2
+ // each thread owns one column of A → one element of y (length n)
3
+ // tiles over x (length m) using shared memory; four independent accumulators
4
+ // let the GPU pipeline A reads across j within each tile (ILP=4)
5
+
6
+ @group(0) @binding(0) var<storage, read> A: array<f32>;
7
+ @group(0) @binding(1) var<storage, read> x: array<f32>;
8
+ @group(0) @binding(2) var<storage, read_write> y: array<f32>;
9
+
10
+ struct Params {
11
+ m: u32,
12
+ n: u32,
13
+ alpha: f32,
14
+ beta: f32,
15
+ incx: u32,
16
+ incy: u32,
17
+ lda: u32,
18
+ }
19
+
20
+ @group(0) @binding(3) var<uniform> params: Params;
21
+
22
+ const WGS: u32 = 64u;
23
+ var<workgroup> x_tile: array<f32, 64>;
24
+
25
+ @compute @workgroup_size(64)
26
+ fn main(
27
+ @builtin(global_invocation_id) gid: vec3u,
28
+ @builtin(local_invocation_id) lid: vec3u,
29
+ ) {
30
+ // each thread owns column col of A → output y[col]
31
+ let col = gid.x;
32
+ // tile over x (length m, the rows of A)
33
+ let m_floor = (params.m / WGS) * WGS;
34
+ var acc0: f32 = 0.0;
35
+ var acc1: f32 = 0.0;
36
+ var acc2: f32 = 0.0;
37
+ var acc3: f32 = 0.0;
38
+
39
+ for (var base = 0u; base < m_floor; base += WGS) {
40
+ // cooperative load: all 64 threads fill x_tile with x[base..base+WGS]
41
+ x_tile[lid.x] = x[(base + lid.x) * params.incx];
42
+ workgroupBarrier();
43
+
44
+ // 4-unrolled inner loop: 4 independent A reads let the GPU pipeline
45
+ // global-memory requests within each tile. WGS=64 divides by 4 exactly.
46
+ if (col < params.n) {
47
+ for (var j = 0u; j < WGS; j += 4u) {
48
+ acc0 += A[(base + j ) * params.lda + col] * x_tile[j ];
49
+ acc1 += A[(base + j + 1) * params.lda + col] * x_tile[j + 1];
50
+ acc2 += A[(base + j + 2) * params.lda + col] * x_tile[j + 2];
51
+ acc3 += A[(base + j + 3) * params.lda + col] * x_tile[j + 3];
52
+ }
53
+ }
54
+ workgroupBarrier();
55
+ }
56
+
57
+ if (col < params.n) {
58
+ // remainder: m not divisible by WGS — short loop, single accumulator fine
59
+ for (var k = m_floor; k < params.m; k++) {
60
+ acc0 += A[k * params.lda + col] * x[k * params.incx];
61
+ }
62
+ let yi = col * params.incy;
63
+ y[yi] = params.alpha * (acc0 + acc1 + acc2 + acc3) + params.beta * y[yi];
64
+ }
65
+ }
@@ -0,0 +1,48 @@
1
+ // sger: A := alpha * x * y^T + A (rank-1 update, A is m×n general/dense)
2
+
3
+ @group(0) @binding(0) var<storage, read> x: array<f32>;
4
+ @group(0) @binding(1) var<storage, read> y: array<f32>;
5
+ @group(0) @binding(2) var<storage, read_write> A: array<f32>;
6
+
7
+ struct Params {
8
+ m: u32,
9
+ n: u32,
10
+ alpha: f32,
11
+ incx: u32,
12
+ incy: u32,
13
+ lda: u32,
14
+ }
15
+
16
+ @group(0) @binding(3) var<uniform> params: Params;
17
+
18
+ const WGS: u32 = 64u;
19
+
20
+ @compute @workgroup_size(64)
21
+ fn main(
22
+ @builtin(workgroup_id) wgid: vec3u,
23
+ @builtin(local_invocation_id) lid: vec3u,
24
+ @builtin(num_workgroups) nwg: vec3u,
25
+ ) {
26
+ for (var row = wgid.x; row < params.m; row += nwg.x) {
27
+ let xi = params.alpha * x[row * params.incx];
28
+ let row_base = row * params.lda;
29
+
30
+ // 4-unrolled loop: each iteration issues 4 independent A/y accesses.
31
+ let n4_floor = (params.n / (4u * WGS)) * (4u * WGS);
32
+ for (var col: u32 = lid.x; col < n4_floor; col += 4u * WGS) {
33
+ let idx0 = row_base + col;
34
+ let idx1 = row_base + col + WGS;
35
+ let idx2 = row_base + col + 2u * WGS;
36
+ let idx3 = row_base + col + 3u * WGS;
37
+ A[idx0] = xi * y[ col * params.incy] + A[idx0];
38
+ A[idx1] = xi * y[(col + WGS) * params.incy] + A[idx1];
39
+ A[idx2] = xi * y[(col + 2u * WGS) * params.incy] + A[idx2];
40
+ A[idx3] = xi * y[(col + 3u * WGS) * params.incy] + A[idx3];
41
+ }
42
+ // Scalar tail: at most 3*WGS elements left after the unrolled block.
43
+ for (var col: u32 = n4_floor + lid.x; col < params.n; col += WGS) {
44
+ let idx = row_base + col;
45
+ A[idx] = xi * y[col * params.incy] + A[idx];
46
+ }
47
+ }
48
+ }
@@ -21,12 +21,30 @@ fn main(
21
21
  @builtin(workgroup_id) wgid: vec3u,
22
22
  @builtin(num_workgroups) num_wg: vec3u,
23
23
  ) {
24
- var acc: f32 = 0.0;
25
- for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
24
+ var acc0: f32 = 0.0;
25
+ var acc1: f32 = 0.0;
26
+ var acc2: f32 = 0.0;
27
+ var acc3: f32 = 0.0;
28
+
29
+ let stride = num_wg.x * WGS;
30
+ let n4_floor = (params.n / (4u * stride)) * (4u * stride);
31
+
32
+ for (var id = gid.x; id < n4_floor; id += 4u * stride) {
33
+ let v0 = x[ id * params.x_inc];
34
+ let v1 = x[(id + stride) * params.x_inc];
35
+ let v2 = x[(id + 2u * stride) * params.x_inc];
36
+ let v3 = x[(id + 3u * stride) * params.x_inc];
37
+ acc0 += v0 * v0;
38
+ acc1 += v1 * v1;
39
+ acc2 += v2 * v2;
40
+ acc3 += v3 * v3;
41
+ }
42
+ for (var id = n4_floor + gid.x; id < params.n; id += stride) {
26
43
  let v = x[id * params.x_inc];
27
- acc += v * v;
44
+ acc0 += v * v;
28
45
  }
29
- tile[lid.x] = acc;
46
+
47
+ tile[lid.x] = acc0 + acc1 + acc2 + acc3;
30
48
  workgroupBarrier();
31
49
 
32
50
  for (var s = WGS / 2u; s > 0u; s >>= 1u) {