wgblas 1.1.0 → 1.2.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +5 -1
- package/dist/wgblas.browser.js +200 -124
- package/index.d.mts +1 -0
- package/index.mjs +1 -0
- package/package.json +1 -1
- package/src/classes/GpuMatrix.d.mts +6 -3
- package/src/classes/GpuMatrix.mjs +15 -25
- package/src/classes/GpuVector.d.mts +5 -2
- package/src/classes/GpuVector.mjs +15 -23
- package/src/dasum/dasum.d.mts +7 -4
- package/src/dasum/dasum.mjs +46 -39
- package/src/shaders/browser-shaders.mjs +2 -0
- package/src/shaders/dasum.wgsl +51 -65
- package/src/shaders/f64/dekker.wgsl +99 -0
- package/src/shaders/reduction/sumF64.wgsl +21 -30
- package/src/util/f64.mjs +44 -0
- package/src/util/f64pack.mjs +0 -152
|
@@ -1,49 +1,40 @@
|
|
|
1
|
-
// sum reduction (f64): collapses 2*WGS partial
|
|
2
|
-
// using
|
|
3
|
-
// f32 original this mirrors).
|
|
4
|
-
// dispatch: 1 workgroup of WGS threads.
|
|
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).
|
|
1
|
+
// sum reduction (f64, double-double): collapses 2*WGS partial (hi, lo) pairs
|
|
2
|
+
// into one, using ddAddProtected instead of plain f32 `+` (see
|
|
3
|
+
// reduction/sum.wgsl for the f32 original this mirrors).
|
|
4
|
+
// dispatch: 1 workgroup of WGS threads. partialsHi/partialsLo must have
|
|
5
|
+
// exactly 2*WGS entries each. Concatenated after f64/dekker.wgsl for
|
|
6
|
+
// DD/ddAddProtected (see it for why plain ddAdd isn't safe).
|
|
16
7
|
|
|
17
|
-
@group(0) @binding(0) var<storage, read>
|
|
18
|
-
@group(0) @binding(1) var<storage, read>
|
|
19
|
-
@group(0) @binding(2) var<storage, read_write>
|
|
20
|
-
@group(0) @binding(3) var<storage, read_write>
|
|
8
|
+
@group(0) @binding(0) var<storage, read> partialsHi: array<f32>;
|
|
9
|
+
@group(0) @binding(1) var<storage, read> partialsLo: array<f32>;
|
|
10
|
+
@group(0) @binding(2) var<storage, read_write> resultHi: array<f32, 1>;
|
|
11
|
+
@group(0) @binding(3) var<storage, read_write> resultLo: array<f32, 1>;
|
|
21
12
|
|
|
22
13
|
const WGS: u32 = 64;
|
|
23
14
|
|
|
24
|
-
var<workgroup> tile: array<
|
|
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
|
-
}
|
|
15
|
+
var<workgroup> tile: array<DD, 64>;
|
|
29
16
|
|
|
30
17
|
@compute @workgroup_size(64)
|
|
31
18
|
fn reduce_f64(
|
|
32
19
|
@builtin(local_invocation_id) lid: vec3u,
|
|
33
20
|
) {
|
|
34
21
|
let i = lid.x;
|
|
35
|
-
let a =
|
|
36
|
-
let b =
|
|
37
|
-
tile[i] =
|
|
22
|
+
let a = DD(partialsHi[i], partialsLo[i]);
|
|
23
|
+
let b = DD(partialsHi[i + WGS], partialsLo[i + WGS]);
|
|
24
|
+
tile[i] = ddAddProtected(a, b, i);
|
|
38
25
|
workgroupBarrier();
|
|
39
26
|
|
|
27
|
+
// ddAddProtected must be called unconditionally by every thread.
|
|
40
28
|
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
41
|
-
|
|
29
|
+
let partner = select(i, i + s, i < s);
|
|
30
|
+
let combined = ddAddProtected(tile[i], tile[partner], i);
|
|
31
|
+
workgroupBarrier();
|
|
32
|
+
if (i < s) { tile[i] = combined; }
|
|
42
33
|
workgroupBarrier();
|
|
43
34
|
}
|
|
44
35
|
|
|
45
36
|
if (i == 0u) {
|
|
46
|
-
|
|
47
|
-
|
|
37
|
+
resultHi[0] = tile[0].hi;
|
|
38
|
+
resultLo[0] = tile[0].lo;
|
|
48
39
|
}
|
|
49
40
|
}
|
package/src/util/f64.mjs
ADDED
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
/** @module devdocs/utility-functions/f64 */
|
|
2
|
+
|
|
3
|
+
// Double-double f64 emulation — splits a double into a (hi, lo) pair of f32
|
|
4
|
+
// values with hi+lo approximating the original, hi holding the leading bits
|
|
5
|
+
// and lo the rounding error hi lost on its own. See
|
|
6
|
+
// src/shaders/f64/dekker.wgsl for the GPU-side arithmetic this pairs with
|
|
7
|
+
// (Dekker's algorithm).
|
|
8
|
+
//
|
|
9
|
+
// Not a value-preserving exact split: double-double buys roughly 2x f32's
|
|
10
|
+
// mantissa (~48 bits vs f32's 24), less than real f64's 52-bit mantissa.
|
|
11
|
+
|
|
12
|
+
/**
|
|
13
|
+
* Splits every element of a Float64Array into a (hi, lo) double-double pair.
|
|
14
|
+
* @param {Float64Array} x
|
|
15
|
+
* @returns {{hi: Float32Array, lo: Float32Array}}
|
|
16
|
+
* @public
|
|
17
|
+
*/
|
|
18
|
+
export function splitDoubleDouble(x) {
|
|
19
|
+
const n = x.length;
|
|
20
|
+
const hi = new Float32Array(n);
|
|
21
|
+
const lo = new Float32Array(n);
|
|
22
|
+
for (let i = 0; i < n; i++) {
|
|
23
|
+
const h = Math.fround(x[i]);
|
|
24
|
+
hi[i] = h;
|
|
25
|
+
lo[i] = Math.fround(x[i] - h);
|
|
26
|
+
}
|
|
27
|
+
return { hi, lo };
|
|
28
|
+
}
|
|
29
|
+
|
|
30
|
+
/**
|
|
31
|
+
* Reassembles a (hi, lo) double-double pair back into a Float64Array — the
|
|
32
|
+
* inverse of splitDoubleDouble, used after reading a double-double GPU
|
|
33
|
+
* readback (GpuVector, GpuMatrix, dasum) back to the CPU.
|
|
34
|
+
* @param {Float32Array} hi
|
|
35
|
+
* @param {Float32Array} lo
|
|
36
|
+
* @returns {Float64Array}
|
|
37
|
+
* @public
|
|
38
|
+
*/
|
|
39
|
+
export function mergeDoubleDouble(hi, lo) {
|
|
40
|
+
const n = hi.length;
|
|
41
|
+
const out = new Float64Array(n);
|
|
42
|
+
for (let i = 0; i < n; i++) out[i] = hi[i] + lo[i];
|
|
43
|
+
return out;
|
|
44
|
+
}
|
package/src/util/f64pack.mjs
DELETED
|
@@ -1,152 +0,0 @@
|
|
|
1
|
-
/** @module devdocs/utility-functions/f64pack */
|
|
2
|
-
|
|
3
|
-
// Scratch buffers reused across calls for bit-pattern <-> float reinterpretation.
|
|
4
|
-
// dv8 is read/written exclusively through explicit big-endian DataView calls
|
|
5
|
-
// (never via Float64Array, whose byte order follows host endianness) so hi/lo
|
|
6
|
-
// word extraction is deterministic regardless of host architecture.
|
|
7
|
-
const buf8 = new ArrayBuffer(8);
|
|
8
|
-
const dv8 = new DataView(buf8);
|
|
9
|
-
|
|
10
|
-
const buf4 = new ArrayBuffer(4);
|
|
11
|
-
const u32View = new Uint32Array(buf4);
|
|
12
|
-
const f32View = new Float32Array(buf4);
|
|
13
|
-
|
|
14
|
-
function u32ToF32(bits) {
|
|
15
|
-
u32View[0] = bits >>> 0;
|
|
16
|
-
return f32View[0];
|
|
17
|
-
}
|
|
18
|
-
|
|
19
|
-
function f32ToU32(f) {
|
|
20
|
-
f32View[0] = f;
|
|
21
|
-
return u32View[0];
|
|
22
|
-
}
|
|
23
|
-
|
|
24
|
-
// Shared bit-scramble between an f64's raw fields (sign, 11-bit exponent,
|
|
25
|
-
// mantissa split as a 20-bit hi word + 32-bit lo word) and a [main, aux]
|
|
26
|
-
// pair. f64 has exactly 3 more exponent bits and 29 more mantissa bits than
|
|
27
|
-
// f32 — 3 + 29 = 32, i.e. exactly one f32's worth of bits, all of which fit
|
|
28
|
-
// in `aux`:
|
|
29
|
-
// aux.sign + aux.exponent[7:6] (3 bits) = f64's low 3 exponent bits
|
|
30
|
-
// aux.exponent[5:0] (6 bits) = top 6 of f64's low 29 mantissa bits
|
|
31
|
-
// aux.mantissa (23 bits) = bottom 23 of f64's low 29 mantissa bits
|
|
32
|
-
//
|
|
33
|
-
// `aux` is deliberately kept as a raw u32 integer, never converted to an
|
|
34
|
-
// actual float32 value (unlike `main`, which genuinely is a float32 for
|
|
35
|
-
// reuse in float-typed math like abs()). Any combination of the source bits
|
|
36
|
-
// above can land on aux.exponent === 0xff — i.e. aux's bit pattern reads as
|
|
37
|
-
// a NaN or Infinity if it's ever treated as a real float — and this can
|
|
38
|
-
// happen not just for unusual inputs but for the RESULT of an ordinary
|
|
39
|
-
// addition (see f64add.wgsl's computeSum), so it can't be filtered out at
|
|
40
|
-
// the input boundary. A JS engine or GPU canonicalizes a NaN bit pattern
|
|
41
|
-
// (sets its mantissa's quiet bit) the moment it passes through an actual
|
|
42
|
-
// Float32Array or f32-typed GPU buffer/register as a VALUE — silently
|
|
43
|
-
// corrupting the payload. Since aux never needs float semantics anywhere
|
|
44
|
-
// (WGSL only ever bitcasts it back to u32 for decode()), storing/uploading/
|
|
45
|
-
// reading it exclusively as Uint32Array bits (see GpuVector, buffer.mjs)
|
|
46
|
-
// sidesteps the whole class of corruption rather than working around it.
|
|
47
|
-
function fieldsToPacked(sign, rawExp, mantissaHi, lo) {
|
|
48
|
-
const expMain = rawExp >>> 3; // top 8 bits -> main's exponent
|
|
49
|
-
const expExtra = rawExp & 0x7; // bottom 3 bits -> stashed in aux
|
|
50
|
-
|
|
51
|
-
const mantTop3 = lo >>> 29; // top 3 bits of the low mantissa word
|
|
52
|
-
const mantMain = (mantissaHi << 3) | mantTop3; // 20 + 3 = 23 bits -> main's mantissa
|
|
53
|
-
const mantExtra29 = lo & 0x1fffffff; // bottom 29 bits -> stashed in aux
|
|
54
|
-
|
|
55
|
-
const mainBits = (sign << 31) | (expMain << 23) | mantMain;
|
|
56
|
-
|
|
57
|
-
const auxSign = (expExtra >>> 2) & 0x1;
|
|
58
|
-
const auxExpTop2 = expExtra & 0x3;
|
|
59
|
-
const auxExpBot6 = mantExtra29 >>> 23; // top 6 of the 29
|
|
60
|
-
const auxMant23 = mantExtra29 & 0x7fffff; // bottom 23 of the 29
|
|
61
|
-
const auxExp8 = (auxExpTop2 << 6) | auxExpBot6;
|
|
62
|
-
|
|
63
|
-
const auxBits = ((auxSign << 31) | (auxExp8 << 23) | auxMant23) >>> 0;
|
|
64
|
-
|
|
65
|
-
return [u32ToF32(mainBits), auxBits];
|
|
66
|
-
}
|
|
67
|
-
|
|
68
|
-
function packedToFields(main, auxBits) {
|
|
69
|
-
const mainBits = f32ToU32(main);
|
|
70
|
-
auxBits = auxBits >>> 0;
|
|
71
|
-
|
|
72
|
-
const sign = mainBits >>> 31;
|
|
73
|
-
const expMain = (mainBits >>> 23) & 0xff;
|
|
74
|
-
const mantMain = mainBits & 0x7fffff; // 23 bits
|
|
75
|
-
|
|
76
|
-
const auxSign = auxBits >>> 31;
|
|
77
|
-
const auxExp8 = (auxBits >>> 23) & 0xff;
|
|
78
|
-
const auxMant23 = auxBits & 0x7fffff;
|
|
79
|
-
|
|
80
|
-
const expExtra = (auxSign << 2) | (auxExp8 >>> 6); // 3 bits
|
|
81
|
-
const mantExtra29 = ((auxExp8 & 0x3f) << 23) | auxMant23; // 29 bits
|
|
82
|
-
|
|
83
|
-
const rawExp = (expMain << 3) | expExtra; // 11 bits
|
|
84
|
-
const mantissaHi = mantMain >>> 3; // top 20 bits
|
|
85
|
-
const mantTop3 = mantMain & 0x7; // bottom 3 bits of mantMain
|
|
86
|
-
|
|
87
|
-
const lo = ((mantTop3 << 29) | mantExtra29) >>> 0;
|
|
88
|
-
|
|
89
|
-
return { sign, rawExp, mantissaHi, lo };
|
|
90
|
-
}
|
|
91
|
-
|
|
92
|
-
// main (unlike aux) genuinely is a float32 value — it gets reinterpreted with
|
|
93
|
-
// real float ops (e.g. abs() in dasum.wgsl), so it can't be transported as
|
|
94
|
-
// raw u32 the way aux is. That leaves a narrow residual gap: main.exponent is
|
|
95
|
-
// `rawExp >>> 3`, which hits 0xff (main's own NaN/Infinity pattern) whenever
|
|
96
|
-
// the f64's raw exponent is >= 2040 — i.e. |value| >= ~1.4e306, plus actual
|
|
97
|
-
// ±Infinity/NaN (rawExp === 2047). Any real Float32Array/GPU round-trip of
|
|
98
|
-
// such a main would get silently NaN-canonicalized, the same corruption class
|
|
99
|
-
// aux was fixed for. Since main can't take that fix, reject the range instead
|
|
100
|
-
// of silently corrupting it.
|
|
101
|
-
const MIN_UNSAFE_RAW_EXP = 2040;
|
|
102
|
-
|
|
103
|
-
/**
|
|
104
|
-
* Packs a double into a [main, aux] pair for storage/transfer where only f32
|
|
105
|
-
* is available (e.g. WGSL, which has no f64 type). This is a raw bit
|
|
106
|
-
* repacking, not a value-preserving numeric split — neither half is a
|
|
107
|
-
* meaningful float on its own; only `unpackF64(main, aux)` reconstructs the
|
|
108
|
-
* original value.
|
|
109
|
-
* @param {number} value - finite double with |value| < ~1.4e306 (see
|
|
110
|
-
* MIN_UNSAFE_RAW_EXP above); larger magnitudes, ±Infinity, and NaN would
|
|
111
|
-
* produce a `main` whose bit pattern is itself NaN/Infinity-shaped, which
|
|
112
|
-
* silently corrupts on any real float32 round-trip, so they're rejected.
|
|
113
|
-
* @returns {[number, number]} `[main, aux]` — main is a float32 value, aux is
|
|
114
|
-
* a raw uint32 bit pattern (0 to 2^32-1). Store/transport aux exclusively
|
|
115
|
-
* via Uint32Array/`array<u32>` — never as an actual float value — see the
|
|
116
|
-
* comment above `fieldsToPacked`.
|
|
117
|
-
*/
|
|
118
|
-
export function packF64(value) {
|
|
119
|
-
dv8.setFloat64(0, value, false);
|
|
120
|
-
const hi = dv8.getUint32(0, false); // sign(1) + exponent(11) + mantissa_hi(20)
|
|
121
|
-
const lo = dv8.getUint32(4, false); // mantissa_lo(32)
|
|
122
|
-
|
|
123
|
-
const sign = hi >>> 31;
|
|
124
|
-
const rawExp = (hi >>> 20) & 0x7ff; // 11 bits
|
|
125
|
-
const mantissaHi = hi & 0xfffff; // 20 bits
|
|
126
|
-
|
|
127
|
-
if (rawExp >= MIN_UNSAFE_RAW_EXP) {
|
|
128
|
-
throw new RangeError(
|
|
129
|
-
`packF64: |${value}| is too large to pack safely (must be finite with ` +
|
|
130
|
-
`magnitude below ~1.4e306); main's bit pattern would itself be NaN/` +
|
|
131
|
-
`Infinity-shaped and get silently corrupted by any real float32 round-trip`,
|
|
132
|
-
);
|
|
133
|
-
}
|
|
134
|
-
|
|
135
|
-
return fieldsToPacked(sign, rawExp, mantissaHi, lo);
|
|
136
|
-
}
|
|
137
|
-
|
|
138
|
-
/**
|
|
139
|
-
* Reconstructs the original double from the [main, aux] pair produced by
|
|
140
|
-
* `packF64`. Bit-exact inverse.
|
|
141
|
-
* @param {number} main - float32 value
|
|
142
|
-
* @param {number} aux - raw uint32 bit pattern (as read from a Uint32Array/`array<u32>`)
|
|
143
|
-
* @returns {number}
|
|
144
|
-
*/
|
|
145
|
-
export function unpackF64(main, aux) {
|
|
146
|
-
const { sign, rawExp, mantissaHi, lo } = packedToFields(main, aux);
|
|
147
|
-
const hi = ((sign << 31) | (rawExp << 20) | mantissaHi) >>> 0;
|
|
148
|
-
|
|
149
|
-
dv8.setUint32(0, hi, false);
|
|
150
|
-
dv8.setUint32(4, lo, false);
|
|
151
|
-
return dv8.getFloat64(0, false);
|
|
152
|
-
}
|