wgblas 2.0.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 +20 -18
- package/dist/wgblas.browser.js +2172 -1174
- package/index.d.mts +49 -44
- package/index.mjs +11 -0
- package/package.json +133 -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 +126 -17
- package/src/classes/GpuVector.d.mts +36 -40
- package/src/classes/GpuVector.mjs +66 -11
- package/src/cscal/cscal.d.mts +47 -0
- package/src/cscal/cscal.mjs +98 -0
- package/src/dasum/dasum.d.mts +4 -4
- package/src/dasum/dasum.mjs +38 -20
- 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/devdocs.mjs +13 -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.d.mts +20 -2
- package/src/idamax/idamax.mjs +56 -24
- package/src/init.mjs +117 -56
- package/src/isamax/isamax.d.mts +20 -2
- package/src/isamax/isamax.mjs +21 -16
- package/src/random/random.d.mts +37 -39
- package/src/random/random.mjs +39 -7
- package/src/sasum/sasum.d.mts +2 -2
- package/src/sasum/sasum.mjs +20 -16
- package/src/saxpy/saxpy.d.mts +2 -2
- package/src/saxpy/saxpy.mjs +14 -11
- package/src/scopy/scopy.d.mts +2 -2
- package/src/scopy/scopy.mjs +13 -9
- package/src/sdot/sdot.d.mts +2 -2
- package/src/sdot/sdot.mjs +21 -17
- package/src/sgemm/sgemm.d.mts +2 -2
- package/src/sgemm/sgemm.mjs +109 -40
- package/src/sgemmtr/sgemmtr.d.mts +3 -2
- package/src/sgemmtr/sgemmtr.mjs +98 -40
- package/src/sgemv/sgemv.d.mts +2 -2
- package/src/sgemv/sgemv.mjs +69 -41
- package/src/sger/sger.d.mts +2 -2
- package/src/sger/sger.mjs +43 -19
- 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 +233 -14
- package/src/shaders/reduction/scaledSum.wgsl +65 -0
- package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
- package/src/shaders/sgemm_large.wgsl +107 -18
- package/src/shaders/sgemm_small.wgsl +115 -15
- package/src/shaders/sgemmtr_large.wgsl +4 -1
- package/src/shaders/sgemmtr_small.wgsl +4 -1
- package/src/shaders/sgemv_n.wgsl +3 -1
- package/src/shaders/sgemv_t.wgsl +3 -1
- package/src/shaders/snrm2.wgsl +72 -23
- package/src/shaders/ssymv.wgsl +3 -1
- package/src/snrm2/snrm2.d.mts +2 -2
- package/src/snrm2/snrm2.mjs +41 -23
- package/src/srot/srot.d.mts +2 -4
- package/src/srot/srot.mjs +16 -11
- package/src/srotm/srotm.d.mts +2 -4
- package/src/srotm/srotm.mjs +17 -11
- package/src/sscal/sscal.d.mts +3 -3
- package/src/sscal/sscal.mjs +14 -12
- package/src/sswap/sswap.d.mts +2 -2
- package/src/sswap/sswap.mjs +18 -10
- package/src/ssymm/ssymm.d.mts +5 -4
- package/src/ssymm/ssymm.mjs +150 -54
- package/src/ssymv/ssymv.d.mts +2 -2
- package/src/ssymv/ssymv.mjs +47 -26
- package/src/ssyr/ssyr.d.mts +2 -2
- package/src/ssyr/ssyr.mjs +38 -17
- package/src/ssyr2/ssyr2.d.mts +2 -2
- package/src/ssyr2/ssyr2.mjs +48 -21
- package/src/ssyr2k/ssyr2k.d.mts +3 -2
- package/src/ssyr2k/ssyr2k.mjs +140 -62
- package/src/ssyrk/ssyrk.d.mts +3 -2
- package/src/ssyrk/ssyrk.mjs +91 -39
- package/src/strmm/strmm.d.mts +5 -4
- package/src/strmm/strmm.mjs +174 -60
- package/src/strmv/strmv.d.mts +2 -2
- package/src/strmv/strmv.mjs +42 -20
- package/src/strsm/strsm.d.mts +6 -4
- package/src/strsm/strsm.mjs +438 -174
- package/src/strsv/strsv.d.mts +5 -3
- package/src/strsv/strsv.mjs +89 -34
- package/src/util/benchmark.mjs +9 -9
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +139 -24
- package/src/util/complex.mjs +87 -0
- package/src/util/compute.mjs +19 -16
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +49 -0
- package/src/util/pipeline.mjs +44 -10
- package/src/util/workgroup.mjs +72 -7
- package/src/shaders/browser-shaders.mjs +0 -81
- package/src/shaders/f64add.wgsl +0 -281
|
@@ -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
|
+
}
|
|
@@ -0,0 +1,99 @@
|
|
|
1
|
+
// drotm: applies a modified Givens rotation H to vectors x and y — double-
|
|
2
|
+
// double (Dekker) f64 emulation of srotm. paramHi/paramLo[0] = flag: -1
|
|
3
|
+
// (full H), 0 (unit diagonal), 1 (unit off-diagonal). param = [ flag, h11,
|
|
4
|
+
// h21, h12, h22 ], each entry an f32 (hi, lo) pair.
|
|
5
|
+
// flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
|
|
6
|
+
//
|
|
7
|
+
// h11/h12/h21/h22 are resolved once, outside the loop, from the (uniform
|
|
8
|
+
// across every thread) flag — same shape as srot's c/s, so no barrier is
|
|
9
|
+
// needed for that selection itself. Each element then costs four
|
|
10
|
+
// ddMulProtected + two ddAddProtected, same as drot.
|
|
11
|
+
|
|
12
|
+
@group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
|
|
13
|
+
@group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
|
|
14
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
15
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
16
|
+
@group(0) @binding(4) var<storage, read> paramHi: array<f32>;
|
|
17
|
+
@group(0) @binding(5) var<storage, read> paramLo: array<f32>;
|
|
18
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
19
|
+
|
|
20
|
+
struct Params {
|
|
21
|
+
n: u32,
|
|
22
|
+
x_inc: u32,
|
|
23
|
+
y_inc: u32,
|
|
24
|
+
}
|
|
25
|
+
|
|
26
|
+
const WGS: u32 = 64;
|
|
27
|
+
const ONE: DD = DD(1.0, 0.0);
|
|
28
|
+
const NEG_ONE: DD = DD(-1.0, 0.0);
|
|
29
|
+
|
|
30
|
+
@compute @workgroup_size(64)
|
|
31
|
+
fn drotm_main(
|
|
32
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
33
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
34
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
35
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
36
|
+
) {
|
|
37
|
+
let flag = paramHi[0]; // exact small integer (-1, 0, or 1) — lo is always 0
|
|
38
|
+
|
|
39
|
+
var h11: DD; var h12: DD;
|
|
40
|
+
var h21: DD; var h22: DD;
|
|
41
|
+
|
|
42
|
+
if (flag == -1.0) {
|
|
43
|
+
// full 2x2 matrix
|
|
44
|
+
h11 = DD(paramHi[1], paramLo[1]); h21 = DD(paramHi[2], paramLo[2]);
|
|
45
|
+
h12 = DD(paramHi[3], paramLo[3]); h22 = DD(paramHi[4], paramLo[4]);
|
|
46
|
+
} else if (flag == 0.0) {
|
|
47
|
+
// diagonal fixed at 1
|
|
48
|
+
h11 = ONE; h21 = DD(paramHi[2], paramLo[2]);
|
|
49
|
+
h12 = DD(paramHi[3], paramLo[3]); h22 = ONE;
|
|
50
|
+
} else {
|
|
51
|
+
// flag == 1.0: off-diagonal fixed at +1 / -1
|
|
52
|
+
h11 = DD(paramHi[1], paramLo[1]); h21 = NEG_ONE;
|
|
53
|
+
h12 = ONE; h22 = DD(paramHi[4], paramLo[4]);
|
|
54
|
+
}
|
|
55
|
+
|
|
56
|
+
let stride = num_wg.x * WGS;
|
|
57
|
+
|
|
58
|
+
let n_floor = (params.n / stride) * stride;
|
|
59
|
+
let mainIters = n_floor / stride;
|
|
60
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
61
|
+
let id = gid.x + iter * stride;
|
|
62
|
+
let ix = id * params.x_inc;
|
|
63
|
+
let iy = id * params.y_inc;
|
|
64
|
+
let xi = DD(xHi[ix], xLo[ix]);
|
|
65
|
+
let yi = DD(yHi[iy], yLo[iy]);
|
|
66
|
+
let xNew = ddAddProtected(ddMulProtected(h11, xi, lid.x), ddMulProtected(h12, yi, lid.x), lid.x);
|
|
67
|
+
let yNew = ddAddProtected(ddMulProtected(h21, xi, lid.x), ddMulProtected(h22, yi, lid.x), lid.x);
|
|
68
|
+
xHi[ix] = xNew.hi;
|
|
69
|
+
xLo[ix] = xNew.lo;
|
|
70
|
+
yHi[iy] = yNew.hi;
|
|
71
|
+
yLo[iy] = yNew.lo;
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
// Tail is ragged (0 or 1 extra per thread) — pad to this workgroup's worst
|
|
75
|
+
// case so every thread in the workgroup still calls ddMulProtected/
|
|
76
|
+
// ddAddProtected the same number of times (their barriers need that),
|
|
77
|
+
// masking only the write.
|
|
78
|
+
let wgBaseGid = wgid.x * WGS;
|
|
79
|
+
var tailIters = 0u;
|
|
80
|
+
if (n_floor + wgBaseGid < params.n) {
|
|
81
|
+
tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
|
|
82
|
+
}
|
|
83
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
84
|
+
let id = n_floor + gid.x + iter * stride;
|
|
85
|
+
let valid = id < params.n;
|
|
86
|
+
let ix = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
87
|
+
let iy = select(0u, id * params.y_inc, valid);
|
|
88
|
+
let xi = DD(xHi[ix], xLo[ix]);
|
|
89
|
+
let yi = DD(yHi[iy], yLo[iy]);
|
|
90
|
+
let xNew = ddAddProtected(ddMulProtected(h11, xi, lid.x), ddMulProtected(h12, yi, lid.x), lid.x);
|
|
91
|
+
let yNew = ddAddProtected(ddMulProtected(h21, xi, lid.x), ddMulProtected(h22, yi, lid.x), lid.x);
|
|
92
|
+
if (valid) {
|
|
93
|
+
xHi[ix] = xNew.hi;
|
|
94
|
+
xLo[ix] = xNew.lo;
|
|
95
|
+
yHi[iy] = yNew.hi;
|
|
96
|
+
yLo[iy] = yNew.lo;
|
|
97
|
+
}
|
|
98
|
+
}
|
|
99
|
+
}
|
|
@@ -0,0 +1,60 @@
|
|
|
1
|
+
// dscal: x := alpha * x, double-double (Dekker) f64 emulation of sscal.
|
|
2
|
+
// alpha and x are each an f32 (hi, lo) pair. See f64/utils/multiply.wgsl for
|
|
3
|
+
// ddMulProtected and why plain ddMulRaw isn't safe without a renormalizing
|
|
4
|
+
// barrier — that barrier needs a provably uniform loop trip count across
|
|
5
|
+
// every thread in the workgroup, so (like dasum.wgsl's reduction loop) this
|
|
6
|
+
// splits into a uniform main pass plus a ragged, select-masked tail rather
|
|
7
|
+
// than a plain `id < params.n` grid-stride loop.
|
|
8
|
+
|
|
9
|
+
@group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
|
|
10
|
+
@group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
|
|
11
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
12
|
+
|
|
13
|
+
struct Params {
|
|
14
|
+
n: u32,
|
|
15
|
+
alphaHi: f32,
|
|
16
|
+
alphaLo: f32,
|
|
17
|
+
x_inc: u32,
|
|
18
|
+
}
|
|
19
|
+
|
|
20
|
+
const WGS: u32 = 64;
|
|
21
|
+
|
|
22
|
+
@compute @workgroup_size(64)
|
|
23
|
+
fn dscal_main(
|
|
24
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
25
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
26
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
27
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
28
|
+
) {
|
|
29
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
30
|
+
let stride = num_wg.x * WGS;
|
|
31
|
+
|
|
32
|
+
let n_floor = (params.n / stride) * stride;
|
|
33
|
+
let mainIters = n_floor / stride;
|
|
34
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
35
|
+
let id = gid.x + iter * stride;
|
|
36
|
+
let i = id * params.x_inc;
|
|
37
|
+
let result = ddMulProtected(alpha, DD(xHi[i], xLo[i]), lid.x);
|
|
38
|
+
xHi[i] = result.hi;
|
|
39
|
+
xLo[i] = result.lo;
|
|
40
|
+
}
|
|
41
|
+
|
|
42
|
+
// Tail is ragged (0 or 1 extra per thread) — pad to this workgroup's worst
|
|
43
|
+
// case so every thread in the workgroup still calls ddMulProtected the
|
|
44
|
+
// same number of times (its barrier needs that), masking only the write.
|
|
45
|
+
let wgBaseGid = wgid.x * WGS;
|
|
46
|
+
var tailIters = 0u;
|
|
47
|
+
if (n_floor + wgBaseGid < params.n) {
|
|
48
|
+
tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
|
|
49
|
+
}
|
|
50
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
51
|
+
let id = n_floor + gid.x + iter * stride;
|
|
52
|
+
let valid = id < params.n;
|
|
53
|
+
let i = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
54
|
+
let result = ddMulProtected(alpha, DD(xHi[i], xLo[i]), lid.x);
|
|
55
|
+
if (valid) {
|
|
56
|
+
xHi[i] = result.hi;
|
|
57
|
+
xLo[i] = result.lo;
|
|
58
|
+
}
|
|
59
|
+
}
|
|
60
|
+
}
|
|
@@ -0,0 +1,38 @@
|
|
|
1
|
+
// dswap: x <-> y, double-double (Dekker) f64 emulation of sswap. A swap is
|
|
2
|
+
// pure data movement — hi and lo are exchanged verbatim, with no arithmetic
|
|
3
|
+
// 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 sswap.wgsl itself.
|
|
7
|
+
|
|
8
|
+
@group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
|
|
9
|
+
@group(0) @binding(1) var<storage, read_write> 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
|
+
let tempHi = xHi[ix];
|
|
32
|
+
let tempLo = xLo[ix];
|
|
33
|
+
xHi[ix] = yHi[iy];
|
|
34
|
+
xLo[ix] = yLo[iy];
|
|
35
|
+
yHi[iy] = tempHi;
|
|
36
|
+
yLo[iy] = tempLo;
|
|
37
|
+
}
|
|
38
|
+
}
|
|
@@ -75,3 +75,9 @@ fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
|
75
75
|
let loSum = a.lo + b.lo;
|
|
76
76
|
return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
|
|
77
77
|
}
|
|
78
|
+
|
|
79
|
+
// Double-double subtraction — a - b, via exact negation (a sign-bit flip,
|
|
80
|
+
// no rounding) then ddAddProtected. Same protection contract.
|
|
81
|
+
fn ddSubProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
82
|
+
return ddAddProtected(a, DD(negf(b.hi), negf(b.lo)), threadSlot);
|
|
83
|
+
}
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
// Requires f64/dekker.wgsl concatenated first for the DD struct, and
|
|
2
|
+
// f64/utils/add.wgsl (ddSubProtected/ddAddProtected/negf) and
|
|
3
|
+
// f64/utils/multiply.wgsl (ddMulProtected).
|
|
4
|
+
//
|
|
5
|
+
// Verified empirically on real hardware (NVIDIA GTX 1650 + Intel Mesa
|
|
6
|
+
// integrated GPU, ~8000 random trials each plus explicit edge cases): no
|
|
7
|
+
// compiler-reassociation-style corruption of the kind that broke twoSum/
|
|
8
|
+
// twoProd (see add.wgsl's/multiply.wgsl's own headers) — every failure
|
|
9
|
+
// found was a genuine algorithm gap, not a driver miscompile, and both are
|
|
10
|
+
// fixed below (the b.hi==0.0 guard). One real, expected-shape difference
|
|
11
|
+
// from every other protected op here: the low-power backend's observed
|
|
12
|
+
// forward-error factor for this op specifically runs noticeably higher
|
|
13
|
+
// (~7-14000x eps, vs ~3-9x on high-performance) than ddSqrtProtected's
|
|
14
|
+
// (~3-4x on both) — division inherently amplifies input imprecision more
|
|
15
|
+
// than a sum/product does, so a real routine built on this needs its own
|
|
16
|
+
// backend-calibrated threshold, same as every other f64 arithmetic routine
|
|
17
|
+
// in this codebase (see e.g. tests/drot/src/test.drot.js's THRESHOLDS).
|
|
18
|
+
//
|
|
19
|
+
// One Newton-style long-division refinement (Bailey/QD-style): q1 = a.hi /
|
|
20
|
+
// b.hi is a plain f32 quotient, accurate to ~24 bits. Computing the residual
|
|
21
|
+
// a - q1*b in DD arithmetic (not f32) recovers the bits q1 lost, and a
|
|
22
|
+
// second plain division of that residual resolves them into a correction
|
|
23
|
+
// term — combining q1 + q2 gives roughly double a lone f32 divide's
|
|
24
|
+
// precision, matching this scheme's ~48-bit double-double target (already
|
|
25
|
+
// short of real f64's 52 bits, so a second refinement step would chase
|
|
26
|
+
// precision this representation has no room for).
|
|
27
|
+
// b.hi == 0.0 makes q1 = a.hi/0.0 already the IEEE-754-correct answer
|
|
28
|
+
// (±Infinity, or NaN for 0/0) via plain float division, but the refinement
|
|
29
|
+
// below would corrupt it: p1 = q1*b multiplies that Infinity by a zero
|
|
30
|
+
// divisor, and Infinity*0 is NaN by definition, poisoning everything after.
|
|
31
|
+
// Substituting a safe non-zero denominator via select() — rather than
|
|
32
|
+
// branching/returning early — keeps every thread calling ddMulProtected/
|
|
33
|
+
// ddSubProtected/ddAddProtected unconditionally, which their internal
|
|
34
|
+
// workgroupBarrier() requires; only the final result is selected between
|
|
35
|
+
// the refined value and q1's own already-correct answer.
|
|
36
|
+
fn ddDivProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
37
|
+
let bIsZero = b.hi == 0.0;
|
|
38
|
+
let q1 = a.hi / b.hi;
|
|
39
|
+
let safeB = DD(select(b.hi, 1.0, bIsZero), select(b.lo, 0.0, bIsZero));
|
|
40
|
+
let p1 = ddMulProtected(DD(q1, 0.0), safeB, threadSlot);
|
|
41
|
+
let r1 = ddSubProtected(a, p1, threadSlot);
|
|
42
|
+
let q2 = r1.hi / safeB.hi;
|
|
43
|
+
let refined = ddAddProtected(DD(q1, 0.0), DD(q2, 0.0), threadSlot);
|
|
44
|
+
return DD(select(refined.hi, q1, bIsZero), select(refined.lo, 0.0, bIsZero));
|
|
45
|
+
}
|
|
@@ -15,7 +15,11 @@ const SPLIT_CONST: f32 = 4097.0;
|
|
|
15
15
|
|
|
16
16
|
fn bitSplit(a: f32) -> DD {
|
|
17
17
|
let bits = bitcast<u32>(a);
|
|
18
|
-
|
|
18
|
+
// Top 11 mantissa bits, so hi carries 12 significant bits with the implicit
|
|
19
|
+
// leading 1 — the halves are multiplied pairwise and f32 holds 24, so a
|
|
20
|
+
// wider split rounds those products and the "exact" error term goes wrong.
|
|
21
|
+
// Matches SPLIT_CONST = 2^12+1 used by the Veltkamp path below.
|
|
22
|
+
let hiBits = bits & 0xFFFFF000u;
|
|
19
23
|
let hi = bitcast<f32>(hiBits);
|
|
20
24
|
let lo = fsub(a, hi); // exact by Sterbenz's lemma (hi, a share an exponent, are close)
|
|
21
25
|
return DD(hi, lo);
|
|
@@ -63,19 +67,24 @@ fn twoProdFma(a: f32, b: f32) -> DD {
|
|
|
63
67
|
|
|
64
68
|
// DD × DD product (Dekker/Bailey): twoProdBit(a.hi, b.hi) already captures
|
|
65
69
|
// the dominant term to full DD precision, and the cross terms are below the
|
|
66
|
-
// ~48-bit floor anyway, so folding them in with plain f32 loses nothing
|
|
67
|
-
//
|
|
68
|
-
//
|
|
69
|
-
//
|
|
70
|
-
//
|
|
71
|
-
//
|
|
72
|
-
|
|
70
|
+
// ~48-bit floor anyway, so folding them in with plain f32 loses nothing.
|
|
71
|
+
//
|
|
72
|
+
// Another real compiler bug, distinct from add.wgsl's twoSum one — confirmed
|
|
73
|
+
// on Intel Mesa ANV: when p.lo feeds straight into `crossAndLo` unobserved,
|
|
74
|
+
// the compiler folds it away entirely. Materializing p.lo itself through
|
|
75
|
+
// workgroup memory + workgroupBarrier() (like twoSumProtected does for its
|
|
76
|
+
// sum) is what fixes it, so ddMulRaw now takes threadSlot and always pays
|
|
77
|
+
// that barrier — no longer a plain unprotected batchable helper.
|
|
78
|
+
fn ddMulRaw(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
73
79
|
let p = twoProdBit(a.hi, b.hi);
|
|
74
|
-
|
|
80
|
+
dekkerScratch[threadSlot] = p.lo;
|
|
81
|
+
workgroupBarrier();
|
|
82
|
+
let pLo = dekkerScratch[threadSlot];
|
|
83
|
+
let crossAndLo = pLo + (a.hi * b.lo + a.lo * b.hi);
|
|
75
84
|
return DD(p.hi, crossAndLo);
|
|
76
85
|
}
|
|
77
86
|
|
|
78
87
|
fn ddMulProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
79
|
-
let raw = ddMulRaw(a, b);
|
|
88
|
+
let raw = ddMulRaw(a, b, threadSlot);
|
|
80
89
|
return fastTwoSumProtected(raw.hi, raw.lo, threadSlot);
|
|
81
90
|
}
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
// Requires f64/dekker.wgsl concatenated first for the DD struct, and
|
|
2
|
+
// f64/utils/add.wgsl (ddSubProtected/ddAddProtected) and
|
|
3
|
+
// f64/utils/multiply.wgsl (twoProdBit — squaring a plain f32 needs no
|
|
4
|
+
// barrier, per multiply.wgsl's own note that twoProdBit is universally safe
|
|
5
|
+
// unprotected).
|
|
6
|
+
//
|
|
7
|
+
// Verified empirically on real hardware (NVIDIA GTX 1650 + Intel Mesa
|
|
8
|
+
// integrated GPU, ~8000 random trials each plus explicit edge cases): no
|
|
9
|
+
// compiler-reassociation-style corruption of the kind that broke twoSum/
|
|
10
|
+
// twoProd (see add.wgsl's/multiply.wgsl's own headers) — every failure
|
|
11
|
+
// found was a genuine algorithm gap (the a.hi==0.0 case below), not a
|
|
12
|
+
// driver miscompile, and is fixed. Observed forward-error factor against a
|
|
13
|
+
// true f64 reference stayed ~3-4x eps on both backends across every random
|
|
14
|
+
// trial — noticeably tighter than ddDivProtected's own low-power spread
|
|
15
|
+
// (see divide.wgsl's header), since sqrt has no denominator to be unlucky
|
|
16
|
+
// about.
|
|
17
|
+
//
|
|
18
|
+
// One Newton refinement step (the classic extended-precision sqrt trick):
|
|
19
|
+
// x0 = sqrt(a.hi) is a plain f32 approximation; the residual a - x0^2,
|
|
20
|
+
// computed in DD arithmetic, recovers what x0 lost, and linearizing sqrt
|
|
21
|
+
// around x0 (dividing that residual by 2*x0) gives a correction term
|
|
22
|
+
// roughly doubling the precision — same ~48-bit target as ddDivProtected,
|
|
23
|
+
// so one step is enough.
|
|
24
|
+
//
|
|
25
|
+
// Undefined for a.hi < 0.0, same as plain sqrt() — callers must guard
|
|
26
|
+
// themselves; this never checks.
|
|
27
|
+
//
|
|
28
|
+
// a.hi == 0.0 (a genuinely zero input, not an underflowed one — zero is
|
|
29
|
+
// exactly representable in f32, unlike this scheme's real range limits;
|
|
30
|
+
// see splitDoubleDouble's own doc comment) makes x0 = sqrt(0) = 0, and the
|
|
31
|
+
// correction step would divide by 2*x0 = 0. Substituting a safe non-zero
|
|
32
|
+
// denominator via select() — rather than branching/returning early — keeps
|
|
33
|
+
// every thread calling ddSubProtected/ddAddProtected unconditionally, which
|
|
34
|
+
// their internal workgroupBarrier() requires; only the final result is
|
|
35
|
+
// selected between the computed value and the exact DD(0,0) answer.
|
|
36
|
+
fn ddSqrtProtected(a: DD, threadSlot: u32) -> DD {
|
|
37
|
+
let isZero = a.hi == 0.0;
|
|
38
|
+
let x0 = sqrt(a.hi);
|
|
39
|
+
let x0sq = twoProdBit(x0, x0);
|
|
40
|
+
let r = ddSubProtected(a, x0sq, threadSlot);
|
|
41
|
+
let safeDenom = select(2.0 * x0, 1.0, isZero);
|
|
42
|
+
let correction = r.hi / safeDenom;
|
|
43
|
+
let result = ddAddProtected(DD(x0, 0.0), DD(correction, 0.0), threadSlot);
|
|
44
|
+
return DD(select(result.hi, 0.0, isZero), select(result.lo, 0.0, isZero));
|
|
45
|
+
}
|