wgblas 2.1.0 → 2.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 +2 -0
- package/dist/wgblas.browser.js +1007 -43
- package/index.d.mts +26 -53
- package/index.mjs +11 -0
- package/package.json +132 -63
- package/src/classes/Complex32.d.mts +43 -0
- package/src/classes/Complex32.mjs +82 -0
- package/src/classes/Complex64.d.mts +44 -0
- package/src/classes/Complex64.mjs +76 -0
- package/src/classes/GpuMatrix.d.mts +41 -41
- package/src/classes/GpuMatrix.mjs +112 -10
- package/src/classes/GpuVector.d.mts +36 -40
- package/src/classes/GpuVector.mjs +39 -2
- package/src/cscal/cscal.d.mts +47 -0
- package/src/cscal/cscal.mjs +98 -0
- package/src/dasum/dasum.mjs +31 -15
- package/src/daxpy/daxpy.d.mts +56 -0
- package/src/daxpy/daxpy.mjs +150 -0
- package/src/dcopy/dcopy.d.mts +52 -0
- package/src/dcopy/dcopy.mjs +140 -0
- package/src/ddot/ddot.d.mts +62 -0
- package/src/ddot/ddot.mjs +184 -0
- package/src/dnrm2/dnrm2.d.mts +50 -0
- package/src/dnrm2/dnrm2.mjs +189 -0
- package/src/drot/drot.d.mts +67 -0
- package/src/drot/drot.mjs +170 -0
- package/src/drotm/drotm.d.mts +67 -0
- package/src/drotm/drotm.mjs +171 -0
- package/src/dscal/dscal.d.mts +52 -0
- package/src/dscal/dscal.mjs +119 -0
- package/src/dswap/dswap.d.mts +57 -0
- package/src/dswap/dswap.mjs +155 -0
- package/src/idamax/idamax.mjs +49 -19
- package/src/init.mjs +6 -3
- package/src/isamax/isamax.mjs +17 -14
- package/src/random/random.d.mts +37 -40
- package/src/random/random.mjs +39 -7
- package/src/sasum/sasum.mjs +13 -11
- package/src/saxpy/saxpy.mjs +9 -8
- package/src/scopy/scopy.mjs +8 -6
- package/src/sdot/sdot.mjs +13 -11
- package/src/sgemm/sgemm.d.mts +2 -2
- package/src/sgemm/sgemm.mjs +91 -35
- package/src/sgemmtr/sgemmtr.d.mts +3 -2
- package/src/sgemmtr/sgemmtr.mjs +92 -35
- package/src/sgemv/sgemv.d.mts +2 -2
- package/src/sgemv/sgemv.mjs +41 -25
- package/src/sger/sger.d.mts +2 -2
- package/src/sger/sger.mjs +38 -16
- package/src/shaders/__test_pipeline_a.wgsl +3 -0
- package/src/shaders/__test_pipeline_b.wgsl +2 -0
- package/src/shaders/cscal.wgsl +33 -0
- package/src/shaders/daxpy.wgsl +66 -0
- package/src/shaders/dcopy.wgsl +34 -0
- package/src/shaders/ddot.wgsl +106 -0
- package/src/shaders/dnrm2.wgsl +167 -0
- package/src/shaders/drot.wgsl +81 -0
- package/src/shaders/drotm.wgsl +99 -0
- package/src/shaders/dscal.wgsl +60 -0
- package/src/shaders/dswap.wgsl +38 -0
- package/src/shaders/f64/utils/add.wgsl +6 -0
- package/src/shaders/f64/utils/divide.wgsl +45 -0
- package/src/shaders/f64/utils/multiply.wgsl +19 -10
- package/src/shaders/f64/utils/sqrt.wgsl +45 -0
- package/src/shaders/index.mjs +69 -0
- package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
- package/src/snrm2/snrm2.mjs +20 -14
- package/src/srot/srot.mjs +10 -7
- package/src/srotm/srotm.mjs +9 -12
- package/src/sscal/sscal.mjs +7 -7
- package/src/sswap/sswap.mjs +14 -8
- package/src/ssymm/ssymm.d.mts +5 -4
- package/src/ssymm/ssymm.mjs +142 -55
- package/src/ssymv/ssymv.d.mts +2 -2
- package/src/ssymv/ssymv.mjs +42 -23
- package/src/ssyr/ssyr.d.mts +2 -2
- package/src/ssyr/ssyr.mjs +34 -15
- package/src/ssyr2/ssyr2.d.mts +2 -2
- package/src/ssyr2/ssyr2.mjs +43 -18
- package/src/ssyr2k/ssyr2k.d.mts +3 -2
- package/src/ssyr2k/ssyr2k.mjs +132 -55
- package/src/ssyrk/ssyrk.d.mts +3 -2
- package/src/ssyrk/ssyrk.mjs +84 -33
- package/src/strmm/strmm.d.mts +5 -4
- package/src/strmm/strmm.mjs +153 -54
- package/src/strmv/strmv.d.mts +2 -2
- package/src/strmv/strmv.mjs +37 -17
- package/src/strsm/strsm.d.mts +6 -4
- package/src/strsm/strsm.mjs +418 -172
- package/src/strsv/strsv.d.mts +5 -3
- package/src/strsv/strsv.mjs +82 -31
- package/src/util/benchmark.mjs +5 -3
- package/src/util/buffer.mjs +33 -12
- package/src/util/complex.mjs +87 -0
- package/src/util/compute.mjs +14 -8
- package/src/util/device.mjs +18 -3
- package/src/util/pipeline.mjs +40 -5
- package/src/util/workgroup.mjs +23 -6
- package/src/shaders/f64add.wgsl +0 -281
package/src/sger/sger.mjs
CHANGED
|
@@ -11,18 +11,28 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
13
13
|
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
14
|
-
import { requireSameDevice } from "../util/device.mjs";
|
|
14
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
15
15
|
|
|
16
|
-
export async function sger(
|
|
16
|
+
export async function sger(
|
|
17
|
+
device,
|
|
18
|
+
m,
|
|
19
|
+
n,
|
|
20
|
+
alpha,
|
|
21
|
+
x,
|
|
22
|
+
incx,
|
|
23
|
+
y,
|
|
24
|
+
incy,
|
|
25
|
+
A,
|
|
26
|
+
lda,
|
|
27
|
+
layout = "row-major",
|
|
28
|
+
) {
|
|
17
29
|
const AIsGpu = A instanceof GpuMatrix;
|
|
18
30
|
|
|
19
|
-
|
|
20
|
-
throw new Error("device must be a GPUDevice.");
|
|
31
|
+
requireGpuDevice(device);
|
|
21
32
|
requireSameDevice(device, "sger", { A, x, y });
|
|
22
33
|
if (layout !== "row-major" && layout !== "column-major")
|
|
23
34
|
throw new Error("layout must be 'row-major' or 'column-major'.");
|
|
24
|
-
if (typeof alpha !== "number")
|
|
25
|
-
throw new Error("alpha must be a number.");
|
|
35
|
+
if (typeof alpha !== "number") throw new Error("alpha must be a number.");
|
|
26
36
|
if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
|
|
27
37
|
if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
|
|
28
38
|
if (
|
|
@@ -79,9 +89,13 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
|
|
|
79
89
|
"A does not have enough elements for the given m, n, and lda.",
|
|
80
90
|
);
|
|
81
91
|
if (x.length < (m - 1) * incx + 1)
|
|
82
|
-
throw new Error(
|
|
92
|
+
throw new Error(
|
|
93
|
+
"x does not have enough elements for the given m and incx.",
|
|
94
|
+
);
|
|
83
95
|
if (y.length < (n - 1) * incy + 1)
|
|
84
|
-
throw new Error(
|
|
96
|
+
throw new Error(
|
|
97
|
+
"y does not have enough elements for the given n and incy.",
|
|
98
|
+
);
|
|
85
99
|
|
|
86
100
|
const pipeline = await getPipeline(device, "sger");
|
|
87
101
|
|
|
@@ -94,14 +108,15 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
|
|
|
94
108
|
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "sger-x", false);
|
|
95
109
|
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "sger-y", false);
|
|
96
110
|
ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "sger-A", true);
|
|
97
|
-
paramsBuffer = createParamsBuffer(
|
|
111
|
+
paramsBuffer = createParamsBuffer(
|
|
112
|
+
device,
|
|
98
113
|
[
|
|
99
|
-
{ value: m,
|
|
100
|
-
{ value: n,
|
|
114
|
+
{ value: m, type: "u32" },
|
|
115
|
+
{ value: n, type: "u32" },
|
|
101
116
|
{ value: alpha, type: "f32" },
|
|
102
|
-
{ value: incx,
|
|
103
|
-
{ value: incy,
|
|
104
|
-
{ value: lda,
|
|
117
|
+
{ value: incx, type: "u32" },
|
|
118
|
+
{ value: incy, type: "u32" },
|
|
119
|
+
{ value: lda, type: "u32" },
|
|
105
120
|
],
|
|
106
121
|
"sger-params",
|
|
107
122
|
);
|
|
@@ -116,8 +131,15 @@ export async function sger(device, m, n, alpha, x, incx, y, incy, A, lda, layout
|
|
|
116
131
|
// One workgroup per row of A; clamped to device limit — the shader's
|
|
117
132
|
// grid-stride loop handles remaining rows when m > dispatch count.
|
|
118
133
|
const wgCount = Math.min(m, device.limits.maxComputeWorkgroupsPerDimension);
|
|
119
|
-
const { commandEncoder, ts } = runComputePass(
|
|
120
|
-
|
|
134
|
+
const { commandEncoder, ts } = runComputePass(
|
|
135
|
+
device,
|
|
136
|
+
pipeline,
|
|
137
|
+
bindGroup,
|
|
138
|
+
wgCount,
|
|
139
|
+
);
|
|
140
|
+
const readBuffer = AIsGpu
|
|
141
|
+
? null
|
|
142
|
+
: stageReadback(device, commandEncoder, ABuffer);
|
|
121
143
|
|
|
122
144
|
submit(device, commandEncoder);
|
|
123
145
|
|
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
// cscal: x := alpha * x, complex. x is one interleaved f32 array
|
|
2
|
+
// (re0, im0, re1, im1, ...), matching Complex32Array/GpuVector's storage
|
|
3
|
+
// (and cuBLAS's cuComplex / stdlib's Complex64Array) — no repacking needed
|
|
4
|
+
// between JS and GPU.
|
|
5
|
+
// (alphaRe + i*alphaIm)(re + i*im) = (alphaRe*re - alphaIm*im) + i*(alphaRe*im + alphaIm*re)
|
|
6
|
+
|
|
7
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
8
|
+
|
|
9
|
+
struct Params {
|
|
10
|
+
n: u32,
|
|
11
|
+
alphaRe: f32,
|
|
12
|
+
alphaIm: f32,
|
|
13
|
+
x_inc: u32,
|
|
14
|
+
}
|
|
15
|
+
|
|
16
|
+
@group(0) @binding(1) var<uniform> params: Params;
|
|
17
|
+
|
|
18
|
+
const WGS: u32 = 64;
|
|
19
|
+
|
|
20
|
+
@compute @workgroup_size(64)
|
|
21
|
+
fn main(
|
|
22
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
23
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
24
|
+
) {
|
|
25
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
26
|
+
let base = 2u * id * params.x_inc;
|
|
27
|
+
// Both new parts need both old parts, so capture them before either write.
|
|
28
|
+
let re = x[base];
|
|
29
|
+
let im = x[base + 1u];
|
|
30
|
+
x[base] = params.alphaRe * re - params.alphaIm * im;
|
|
31
|
+
x[base + 1u] = params.alphaRe * im + params.alphaIm * re;
|
|
32
|
+
}
|
|
33
|
+
}
|
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
// daxpy: y := alpha * x + y, double-double (Dekker) f64 emulation of saxpy.
|
|
2
|
+
// Each element costs one ddMulProtected (alpha*x[i]) then one ddAddProtected
|
|
3
|
+
// (+ y[i]) — the same two-protected-op shape ddot spends per term, applied
|
|
4
|
+
// straight to the output instead of folded into a reduction. See dscal.wgsl
|
|
5
|
+
// for why this is a uniform main pass plus a ragged, select-masked tail
|
|
6
|
+
// rather than a plain `id < params.n` grid-stride loop.
|
|
7
|
+
|
|
8
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
9
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
10
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
11
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
12
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
13
|
+
|
|
14
|
+
struct Params {
|
|
15
|
+
n: u32,
|
|
16
|
+
alphaHi: f32,
|
|
17
|
+
alphaLo: f32,
|
|
18
|
+
x_inc: u32,
|
|
19
|
+
y_inc: u32,
|
|
20
|
+
}
|
|
21
|
+
|
|
22
|
+
const WGS: u32 = 64;
|
|
23
|
+
|
|
24
|
+
@compute @workgroup_size(64)
|
|
25
|
+
fn daxpy_main(
|
|
26
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
27
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
28
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
29
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
30
|
+
) {
|
|
31
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
32
|
+
let stride = num_wg.x * WGS;
|
|
33
|
+
|
|
34
|
+
let n_floor = (params.n / stride) * stride;
|
|
35
|
+
let mainIters = n_floor / stride;
|
|
36
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
37
|
+
let id = gid.x + iter * stride;
|
|
38
|
+
let ix = id * params.x_inc;
|
|
39
|
+
let iy = id * params.y_inc;
|
|
40
|
+
let prod = ddMulProtected(alpha, DD(xHi[ix], xLo[ix]), lid.x);
|
|
41
|
+
let result = ddAddProtected(prod, DD(yHi[iy], yLo[iy]), lid.x);
|
|
42
|
+
yHi[iy] = result.hi;
|
|
43
|
+
yLo[iy] = result.lo;
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
// Tail is ragged (0 or 1 extra per thread) — pad to this workgroup's worst
|
|
47
|
+
// case so every thread still calls ddMulProtected/ddAddProtected the same
|
|
48
|
+
// number of times (their barriers need that), masking only the write.
|
|
49
|
+
let wgBaseGid = wgid.x * WGS;
|
|
50
|
+
var tailIters = 0u;
|
|
51
|
+
if (n_floor + wgBaseGid < params.n) {
|
|
52
|
+
tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
|
|
53
|
+
}
|
|
54
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
55
|
+
let id = n_floor + gid.x + iter * stride;
|
|
56
|
+
let valid = id < params.n;
|
|
57
|
+
let ix = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
58
|
+
let iy = select(0u, id * params.y_inc, valid);
|
|
59
|
+
let prod = ddMulProtected(alpha, DD(xHi[ix], xLo[ix]), lid.x);
|
|
60
|
+
let result = ddAddProtected(prod, DD(yHi[iy], yLo[iy]), lid.x);
|
|
61
|
+
if (valid) {
|
|
62
|
+
yHi[iy] = result.hi;
|
|
63
|
+
yLo[iy] = result.lo;
|
|
64
|
+
}
|
|
65
|
+
}
|
|
66
|
+
}
|
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
// dcopy: y = x, double-double (Dekker) f64 emulation of scopy. A copy is
|
|
2
|
+
// pure data movement — hi and lo are transferred verbatim, with no
|
|
3
|
+
// arithmetic at all — so (unlike dscal/daxpy/ddot) this needs no
|
|
4
|
+
// ddMulProtected/ddAddProtected renormalizing barrier, and so no
|
|
5
|
+
// ragged-tail-with-barrier split; a plain grid-stride loop is safe, same
|
|
6
|
+
// shape as scopy.wgsl itself.
|
|
7
|
+
|
|
8
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
9
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
10
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
11
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
12
|
+
|
|
13
|
+
struct Params {
|
|
14
|
+
n: u32,
|
|
15
|
+
x_inc: u32,
|
|
16
|
+
y_inc: u32,
|
|
17
|
+
}
|
|
18
|
+
|
|
19
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
20
|
+
|
|
21
|
+
const WGS: u32 = 64;
|
|
22
|
+
|
|
23
|
+
@compute @workgroup_size(64)
|
|
24
|
+
fn main(
|
|
25
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
26
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
27
|
+
) {
|
|
28
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
29
|
+
let ix = id * params.x_inc;
|
|
30
|
+
let iy = id * params.y_inc;
|
|
31
|
+
yHi[iy] = xHi[ix];
|
|
32
|
+
yLo[iy] = xLo[ix];
|
|
33
|
+
}
|
|
34
|
+
}
|
|
@@ -0,0 +1,106 @@
|
|
|
1
|
+
// ddot: sum(x[i] * y[i]), double-double (Dekker). Same ILP=4 shape as
|
|
2
|
+
// dasum.wgsl, which this mirrors closely — the only structural difference is
|
|
3
|
+
// a second input vector and a product where dasum takes an absolute value.
|
|
4
|
+
//
|
|
5
|
+
// See f64/utils/add.wgsl for ddAddProtected and why plain ddAdd isn't safe,
|
|
6
|
+
// and f64/utils/multiply.wgsl for ddMulProtected. The multiply itself
|
|
7
|
+
// (twoProdBit) needs no barrier; only its final renormalisation does, which
|
|
8
|
+
// is why each element costs two protected ops here against dasum's one.
|
|
9
|
+
|
|
10
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
11
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
12
|
+
@group(0) @binding(2) var<storage, read> yHi: array<f32>;
|
|
13
|
+
@group(0) @binding(3) var<storage, read> yLo: array<f32>;
|
|
14
|
+
@group(0) @binding(4) var<storage, read_write> partialsHi: array<f32>;
|
|
15
|
+
@group(0) @binding(5) var<storage, read_write> partialsLo: array<f32>;
|
|
16
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
17
|
+
|
|
18
|
+
struct Params {
|
|
19
|
+
n: u32,
|
|
20
|
+
x_inc: u32,
|
|
21
|
+
y_inc: u32,
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
const WGS: u32 = 64;
|
|
25
|
+
|
|
26
|
+
var<workgroup> tile: array<DD, 64>;
|
|
27
|
+
|
|
28
|
+
@compute @workgroup_size(64)
|
|
29
|
+
fn ddot_main(
|
|
30
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
31
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
32
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
33
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
34
|
+
) {
|
|
35
|
+
var acc0 = DD(0.0, 0.0);
|
|
36
|
+
var acc1 = DD(0.0, 0.0);
|
|
37
|
+
var acc2 = DD(0.0, 0.0);
|
|
38
|
+
var acc3 = DD(0.0, 0.0);
|
|
39
|
+
|
|
40
|
+
let stride = num_wg.x * WGS;
|
|
41
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
42
|
+
|
|
43
|
+
// Same trip count for every thread, but driven by a counter, not `id`
|
|
44
|
+
// itself (the protected ops' barriers need a provably-uniform loop bound).
|
|
45
|
+
let mainIters = n4_floor / (4u * stride);
|
|
46
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
47
|
+
let id = gid.x + iter * 4u * stride;
|
|
48
|
+
let d0 = id;
|
|
49
|
+
let d1 = id + stride;
|
|
50
|
+
let d2 = id + 2u * stride;
|
|
51
|
+
let d3 = id + 3u * stride;
|
|
52
|
+
|
|
53
|
+
let p0 = ddMulProtected(DD(xHi[d0 * params.x_inc], xLo[d0 * params.x_inc]),
|
|
54
|
+
DD(yHi[d0 * params.y_inc], yLo[d0 * params.y_inc]), lid.x);
|
|
55
|
+
let p1 = ddMulProtected(DD(xHi[d1 * params.x_inc], xLo[d1 * params.x_inc]),
|
|
56
|
+
DD(yHi[d1 * params.y_inc], yLo[d1 * params.y_inc]), lid.x);
|
|
57
|
+
let p2 = ddMulProtected(DD(xHi[d2 * params.x_inc], xLo[d2 * params.x_inc]),
|
|
58
|
+
DD(yHi[d2 * params.y_inc], yLo[d2 * params.y_inc]), lid.x);
|
|
59
|
+
let p3 = ddMulProtected(DD(xHi[d3 * params.x_inc], xLo[d3 * params.x_inc]),
|
|
60
|
+
DD(yHi[d3 * params.y_inc], yLo[d3 * params.y_inc]), lid.x);
|
|
61
|
+
|
|
62
|
+
acc0 = ddAddProtected(acc0, p0, lid.x);
|
|
63
|
+
acc1 = ddAddProtected(acc1, p1, lid.x);
|
|
64
|
+
acc2 = ddAddProtected(acc2, p2, lid.x);
|
|
65
|
+
acc3 = ddAddProtected(acc3, p3, lid.x);
|
|
66
|
+
}
|
|
67
|
+
|
|
68
|
+
// Tail is ragged (0-3 extra per thread) — pad to this workgroup's worst case.
|
|
69
|
+
// Out-of-range lanes still run the multiply (it carries a barrier, so every
|
|
70
|
+
// thread must reach it) against index 0, then mask the result to zero.
|
|
71
|
+
let wgBaseGid = wgid.x * WGS;
|
|
72
|
+
var tailIters = 0u;
|
|
73
|
+
if (n4_floor + wgBaseGid < params.n) {
|
|
74
|
+
tailIters = (params.n - 1u - n4_floor - wgBaseGid) / stride + 1u;
|
|
75
|
+
}
|
|
76
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
77
|
+
let id = n4_floor + gid.x + iter * stride;
|
|
78
|
+
let valid = id < params.n;
|
|
79
|
+
let ix = select(0u, id * params.x_inc, valid);
|
|
80
|
+
let iy = select(0u, id * params.y_inc, valid);
|
|
81
|
+
let prod = ddMulProtected(DD(xHi[ix], xLo[ix]), DD(yHi[iy], yLo[iy]), lid.x);
|
|
82
|
+
// select() has no DD overload
|
|
83
|
+
let contribution = DD(select(0.0, prod.hi, valid), select(0.0, prod.lo, valid));
|
|
84
|
+
acc0 = ddAddProtected(acc0, contribution, lid.x);
|
|
85
|
+
}
|
|
86
|
+
|
|
87
|
+
let combined01 = ddAddProtected(acc0, acc1, lid.x);
|
|
88
|
+
let combined23 = ddAddProtected(acc2, acc3, lid.x);
|
|
89
|
+
tile[lid.x] = ddAddProtected(combined01, combined23, lid.x);
|
|
90
|
+
workgroupBarrier();
|
|
91
|
+
|
|
92
|
+
// Inactive threads combine against a throwaway partner and discard it
|
|
93
|
+
// (ddAddProtected must be called unconditionally by every thread).
|
|
94
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
95
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
96
|
+
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
97
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
98
|
+
if (lid.x < s) { tile[lid.x] = combined; }
|
|
99
|
+
workgroupBarrier();
|
|
100
|
+
}
|
|
101
|
+
|
|
102
|
+
if (lid.x == 0u) {
|
|
103
|
+
partialsHi[wgid.x] = tile[0].hi;
|
|
104
|
+
partialsLo[wgid.x] = tile[0].lo;
|
|
105
|
+
}
|
|
106
|
+
}
|
|
@@ -0,0 +1,167 @@
|
|
|
1
|
+
// dnrm2: result = sqrt(sum(x[i] * x[i])), double-double (Dekker) f64
|
|
2
|
+
// emulation of snrm2 — same scaled accumulation (Blue's algorithm), just
|
|
3
|
+
// with `scale`/`ssq` as DD pairs (via ddDivProtected/ddMulProtected/
|
|
4
|
+
// ddAddProtected/ddSqrtProtected) instead of plain f32. Squaring still
|
|
5
|
+
// saturates an f32 hi component above ~1.8e19 regardless of DD precision
|
|
6
|
+
// (DD widens the mantissa, not the exponent range), so the scaling is
|
|
7
|
+
// still needed here for the same reason it was in snrm2.
|
|
8
|
+
//
|
|
9
|
+
// snrm2.wgsl's ssqAccum/ssqMerge each branch on which operand is bigger —
|
|
10
|
+
// can't carry over directly, since a protected op's workgroupBarrier()
|
|
11
|
+
// needs every thread to reach the same call site, and here different
|
|
12
|
+
// threads could take different branches. Both formulas are computed
|
|
13
|
+
// unconditionally below; only the final combine (`ddSelect`) differs per
|
|
14
|
+
// thread — same fix shape as drot's/drotm's own per-dispatch flags, just
|
|
15
|
+
// applied to a per-element branch instead.
|
|
16
|
+
//
|
|
17
|
+
// pass 1 dispatches 2*WGS workgroups; pass 2 (reduction/scaledSumF64.wgsl)
|
|
18
|
+
// duplicates ssqAccumProtected/ssqMergeProtected rather than sharing them,
|
|
19
|
+
// same as the plain-f32 pair already does.
|
|
20
|
+
|
|
21
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
22
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
23
|
+
@group(0) @binding(2) var<storage, read_write> partialsScaleHi: array<f32>;
|
|
24
|
+
@group(0) @binding(3) var<storage, read_write> partialsScaleLo: array<f32>;
|
|
25
|
+
@group(0) @binding(4) var<storage, read_write> partialsSsqHi: array<f32>;
|
|
26
|
+
@group(0) @binding(5) var<storage, read_write> partialsSsqLo: array<f32>;
|
|
27
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
28
|
+
|
|
29
|
+
struct Params {
|
|
30
|
+
n: u32,
|
|
31
|
+
x_inc: u32,
|
|
32
|
+
}
|
|
33
|
+
|
|
34
|
+
const WGS: u32 = 64;
|
|
35
|
+
|
|
36
|
+
struct ScaleSsq {
|
|
37
|
+
scale: DD,
|
|
38
|
+
ssq: DD,
|
|
39
|
+
}
|
|
40
|
+
|
|
41
|
+
fn ddSelect(a: DD, b: DD, cond: bool) -> DD {
|
|
42
|
+
return DD(select(a.hi, b.hi, cond), select(a.lo, b.lo, cond));
|
|
43
|
+
}
|
|
44
|
+
|
|
45
|
+
// Folds one more |value| (DD) into a running (scale, ssq) pair — branch-free,
|
|
46
|
+
// see file header. `bigger`/`smaller` name the two operands by magnitude
|
|
47
|
+
// (not by which one was "acc" vs "new"), and biggerIsZero==true only when
|
|
48
|
+
// both scale and absxi are still exactly zero (the very first zero
|
|
49
|
+
// elements, before any nonzero value has been seen) — substituting a safe
|
|
50
|
+
// denominator there avoids a 0/0 without needing a separate branch/return;
|
|
51
|
+
// the arithmetic already reduces to a correct no-op in that case.
|
|
52
|
+
fn ssqAccumProtected(acc: ScaleSsq, absxi: DD, threadSlot: u32) -> ScaleSsq {
|
|
53
|
+
let isBigger = ddGreater(absxi, acc.scale);
|
|
54
|
+
let bigger = ddSelect(acc.scale, absxi, isBigger);
|
|
55
|
+
let smaller = ddSelect(absxi, acc.scale, isBigger);
|
|
56
|
+
let biggerIsZero = bigger.hi == 0.0;
|
|
57
|
+
let safeBigger = ddSelect(bigger, DD(1.0, 0.0), biggerIsZero);
|
|
58
|
+
let r = ddDivProtected(smaller, safeBigger, threadSlot);
|
|
59
|
+
let rsq = ddMulProtected(r, r, threadSlot);
|
|
60
|
+
let ssqTimesRsq = ddMulProtected(acc.ssq, rsq, threadSlot);
|
|
61
|
+
let sumIfBigger = ddAddProtected(DD(1.0, 0.0), ssqTimesRsq, threadSlot);
|
|
62
|
+
let sumIfNotBigger = ddAddProtected(acc.ssq, rsq, threadSlot);
|
|
63
|
+
let newSsq = ddSelect(sumIfNotBigger, sumIfBigger, isBigger);
|
|
64
|
+
return ScaleSsq(bigger, newSsq);
|
|
65
|
+
}
|
|
66
|
+
|
|
67
|
+
// Associative merge of two independent (scale, ssq) partials — same
|
|
68
|
+
// branch-free shape, for combining ILP lanes and the tree reduction.
|
|
69
|
+
fn ssqMergeProtected(a: ScaleSsq, b: ScaleSsq, threadSlot: u32) -> ScaleSsq {
|
|
70
|
+
let isBigger = !ddGreater(b.scale, a.scale); // a.scale >= b.scale
|
|
71
|
+
let bigger = ddSelect(b.scale, a.scale, isBigger);
|
|
72
|
+
let smaller = ddSelect(a.scale, b.scale, isBigger);
|
|
73
|
+
let biggerSsq = ddSelect(b.ssq, a.ssq, isBigger);
|
|
74
|
+
let smallerSsq = ddSelect(a.ssq, b.ssq, isBigger);
|
|
75
|
+
let biggerIsZero = bigger.hi == 0.0;
|
|
76
|
+
let safeBigger = ddSelect(bigger, DD(1.0, 0.0), biggerIsZero);
|
|
77
|
+
let r = ddDivProtected(smaller, safeBigger, threadSlot);
|
|
78
|
+
let rsq = ddMulProtected(r, r, threadSlot);
|
|
79
|
+
let smallerSsqTimesRsq = ddMulProtected(smallerSsq, rsq, threadSlot);
|
|
80
|
+
let newSsq = ddAddProtected(biggerSsq, smallerSsqTimesRsq, threadSlot);
|
|
81
|
+
return ScaleSsq(bigger, newSsq);
|
|
82
|
+
}
|
|
83
|
+
|
|
84
|
+
var<workgroup> tileScaleHi: array<f32, 64>;
|
|
85
|
+
var<workgroup> tileScaleLo: array<f32, 64>;
|
|
86
|
+
var<workgroup> tileSsqHi: array<f32, 64>;
|
|
87
|
+
var<workgroup> tileSsqLo: array<f32, 64>;
|
|
88
|
+
|
|
89
|
+
@compute @workgroup_size(64)
|
|
90
|
+
fn dnrm2_main(
|
|
91
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
92
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
93
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
94
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
95
|
+
) {
|
|
96
|
+
var acc0 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
|
|
97
|
+
var acc1 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
|
|
98
|
+
var acc2 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
|
|
99
|
+
var acc3 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
|
|
100
|
+
|
|
101
|
+
let stride = num_wg.x * WGS;
|
|
102
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
103
|
+
|
|
104
|
+
// Same trip count for every thread, driven by a counter (protected ops'
|
|
105
|
+
// barriers need a provably-uniform loop bound) — see dasum.wgsl.
|
|
106
|
+
let mainIters = n4_floor / (4u * stride);
|
|
107
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
108
|
+
let id = gid.x + iter * 4u * stride;
|
|
109
|
+
let i0 = id * params.x_inc;
|
|
110
|
+
let i1 = (id + stride) * params.x_inc;
|
|
111
|
+
let i2 = (id + 2u * stride) * params.x_inc;
|
|
112
|
+
let i3 = (id + 3u * stride) * params.x_inc;
|
|
113
|
+
acc0 = ssqAccumProtected(acc0, ddAbs(DD(xHi[i0], xLo[i0])), lid.x);
|
|
114
|
+
acc1 = ssqAccumProtected(acc1, ddAbs(DD(xHi[i1], xLo[i1])), lid.x);
|
|
115
|
+
acc2 = ssqAccumProtected(acc2, ddAbs(DD(xHi[i2], xLo[i2])), lid.x);
|
|
116
|
+
acc3 = ssqAccumProtected(acc3, ddAbs(DD(xHi[i3], xLo[i3])), lid.x);
|
|
117
|
+
}
|
|
118
|
+
|
|
119
|
+
// Tail is ragged (0-3 extra per thread) — pad to this workgroup's worst
|
|
120
|
+
// case, masking an invalid element to exactly 0 (contributes nothing).
|
|
121
|
+
let wgBaseGid = wgid.x * WGS;
|
|
122
|
+
var tailIters = 0u;
|
|
123
|
+
if (n4_floor + wgBaseGid < params.n) {
|
|
124
|
+
tailIters = (params.n - 1u - n4_floor - wgBaseGid) / stride + 1u;
|
|
125
|
+
}
|
|
126
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
127
|
+
let id = n4_floor + gid.x + iter * stride;
|
|
128
|
+
let valid = id < params.n;
|
|
129
|
+
let i = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
130
|
+
let loaded = ddAbs(DD(xHi[i], xLo[i]));
|
|
131
|
+
let contribution = DD(select(0.0, loaded.hi, valid), select(0.0, loaded.lo, valid));
|
|
132
|
+
acc0 = ssqAccumProtected(acc0, contribution, lid.x);
|
|
133
|
+
}
|
|
134
|
+
|
|
135
|
+
let combined01 = ssqMergeProtected(acc0, acc1, lid.x);
|
|
136
|
+
let combined23 = ssqMergeProtected(acc2, acc3, lid.x);
|
|
137
|
+
let combined = ssqMergeProtected(combined01, combined23, lid.x);
|
|
138
|
+
tileScaleHi[lid.x] = combined.scale.hi;
|
|
139
|
+
tileScaleLo[lid.x] = combined.scale.lo;
|
|
140
|
+
tileSsqHi[lid.x] = combined.ssq.hi;
|
|
141
|
+
tileSsqLo[lid.x] = combined.ssq.lo;
|
|
142
|
+
workgroupBarrier();
|
|
143
|
+
|
|
144
|
+
// Inactive threads merge against a throwaway partner and discard it
|
|
145
|
+
// (ssqMergeProtected must be called unconditionally by every thread).
|
|
146
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
147
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
148
|
+
let a = ScaleSsq(DD(tileScaleHi[lid.x], tileScaleLo[lid.x]), DD(tileSsqHi[lid.x], tileSsqLo[lid.x]));
|
|
149
|
+
let b = ScaleSsq(DD(tileScaleHi[partner], tileScaleLo[partner]), DD(tileSsqHi[partner], tileSsqLo[partner]));
|
|
150
|
+
let merged = ssqMergeProtected(a, b, lid.x);
|
|
151
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
152
|
+
if (lid.x < s) {
|
|
153
|
+
tileScaleHi[lid.x] = merged.scale.hi;
|
|
154
|
+
tileScaleLo[lid.x] = merged.scale.lo;
|
|
155
|
+
tileSsqHi[lid.x] = merged.ssq.hi;
|
|
156
|
+
tileSsqLo[lid.x] = merged.ssq.lo;
|
|
157
|
+
}
|
|
158
|
+
workgroupBarrier();
|
|
159
|
+
}
|
|
160
|
+
|
|
161
|
+
if (lid.x == 0u) {
|
|
162
|
+
partialsScaleHi[wgid.x] = tileScaleHi[0];
|
|
163
|
+
partialsScaleLo[wgid.x] = tileScaleLo[0];
|
|
164
|
+
partialsSsqHi[wgid.x] = tileSsqHi[0];
|
|
165
|
+
partialsSsqLo[wgid.x] = tileSsqLo[0];
|
|
166
|
+
}
|
|
167
|
+
}
|
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
// drot: x = c*x + s*y, y = -s*x + c*y — double-double (Dekker) f64 emulation
|
|
2
|
+
// of srot. c, s, x, and y are each split into an f32 (hi, lo) pair. Each
|
|
3
|
+
// element costs four ddMulProtected (c*x, s*y, -s*x, c*y) then two
|
|
4
|
+
// ddAddProtected (the two sums) — negS is computed once outside the loop
|
|
5
|
+
// via bitcast negation (exact, no rounding, so no barrier needed there)
|
|
6
|
+
// rather than adding a DD-subtract helper. See dscal.wgsl for why this is a
|
|
7
|
+
// uniform main pass plus a ragged, select-masked tail rather than a plain
|
|
8
|
+
// `id < params.n` grid-stride loop.
|
|
9
|
+
|
|
10
|
+
@group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
|
|
11
|
+
@group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
|
|
12
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
13
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
14
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
15
|
+
|
|
16
|
+
struct Params {
|
|
17
|
+
n: u32,
|
|
18
|
+
cHi: f32,
|
|
19
|
+
cLo: f32,
|
|
20
|
+
sHi: f32,
|
|
21
|
+
sLo: f32,
|
|
22
|
+
x_inc: u32,
|
|
23
|
+
y_inc: u32,
|
|
24
|
+
}
|
|
25
|
+
|
|
26
|
+
const WGS: u32 = 64;
|
|
27
|
+
|
|
28
|
+
@compute @workgroup_size(64)
|
|
29
|
+
fn drot_main(
|
|
30
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
31
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
32
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
33
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
34
|
+
) {
|
|
35
|
+
let c = DD(params.cHi, params.cLo);
|
|
36
|
+
let s = DD(params.sHi, params.sLo);
|
|
37
|
+
let negS = DD(negf(params.sHi), negf(params.sLo));
|
|
38
|
+
let stride = num_wg.x * WGS;
|
|
39
|
+
|
|
40
|
+
let n_floor = (params.n / stride) * stride;
|
|
41
|
+
let mainIters = n_floor / stride;
|
|
42
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
43
|
+
let id = gid.x + iter * stride;
|
|
44
|
+
let ix = id * params.x_inc;
|
|
45
|
+
let iy = id * params.y_inc;
|
|
46
|
+
let xi = DD(xHi[ix], xLo[ix]);
|
|
47
|
+
let yi = DD(yHi[iy], yLo[iy]);
|
|
48
|
+
let xNew = ddAddProtected(ddMulProtected(c, xi, lid.x), ddMulProtected(s, yi, lid.x), lid.x);
|
|
49
|
+
let yNew = ddAddProtected(ddMulProtected(negS, xi, lid.x), ddMulProtected(c, yi, lid.x), lid.x);
|
|
50
|
+
xHi[ix] = xNew.hi;
|
|
51
|
+
xLo[ix] = xNew.lo;
|
|
52
|
+
yHi[iy] = yNew.hi;
|
|
53
|
+
yLo[iy] = yNew.lo;
|
|
54
|
+
}
|
|
55
|
+
|
|
56
|
+
// Tail is ragged (0 or 1 extra per thread) — pad to this workgroup's worst
|
|
57
|
+
// case so every thread in the workgroup still calls ddMulProtected/
|
|
58
|
+
// ddAddProtected the same number of times (their barriers need that),
|
|
59
|
+
// masking only the write.
|
|
60
|
+
let wgBaseGid = wgid.x * WGS;
|
|
61
|
+
var tailIters = 0u;
|
|
62
|
+
if (n_floor + wgBaseGid < params.n) {
|
|
63
|
+
tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
|
|
64
|
+
}
|
|
65
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
66
|
+
let id = n_floor + gid.x + iter * stride;
|
|
67
|
+
let valid = id < params.n;
|
|
68
|
+
let ix = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
69
|
+
let iy = select(0u, id * params.y_inc, valid);
|
|
70
|
+
let xi = DD(xHi[ix], xLo[ix]);
|
|
71
|
+
let yi = DD(yHi[iy], yLo[iy]);
|
|
72
|
+
let xNew = ddAddProtected(ddMulProtected(c, xi, lid.x), ddMulProtected(s, yi, lid.x), lid.x);
|
|
73
|
+
let yNew = ddAddProtected(ddMulProtected(negS, xi, lid.x), ddMulProtected(c, yi, lid.x), lid.x);
|
|
74
|
+
if (valid) {
|
|
75
|
+
xHi[ix] = xNew.hi;
|
|
76
|
+
xLo[ix] = xNew.lo;
|
|
77
|
+
yHi[iy] = yNew.hi;
|
|
78
|
+
yLo[iy] = yNew.lo;
|
|
79
|
+
}
|
|
80
|
+
}
|
|
81
|
+
}
|