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.
@@ -1,49 +1,40 @@
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).
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> 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>;
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<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
- }
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 = Packed(partialsMain[i], partialsAux[i]);
36
- let b = Packed(partialsMain[i + WGS], partialsAux[i + WGS]);
37
- tile[i] = addPair(a, b);
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
- if (i < s) { tile[i] = addPair(tile[i], tile[i + s]); }
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
- resultMain[0] = tile[0].main;
47
- resultAux[0] = tile[0].aux;
37
+ resultHi[0] = tile[0].hi;
38
+ resultLo[0] = tile[0].lo;
48
39
  }
49
40
  }
@@ -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
+ }
@@ -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
- }