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.
- package/README.md +3 -0
- package/dist/wgblas.browser.js +1249 -37
- package/index.d.mts +9 -0
- package/index.mjs +9 -0
- package/package.json +47 -1
- package/src/classes/GpuMatrix.d.mts +98 -0
- package/src/classes/GpuMatrix.mjs +109 -0
- package/src/classes/GpuVector.d.mts +11 -7
- package/src/classes/GpuVector.mjs +31 -5
- package/src/dasum/dasum.d.mts +48 -0
- package/src/dasum/dasum.mjs +121 -0
- package/src/init.mjs +3 -1
- package/src/isamax/isamax.mjs +66 -67
- package/src/random/random.d.mts +37 -0
- package/src/random/random.mjs +17 -0
- package/src/sasum/sasum.mjs +55 -51
- package/src/saxpy/saxpy.mjs +48 -35
- package/src/scopy/scopy.mjs +43 -32
- package/src/sdot/sdot.mjs +60 -55
- package/src/sgemv/sgemv.d.mts +127 -0
- package/src/sgemv/sgemv.mjs +148 -0
- package/src/sger/sger.d.mts +111 -0
- package/src/sger/sger.mjs +136 -0
- package/src/shaders/browser-shaders.mjs +26 -0
- package/src/shaders/dasum.wgsl +98 -0
- package/src/shaders/f64add.wgsl +281 -0
- package/src/shaders/isamax.wgsl +32 -9
- package/src/shaders/reduction/sumF64.wgsl +49 -0
- package/src/shaders/sasum.wgsl +18 -4
- package/src/shaders/sdot.wgsl +18 -4
- package/src/shaders/sgemv_n.wgsl +75 -0
- package/src/shaders/sgemv_t.wgsl +65 -0
- package/src/shaders/sger.wgsl +48 -0
- package/src/shaders/snrm2.wgsl +22 -4
- package/src/shaders/ssymv.wgsl +69 -0
- package/src/shaders/ssyr.wgsl +60 -0
- package/src/shaders/ssyr2.wgsl +63 -0
- package/src/shaders/strmv.wgsl +103 -0
- package/src/shaders/strsv_apply_inverse.wgsl +47 -0
- package/src/shaders/strsv_invert_block.wgsl +109 -0
- package/src/shaders/strsv_update.wgsl +75 -0
- package/src/snrm2/snrm2.mjs +56 -52
- package/src/srot/srot.mjs +57 -41
- package/src/srotm/srotm.mjs +54 -38
- package/src/sscal/sscal.mjs +43 -32
- package/src/sswap/sswap.mjs +49 -34
- package/src/ssymv/ssymv.d.mts +117 -0
- package/src/ssymv/ssymv.mjs +135 -0
- package/src/ssyr/ssyr.d.mts +100 -0
- package/src/ssyr/ssyr.mjs +106 -0
- package/src/ssyr2/ssyr2.d.mts +112 -0
- package/src/ssyr2/ssyr2.mjs +130 -0
- package/src/strmv/strmv.d.mts +117 -0
- package/src/strmv/strmv.mjs +138 -0
- package/src/strsv/strsv.d.mts +106 -0
- package/src/strsv/strsv.mjs +207 -0
- package/src/util/benchmark.mjs +1 -1
- package/src/util/bindgroup.mjs +14 -10
- package/src/util/buffer.mjs +7 -2
- package/src/util/compute.mjs +41 -15
- package/src/util/f64pack.mjs +152 -0
- package/src/util/pipeline.mjs +32 -17
- package/src/util/result.mjs +8 -4
- 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
|
+
}
|
package/src/shaders/isamax.wgsl
CHANGED
|
@@ -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
|
|
29
|
-
var
|
|
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
|
-
|
|
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 >
|
|
34
|
-
|
|
35
|
-
|
|
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] =
|
|
40
|
-
tile_idx[lid.x] =
|
|
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
|
+
}
|
package/src/shaders/sasum.wgsl
CHANGED
|
@@ -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
|
|
25
|
-
|
|
26
|
-
|
|
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
|
-
|
|
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) {
|
package/src/shaders/sdot.wgsl
CHANGED
|
@@ -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
|
|
27
|
-
|
|
28
|
-
|
|
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
|
-
|
|
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
|
+
}
|
package/src/shaders/snrm2.wgsl
CHANGED
|
@@ -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
|
|
25
|
-
|
|
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
|
-
|
|
44
|
+
acc0 += v * v;
|
|
28
45
|
}
|
|
29
|
-
|
|
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) {
|