wgblas 0.1.2 → 1.0.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 +1078 -37
- package/index.d.mts +6 -0
- package/index.mjs +6 -0
- package/package.json +32 -1
- package/src/classes/GpuMatrix.d.mts +85 -0
- package/src/classes/GpuMatrix.mjs +91 -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 +118 -0
- package/src/sgemv/sgemv.mjs +141 -0
- package/src/shaders/browser-shaders.mjs +20 -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/snrm2.wgsl +22 -4
- package/src/shaders/ssymv.wgsl +69 -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 +109 -0
- package/src/ssymv/ssymv.mjs +130 -0
- package/src/strmv/strmv.d.mts +109 -0
- package/src/strmv/strmv.mjs +132 -0
- package/src/strsv/strsv.d.mts +98 -0
- package/src/strsv/strsv.mjs +212 -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,118 @@
|
|
|
1
|
+
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
|
+
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
3
|
+
|
|
4
|
+
/**
|
|
5
|
+
* Performs the matrix-vector operation y = alpha * op(A) * x + beta * y
|
|
6
|
+
*
|
|
7
|
+
* - `trans='no-transpose'`: op(A) = A, x is length n, y is length m
|
|
8
|
+
* - `trans='transpose'`: op(A) = A^T, x is length m, y is length n
|
|
9
|
+
*
|
|
10
|
+
* A is an m×n matrix stored in row-major order. `lda` is the leading dimension
|
|
11
|
+
* (number of floats between the start of consecutive rows — must be >= n).
|
|
12
|
+
*
|
|
13
|
+
* {@includeCode ../../examples/sgemv/sgemv.js}
|
|
14
|
+
*
|
|
15
|
+
* **Browser (standalone HTML):**
|
|
16
|
+
* {@includeCode ../../examples/sgemv/web/sgemv.html}
|
|
17
|
+
*
|
|
18
|
+
* @param device - GPUDevice from `init()`
|
|
19
|
+
* @param trans - `'no-transpose'` for A, `'transpose'` for A^T
|
|
20
|
+
* @param m - number of rows in A
|
|
21
|
+
* @param n - number of columns in A
|
|
22
|
+
* @param alpha - scalar multiplier for op(A)*x
|
|
23
|
+
* @param A - Float32Array or GpuMatrix, row-major, at least (m-1)*lda+n elements
|
|
24
|
+
* @param lda - leading dimension of A (>= n)
|
|
25
|
+
* @param x - Float32Array input vector
|
|
26
|
+
* @param incx - stride for x (must be a positive integer)
|
|
27
|
+
* @param beta - scalar multiplier for y
|
|
28
|
+
* @param y - Float32Array input/output vector
|
|
29
|
+
* @param incy - stride for y (must be a positive integer)
|
|
30
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemv/sgemv.mjs#L16">Source code: sgemv.mjs (L16)</a>
|
|
31
|
+
* @category BLAS Level 2
|
|
32
|
+
*/
|
|
33
|
+
export declare function sgemv(
|
|
34
|
+
device: GPUDevice,
|
|
35
|
+
trans: 'no-transpose' | 'transpose',
|
|
36
|
+
m: number,
|
|
37
|
+
n: number,
|
|
38
|
+
alpha: number,
|
|
39
|
+
A: Float32Array | GpuMatrix,
|
|
40
|
+
lda: number,
|
|
41
|
+
x: Float32Array,
|
|
42
|
+
incx: number,
|
|
43
|
+
beta: number,
|
|
44
|
+
y: Float32Array,
|
|
45
|
+
incy: number,
|
|
46
|
+
): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
|
|
47
|
+
|
|
48
|
+
/**
|
|
49
|
+
* Performs the matrix-vector operation y = alpha * op(A) * x + beta * y
|
|
50
|
+
*
|
|
51
|
+
* A is kept GPU-resident; x and y are CPU Float32Arrays.
|
|
52
|
+
*
|
|
53
|
+
* @param device - GPUDevice from `init()`
|
|
54
|
+
* @param trans - `'no-transpose'` for A, `'transpose'` for A^T
|
|
55
|
+
* @param m - number of rows in A
|
|
56
|
+
* @param n - number of columns in A
|
|
57
|
+
* @param alpha - scalar multiplier for op(A)*x
|
|
58
|
+
* @param A - GpuMatrix, row-major, GPU-resident
|
|
59
|
+
* @param lda - leading dimension of A (must equal A.lda)
|
|
60
|
+
* @param x - Float32Array input vector
|
|
61
|
+
* @param incx - stride for x (must be a positive integer)
|
|
62
|
+
* @param beta - scalar multiplier for y
|
|
63
|
+
* @param y - Float32Array input/output vector
|
|
64
|
+
* @param incy - stride for y (must be a positive integer)
|
|
65
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemv/sgemv.mjs#L16">Source code: sgemv.mjs (L16)</a>
|
|
66
|
+
* @category BLAS Level 2
|
|
67
|
+
*/
|
|
68
|
+
export declare function sgemv(
|
|
69
|
+
device: GPUDevice,
|
|
70
|
+
trans: 'no-transpose' | 'transpose',
|
|
71
|
+
m: number,
|
|
72
|
+
n: number,
|
|
73
|
+
alpha: number,
|
|
74
|
+
A: GpuMatrix,
|
|
75
|
+
lda: number,
|
|
76
|
+
x: Float32Array,
|
|
77
|
+
incx: number,
|
|
78
|
+
beta: number,
|
|
79
|
+
y: Float32Array,
|
|
80
|
+
incy: number,
|
|
81
|
+
): Promise<{ y: Float32Array; gpuTimeMs?: number }>;
|
|
82
|
+
|
|
83
|
+
/**
|
|
84
|
+
* Performs the matrix-vector operation y = alpha * op(A) * x + beta * y
|
|
85
|
+
*
|
|
86
|
+
* x and y are kept resident on the GPU. A must be a GpuMatrix.
|
|
87
|
+
*
|
|
88
|
+
* {@includeCode ../../examples/sgemv/gpuvec.sgemv.js}
|
|
89
|
+
*
|
|
90
|
+
* @param device - GPUDevice from `init()`
|
|
91
|
+
* @param trans - `'no-transpose'` for A, `'transpose'` for A^T
|
|
92
|
+
* @param m - number of rows in A
|
|
93
|
+
* @param n - number of columns in A
|
|
94
|
+
* @param alpha - scalar multiplier for op(A)*x
|
|
95
|
+
* @param A - GpuMatrix, row-major
|
|
96
|
+
* @param lda - leading dimension of A (>= n)
|
|
97
|
+
* @param x - GpuVector input vector (not mutated)
|
|
98
|
+
* @param incx - stride for x (must be a positive integer)
|
|
99
|
+
* @param beta - scalar multiplier for y
|
|
100
|
+
* @param y - GpuVector input/output vector (mutated in place)
|
|
101
|
+
* @param incy - stride for y (must be a positive integer)
|
|
102
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/sgemv/sgemv.mjs#L16">Source code: sgemv.mjs (L16)</a>
|
|
103
|
+
* @category BLAS Level 2
|
|
104
|
+
*/
|
|
105
|
+
export declare function sgemv(
|
|
106
|
+
device: GPUDevice,
|
|
107
|
+
trans: 'no-transpose' | 'transpose',
|
|
108
|
+
m: number,
|
|
109
|
+
n: number,
|
|
110
|
+
alpha: number,
|
|
111
|
+
A: GpuMatrix,
|
|
112
|
+
lda: number,
|
|
113
|
+
x: GpuVector,
|
|
114
|
+
incx: number,
|
|
115
|
+
beta: number,
|
|
116
|
+
y: GpuVector,
|
|
117
|
+
incy: number,
|
|
118
|
+
): Promise<{ gpuTimeMs?: number }>;
|
|
@@ -0,0 +1,141 @@
|
|
|
1
|
+
import {
|
|
2
|
+
uploadBuffer,
|
|
3
|
+
createParamsBuffer,
|
|
4
|
+
stageReadback,
|
|
5
|
+
destroyBuffers,
|
|
6
|
+
} from "../util/buffer.mjs";
|
|
7
|
+
import { createBindGroup } from "../util/bindgroup.mjs";
|
|
8
|
+
import { runComputePass, submit } from "../util/compute.mjs";
|
|
9
|
+
import { extractResult } from "../util/result.mjs";
|
|
10
|
+
import { extractTimestamp } from "../util/benchmark.mjs";
|
|
11
|
+
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
|
+
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
|
+
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
15
|
+
|
|
16
|
+
export async function sgemv(device, trans, m, n, alpha, A, lda, x, incx, beta, y, incy) {
|
|
17
|
+
const xIsGpu = x instanceof GpuVector;
|
|
18
|
+
const yIsGpu = y instanceof GpuVector;
|
|
19
|
+
const AIsGpu = A instanceof GpuMatrix;
|
|
20
|
+
const isNoTrans = trans === "no-transpose";
|
|
21
|
+
|
|
22
|
+
if (!(device instanceof GPUDevice))
|
|
23
|
+
throw new Error("device must be a GPUDevice.");
|
|
24
|
+
if (!isNoTrans && trans !== "transpose")
|
|
25
|
+
throw new Error("trans must be 'no-transpose' or 'transpose'.");
|
|
26
|
+
if (
|
|
27
|
+
!Number.isInteger(m) ||
|
|
28
|
+
!Number.isInteger(n) ||
|
|
29
|
+
!Number.isInteger(incx) ||
|
|
30
|
+
!Number.isInteger(incy) ||
|
|
31
|
+
!Number.isInteger(lda)
|
|
32
|
+
)
|
|
33
|
+
throw new Error("m, n, incx, incy, and lda must be integers.");
|
|
34
|
+
if (typeof alpha !== "number")
|
|
35
|
+
throw new Error("alpha must be a number.");
|
|
36
|
+
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
37
|
+
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
38
|
+
if (typeof beta !== "number")
|
|
39
|
+
throw new Error("beta must be a number.");
|
|
40
|
+
if (Number.isNaN(beta)) throw new Error("beta must not be NaN.");
|
|
41
|
+
if (!Number.isFinite(beta)) throw new Error("beta must be finite.");
|
|
42
|
+
if (incx <= 0 || incy <= 0)
|
|
43
|
+
throw new Error("incx and incy must be positive.");
|
|
44
|
+
if (lda < n) throw new Error("lda must be >= n.");
|
|
45
|
+
if (!AIsGpu && !(A instanceof Float32Array))
|
|
46
|
+
throw new Error("A must be a Float32Array or GpuMatrix.");
|
|
47
|
+
if (!xIsGpu && !(x instanceof Float32Array))
|
|
48
|
+
throw new Error("x must be a Float32Array or GpuVector.");
|
|
49
|
+
if (!yIsGpu && !(y instanceof Float32Array))
|
|
50
|
+
throw new Error("y must be a Float32Array or GpuVector.");
|
|
51
|
+
if (xIsGpu !== yIsGpu)
|
|
52
|
+
throw new Error(
|
|
53
|
+
"x and y must be the same type (both Float32Array or both GpuVector).",
|
|
54
|
+
);
|
|
55
|
+
if (xIsGpu && !AIsGpu)
|
|
56
|
+
throw new Error(
|
|
57
|
+
"A must be a GpuMatrix when x and y are GpuVectors.",
|
|
58
|
+
);
|
|
59
|
+
if (xIsGpu && x._buf === y._buf)
|
|
60
|
+
throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");
|
|
61
|
+
if (AIsGpu && lda !== A.lda)
|
|
62
|
+
throw new Error("lda must match A.lda when A is a GpuMatrix.");
|
|
63
|
+
if (AIsGpu && (A.rows < m || A.cols < n))
|
|
64
|
+
throw new Error("A is too small for the given m and n.");
|
|
65
|
+
if (m < 0 || n < 0) throw new Error("m and n must be non-negative.");
|
|
66
|
+
if (m === 0 || n === 0) return yIsGpu ? {} : { y };
|
|
67
|
+
|
|
68
|
+
// NoTrans: x has n elements, y has m elements
|
|
69
|
+
// Trans: x has m elements, y has n elements
|
|
70
|
+
const xLen = isNoTrans ? n : m;
|
|
71
|
+
const yLen = isNoTrans ? m : n;
|
|
72
|
+
|
|
73
|
+
if (!AIsGpu && A.length < (m - 1) * lda + n)
|
|
74
|
+
throw new Error(
|
|
75
|
+
"A does not have enough elements for the given m, n, and lda.",
|
|
76
|
+
);
|
|
77
|
+
if (x.length < (xLen - 1) * incx + 1)
|
|
78
|
+
throw new Error(
|
|
79
|
+
"x does not have enough elements for the given dimensions and incx.",
|
|
80
|
+
);
|
|
81
|
+
if (y.length < (yLen - 1) * incy + 1)
|
|
82
|
+
throw new Error(
|
|
83
|
+
"y does not have enough elements for the given dimensions and incy.",
|
|
84
|
+
);
|
|
85
|
+
|
|
86
|
+
const shaderName = isNoTrans ? "sgemv_n" : "sgemv_t";
|
|
87
|
+
const pipeline = await getPipeline(device, shaderName);
|
|
88
|
+
|
|
89
|
+
const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "sgemv-A", false);
|
|
90
|
+
const xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "sgemv-x", false);
|
|
91
|
+
const yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "sgemv-y", true);
|
|
92
|
+
const paramsBuffer = createParamsBuffer(
|
|
93
|
+
[
|
|
94
|
+
{ value: m, type: "u32" },
|
|
95
|
+
{ value: n, type: "u32" },
|
|
96
|
+
{ value: alpha, type: "f32" },
|
|
97
|
+
{ value: beta, type: "f32" },
|
|
98
|
+
{ value: incx, type: "u32" },
|
|
99
|
+
{ value: incy, type: "u32" },
|
|
100
|
+
{ value: lda, type: "u32" },
|
|
101
|
+
],
|
|
102
|
+
"sgemv-params",
|
|
103
|
+
);
|
|
104
|
+
|
|
105
|
+
try {
|
|
106
|
+
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
107
|
+
ABuffer,
|
|
108
|
+
xBuffer,
|
|
109
|
+
yBuffer,
|
|
110
|
+
paramsBuffer,
|
|
111
|
+
]);
|
|
112
|
+
|
|
113
|
+
// NoTrans: one workgroup per row; clamped to device limit — the shader's
|
|
114
|
+
// grid-stride loop handles remaining rows when m > dispatch count.
|
|
115
|
+
// Trans: one thread per output column → dispatch ceil(n/64)
|
|
116
|
+
const wgCount = isNoTrans
|
|
117
|
+
? Math.min(m, device.limits.maxComputeWorkgroupsPerDimension)
|
|
118
|
+
: calcWorkgroups(yLen);
|
|
119
|
+
const { commandEncoder, ts } = runComputePass(pipeline, bindGroup, wgCount);
|
|
120
|
+
const readBuffer = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
121
|
+
|
|
122
|
+
submit(commandEncoder);
|
|
123
|
+
|
|
124
|
+
const gpuTimeMs = await extractTimestamp(ts);
|
|
125
|
+
|
|
126
|
+
if (yIsGpu) {
|
|
127
|
+
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
128
|
+
return {};
|
|
129
|
+
}
|
|
130
|
+
|
|
131
|
+
const result = await extractResult(readBuffer, Float32Array);
|
|
132
|
+
if (gpuTimeMs !== undefined) return { y: result, gpuTimeMs };
|
|
133
|
+
return { y: result };
|
|
134
|
+
} finally {
|
|
135
|
+
if (!AIsGpu) destroyBuffers(ABuffer);
|
|
136
|
+
if (!xIsGpu) destroyBuffers(xBuffer);
|
|
137
|
+
if (!yIsGpu) destroyBuffers(yBuffer);
|
|
138
|
+
destroyBuffers(paramsBuffer);
|
|
139
|
+
|
|
140
|
+
}
|
|
141
|
+
}
|
|
@@ -1,5 +1,6 @@
|
|
|
1
1
|
import argmax from "./reduction/argmax.wgsl";
|
|
2
2
|
import sum from "./reduction/sum.wgsl";
|
|
3
|
+
import sumF64 from "./reduction/sumF64.wgsl";
|
|
3
4
|
import sscal from "./sscal.wgsl";
|
|
4
5
|
import sswap from "./sswap.wgsl";
|
|
5
6
|
import saxpy from "./saxpy.wgsl";
|
|
@@ -10,10 +11,20 @@ import snrm2 from "./snrm2.wgsl";
|
|
|
10
11
|
import srot from "./srot.wgsl";
|
|
11
12
|
import srotm from "./srotm.wgsl";
|
|
12
13
|
import isamax from "./isamax.wgsl";
|
|
14
|
+
import sgemv_n from "./sgemv_n.wgsl";
|
|
15
|
+
import sgemv_t from "./sgemv_t.wgsl";
|
|
16
|
+
import ssymv from "./ssymv.wgsl";
|
|
17
|
+
import strmv from "./strmv.wgsl";
|
|
18
|
+
import f64add from "./f64add.wgsl";
|
|
19
|
+
import dasum from "./dasum.wgsl";
|
|
20
|
+
import strsv_invert_block from "./strsv_invert_block.wgsl";
|
|
21
|
+
import strsv_apply_inverse from "./strsv_apply_inverse.wgsl";
|
|
22
|
+
import strsv_update from "./strsv_update.wgsl";
|
|
13
23
|
|
|
14
24
|
export const shaderSources = {
|
|
15
25
|
"reduction/argmax": argmax,
|
|
16
26
|
"reduction/sum": sum,
|
|
27
|
+
"reduction/sumF64": sumF64,
|
|
17
28
|
sscal,
|
|
18
29
|
sswap,
|
|
19
30
|
saxpy,
|
|
@@ -24,4 +35,13 @@ export const shaderSources = {
|
|
|
24
35
|
srot,
|
|
25
36
|
srotm,
|
|
26
37
|
isamax,
|
|
38
|
+
sgemv_n,
|
|
39
|
+
sgemv_t,
|
|
40
|
+
ssymv,
|
|
41
|
+
strmv,
|
|
42
|
+
f64add,
|
|
43
|
+
dasum,
|
|
44
|
+
strsv_invert_block,
|
|
45
|
+
strsv_apply_inverse,
|
|
46
|
+
strsv_update,
|
|
27
47
|
};
|
|
@@ -0,0 +1,98 @@
|
|
|
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.
|
|
28
|
+
|
|
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;
|
|
34
|
+
|
|
35
|
+
struct Params {
|
|
36
|
+
n: u32,
|
|
37
|
+
x_inc: u32,
|
|
38
|
+
}
|
|
39
|
+
|
|
40
|
+
const WGS: u32 = 64;
|
|
41
|
+
|
|
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
|
+
}
|
|
57
|
+
|
|
58
|
+
@compute @workgroup_size(64)
|
|
59
|
+
fn dasum_main(
|
|
60
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
61
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
62
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
63
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
64
|
+
) {
|
|
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);
|
|
69
|
+
|
|
70
|
+
let stride = num_wg.x * WGS;
|
|
71
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
72
|
+
|
|
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));
|
|
78
|
+
}
|
|
79
|
+
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
80
|
+
acc0 = addFields(acc0, absFields(id * params.x_inc));
|
|
81
|
+
}
|
|
82
|
+
|
|
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);
|
|
87
|
+
workgroupBarrier();
|
|
88
|
+
|
|
89
|
+
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]); }
|
|
91
|
+
workgroupBarrier();
|
|
92
|
+
}
|
|
93
|
+
|
|
94
|
+
if (lid.x == 0u) {
|
|
95
|
+
partialsMain[wgid.x] = tile[0].main;
|
|
96
|
+
partialsAux[wgid.x] = tile[0].aux;
|
|
97
|
+
}
|
|
98
|
+
}
|
|
@@ -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
|
+
}
|