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/index.d.mts CHANGED
@@ -6,6 +6,7 @@ export { GpuMatrix } from "./src/classes/GpuMatrix.mjs";
6
6
  export {
7
7
  randomFloat32Array,
8
8
  randomFloat64Array,
9
+ randomTriangularFloat32Array,
9
10
  } from "./src/random/random.mjs";
10
11
  export { sscal } from "./src/sscal/sscal.mjs";
11
12
  export { sswap } from "./src/sswap/sswap.mjs";
package/index.mjs CHANGED
@@ -4,6 +4,7 @@ export { GpuMatrix } from "./src/classes/GpuMatrix.mjs";
4
4
  export {
5
5
  randomFloat32Array,
6
6
  randomFloat64Array,
7
+ randomTriangularFloat32Array,
7
8
  } from "./src/random/random.mjs";
8
9
  export { sscal } from "./src/sscal/sscal.mjs";
9
10
  export { sswap } from "./src/sswap/sswap.mjs";
package/package.json CHANGED
@@ -1,6 +1,6 @@
1
1
  {
2
2
  "name": "wgblas",
3
- "version": "1.1.0",
3
+ "version": "1.2.0",
4
4
  "description": "BLAS on WebGPU",
5
5
  "type": "module",
6
6
  "main": "index.mjs",
@@ -32,9 +32,12 @@ export declare class GpuMatrix {
32
32
 
33
33
  /**
34
34
  * Uploads a Float32Array or Float64Array matrix to GPU memory, row-major
35
- * or column-major. A Float64Array is packed as two f32s per element (WGSL
36
- * has no f64 type) and stored across two GPU buffers internally; `read()`
37
- * reassembles the original doubles.
35
+ * or column-major. A Float64Array is split into a double-double (hi, lo)
36
+ * f32 pair per element (WGSL has no f64 type) and stored across two GPU
37
+ * buffers internally; `read()` reassembles doubles from these pairs. This
38
+ * gives ~48 bits of mantissa (vs. 24 for a single f32) but less than true
39
+ * f64 precision (52 bits), so round-tripped values are not always
40
+ * bit-exact with the original input.
38
41
  *
39
42
  * `rows`/`cols` always describe the logical shape regardless of layout.
40
43
  * `lda` defaults to `cols` (row-major) or `rows` (column-major) — dense, no
@@ -1,15 +1,12 @@
1
1
  import { getDevice } from "../init.mjs";
2
2
  import { uploadBuffer, stageReadback } from "../util/buffer.mjs";
3
3
  import { extractResult } from "../util/result.mjs";
4
- import { packF64, unpackF64 } from "../util/f64pack.mjs";
4
+ import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
5
5
 
6
6
  export class GpuMatrix {
7
- constructor(buffer, rows, cols, lda, auxBuffer = null, layout = "row-major") {
7
+ constructor(buffer, rows, cols, lda, loBuffer = null, layout = "row-major") {
8
8
  this._buf = buffer;
9
- // Non-null only for Float64Array-backed matrices — see GpuVector for why
10
- // (packF64 splits each element into a "main"/_buf f32 and "aux"/_auxBuf
11
- // raw u32 — never a Float32Array, see f64pack.mjs).
12
- this._auxBuf = auxBuffer;
9
+ this._loBuf = loBuffer; // Non-null only for Float64Array-backed matrices
13
10
  this.rows = rows;
14
11
  this.cols = cols;
15
12
  this.lda = lda;
@@ -50,16 +47,10 @@ export class GpuMatrix {
50
47
 
51
48
  if (data instanceof Float64Array) {
52
49
  const n = outerCount * lda;
53
- const main = new Float32Array(n);
54
- const aux = new Uint32Array(n);
55
- for (let i = 0; i < n; i++) {
56
- const packed = packF64(data[i]);
57
- main[i] = packed[0];
58
- aux[i] = packed[1];
59
- }
60
- const mainBuf = uploadBuffer(main, "gpu-matrix-f64-main", true);
61
- const auxBuf = uploadBuffer(aux, "gpu-matrix-f64-aux", true);
62
- return new GpuMatrix(mainBuf, rows, cols, lda, auxBuf, layout);
50
+ const { hi, lo } = splitDoubleDouble(data.subarray(0, n));
51
+ const hiBuf = uploadBuffer(hi, "gpu-matrix-f64-hi", true);
52
+ const loBuf = uploadBuffer(lo, "gpu-matrix-f64-lo", true);
53
+ return new GpuMatrix(hiBuf, rows, cols, lda, loBuf, layout);
63
54
  }
64
55
 
65
56
  const buf = uploadBuffer(data.subarray(0, outerCount * lda), "gpu-matrix", true);
@@ -76,17 +67,16 @@ export class GpuMatrix {
76
67
  const outerCount = isRowMajor ? this.rows : this.cols;
77
68
  const innerLen = isRowMajor ? this.cols : this.rows;
78
69
 
79
- if (this._auxBuf) {
80
- const encAux = device.createCommandEncoder();
81
- const rbAux = stageReadback(encAux, this._auxBuf);
82
- device.queue.submit([encAux.finish()]);
70
+ if (this._loBuf) {
71
+ const encLo = device.createCommandEncoder();
72
+ const rbLo = stageReadback(encLo, this._loBuf);
73
+ device.queue.submit([encLo.finish()]);
83
74
 
84
- const [main, aux] = await Promise.all([
75
+ const [hi, lo] = await Promise.all([
85
76
  extractResult(rb, Float32Array),
86
- extractResult(rbAux, Uint32Array),
77
+ extractResult(rbLo, Float32Array),
87
78
  ]);
88
- const raw = new Float64Array(outerCount * this.lda);
89
- for (let i = 0; i < raw.length; i++) raw[i] = unpackF64(main[i], aux[i]);
79
+ const raw = mergeDoubleDouble(hi, lo);
90
80
  if (this.lda === innerLen) return raw;
91
81
  const out = new Float64Array(outerCount * innerLen);
92
82
  for (let r = 0; r < outerCount; r++)
@@ -104,6 +94,6 @@ export class GpuMatrix {
104
94
 
105
95
  destroy() {
106
96
  this._buf.destroy();
107
- if (this._auxBuf) this._auxBuf.destroy();
97
+ if (this._loBuf) this._loBuf.destroy();
108
98
  }
109
99
  }
@@ -17,8 +17,11 @@ export declare class GpuVector {
17
17
 
18
18
  /**
19
19
  * Uploads a Float32Array or Float64Array to GPU memory. A Float64Array is
20
- * packed as two f32s per element (WGSL has no f64 type) and stored across
21
- * two GPU buffers internally; `read()` reassembles the original doubles.
20
+ * split into a double-double (hi, lo) f32 pair per element (WGSL has no
21
+ * f64 type) and stored across two GPU buffers internally; `read()`
22
+ * reassembles doubles from these pairs. This gives ~48 bits of mantissa
23
+ * (vs. 24 for a single f32) but less than true f64 precision (52 bits), so
24
+ * round-tripped values are not always bit-exact with the original input.
22
25
  *
23
26
  * @param data - input vector data
24
27
  * @returns GpuVector backed by a GPU buffer
@@ -1,28 +1,22 @@
1
1
  import { getDevice } from "../init.mjs";
2
2
  import { uploadBuffer, stageReadback } from "../util/buffer.mjs";
3
3
  import { extractResult } from "../util/result.mjs";
4
- import { packF64, unpackF64 } from "../util/f64pack.mjs";
4
+ import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
5
5
 
6
6
  export class GpuVector {
7
- constructor(buffer, length, dtype = Float32Array, auxBuffer = null) {
7
+ constructor(buffer, length, dtype = Float32Array, loBuffer = null) {
8
8
  this._buf = buffer;
9
- this._auxBuf = auxBuffer;
9
+ this._loBuf = loBuffer; // Non-null only for Float64Array-backed vectors
10
10
  this.length = length;
11
11
  this.dtype = dtype;
12
12
  }
13
13
 
14
14
  static from(data) {
15
15
  if (data instanceof Float64Array) {
16
- const main = new Float32Array(data.length);
17
- const aux = new Uint32Array(data.length);
18
- for (let i = 0; i < data.length; i++) {
19
- const packed = packF64(data[i]);
20
- main[i] = packed[0];
21
- aux[i] = packed[1];
22
- }
23
- const mainBuf = uploadBuffer(main, "gpu-vector-f64-main", true);
24
- const auxBuf = uploadBuffer(aux, "gpu-vector-f64-aux", true);
25
- return new GpuVector(mainBuf, data.length, Float64Array, auxBuf);
16
+ const { hi, lo } = splitDoubleDouble(data);
17
+ const hiBuf = uploadBuffer(hi, "gpu-vector-f64-hi", true);
18
+ const loBuf = uploadBuffer(lo, "gpu-vector-f64-lo", true);
19
+ return new GpuVector(hiBuf, data.length, Float64Array, loBuf);
26
20
  }
27
21
  if (!(data instanceof Float32Array)) {
28
22
  throw new Error("GpuVector.from expects a Float32Array or Float64Array.");
@@ -37,23 +31,21 @@ export class GpuVector {
37
31
  const rb = stageReadback(enc, this._buf);
38
32
  device.queue.submit([enc.finish()]);
39
33
 
40
- if (!this._auxBuf) return extractResult(rb, this.dtype);
34
+ if (!this._loBuf) return extractResult(rb, this.dtype);
41
35
 
42
- const encAux = device.createCommandEncoder();
43
- const rbAux = stageReadback(encAux, this._auxBuf);
44
- device.queue.submit([encAux.finish()]);
36
+ const encLo = device.createCommandEncoder();
37
+ const rbLo = stageReadback(encLo, this._loBuf);
38
+ device.queue.submit([encLo.finish()]);
45
39
 
46
- const [main, aux] = await Promise.all([
40
+ const [hi, lo] = await Promise.all([
47
41
  extractResult(rb, Float32Array),
48
- extractResult(rbAux, Uint32Array),
42
+ extractResult(rbLo, Float32Array),
49
43
  ]);
50
- const out = new Float64Array(this.length);
51
- for (let i = 0; i < this.length; i++) out[i] = unpackF64(main[i], aux[i]);
52
- return out;
44
+ return mergeDoubleDouble(hi, lo);
53
45
  }
54
46
 
55
47
  destroy() {
56
48
  this._buf.destroy();
57
- if (this._auxBuf) this._auxBuf.destroy();
49
+ if (this._loBuf) this._loBuf.destroy();
58
50
  }
59
51
  }
@@ -1,10 +1,13 @@
1
1
  import { GpuVector } from "../classes/GpuVector.mjs";
2
2
 
3
3
  /**
4
- * Computes the sum of absolute values of a vector of doubles in double
5
- * precision: result = sum(|x[i]|). Each element of `x` is packed into a
6
- * [main, aux] f32 pair (see `packF64`/`GpuVector`) since WGSL has no f64
7
- * type; accumulation is done with f64add.wgsl's IEEE-754 binary64 addition.
4
+ * Computes the sum of absolute values of a vector of doubles in extended
5
+ * precision: result = sum(|x[i]|). Each element of `x` has abs() applied,
6
+ * then is split into a (hi, lo) double-double f32 pair (see
7
+ * `splitDoubleDouble`/`f64.mjs`) since WGSL has no f64 type; accumulation
8
+ * uses Dekker's double-double algorithm (see `shaders/f64/dekker.wgsl`), giving ~48 bits
9
+ * of mantissa — more than a single f32 (24 bits) but less than true f64
10
+ * (52 bits), so results are not bit-exact with a CPU double.
8
11
  *
9
12
  * {@includeCode ../../examples/dasum/dasum.js}
10
13
  *
@@ -4,14 +4,15 @@ import {
4
4
  createResultBuffer,
5
5
  stageReadback,
6
6
  destroyBuffers,
7
+ uploadBuffer,
7
8
  } from "../util/buffer.mjs";
8
9
  import { createBindGroup } from "../util/bindgroup.mjs";
9
10
  import { runComputePass, submit } from "../util/compute.mjs";
10
11
  import { extractTimestamp } from "../util/benchmark.mjs";
11
12
  import { extractResult } from "../util/result.mjs";
12
13
  import { getPipeline } from "../util/pipeline.mjs";
13
- import { unpackF64 } from "../util/f64pack.mjs";
14
14
  import { GpuVector } from "../classes/GpuVector.mjs";
15
+ import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
15
16
 
16
17
  const WGS = 64; // workgroup size
17
18
 
@@ -33,29 +34,34 @@ export async function dasum(device, n, x, incx) {
33
34
  "x does not have enough elements for the given n and incx.",
34
35
  );
35
36
 
36
- // dasum.wgsl/reduction/sumF64.wgsl are each concatenated with f64add.wgsl
37
- // (reusing its decode/encode/computeSum/Packed — WGSL has no #include).
38
- // f64add.wgsl declares no bindings/entry point of its own, so each
39
- // concatenated module has exactly one @compute entry — omitting entryPoint
40
- // here lets getPipeline auto-detect it, the stable path (see pipeline.mjs).
41
- const pipelineMain = await getPipeline(device, ["f64add", "dasum"]);
42
- const pipelineReduce = await getPipeline(device, ["f64add", "reduction/sumF64"]);
37
+ // Concatenated with f64/dekker.wgsl for its DD struct/ddAdd helpers (WGSL
38
+ // has no #include); entryPoint omitted since each module has only one @compute.
39
+ const pipelineMain = await getPipeline(device, ["f64/dekker", "dasum"]);
40
+ const pipelineReduce = await getPipeline(device, ["f64/dekker", "reduction/sumF64"]);
43
41
 
44
- let xVec = null;
45
- let partialsMainBuffer = null;
46
- let partialsAuxBuffer = null;
47
- let resultMainBuffer = null;
48
- let resultAuxBuffer = null;
42
+ let xHiBuffer = null;
43
+ let xLoBuffer = null;
44
+ let partialsHiBuffer = null;
45
+ let partialsLoBuffer = null;
46
+ let resultHiBuffer = null;
47
+ let resultLoBuffer = null;
49
48
  let paramsBuffer = null;
50
- let readMainBuffer = null;
51
- let readAuxBuffer = null;
49
+ let readHiBuffer = null;
50
+ let readLoBuffer = null;
52
51
 
53
52
  try {
54
- xVec = xIsGpu ? x : GpuVector.from(x);
55
- partialsMainBuffer = createStorageBuffer(2 * WGS * 4, "dasum-partialsMain"); // 2*WGS partial sums, main halves
56
- partialsAuxBuffer = createStorageBuffer(2 * WGS * 4, "dasum-partialsAux"); // 2*WGS partial sums, aux halves (raw u32 bits)
57
- resultMainBuffer = createResultBuffer(4, "dasum-result-main"); // final main half
58
- resultAuxBuffer = createResultBuffer(4, "dasum-result-aux"); // final aux half (raw u32 bits)
53
+ if (xIsGpu) {
54
+ xHiBuffer = x._buf;
55
+ xLoBuffer = x._loBuf;
56
+ } else {
57
+ const { hi, lo } = splitDoubleDouble(x.map(Math.abs));
58
+ xHiBuffer = uploadBuffer(hi, "dasum-xHi", false);
59
+ xLoBuffer = uploadBuffer(lo, "dasum-xLo", false);
60
+ }
61
+ partialsHiBuffer = createStorageBuffer(2 * WGS * 4, "dasum-partialsHi");
62
+ partialsLoBuffer = createStorageBuffer(2 * WGS * 4, "dasum-partialsLo");
63
+ resultHiBuffer = createResultBuffer(4, "dasum-result-hi");
64
+ resultLoBuffer = createResultBuffer(4, "dasum-result-lo");
59
65
  paramsBuffer = createParamsBuffer(
60
66
  [
61
67
  { value: n, type: "u32" },
@@ -66,7 +72,7 @@ export async function dasum(device, n, x, incx) {
66
72
 
67
73
  const bgMain = createBindGroup(
68
74
  pipelineMain.getBindGroupLayout(0),
69
- [xVec._buf, xVec._auxBuf, partialsMainBuffer, partialsAuxBuffer, paramsBuffer],
75
+ [xHiBuffer, xLoBuffer, partialsHiBuffer, partialsLoBuffer, paramsBuffer],
70
76
  );
71
77
  const { commandEncoder: enc1, ts: ts1 } = runComputePass(
72
78
  pipelineMain,
@@ -78,44 +84,45 @@ export async function dasum(device, n, x, incx) {
78
84
 
79
85
  const bgReduce = createBindGroup(
80
86
  pipelineReduce.getBindGroupLayout(0),
81
- [partialsMainBuffer, partialsAuxBuffer, resultMainBuffer, resultAuxBuffer],
87
+ [partialsHiBuffer, partialsLoBuffer, resultHiBuffer, resultLoBuffer],
82
88
  );
83
89
  const { commandEncoder: enc2, ts: ts2 } = runComputePass(
84
90
  pipelineReduce,
85
91
  bgReduce,
86
92
  1,
87
93
  ); // dispatch 1 workgroup to reduce the partial sums to a single result
88
- readMainBuffer = stageReadback(enc2, resultMainBuffer);
89
- readAuxBuffer = stageReadback(enc2, resultAuxBuffer);
94
+ readHiBuffer = stageReadback(enc2, resultHiBuffer);
95
+ readLoBuffer = stageReadback(enc2, resultLoBuffer);
90
96
 
91
97
  submit(enc2);
92
98
 
93
- const mainPromise = extractResult(readMainBuffer, Float32Array);
94
- const auxPromise = extractResult(readAuxBuffer, Uint32Array);
95
- readMainBuffer = null; // ownership transferred — extractResult's own finally destroys it
96
- readAuxBuffer = null;
99
+ const hiPromise = extractResult(readHiBuffer, Float32Array);
100
+ const loPromise = extractResult(readLoBuffer, Float32Array);
101
+ readHiBuffer = null; // ownership transferred — extractResult's own finally destroys it
102
+ readLoBuffer = null;
97
103
 
98
- const [gpuTime1, gpuTime2, mainArr, auxArr] = await Promise.all([
104
+ const [gpuTime1, gpuTime2, hiArr, loArr] = await Promise.all([
99
105
  extractTimestamp(ts1),
100
106
  extractTimestamp(ts2),
101
- mainPromise,
102
- auxPromise,
107
+ hiPromise,
108
+ loPromise,
103
109
  ]);
104
110
 
105
111
  // asum is always a scalar readback — both paths return { asum }
106
- const asum = unpackF64(mainArr[0], auxArr[0]);
112
+ const asum = mergeDoubleDouble(hiArr, loArr)[0];
107
113
  if (gpuTime1 !== undefined && gpuTime2 !== undefined)
108
114
  return { asum, gpuTimeMs: gpuTime1 + gpuTime2 };
109
115
  return { asum };
110
116
  } finally {
111
- if (!xIsGpu && xVec) xVec.destroy();
112
- if (partialsMainBuffer) destroyBuffers(partialsMainBuffer);
113
- if (partialsAuxBuffer) destroyBuffers(partialsAuxBuffer);
114
- if (resultMainBuffer) destroyBuffers(resultMainBuffer);
115
- if (resultAuxBuffer) destroyBuffers(resultAuxBuffer);
117
+ if (!xIsGpu && xHiBuffer) destroyBuffers(xHiBuffer);
118
+ if (!xIsGpu && xLoBuffer) destroyBuffers(xLoBuffer);
119
+ if (partialsHiBuffer) destroyBuffers(partialsHiBuffer);
120
+ if (partialsLoBuffer) destroyBuffers(partialsLoBuffer);
121
+ if (resultHiBuffer) destroyBuffers(resultHiBuffer);
122
+ if (resultLoBuffer) destroyBuffers(resultLoBuffer);
116
123
  if (paramsBuffer) destroyBuffers(paramsBuffer);
117
124
  // Only reached if submit(enc2) threw before ownership was transferred above.
118
- if (readMainBuffer) destroyBuffers(readMainBuffer);
119
- if (readAuxBuffer) destroyBuffers(readAuxBuffer);
125
+ if (readHiBuffer) destroyBuffers(readHiBuffer);
126
+ if (readLoBuffer) destroyBuffers(readLoBuffer);
120
127
  }
121
128
  }
@@ -19,6 +19,7 @@ import sger from "./sger.wgsl";
19
19
  import ssyr from "./ssyr.wgsl";
20
20
  import ssyr2 from "./ssyr2.wgsl";
21
21
  import f64add from "./f64add.wgsl";
22
+ import dekker from "./f64/dekker.wgsl";
22
23
  import dasum from "./dasum.wgsl";
23
24
  import strsv_invert_block from "./strsv_invert_block.wgsl";
24
25
  import strsv_apply_inverse from "./strsv_apply_inverse.wgsl";
@@ -46,6 +47,7 @@ export const shaderSources = {
46
47
  ssyr,
47
48
  ssyr2,
48
49
  f64add,
50
+ "f64/dekker": dekker,
49
51
  dasum,
50
52
  strsv_invert_block,
51
53
  strsv_apply_inverse,
@@ -1,36 +1,12 @@
1
- // dasum: result = sum(|x[i]|)
2
- // pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sumF64.wgsl.
3
- // Same structure as sasum.wgsl — every value is now a [main, aux] pair
4
- // (see src/util/f64pack.mjs) and every `+`/`+=` is computeSum via addPair
5
- // instead of plain f32 addition. Concatenated after f64add.wgsl by
6
- // getPipeline (WGSL has no #include), reusing its decode/encode/computeSum/
7
- // addFields and Packed struct — f64add.wgsl declares no bindings and no entry
8
- // point of its own (just helper functions), so bindings here start at 0 and
9
- // the entry point is simply `dasum_main`.
10
- //
11
- // xAux/partialsAux are array<u32>, not array<f32> — aux's bits must never
12
- // pass through an f32-typed storage slot (NaN-bit-pattern corruption risk,
13
- // see f64pack.mjs and the Packed struct comment above decode()/encode() in
14
- // f64add.wgsl); Packed (from f64add.wgsl) keeps aux as u32 in registers/
15
- // workgroup memory too.
16
- //
17
- // Per-thread accumulation (acc0..acc3) stays in DECODED Fields form for the
18
- // entire strided loop below, via addFields — not re-encoded to Packed and
19
- // re-decoded on every single element like a naive version would. Only the
20
- // freshly-loaded x[idx] needs decoding each iteration (unavoidable, it's new
21
- // data every time); the running total never leaves Fields form until the
22
- // four accumulators are combined and encoded exactly once, right before
23
- // writing into workgroup-shared `tile`. The cross-thread reduction tree
24
- // after that still goes through Packed per level (unavoidable — each level
25
- // combines values that live in different threads' registers via shared
26
- // memory), but that's a fixed 6 levels regardless of n, unlike the strided
27
- // loop above whose iteration count scales with n.
1
+ // dasum: sum(|x[i]|), double-double (Dekker). Same ILP=4 shape as sasum.wgsl;
2
+ // see f64/dekker.wgsl for ddAddProtected and why plain ddAdd isn't safe.
3
+ // GpuVector input isn't pre-abs'd, so ddAbs() applies unconditionally below.
28
4
 
29
- @group(0) @binding(0) var<storage, read> xMain: array<f32>;
30
- @group(0) @binding(1) var<storage, read> xAux: array<u32>;
31
- @group(0) @binding(2) var<storage, read_write> partialsMain: array<f32>;
32
- @group(0) @binding(3) var<storage, read_write> partialsAux: array<u32>;
33
- @group(0) @binding(4) var<uniform> params: Params;
5
+ @group(0) @binding(0) var<storage, read> xHi: array<f32>;
6
+ @group(0) @binding(1) var<storage, read> xLo: array<f32>;
7
+ @group(0) @binding(2) var<storage, read_write> partialsHi: array<f32>;
8
+ @group(0) @binding(3) var<storage, read_write> partialsLo: array<f32>;
9
+ @group(0) @binding(4) var<uniform> params: Params;
34
10
 
35
11
  struct Params {
36
12
  n: u32,
@@ -39,21 +15,7 @@ struct Params {
39
15
 
40
16
  const WGS: u32 = 64;
41
17
 
42
- var<workgroup> tile: array<Packed, 64>;
43
-
44
- // a + b, where a/b are [main, aux] pairs — computeSum takes decoded Fields.
45
- // Only used for the cross-thread reduction tree below; the per-thread
46
- // strided loop uses addFields directly instead (see module comment).
47
- fn addPair(a: Packed, b: Packed) -> Packed {
48
- return computeSum(decode(bitcast<u32>(a.main), a.aux), decode(bitcast<u32>(b.main), b.aux));
49
- }
50
-
51
- // |x| for a packed double is abs(main) with aux untouched — only main's
52
- // sign bit carries the double's sign (see fieldsToPacked() in f64pack.mjs).
53
- // Returns decoded Fields directly (not Packed) for the per-thread loop.
54
- fn absFields(idx: u32) -> Fields {
55
- return decode(bitcast<u32>(abs(xMain[idx])), xAux[idx]);
56
- }
18
+ var<workgroup> tile: array<DD, 64>;
57
19
 
58
20
  @compute @workgroup_size(64)
59
21
  fn dasum_main(
@@ -62,37 +24,61 @@ fn dasum_main(
62
24
  @builtin(workgroup_id) wgid: vec3u,
63
25
  @builtin(num_workgroups) num_wg: vec3u,
64
26
  ) {
65
- var acc0: Fields = Fields(0u, 0u, 0u, 0u);
66
- var acc1: Fields = Fields(0u, 0u, 0u, 0u);
67
- var acc2: Fields = Fields(0u, 0u, 0u, 0u);
68
- var acc3: Fields = Fields(0u, 0u, 0u, 0u);
27
+ var acc0 = DD(0.0, 0.0);
28
+ var acc1 = DD(0.0, 0.0);
29
+ var acc2 = DD(0.0, 0.0);
30
+ var acc3 = DD(0.0, 0.0);
69
31
 
70
32
  let stride = num_wg.x * WGS;
71
33
  let n4_floor = (params.n / (4u * stride)) * (4u * stride);
72
34
 
73
- for (var id = gid.x; id < n4_floor; id += 4u * stride) {
74
- acc0 = addFields(acc0, absFields( id * params.x_inc));
75
- acc1 = addFields(acc1, absFields((id + stride) * params.x_inc));
76
- acc2 = addFields(acc2, absFields((id + 2u * stride) * params.x_inc));
77
- acc3 = addFields(acc3, absFields((id + 3u * stride) * params.x_inc));
35
+ // Same trip count for every thread, but driven by a counter, not `id`
36
+ // itself (ddAddProtected's barrier needs a provably-uniform loop bound).
37
+ let mainIters = n4_floor / (4u * stride);
38
+ for (var iter = 0u; iter < mainIters; iter++) {
39
+ let id = gid.x + iter * 4u * stride;
40
+ let i0 = id * params.x_inc;
41
+ let i1 = (id + stride) * params.x_inc;
42
+ let i2 = (id + 2u * stride) * params.x_inc;
43
+ let i3 = (id + 3u * stride) * params.x_inc;
44
+ acc0 = ddAddProtected(acc0, ddAbs(DD(xHi[i0], xLo[i0])), lid.x);
45
+ acc1 = ddAddProtected(acc1, ddAbs(DD(xHi[i1], xLo[i1])), lid.x);
46
+ acc2 = ddAddProtected(acc2, ddAbs(DD(xHi[i2], xLo[i2])), lid.x);
47
+ acc3 = ddAddProtected(acc3, ddAbs(DD(xHi[i3], xLo[i3])), lid.x);
48
+ }
49
+
50
+ // Tail is ragged (0-3 extra per thread) — pad to this workgroup's worst case.
51
+ let wgBaseGid = wgid.x * WGS;
52
+ var tailIters = 0u;
53
+ if (n4_floor + wgBaseGid < params.n) {
54
+ tailIters = (params.n - 1u - n4_floor - wgBaseGid) / stride + 1u;
78
55
  }
79
- for (var id = n4_floor + gid.x; id < params.n; id += stride) {
80
- acc0 = addFields(acc0, absFields(id * params.x_inc));
56
+ for (var iter = 0u; iter < tailIters; iter++) {
57
+ let id = n4_floor + gid.x + iter * stride;
58
+ let valid = id < params.n;
59
+ let i = select(0u, id * params.x_inc, valid);
60
+ let loaded = ddAbs(DD(xHi[i], xLo[i])); // select() has no DD overload
61
+ let contribution = DD(select(0.0, loaded.hi, valid), select(0.0, loaded.lo, valid));
62
+ acc0 = ddAddProtected(acc0, contribution, lid.x);
81
63
  }
82
64
 
83
- // Combine the 4 per-thread accumulators in Fields form too — still no
84
- // encode/decode needed, since none of them have touched Packed yet.
85
- let combined = addFields(addFields(acc0, acc1), addFields(acc2, acc3));
86
- tile[lid.x] = encode(combined.sign, combined.rawExp, combined.mantissaHi, combined.lo);
65
+ let combined01 = ddAddProtected(acc0, acc1, lid.x);
66
+ let combined23 = ddAddProtected(acc2, acc3, lid.x);
67
+ tile[lid.x] = ddAddProtected(combined01, combined23, lid.x);
87
68
  workgroupBarrier();
88
69
 
70
+ // Inactive threads combine against a throwaway partner and discard it
71
+ // (ddAddProtected must be called unconditionally by every thread).
89
72
  for (var s = WGS / 2u; s > 0u; s >>= 1u) {
90
- if (lid.x < s) { tile[lid.x] = addPair(tile[lid.x], tile[lid.x + s]); }
73
+ let partner = select(lid.x, lid.x + s, lid.x < s);
74
+ let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
75
+ workgroupBarrier(); // all threads must read tile[] above before any write below
76
+ if (lid.x < s) { tile[lid.x] = combined; }
91
77
  workgroupBarrier();
92
78
  }
93
79
 
94
80
  if (lid.x == 0u) {
95
- partialsMain[wgid.x] = tile[0].main;
96
- partialsAux[wgid.x] = tile[0].aux;
81
+ partialsHi[wgid.x] = tile[0].hi;
82
+ partialsLo[wgid.x] = tile[0].lo;
97
83
  }
98
84
  }
@@ -0,0 +1,99 @@
1
+ // Double-double arithmetic via Dekker's algorithm — an alternative to
2
+ // f64add.wgsl's bit-exact IEEE-754 emulation. Doesn't touch that path.
3
+ //
4
+ // A double-double number is a pair (hi, lo) of f32 with hi+lo approximating
5
+ // a higher-precision value, hi holding the leading bits and lo the rounding
6
+ // error hi lost. ~48 bits of mantissa vs f32's 24, less than real f64's 52.
7
+ //
8
+ // No bindings, no entry point — a helper library, concatenated with a
9
+ // consumer's own bindings/entry point by getPipeline (WGSL has no #include).
10
+
11
+ struct DD {
12
+ hi: f32,
13
+ lo: f32,
14
+ }
15
+
16
+ // |a| for a double-double pair. Negation is exact (no rounding), so this is
17
+ // just a sign flip on both components — hi alone determines the pair's sign.
18
+ fn ddAbs(a: DD) -> DD {
19
+ if (a.hi < 0.0) {
20
+ return DD(-a.hi, -a.lo);
21
+ }
22
+ return a;
23
+ }
24
+
25
+ // ── A real compiler bug — read before touching anything below ──────────────
26
+ //
27
+ // twoSum/fastTwoSum's error term `e` should be nonzero (that's the point —
28
+ // `s` is rounded). A front-end-optimizer bug zeros it anyway, on both NVIDIA
29
+ // and Mesa (ANV + llvmpipe), via two different mechanisms needing two fixes:
30
+ // bitcast-based subtraction (`fsub`/`negf`, fixes NVIDIA) and materializing
31
+ // the sum through workgroup memory + workgroupBarrier() (fixes Mesa). Only
32
+ // both together (ddAddProtected) is verified correct everywhere — the plain
33
+ // twoSum/fastTwoSum/ddAdd below are reference-only, not safe to use.
34
+ fn negf(x: f32) -> f32 {
35
+ return bitcast<f32>(bitcast<u32>(x) ^ 0x80000000u);
36
+ }
37
+ fn fsub(a: f32, b: f32) -> f32 {
38
+ return a + negf(b);
39
+ }
40
+
41
+ // Knuth/Møller's TwoSum: s = fl(a+b), e = exact rounding error, a+b == s+e.
42
+ // Works for any a, b. UNPROTECTED — see header above.
43
+ fn twoSum(a: f32, b: f32) -> DD {
44
+ let s = a + b;
45
+ let v = s - a;
46
+ let e = (a - (s - v)) + (b - v);
47
+ return DD(s, e);
48
+ }
49
+
50
+ // Dekker's Fast-Two-Sum: same contract, but only correct when |a| >= |b|.
51
+ // UNPROTECTED — see header above.
52
+ fn fastTwoSum(a: f32, b: f32) -> DD {
53
+ let s = a + b;
54
+ let e = b - (s - a);
55
+ return DD(s, e);
56
+ }
57
+
58
+ // Double-double addition (Dekker's Add2). UNPROTECTED — see header above.
59
+ fn ddAdd(a: DD, b: DD) -> DD {
60
+ let s = twoSum(a.hi, b.hi);
61
+ let loSum = a.lo + b.lo;
62
+ return fastTwoSum(s.hi, s.lo + loSum);
63
+ }
64
+
65
+ // ── Protected variants — use these ──────────────────────────────────────────
66
+ //
67
+ // Bitcast subtraction + workgroup-barrier materialization, verified correct
68
+ // on all three backends tested. Costs a real barrier: fine for O(1)-per-
69
+ // thread or O(log n) reduction use, not a long per-element loop. A
70
+ // workgroupBarrier() requires uniform control flow, so:
71
+ // - `threadSlot` must be unique per concurrent caller (e.g. local_invocation_index).
72
+ // - Every thread in the workgroup must call this the same number of times
73
+ // — including ones whose result gets discarded. Compute unconditionally;
74
+ // only the write-back should be conditional.
75
+ var<workgroup> dekkerScratch: array<f32, 64>;
76
+
77
+ fn twoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
78
+ dekkerScratch[threadSlot] = a + b;
79
+ workgroupBarrier();
80
+ let s = dekkerScratch[threadSlot];
81
+ let v = fsub(s, a);
82
+ let e = fsub(a, fsub(s, v)) + fsub(b, v);
83
+ return DD(s, e);
84
+ }
85
+
86
+ fn fastTwoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
87
+ dekkerScratch[threadSlot] = a + b;
88
+ workgroupBarrier();
89
+ let s = dekkerScratch[threadSlot];
90
+ let e = fsub(b, fsub(s, a));
91
+ return DD(s, e);
92
+ }
93
+
94
+ // Protected double-double addition — same contract as ddAdd, but exact.
95
+ fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
96
+ let s = twoSumProtected(a.hi, b.hi, threadSlot);
97
+ let loSum = a.lo + b.lo;
98
+ return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
99
+ }