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
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
|
@@ -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
|
|
36
|
-
* has no f64 type) and stored across two GPU
|
|
37
|
-
* reassembles
|
|
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 {
|
|
4
|
+
import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
|
|
5
5
|
|
|
6
6
|
export class GpuMatrix {
|
|
7
|
-
constructor(buffer, rows, cols, lda,
|
|
7
|
+
constructor(buffer, rows, cols, lda, loBuffer = null, layout = "row-major") {
|
|
8
8
|
this._buf = buffer;
|
|
9
|
-
// Non-null only for Float64Array-backed matrices
|
|
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
|
|
54
|
-
const
|
|
55
|
-
|
|
56
|
-
|
|
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.
|
|
80
|
-
const
|
|
81
|
-
const
|
|
82
|
-
device.queue.submit([
|
|
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 [
|
|
75
|
+
const [hi, lo] = await Promise.all([
|
|
85
76
|
extractResult(rb, Float32Array),
|
|
86
|
-
extractResult(
|
|
77
|
+
extractResult(rbLo, Float32Array),
|
|
87
78
|
]);
|
|
88
|
-
const raw =
|
|
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.
|
|
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
|
-
*
|
|
21
|
-
* two GPU buffers internally; `read()`
|
|
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 {
|
|
4
|
+
import { splitDoubleDouble, mergeDoubleDouble } from "../util/f64.mjs";
|
|
5
5
|
|
|
6
6
|
export class GpuVector {
|
|
7
|
-
constructor(buffer, length, dtype = Float32Array,
|
|
7
|
+
constructor(buffer, length, dtype = Float32Array, loBuffer = null) {
|
|
8
8
|
this._buf = buffer;
|
|
9
|
-
this.
|
|
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
|
|
17
|
-
const
|
|
18
|
-
|
|
19
|
-
|
|
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.
|
|
34
|
+
if (!this._loBuf) return extractResult(rb, this.dtype);
|
|
41
35
|
|
|
42
|
-
const
|
|
43
|
-
const
|
|
44
|
-
device.queue.submit([
|
|
36
|
+
const encLo = device.createCommandEncoder();
|
|
37
|
+
const rbLo = stageReadback(encLo, this._loBuf);
|
|
38
|
+
device.queue.submit([encLo.finish()]);
|
|
45
39
|
|
|
46
|
-
const [
|
|
40
|
+
const [hi, lo] = await Promise.all([
|
|
47
41
|
extractResult(rb, Float32Array),
|
|
48
|
-
extractResult(
|
|
42
|
+
extractResult(rbLo, Float32Array),
|
|
49
43
|
]);
|
|
50
|
-
|
|
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.
|
|
49
|
+
if (this._loBuf) this._loBuf.destroy();
|
|
58
50
|
}
|
|
59
51
|
}
|
package/src/dasum/dasum.d.mts
CHANGED
|
@@ -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
|
|
5
|
-
* precision: result = sum(|x[i]|). Each element of `x`
|
|
6
|
-
*
|
|
7
|
-
*
|
|
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
|
*
|
package/src/dasum/dasum.mjs
CHANGED
|
@@ -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
|
-
//
|
|
37
|
-
//
|
|
38
|
-
|
|
39
|
-
|
|
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
|
|
45
|
-
let
|
|
46
|
-
let
|
|
47
|
-
let
|
|
48
|
-
let
|
|
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
|
|
51
|
-
let
|
|
49
|
+
let readHiBuffer = null;
|
|
50
|
+
let readLoBuffer = null;
|
|
52
51
|
|
|
53
52
|
try {
|
|
54
|
-
|
|
55
|
-
|
|
56
|
-
|
|
57
|
-
|
|
58
|
-
|
|
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
|
-
[
|
|
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
|
-
[
|
|
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
|
-
|
|
89
|
-
|
|
94
|
+
readHiBuffer = stageReadback(enc2, resultHiBuffer);
|
|
95
|
+
readLoBuffer = stageReadback(enc2, resultLoBuffer);
|
|
90
96
|
|
|
91
97
|
submit(enc2);
|
|
92
98
|
|
|
93
|
-
const
|
|
94
|
-
const
|
|
95
|
-
|
|
96
|
-
|
|
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,
|
|
104
|
+
const [gpuTime1, gpuTime2, hiArr, loArr] = await Promise.all([
|
|
99
105
|
extractTimestamp(ts1),
|
|
100
106
|
extractTimestamp(ts2),
|
|
101
|
-
|
|
102
|
-
|
|
107
|
+
hiPromise,
|
|
108
|
+
loPromise,
|
|
103
109
|
]);
|
|
104
110
|
|
|
105
111
|
// asum is always a scalar readback — both paths return { asum }
|
|
106
|
-
const asum =
|
|
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 &&
|
|
112
|
-
if (
|
|
113
|
-
if (
|
|
114
|
-
if (
|
|
115
|
-
if (
|
|
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 (
|
|
119
|
-
if (
|
|
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,
|
package/src/shaders/dasum.wgsl
CHANGED
|
@@ -1,36 +1,12 @@
|
|
|
1
|
-
// dasum:
|
|
2
|
-
//
|
|
3
|
-
//
|
|
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>
|
|
30
|
-
@group(0) @binding(1) var<storage, read>
|
|
31
|
-
@group(0) @binding(2) var<storage, read_write>
|
|
32
|
-
@group(0) @binding(3) var<storage, read_write>
|
|
33
|
-
@group(0) @binding(4) var<uniform> 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<
|
|
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
|
|
66
|
-
var acc1
|
|
67
|
-
var acc2
|
|
68
|
-
var acc3
|
|
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
|
-
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
|
|
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
|
|
80
|
-
|
|
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
|
-
|
|
84
|
-
|
|
85
|
-
|
|
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
|
-
|
|
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
|
-
|
|
96
|
-
|
|
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
|
+
}
|