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
package/dist/wgblas.browser.js
CHANGED
|
@@ -1,173 +1,104 @@
|
|
|
1
|
-
var wgblas=(()=>{var
|
|
2
|
-
// dispatch: 1 workgroup of WGS threads.
|
|
3
|
-
// partials_val and partials_idx must have exactly 2*WGS entries.
|
|
4
|
-
|
|
5
|
-
@group(0) @binding(0) var<storage, read> partials_val: array<f32>;
|
|
6
|
-
@group(0) @binding(1) var<storage, read> partials_idx: array<u32>;
|
|
7
|
-
@group(0) @binding(2) var<storage, read_write> result: array<u32>;
|
|
8
|
-
|
|
9
|
-
const WGS: u32 = 64;
|
|
10
|
-
|
|
11
|
-
var<workgroup> tile_val: array<f32, 64>;
|
|
12
|
-
var<workgroup> tile_idx: array<u32, 64>;
|
|
13
|
-
|
|
14
|
-
@compute @workgroup_size(64)
|
|
15
|
-
fn reduce(
|
|
16
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
17
|
-
) {
|
|
18
|
-
let i = lid.x;
|
|
19
|
-
let a_val = partials_val[i];
|
|
20
|
-
let b_val = partials_val[i + WGS];
|
|
21
|
-
if (b_val > a_val || (b_val == a_val && partials_idx[i + WGS] < partials_idx[i])) {
|
|
22
|
-
tile_val[i] = b_val;
|
|
23
|
-
tile_idx[i] = partials_idx[i + WGS];
|
|
24
|
-
} else {
|
|
25
|
-
tile_val[i] = a_val;
|
|
26
|
-
tile_idx[i] = partials_idx[i];
|
|
27
|
-
}
|
|
28
|
-
workgroupBarrier();
|
|
1
|
+
var wgblas=(()=>{var ka=Object.create;var pe=Object.defineProperty;var Da=Object.getOwnPropertyDescriptor;var Na=Object.getOwnPropertyNames;var Pa=Object.getPrototypeOf,Ma=Object.prototype.hasOwnProperty;var ge=(r=>typeof require<"u"?require:typeof Proxy<"u"?new Proxy(r,{get:(t,e)=>(typeof require<"u"?require:t)[e]}):r)(function(r){if(typeof require<"u")return require.apply(this,arguments);throw Error('Dynamic require of "'+r+'" is not supported')});var O=(r,t,e)=>()=>{if(e)throw e[0];try{return r&&(t=r(r=0)),t}catch(o){throw e=[o],o}};var qe=(r,t)=>{for(var e in t)pe(r,e,{get:t[e],enumerable:!0})},Te=(r,t,e,o)=>{if(t&&typeof t=="object"||typeof t=="function")for(let a of Na(t))!Ma.call(r,a)&&a!==e&&pe(r,a,{get:()=>t[a],enumerable:!(o=Da(t,a))||o.enumerable});return r};var we=(r,t,e)=>(e=r!=null?ka(Pa(r)):{},Te(t||!r||!r.__esModule?pe(e,"default",{value:r,enumerable:!0}):e,r)),Ia=r=>Te(pe({},"__esModule",{value:!0}),r);var ke,$e=O(()=>{ke=`// sscal: x = alpha * x
|
|
29
2
|
|
|
30
|
-
|
|
31
|
-
if (i < s) {
|
|
32
|
-
let c_val = tile_val[i];
|
|
33
|
-
let d_val = tile_val[i + s];
|
|
34
|
-
if (d_val > c_val || (d_val == c_val && tile_idx[i + s] < tile_idx[i])) {
|
|
35
|
-
tile_val[i] = d_val;
|
|
36
|
-
tile_idx[i] = tile_idx[i + s];
|
|
37
|
-
}
|
|
38
|
-
}
|
|
39
|
-
workgroupBarrier();
|
|
40
|
-
}
|
|
3
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
41
4
|
|
|
42
|
-
|
|
5
|
+
struct Params {
|
|
6
|
+
n: u32,
|
|
7
|
+
alpha: f32,
|
|
8
|
+
x_inc: u32,
|
|
43
9
|
}
|
|
44
|
-
`});var be,ge=O(()=>{be=`// amax reduction (f64, double-double): collapses 2*WGS (value, index) pairs
|
|
45
|
-
// into one index, using ddGreater/ddEqual instead of plain f32 \`>\`/\`==\` (see
|
|
46
|
-
// reduction/argmax.wgsl for the f32 original this mirrors).
|
|
47
|
-
// dispatch: 1 workgroup of WGS threads. partialsValHi/partialsValLo/
|
|
48
|
-
// partialsIdx must have exactly 2*WGS entries each. Concatenated after
|
|
49
|
-
// f64/dekker.wgsl (DD struct), f64/utils/greater.wgsl (ddGreater), and
|
|
50
|
-
// f64/utils/equal.wgsl (ddEqual).
|
|
51
10
|
|
|
52
|
-
@group(0) @binding(
|
|
53
|
-
@group(0) @binding(1) var<storage, read> partialsValLo: array<f32>;
|
|
54
|
-
@group(0) @binding(2) var<storage, read> partialsIdx: array<u32>;
|
|
55
|
-
@group(0) @binding(3) var<storage, read_write> result: array<u32>;
|
|
11
|
+
@group(0) @binding(1) var<uniform> params: Params;
|
|
56
12
|
|
|
57
13
|
const WGS: u32 = 64;
|
|
58
14
|
|
|
59
|
-
var<workgroup> tile_val: array<DD, 64>;
|
|
60
|
-
var<workgroup> tile_idx: array<u32, 64>;
|
|
61
|
-
|
|
62
15
|
@compute @workgroup_size(64)
|
|
63
|
-
fn
|
|
64
|
-
@builtin(
|
|
16
|
+
fn main(
|
|
17
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
18
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
65
19
|
) {
|
|
66
|
-
|
|
67
|
-
|
|
68
|
-
let b_val = DD(partialsValHi[i + WGS], partialsValLo[i + WGS]);
|
|
69
|
-
if (ddGreater(b_val, a_val) ||
|
|
70
|
-
(ddEqual(b_val, a_val) && partialsIdx[i + WGS] < partialsIdx[i])) {
|
|
71
|
-
tile_val[i] = b_val;
|
|
72
|
-
tile_idx[i] = partialsIdx[i + WGS];
|
|
73
|
-
} else {
|
|
74
|
-
tile_val[i] = a_val;
|
|
75
|
-
tile_idx[i] = partialsIdx[i];
|
|
20
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
21
|
+
x[id * params.x_inc] = params.alpha * x[id * params.x_inc];
|
|
76
22
|
}
|
|
77
|
-
|
|
23
|
+
}
|
|
24
|
+
`});var Qe,Ze=O(()=>{Qe=`// cscal: x := alpha * x, complex. x is one interleaved f32 array
|
|
25
|
+
// (re0, im0, re1, im1, ...), matching Complex32Array/GpuVector's storage
|
|
26
|
+
// (and cuBLAS's cuComplex / stdlib's Complex64Array) \u2014 no repacking needed
|
|
27
|
+
// between JS and GPU.
|
|
28
|
+
// (alphaRe + i*alphaIm)(re + i*im) = (alphaRe*re - alphaIm*im) + i*(alphaRe*im + alphaIm*re)
|
|
78
29
|
|
|
79
|
-
|
|
80
|
-
if (i < s) {
|
|
81
|
-
let c_val = tile_val[i];
|
|
82
|
-
let d_val = tile_val[i + s];
|
|
83
|
-
if (ddGreater(d_val, c_val) ||
|
|
84
|
-
(ddEqual(d_val, c_val) && tile_idx[i + s] < tile_idx[i])) {
|
|
85
|
-
tile_val[i] = d_val;
|
|
86
|
-
tile_idx[i] = tile_idx[i + s];
|
|
87
|
-
}
|
|
88
|
-
}
|
|
89
|
-
workgroupBarrier();
|
|
90
|
-
}
|
|
30
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
91
31
|
|
|
92
|
-
|
|
32
|
+
struct Params {
|
|
33
|
+
n: u32,
|
|
34
|
+
alphaRe: f32,
|
|
35
|
+
alphaIm: f32,
|
|
36
|
+
x_inc: u32,
|
|
93
37
|
}
|
|
94
|
-
`});var xe,he=O(()=>{xe=`// sum reduction: collapses 2*WGS partials into one scalar.
|
|
95
|
-
// dispatch: 1 workgroup of WGS threads.
|
|
96
|
-
// partials must have exactly 2*WGS entries.
|
|
97
38
|
|
|
98
|
-
@group(0) @binding(
|
|
99
|
-
@group(0) @binding(1) var<storage, read_write> result: array<f32>;
|
|
39
|
+
@group(0) @binding(1) var<uniform> params: Params;
|
|
100
40
|
|
|
101
41
|
const WGS: u32 = 64;
|
|
102
42
|
|
|
103
|
-
var<workgroup> tile: array<f32, 64>;
|
|
104
|
-
|
|
105
43
|
@compute @workgroup_size(64)
|
|
106
|
-
fn
|
|
107
|
-
@builtin(
|
|
44
|
+
fn main(
|
|
45
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
46
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
108
47
|
) {
|
|
109
|
-
|
|
110
|
-
|
|
111
|
-
|
|
112
|
-
|
|
113
|
-
|
|
114
|
-
|
|
115
|
-
|
|
48
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
49
|
+
let base = 2u * id * params.x_inc;
|
|
50
|
+
// Both new parts need both old parts, so capture them before either write.
|
|
51
|
+
let re = x[base];
|
|
52
|
+
let im = x[base + 1u];
|
|
53
|
+
x[base] = params.alphaRe * re - params.alphaIm * im;
|
|
54
|
+
x[base + 1u] = params.alphaRe * im + params.alphaIm * re;
|
|
116
55
|
}
|
|
56
|
+
}
|
|
57
|
+
`});var rt,Je=O(()=>{rt=`// sswap: x <-> y
|
|
117
58
|
|
|
118
|
-
|
|
59
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
60
|
+
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
61
|
+
|
|
62
|
+
struct Params {
|
|
63
|
+
n: u32,
|
|
64
|
+
x_inc: u32,
|
|
65
|
+
y_inc: u32,
|
|
119
66
|
}
|
|
120
|
-
`});var ye,ve=O(()=>{ye=`// sum reduction (f64, double-double): collapses 2*WGS partial (hi, lo) pairs
|
|
121
|
-
// into one, using ddAddProtected instead of plain f32 \`+\` (see
|
|
122
|
-
// reduction/sum.wgsl for the f32 original this mirrors).
|
|
123
|
-
// dispatch: 1 workgroup of WGS threads. partialsHi/partialsLo must have
|
|
124
|
-
// exactly 2*WGS entries each. Concatenated after f64/dekker.wgsl (DD struct)
|
|
125
|
-
// and f64/utils/add.wgsl (ddAddProtected \u2014 see it for why plain ddAdd isn't safe).
|
|
126
67
|
|
|
127
|
-
@group(0) @binding(
|
|
128
|
-
@group(0) @binding(1) var<storage, read> partialsLo: array<f32>;
|
|
129
|
-
@group(0) @binding(2) var<storage, read_write> resultHi: array<f32, 1>;
|
|
130
|
-
@group(0) @binding(3) var<storage, read_write> resultLo: array<f32, 1>;
|
|
68
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
131
69
|
|
|
132
70
|
const WGS: u32 = 64;
|
|
133
71
|
|
|
134
|
-
var<workgroup> tile: array<DD, 64>;
|
|
135
|
-
|
|
136
72
|
@compute @workgroup_size(64)
|
|
137
|
-
fn
|
|
138
|
-
@builtin(
|
|
73
|
+
fn main(
|
|
74
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
75
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
139
76
|
) {
|
|
140
|
-
|
|
141
|
-
|
|
142
|
-
|
|
143
|
-
|
|
144
|
-
workgroupBarrier();
|
|
145
|
-
|
|
146
|
-
// ddAddProtected must be called unconditionally by every thread.
|
|
147
|
-
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
148
|
-
let partner = select(i, i + s, i < s);
|
|
149
|
-
let combined = ddAddProtected(tile[i], tile[partner], i);
|
|
150
|
-
workgroupBarrier();
|
|
151
|
-
if (i < s) { tile[i] = combined; }
|
|
152
|
-
workgroupBarrier();
|
|
153
|
-
}
|
|
154
|
-
|
|
155
|
-
if (i == 0u) {
|
|
156
|
-
resultHi[0] = tile[0].hi;
|
|
157
|
-
resultLo[0] = tile[0].lo;
|
|
77
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
78
|
+
let temp = x[id * params.x_inc];
|
|
79
|
+
x[id * params.x_inc] = y[id * params.y_inc];
|
|
80
|
+
y[id * params.y_inc] = temp;
|
|
158
81
|
}
|
|
159
82
|
}
|
|
160
|
-
`});var
|
|
83
|
+
`});var tt,et=O(()=>{tt=`// dswap: x <-> y, double-double (Dekker) f64 emulation of sswap. A swap is
|
|
84
|
+
// pure data movement \u2014 hi and lo are exchanged verbatim, with no arithmetic
|
|
85
|
+
// at all \u2014 so (unlike dscal/daxpy/ddot) this needs no
|
|
86
|
+
// ddMulProtected/ddAddProtected renormalizing barrier, and so no
|
|
87
|
+
// ragged-tail-with-barrier split; a plain grid-stride loop is safe, same
|
|
88
|
+
// shape as sswap.wgsl itself.
|
|
161
89
|
|
|
162
|
-
@group(0) @binding(0) var<storage, read_write>
|
|
90
|
+
@group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
|
|
91
|
+
@group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
|
|
92
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
93
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
163
94
|
|
|
164
95
|
struct Params {
|
|
165
96
|
n: u32,
|
|
166
|
-
alpha: f32,
|
|
167
97
|
x_inc: u32,
|
|
98
|
+
y_inc: u32,
|
|
168
99
|
}
|
|
169
100
|
|
|
170
|
-
@group(0) @binding(
|
|
101
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
171
102
|
|
|
172
103
|
const WGS: u32 = 64;
|
|
173
104
|
|
|
@@ -177,16 +108,24 @@ fn main(
|
|
|
177
108
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
178
109
|
) {
|
|
179
110
|
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
180
|
-
|
|
111
|
+
let ix = id * params.x_inc;
|
|
112
|
+
let iy = id * params.y_inc;
|
|
113
|
+
let tempHi = xHi[ix];
|
|
114
|
+
let tempLo = xLo[ix];
|
|
115
|
+
xHi[ix] = yHi[iy];
|
|
116
|
+
xLo[ix] = yLo[iy];
|
|
117
|
+
yHi[iy] = tempHi;
|
|
118
|
+
yLo[iy] = tempLo;
|
|
181
119
|
}
|
|
182
120
|
}
|
|
183
|
-
`});var
|
|
121
|
+
`});var at,ot=O(()=>{at=`// saxpy: y = alpha * x + y
|
|
184
122
|
|
|
185
|
-
@group(0) @binding(0) var<storage,
|
|
123
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
186
124
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
187
125
|
|
|
188
126
|
struct Params {
|
|
189
127
|
n: u32,
|
|
128
|
+
alpha: f32,
|
|
190
129
|
x_inc: u32,
|
|
191
130
|
y_inc: u32,
|
|
192
131
|
}
|
|
@@ -201,19 +140,16 @@ fn main(
|
|
|
201
140
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
202
141
|
) {
|
|
203
142
|
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
204
|
-
|
|
205
|
-
x[id * params.x_inc] = y[id * params.y_inc];
|
|
206
|
-
y[id * params.y_inc] = temp;
|
|
143
|
+
y[id * params.y_inc] = params.alpha * x[id * params.x_inc] + y[id * params.y_inc];
|
|
207
144
|
}
|
|
208
145
|
}
|
|
209
|
-
`});var
|
|
146
|
+
`});var st,it=O(()=>{st=`// scopy: y = x
|
|
210
147
|
|
|
211
148
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
212
149
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
213
150
|
|
|
214
151
|
struct Params {
|
|
215
152
|
n: u32,
|
|
216
|
-
alpha: f32,
|
|
217
153
|
x_inc: u32,
|
|
218
154
|
y_inc: u32,
|
|
219
155
|
}
|
|
@@ -228,13 +164,20 @@ fn main(
|
|
|
228
164
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
229
165
|
) {
|
|
230
166
|
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
231
|
-
y[id * params.y_inc] =
|
|
167
|
+
y[id * params.y_inc] = x[id * params.x_inc];
|
|
232
168
|
}
|
|
233
169
|
}
|
|
234
|
-
`});var
|
|
170
|
+
`});var lt,nt=O(()=>{lt=`// dcopy: y = x, double-double (Dekker) f64 emulation of scopy. A copy is
|
|
171
|
+
// pure data movement \u2014 hi and lo are transferred verbatim, with no
|
|
172
|
+
// arithmetic at all \u2014 so (unlike dscal/daxpy/ddot) this needs no
|
|
173
|
+
// ddMulProtected/ddAddProtected renormalizing barrier, and so no
|
|
174
|
+
// ragged-tail-with-barrier split; a plain grid-stride loop is safe, same
|
|
175
|
+
// shape as scopy.wgsl itself.
|
|
235
176
|
|
|
236
|
-
@group(0) @binding(0) var<storage, read>
|
|
237
|
-
@group(0) @binding(1) var<storage,
|
|
177
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
178
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
179
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
180
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
238
181
|
|
|
239
182
|
struct Params {
|
|
240
183
|
n: u32,
|
|
@@ -242,7 +185,7 @@ struct Params {
|
|
|
242
185
|
y_inc: u32,
|
|
243
186
|
}
|
|
244
187
|
|
|
245
|
-
@group(0) @binding(
|
|
188
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
246
189
|
|
|
247
190
|
const WGS: u32 = 64;
|
|
248
191
|
|
|
@@ -252,10 +195,13 @@ fn main(
|
|
|
252
195
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
253
196
|
) {
|
|
254
197
|
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
255
|
-
|
|
198
|
+
let ix = id * params.x_inc;
|
|
199
|
+
let iy = id * params.y_inc;
|
|
200
|
+
yHi[iy] = xHi[ix];
|
|
201
|
+
yLo[iy] = xLo[ix];
|
|
256
202
|
}
|
|
257
203
|
}
|
|
258
|
-
`});var
|
|
204
|
+
`});var ft,ut=O(()=>{ft=`// sdot: result = sum(x[i] * y[i])
|
|
259
205
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sum.wgsl.
|
|
260
206
|
|
|
261
207
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -308,7 +254,33 @@ fn main(
|
|
|
308
254
|
|
|
309
255
|
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
310
256
|
}
|
|
311
|
-
`});var
|
|
257
|
+
`});var De,mt=O(()=>{De=`// sum reduction: collapses 2*WGS partials into one scalar.
|
|
258
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
259
|
+
// partials must have exactly 2*WGS entries.
|
|
260
|
+
|
|
261
|
+
@group(0) @binding(0) var<storage, read> partials: array<f32>;
|
|
262
|
+
@group(0) @binding(1) var<storage, read_write> result: array<f32>;
|
|
263
|
+
|
|
264
|
+
const WGS: u32 = 64;
|
|
265
|
+
|
|
266
|
+
var<workgroup> tile: array<f32, 64>;
|
|
267
|
+
|
|
268
|
+
@compute @workgroup_size(64)
|
|
269
|
+
fn reduce(
|
|
270
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
271
|
+
) {
|
|
272
|
+
let i = lid.x;
|
|
273
|
+
tile[i] = partials[i] + partials[i + WGS];
|
|
274
|
+
workgroupBarrier();
|
|
275
|
+
|
|
276
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
277
|
+
if (i < s) { tile[i] += tile[i + s]; }
|
|
278
|
+
workgroupBarrier();
|
|
279
|
+
}
|
|
280
|
+
|
|
281
|
+
if (i == 0u) { result[0] = tile[0]; }
|
|
282
|
+
}
|
|
283
|
+
`});var ct,dt=O(()=>{ct=`// sasum: result = sum(|x[i]|)
|
|
312
284
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/abssum.wgsl.
|
|
313
285
|
|
|
314
286
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -359,12 +331,25 @@ fn main(
|
|
|
359
331
|
|
|
360
332
|
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
361
333
|
}
|
|
362
|
-
`});var
|
|
363
|
-
//
|
|
364
|
-
|
|
365
|
-
|
|
366
|
-
|
|
367
|
-
|
|
334
|
+
`});var gt,pt=O(()=>{gt=`// snrm2: result = sqrt(sum(x[i] * x[i])), computed via scaled accumulation
|
|
335
|
+
// (Blue's algorithm / reference BLAS's SLASSQ) rather than naive squaring \u2014
|
|
336
|
+
// naive \`sum += x_i * x_i\` overflows to inf for |x_i| \u2273 1.8e19 (f32's
|
|
337
|
+
// squaring range is only sqrt(f32_max)) and loses precision on tiny
|
|
338
|
+
// magnitudes squaring into the denormal range. Running state is (scale,
|
|
339
|
+
// ssq) with true-sum-of-squares == scale\xB2 \xB7 ssq: scale tracks the largest
|
|
340
|
+
// |x_i| seen so far, and every other contribution is expressed *relative
|
|
341
|
+
// to* scale (never squared in absolute terms), so ssq stays near 1
|
|
342
|
+
// regardless of x's magnitude range. Merging two independent partials
|
|
343
|
+
// (ssqMerge) is associative, so this composes with the same 4-way-ILP +
|
|
344
|
+
// tree-reduction shape every other Level 1 reduction here uses \u2014 see
|
|
345
|
+
// reduction/scaledSum.wgsl for the pass-2 counterpart, which finishes with
|
|
346
|
+
// scale\xB7sqrt(ssq).
|
|
347
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/scaledSum.wgsl.
|
|
348
|
+
|
|
349
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
350
|
+
@group(0) @binding(1) var<storage, read_write> partialsScale: array<f32>;
|
|
351
|
+
@group(0) @binding(2) var<storage, read_write> partialsSsq: array<f32>;
|
|
352
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
368
353
|
|
|
369
354
|
struct Params {
|
|
370
355
|
n: u32,
|
|
@@ -373,7 +358,36 @@ struct Params {
|
|
|
373
358
|
|
|
374
359
|
const WGS: u32 = 64;
|
|
375
360
|
|
|
376
|
-
|
|
361
|
+
struct ScaleSsq {
|
|
362
|
+
scale: f32,
|
|
363
|
+
ssq: f32,
|
|
364
|
+
}
|
|
365
|
+
|
|
366
|
+
// Folds one more |value| into a running (scale, ssq) pair.
|
|
367
|
+
fn ssqAccum(acc: ScaleSsq, absxi: f32) -> ScaleSsq {
|
|
368
|
+
if (absxi == 0.0) { return acc; }
|
|
369
|
+
if (absxi > acc.scale) {
|
|
370
|
+
let r = acc.scale / absxi; // 0/absxi == 0 on the first nonzero value \u2014 safe
|
|
371
|
+
return ScaleSsq(absxi, 1.0 + acc.ssq * r * r);
|
|
372
|
+
}
|
|
373
|
+
let r = absxi / acc.scale; // reached only once acc.scale > 0 (absxi <= acc.scale and absxi > 0)
|
|
374
|
+
return ScaleSsq(acc.scale, acc.ssq + r * r);
|
|
375
|
+
}
|
|
376
|
+
|
|
377
|
+
// Associative merge of two independent (scale, ssq) partials \u2014 lets this
|
|
378
|
+
// compose with a tree reduction exactly like a plain sum would.
|
|
379
|
+
fn ssqMerge(a: ScaleSsq, b: ScaleSsq) -> ScaleSsq {
|
|
380
|
+
if (a.scale == 0.0 && b.scale == 0.0) { return ScaleSsq(0.0, 1.0); }
|
|
381
|
+
if (a.scale >= b.scale) {
|
|
382
|
+
let r = b.scale / a.scale;
|
|
383
|
+
return ScaleSsq(a.scale, a.ssq + b.ssq * r * r);
|
|
384
|
+
}
|
|
385
|
+
let r = a.scale / b.scale;
|
|
386
|
+
return ScaleSsq(b.scale, b.ssq + a.ssq * r * r);
|
|
387
|
+
}
|
|
388
|
+
|
|
389
|
+
var<workgroup> tileScale: array<f32, 64>;
|
|
390
|
+
var<workgroup> tileSsq: array<f32, 64>;
|
|
377
391
|
|
|
378
392
|
@compute @workgroup_size(64)
|
|
379
393
|
fn main(
|
|
@@ -382,119 +396,112 @@ fn main(
|
|
|
382
396
|
@builtin(workgroup_id) wgid: vec3u,
|
|
383
397
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
384
398
|
) {
|
|
385
|
-
var acc0
|
|
386
|
-
var acc1
|
|
387
|
-
var acc2
|
|
388
|
-
var acc3
|
|
399
|
+
var acc0 = ScaleSsq(0.0, 1.0);
|
|
400
|
+
var acc1 = ScaleSsq(0.0, 1.0);
|
|
401
|
+
var acc2 = ScaleSsq(0.0, 1.0);
|
|
402
|
+
var acc3 = ScaleSsq(0.0, 1.0);
|
|
389
403
|
|
|
390
404
|
let stride = num_wg.x * WGS;
|
|
391
405
|
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
392
406
|
|
|
393
407
|
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
394
|
-
|
|
395
|
-
|
|
396
|
-
|
|
397
|
-
|
|
398
|
-
acc0 += v0 * v0;
|
|
399
|
-
acc1 += v1 * v1;
|
|
400
|
-
acc2 += v2 * v2;
|
|
401
|
-
acc3 += v3 * v3;
|
|
408
|
+
acc0 = ssqAccum(acc0, abs(x[ id * params.x_inc]));
|
|
409
|
+
acc1 = ssqAccum(acc1, abs(x[(id + stride) * params.x_inc]));
|
|
410
|
+
acc2 = ssqAccum(acc2, abs(x[(id + 2u * stride) * params.x_inc]));
|
|
411
|
+
acc3 = ssqAccum(acc3, abs(x[(id + 3u * stride) * params.x_inc]));
|
|
402
412
|
}
|
|
403
413
|
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
404
|
-
|
|
405
|
-
acc0 += v * v;
|
|
414
|
+
acc0 = ssqAccum(acc0, abs(x[id * params.x_inc]));
|
|
406
415
|
}
|
|
407
416
|
|
|
408
|
-
|
|
417
|
+
let combined = ssqMerge(ssqMerge(acc0, acc1), ssqMerge(acc2, acc3));
|
|
418
|
+
tileScale[lid.x] = combined.scale;
|
|
419
|
+
tileSsq[lid.x] = combined.ssq;
|
|
409
420
|
workgroupBarrier();
|
|
410
421
|
|
|
411
422
|
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
412
|
-
if (lid.x < s) {
|
|
423
|
+
if (lid.x < s) {
|
|
424
|
+
let merged = ssqMerge(
|
|
425
|
+
ScaleSsq(tileScale[lid.x], tileSsq[lid.x]),
|
|
426
|
+
ScaleSsq(tileScale[lid.x + s], tileSsq[lid.x + s]),
|
|
427
|
+
);
|
|
428
|
+
tileScale[lid.x] = merged.scale;
|
|
429
|
+
tileSsq[lid.x] = merged.ssq;
|
|
430
|
+
}
|
|
413
431
|
workgroupBarrier();
|
|
414
432
|
}
|
|
415
433
|
|
|
416
|
-
if (lid.x == 0u) {
|
|
434
|
+
if (lid.x == 0u) {
|
|
435
|
+
partialsScale[wgid.x] = tileScale[0];
|
|
436
|
+
partialsSsq[wgid.x] = tileSsq[0];
|
|
437
|
+
}
|
|
417
438
|
}
|
|
418
|
-
`});var
|
|
439
|
+
`});var ht,wt=O(()=>{ht=`// scaledSum reduction: collapses 2*WGS (scale, ssq) partials from
|
|
440
|
+
// snrm2.wgsl into the final norm \u2014 sqrt(scale\xB2 \xB7 ssq) == scale \xB7 sqrt(ssq).
|
|
441
|
+
// Mirrors reduction/sum.wgsl's shape exactly, merging via ssqMerge (see
|
|
442
|
+
// snrm2.wgsl for the derivation) instead of plain \`+\`, and taking the final
|
|
443
|
+
// sqrt here rather than on the CPU \u2014 unlike sasum/sdot's plain sum, "sum of
|
|
444
|
+
// squares" isn't a meaningful standalone value to hand back, only
|
|
445
|
+
// scale\xB7sqrt(ssq) is.
|
|
446
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
447
|
+
// partialsScale/partialsSsq must have exactly 2*WGS entries each.
|
|
419
448
|
|
|
420
|
-
@group(0) @binding(0) var<storage,
|
|
421
|
-
@group(0) @binding(1) var<storage,
|
|
449
|
+
@group(0) @binding(0) var<storage, read> partialsScale: array<f32>;
|
|
450
|
+
@group(0) @binding(1) var<storage, read> partialsSsq: array<f32>;
|
|
451
|
+
@group(0) @binding(2) var<storage, read_write> result: array<f32>;
|
|
422
452
|
|
|
423
|
-
|
|
424
|
-
|
|
425
|
-
|
|
426
|
-
|
|
427
|
-
|
|
428
|
-
|
|
453
|
+
const WGS: u32 = 64;
|
|
454
|
+
|
|
455
|
+
// True sum-of-squares represented so far == scale\xB2 \xB7 ssq \u2014 see snrm2.wgsl.
|
|
456
|
+
struct ScaleSsq {
|
|
457
|
+
scale: f32,
|
|
458
|
+
ssq: f32,
|
|
429
459
|
}
|
|
430
460
|
|
|
431
|
-
|
|
461
|
+
// Associative merge of two independent (scale, ssq) partials.
|
|
462
|
+
fn ssqMerge(a: ScaleSsq, b: ScaleSsq) -> ScaleSsq {
|
|
463
|
+
if (a.scale == 0.0 && b.scale == 0.0) { return ScaleSsq(0.0, 1.0); }
|
|
464
|
+
if (a.scale >= b.scale) {
|
|
465
|
+
let r = b.scale / a.scale;
|
|
466
|
+
return ScaleSsq(a.scale, a.ssq + b.ssq * r * r);
|
|
467
|
+
}
|
|
468
|
+
let r = a.scale / b.scale;
|
|
469
|
+
return ScaleSsq(b.scale, b.ssq + a.ssq * r * r);
|
|
470
|
+
}
|
|
432
471
|
|
|
433
|
-
|
|
472
|
+
var<workgroup> tileScale: array<f32, 64>;
|
|
473
|
+
var<workgroup> tileSsq: array<f32, 64>;
|
|
434
474
|
|
|
435
475
|
@compute @workgroup_size(64)
|
|
436
|
-
fn
|
|
437
|
-
@builtin(
|
|
438
|
-
@builtin(num_workgroups) num_wg: vec3u,
|
|
476
|
+
fn reduce_scaled(
|
|
477
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
439
478
|
) {
|
|
440
|
-
|
|
441
|
-
|
|
442
|
-
|
|
443
|
-
|
|
444
|
-
|
|
445
|
-
|
|
446
|
-
|
|
447
|
-
|
|
448
|
-
// param[0] = flag: -1 (full H), 0 (unit diagonal), 1 (unit off-diagonal)
|
|
449
|
-
// param = [ flag, h11, h21, h12, h22 ]
|
|
450
|
-
// flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
|
|
451
|
-
|
|
452
|
-
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
453
|
-
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
454
|
-
@group(0) @binding(2) var<storage, read> param: array<f32>;
|
|
455
|
-
|
|
456
|
-
struct Params {
|
|
457
|
-
n: u32,
|
|
458
|
-
x_inc: u32,
|
|
459
|
-
y_inc: u32,
|
|
460
|
-
}
|
|
461
|
-
|
|
462
|
-
@group(0) @binding(3) var<uniform> params: Params;
|
|
463
|
-
|
|
464
|
-
const WGS: u32 = 64;
|
|
465
|
-
|
|
466
|
-
@compute @workgroup_size(64)
|
|
467
|
-
fn main(
|
|
468
|
-
@builtin(global_invocation_id) gid: vec3u,
|
|
469
|
-
@builtin(num_workgroups) num_wg: vec3u,
|
|
470
|
-
) {
|
|
471
|
-
let flag = param[0];
|
|
472
|
-
|
|
473
|
-
var h11: f32; var h12: f32;
|
|
474
|
-
var h21: f32; var h22: f32;
|
|
479
|
+
let i = lid.x;
|
|
480
|
+
let merged0 = ssqMerge(
|
|
481
|
+
ScaleSsq(partialsScale[i], partialsSsq[i]),
|
|
482
|
+
ScaleSsq(partialsScale[i + WGS], partialsSsq[i + WGS]),
|
|
483
|
+
);
|
|
484
|
+
tileScale[i] = merged0.scale;
|
|
485
|
+
tileSsq[i] = merged0.ssq;
|
|
486
|
+
workgroupBarrier();
|
|
475
487
|
|
|
476
|
-
|
|
477
|
-
|
|
478
|
-
|
|
479
|
-
|
|
480
|
-
|
|
481
|
-
|
|
482
|
-
|
|
483
|
-
|
|
484
|
-
|
|
485
|
-
|
|
486
|
-
h11 = param[1]; h21 = -1.0;
|
|
487
|
-
h12 = 1.0; h22 = param[4];
|
|
488
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
489
|
+
if (i < s) {
|
|
490
|
+
let merged = ssqMerge(
|
|
491
|
+
ScaleSsq(tileScale[i], tileSsq[i]),
|
|
492
|
+
ScaleSsq(tileScale[i + s], tileSsq[i + s]),
|
|
493
|
+
);
|
|
494
|
+
tileScale[i] = merged.scale;
|
|
495
|
+
tileSsq[i] = merged.ssq;
|
|
496
|
+
}
|
|
497
|
+
workgroupBarrier();
|
|
488
498
|
}
|
|
489
499
|
|
|
490
|
-
|
|
491
|
-
|
|
492
|
-
let yi = y[id * params.y_inc];
|
|
493
|
-
x[id * params.x_inc] = h11 * xi + h12 * yi;
|
|
494
|
-
y[id * params.y_inc] = h21 * xi + h22 * yi;
|
|
500
|
+
if (i == 0u) {
|
|
501
|
+
result[0] = tileScale[0] * sqrt(tileSsq[0]);
|
|
495
502
|
}
|
|
496
503
|
}
|
|
497
|
-
`});var
|
|
504
|
+
`});var yt,bt=O(()=>{yt=`// isamax: returns index of element with largest absolute value
|
|
498
505
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/argmax.wgsl.
|
|
499
506
|
|
|
500
507
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -576,1076 +583,1700 @@ fn main(
|
|
|
576
583
|
partials_idx[wgid.x] = tile_idx[0];
|
|
577
584
|
}
|
|
578
585
|
}
|
|
579
|
-
`});var
|
|
580
|
-
//
|
|
581
|
-
//
|
|
582
|
-
// still covers all rows when m exceeds maxComputeWorkgroupsPerDimension.
|
|
583
|
-
// Threads stride through A[row, :] and x with coalesced reads (consecutive
|
|
584
|
-
// threads \u2192 consecutive addresses). Four independent accumulators let the GPU
|
|
585
|
-
// pipeline memory requests across iterations (ILP=4), hiding the
|
|
586
|
-
// global-memory latency.
|
|
587
|
-
|
|
588
|
-
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
589
|
-
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
590
|
-
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
586
|
+
`});var vt,xt=O(()=>{vt=`// amax reduction: collapses 2*WGS (value, index) pairs into one index.
|
|
587
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
588
|
+
// partials_val and partials_idx must have exactly 2*WGS entries.
|
|
591
589
|
|
|
592
|
-
|
|
593
|
-
|
|
594
|
-
|
|
595
|
-
alpha: f32,
|
|
596
|
-
beta: f32,
|
|
597
|
-
incx: u32,
|
|
598
|
-
incy: u32,
|
|
599
|
-
lda: u32,
|
|
600
|
-
}
|
|
590
|
+
@group(0) @binding(0) var<storage, read> partials_val: array<f32>;
|
|
591
|
+
@group(0) @binding(1) var<storage, read> partials_idx: array<u32>;
|
|
592
|
+
@group(0) @binding(2) var<storage, read_write> result: array<u32>;
|
|
601
593
|
|
|
602
|
-
|
|
594
|
+
const WGS: u32 = 64;
|
|
603
595
|
|
|
604
|
-
|
|
605
|
-
var<workgroup>
|
|
596
|
+
var<workgroup> tile_val: array<f32, 64>;
|
|
597
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
606
598
|
|
|
607
599
|
@compute @workgroup_size(64)
|
|
608
|
-
fn
|
|
609
|
-
@builtin(
|
|
610
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
611
|
-
@builtin(num_workgroups) nwg: vec3u,
|
|
600
|
+
fn reduce(
|
|
601
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
612
602
|
) {
|
|
613
|
-
|
|
614
|
-
|
|
615
|
-
|
|
616
|
-
|
|
617
|
-
|
|
618
|
-
|
|
619
|
-
|
|
620
|
-
|
|
621
|
-
|
|
622
|
-
|
|
623
|
-
|
|
624
|
-
let n4_floor = (params.n / (4u * WGS)) * (4u * WGS);
|
|
625
|
-
for (var j: u32 = lid.x; j < n4_floor; j += 4u * WGS) {
|
|
626
|
-
acc0 += A[row_base + j ] * x[ j * params.incx];
|
|
627
|
-
acc1 += A[row_base + j + WGS ] * x[(j + WGS) * params.incx];
|
|
628
|
-
acc2 += A[row_base + j + 2u * WGS ] * x[(j + 2u * WGS) * params.incx];
|
|
629
|
-
acc3 += A[row_base + j + 3u * WGS ] * x[(j + 3u * WGS) * params.incx];
|
|
630
|
-
}
|
|
631
|
-
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
632
|
-
for (var j: u32 = n4_floor + lid.x; j < params.n; j += WGS) {
|
|
633
|
-
acc0 += A[row_base + j] * x[j * params.incx];
|
|
634
|
-
}
|
|
603
|
+
let i = lid.x;
|
|
604
|
+
let a_val = partials_val[i];
|
|
605
|
+
let b_val = partials_val[i + WGS];
|
|
606
|
+
if (b_val > a_val || (b_val == a_val && partials_idx[i + WGS] < partials_idx[i])) {
|
|
607
|
+
tile_val[i] = b_val;
|
|
608
|
+
tile_idx[i] = partials_idx[i + WGS];
|
|
609
|
+
} else {
|
|
610
|
+
tile_val[i] = a_val;
|
|
611
|
+
tile_idx[i] = partials_idx[i];
|
|
612
|
+
}
|
|
613
|
+
workgroupBarrier();
|
|
635
614
|
|
|
636
|
-
|
|
637
|
-
|
|
638
|
-
|
|
639
|
-
|
|
640
|
-
if
|
|
641
|
-
|
|
615
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
616
|
+
if (i < s) {
|
|
617
|
+
let c_val = tile_val[i];
|
|
618
|
+
let d_val = tile_val[i + s];
|
|
619
|
+
if (d_val > c_val || (d_val == c_val && tile_idx[i + s] < tile_idx[i])) {
|
|
620
|
+
tile_val[i] = d_val;
|
|
621
|
+
tile_idx[i] = tile_idx[i + s];
|
|
642
622
|
}
|
|
643
|
-
workgroupBarrier();
|
|
644
|
-
}
|
|
645
|
-
|
|
646
|
-
if lid.x == 0u {
|
|
647
|
-
let yi = row * params.incy;
|
|
648
|
-
y[yi] = params.alpha * scratch[0] + params.beta * y[yi];
|
|
649
623
|
}
|
|
650
|
-
// All 64 threads must agree before the next row reuses scratch[].
|
|
651
624
|
workgroupBarrier();
|
|
652
625
|
}
|
|
626
|
+
|
|
627
|
+
if (i == 0u) { result[0] = tile_idx[0]; }
|
|
653
628
|
}
|
|
654
|
-
`});var
|
|
655
|
-
//
|
|
656
|
-
//
|
|
657
|
-
//
|
|
629
|
+
`});var Kr,_t=O(()=>{Kr=`// Double-double arithmetic via Dekker's algorithm \u2014 an alternative to
|
|
630
|
+
// f64add.wgsl's bit-exact IEEE-754 emulation. Doesn't touch that path.
|
|
631
|
+
//
|
|
632
|
+
// A double-double number is a pair (hi, lo) of f32 with hi+lo approximating
|
|
633
|
+
// a higher-precision value, hi holding the leading bits and lo the rounding
|
|
634
|
+
// error hi lost. ~48 bits of mantissa vs f32's 24, less than real f64's 52.
|
|
635
|
+
//
|
|
636
|
+
// No bindings, no entry point \u2014 a helper library, concatenated with a
|
|
637
|
+
// consumer's own bindings/entry point by getPipeline (WGSL has no #include).
|
|
638
|
+
// The DD struct lives here \u2014 abs.wgsl/add.wgsl/greater.wgsl/equal.wgsl all
|
|
639
|
+
// use it but don't redefine it (WGSL errors on duplicate struct definitions
|
|
640
|
+
// once concatenated), so any consumer using those must concatenate this
|
|
641
|
+
// file too, first.
|
|
658
642
|
|
|
659
|
-
|
|
660
|
-
|
|
661
|
-
|
|
643
|
+
struct DD {
|
|
644
|
+
hi: f32,
|
|
645
|
+
lo: f32,
|
|
646
|
+
}
|
|
647
|
+
`});var ye,Bt=O(()=>{ye=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
662
648
|
|
|
663
|
-
|
|
664
|
-
|
|
665
|
-
|
|
666
|
-
|
|
667
|
-
|
|
668
|
-
|
|
669
|
-
|
|
670
|
-
lda: u32,
|
|
649
|
+
// |a| for a double-double pair. Negation is exact (no rounding), so this is
|
|
650
|
+
// just a sign flip on both components \u2014 hi alone determines the pair's sign.
|
|
651
|
+
fn ddAbs(a: DD) -> DD {
|
|
652
|
+
if (a.hi < 0.0) {
|
|
653
|
+
return DD(-a.hi, -a.lo);
|
|
654
|
+
}
|
|
655
|
+
return a;
|
|
671
656
|
}
|
|
657
|
+
`});var zr,At=O(()=>{zr=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
672
658
|
|
|
673
|
-
|
|
659
|
+
// \u2500\u2500 A real compiler bug \u2014 read before touching anything below \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500
|
|
660
|
+
//
|
|
661
|
+
// twoSum/fastTwoSum's error term \`e\` should be nonzero (that's the point \u2014
|
|
662
|
+
// \`s\` is rounded). A front-end-optimizer bug zeros it anyway, on both NVIDIA
|
|
663
|
+
// and Mesa (ANV + llvmpipe), via two different mechanisms needing two fixes:
|
|
664
|
+
// bitcast-based subtraction (\`fsub\`/\`negf\`, fixes NVIDIA) and materializing
|
|
665
|
+
// the sum through workgroup memory + workgroupBarrier() (fixes Mesa). Only
|
|
666
|
+
// both together (ddAddProtected) is verified correct everywhere \u2014 the plain
|
|
667
|
+
// twoSum/fastTwoSum/ddAdd below are reference-only, not safe to use.
|
|
668
|
+
fn negf(x: f32) -> f32 {
|
|
669
|
+
return bitcast<f32>(bitcast<u32>(x) ^ 0x80000000u);
|
|
670
|
+
}
|
|
671
|
+
fn fsub(a: f32, b: f32) -> f32 {
|
|
672
|
+
return a + negf(b);
|
|
673
|
+
}
|
|
674
674
|
|
|
675
|
-
|
|
676
|
-
|
|
675
|
+
// Knuth/M\xF8ller's TwoSum: s = fl(a+b), e = exact rounding error, a+b == s+e.
|
|
676
|
+
// Works for any a, b. UNPROTECTED \u2014 see header above.
|
|
677
|
+
fn twoSum(a: f32, b: f32) -> DD {
|
|
678
|
+
let s = a + b;
|
|
679
|
+
let v = s - a;
|
|
680
|
+
let e = (a - (s - v)) + (b - v);
|
|
681
|
+
return DD(s, e);
|
|
682
|
+
}
|
|
677
683
|
|
|
678
|
-
|
|
679
|
-
|
|
680
|
-
|
|
681
|
-
|
|
682
|
-
)
|
|
683
|
-
|
|
684
|
-
|
|
685
|
-
// tile over x (length m, the rows of A)
|
|
686
|
-
let m_floor = (params.m / WGS) * WGS;
|
|
687
|
-
var acc0: f32 = 0.0;
|
|
688
|
-
var acc1: f32 = 0.0;
|
|
689
|
-
var acc2: f32 = 0.0;
|
|
690
|
-
var acc3: f32 = 0.0;
|
|
684
|
+
// Dekker's Fast-Two-Sum: same contract, but only correct when |a| >= |b|.
|
|
685
|
+
// UNPROTECTED \u2014 see header above.
|
|
686
|
+
fn fastTwoSum(a: f32, b: f32) -> DD {
|
|
687
|
+
let s = a + b;
|
|
688
|
+
let e = b - (s - a);
|
|
689
|
+
return DD(s, e);
|
|
690
|
+
}
|
|
691
691
|
|
|
692
|
-
|
|
693
|
-
|
|
694
|
-
|
|
695
|
-
|
|
692
|
+
// Double-double addition (Dekker's Add2). UNPROTECTED \u2014 see header above.
|
|
693
|
+
fn ddAdd(a: DD, b: DD) -> DD {
|
|
694
|
+
let s = twoSum(a.hi, b.hi);
|
|
695
|
+
let loSum = a.lo + b.lo;
|
|
696
|
+
return fastTwoSum(s.hi, s.lo + loSum);
|
|
697
|
+
}
|
|
696
698
|
|
|
697
|
-
|
|
698
|
-
|
|
699
|
-
|
|
700
|
-
|
|
701
|
-
|
|
702
|
-
|
|
703
|
-
|
|
704
|
-
|
|
705
|
-
|
|
706
|
-
|
|
707
|
-
|
|
708
|
-
}
|
|
699
|
+
// \u2500\u2500 Protected variants \u2014 use these \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500
|
|
700
|
+
//
|
|
701
|
+
// Bitcast subtraction + workgroup-barrier materialization, verified correct
|
|
702
|
+
// on all three backends tested. Costs a real barrier: fine for O(1)-per-
|
|
703
|
+
// thread or O(log n) reduction use, not a long per-element loop. A
|
|
704
|
+
// workgroupBarrier() requires uniform control flow, so:
|
|
705
|
+
// - \`threadSlot\` must be unique per concurrent caller (e.g. local_invocation_index).
|
|
706
|
+
// - Every thread in the workgroup must call this the same number of times
|
|
707
|
+
// \u2014 including ones whose result gets discarded. Compute unconditionally;
|
|
708
|
+
// only the write-back should be conditional.
|
|
709
|
+
var<workgroup> dekkerScratch: array<f32, 64>;
|
|
709
710
|
|
|
710
|
-
|
|
711
|
-
|
|
712
|
-
|
|
713
|
-
|
|
714
|
-
|
|
715
|
-
|
|
716
|
-
|
|
717
|
-
}
|
|
711
|
+
fn twoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
|
|
712
|
+
dekkerScratch[threadSlot] = a + b;
|
|
713
|
+
workgroupBarrier();
|
|
714
|
+
let s = dekkerScratch[threadSlot];
|
|
715
|
+
let v = fsub(s, a);
|
|
716
|
+
let e = fsub(a, fsub(s, v)) + fsub(b, v);
|
|
717
|
+
return DD(s, e);
|
|
718
718
|
}
|
|
719
|
-
`});var qe,ze=O(()=>{qe=`// ssymv: y = alpha * A * x + beta * y
|
|
720
|
-
// A is n\xD7n symmetric, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
721
|
-
// The logical matrix is fully dense (symmetric), so each row's dot product
|
|
722
|
-
// sums over all n columns; entries on the unstored side of the diagonal are
|
|
723
|
-
// fetched from their mirror position (A[i,j] == A[j,i]).
|
|
724
|
-
// One workgroup per row, grid-stride outer loop.
|
|
725
719
|
|
|
726
|
-
|
|
727
|
-
|
|
728
|
-
|
|
720
|
+
fn fastTwoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
|
|
721
|
+
dekkerScratch[threadSlot] = a + b;
|
|
722
|
+
workgroupBarrier();
|
|
723
|
+
let s = dekkerScratch[threadSlot];
|
|
724
|
+
let e = fsub(b, fsub(s, a));
|
|
725
|
+
return DD(s, e);
|
|
726
|
+
}
|
|
727
|
+
|
|
728
|
+
// Protected double-double addition \u2014 same contract as ddAdd, but exact.
|
|
729
|
+
fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
730
|
+
let s = twoSumProtected(a.hi, b.hi, threadSlot);
|
|
731
|
+
let loSum = a.lo + b.lo;
|
|
732
|
+
return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
|
|
733
|
+
}
|
|
734
|
+
|
|
735
|
+
// Double-double subtraction \u2014 a - b, via exact negation (a sign-bit flip,
|
|
736
|
+
// no rounding) then ddAddProtected. Same protection contract.
|
|
737
|
+
fn ddSubProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
738
|
+
return ddAddProtected(a, DD(negf(b.hi), negf(b.lo)), threadSlot);
|
|
739
|
+
}
|
|
740
|
+
`});var Gt,St=O(()=>{Gt=`// dasum: sum(|x[i]|), double-double (Dekker). Same ILP=4 shape as sasum.wgsl;
|
|
741
|
+
// see f64/utils/add.wgsl for ddAddProtected and why plain ddAdd isn't safe.
|
|
742
|
+
// GpuVector input isn't pre-abs'd, so ddAbs() (f64/utils/abs.wgsl) applies
|
|
743
|
+
// unconditionally below.
|
|
744
|
+
|
|
745
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
746
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
747
|
+
@group(0) @binding(2) var<storage, read_write> partialsHi: array<f32>;
|
|
748
|
+
@group(0) @binding(3) var<storage, read_write> partialsLo: array<f32>;
|
|
749
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
729
750
|
|
|
730
751
|
struct Params {
|
|
731
752
|
n: u32,
|
|
732
|
-
|
|
733
|
-
beta: f32,
|
|
734
|
-
incx: u32,
|
|
735
|
-
incy: u32,
|
|
736
|
-
lda: u32,
|
|
737
|
-
uplo: u32, // 0 = lower, 1 = upper
|
|
753
|
+
x_inc: u32,
|
|
738
754
|
}
|
|
739
755
|
|
|
740
|
-
|
|
756
|
+
const WGS: u32 = 64;
|
|
741
757
|
|
|
742
|
-
|
|
743
|
-
var<workgroup> scratch: array<f32, 64>;
|
|
758
|
+
var<workgroup> tile: array<DD, 64>;
|
|
744
759
|
|
|
745
760
|
@compute @workgroup_size(64)
|
|
746
|
-
fn
|
|
747
|
-
@builtin(
|
|
748
|
-
@builtin(local_invocation_id)
|
|
749
|
-
@builtin(
|
|
761
|
+
fn dasum_main(
|
|
762
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
763
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
764
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
765
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
750
766
|
) {
|
|
751
|
-
|
|
752
|
-
|
|
767
|
+
var acc0 = DD(0.0, 0.0);
|
|
768
|
+
var acc1 = DD(0.0, 0.0);
|
|
769
|
+
var acc2 = DD(0.0, 0.0);
|
|
770
|
+
var acc3 = DD(0.0, 0.0);
|
|
753
771
|
|
|
754
|
-
|
|
755
|
-
|
|
756
|
-
var aVal: f32;
|
|
757
|
-
if params.uplo == 0u {
|
|
758
|
-
// Lower: A[i,j] stored at A[i*lda+j] for j \u2264 i, mirrored from A[j*lda+i] otherwise
|
|
759
|
-
if j <= i {
|
|
760
|
-
aVal = A[i * params.lda + j];
|
|
761
|
-
} else {
|
|
762
|
-
aVal = A[j * params.lda + i];
|
|
763
|
-
}
|
|
764
|
-
} else {
|
|
765
|
-
// Upper: A[i,j] stored at A[i*lda+j] for j \u2265 i, mirrored from A[j*lda+i] otherwise
|
|
766
|
-
if j >= i {
|
|
767
|
-
aVal = A[i * params.lda + j];
|
|
768
|
-
} else {
|
|
769
|
-
aVal = A[j * params.lda + i];
|
|
770
|
-
}
|
|
771
|
-
}
|
|
772
|
-
acc += aVal * x[j * params.incx];
|
|
773
|
-
}
|
|
772
|
+
let stride = num_wg.x * WGS;
|
|
773
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
774
774
|
|
|
775
|
-
|
|
776
|
-
|
|
775
|
+
// Same trip count for every thread, but driven by a counter, not \`id\`
|
|
776
|
+
// itself (ddAddProtected's barrier needs a provably-uniform loop bound).
|
|
777
|
+
let mainIters = n4_floor / (4u * stride);
|
|
778
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
779
|
+
let id = gid.x + iter * 4u * stride;
|
|
780
|
+
let i0 = id * params.x_inc;
|
|
781
|
+
let i1 = (id + stride) * params.x_inc;
|
|
782
|
+
let i2 = (id + 2u * stride) * params.x_inc;
|
|
783
|
+
let i3 = (id + 3u * stride) * params.x_inc;
|
|
784
|
+
acc0 = ddAddProtected(acc0, ddAbs(DD(xHi[i0], xLo[i0])), lid.x);
|
|
785
|
+
acc1 = ddAddProtected(acc1, ddAbs(DD(xHi[i1], xLo[i1])), lid.x);
|
|
786
|
+
acc2 = ddAddProtected(acc2, ddAbs(DD(xHi[i2], xLo[i2])), lid.x);
|
|
787
|
+
acc3 = ddAddProtected(acc3, ddAbs(DD(xHi[i3], xLo[i3])), lid.x);
|
|
788
|
+
}
|
|
789
|
+
|
|
790
|
+
// Tail is ragged (0-3 extra per thread) \u2014 pad to this workgroup's worst case.
|
|
791
|
+
let wgBaseGid = wgid.x * WGS;
|
|
792
|
+
var tailIters = 0u;
|
|
793
|
+
if (n4_floor + wgBaseGid < params.n) {
|
|
794
|
+
tailIters = (params.n - 1u - n4_floor - wgBaseGid) / stride + 1u;
|
|
795
|
+
}
|
|
796
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
797
|
+
let id = n4_floor + gid.x + iter * stride;
|
|
798
|
+
let valid = id < params.n;
|
|
799
|
+
let i = select(0u, id * params.x_inc, valid);
|
|
800
|
+
let loaded = ddAbs(DD(xHi[i], xLo[i])); // select() has no DD overload
|
|
801
|
+
let contribution = DD(select(0.0, loaded.hi, valid), select(0.0, loaded.lo, valid));
|
|
802
|
+
acc0 = ddAddProtected(acc0, contribution, lid.x);
|
|
803
|
+
}
|
|
804
|
+
|
|
805
|
+
let combined01 = ddAddProtected(acc0, acc1, lid.x);
|
|
806
|
+
let combined23 = ddAddProtected(acc2, acc3, lid.x);
|
|
807
|
+
tile[lid.x] = ddAddProtected(combined01, combined23, lid.x);
|
|
808
|
+
workgroupBarrier();
|
|
809
|
+
|
|
810
|
+
// Inactive threads combine against a throwaway partner and discard it
|
|
811
|
+
// (ddAddProtected must be called unconditionally by every thread).
|
|
812
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
813
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
814
|
+
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
815
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
816
|
+
if (lid.x < s) { tile[lid.x] = combined; }
|
|
777
817
|
workgroupBarrier();
|
|
778
|
-
|
|
779
|
-
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
780
|
-
workgroupBarrier();
|
|
781
|
-
}
|
|
818
|
+
}
|
|
782
819
|
|
|
783
|
-
|
|
784
|
-
|
|
785
|
-
|
|
820
|
+
if (lid.x == 0u) {
|
|
821
|
+
partialsHi[wgid.x] = tile[0].hi;
|
|
822
|
+
partialsLo[wgid.x] = tile[0].lo;
|
|
786
823
|
}
|
|
787
824
|
}
|
|
788
|
-
`});var
|
|
789
|
-
//
|
|
790
|
-
//
|
|
791
|
-
//
|
|
792
|
-
//
|
|
825
|
+
`});var Ne,Et=O(()=>{Ne=`// sum reduction (f64, double-double): collapses 2*WGS partial (hi, lo) pairs
|
|
826
|
+
// into one, using ddAddProtected instead of plain f32 \`+\` (see
|
|
827
|
+
// reduction/sum.wgsl for the f32 original this mirrors).
|
|
828
|
+
// dispatch: 1 workgroup of WGS threads. partialsHi/partialsLo must have
|
|
829
|
+
// exactly 2*WGS entries each. Concatenated after f64/dekker.wgsl (DD struct)
|
|
830
|
+
// and f64/utils/add.wgsl (ddAddProtected \u2014 see it for why plain ddAdd isn't safe).
|
|
793
831
|
|
|
794
|
-
@group(0) @binding(0) var<storage, read>
|
|
795
|
-
@group(0) @binding(1) var<storage, read>
|
|
796
|
-
@group(0) @binding(2) var<storage, read_write>
|
|
832
|
+
@group(0) @binding(0) var<storage, read> partialsHi: array<f32>;
|
|
833
|
+
@group(0) @binding(1) var<storage, read> partialsLo: array<f32>;
|
|
834
|
+
@group(0) @binding(2) var<storage, read_write> resultHi: array<f32, 1>;
|
|
835
|
+
@group(0) @binding(3) var<storage, read_write> resultLo: array<f32, 1>;
|
|
836
|
+
|
|
837
|
+
const WGS: u32 = 64;
|
|
838
|
+
|
|
839
|
+
var<workgroup> tile: array<DD, 64>;
|
|
840
|
+
|
|
841
|
+
@compute @workgroup_size(64)
|
|
842
|
+
fn reduce_f64(
|
|
843
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
844
|
+
) {
|
|
845
|
+
let i = lid.x;
|
|
846
|
+
let a = DD(partialsHi[i], partialsLo[i]);
|
|
847
|
+
let b = DD(partialsHi[i + WGS], partialsLo[i + WGS]);
|
|
848
|
+
tile[i] = ddAddProtected(a, b, i);
|
|
849
|
+
workgroupBarrier();
|
|
850
|
+
|
|
851
|
+
// ddAddProtected must be called unconditionally by every thread.
|
|
852
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
853
|
+
let partner = select(i, i + s, i < s);
|
|
854
|
+
let combined = ddAddProtected(tile[i], tile[partner], i);
|
|
855
|
+
workgroupBarrier();
|
|
856
|
+
if (i < s) { tile[i] = combined; }
|
|
857
|
+
workgroupBarrier();
|
|
858
|
+
}
|
|
859
|
+
|
|
860
|
+
if (i == 0u) {
|
|
861
|
+
resultHi[0] = tile[0].hi;
|
|
862
|
+
resultLo[0] = tile[0].lo;
|
|
863
|
+
}
|
|
864
|
+
}
|
|
865
|
+
`});var $r,kt=O(()=>{$r=`// Requires f64/dekker.wgsl concatenated first for the DD struct, and
|
|
866
|
+
// f64/utils/add.wgsl for fsub/negf (bitcast-based subtraction/negation) and,
|
|
867
|
+
// for ddMulProtected at the bottom, fastTwoSumProtected.
|
|
868
|
+
//
|
|
869
|
+
// Use twoProdBit \u2014 verified universal (0 corrupting failures across 3000+
|
|
870
|
+
// random trials on NVIDIA/Intel-Mesa-ANV/llvmpipe), no barrier protection
|
|
871
|
+
// needed. The classic approaches below (twoProd, twoProdFma) each fail on
|
|
872
|
+
// one backend in a way barrier materialization doesn't fix; twoProdBit
|
|
873
|
+
// sidesteps the bug instead by deriving the split via bitcast+bitmask
|
|
874
|
+
// rather than an arithmetic identity, leaving nothing for a reassociating
|
|
875
|
+
// compiler to fold. Intel Mesa ANV shows frequent last-bit-only diffs from
|
|
876
|
+
// strict ground truth (never data-corrupting) \u2014 consistent with the driver
|
|
877
|
+
// legitimately auto-fusing \`x - y*z\` into hardware FMA.
|
|
878
|
+
const SPLIT_CONST: f32 = 4097.0;
|
|
879
|
+
|
|
880
|
+
fn bitSplit(a: f32) -> DD {
|
|
881
|
+
let bits = bitcast<u32>(a);
|
|
882
|
+
// Top 11 mantissa bits, so hi carries 12 significant bits with the implicit
|
|
883
|
+
// leading 1 \u2014 the halves are multiplied pairwise and f32 holds 24, so a
|
|
884
|
+
// wider split rounds those products and the "exact" error term goes wrong.
|
|
885
|
+
// Matches SPLIT_CONST = 2^12+1 used by the Veltkamp path below.
|
|
886
|
+
let hiBits = bits & 0xFFFFF000u;
|
|
887
|
+
let hi = bitcast<f32>(hiBits);
|
|
888
|
+
let lo = fsub(a, hi); // exact by Sterbenz's lemma (hi, a share an exponent, are close)
|
|
889
|
+
return DD(hi, lo);
|
|
890
|
+
}
|
|
891
|
+
|
|
892
|
+
fn twoProdBit(a: f32, b: f32) -> DD {
|
|
893
|
+
let s = a * b;
|
|
894
|
+
let aSplit = bitSplit(a);
|
|
895
|
+
let bSplit = bitSplit(b);
|
|
896
|
+
let e = fsub(fsub(fsub(fsub(s, aSplit.hi * bSplit.hi), aSplit.lo * bSplit.hi), aSplit.hi * bSplit.lo), aSplit.lo * bSplit.lo);
|
|
897
|
+
return DD(s, negf(e));
|
|
898
|
+
}
|
|
899
|
+
|
|
900
|
+
// \u2500\u2500 Unsafe historical reference \u2014 do not use \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500
|
|
901
|
+
// Both broken on one backend, confirmed via isolated cross-driver testing,
|
|
902
|
+
// NOT fixed by barrier materialization (unlike addition's bug):
|
|
903
|
+
// - veltkampSplit/twoProd (Dekker's original): fails on NVIDIA \u2014 compiler
|
|
904
|
+
// folds \`hi = c - (c - a)\` to \`= a\` straight through fsub/negf, even
|
|
905
|
+
// with every intermediate barrier-materialized (11/11 fail, worse than
|
|
906
|
+
// unprotected's 6/11).
|
|
907
|
+
// - twoProdFma (Ogita/Rump/Oishi): fails on llvmpipe \u2014 its software fma()
|
|
908
|
+
// likely isn't genuinely fused, making \`fma(a,b,-(a*b))\` correctly (not
|
|
909
|
+
// buggily) zero. Materializing \`s\` doesn't change this.
|
|
910
|
+
fn veltkampSplit(a: f32) -> DD {
|
|
911
|
+
let c = SPLIT_CONST * a;
|
|
912
|
+
let big = fsub(c, a);
|
|
913
|
+
let hi = fsub(c, big);
|
|
914
|
+
let lo = fsub(a, hi);
|
|
915
|
+
return DD(hi, lo);
|
|
916
|
+
}
|
|
917
|
+
|
|
918
|
+
fn twoProd(a: f32, b: f32) -> DD {
|
|
919
|
+
let s = a * b;
|
|
920
|
+
let aSplit = veltkampSplit(a);
|
|
921
|
+
let bSplit = veltkampSplit(b);
|
|
922
|
+
let e = fsub(fsub(fsub(fsub(s, aSplit.hi * bSplit.hi), aSplit.lo * bSplit.hi), aSplit.hi * bSplit.lo), aSplit.lo * bSplit.lo);
|
|
923
|
+
return DD(s, negf(e));
|
|
924
|
+
}
|
|
925
|
+
|
|
926
|
+
fn twoProdFma(a: f32, b: f32) -> DD {
|
|
927
|
+
let s = a * b;
|
|
928
|
+
let e = fma(a, b, negf(s));
|
|
929
|
+
return DD(s, e);
|
|
930
|
+
}
|
|
931
|
+
|
|
932
|
+
// DD \xD7 DD product (Dekker/Bailey): twoProdBit(a.hi, b.hi) already captures
|
|
933
|
+
// the dominant term to full DD precision, and the cross terms are below the
|
|
934
|
+
// ~48-bit floor anyway, so folding them in with plain f32 loses nothing.
|
|
935
|
+
//
|
|
936
|
+
// Another real compiler bug, distinct from add.wgsl's twoSum one \u2014 confirmed
|
|
937
|
+
// on Intel Mesa ANV: when p.lo feeds straight into \`crossAndLo\` unobserved,
|
|
938
|
+
// the compiler folds it away entirely. Materializing p.lo itself through
|
|
939
|
+
// workgroup memory + workgroupBarrier() (like twoSumProtected does for its
|
|
940
|
+
// sum) is what fixes it, so ddMulRaw now takes threadSlot and always pays
|
|
941
|
+
// that barrier \u2014 no longer a plain unprotected batchable helper.
|
|
942
|
+
fn ddMulRaw(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
943
|
+
let p = twoProdBit(a.hi, b.hi);
|
|
944
|
+
dekkerScratch[threadSlot] = p.lo;
|
|
945
|
+
workgroupBarrier();
|
|
946
|
+
let pLo = dekkerScratch[threadSlot];
|
|
947
|
+
let crossAndLo = pLo + (a.hi * b.lo + a.lo * b.hi);
|
|
948
|
+
return DD(p.hi, crossAndLo);
|
|
949
|
+
}
|
|
950
|
+
|
|
951
|
+
fn ddMulProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
952
|
+
let raw = ddMulRaw(a, b, threadSlot);
|
|
953
|
+
return fastTwoSumProtected(raw.hi, raw.lo, threadSlot);
|
|
954
|
+
}
|
|
955
|
+
`});var Nt,Dt=O(()=>{Nt=`// ddot: sum(x[i] * y[i]), double-double (Dekker). Same ILP=4 shape as
|
|
956
|
+
// dasum.wgsl, which this mirrors closely \u2014 the only structural difference is
|
|
957
|
+
// a second input vector and a product where dasum takes an absolute value.
|
|
958
|
+
//
|
|
959
|
+
// See f64/utils/add.wgsl for ddAddProtected and why plain ddAdd isn't safe,
|
|
960
|
+
// and f64/utils/multiply.wgsl for ddMulProtected. The multiply itself
|
|
961
|
+
// (twoProdBit) needs no barrier; only its final renormalisation does, which
|
|
962
|
+
// is why each element costs two protected ops here against dasum's one.
|
|
963
|
+
|
|
964
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
965
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
966
|
+
@group(0) @binding(2) var<storage, read> yHi: array<f32>;
|
|
967
|
+
@group(0) @binding(3) var<storage, read> yLo: array<f32>;
|
|
968
|
+
@group(0) @binding(4) var<storage, read_write> partialsHi: array<f32>;
|
|
969
|
+
@group(0) @binding(5) var<storage, read_write> partialsLo: array<f32>;
|
|
970
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
797
971
|
|
|
798
972
|
struct Params {
|
|
799
973
|
n: u32,
|
|
800
|
-
|
|
801
|
-
|
|
802
|
-
lda: u32,
|
|
803
|
-
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
804
|
-
uplo: u32, // 0 = lower, 1 = upper
|
|
805
|
-
diag: u32, // 0 = non-unit, 1 = unit
|
|
974
|
+
x_inc: u32,
|
|
975
|
+
y_inc: u32,
|
|
806
976
|
}
|
|
807
977
|
|
|
808
|
-
|
|
978
|
+
const WGS: u32 = 64;
|
|
809
979
|
|
|
810
|
-
|
|
811
|
-
var<workgroup> scratch: array<f32, 64>;
|
|
980
|
+
var<workgroup> tile: array<DD, 64>;
|
|
812
981
|
|
|
813
982
|
@compute @workgroup_size(64)
|
|
814
|
-
fn
|
|
815
|
-
@builtin(
|
|
816
|
-
@builtin(local_invocation_id)
|
|
817
|
-
@builtin(
|
|
983
|
+
fn ddot_main(
|
|
984
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
985
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
986
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
987
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
818
988
|
) {
|
|
819
|
-
|
|
820
|
-
|
|
989
|
+
var acc0 = DD(0.0, 0.0);
|
|
990
|
+
var acc1 = DD(0.0, 0.0);
|
|
991
|
+
var acc2 = DD(0.0, 0.0);
|
|
992
|
+
var acc3 = DD(0.0, 0.0);
|
|
821
993
|
|
|
822
|
-
|
|
823
|
-
|
|
824
|
-
if params.uplo == 0u {
|
|
825
|
-
// Lower: A[i,j] stored at A[i*lda+j] for j \u2264 i
|
|
826
|
-
for (var j = lid.x; j <= i; j += WGS) {
|
|
827
|
-
var aVal: f32;
|
|
828
|
-
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
829
|
-
if params.diag == 1u && j == i {
|
|
830
|
-
aVal = 1.0;
|
|
831
|
-
} else if ( j <= i ) {
|
|
832
|
-
aVal = A[i * params.lda + j];
|
|
833
|
-
}
|
|
834
|
-
acc += aVal * x[j * params.incx];
|
|
835
|
-
}
|
|
836
|
-
} else {
|
|
837
|
-
// Upper: A[i,j] stored at A[i*lda+j] for j \u2265 i
|
|
838
|
-
for (var j = i + lid.x; j < params.n; j += WGS) {
|
|
839
|
-
var aVal: f32;
|
|
840
|
-
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
841
|
-
if params.diag == 1u && j == i {
|
|
842
|
-
aVal = 1.0;
|
|
843
|
-
} else if ( j >= i ) {
|
|
844
|
-
aVal = A[i * params.lda + j];
|
|
845
|
-
}
|
|
846
|
-
acc += aVal * x[j * params.incx];
|
|
847
|
-
}
|
|
848
|
-
}
|
|
849
|
-
} else {
|
|
850
|
-
// Transpose: y[i] = \u03A3_j A[j,i] * x[j]
|
|
851
|
-
if params.uplo == 0u {
|
|
852
|
-
// Lower: A[j,i] stored at A[j*lda+i] for j \u2265 i
|
|
853
|
-
for (var j = i + lid.x; j < params.n; j += WGS) {
|
|
854
|
-
var aVal: f32;
|
|
855
|
-
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
856
|
-
if params.diag == 1u && j == i {
|
|
857
|
-
aVal = 1.0;
|
|
858
|
-
} else if ( j >= i ) {
|
|
859
|
-
aVal = A[j * params.lda + i];
|
|
860
|
-
}
|
|
861
|
-
acc += aVal * x[j * params.incx];
|
|
862
|
-
}
|
|
863
|
-
} else {
|
|
864
|
-
// Upper: A[j,i] stored at A[j*lda+i] for j \u2264 i
|
|
865
|
-
for (var j = lid.x; j <= i; j += WGS) {
|
|
866
|
-
var aVal: f32;
|
|
867
|
-
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
868
|
-
if params.diag == 1u && j == i {
|
|
869
|
-
aVal = 1.0;
|
|
870
|
-
} else if ( j <= i ) {
|
|
871
|
-
aVal = A[j * params.lda + i];
|
|
872
|
-
}
|
|
873
|
-
acc += aVal * x[j * params.incx];
|
|
874
|
-
}
|
|
875
|
-
}
|
|
876
|
-
}
|
|
994
|
+
let stride = num_wg.x * WGS;
|
|
995
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
877
996
|
|
|
878
|
-
|
|
879
|
-
|
|
997
|
+
// Same trip count for every thread, but driven by a counter, not \`id\`
|
|
998
|
+
// itself (the protected ops' barriers need a provably-uniform loop bound).
|
|
999
|
+
let mainIters = n4_floor / (4u * stride);
|
|
1000
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
1001
|
+
let id = gid.x + iter * 4u * stride;
|
|
1002
|
+
let d0 = id;
|
|
1003
|
+
let d1 = id + stride;
|
|
1004
|
+
let d2 = id + 2u * stride;
|
|
1005
|
+
let d3 = id + 3u * stride;
|
|
1006
|
+
|
|
1007
|
+
let p0 = ddMulProtected(DD(xHi[d0 * params.x_inc], xLo[d0 * params.x_inc]),
|
|
1008
|
+
DD(yHi[d0 * params.y_inc], yLo[d0 * params.y_inc]), lid.x);
|
|
1009
|
+
let p1 = ddMulProtected(DD(xHi[d1 * params.x_inc], xLo[d1 * params.x_inc]),
|
|
1010
|
+
DD(yHi[d1 * params.y_inc], yLo[d1 * params.y_inc]), lid.x);
|
|
1011
|
+
let p2 = ddMulProtected(DD(xHi[d2 * params.x_inc], xLo[d2 * params.x_inc]),
|
|
1012
|
+
DD(yHi[d2 * params.y_inc], yLo[d2 * params.y_inc]), lid.x);
|
|
1013
|
+
let p3 = ddMulProtected(DD(xHi[d3 * params.x_inc], xLo[d3 * params.x_inc]),
|
|
1014
|
+
DD(yHi[d3 * params.y_inc], yLo[d3 * params.y_inc]), lid.x);
|
|
1015
|
+
|
|
1016
|
+
acc0 = ddAddProtected(acc0, p0, lid.x);
|
|
1017
|
+
acc1 = ddAddProtected(acc1, p1, lid.x);
|
|
1018
|
+
acc2 = ddAddProtected(acc2, p2, lid.x);
|
|
1019
|
+
acc3 = ddAddProtected(acc3, p3, lid.x);
|
|
1020
|
+
}
|
|
1021
|
+
|
|
1022
|
+
// Tail is ragged (0-3 extra per thread) \u2014 pad to this workgroup's worst case.
|
|
1023
|
+
// Out-of-range lanes still run the multiply (it carries a barrier, so every
|
|
1024
|
+
// thread must reach it) against index 0, then mask the result to zero.
|
|
1025
|
+
let wgBaseGid = wgid.x * WGS;
|
|
1026
|
+
var tailIters = 0u;
|
|
1027
|
+
if (n4_floor + wgBaseGid < params.n) {
|
|
1028
|
+
tailIters = (params.n - 1u - n4_floor - wgBaseGid) / stride + 1u;
|
|
1029
|
+
}
|
|
1030
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
1031
|
+
let id = n4_floor + gid.x + iter * stride;
|
|
1032
|
+
let valid = id < params.n;
|
|
1033
|
+
let ix = select(0u, id * params.x_inc, valid);
|
|
1034
|
+
let iy = select(0u, id * params.y_inc, valid);
|
|
1035
|
+
let prod = ddMulProtected(DD(xHi[ix], xLo[ix]), DD(yHi[iy], yLo[iy]), lid.x);
|
|
1036
|
+
// select() has no DD overload
|
|
1037
|
+
let contribution = DD(select(0.0, prod.hi, valid), select(0.0, prod.lo, valid));
|
|
1038
|
+
acc0 = ddAddProtected(acc0, contribution, lid.x);
|
|
1039
|
+
}
|
|
1040
|
+
|
|
1041
|
+
let combined01 = ddAddProtected(acc0, acc1, lid.x);
|
|
1042
|
+
let combined23 = ddAddProtected(acc2, acc3, lid.x);
|
|
1043
|
+
tile[lid.x] = ddAddProtected(combined01, combined23, lid.x);
|
|
1044
|
+
workgroupBarrier();
|
|
1045
|
+
|
|
1046
|
+
// Inactive threads combine against a throwaway partner and discard it
|
|
1047
|
+
// (ddAddProtected must be called unconditionally by every thread).
|
|
1048
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
1049
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
1050
|
+
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
1051
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
1052
|
+
if (lid.x < s) { tile[lid.x] = combined; }
|
|
880
1053
|
workgroupBarrier();
|
|
881
|
-
|
|
882
|
-
|
|
883
|
-
|
|
1054
|
+
}
|
|
1055
|
+
|
|
1056
|
+
if (lid.x == 0u) {
|
|
1057
|
+
partialsHi[wgid.x] = tile[0].hi;
|
|
1058
|
+
partialsLo[wgid.x] = tile[0].lo;
|
|
1059
|
+
}
|
|
1060
|
+
}
|
|
1061
|
+
`});var Mt,Pt=O(()=>{Mt=`// dscal: x := alpha * x, double-double (Dekker) f64 emulation of sscal.
|
|
1062
|
+
// alpha and x are each an f32 (hi, lo) pair. See f64/utils/multiply.wgsl for
|
|
1063
|
+
// ddMulProtected and why plain ddMulRaw isn't safe without a renormalizing
|
|
1064
|
+
// barrier \u2014 that barrier needs a provably uniform loop trip count across
|
|
1065
|
+
// every thread in the workgroup, so (like dasum.wgsl's reduction loop) this
|
|
1066
|
+
// splits into a uniform main pass plus a ragged, select-masked tail rather
|
|
1067
|
+
// than a plain \`id < params.n\` grid-stride loop.
|
|
1068
|
+
|
|
1069
|
+
@group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
|
|
1070
|
+
@group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
|
|
1071
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
1072
|
+
|
|
1073
|
+
struct Params {
|
|
1074
|
+
n: u32,
|
|
1075
|
+
alphaHi: f32,
|
|
1076
|
+
alphaLo: f32,
|
|
1077
|
+
x_inc: u32,
|
|
1078
|
+
}
|
|
1079
|
+
|
|
1080
|
+
const WGS: u32 = 64;
|
|
1081
|
+
|
|
1082
|
+
@compute @workgroup_size(64)
|
|
1083
|
+
fn dscal_main(
|
|
1084
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
1085
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1086
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1087
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
1088
|
+
) {
|
|
1089
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
1090
|
+
let stride = num_wg.x * WGS;
|
|
1091
|
+
|
|
1092
|
+
let n_floor = (params.n / stride) * stride;
|
|
1093
|
+
let mainIters = n_floor / stride;
|
|
1094
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
1095
|
+
let id = gid.x + iter * stride;
|
|
1096
|
+
let i = id * params.x_inc;
|
|
1097
|
+
let result = ddMulProtected(alpha, DD(xHi[i], xLo[i]), lid.x);
|
|
1098
|
+
xHi[i] = result.hi;
|
|
1099
|
+
xLo[i] = result.lo;
|
|
1100
|
+
}
|
|
1101
|
+
|
|
1102
|
+
// Tail is ragged (0 or 1 extra per thread) \u2014 pad to this workgroup's worst
|
|
1103
|
+
// case so every thread in the workgroup still calls ddMulProtected the
|
|
1104
|
+
// same number of times (its barrier needs that), masking only the write.
|
|
1105
|
+
let wgBaseGid = wgid.x * WGS;
|
|
1106
|
+
var tailIters = 0u;
|
|
1107
|
+
if (n_floor + wgBaseGid < params.n) {
|
|
1108
|
+
tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
|
|
1109
|
+
}
|
|
1110
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
1111
|
+
let id = n_floor + gid.x + iter * stride;
|
|
1112
|
+
let valid = id < params.n;
|
|
1113
|
+
let i = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
1114
|
+
let result = ddMulProtected(alpha, DD(xHi[i], xLo[i]), lid.x);
|
|
1115
|
+
if (valid) {
|
|
1116
|
+
xHi[i] = result.hi;
|
|
1117
|
+
xLo[i] = result.lo;
|
|
884
1118
|
}
|
|
1119
|
+
}
|
|
1120
|
+
}
|
|
1121
|
+
`});var Lt,It=O(()=>{Lt=`// daxpy: y := alpha * x + y, double-double (Dekker) f64 emulation of saxpy.
|
|
1122
|
+
// Each element costs one ddMulProtected (alpha*x[i]) then one ddAddProtected
|
|
1123
|
+
// (+ y[i]) \u2014 the same two-protected-op shape ddot spends per term, applied
|
|
1124
|
+
// straight to the output instead of folded into a reduction. See dscal.wgsl
|
|
1125
|
+
// for why this is a uniform main pass plus a ragged, select-masked tail
|
|
1126
|
+
// rather than a plain \`id < params.n\` grid-stride loop.
|
|
885
1127
|
|
|
886
|
-
|
|
887
|
-
|
|
1128
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
1129
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
1130
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
1131
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
1132
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
1133
|
+
|
|
1134
|
+
struct Params {
|
|
1135
|
+
n: u32,
|
|
1136
|
+
alphaHi: f32,
|
|
1137
|
+
alphaLo: f32,
|
|
1138
|
+
x_inc: u32,
|
|
1139
|
+
y_inc: u32,
|
|
1140
|
+
}
|
|
1141
|
+
|
|
1142
|
+
const WGS: u32 = 64;
|
|
1143
|
+
|
|
1144
|
+
@compute @workgroup_size(64)
|
|
1145
|
+
fn daxpy_main(
|
|
1146
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
1147
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1148
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1149
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
1150
|
+
) {
|
|
1151
|
+
let alpha = DD(params.alphaHi, params.alphaLo);
|
|
1152
|
+
let stride = num_wg.x * WGS;
|
|
1153
|
+
|
|
1154
|
+
let n_floor = (params.n / stride) * stride;
|
|
1155
|
+
let mainIters = n_floor / stride;
|
|
1156
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
1157
|
+
let id = gid.x + iter * stride;
|
|
1158
|
+
let ix = id * params.x_inc;
|
|
1159
|
+
let iy = id * params.y_inc;
|
|
1160
|
+
let prod = ddMulProtected(alpha, DD(xHi[ix], xLo[ix]), lid.x);
|
|
1161
|
+
let result = ddAddProtected(prod, DD(yHi[iy], yLo[iy]), lid.x);
|
|
1162
|
+
yHi[iy] = result.hi;
|
|
1163
|
+
yLo[iy] = result.lo;
|
|
1164
|
+
}
|
|
1165
|
+
|
|
1166
|
+
// Tail is ragged (0 or 1 extra per thread) \u2014 pad to this workgroup's worst
|
|
1167
|
+
// case so every thread still calls ddMulProtected/ddAddProtected the same
|
|
1168
|
+
// number of times (their barriers need that), masking only the write.
|
|
1169
|
+
let wgBaseGid = wgid.x * WGS;
|
|
1170
|
+
var tailIters = 0u;
|
|
1171
|
+
if (n_floor + wgBaseGid < params.n) {
|
|
1172
|
+
tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
|
|
1173
|
+
}
|
|
1174
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
1175
|
+
let id = n_floor + gid.x + iter * stride;
|
|
1176
|
+
let valid = id < params.n;
|
|
1177
|
+
let ix = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
1178
|
+
let iy = select(0u, id * params.y_inc, valid);
|
|
1179
|
+
let prod = ddMulProtected(alpha, DD(xHi[ix], xLo[ix]), lid.x);
|
|
1180
|
+
let result = ddAddProtected(prod, DD(yHi[iy], yLo[iy]), lid.x);
|
|
1181
|
+
if (valid) {
|
|
1182
|
+
yHi[iy] = result.hi;
|
|
1183
|
+
yLo[iy] = result.lo;
|
|
1184
|
+
}
|
|
1185
|
+
}
|
|
1186
|
+
}
|
|
1187
|
+
`});var Pe,Rt=O(()=>{Pe=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
1188
|
+
|
|
1189
|
+
// a > b for double-double pairs. hi dominates (|lo| <= ulp(hi)/2 always), so
|
|
1190
|
+
// comparing hi alone is correct except on an exact hi tie, when lo breaks it.
|
|
1191
|
+
// A plain comparison, not a rounding-identity subtraction \u2014 no reassociation
|
|
1192
|
+
// risk, so unlike twoSum/fastTwoSum this needs no protection.
|
|
1193
|
+
fn ddGreater(a: DD, b: DD) -> bool {
|
|
1194
|
+
if (a.hi != b.hi) {
|
|
1195
|
+
return a.hi > b.hi;
|
|
1196
|
+
}
|
|
1197
|
+
return a.lo > b.lo;
|
|
1198
|
+
}
|
|
1199
|
+
`});var Tt,qt=O(()=>{Tt=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
1200
|
+
|
|
1201
|
+
// a == b for double-double pairs \u2014 exact field equality, no rounding
|
|
1202
|
+
// involved, so (like ddGreater) this needs no protection.
|
|
1203
|
+
fn ddEqual(a: DD, b: DD) -> bool {
|
|
1204
|
+
return a.hi == b.hi && a.lo == b.lo;
|
|
1205
|
+
}
|
|
1206
|
+
`});var Ft,Ct=O(()=>{Ft=`// idamax: returns index of element with largest absolute value (f64, double-double)
|
|
1207
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/argmaxF64.wgsl.
|
|
1208
|
+
// Concatenated after f64/dekker.wgsl (DD struct), f64/utils/abs.wgsl (ddAbs),
|
|
1209
|
+
// f64/utils/greater.wgsl (ddGreater), and f64/utils/equal.wgsl (ddEqual).
|
|
1210
|
+
|
|
1211
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
1212
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
1213
|
+
@group(0) @binding(2) var<storage, read_write> partialsValHi: array<f32>;
|
|
1214
|
+
@group(0) @binding(3) var<storage, read_write> partialsValLo: array<f32>;
|
|
1215
|
+
@group(0) @binding(4) var<storage, read_write> partialsIdx: array<u32>;
|
|
1216
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
1217
|
+
|
|
1218
|
+
struct Params {
|
|
1219
|
+
n: u32,
|
|
1220
|
+
x_inc: u32,
|
|
1221
|
+
}
|
|
1222
|
+
|
|
1223
|
+
const WGS: u32 = 64;
|
|
1224
|
+
|
|
1225
|
+
var<workgroup> tile_val: array<DD, 64>;
|
|
1226
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
1227
|
+
|
|
1228
|
+
@compute @workgroup_size(64)
|
|
1229
|
+
fn idamax_main(
|
|
1230
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
1231
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1232
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1233
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
1234
|
+
) {
|
|
1235
|
+
// DD(-1.0, 0.0) is a safe sentinel: any |x[i]| >= 0 beats it,
|
|
1236
|
+
// so workgroups with no elements lose gracefully in the epilogue.
|
|
1237
|
+
var best_val0 = DD(-1.0, 0.0); var best_idx0: u32 = 0u;
|
|
1238
|
+
var best_val1 = DD(-1.0, 0.0); var best_idx1: u32 = 0u;
|
|
1239
|
+
var best_val2 = DD(-1.0, 0.0); var best_idx2: u32 = 0u;
|
|
1240
|
+
var best_val3 = DD(-1.0, 0.0); var best_idx3: u32 = 0u;
|
|
1241
|
+
|
|
1242
|
+
let stride = num_wg.x * WGS;
|
|
1243
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
1244
|
+
|
|
1245
|
+
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
1246
|
+
let i0 = id * params.x_inc;
|
|
1247
|
+
let i1 = (id + stride) * params.x_inc;
|
|
1248
|
+
let i2 = (id + 2u * stride) * params.x_inc;
|
|
1249
|
+
let i3 = (id + 3u * stride) * params.x_inc;
|
|
1250
|
+
let v0 = ddAbs(DD(xHi[i0], xLo[i0]));
|
|
1251
|
+
let v1 = ddAbs(DD(xHi[i1], xLo[i1]));
|
|
1252
|
+
let v2 = ddAbs(DD(xHi[i2], xLo[i2]));
|
|
1253
|
+
let v3 = ddAbs(DD(xHi[i3], xLo[i3]));
|
|
1254
|
+
if (ddGreater(v0, best_val0)) { best_val0 = v0; best_idx0 = id; }
|
|
1255
|
+
if (ddGreater(v1, best_val1)) { best_val1 = v1; best_idx1 = id + stride; }
|
|
1256
|
+
if (ddGreater(v2, best_val2)) { best_val2 = v2; best_idx2 = id + 2u * stride; }
|
|
1257
|
+
if (ddGreater(v3, best_val3)) { best_val3 = v3; best_idx3 = id + 3u * stride; }
|
|
1258
|
+
}
|
|
1259
|
+
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
1260
|
+
let i = id * params.x_inc;
|
|
1261
|
+
let v = ddAbs(DD(xHi[i], xLo[i]));
|
|
1262
|
+
if (ddGreater(v, best_val0)) { best_val0 = v; best_idx0 = id; }
|
|
1263
|
+
}
|
|
1264
|
+
|
|
1265
|
+
// merge 4 independent lanes; prefer lower index on tie (first occurrence wins)
|
|
1266
|
+
if (ddGreater(best_val1, best_val0) ||
|
|
1267
|
+
(ddEqual(best_val1, best_val0) && best_idx1 < best_idx0)) {
|
|
1268
|
+
best_val0 = best_val1; best_idx0 = best_idx1;
|
|
1269
|
+
}
|
|
1270
|
+
if (ddGreater(best_val2, best_val0) ||
|
|
1271
|
+
(ddEqual(best_val2, best_val0) && best_idx2 < best_idx0)) {
|
|
1272
|
+
best_val0 = best_val2; best_idx0 = best_idx2;
|
|
1273
|
+
}
|
|
1274
|
+
if (ddGreater(best_val3, best_val0) ||
|
|
1275
|
+
(ddEqual(best_val3, best_val0) && best_idx3 < best_idx0)) {
|
|
1276
|
+
best_val0 = best_val3; best_idx0 = best_idx3;
|
|
1277
|
+
}
|
|
1278
|
+
|
|
1279
|
+
tile_val[lid.x] = best_val0;
|
|
1280
|
+
tile_idx[lid.x] = best_idx0;
|
|
1281
|
+
workgroupBarrier();
|
|
1282
|
+
|
|
1283
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
1284
|
+
if (lid.x < s) {
|
|
1285
|
+
let a_val = tile_val[lid.x];
|
|
1286
|
+
let b_val = tile_val[lid.x + s];
|
|
1287
|
+
if (ddGreater(b_val, a_val) ||
|
|
1288
|
+
(ddEqual(b_val, a_val) && tile_idx[lid.x + s] < tile_idx[lid.x])) {
|
|
1289
|
+
tile_val[lid.x] = b_val;
|
|
1290
|
+
tile_idx[lid.x] = tile_idx[lid.x + s];
|
|
1291
|
+
}
|
|
1292
|
+
}
|
|
1293
|
+
workgroupBarrier();
|
|
1294
|
+
}
|
|
1295
|
+
|
|
1296
|
+
if (lid.x == 0u) {
|
|
1297
|
+
partialsValHi[wgid.x] = tile_val[0].hi;
|
|
1298
|
+
partialsValLo[wgid.x] = tile_val[0].lo;
|
|
1299
|
+
partialsIdx[wgid.x] = tile_idx[0];
|
|
1300
|
+
}
|
|
1301
|
+
}
|
|
1302
|
+
`});var Wt,jt=O(()=>{Wt=`// amax reduction (f64, double-double): collapses 2*WGS (value, index) pairs
|
|
1303
|
+
// into one index, using ddGreater/ddEqual instead of plain f32 \`>\`/\`==\` (see
|
|
1304
|
+
// reduction/argmax.wgsl for the f32 original this mirrors).
|
|
1305
|
+
// dispatch: 1 workgroup of WGS threads. partialsValHi/partialsValLo/
|
|
1306
|
+
// partialsIdx must have exactly 2*WGS entries each. Concatenated after
|
|
1307
|
+
// f64/dekker.wgsl (DD struct), f64/utils/greater.wgsl (ddGreater), and
|
|
1308
|
+
// f64/utils/equal.wgsl (ddEqual).
|
|
1309
|
+
|
|
1310
|
+
@group(0) @binding(0) var<storage, read> partialsValHi: array<f32>;
|
|
1311
|
+
@group(0) @binding(1) var<storage, read> partialsValLo: array<f32>;
|
|
1312
|
+
@group(0) @binding(2) var<storage, read> partialsIdx: array<u32>;
|
|
1313
|
+
@group(0) @binding(3) var<storage, read_write> result: array<u32>;
|
|
1314
|
+
|
|
1315
|
+
const WGS: u32 = 64;
|
|
1316
|
+
|
|
1317
|
+
var<workgroup> tile_val: array<DD, 64>;
|
|
1318
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
1319
|
+
|
|
1320
|
+
@compute @workgroup_size(64)
|
|
1321
|
+
fn reduce_f64(
|
|
1322
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1323
|
+
) {
|
|
1324
|
+
let i = lid.x;
|
|
1325
|
+
let a_val = DD(partialsValHi[i], partialsValLo[i]);
|
|
1326
|
+
let b_val = DD(partialsValHi[i + WGS], partialsValLo[i + WGS]);
|
|
1327
|
+
if (ddGreater(b_val, a_val) ||
|
|
1328
|
+
(ddEqual(b_val, a_val) && partialsIdx[i + WGS] < partialsIdx[i])) {
|
|
1329
|
+
tile_val[i] = b_val;
|
|
1330
|
+
tile_idx[i] = partialsIdx[i + WGS];
|
|
1331
|
+
} else {
|
|
1332
|
+
tile_val[i] = a_val;
|
|
1333
|
+
tile_idx[i] = partialsIdx[i];
|
|
1334
|
+
}
|
|
1335
|
+
workgroupBarrier();
|
|
1336
|
+
|
|
1337
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
1338
|
+
if (i < s) {
|
|
1339
|
+
let c_val = tile_val[i];
|
|
1340
|
+
let d_val = tile_val[i + s];
|
|
1341
|
+
if (ddGreater(d_val, c_val) ||
|
|
1342
|
+
(ddEqual(d_val, c_val) && tile_idx[i + s] < tile_idx[i])) {
|
|
1343
|
+
tile_val[i] = d_val;
|
|
1344
|
+
tile_idx[i] = tile_idx[i + s];
|
|
1345
|
+
}
|
|
888
1346
|
}
|
|
1347
|
+
workgroupBarrier();
|
|
889
1348
|
}
|
|
1349
|
+
|
|
1350
|
+
if (i == 0u) { result[0] = tile_idx[0]; }
|
|
890
1351
|
}
|
|
891
|
-
`});var
|
|
1352
|
+
`});var Ot,Ht=O(()=>{Ot=`// srot: x = c*x + s*y, y = -s*x + c*y
|
|
892
1353
|
|
|
893
|
-
@group(0) @binding(0) var<storage,
|
|
894
|
-
@group(0) @binding(1) var<storage,
|
|
895
|
-
@group(0) @binding(2) var<storage, read_write> A: array<f32>;
|
|
1354
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
1355
|
+
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
896
1356
|
|
|
897
1357
|
struct Params {
|
|
898
|
-
m: u32,
|
|
899
1358
|
n: u32,
|
|
900
|
-
|
|
901
|
-
|
|
902
|
-
|
|
903
|
-
|
|
1359
|
+
c: f32,
|
|
1360
|
+
s: f32,
|
|
1361
|
+
x_inc: u32,
|
|
1362
|
+
y_inc: u32,
|
|
904
1363
|
}
|
|
905
1364
|
|
|
906
|
-
@group(0) @binding(
|
|
1365
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
907
1366
|
|
|
908
|
-
const WGS: u32 =
|
|
1367
|
+
const WGS: u32 = 64;
|
|
909
1368
|
|
|
910
1369
|
@compute @workgroup_size(64)
|
|
911
1370
|
fn main(
|
|
912
|
-
@builtin(
|
|
913
|
-
@builtin(
|
|
914
|
-
@builtin(num_workgroups) nwg: vec3u,
|
|
1371
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
1372
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
915
1373
|
) {
|
|
916
|
-
for (var
|
|
917
|
-
let xi =
|
|
918
|
-
let
|
|
919
|
-
|
|
920
|
-
|
|
921
|
-
let n4_floor = (params.n / (4u * WGS)) * (4u * WGS);
|
|
922
|
-
for (var col: u32 = lid.x; col < n4_floor; col += 4u * WGS) {
|
|
923
|
-
let idx0 = row_base + col;
|
|
924
|
-
let idx1 = row_base + col + WGS;
|
|
925
|
-
let idx2 = row_base + col + 2u * WGS;
|
|
926
|
-
let idx3 = row_base + col + 3u * WGS;
|
|
927
|
-
A[idx0] = xi * y[ col * params.incy] + A[idx0];
|
|
928
|
-
A[idx1] = xi * y[(col + WGS) * params.incy] + A[idx1];
|
|
929
|
-
A[idx2] = xi * y[(col + 2u * WGS) * params.incy] + A[idx2];
|
|
930
|
-
A[idx3] = xi * y[(col + 3u * WGS) * params.incy] + A[idx3];
|
|
931
|
-
}
|
|
932
|
-
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
933
|
-
for (var col: u32 = n4_floor + lid.x; col < params.n; col += WGS) {
|
|
934
|
-
let idx = row_base + col;
|
|
935
|
-
A[idx] = xi * y[col * params.incy] + A[idx];
|
|
936
|
-
}
|
|
1374
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
1375
|
+
let xi = x[id * params.x_inc];
|
|
1376
|
+
let yi = y[id * params.y_inc];
|
|
1377
|
+
x[id * params.x_inc] = params.c * xi + params.s * yi;
|
|
1378
|
+
y[id * params.y_inc] = -params.s * xi + params.c * yi;
|
|
937
1379
|
}
|
|
938
1380
|
}
|
|
939
|
-
`});var
|
|
940
|
-
//
|
|
941
|
-
//
|
|
1381
|
+
`});var Kt,Vt=O(()=>{Kt=`// drot: x = c*x + s*y, y = -s*x + c*y \u2014 double-double (Dekker) f64 emulation
|
|
1382
|
+
// of srot. c, s, x, and y are each split into an f32 (hi, lo) pair. Each
|
|
1383
|
+
// element costs four ddMulProtected (c*x, s*y, -s*x, c*y) then two
|
|
1384
|
+
// ddAddProtected (the two sums) \u2014 negS is computed once outside the loop
|
|
1385
|
+
// via bitcast negation (exact, no rounding, so no barrier needed there)
|
|
1386
|
+
// rather than adding a DD-subtract helper. See dscal.wgsl for why this is a
|
|
1387
|
+
// uniform main pass plus a ragged, select-masked tail rather than a plain
|
|
1388
|
+
// \`id < params.n\` grid-stride loop.
|
|
942
1389
|
|
|
943
|
-
@group(0) @binding(0) var<storage,
|
|
944
|
-
@group(0) @binding(1) var<storage, read_write>
|
|
1390
|
+
@group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
|
|
1391
|
+
@group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
|
|
1392
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
1393
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
1394
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
945
1395
|
|
|
946
1396
|
struct Params {
|
|
947
1397
|
n: u32,
|
|
948
|
-
|
|
949
|
-
|
|
950
|
-
|
|
951
|
-
|
|
1398
|
+
cHi: f32,
|
|
1399
|
+
cLo: f32,
|
|
1400
|
+
sHi: f32,
|
|
1401
|
+
sLo: f32,
|
|
1402
|
+
x_inc: u32,
|
|
1403
|
+
y_inc: u32,
|
|
952
1404
|
}
|
|
953
1405
|
|
|
954
|
-
|
|
955
|
-
|
|
956
|
-
const WGS: u32 = 64u;
|
|
1406
|
+
const WGS: u32 = 64;
|
|
957
1407
|
|
|
958
1408
|
@compute @workgroup_size(64)
|
|
959
|
-
fn
|
|
960
|
-
@builtin(
|
|
961
|
-
@builtin(local_invocation_id)
|
|
962
|
-
@builtin(
|
|
1409
|
+
fn drot_main(
|
|
1410
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
1411
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1412
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1413
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
963
1414
|
) {
|
|
964
|
-
|
|
965
|
-
|
|
966
|
-
|
|
967
|
-
|
|
968
|
-
// Stored-triangle column range for this row: lower [0,row], upper [row,n).
|
|
969
|
-
var colStart: u32;
|
|
970
|
-
var colEnd: u32;
|
|
971
|
-
if params.uplo == 1u {
|
|
972
|
-
colStart = row;
|
|
973
|
-
colEnd = params.n;
|
|
974
|
-
} else {
|
|
975
|
-
colStart = 0u;
|
|
976
|
-
colEnd = row + 1u;
|
|
977
|
-
}
|
|
1415
|
+
let c = DD(params.cHi, params.cLo);
|
|
1416
|
+
let s = DD(params.sHi, params.sLo);
|
|
1417
|
+
let negS = DD(negf(params.sHi), negf(params.sLo));
|
|
1418
|
+
let stride = num_wg.x * WGS;
|
|
978
1419
|
|
|
979
|
-
|
|
980
|
-
|
|
981
|
-
|
|
982
|
-
|
|
983
|
-
|
|
984
|
-
|
|
985
|
-
|
|
986
|
-
|
|
987
|
-
|
|
988
|
-
|
|
989
|
-
|
|
990
|
-
|
|
991
|
-
|
|
992
|
-
|
|
993
|
-
|
|
994
|
-
|
|
995
|
-
|
|
1420
|
+
let n_floor = (params.n / stride) * stride;
|
|
1421
|
+
let mainIters = n_floor / stride;
|
|
1422
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
1423
|
+
let id = gid.x + iter * stride;
|
|
1424
|
+
let ix = id * params.x_inc;
|
|
1425
|
+
let iy = id * params.y_inc;
|
|
1426
|
+
let xi = DD(xHi[ix], xLo[ix]);
|
|
1427
|
+
let yi = DD(yHi[iy], yLo[iy]);
|
|
1428
|
+
let xNew = ddAddProtected(ddMulProtected(c, xi, lid.x), ddMulProtected(s, yi, lid.x), lid.x);
|
|
1429
|
+
let yNew = ddAddProtected(ddMulProtected(negS, xi, lid.x), ddMulProtected(c, yi, lid.x), lid.x);
|
|
1430
|
+
xHi[ix] = xNew.hi;
|
|
1431
|
+
xLo[ix] = xNew.lo;
|
|
1432
|
+
yHi[iy] = yNew.hi;
|
|
1433
|
+
yLo[iy] = yNew.lo;
|
|
1434
|
+
}
|
|
1435
|
+
|
|
1436
|
+
// Tail is ragged (0 or 1 extra per thread) \u2014 pad to this workgroup's worst
|
|
1437
|
+
// case so every thread in the workgroup still calls ddMulProtected/
|
|
1438
|
+
// ddAddProtected the same number of times (their barriers need that),
|
|
1439
|
+
// masking only the write.
|
|
1440
|
+
let wgBaseGid = wgid.x * WGS;
|
|
1441
|
+
var tailIters = 0u;
|
|
1442
|
+
if (n_floor + wgBaseGid < params.n) {
|
|
1443
|
+
tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
|
|
1444
|
+
}
|
|
1445
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
1446
|
+
let id = n_floor + gid.x + iter * stride;
|
|
1447
|
+
let valid = id < params.n;
|
|
1448
|
+
let ix = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
1449
|
+
let iy = select(0u, id * params.y_inc, valid);
|
|
1450
|
+
let xi = DD(xHi[ix], xLo[ix]);
|
|
1451
|
+
let yi = DD(yHi[iy], yLo[iy]);
|
|
1452
|
+
let xNew = ddAddProtected(ddMulProtected(c, xi, lid.x), ddMulProtected(s, yi, lid.x), lid.x);
|
|
1453
|
+
let yNew = ddAddProtected(ddMulProtected(negS, xi, lid.x), ddMulProtected(c, yi, lid.x), lid.x);
|
|
1454
|
+
if (valid) {
|
|
1455
|
+
xHi[ix] = xNew.hi;
|
|
1456
|
+
xLo[ix] = xNew.lo;
|
|
1457
|
+
yHi[iy] = yNew.hi;
|
|
1458
|
+
yLo[iy] = yNew.lo;
|
|
996
1459
|
}
|
|
997
1460
|
}
|
|
998
1461
|
}
|
|
999
|
-
`});var
|
|
1000
|
-
//
|
|
1001
|
-
//
|
|
1462
|
+
`});var Ut,zt=O(()=>{Ut=`// srotm: applies modified Givens rotation H to vectors x and y.
|
|
1463
|
+
// param[0] = flag: -1 (full H), 0 (unit diagonal), 1 (unit off-diagonal)
|
|
1464
|
+
// param = [ flag, h11, h21, h12, h22 ]
|
|
1465
|
+
// flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
|
|
1002
1466
|
|
|
1003
|
-
@group(0) @binding(0) var<storage,
|
|
1004
|
-
@group(0) @binding(1) var<storage,
|
|
1005
|
-
@group(0) @binding(2) var<storage,
|
|
1467
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
1468
|
+
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
1469
|
+
@group(0) @binding(2) var<storage, read> param: array<f32>;
|
|
1006
1470
|
|
|
1007
1471
|
struct Params {
|
|
1008
1472
|
n: u32,
|
|
1009
|
-
|
|
1010
|
-
|
|
1011
|
-
incy: u32,
|
|
1012
|
-
lda: u32,
|
|
1013
|
-
uplo: u32, // 0 = lower, 1 = upper
|
|
1473
|
+
x_inc: u32,
|
|
1474
|
+
y_inc: u32,
|
|
1014
1475
|
}
|
|
1015
1476
|
|
|
1016
1477
|
@group(0) @binding(3) var<uniform> params: Params;
|
|
1017
1478
|
|
|
1018
|
-
const WGS: u32 =
|
|
1479
|
+
const WGS: u32 = 64;
|
|
1019
1480
|
|
|
1020
1481
|
@compute @workgroup_size(64)
|
|
1021
1482
|
fn main(
|
|
1022
|
-
@builtin(
|
|
1023
|
-
@builtin(
|
|
1024
|
-
@builtin(num_workgroups) nwg: vec3u,
|
|
1483
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
1484
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
1025
1485
|
) {
|
|
1026
|
-
|
|
1027
|
-
let xi = params.alpha * x[row * params.incx];
|
|
1028
|
-
let yi = params.alpha * y[row * params.incy];
|
|
1029
|
-
let row_base = row * params.lda;
|
|
1486
|
+
let flag = param[0];
|
|
1030
1487
|
|
|
1031
|
-
|
|
1032
|
-
|
|
1033
|
-
var colEnd: u32;
|
|
1034
|
-
if params.uplo == 1u {
|
|
1035
|
-
colStart = row;
|
|
1036
|
-
colEnd = params.n;
|
|
1037
|
-
} else {
|
|
1038
|
-
colStart = 0u;
|
|
1039
|
-
colEnd = row + 1u;
|
|
1040
|
-
}
|
|
1488
|
+
var h11: f32; var h12: f32;
|
|
1489
|
+
var h21: f32; var h22: f32;
|
|
1041
1490
|
|
|
1042
|
-
|
|
1043
|
-
|
|
1044
|
-
|
|
1045
|
-
|
|
1046
|
-
|
|
1047
|
-
|
|
1048
|
-
|
|
1049
|
-
|
|
1050
|
-
|
|
1051
|
-
|
|
1052
|
-
|
|
1053
|
-
|
|
1054
|
-
|
|
1055
|
-
|
|
1056
|
-
|
|
1057
|
-
|
|
1058
|
-
|
|
1059
|
-
|
|
1491
|
+
if (flag == -1.0) {
|
|
1492
|
+
// full 2x2 matrix
|
|
1493
|
+
h11 = param[1]; h21 = param[2];
|
|
1494
|
+
h12 = param[3]; h22 = param[4];
|
|
1495
|
+
} else if (flag == 0.0) {
|
|
1496
|
+
// diagonal fixed at 1
|
|
1497
|
+
h11 = 1.0; h21 = param[2];
|
|
1498
|
+
h12 = param[3]; h22 = 1.0;
|
|
1499
|
+
} else if (flag == 1.0) {
|
|
1500
|
+
// flag == 1.0: off-diagonal fixed at +1 / -1
|
|
1501
|
+
h11 = param[1]; h21 = -1.0;
|
|
1502
|
+
h12 = 1.0; h22 = param[4];
|
|
1503
|
+
}
|
|
1504
|
+
|
|
1505
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
1506
|
+
let xi = x[id * params.x_inc];
|
|
1507
|
+
let yi = y[id * params.y_inc];
|
|
1508
|
+
x[id * params.x_inc] = h11 * xi + h12 * yi;
|
|
1509
|
+
y[id * params.y_inc] = h21 * xi + h22 * yi;
|
|
1060
1510
|
}
|
|
1061
1511
|
}
|
|
1062
|
-
`});var
|
|
1063
|
-
//
|
|
1064
|
-
//
|
|
1065
|
-
//
|
|
1512
|
+
`});var Xt,Yt=O(()=>{Xt=`// drotm: applies a modified Givens rotation H to vectors x and y \u2014 double-
|
|
1513
|
+
// double (Dekker) f64 emulation of srotm. paramHi/paramLo[0] = flag: -1
|
|
1514
|
+
// (full H), 0 (unit diagonal), 1 (unit off-diagonal). param = [ flag, h11,
|
|
1515
|
+
// h21, h12, h22 ], each entry an f32 (hi, lo) pair.
|
|
1516
|
+
// flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
|
|
1066
1517
|
//
|
|
1067
|
-
//
|
|
1068
|
-
//
|
|
1069
|
-
//
|
|
1070
|
-
//
|
|
1071
|
-
|
|
1072
|
-
|
|
1073
|
-
|
|
1074
|
-
|
|
1075
|
-
|
|
1076
|
-
|
|
1518
|
+
// h11/h12/h21/h22 are resolved once, outside the loop, from the (uniform
|
|
1519
|
+
// across every thread) flag \u2014 same shape as srot's c/s, so no barrier is
|
|
1520
|
+
// needed for that selection itself. Each element then costs four
|
|
1521
|
+
// ddMulProtected + two ddAddProtected, same as drot.
|
|
1522
|
+
|
|
1523
|
+
@group(0) @binding(0) var<storage, read_write> xHi: array<f32>;
|
|
1524
|
+
@group(0) @binding(1) var<storage, read_write> xLo: array<f32>;
|
|
1525
|
+
@group(0) @binding(2) var<storage, read_write> yHi: array<f32>;
|
|
1526
|
+
@group(0) @binding(3) var<storage, read_write> yLo: array<f32>;
|
|
1527
|
+
@group(0) @binding(4) var<storage, read> paramHi: array<f32>;
|
|
1528
|
+
@group(0) @binding(5) var<storage, read> paramLo: array<f32>;
|
|
1529
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
1077
1530
|
|
|
1078
|
-
struct
|
|
1079
|
-
|
|
1080
|
-
|
|
1081
|
-
|
|
1082
|
-
lo: u32, // 32 bits
|
|
1531
|
+
struct Params {
|
|
1532
|
+
n: u32,
|
|
1533
|
+
x_inc: u32,
|
|
1534
|
+
y_inc: u32,
|
|
1083
1535
|
}
|
|
1084
1536
|
|
|
1085
|
-
|
|
1086
|
-
|
|
1087
|
-
|
|
1088
|
-
// slot canonicalizes/corrupts that on any round-trip). See f64pack.mjs's
|
|
1089
|
-
// comment above fieldsToPacked, and dasum.wgsl's xAux/partialsAux bindings.
|
|
1090
|
-
struct Packed {
|
|
1091
|
-
main: f32,
|
|
1092
|
-
aux: u32,
|
|
1093
|
-
}
|
|
1537
|
+
const WGS: u32 = 64;
|
|
1538
|
+
const ONE: DD = DD(1.0, 0.0);
|
|
1539
|
+
const NEG_ONE: DD = DD(-1.0, 0.0);
|
|
1094
1540
|
|
|
1095
|
-
|
|
1096
|
-
fn
|
|
1097
|
-
|
|
1098
|
-
|
|
1099
|
-
|
|
1541
|
+
@compute @workgroup_size(64)
|
|
1542
|
+
fn drotm_main(
|
|
1543
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
1544
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1545
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1546
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
1547
|
+
) {
|
|
1548
|
+
let flag = paramHi[0]; // exact small integer (-1, 0, or 1) \u2014 lo is always 0
|
|
1100
1549
|
|
|
1101
|
-
|
|
1102
|
-
|
|
1103
|
-
|
|
1550
|
+
var h11: DD; var h12: DD;
|
|
1551
|
+
var h21: DD; var h22: DD;
|
|
1552
|
+
|
|
1553
|
+
if (flag == -1.0) {
|
|
1554
|
+
// full 2x2 matrix
|
|
1555
|
+
h11 = DD(paramHi[1], paramLo[1]); h21 = DD(paramHi[2], paramLo[2]);
|
|
1556
|
+
h12 = DD(paramHi[3], paramLo[3]); h22 = DD(paramHi[4], paramLo[4]);
|
|
1557
|
+
} else if (flag == 0.0) {
|
|
1558
|
+
// diagonal fixed at 1
|
|
1559
|
+
h11 = ONE; h21 = DD(paramHi[2], paramLo[2]);
|
|
1560
|
+
h12 = DD(paramHi[3], paramLo[3]); h22 = ONE;
|
|
1561
|
+
} else {
|
|
1562
|
+
// flag == 1.0: off-diagonal fixed at +1 / -1
|
|
1563
|
+
h11 = DD(paramHi[1], paramLo[1]); h21 = NEG_ONE;
|
|
1564
|
+
h12 = ONE; h22 = DD(paramHi[4], paramLo[4]);
|
|
1565
|
+
}
|
|
1104
1566
|
|
|
1105
|
-
let
|
|
1106
|
-
let mantExtra29 = ((auxExp8 & 0x3fu) << 23u) | auxMant23;
|
|
1567
|
+
let stride = num_wg.x * WGS;
|
|
1107
1568
|
|
|
1108
|
-
let
|
|
1109
|
-
let
|
|
1110
|
-
|
|
1111
|
-
|
|
1569
|
+
let n_floor = (params.n / stride) * stride;
|
|
1570
|
+
let mainIters = n_floor / stride;
|
|
1571
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
1572
|
+
let id = gid.x + iter * stride;
|
|
1573
|
+
let ix = id * params.x_inc;
|
|
1574
|
+
let iy = id * params.y_inc;
|
|
1575
|
+
let xi = DD(xHi[ix], xLo[ix]);
|
|
1576
|
+
let yi = DD(yHi[iy], yLo[iy]);
|
|
1577
|
+
let xNew = ddAddProtected(ddMulProtected(h11, xi, lid.x), ddMulProtected(h12, yi, lid.x), lid.x);
|
|
1578
|
+
let yNew = ddAddProtected(ddMulProtected(h21, xi, lid.x), ddMulProtected(h22, yi, lid.x), lid.x);
|
|
1579
|
+
xHi[ix] = xNew.hi;
|
|
1580
|
+
xLo[ix] = xNew.lo;
|
|
1581
|
+
yHi[iy] = yNew.hi;
|
|
1582
|
+
yLo[iy] = yNew.lo;
|
|
1583
|
+
}
|
|
1584
|
+
|
|
1585
|
+
// Tail is ragged (0 or 1 extra per thread) \u2014 pad to this workgroup's worst
|
|
1586
|
+
// case so every thread in the workgroup still calls ddMulProtected/
|
|
1587
|
+
// ddAddProtected the same number of times (their barriers need that),
|
|
1588
|
+
// masking only the write.
|
|
1589
|
+
let wgBaseGid = wgid.x * WGS;
|
|
1590
|
+
var tailIters = 0u;
|
|
1591
|
+
if (n_floor + wgBaseGid < params.n) {
|
|
1592
|
+
tailIters = (params.n - 1u - n_floor - wgBaseGid) / stride + 1u;
|
|
1593
|
+
}
|
|
1594
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
1595
|
+
let id = n_floor + gid.x + iter * stride;
|
|
1596
|
+
let valid = id < params.n;
|
|
1597
|
+
let ix = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
1598
|
+
let iy = select(0u, id * params.y_inc, valid);
|
|
1599
|
+
let xi = DD(xHi[ix], xLo[ix]);
|
|
1600
|
+
let yi = DD(yHi[iy], yLo[iy]);
|
|
1601
|
+
let xNew = ddAddProtected(ddMulProtected(h11, xi, lid.x), ddMulProtected(h12, yi, lid.x), lid.x);
|
|
1602
|
+
let yNew = ddAddProtected(ddMulProtected(h21, xi, lid.x), ddMulProtected(h22, yi, lid.x), lid.x);
|
|
1603
|
+
if (valid) {
|
|
1604
|
+
xHi[ix] = xNew.hi;
|
|
1605
|
+
xLo[ix] = xNew.lo;
|
|
1606
|
+
yHi[iy] = yNew.hi;
|
|
1607
|
+
yLo[iy] = yNew.lo;
|
|
1608
|
+
}
|
|
1609
|
+
}
|
|
1610
|
+
}
|
|
1611
|
+
`});var Zt,$t=O(()=>{Zt=`// Requires f64/dekker.wgsl concatenated first for the DD struct, and
|
|
1612
|
+
// f64/utils/add.wgsl (ddSubProtected/ddAddProtected/negf) and
|
|
1613
|
+
// f64/utils/multiply.wgsl (ddMulProtected).
|
|
1614
|
+
//
|
|
1615
|
+
// Verified empirically on real hardware (NVIDIA GTX 1650 + Intel Mesa
|
|
1616
|
+
// integrated GPU, ~8000 random trials each plus explicit edge cases): no
|
|
1617
|
+
// compiler-reassociation-style corruption of the kind that broke twoSum/
|
|
1618
|
+
// twoProd (see add.wgsl's/multiply.wgsl's own headers) \u2014 every failure
|
|
1619
|
+
// found was a genuine algorithm gap, not a driver miscompile, and both are
|
|
1620
|
+
// fixed below (the b.hi==0.0 guard). One real, expected-shape difference
|
|
1621
|
+
// from every other protected op here: the low-power backend's observed
|
|
1622
|
+
// forward-error factor for this op specifically runs noticeably higher
|
|
1623
|
+
// (~7-14000x eps, vs ~3-9x on high-performance) than ddSqrtProtected's
|
|
1624
|
+
// (~3-4x on both) \u2014 division inherently amplifies input imprecision more
|
|
1625
|
+
// than a sum/product does, so a real routine built on this needs its own
|
|
1626
|
+
// backend-calibrated threshold, same as every other f64 arithmetic routine
|
|
1627
|
+
// in this codebase (see e.g. tests/drot/src/test.drot.js's THRESHOLDS).
|
|
1628
|
+
//
|
|
1629
|
+
// One Newton-style long-division refinement (Bailey/QD-style): q1 = a.hi /
|
|
1630
|
+
// b.hi is a plain f32 quotient, accurate to ~24 bits. Computing the residual
|
|
1631
|
+
// a - q1*b in DD arithmetic (not f32) recovers the bits q1 lost, and a
|
|
1632
|
+
// second plain division of that residual resolves them into a correction
|
|
1633
|
+
// term \u2014 combining q1 + q2 gives roughly double a lone f32 divide's
|
|
1634
|
+
// precision, matching this scheme's ~48-bit double-double target (already
|
|
1635
|
+
// short of real f64's 52 bits, so a second refinement step would chase
|
|
1636
|
+
// precision this representation has no room for).
|
|
1637
|
+
// b.hi == 0.0 makes q1 = a.hi/0.0 already the IEEE-754-correct answer
|
|
1638
|
+
// (\xB1Infinity, or NaN for 0/0) via plain float division, but the refinement
|
|
1639
|
+
// below would corrupt it: p1 = q1*b multiplies that Infinity by a zero
|
|
1640
|
+
// divisor, and Infinity*0 is NaN by definition, poisoning everything after.
|
|
1641
|
+
// Substituting a safe non-zero denominator via select() \u2014 rather than
|
|
1642
|
+
// branching/returning early \u2014 keeps every thread calling ddMulProtected/
|
|
1643
|
+
// ddSubProtected/ddAddProtected unconditionally, which their internal
|
|
1644
|
+
// workgroupBarrier() requires; only the final result is selected between
|
|
1645
|
+
// the refined value and q1's own already-correct answer.
|
|
1646
|
+
fn ddDivProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
1647
|
+
let bIsZero = b.hi == 0.0;
|
|
1648
|
+
let q1 = a.hi / b.hi;
|
|
1649
|
+
let safeB = DD(select(b.hi, 1.0, bIsZero), select(b.lo, 0.0, bIsZero));
|
|
1650
|
+
let p1 = ddMulProtected(DD(q1, 0.0), safeB, threadSlot);
|
|
1651
|
+
let r1 = ddSubProtected(a, p1, threadSlot);
|
|
1652
|
+
let q2 = r1.hi / safeB.hi;
|
|
1653
|
+
let refined = ddAddProtected(DD(q1, 0.0), DD(q2, 0.0), threadSlot);
|
|
1654
|
+
return DD(select(refined.hi, q1, bIsZero), select(refined.lo, 0.0, bIsZero));
|
|
1655
|
+
}
|
|
1656
|
+
`});var Jt,Qt=O(()=>{Jt=`// Requires f64/dekker.wgsl concatenated first for the DD struct, and
|
|
1657
|
+
// f64/utils/add.wgsl (ddSubProtected/ddAddProtected) and
|
|
1658
|
+
// f64/utils/multiply.wgsl (twoProdBit \u2014 squaring a plain f32 needs no
|
|
1659
|
+
// barrier, per multiply.wgsl's own note that twoProdBit is universally safe
|
|
1660
|
+
// unprotected).
|
|
1661
|
+
//
|
|
1662
|
+
// Verified empirically on real hardware (NVIDIA GTX 1650 + Intel Mesa
|
|
1663
|
+
// integrated GPU, ~8000 random trials each plus explicit edge cases): no
|
|
1664
|
+
// compiler-reassociation-style corruption of the kind that broke twoSum/
|
|
1665
|
+
// twoProd (see add.wgsl's/multiply.wgsl's own headers) \u2014 every failure
|
|
1666
|
+
// found was a genuine algorithm gap (the a.hi==0.0 case below), not a
|
|
1667
|
+
// driver miscompile, and is fixed. Observed forward-error factor against a
|
|
1668
|
+
// true f64 reference stayed ~3-4x eps on both backends across every random
|
|
1669
|
+
// trial \u2014 noticeably tighter than ddDivProtected's own low-power spread
|
|
1670
|
+
// (see divide.wgsl's header), since sqrt has no denominator to be unlucky
|
|
1671
|
+
// about.
|
|
1672
|
+
//
|
|
1673
|
+
// One Newton refinement step (the classic extended-precision sqrt trick):
|
|
1674
|
+
// x0 = sqrt(a.hi) is a plain f32 approximation; the residual a - x0^2,
|
|
1675
|
+
// computed in DD arithmetic, recovers what x0 lost, and linearizing sqrt
|
|
1676
|
+
// around x0 (dividing that residual by 2*x0) gives a correction term
|
|
1677
|
+
// roughly doubling the precision \u2014 same ~48-bit target as ddDivProtected,
|
|
1678
|
+
// so one step is enough.
|
|
1679
|
+
//
|
|
1680
|
+
// Undefined for a.hi < 0.0, same as plain sqrt() \u2014 callers must guard
|
|
1681
|
+
// themselves; this never checks.
|
|
1682
|
+
//
|
|
1683
|
+
// a.hi == 0.0 (a genuinely zero input, not an underflowed one \u2014 zero is
|
|
1684
|
+
// exactly representable in f32, unlike this scheme's real range limits;
|
|
1685
|
+
// see splitDoubleDouble's own doc comment) makes x0 = sqrt(0) = 0, and the
|
|
1686
|
+
// correction step would divide by 2*x0 = 0. Substituting a safe non-zero
|
|
1687
|
+
// denominator via select() \u2014 rather than branching/returning early \u2014 keeps
|
|
1688
|
+
// every thread calling ddSubProtected/ddAddProtected unconditionally, which
|
|
1689
|
+
// their internal workgroupBarrier() requires; only the final result is
|
|
1690
|
+
// selected between the computed value and the exact DD(0,0) answer.
|
|
1691
|
+
fn ddSqrtProtected(a: DD, threadSlot: u32) -> DD {
|
|
1692
|
+
let isZero = a.hi == 0.0;
|
|
1693
|
+
let x0 = sqrt(a.hi);
|
|
1694
|
+
let x0sq = twoProdBit(x0, x0);
|
|
1695
|
+
let r = ddSubProtected(a, x0sq, threadSlot);
|
|
1696
|
+
let safeDenom = select(2.0 * x0, 1.0, isZero);
|
|
1697
|
+
let correction = r.hi / safeDenom;
|
|
1698
|
+
let result = ddAddProtected(DD(x0, 0.0), DD(correction, 0.0), threadSlot);
|
|
1699
|
+
return DD(select(result.hi, 0.0, isZero), select(result.lo, 0.0, isZero));
|
|
1700
|
+
}
|
|
1701
|
+
`});var eo,ro=O(()=>{eo=`// dnrm2: result = sqrt(sum(x[i] * x[i])), double-double (Dekker) f64
|
|
1702
|
+
// emulation of snrm2 \u2014 same scaled accumulation (Blue's algorithm), just
|
|
1703
|
+
// with \`scale\`/\`ssq\` as DD pairs (via ddDivProtected/ddMulProtected/
|
|
1704
|
+
// ddAddProtected/ddSqrtProtected) instead of plain f32. Squaring still
|
|
1705
|
+
// saturates an f32 hi component above ~1.8e19 regardless of DD precision
|
|
1706
|
+
// (DD widens the mantissa, not the exponent range), so the scaling is
|
|
1707
|
+
// still needed here for the same reason it was in snrm2.
|
|
1708
|
+
//
|
|
1709
|
+
// snrm2.wgsl's ssqAccum/ssqMerge each branch on which operand is bigger \u2014
|
|
1710
|
+
// can't carry over directly, since a protected op's workgroupBarrier()
|
|
1711
|
+
// needs every thread to reach the same call site, and here different
|
|
1712
|
+
// threads could take different branches. Both formulas are computed
|
|
1713
|
+
// unconditionally below; only the final combine (\`ddSelect\`) differs per
|
|
1714
|
+
// thread \u2014 same fix shape as drot's/drotm's own per-dispatch flags, just
|
|
1715
|
+
// applied to a per-element branch instead.
|
|
1716
|
+
//
|
|
1717
|
+
// pass 1 dispatches 2*WGS workgroups; pass 2 (reduction/scaledSumF64.wgsl)
|
|
1718
|
+
// duplicates ssqAccumProtected/ssqMergeProtected rather than sharing them,
|
|
1719
|
+
// same as the plain-f32 pair already does.
|
|
1720
|
+
|
|
1721
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
1722
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
1723
|
+
@group(0) @binding(2) var<storage, read_write> partialsScaleHi: array<f32>;
|
|
1724
|
+
@group(0) @binding(3) var<storage, read_write> partialsScaleLo: array<f32>;
|
|
1725
|
+
@group(0) @binding(4) var<storage, read_write> partialsSsqHi: array<f32>;
|
|
1726
|
+
@group(0) @binding(5) var<storage, read_write> partialsSsqLo: array<f32>;
|
|
1727
|
+
@group(0) @binding(6) var<uniform> params: Params;
|
|
1112
1728
|
|
|
1113
|
-
|
|
1729
|
+
struct Params {
|
|
1730
|
+
n: u32,
|
|
1731
|
+
x_inc: u32,
|
|
1114
1732
|
}
|
|
1115
1733
|
|
|
1116
|
-
|
|
1117
|
-
|
|
1118
|
-
|
|
1119
|
-
|
|
1734
|
+
const WGS: u32 = 64;
|
|
1735
|
+
|
|
1736
|
+
struct ScaleSsq {
|
|
1737
|
+
scale: DD,
|
|
1738
|
+
ssq: DD,
|
|
1739
|
+
}
|
|
1740
|
+
|
|
1741
|
+
fn ddSelect(a: DD, b: DD, cond: bool) -> DD {
|
|
1742
|
+
return DD(select(a.hi, b.hi, cond), select(a.lo, b.lo, cond));
|
|
1743
|
+
}
|
|
1744
|
+
|
|
1745
|
+
// Folds one more |value| (DD) into a running (scale, ssq) pair \u2014 branch-free,
|
|
1746
|
+
// see file header. \`bigger\`/\`smaller\` name the two operands by magnitude
|
|
1747
|
+
// (not by which one was "acc" vs "new"), and biggerIsZero==true only when
|
|
1748
|
+
// both scale and absxi are still exactly zero (the very first zero
|
|
1749
|
+
// elements, before any nonzero value has been seen) \u2014 substituting a safe
|
|
1750
|
+
// denominator there avoids a 0/0 without needing a separate branch/return;
|
|
1751
|
+
// the arithmetic already reduces to a correct no-op in that case.
|
|
1752
|
+
fn ssqAccumProtected(acc: ScaleSsq, absxi: DD, threadSlot: u32) -> ScaleSsq {
|
|
1753
|
+
let isBigger = ddGreater(absxi, acc.scale);
|
|
1754
|
+
let bigger = ddSelect(acc.scale, absxi, isBigger);
|
|
1755
|
+
let smaller = ddSelect(absxi, acc.scale, isBigger);
|
|
1756
|
+
let biggerIsZero = bigger.hi == 0.0;
|
|
1757
|
+
let safeBigger = ddSelect(bigger, DD(1.0, 0.0), biggerIsZero);
|
|
1758
|
+
let r = ddDivProtected(smaller, safeBigger, threadSlot);
|
|
1759
|
+
let rsq = ddMulProtected(r, r, threadSlot);
|
|
1760
|
+
let ssqTimesRsq = ddMulProtected(acc.ssq, rsq, threadSlot);
|
|
1761
|
+
let sumIfBigger = ddAddProtected(DD(1.0, 0.0), ssqTimesRsq, threadSlot);
|
|
1762
|
+
let sumIfNotBigger = ddAddProtected(acc.ssq, rsq, threadSlot);
|
|
1763
|
+
let newSsq = ddSelect(sumIfNotBigger, sumIfBigger, isBigger);
|
|
1764
|
+
return ScaleSsq(bigger, newSsq);
|
|
1765
|
+
}
|
|
1766
|
+
|
|
1767
|
+
// Associative merge of two independent (scale, ssq) partials \u2014 same
|
|
1768
|
+
// branch-free shape, for combining ILP lanes and the tree reduction.
|
|
1769
|
+
fn ssqMergeProtected(a: ScaleSsq, b: ScaleSsq, threadSlot: u32) -> ScaleSsq {
|
|
1770
|
+
let isBigger = !ddGreater(b.scale, a.scale); // a.scale >= b.scale
|
|
1771
|
+
let bigger = ddSelect(b.scale, a.scale, isBigger);
|
|
1772
|
+
let smaller = ddSelect(a.scale, b.scale, isBigger);
|
|
1773
|
+
let biggerSsq = ddSelect(b.ssq, a.ssq, isBigger);
|
|
1774
|
+
let smallerSsq = ddSelect(a.ssq, b.ssq, isBigger);
|
|
1775
|
+
let biggerIsZero = bigger.hi == 0.0;
|
|
1776
|
+
let safeBigger = ddSelect(bigger, DD(1.0, 0.0), biggerIsZero);
|
|
1777
|
+
let r = ddDivProtected(smaller, safeBigger, threadSlot);
|
|
1778
|
+
let rsq = ddMulProtected(r, r, threadSlot);
|
|
1779
|
+
let smallerSsqTimesRsq = ddMulProtected(smallerSsq, rsq, threadSlot);
|
|
1780
|
+
let newSsq = ddAddProtected(biggerSsq, smallerSsqTimesRsq, threadSlot);
|
|
1781
|
+
return ScaleSsq(bigger, newSsq);
|
|
1782
|
+
}
|
|
1783
|
+
|
|
1784
|
+
var<workgroup> tileScaleHi: array<f32, 64>;
|
|
1785
|
+
var<workgroup> tileScaleLo: array<f32, 64>;
|
|
1786
|
+
var<workgroup> tileSsqHi: array<f32, 64>;
|
|
1787
|
+
var<workgroup> tileSsqLo: array<f32, 64>;
|
|
1788
|
+
|
|
1789
|
+
@compute @workgroup_size(64)
|
|
1790
|
+
fn dnrm2_main(
|
|
1791
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
1792
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1793
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1794
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
1795
|
+
) {
|
|
1796
|
+
var acc0 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
|
|
1797
|
+
var acc1 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
|
|
1798
|
+
var acc2 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
|
|
1799
|
+
var acc3 = ScaleSsq(DD(0.0, 0.0), DD(1.0, 0.0));
|
|
1800
|
+
|
|
1801
|
+
let stride = num_wg.x * WGS;
|
|
1802
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
1120
1803
|
|
|
1121
|
-
|
|
1122
|
-
|
|
1123
|
-
let
|
|
1804
|
+
// Same trip count for every thread, driven by a counter (protected ops'
|
|
1805
|
+
// barriers need a provably-uniform loop bound) \u2014 see dasum.wgsl.
|
|
1806
|
+
let mainIters = n4_floor / (4u * stride);
|
|
1807
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
1808
|
+
let id = gid.x + iter * 4u * stride;
|
|
1809
|
+
let i0 = id * params.x_inc;
|
|
1810
|
+
let i1 = (id + stride) * params.x_inc;
|
|
1811
|
+
let i2 = (id + 2u * stride) * params.x_inc;
|
|
1812
|
+
let i3 = (id + 3u * stride) * params.x_inc;
|
|
1813
|
+
acc0 = ssqAccumProtected(acc0, ddAbs(DD(xHi[i0], xLo[i0])), lid.x);
|
|
1814
|
+
acc1 = ssqAccumProtected(acc1, ddAbs(DD(xHi[i1], xLo[i1])), lid.x);
|
|
1815
|
+
acc2 = ssqAccumProtected(acc2, ddAbs(DD(xHi[i2], xLo[i2])), lid.x);
|
|
1816
|
+
acc3 = ssqAccumProtected(acc3, ddAbs(DD(xHi[i3], xLo[i3])), lid.x);
|
|
1817
|
+
}
|
|
1124
1818
|
|
|
1125
|
-
|
|
1819
|
+
// Tail is ragged (0-3 extra per thread) \u2014 pad to this workgroup's worst
|
|
1820
|
+
// case, masking an invalid element to exactly 0 (contributes nothing).
|
|
1821
|
+
let wgBaseGid = wgid.x * WGS;
|
|
1822
|
+
var tailIters = 0u;
|
|
1823
|
+
if (n4_floor + wgBaseGid < params.n) {
|
|
1824
|
+
tailIters = (params.n - 1u - n4_floor - wgBaseGid) / stride + 1u;
|
|
1825
|
+
}
|
|
1826
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
1827
|
+
let id = n4_floor + gid.x + iter * stride;
|
|
1828
|
+
let valid = id < params.n;
|
|
1829
|
+
let i = select(0u, id * params.x_inc, valid); // index 0 always in-bounds
|
|
1830
|
+
let loaded = ddAbs(DD(xHi[i], xLo[i]));
|
|
1831
|
+
let contribution = DD(select(0.0, loaded.hi, valid), select(0.0, loaded.lo, valid));
|
|
1832
|
+
acc0 = ssqAccumProtected(acc0, contribution, lid.x);
|
|
1833
|
+
}
|
|
1126
1834
|
|
|
1127
|
-
let
|
|
1128
|
-
let
|
|
1129
|
-
let
|
|
1130
|
-
|
|
1131
|
-
|
|
1835
|
+
let combined01 = ssqMergeProtected(acc0, acc1, lid.x);
|
|
1836
|
+
let combined23 = ssqMergeProtected(acc2, acc3, lid.x);
|
|
1837
|
+
let combined = ssqMergeProtected(combined01, combined23, lid.x);
|
|
1838
|
+
tileScaleHi[lid.x] = combined.scale.hi;
|
|
1839
|
+
tileScaleLo[lid.x] = combined.scale.lo;
|
|
1840
|
+
tileSsqHi[lid.x] = combined.ssq.hi;
|
|
1841
|
+
tileSsqLo[lid.x] = combined.ssq.lo;
|
|
1842
|
+
workgroupBarrier();
|
|
1132
1843
|
|
|
1133
|
-
|
|
1844
|
+
// Inactive threads merge against a throwaway partner and discard it
|
|
1845
|
+
// (ssqMergeProtected must be called unconditionally by every thread).
|
|
1846
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
1847
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
1848
|
+
let a = ScaleSsq(DD(tileScaleHi[lid.x], tileScaleLo[lid.x]), DD(tileSsqHi[lid.x], tileSsqLo[lid.x]));
|
|
1849
|
+
let b = ScaleSsq(DD(tileScaleHi[partner], tileScaleLo[partner]), DD(tileSsqHi[partner], tileSsqLo[partner]));
|
|
1850
|
+
let merged = ssqMergeProtected(a, b, lid.x);
|
|
1851
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
1852
|
+
if (lid.x < s) {
|
|
1853
|
+
tileScaleHi[lid.x] = merged.scale.hi;
|
|
1854
|
+
tileScaleLo[lid.x] = merged.scale.lo;
|
|
1855
|
+
tileSsqHi[lid.x] = merged.ssq.hi;
|
|
1856
|
+
tileSsqLo[lid.x] = merged.ssq.lo;
|
|
1857
|
+
}
|
|
1858
|
+
workgroupBarrier();
|
|
1859
|
+
}
|
|
1134
1860
|
|
|
1135
|
-
|
|
1136
|
-
|
|
1861
|
+
if (lid.x == 0u) {
|
|
1862
|
+
partialsScaleHi[wgid.x] = tileScaleHi[0];
|
|
1863
|
+
partialsScaleLo[wgid.x] = tileScaleLo[0];
|
|
1864
|
+
partialsSsqHi[wgid.x] = tileSsqHi[0];
|
|
1865
|
+
partialsSsqLo[wgid.x] = tileSsqLo[0];
|
|
1866
|
+
}
|
|
1867
|
+
}
|
|
1868
|
+
`});var oo,to=O(()=>{oo=`// scaledSum reduction (f64, double-double): collapses 2*WGS (scale, ssq) DD
|
|
1869
|
+
// partials from dnrm2.wgsl into the final norm \u2014 sqrt(scale\xB2 \xB7 ssq) ==
|
|
1870
|
+
// scale \xB7 sqrt(ssq), via ddMulProtected/ddSqrtProtected. Mirrors
|
|
1871
|
+
// reduction/scaledSum.wgsl's shape exactly; ssqMergeProtected is duplicated
|
|
1872
|
+
// from dnrm2.wgsl rather than shared via f64/utils/ \u2014 see that file's own
|
|
1873
|
+
// header for why (same convention the f32 pair already uses).
|
|
1874
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
1875
|
+
// partialsScale*/partialsSsq* must have exactly 2*WGS entries each.
|
|
1137
1876
|
|
|
1138
|
-
|
|
1139
|
-
|
|
1877
|
+
@group(0) @binding(0) var<storage, read> partialsScaleHi: array<f32>;
|
|
1878
|
+
@group(0) @binding(1) var<storage, read> partialsScaleLo: array<f32>;
|
|
1879
|
+
@group(0) @binding(2) var<storage, read> partialsSsqHi: array<f32>;
|
|
1880
|
+
@group(0) @binding(3) var<storage, read> partialsSsqLo: array<f32>;
|
|
1881
|
+
@group(0) @binding(4) var<storage, read_write> resultHi: array<f32, 1>;
|
|
1882
|
+
@group(0) @binding(5) var<storage, read_write> resultLo: array<f32, 1>;
|
|
1140
1883
|
|
|
1141
|
-
|
|
1142
|
-
// a returned sticky flag \u2014 used only for the (potentially huge) exponent
|
|
1143
|
-
// alignment shift, where exact bits can't all be kept.
|
|
1144
|
-
fn shr_sticky(hi: u32, lo: u32, n: u32) -> Shifted {
|
|
1145
|
-
if (n == 0u) {
|
|
1146
|
-
return Shifted(hi, lo, 0u);
|
|
1147
|
-
}
|
|
1148
|
-
if (n >= 64u) {
|
|
1149
|
-
return Shifted(0u, 0u, select(0u, 1u, hi != 0u || lo != 0u));
|
|
1150
|
-
}
|
|
1151
|
-
if (n < 32u) {
|
|
1152
|
-
let stickyBits = lo & ((1u << n) - 1u);
|
|
1153
|
-
let newLo = (lo >> n) | (hi << (32u - n));
|
|
1154
|
-
let newHi = hi >> n;
|
|
1155
|
-
return Shifted(newHi, newLo, select(0u, 1u, stickyBits != 0u));
|
|
1156
|
-
}
|
|
1157
|
-
if (n == 32u) {
|
|
1158
|
-
return Shifted(0u, hi, select(0u, 1u, lo != 0u));
|
|
1159
|
-
}
|
|
1160
|
-
let m = n - 32u;
|
|
1161
|
-
let stickyBits = lo | (hi & ((1u << m) - 1u));
|
|
1162
|
-
let newLo = hi >> m;
|
|
1163
|
-
return Shifted(0u, newLo, select(0u, 1u, stickyBits != 0u));
|
|
1164
|
-
}
|
|
1884
|
+
const WGS: u32 = 64;
|
|
1165
1885
|
|
|
1166
|
-
|
|
1167
|
-
|
|
1168
|
-
|
|
1169
|
-
fn shl(hi: u32, lo: u32, n: u32) -> Pair {
|
|
1170
|
-
if (n == 0u) {
|
|
1171
|
-
return Pair(hi, lo);
|
|
1172
|
-
}
|
|
1173
|
-
if (n < 32u) {
|
|
1174
|
-
let newHi = (hi << n) | (lo >> (32u - n));
|
|
1175
|
-
let newLo = lo << n;
|
|
1176
|
-
return Pair(newHi, newLo);
|
|
1177
|
-
}
|
|
1178
|
-
if (n == 32u) {
|
|
1179
|
-
return Pair(lo, 0u);
|
|
1180
|
-
}
|
|
1181
|
-
let m = n - 32u;
|
|
1182
|
-
return Pair(lo << m, 0u);
|
|
1886
|
+
struct ScaleSsq {
|
|
1887
|
+
scale: DD,
|
|
1888
|
+
ssq: DD,
|
|
1183
1889
|
}
|
|
1184
1890
|
|
|
1185
|
-
fn
|
|
1186
|
-
|
|
1187
|
-
let carry = select(0u, 1u, sumLo < aLo); // wrapped around -> there was a carry
|
|
1188
|
-
let sumHi = aHi + bHi + carry;
|
|
1189
|
-
return Pair(sumHi, sumLo);
|
|
1891
|
+
fn ddSelect(a: DD, b: DD, cond: bool) -> DD {
|
|
1892
|
+
return DD(select(a.hi, b.hi, cond), select(a.lo, b.lo, cond));
|
|
1190
1893
|
}
|
|
1191
1894
|
|
|
1192
|
-
//
|
|
1193
|
-
|
|
1194
|
-
|
|
1195
|
-
let
|
|
1196
|
-
let
|
|
1197
|
-
|
|
1895
|
+
// Associative merge of two independent (scale, ssq) partials \u2014 see
|
|
1896
|
+
// dnrm2.wgsl for the derivation and why this is branch-free.
|
|
1897
|
+
fn ssqMergeProtected(a: ScaleSsq, b: ScaleSsq, threadSlot: u32) -> ScaleSsq {
|
|
1898
|
+
let isBigger = !ddGreater(b.scale, a.scale); // a.scale >= b.scale
|
|
1899
|
+
let bigger = ddSelect(b.scale, a.scale, isBigger);
|
|
1900
|
+
let smaller = ddSelect(a.scale, b.scale, isBigger);
|
|
1901
|
+
let biggerSsq = ddSelect(b.ssq, a.ssq, isBigger);
|
|
1902
|
+
let smallerSsq = ddSelect(a.ssq, b.ssq, isBigger);
|
|
1903
|
+
let biggerIsZero = bigger.hi == 0.0;
|
|
1904
|
+
let safeBigger = ddSelect(bigger, DD(1.0, 0.0), biggerIsZero);
|
|
1905
|
+
let r = ddDivProtected(smaller, safeBigger, threadSlot);
|
|
1906
|
+
let rsq = ddMulProtected(r, r, threadSlot);
|
|
1907
|
+
let smallerSsqTimesRsq = ddMulProtected(smallerSsq, rsq, threadSlot);
|
|
1908
|
+
let newSsq = ddAddProtected(biggerSsq, smallerSsqTimesRsq, threadSlot);
|
|
1909
|
+
return ScaleSsq(bigger, newSsq);
|
|
1198
1910
|
}
|
|
1199
1911
|
|
|
1200
|
-
|
|
1201
|
-
|
|
1202
|
-
|
|
1912
|
+
var<workgroup> tileScaleHi: array<f32, 64>;
|
|
1913
|
+
var<workgroup> tileScaleLo: array<f32, 64>;
|
|
1914
|
+
var<workgroup> tileSsqHi: array<f32, 64>;
|
|
1915
|
+
var<workgroup> tileSsqLo: array<f32, 64>;
|
|
1203
1916
|
|
|
1204
|
-
|
|
1205
|
-
|
|
1206
|
-
|
|
1207
|
-
|
|
1208
|
-
|
|
1209
|
-
|
|
1210
|
-
|
|
1211
|
-
let
|
|
1212
|
-
|
|
1213
|
-
|
|
1214
|
-
|
|
1215
|
-
|
|
1917
|
+
@compute @workgroup_size(64)
|
|
1918
|
+
fn reduce_scaled_f64(
|
|
1919
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1920
|
+
) {
|
|
1921
|
+
let i = lid.x;
|
|
1922
|
+
let a = ScaleSsq(DD(partialsScaleHi[i], partialsScaleLo[i]), DD(partialsSsqHi[i], partialsSsqLo[i]));
|
|
1923
|
+
let b = ScaleSsq(DD(partialsScaleHi[i + WGS], partialsScaleLo[i + WGS]), DD(partialsSsqHi[i + WGS], partialsSsqLo[i + WGS]));
|
|
1924
|
+
let merged0 = ssqMergeProtected(a, b, i);
|
|
1925
|
+
tileScaleHi[i] = merged0.scale.hi;
|
|
1926
|
+
tileScaleLo[i] = merged0.scale.lo;
|
|
1927
|
+
tileSsqHi[i] = merged0.ssq.hi;
|
|
1928
|
+
tileSsqLo[i] = merged0.ssq.lo;
|
|
1929
|
+
workgroupBarrier();
|
|
1216
1930
|
|
|
1217
|
-
|
|
1218
|
-
|
|
1219
|
-
|
|
1220
|
-
|
|
1221
|
-
|
|
1931
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
1932
|
+
let partner = select(i, i + s, i < s);
|
|
1933
|
+
let ai = ScaleSsq(DD(tileScaleHi[i], tileScaleLo[i]), DD(tileSsqHi[i], tileSsqLo[i]));
|
|
1934
|
+
let bi = ScaleSsq(DD(tileScaleHi[partner], tileScaleLo[partner]), DD(tileSsqHi[partner], tileSsqLo[partner]));
|
|
1935
|
+
let merged = ssqMergeProtected(ai, bi, i);
|
|
1936
|
+
workgroupBarrier();
|
|
1937
|
+
if (i < s) {
|
|
1938
|
+
tileScaleHi[i] = merged.scale.hi;
|
|
1939
|
+
tileScaleLo[i] = merged.scale.lo;
|
|
1940
|
+
tileSsqHi[i] = merged.ssq.hi;
|
|
1941
|
+
tileSsqLo[i] = merged.ssq.lo;
|
|
1222
1942
|
}
|
|
1223
|
-
|
|
1224
|
-
}
|
|
1225
|
-
if (aIsInf) { return Fields(a.sign, EXP_ALL_ONES, 0u, 0u); }
|
|
1226
|
-
if (bIsInf) { return Fields(b.sign, EXP_ALL_ONES, 0u, 0u); }
|
|
1227
|
-
|
|
1228
|
-
let aIsZero = a.rawExp == 0u && a.mantissaHi == 0u && a.lo == 0u;
|
|
1229
|
-
let bIsZero = b.rawExp == 0u && b.mantissaHi == 0u && b.lo == 0u;
|
|
1230
|
-
if (aIsZero && bIsZero) {
|
|
1231
|
-
return Fields(a.sign & b.sign, 0u, 0u, 0u);
|
|
1232
|
-
}
|
|
1233
|
-
if (aIsZero) { return Fields(b.sign, b.rawExp, b.mantissaHi, b.lo); }
|
|
1234
|
-
if (bIsZero) { return Fields(a.sign, a.rawExp, a.mantissaHi, a.lo); }
|
|
1235
|
-
|
|
1236
|
-
// Effective (unbiased) exponent \u2014 subnormals share the smallest normal
|
|
1237
|
-
// exponent for alignment purposes and have no implicit leading 1.
|
|
1238
|
-
var expA = i32(a.rawExp) - BIAS;
|
|
1239
|
-
if (a.rawExp == 0u) { expA = 1 - BIAS; }
|
|
1240
|
-
var expB = i32(b.rawExp) - BIAS;
|
|
1241
|
-
if (b.rawExp == 0u) { expB = 1 - BIAS; }
|
|
1242
|
-
|
|
1243
|
-
let implicitA = select(0u, 1u, a.rawExp != 0u);
|
|
1244
|
-
let implicitB = select(0u, 1u, b.rawExp != 0u);
|
|
1245
|
-
|
|
1246
|
-
// Widen each 53-bit significand (implicit + 52-bit mantissa) by 3 zero
|
|
1247
|
-
// bits at the bottom \u2014 room for guard/round/sticky once alignment shifts happen.
|
|
1248
|
-
let sigHiA = (implicitA << 23u) | (a.mantissaHi << 3u) | (a.lo >> 29u);
|
|
1249
|
-
let sigLoA = a.lo << 3u;
|
|
1250
|
-
let sigHiB = (implicitB << 23u) | (b.mantissaHi << 3u) | (b.lo >> 29u);
|
|
1251
|
-
let sigLoB = b.lo << 3u;
|
|
1252
|
-
|
|
1253
|
-
// P = the operand with the larger exponent (Q = the other); on a tie, P =
|
|
1254
|
-
// whichever has the larger significand \u2014 keeps subtraction below always
|
|
1255
|
-
// non-negative without needing signed magnitudes.
|
|
1256
|
-
var signP: u32; var expP: i32; var sigHiP: u32; var sigLoP: u32;
|
|
1257
|
-
var signQ: u32; var expQ: i32; var sigHiQ: u32; var sigLoQ: u32;
|
|
1258
|
-
if (expA > expB || (expA == expB && ge64(sigHiA, sigLoA, sigHiB, sigLoB))) {
|
|
1259
|
-
signP = a.sign; expP = expA; sigHiP = sigHiA; sigLoP = sigLoA;
|
|
1260
|
-
signQ = b.sign; expQ = expB; sigHiQ = sigHiB; sigLoQ = sigLoB;
|
|
1261
|
-
} else {
|
|
1262
|
-
signP = b.sign; expP = expB; sigHiP = sigHiB; sigLoP = sigLoB;
|
|
1263
|
-
signQ = a.sign; expQ = expA; sigHiQ = sigHiA; sigLoQ = sigLoA;
|
|
1943
|
+
workgroupBarrier();
|
|
1264
1944
|
}
|
|
1265
1945
|
|
|
1266
|
-
|
|
1267
|
-
|
|
1268
|
-
|
|
1269
|
-
|
|
1270
|
-
|
|
1271
|
-
|
|
1272
|
-
|
|
1273
|
-
|
|
1274
|
-
|
|
1275
|
-
|
|
1276
|
-
|
|
1277
|
-
|
|
1946
|
+
// ddSqrtProtected/ddMulProtected's own workgroupBarrier()s need every
|
|
1947
|
+
// thread to call them \u2014 every thread redundantly computes the same final
|
|
1948
|
+
// scale\xB7sqrt(ssq) from tile[0] (still visible to all after the reduction
|
|
1949
|
+
// above), and only the write-back is conditional. Guarding the calls
|
|
1950
|
+
// themselves behind \`if (i == 0u)\` (as the plain-f32 original safely
|
|
1951
|
+
// does with its unprotected \`sqrt()\`) would leave 63 threads never
|
|
1952
|
+
// reaching a barrier the one remaining thread still needs.
|
|
1953
|
+
let scale = DD(tileScaleHi[0], tileScaleLo[0]);
|
|
1954
|
+
let ssq = DD(tileSsqHi[0], tileSsqLo[0]);
|
|
1955
|
+
let result = ddMulProtected(scale, ddSqrtProtected(ssq, i), i);
|
|
1956
|
+
if (i == 0u) {
|
|
1957
|
+
resultHi[0] = result.hi;
|
|
1958
|
+
resultLo[0] = result.lo;
|
|
1278
1959
|
}
|
|
1960
|
+
}
|
|
1961
|
+
`});var io,ao=O(()=>{io=`// sgemv_n: y = alpha * A * x + beta * y (A is m\xD7n row-major, no-transpose)
|
|
1962
|
+
//
|
|
1963
|
+
// One workgroup per output row, with a grid-stride outer loop so the shader
|
|
1964
|
+
// still covers all rows when m exceeds maxComputeWorkgroupsPerDimension.
|
|
1965
|
+
// Threads stride through A[row, :] and x with coalesced reads (consecutive
|
|
1966
|
+
// threads \u2192 consecutive addresses). Four independent accumulators let the GPU
|
|
1967
|
+
// pipeline memory requests across iterations (ILP=4), hiding the
|
|
1968
|
+
// global-memory latency.
|
|
1279
1969
|
|
|
1280
|
-
|
|
1281
|
-
|
|
1282
|
-
|
|
1970
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1971
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
1972
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
1283
1973
|
|
|
1284
|
-
|
|
1285
|
-
|
|
1974
|
+
struct Params {
|
|
1975
|
+
m: u32,
|
|
1976
|
+
n: u32,
|
|
1977
|
+
alpha: f32,
|
|
1978
|
+
beta: f32,
|
|
1979
|
+
incx: u32,
|
|
1980
|
+
incy: u32,
|
|
1981
|
+
lda: u32,
|
|
1982
|
+
}
|
|
1286
1983
|
|
|
1287
|
-
|
|
1288
|
-
if (sumHi != 0u) {
|
|
1289
|
-
leadPos = 32 + i32(31u - countLeadingZeros(sumHi));
|
|
1290
|
-
} else {
|
|
1291
|
-
leadPos = i32(31u - countLeadingZeros(sumLo));
|
|
1292
|
-
}
|
|
1293
|
-
let tentativeExp = leadPos + commonExp2;
|
|
1294
|
-
var targetLSBScale = tentativeExp - 52;
|
|
1295
|
-
if (tentativeExp < -1022) { targetLSBScale = -1074; }
|
|
1296
|
-
let shiftAmt = targetLSBScale - commonExp2;
|
|
1984
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1297
1985
|
|
|
1298
|
-
|
|
1299
|
-
|
|
1300
|
-
let sh = shl(sumHi, sumLo, u32(-shiftAmt)); // exact \u2014 cancellation only, never loses bits
|
|
1301
|
-
keepHi = sh.hi; keepLo = sh.lo;
|
|
1302
|
-
} else {
|
|
1303
|
-
// Only reached without cancellation (same-sign add, or a tied-exponent
|
|
1304
|
-
// subtract with no shrinkage) \u2014 shiftAmt here is always exactly 3 or 4,
|
|
1305
|
-
// so the dropped bits are fully known from sumLo directly (no sticky
|
|
1306
|
-
// approximation needed, unlike the Q-alignment shift above).
|
|
1307
|
-
let n = u32(shiftAmt);
|
|
1308
|
-
let remainder = sumLo & ((1u << n) - 1u);
|
|
1309
|
-
let halfway = 1u << (n - 1u);
|
|
1310
|
-
let sh = shr_sticky(sumHi, sumLo, n);
|
|
1311
|
-
keepHi = sh.hi; keepLo = sh.lo;
|
|
1312
|
-
if (remainder > halfway || (remainder == halfway && (keepLo & 1u) != 0u)) {
|
|
1313
|
-
let inc = add64(keepHi, keepLo, 0u, 1u);
|
|
1314
|
-
keepHi = inc.hi; keepLo = inc.lo;
|
|
1315
|
-
}
|
|
1316
|
-
}
|
|
1986
|
+
const WGS: u32 = 64u;
|
|
1987
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
1317
1988
|
|
|
1318
|
-
|
|
1319
|
-
|
|
1320
|
-
|
|
1321
|
-
|
|
1322
|
-
|
|
1323
|
-
|
|
1989
|
+
@compute @workgroup_size(64)
|
|
1990
|
+
fn main(
|
|
1991
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1992
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1993
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
1994
|
+
) {
|
|
1995
|
+
// Grid-stride loop: each workgroup handles ceil(m / nwg.x) rows.
|
|
1996
|
+
for (var row = wgid.x; row < params.m; row += nwg.x) {
|
|
1997
|
+
let row_base = row * params.lda;
|
|
1998
|
+
var acc0: f32 = 0.0;
|
|
1999
|
+
var acc1: f32 = 0.0;
|
|
2000
|
+
var acc2: f32 = 0.0;
|
|
2001
|
+
var acc3: f32 = 0.0;
|
|
1324
2002
|
|
|
1325
|
-
|
|
1326
|
-
|
|
1327
|
-
|
|
1328
|
-
let
|
|
1329
|
-
|
|
1330
|
-
|
|
2003
|
+
// 4-unrolled loop: each iteration issues 4 independent loads for A and x.
|
|
2004
|
+
// The accumulators are independent so the GPU can overlap the memory
|
|
2005
|
+
// requests rather than serialising them behind a dependency chain.
|
|
2006
|
+
let n4_floor = (params.n / (4u * WGS)) * (4u * WGS);
|
|
2007
|
+
for (var j: u32 = lid.x; j < n4_floor; j += 4u * WGS) {
|
|
2008
|
+
acc0 += A[row_base + j ] * x[ j * params.incx];
|
|
2009
|
+
acc1 += A[row_base + j + WGS ] * x[(j + WGS) * params.incx];
|
|
2010
|
+
acc2 += A[row_base + j + 2u * WGS ] * x[(j + 2u * WGS) * params.incx];
|
|
2011
|
+
acc3 += A[row_base + j + 3u * WGS ] * x[(j + 3u * WGS) * params.incx];
|
|
2012
|
+
}
|
|
2013
|
+
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
2014
|
+
for (var j: u32 = n4_floor + lid.x; j < params.n; j += WGS) {
|
|
2015
|
+
acc0 += A[row_base + j] * x[j * params.incx];
|
|
1331
2016
|
}
|
|
1332
|
-
return Fields(resultSign, u32(rawExpFinal), keepHi & 0xfffffu, keepLo);
|
|
1333
|
-
}
|
|
1334
|
-
return Fields(resultSign, 0u, keepHi & 0xfffffu, keepLo); // subnormal result
|
|
1335
|
-
}
|
|
1336
|
-
|
|
1337
|
-
// Packed-in/Packed-out convenience wrapper around addFields \u2014 encodes once,
|
|
1338
|
-
// after the math, rather than addFields itself needing to know about Packed.
|
|
1339
|
-
fn computeSum(a: Fields, b: Fields) -> Packed {
|
|
1340
|
-
let f = addFields(a, b);
|
|
1341
|
-
return encode(f.sign, f.rawExp, f.mantissaHi, f.lo);
|
|
1342
|
-
}
|
|
1343
|
-
`});var it,at=O(()=>{it=`// Double-double arithmetic via Dekker's algorithm \u2014 an alternative to
|
|
1344
|
-
// f64add.wgsl's bit-exact IEEE-754 emulation. Doesn't touch that path.
|
|
1345
|
-
//
|
|
1346
|
-
// A double-double number is a pair (hi, lo) of f32 with hi+lo approximating
|
|
1347
|
-
// a higher-precision value, hi holding the leading bits and lo the rounding
|
|
1348
|
-
// error hi lost. ~48 bits of mantissa vs f32's 24, less than real f64's 52.
|
|
1349
|
-
//
|
|
1350
|
-
// No bindings, no entry point \u2014 a helper library, concatenated with a
|
|
1351
|
-
// consumer's own bindings/entry point by getPipeline (WGSL has no #include).
|
|
1352
|
-
// The DD struct lives here \u2014 abs.wgsl/add.wgsl/greater.wgsl/equal.wgsl all
|
|
1353
|
-
// use it but don't redefine it (WGSL errors on duplicate struct definitions
|
|
1354
|
-
// once concatenated), so any consumer using those must concatenate this
|
|
1355
|
-
// file too, first.
|
|
1356
2017
|
|
|
1357
|
-
|
|
1358
|
-
|
|
1359
|
-
|
|
1360
|
-
|
|
1361
|
-
|
|
2018
|
+
// Parallel reduction: 64 \u2192 32 \u2192 16 \u2192 8 \u2192 4 \u2192 2 \u2192 1
|
|
2019
|
+
scratch[lid.x] = acc0 + acc1 + acc2 + acc3;
|
|
2020
|
+
workgroupBarrier();
|
|
2021
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
2022
|
+
if lid.x < stride {
|
|
2023
|
+
scratch[lid.x] += scratch[lid.x + stride];
|
|
2024
|
+
}
|
|
2025
|
+
workgroupBarrier();
|
|
2026
|
+
}
|
|
1362
2027
|
|
|
1363
|
-
|
|
1364
|
-
|
|
1365
|
-
|
|
1366
|
-
|
|
1367
|
-
|
|
2028
|
+
if lid.x == 0u {
|
|
2029
|
+
let yi = row * params.incy;
|
|
2030
|
+
// BLAS beta==0 semantics: y is written, not accumulated \u2014 must not read y.
|
|
2031
|
+
let acc = params.alpha * scratch[0];
|
|
2032
|
+
y[yi] = select(acc, acc + params.beta * y[yi], params.beta != 0.0);
|
|
2033
|
+
}
|
|
2034
|
+
// All 64 threads must agree before the next row reuses scratch[].
|
|
2035
|
+
workgroupBarrier();
|
|
1368
2036
|
}
|
|
1369
|
-
return a;
|
|
1370
|
-
}
|
|
1371
|
-
`});var lt,ut=O(()=>{lt=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
1372
|
-
|
|
1373
|
-
// \u2500\u2500 A real compiler bug \u2014 read before touching anything below \u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500\u2500
|
|
1374
|
-
//
|
|
1375
|
-
// twoSum/fastTwoSum's error term \`e\` should be nonzero (that's the point \u2014
|
|
1376
|
-
// \`s\` is rounded). A front-end-optimizer bug zeros it anyway, on both NVIDIA
|
|
1377
|
-
// and Mesa (ANV + llvmpipe), via two different mechanisms needing two fixes:
|
|
1378
|
-
// bitcast-based subtraction (\`fsub\`/\`negf\`, fixes NVIDIA) and materializing
|
|
1379
|
-
// the sum through workgroup memory + workgroupBarrier() (fixes Mesa). Only
|
|
1380
|
-
// both together (ddAddProtected) is verified correct everywhere \u2014 the plain
|
|
1381
|
-
// twoSum/fastTwoSum/ddAdd below are reference-only, not safe to use.
|
|
1382
|
-
fn negf(x: f32) -> f32 {
|
|
1383
|
-
return bitcast<f32>(bitcast<u32>(x) ^ 0x80000000u);
|
|
1384
|
-
}
|
|
1385
|
-
fn fsub(a: f32, b: f32) -> f32 {
|
|
1386
|
-
return a + negf(b);
|
|
1387
|
-
}
|
|
1388
|
-
|
|
1389
|
-
// Knuth/M\xF8ller's TwoSum: s = fl(a+b), e = exact rounding error, a+b == s+e.
|
|
1390
|
-
// Works for any a, b. UNPROTECTED \u2014 see header above.
|
|
1391
|
-
fn twoSum(a: f32, b: f32) -> DD {
|
|
1392
|
-
let s = a + b;
|
|
1393
|
-
let v = s - a;
|
|
1394
|
-
let e = (a - (s - v)) + (b - v);
|
|
1395
|
-
return DD(s, e);
|
|
1396
2037
|
}
|
|
2038
|
+
`});var no,so=O(()=>{no=`// sgemv_t: y = alpha * A^T * x + beta * y (A is m\xD7n row-major, transposed)
|
|
2039
|
+
// each thread owns one column of A \u2192 one element of y (length n)
|
|
2040
|
+
// tiles over x (length m) using shared memory; four independent accumulators
|
|
2041
|
+
// let the GPU pipeline A reads across j within each tile (ILP=4)
|
|
1397
2042
|
|
|
1398
|
-
|
|
1399
|
-
|
|
1400
|
-
|
|
1401
|
-
let s = a + b;
|
|
1402
|
-
let e = b - (s - a);
|
|
1403
|
-
return DD(s, e);
|
|
1404
|
-
}
|
|
2043
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
2044
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
2045
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
1405
2046
|
|
|
1406
|
-
|
|
1407
|
-
|
|
1408
|
-
|
|
1409
|
-
|
|
1410
|
-
|
|
2047
|
+
struct Params {
|
|
2048
|
+
m: u32,
|
|
2049
|
+
n: u32,
|
|
2050
|
+
alpha: f32,
|
|
2051
|
+
beta: f32,
|
|
2052
|
+
incx: u32,
|
|
2053
|
+
incy: u32,
|
|
2054
|
+
lda: u32,
|
|
1411
2055
|
}
|
|
1412
2056
|
|
|
1413
|
-
|
|
1414
|
-
//
|
|
1415
|
-
// Bitcast subtraction + workgroup-barrier materialization, verified correct
|
|
1416
|
-
// on all three backends tested. Costs a real barrier: fine for O(1)-per-
|
|
1417
|
-
// thread or O(log n) reduction use, not a long per-element loop. A
|
|
1418
|
-
// workgroupBarrier() requires uniform control flow, so:
|
|
1419
|
-
// - \`threadSlot\` must be unique per concurrent caller (e.g. local_invocation_index).
|
|
1420
|
-
// - Every thread in the workgroup must call this the same number of times
|
|
1421
|
-
// \u2014 including ones whose result gets discarded. Compute unconditionally;
|
|
1422
|
-
// only the write-back should be conditional.
|
|
1423
|
-
var<workgroup> dekkerScratch: array<f32, 64>;
|
|
2057
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1424
2058
|
|
|
1425
|
-
|
|
1426
|
-
|
|
1427
|
-
workgroupBarrier();
|
|
1428
|
-
let s = dekkerScratch[threadSlot];
|
|
1429
|
-
let v = fsub(s, a);
|
|
1430
|
-
let e = fsub(a, fsub(s, v)) + fsub(b, v);
|
|
1431
|
-
return DD(s, e);
|
|
1432
|
-
}
|
|
2059
|
+
const WGS: u32 = 64u;
|
|
2060
|
+
var<workgroup> x_tile: array<f32, 64>;
|
|
1433
2061
|
|
|
1434
|
-
|
|
1435
|
-
|
|
1436
|
-
|
|
1437
|
-
|
|
1438
|
-
|
|
1439
|
-
|
|
1440
|
-
|
|
2062
|
+
@compute @workgroup_size(64)
|
|
2063
|
+
fn main(
|
|
2064
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
2065
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
2066
|
+
) {
|
|
2067
|
+
// each thread owns column col of A \u2192 output y[col]
|
|
2068
|
+
let col = gid.x;
|
|
2069
|
+
// tile over x (length m, the rows of A)
|
|
2070
|
+
let m_floor = (params.m / WGS) * WGS;
|
|
2071
|
+
var acc0: f32 = 0.0;
|
|
2072
|
+
var acc1: f32 = 0.0;
|
|
2073
|
+
var acc2: f32 = 0.0;
|
|
2074
|
+
var acc3: f32 = 0.0;
|
|
1441
2075
|
|
|
1442
|
-
|
|
1443
|
-
|
|
1444
|
-
|
|
1445
|
-
|
|
1446
|
-
return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
|
|
1447
|
-
}
|
|
1448
|
-
`});var mt,ft=O(()=>{mt=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
2076
|
+
for (var base = 0u; base < m_floor; base += WGS) {
|
|
2077
|
+
// cooperative load: all 64 threads fill x_tile with x[base..base+WGS]
|
|
2078
|
+
x_tile[lid.x] = x[(base + lid.x) * params.incx];
|
|
2079
|
+
workgroupBarrier();
|
|
1449
2080
|
|
|
1450
|
-
//
|
|
1451
|
-
//
|
|
1452
|
-
|
|
1453
|
-
|
|
1454
|
-
|
|
1455
|
-
|
|
1456
|
-
|
|
2081
|
+
// 4-unrolled inner loop: 4 independent A reads let the GPU pipeline
|
|
2082
|
+
// global-memory requests within each tile. WGS=64 divides by 4 exactly.
|
|
2083
|
+
if (col < params.n) {
|
|
2084
|
+
for (var j = 0u; j < WGS; j += 4u) {
|
|
2085
|
+
acc0 += A[(base + j ) * params.lda + col] * x_tile[j ];
|
|
2086
|
+
acc1 += A[(base + j + 1) * params.lda + col] * x_tile[j + 1];
|
|
2087
|
+
acc2 += A[(base + j + 2) * params.lda + col] * x_tile[j + 2];
|
|
2088
|
+
acc3 += A[(base + j + 3) * params.lda + col] * x_tile[j + 3];
|
|
2089
|
+
}
|
|
2090
|
+
}
|
|
2091
|
+
workgroupBarrier();
|
|
1457
2092
|
}
|
|
1458
|
-
return a.lo > b.lo;
|
|
1459
|
-
}
|
|
1460
|
-
`});var dt,ct=O(()=>{dt=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
1461
2093
|
|
|
1462
|
-
|
|
1463
|
-
//
|
|
1464
|
-
|
|
1465
|
-
|
|
2094
|
+
if (col < params.n) {
|
|
2095
|
+
// remainder: m not divisible by WGS \u2014 short loop, single accumulator fine
|
|
2096
|
+
for (var k = m_floor; k < params.m; k++) {
|
|
2097
|
+
acc0 += A[k * params.lda + col] * x[k * params.incx];
|
|
2098
|
+
}
|
|
2099
|
+
let yi = col * params.incy;
|
|
2100
|
+
// BLAS beta==0 semantics: y is written, not accumulated \u2014 must not read y.
|
|
2101
|
+
let acc = params.alpha * (acc0 + acc1 + acc2 + acc3);
|
|
2102
|
+
y[yi] = select(acc, acc + params.beta * y[yi], params.beta != 0.0);
|
|
2103
|
+
}
|
|
1466
2104
|
}
|
|
1467
|
-
`});var
|
|
1468
|
-
//
|
|
1469
|
-
//
|
|
1470
|
-
//
|
|
2105
|
+
`});var uo,lo=O(()=>{uo=`// ssymv: y = alpha * A * x + beta * y
|
|
2106
|
+
// A is n\xD7n symmetric, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
2107
|
+
// The logical matrix is fully dense (symmetric), so each row's dot product
|
|
2108
|
+
// sums over all n columns; entries on the unstored side of the diagonal are
|
|
2109
|
+
// fetched from their mirror position (A[i,j] == A[j,i]).
|
|
2110
|
+
// One workgroup per row, grid-stride outer loop.
|
|
1471
2111
|
|
|
1472
|
-
@group(0) @binding(0) var<storage, read>
|
|
1473
|
-
@group(0) @binding(1) var<storage, read>
|
|
1474
|
-
@group(0) @binding(2) var<storage, read_write>
|
|
1475
|
-
@group(0) @binding(3) var<storage, read_write> partialsLo: array<f32>;
|
|
1476
|
-
@group(0) @binding(4) var<uniform> params: Params;
|
|
2112
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
2113
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
2114
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
1477
2115
|
|
|
1478
2116
|
struct Params {
|
|
1479
2117
|
n: u32,
|
|
1480
|
-
|
|
2118
|
+
alpha: f32,
|
|
2119
|
+
beta: f32,
|
|
2120
|
+
incx: u32,
|
|
2121
|
+
incy: u32,
|
|
2122
|
+
lda: u32,
|
|
2123
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
1481
2124
|
}
|
|
1482
2125
|
|
|
1483
|
-
|
|
2126
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1484
2127
|
|
|
1485
|
-
|
|
2128
|
+
const WGS: u32 = 64u;
|
|
2129
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
1486
2130
|
|
|
1487
2131
|
@compute @workgroup_size(64)
|
|
1488
|
-
fn
|
|
1489
|
-
@builtin(
|
|
1490
|
-
@builtin(local_invocation_id)
|
|
1491
|
-
@builtin(
|
|
1492
|
-
@builtin(num_workgroups) num_wg: vec3u,
|
|
2132
|
+
fn main(
|
|
2133
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
2134
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
2135
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
1493
2136
|
) {
|
|
1494
|
-
var
|
|
1495
|
-
|
|
1496
|
-
var acc2 = DD(0.0, 0.0);
|
|
1497
|
-
var acc3 = DD(0.0, 0.0);
|
|
1498
|
-
|
|
1499
|
-
let stride = num_wg.x * WGS;
|
|
1500
|
-
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
1501
|
-
|
|
1502
|
-
// Same trip count for every thread, but driven by a counter, not \`id\`
|
|
1503
|
-
// itself (ddAddProtected's barrier needs a provably-uniform loop bound).
|
|
1504
|
-
let mainIters = n4_floor / (4u * stride);
|
|
1505
|
-
for (var iter = 0u; iter < mainIters; iter++) {
|
|
1506
|
-
let id = gid.x + iter * 4u * stride;
|
|
1507
|
-
let i0 = id * params.x_inc;
|
|
1508
|
-
let i1 = (id + stride) * params.x_inc;
|
|
1509
|
-
let i2 = (id + 2u * stride) * params.x_inc;
|
|
1510
|
-
let i3 = (id + 3u * stride) * params.x_inc;
|
|
1511
|
-
acc0 = ddAddProtected(acc0, ddAbs(DD(xHi[i0], xLo[i0])), lid.x);
|
|
1512
|
-
acc1 = ddAddProtected(acc1, ddAbs(DD(xHi[i1], xLo[i1])), lid.x);
|
|
1513
|
-
acc2 = ddAddProtected(acc2, ddAbs(DD(xHi[i2], xLo[i2])), lid.x);
|
|
1514
|
-
acc3 = ddAddProtected(acc3, ddAbs(DD(xHi[i3], xLo[i3])), lid.x);
|
|
1515
|
-
}
|
|
1516
|
-
|
|
1517
|
-
// Tail is ragged (0-3 extra per thread) \u2014 pad to this workgroup's worst case.
|
|
1518
|
-
let wgBaseGid = wgid.x * WGS;
|
|
1519
|
-
var tailIters = 0u;
|
|
1520
|
-
if (n4_floor + wgBaseGid < params.n) {
|
|
1521
|
-
tailIters = (params.n - 1u - n4_floor - wgBaseGid) / stride + 1u;
|
|
1522
|
-
}
|
|
1523
|
-
for (var iter = 0u; iter < tailIters; iter++) {
|
|
1524
|
-
let id = n4_floor + gid.x + iter * stride;
|
|
1525
|
-
let valid = id < params.n;
|
|
1526
|
-
let i = select(0u, id * params.x_inc, valid);
|
|
1527
|
-
let loaded = ddAbs(DD(xHi[i], xLo[i])); // select() has no DD overload
|
|
1528
|
-
let contribution = DD(select(0.0, loaded.hi, valid), select(0.0, loaded.lo, valid));
|
|
1529
|
-
acc0 = ddAddProtected(acc0, contribution, lid.x);
|
|
1530
|
-
}
|
|
2137
|
+
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
2138
|
+
var acc = 0.0f;
|
|
1531
2139
|
|
|
1532
|
-
|
|
1533
|
-
|
|
1534
|
-
|
|
1535
|
-
|
|
2140
|
+
// y[i] = \u03A3_j A[i,j] * x[j]
|
|
2141
|
+
for (var j = lid.x; j < params.n; j += WGS) {
|
|
2142
|
+
var aVal: f32;
|
|
2143
|
+
if params.uplo == 0u {
|
|
2144
|
+
// Lower: A[i,j] stored at A[i*lda+j] for j \u2264 i, mirrored from A[j*lda+i] otherwise
|
|
2145
|
+
if j <= i {
|
|
2146
|
+
aVal = A[i * params.lda + j];
|
|
2147
|
+
} else {
|
|
2148
|
+
aVal = A[j * params.lda + i];
|
|
2149
|
+
}
|
|
2150
|
+
} else {
|
|
2151
|
+
// Upper: A[i,j] stored at A[i*lda+j] for j \u2265 i, mirrored from A[j*lda+i] otherwise
|
|
2152
|
+
if j >= i {
|
|
2153
|
+
aVal = A[i * params.lda + j];
|
|
2154
|
+
} else {
|
|
2155
|
+
aVal = A[j * params.lda + i];
|
|
2156
|
+
}
|
|
2157
|
+
}
|
|
2158
|
+
acc += aVal * x[j * params.incx];
|
|
2159
|
+
}
|
|
1536
2160
|
|
|
1537
|
-
|
|
1538
|
-
|
|
1539
|
-
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
1540
|
-
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
1541
|
-
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
1542
|
-
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
1543
|
-
if (lid.x < s) { tile[lid.x] = combined; }
|
|
2161
|
+
// Parallel reduction: 64 \u2192 1
|
|
2162
|
+
scratch[lid.x] = acc;
|
|
1544
2163
|
workgroupBarrier();
|
|
1545
|
-
|
|
2164
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
2165
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
2166
|
+
workgroupBarrier();
|
|
2167
|
+
}
|
|
1546
2168
|
|
|
1547
|
-
|
|
1548
|
-
|
|
1549
|
-
|
|
2169
|
+
if lid.x == 0u {
|
|
2170
|
+
// BLAS beta==0 semantics: y is written, not accumulated \u2014 must not read y.
|
|
2171
|
+
let acc = params.alpha * scratch[0];
|
|
2172
|
+
y[i * params.incy] = select(acc, acc + params.beta * y[i * params.incy], params.beta != 0.0);
|
|
2173
|
+
}
|
|
1550
2174
|
}
|
|
1551
2175
|
}
|
|
1552
|
-
`});var
|
|
1553
|
-
//
|
|
1554
|
-
//
|
|
1555
|
-
//
|
|
2176
|
+
`});var mo,fo=O(()=>{mo=`// strmv: y = op(A) * x
|
|
2177
|
+
// A is n\xD7n triangular, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
2178
|
+
// op(A) is A (trans=0) or A^T (trans=1).
|
|
2179
|
+
// diag=1 (unit) treats the diagonal as 1 without reading A's diagonal values.
|
|
2180
|
+
// One workgroup per row, grid-stride outer loop.
|
|
1556
2181
|
|
|
1557
|
-
@group(0) @binding(0) var<storage, read>
|
|
1558
|
-
@group(0) @binding(1) var<storage, read>
|
|
1559
|
-
@group(0) @binding(2) var<storage, read_write>
|
|
1560
|
-
@group(0) @binding(3) var<storage, read_write> partialsValLo: array<f32>;
|
|
1561
|
-
@group(0) @binding(4) var<storage, read_write> partialsIdx: array<u32>;
|
|
1562
|
-
@group(0) @binding(5) var<uniform> params: Params;
|
|
2182
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
2183
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
2184
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
1563
2185
|
|
|
1564
2186
|
struct Params {
|
|
1565
2187
|
n: u32,
|
|
1566
|
-
|
|
2188
|
+
incx: u32,
|
|
2189
|
+
incy: u32,
|
|
2190
|
+
lda: u32,
|
|
2191
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
2192
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
2193
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
1567
2194
|
}
|
|
1568
2195
|
|
|
1569
|
-
|
|
2196
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1570
2197
|
|
|
1571
|
-
|
|
1572
|
-
var<workgroup>
|
|
2198
|
+
const WGS: u32 = 64u;
|
|
2199
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
1573
2200
|
|
|
1574
2201
|
@compute @workgroup_size(64)
|
|
1575
|
-
fn
|
|
1576
|
-
@builtin(
|
|
1577
|
-
@builtin(local_invocation_id)
|
|
1578
|
-
@builtin(
|
|
1579
|
-
@builtin(num_workgroups) num_wg: vec3u,
|
|
2202
|
+
fn main(
|
|
2203
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
2204
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
2205
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
1580
2206
|
) {
|
|
1581
|
-
|
|
1582
|
-
|
|
1583
|
-
var best_val0 = DD(-1.0, 0.0); var best_idx0: u32 = 0u;
|
|
1584
|
-
var best_val1 = DD(-1.0, 0.0); var best_idx1: u32 = 0u;
|
|
1585
|
-
var best_val2 = DD(-1.0, 0.0); var best_idx2: u32 = 0u;
|
|
1586
|
-
var best_val3 = DD(-1.0, 0.0); var best_idx3: u32 = 0u;
|
|
1587
|
-
|
|
1588
|
-
let stride = num_wg.x * WGS;
|
|
1589
|
-
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
1590
|
-
|
|
1591
|
-
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
1592
|
-
let i0 = id * params.x_inc;
|
|
1593
|
-
let i1 = (id + stride) * params.x_inc;
|
|
1594
|
-
let i2 = (id + 2u * stride) * params.x_inc;
|
|
1595
|
-
let i3 = (id + 3u * stride) * params.x_inc;
|
|
1596
|
-
let v0 = ddAbs(DD(xHi[i0], xLo[i0]));
|
|
1597
|
-
let v1 = ddAbs(DD(xHi[i1], xLo[i1]));
|
|
1598
|
-
let v2 = ddAbs(DD(xHi[i2], xLo[i2]));
|
|
1599
|
-
let v3 = ddAbs(DD(xHi[i3], xLo[i3]));
|
|
1600
|
-
if (ddGreater(v0, best_val0)) { best_val0 = v0; best_idx0 = id; }
|
|
1601
|
-
if (ddGreater(v1, best_val1)) { best_val1 = v1; best_idx1 = id + stride; }
|
|
1602
|
-
if (ddGreater(v2, best_val2)) { best_val2 = v2; best_idx2 = id + 2u * stride; }
|
|
1603
|
-
if (ddGreater(v3, best_val3)) { best_val3 = v3; best_idx3 = id + 3u * stride; }
|
|
1604
|
-
}
|
|
1605
|
-
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
1606
|
-
let i = id * params.x_inc;
|
|
1607
|
-
let v = ddAbs(DD(xHi[i], xLo[i]));
|
|
1608
|
-
if (ddGreater(v, best_val0)) { best_val0 = v; best_idx0 = id; }
|
|
1609
|
-
}
|
|
1610
|
-
|
|
1611
|
-
// merge 4 independent lanes; prefer lower index on tie (first occurrence wins)
|
|
1612
|
-
if (ddGreater(best_val1, best_val0) ||
|
|
1613
|
-
(ddEqual(best_val1, best_val0) && best_idx1 < best_idx0)) {
|
|
1614
|
-
best_val0 = best_val1; best_idx0 = best_idx1;
|
|
1615
|
-
}
|
|
1616
|
-
if (ddGreater(best_val2, best_val0) ||
|
|
1617
|
-
(ddEqual(best_val2, best_val0) && best_idx2 < best_idx0)) {
|
|
1618
|
-
best_val0 = best_val2; best_idx0 = best_idx2;
|
|
1619
|
-
}
|
|
1620
|
-
if (ddGreater(best_val3, best_val0) ||
|
|
1621
|
-
(ddEqual(best_val3, best_val0) && best_idx3 < best_idx0)) {
|
|
1622
|
-
best_val0 = best_val3; best_idx0 = best_idx3;
|
|
1623
|
-
}
|
|
1624
|
-
|
|
1625
|
-
tile_val[lid.x] = best_val0;
|
|
1626
|
-
tile_idx[lid.x] = best_idx0;
|
|
1627
|
-
workgroupBarrier();
|
|
2207
|
+
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
2208
|
+
var acc = 0.0f;
|
|
1628
2209
|
|
|
1629
|
-
|
|
1630
|
-
|
|
1631
|
-
|
|
1632
|
-
|
|
1633
|
-
|
|
1634
|
-
|
|
1635
|
-
|
|
1636
|
-
|
|
2210
|
+
if params.trans == 0u {
|
|
2211
|
+
// No-transpose: y[i] = \u03A3_j A[i,j] * x[j]
|
|
2212
|
+
if params.uplo == 0u {
|
|
2213
|
+
// Lower: A[i,j] stored at A[i*lda+j] for j \u2264 i
|
|
2214
|
+
for (var j = lid.x; j <= i; j += WGS) {
|
|
2215
|
+
var aVal: f32;
|
|
2216
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
2217
|
+
if params.diag == 1u && j == i {
|
|
2218
|
+
aVal = 1.0;
|
|
2219
|
+
} else if ( j <= i ) {
|
|
2220
|
+
aVal = A[i * params.lda + j];
|
|
2221
|
+
}
|
|
2222
|
+
acc += aVal * x[j * params.incx];
|
|
2223
|
+
}
|
|
2224
|
+
} else {
|
|
2225
|
+
// Upper: A[i,j] stored at A[i*lda+j] for j \u2265 i
|
|
2226
|
+
for (var j = i + lid.x; j < params.n; j += WGS) {
|
|
2227
|
+
var aVal: f32;
|
|
2228
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
2229
|
+
if params.diag == 1u && j == i {
|
|
2230
|
+
aVal = 1.0;
|
|
2231
|
+
} else if ( j >= i ) {
|
|
2232
|
+
aVal = A[i * params.lda + j];
|
|
2233
|
+
}
|
|
2234
|
+
acc += aVal * x[j * params.incx];
|
|
2235
|
+
}
|
|
2236
|
+
}
|
|
2237
|
+
} else {
|
|
2238
|
+
// Transpose: y[i] = \u03A3_j A[j,i] * x[j]
|
|
2239
|
+
if params.uplo == 0u {
|
|
2240
|
+
// Lower: A[j,i] stored at A[j*lda+i] for j \u2265 i
|
|
2241
|
+
for (var j = i + lid.x; j < params.n; j += WGS) {
|
|
2242
|
+
var aVal: f32;
|
|
2243
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
2244
|
+
if params.diag == 1u && j == i {
|
|
2245
|
+
aVal = 1.0;
|
|
2246
|
+
} else if ( j >= i ) {
|
|
2247
|
+
aVal = A[j * params.lda + i];
|
|
2248
|
+
}
|
|
2249
|
+
acc += aVal * x[j * params.incx];
|
|
2250
|
+
}
|
|
2251
|
+
} else {
|
|
2252
|
+
// Upper: A[j,i] stored at A[j*lda+i] for j \u2264 i
|
|
2253
|
+
for (var j = lid.x; j <= i; j += WGS) {
|
|
2254
|
+
var aVal: f32;
|
|
2255
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
2256
|
+
if params.diag == 1u && j == i {
|
|
2257
|
+
aVal = 1.0;
|
|
2258
|
+
} else if ( j <= i ) {
|
|
2259
|
+
aVal = A[j * params.lda + i];
|
|
2260
|
+
}
|
|
2261
|
+
acc += aVal * x[j * params.incx];
|
|
2262
|
+
}
|
|
1637
2263
|
}
|
|
1638
2264
|
}
|
|
2265
|
+
|
|
2266
|
+
// Parallel reduction: 64 \u2192 1
|
|
2267
|
+
scratch[lid.x] = acc;
|
|
1639
2268
|
workgroupBarrier();
|
|
1640
|
-
|
|
2269
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
2270
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
2271
|
+
workgroupBarrier();
|
|
2272
|
+
}
|
|
1641
2273
|
|
|
1642
|
-
|
|
1643
|
-
|
|
1644
|
-
|
|
1645
|
-
partialsIdx[wgid.x] = tile_idx[0];
|
|
2274
|
+
if lid.x == 0u {
|
|
2275
|
+
y[ i * params.incy ] = scratch[0];
|
|
2276
|
+
}
|
|
1646
2277
|
}
|
|
1647
2278
|
}
|
|
1648
|
-
`});var
|
|
2279
|
+
`});var Me,co=O(()=>{Me=`// strsv_invert_block: computes ONE column (workgroup_id.x) of ONE block's
|
|
1649
2280
|
// (workgroup_id.y) explicit inverse, via the same one-row-at-a-time
|
|
1650
2281
|
// substitution as strsv_block.wgsl, but solving against a unit basis vector
|
|
1651
2282
|
// e_col instead of the real right-hand side, and writing to a dense
|
|
@@ -1754,7 +2385,7 @@ fn strsv_invert_block_main(
|
|
|
1754
2385
|
workgroupBarrier();
|
|
1755
2386
|
}
|
|
1756
2387
|
}
|
|
1757
|
-
`});var
|
|
2388
|
+
`});var go,po=O(()=>{go=`// strsv_apply_inverse: given a precomputed block inverse (from
|
|
1758
2389
|
// strsv_invert_block.wgsl), computes this block's solution as a dense
|
|
1759
2390
|
// matrix-vector multiply against the block's current remainder in x \u2014
|
|
1760
2391
|
// replacing what the old strsv_block.wgsl did via a genuinely sequential,
|
|
@@ -1801,7 +2432,7 @@ fn strsv_apply_inverse_main(@builtin(local_invocation_id) lid: vec3u) {
|
|
|
1801
2432
|
}
|
|
1802
2433
|
x[(params.blockStart + lid.x) * params.incx] = acc;
|
|
1803
2434
|
}
|
|
1804
|
-
`});var
|
|
2435
|
+
`});var ho,wo=O(()=>{ho=`// strsv_update: subtracts a solved block's contribution from every
|
|
1805
2436
|
// remaining row in parallel (one workgroup per row, like strmv.wgsl) \u2014
|
|
1806
2437
|
// this is what turns strsv's O(n) sequential stages into O(n/blockSize).
|
|
1807
2438
|
// No diag/masking needed: this region never touches the diagonal.
|
|
@@ -1876,13 +2507,191 @@ fn strsv_update_main(
|
|
|
1876
2507
|
workgroupBarrier();
|
|
1877
2508
|
}
|
|
1878
2509
|
}
|
|
1879
|
-
`});var
|
|
2510
|
+
`});var yo,bo=O(()=>{yo=`// sger: A := alpha * x * y^T + A (rank-1 update, A is m\xD7n general/dense)
|
|
2511
|
+
|
|
2512
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
2513
|
+
@group(0) @binding(1) var<storage, read> y: array<f32>;
|
|
2514
|
+
@group(0) @binding(2) var<storage, read_write> A: array<f32>;
|
|
2515
|
+
|
|
2516
|
+
struct Params {
|
|
2517
|
+
m: u32,
|
|
2518
|
+
n: u32,
|
|
2519
|
+
alpha: f32,
|
|
2520
|
+
incx: u32,
|
|
2521
|
+
incy: u32,
|
|
2522
|
+
lda: u32,
|
|
2523
|
+
}
|
|
2524
|
+
|
|
2525
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
2526
|
+
|
|
2527
|
+
const WGS: u32 = 64u;
|
|
2528
|
+
|
|
2529
|
+
@compute @workgroup_size(64)
|
|
2530
|
+
fn main(
|
|
2531
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
2532
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
2533
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
2534
|
+
) {
|
|
2535
|
+
for (var row = wgid.x; row < params.m; row += nwg.x) {
|
|
2536
|
+
let xi = params.alpha * x[row * params.incx];
|
|
2537
|
+
let row_base = row * params.lda;
|
|
2538
|
+
|
|
2539
|
+
// 4-unrolled loop: each iteration issues 4 independent A/y accesses.
|
|
2540
|
+
let n4_floor = (params.n / (4u * WGS)) * (4u * WGS);
|
|
2541
|
+
for (var col: u32 = lid.x; col < n4_floor; col += 4u * WGS) {
|
|
2542
|
+
let idx0 = row_base + col;
|
|
2543
|
+
let idx1 = row_base + col + WGS;
|
|
2544
|
+
let idx2 = row_base + col + 2u * WGS;
|
|
2545
|
+
let idx3 = row_base + col + 3u * WGS;
|
|
2546
|
+
A[idx0] = xi * y[ col * params.incy] + A[idx0];
|
|
2547
|
+
A[idx1] = xi * y[(col + WGS) * params.incy] + A[idx1];
|
|
2548
|
+
A[idx2] = xi * y[(col + 2u * WGS) * params.incy] + A[idx2];
|
|
2549
|
+
A[idx3] = xi * y[(col + 3u * WGS) * params.incy] + A[idx3];
|
|
2550
|
+
}
|
|
2551
|
+
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
2552
|
+
for (var col: u32 = n4_floor + lid.x; col < params.n; col += WGS) {
|
|
2553
|
+
let idx = row_base + col;
|
|
2554
|
+
A[idx] = xi * y[col * params.incy] + A[idx];
|
|
2555
|
+
}
|
|
2556
|
+
}
|
|
2557
|
+
}
|
|
2558
|
+
`});var vo,xo=O(()=>{vo=`// ssyr: A := alpha * x * x^T + A (symmetric rank-1 update)
|
|
2559
|
+
// A is n\xD7n symmetric; only the triangle specified by uplo is referenced/updated,
|
|
2560
|
+
// the other triangle is implied by symmetry (not touched).
|
|
2561
|
+
|
|
2562
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
2563
|
+
@group(0) @binding(1) var<storage, read_write> A: array<f32>;
|
|
2564
|
+
|
|
2565
|
+
struct Params {
|
|
2566
|
+
n: u32,
|
|
2567
|
+
alpha: f32,
|
|
2568
|
+
incx: u32,
|
|
2569
|
+
lda: u32,
|
|
2570
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
2571
|
+
}
|
|
2572
|
+
|
|
2573
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
2574
|
+
|
|
2575
|
+
const WGS: u32 = 64u;
|
|
2576
|
+
|
|
2577
|
+
@compute @workgroup_size(64)
|
|
2578
|
+
fn main(
|
|
2579
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
2580
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
2581
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
2582
|
+
) {
|
|
2583
|
+
for (var row = wgid.x; row < params.n; row += nwg.x) {
|
|
2584
|
+
let xi = params.alpha * x[row * params.incx];
|
|
2585
|
+
let row_base = row * params.lda;
|
|
2586
|
+
|
|
2587
|
+
// Stored-triangle column range for this row: lower [0,row], upper [row,n).
|
|
2588
|
+
var colStart: u32;
|
|
2589
|
+
var colEnd: u32;
|
|
2590
|
+
if params.uplo == 1u {
|
|
2591
|
+
colStart = row;
|
|
2592
|
+
colEnd = params.n;
|
|
2593
|
+
} else {
|
|
2594
|
+
colStart = 0u;
|
|
2595
|
+
colEnd = row + 1u;
|
|
2596
|
+
}
|
|
2597
|
+
|
|
2598
|
+
// 4-unrolled loop over the stored range.
|
|
2599
|
+
let rangeLen = colEnd - colStart;
|
|
2600
|
+
let n4_floor = colStart + (rangeLen / (4u * WGS)) * (4u * WGS);
|
|
2601
|
+
for (var col: u32 = colStart + lid.x; col < n4_floor; col += 4u * WGS) {
|
|
2602
|
+
let idx0 = row_base + col;
|
|
2603
|
+
let idx1 = row_base + col + WGS;
|
|
2604
|
+
let idx2 = row_base + col + 2u * WGS;
|
|
2605
|
+
let idx3 = row_base + col + 3u * WGS;
|
|
2606
|
+
A[idx0] = xi * x[ col * params.incx] + A[idx0];
|
|
2607
|
+
A[idx1] = xi * x[(col + WGS) * params.incx] + A[idx1];
|
|
2608
|
+
A[idx2] = xi * x[(col + 2u * WGS) * params.incx] + A[idx2];
|
|
2609
|
+
A[idx3] = xi * x[(col + 3u * WGS) * params.incx] + A[idx3];
|
|
2610
|
+
}
|
|
2611
|
+
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
2612
|
+
for (var col: u32 = n4_floor + lid.x; col < colEnd; col += WGS) {
|
|
2613
|
+
let idx = row_base + col;
|
|
2614
|
+
A[idx] = xi * x[col * params.incx] + A[idx];
|
|
2615
|
+
}
|
|
2616
|
+
}
|
|
2617
|
+
}
|
|
2618
|
+
`});var Bo,_o=O(()=>{Bo=`// ssyr2: A := alpha * x * y^T + alpha * y * x^T + A (symmetric rank-2 update)
|
|
2619
|
+
// A is n\xD7n symmetric; only the triangle specified by uplo is referenced/updated,
|
|
2620
|
+
// the other triangle is implied by symmetry (not touched).
|
|
2621
|
+
|
|
2622
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
2623
|
+
@group(0) @binding(1) var<storage, read> y: array<f32>;
|
|
2624
|
+
@group(0) @binding(2) var<storage, read_write> A: array<f32>;
|
|
2625
|
+
|
|
2626
|
+
struct Params {
|
|
2627
|
+
n: u32,
|
|
2628
|
+
alpha: f32,
|
|
2629
|
+
incx: u32,
|
|
2630
|
+
incy: u32,
|
|
2631
|
+
lda: u32,
|
|
2632
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
2633
|
+
}
|
|
2634
|
+
|
|
2635
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
2636
|
+
|
|
2637
|
+
const WGS: u32 = 64u;
|
|
2638
|
+
|
|
2639
|
+
@compute @workgroup_size(64)
|
|
2640
|
+
fn main(
|
|
2641
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
2642
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
2643
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
2644
|
+
) {
|
|
2645
|
+
for (var row = wgid.x; row < params.n; row += nwg.x) {
|
|
2646
|
+
let xi = params.alpha * x[row * params.incx];
|
|
2647
|
+
let yi = params.alpha * y[row * params.incy];
|
|
2648
|
+
let row_base = row * params.lda;
|
|
2649
|
+
|
|
2650
|
+
// Stored-triangle column range for this row: lower [0,row], upper [row,n).
|
|
2651
|
+
var colStart: u32;
|
|
2652
|
+
var colEnd: u32;
|
|
2653
|
+
if params.uplo == 1u {
|
|
2654
|
+
colStart = row;
|
|
2655
|
+
colEnd = params.n;
|
|
2656
|
+
} else {
|
|
2657
|
+
colStart = 0u;
|
|
2658
|
+
colEnd = row + 1u;
|
|
2659
|
+
}
|
|
2660
|
+
|
|
2661
|
+
// 4-unrolled loop over the stored range.
|
|
2662
|
+
let rangeLen = colEnd - colStart;
|
|
2663
|
+
let n4_floor = colStart + (rangeLen / (4u * WGS)) * (4u * WGS);
|
|
2664
|
+
for (var col: u32 = colStart + lid.x; col < n4_floor; col += 4u * WGS) {
|
|
2665
|
+
let idx0 = row_base + col;
|
|
2666
|
+
let idx1 = row_base + col + WGS;
|
|
2667
|
+
let idx2 = row_base + col + 2u * WGS;
|
|
2668
|
+
let idx3 = row_base + col + 3u * WGS;
|
|
2669
|
+
A[idx0] = xi * y[ col * params.incy] + yi * x[ col * params.incx] + A[idx0];
|
|
2670
|
+
A[idx1] = xi * y[(col + WGS) * params.incy] + yi * x[(col + WGS) * params.incx] + A[idx1];
|
|
2671
|
+
A[idx2] = xi * y[(col + 2u * WGS) * params.incy] + yi * x[(col + 2u * WGS) * params.incx] + A[idx2];
|
|
2672
|
+
A[idx3] = xi * y[(col + 3u * WGS) * params.incy] + yi * x[(col + 3u * WGS) * params.incx] + A[idx3];
|
|
2673
|
+
}
|
|
2674
|
+
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
2675
|
+
for (var col: u32 = n4_floor + lid.x; col < colEnd; col += WGS) {
|
|
2676
|
+
let idx = row_base + col;
|
|
2677
|
+
A[idx] = xi * y[col * params.incy] + yi * x[col * params.incx] + A[idx];
|
|
2678
|
+
}
|
|
2679
|
+
}
|
|
2680
|
+
}
|
|
2681
|
+
`});var ue,Ao=O(()=>{ue=`// sgemm_small: C = alpha * op(A) * op(B) + beta * C \u2014 small-tile half of
|
|
1880
2682
|
// the two-tier autotuned dispatch (see sgemm.mjs and sgemm_large.wgsl).
|
|
1881
2683
|
// BM=BN=32, BK=8, TM=TN=2 \u2014 wins over the large tile below a 6x6=36
|
|
1882
2684
|
// workgroup grid of 64-tiles, where the large tile doesn't have enough
|
|
1883
2685
|
// workgroups to fill the GPU. Same structure as sgemm_large.wgsl (2D
|
|
1884
2686
|
// register-blocked, shared-memory-tiled), just smaller.
|
|
1885
2687
|
//
|
|
2688
|
+
// A and B are bound twice \u2014 scalar array<f32> and array<vec4<f32>> views of
|
|
2689
|
+
// the same GPUBuffer (see vec4ViewBinding) \u2014 so each tile load can issue
|
|
2690
|
+
// 16-byte vector reads along op(A)/op(B)'s contiguous dimension when the
|
|
2691
|
+
// stride allows it (stride % 4 == 0 keeps every row base 16-byte aligned).
|
|
2692
|
+
// NUM_THREADS (256) exceeds some small-tile load shapes, so the vectorized
|
|
2693
|
+
// paths whose lane count doesn't tile exactly guard their As/Bs stores.
|
|
2694
|
+
//
|
|
1886
2695
|
// col mapped to gid.x for coalesced B/C access (row-major: col contiguous).
|
|
1887
2696
|
|
|
1888
2697
|
const BM: u32 = 32u;
|
|
@@ -1896,9 +2705,11 @@ const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 256
|
|
|
1896
2705
|
const STRIDE_A: u32 = NUM_THREADS / BK;
|
|
1897
2706
|
const STRIDE_B: u32 = NUM_THREADS / BN;
|
|
1898
2707
|
|
|
1899
|
-
@group(0) @binding(0) var<storage, read> A:
|
|
1900
|
-
@group(0) @binding(1) var<storage, read>
|
|
1901
|
-
@group(0) @binding(2) var<storage,
|
|
2708
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
2709
|
+
@group(0) @binding(1) var<storage, read> A4: array<vec4<f32>>;
|
|
2710
|
+
@group(0) @binding(2) var<storage, read> B: array<f32>;
|
|
2711
|
+
@group(0) @binding(3) var<storage, read> B4: array<vec4<f32>>;
|
|
2712
|
+
@group(0) @binding(4) var<storage, read_write> C: array<f32>;
|
|
1902
2713
|
|
|
1903
2714
|
struct Params {
|
|
1904
2715
|
m: u32,
|
|
@@ -1911,9 +2722,11 @@ struct Params {
|
|
|
1911
2722
|
ldc: u32,
|
|
1912
2723
|
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
1913
2724
|
transB: u32,
|
|
2725
|
+
useVecA: u32, // 1 = A's vec4 view covers every in-bounds element (see vec4Usable)
|
|
2726
|
+
useVecB: u32, // 1 = B's vec4 view covers every in-bounds element
|
|
1914
2727
|
}
|
|
1915
2728
|
|
|
1916
|
-
@group(0) @binding(
|
|
2729
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
1917
2730
|
|
|
1918
2731
|
var<workgroup> As: array<f32, BM * BK>;
|
|
1919
2732
|
var<workgroup> Bs: array<f32, BK * BN>;
|
|
@@ -1943,17 +2756,103 @@ fn main(
|
|
|
1943
2756
|
|
|
1944
2757
|
let numTiles = (params.k + BK - 1u) / BK;
|
|
1945
2758
|
for (var t = 0u; t < numTiles; t++) {
|
|
1946
|
-
|
|
1947
|
-
|
|
1948
|
-
|
|
1949
|
-
|
|
1950
|
-
|
|
2759
|
+
// \u2500\u2500 Load the BM\xD7BK A tile into As (vectorized along op(A)'s fast dim
|
|
2760
|
+
// when lda allows; every branch here is dispatch-uniform) \u2500\u2500
|
|
2761
|
+
if (params.useVecA == 1u && params.transA == 0u) {
|
|
2762
|
+
// No-transpose: columns contiguous, one vec4 per thread. NUM_THREADS
|
|
2763
|
+
// spans BM/4\xD7(BK/4) several times over \u2014 guard the store.
|
|
2764
|
+
let r4 = tid / (BK / 4u);
|
|
2765
|
+
let c4 = tid % (BK / 4u);
|
|
2766
|
+
if (r4 < BM) {
|
|
2767
|
+
let gRow = blockRow + r4;
|
|
2768
|
+
let gCol = t * BK + c4 * 4u;
|
|
2769
|
+
var v = A4[(gRow * params.lda + gCol) / 4u];
|
|
2770
|
+
let rowOK = gRow < params.m;
|
|
2771
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.k);
|
|
2772
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.k);
|
|
2773
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.k);
|
|
2774
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.k);
|
|
2775
|
+
As[r4 * BK + c4 * 4u] = v.x;
|
|
2776
|
+
As[r4 * BK + c4 * 4u + 1u] = v.y;
|
|
2777
|
+
As[r4 * BK + c4 * 4u + 2u] = v.z;
|
|
2778
|
+
As[r4 * BK + c4 * 4u + 3u] = v.w;
|
|
2779
|
+
}
|
|
2780
|
+
} else if (params.useVecA == 1u && params.transA != 0u) {
|
|
2781
|
+
// Transpose: rows contiguous within a column. NUM_THREADS over-spans
|
|
2782
|
+
// the BK-column tile \u2014 guard the store.
|
|
2783
|
+
let r4 = tid % (BM / 4u);
|
|
2784
|
+
let c = tid / (BM / 4u);
|
|
2785
|
+
if (c < BK) {
|
|
2786
|
+
let gRow = blockRow + r4 * 4u;
|
|
2787
|
+
let gCol = t * BK + c;
|
|
2788
|
+
var v = A4[(gCol * params.lda + gRow) / 4u];
|
|
2789
|
+
let colOK = gCol < params.k;
|
|
2790
|
+
v.x = select(0.0, v.x, colOK && gRow < params.m);
|
|
2791
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.m);
|
|
2792
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.m);
|
|
2793
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.m);
|
|
2794
|
+
As[(r4 * 4u) * BK + c] = v.x;
|
|
2795
|
+
As[(r4 * 4u + 1u) * BK + c] = v.y;
|
|
2796
|
+
As[(r4 * 4u + 2u) * BK + c] = v.z;
|
|
2797
|
+
As[(r4 * 4u + 3u) * BK + c] = v.w;
|
|
2798
|
+
}
|
|
2799
|
+
} else {
|
|
2800
|
+
// Scalar fallback: odd stride or unhandled orientation.
|
|
2801
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
2802
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
2803
|
+
let gColA = t * BK + innerColA;
|
|
2804
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
2805
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
2806
|
+
}
|
|
1951
2807
|
}
|
|
1952
|
-
|
|
1953
|
-
|
|
1954
|
-
|
|
1955
|
-
|
|
1956
|
-
|
|
2808
|
+
|
|
2809
|
+
// \u2500\u2500 Load the BK\xD7BN B tile into Bs \u2500\u2500
|
|
2810
|
+
if (params.useVecB == 1u && params.transB == 0u) {
|
|
2811
|
+
// No-transpose: columns contiguous, one vec4 per thread. NUM_THREADS
|
|
2812
|
+
// over-spans the BK-row tile \u2014 guard the store.
|
|
2813
|
+
let r = tid / (BN / 4u);
|
|
2814
|
+
let c4 = tid % (BN / 4u);
|
|
2815
|
+
if (r < BK) {
|
|
2816
|
+
let gRow = t * BK + r;
|
|
2817
|
+
let gCol = blockCol + c4 * 4u;
|
|
2818
|
+
var v = B4[(gRow * params.ldb + gCol) / 4u];
|
|
2819
|
+
let rowOK = gRow < params.k;
|
|
2820
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.n);
|
|
2821
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.n);
|
|
2822
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.n);
|
|
2823
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.n);
|
|
2824
|
+
Bs[r * BN + c4 * 4u] = v.x;
|
|
2825
|
+
Bs[r * BN + c4 * 4u + 1u] = v.y;
|
|
2826
|
+
Bs[r * BN + c4 * 4u + 2u] = v.z;
|
|
2827
|
+
Bs[r * BN + c4 * 4u + 3u] = v.w;
|
|
2828
|
+
}
|
|
2829
|
+
} else if (params.useVecB == 1u && params.transB != 0u) {
|
|
2830
|
+
// Transpose: rows contiguous within a column, one vec4 per thread \u2014
|
|
2831
|
+
// NUM_THREADS over-spans the 32-column tile, so guard the store.
|
|
2832
|
+
let r4 = tid % (BK / 4u);
|
|
2833
|
+
let c = tid / (BK / 4u);
|
|
2834
|
+
if (c < BN) {
|
|
2835
|
+
let gRow = t * BK + r4 * 4u;
|
|
2836
|
+
let gCol = blockCol + c;
|
|
2837
|
+
var v = B4[(gCol * params.ldb + gRow) / 4u];
|
|
2838
|
+
let colOK = gCol < params.n;
|
|
2839
|
+
v.x = select(0.0, v.x, colOK && gRow < params.k);
|
|
2840
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.k);
|
|
2841
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.k);
|
|
2842
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.k);
|
|
2843
|
+
Bs[(r4 * 4u) * BN + c] = v.x;
|
|
2844
|
+
Bs[(r4 * 4u + 1u) * BN + c] = v.y;
|
|
2845
|
+
Bs[(r4 * 4u + 2u) * BN + c] = v.z;
|
|
2846
|
+
Bs[(r4 * 4u + 3u) * BN + c] = v.w;
|
|
2847
|
+
}
|
|
2848
|
+
} else {
|
|
2849
|
+
// Scalar fallback.
|
|
2850
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
2851
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
2852
|
+
let gColB = blockCol + innerColB;
|
|
2853
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
2854
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
2855
|
+
}
|
|
1957
2856
|
}
|
|
1958
2857
|
|
|
1959
2858
|
workgroupBarrier();
|
|
@@ -1982,13 +2881,16 @@ fn main(
|
|
|
1982
2881
|
let col = blockCol + threadCol * TN + resIdxN;
|
|
1983
2882
|
if (col < params.n) {
|
|
1984
2883
|
let cIdx = row * params.ldc + col;
|
|
1985
|
-
|
|
2884
|
+
// BLAS beta==0 semantics: C is written, not accumulated \u2014 must not
|
|
2885
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
2886
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
2887
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
1986
2888
|
}
|
|
1987
2889
|
}
|
|
1988
2890
|
}
|
|
1989
2891
|
}
|
|
1990
2892
|
}
|
|
1991
|
-
`});var
|
|
2893
|
+
`});var fe,So=O(()=>{fe=`// sgemm_large: C = alpha * op(A) * op(B) + beta * C \u2014 large-tile half of
|
|
1992
2894
|
// the two-tier autotuned dispatch (see sgemm.mjs and sgemm_small.wgsl).
|
|
1993
2895
|
// BM=BN=64, BK=8, TM=8, TN=4 (128 threads/workgroup) \u2014 the kernel 9
|
|
1994
2896
|
// autotuning winner (temp/autotune_sweep.mjs, temp/gen_sweep_kernel.mjs,
|
|
@@ -1996,9 +2898,13 @@ fn main(
|
|
|
1996
2898
|
// single-tier baseline at n=512, +84% at n=1024. But BM=64 loses to BM=32
|
|
1997
2899
|
// below a 6x6=36 workgroup grid (not enough workgroups to fill the GPU at
|
|
1998
2900
|
// that tile size), hence the two-tier split rather than one global config.
|
|
1999
|
-
//
|
|
2000
|
-
//
|
|
2001
|
-
//
|
|
2901
|
+
//
|
|
2902
|
+
// A and B are bound twice \u2014 scalar array<f32> and array<vec4<f32>> views of
|
|
2903
|
+
// the same GPUBuffer (see vec4ViewBinding) \u2014 so each tile load can issue
|
|
2904
|
+
// 16-byte vector reads along op(A)/op(B)'s contiguous dimension when the
|
|
2905
|
+
// stride allows it (stride % 4 == 0 keeps every row base 16-byte aligned).
|
|
2906
|
+
// Transposed or odd-stride operands take the scalar path; both paths
|
|
2907
|
+
// zero-fill out-of-bounds components identically.
|
|
2002
2908
|
|
|
2003
2909
|
const BM: u32 = 64u;
|
|
2004
2910
|
const BN: u32 = 64u;
|
|
@@ -2011,9 +2917,11 @@ const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 128
|
|
|
2011
2917
|
const STRIDE_A: u32 = NUM_THREADS / BK; // rows of As covered per load-loop step
|
|
2012
2918
|
const STRIDE_B: u32 = NUM_THREADS / BN; // rows of Bs covered per load-loop step
|
|
2013
2919
|
|
|
2014
|
-
@group(0) @binding(0) var<storage, read> A:
|
|
2015
|
-
@group(0) @binding(1) var<storage, read>
|
|
2016
|
-
@group(0) @binding(2) var<storage,
|
|
2920
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
2921
|
+
@group(0) @binding(1) var<storage, read> A4: array<vec4<f32>>;
|
|
2922
|
+
@group(0) @binding(2) var<storage, read> B: array<f32>;
|
|
2923
|
+
@group(0) @binding(3) var<storage, read> B4: array<vec4<f32>>;
|
|
2924
|
+
@group(0) @binding(4) var<storage, read_write> C: array<f32>;
|
|
2017
2925
|
|
|
2018
2926
|
struct Params {
|
|
2019
2927
|
m: u32,
|
|
@@ -2026,9 +2934,11 @@ struct Params {
|
|
|
2026
2934
|
ldc: u32,
|
|
2027
2935
|
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
2028
2936
|
transB: u32,
|
|
2937
|
+
useVecA: u32, // 1 = A's vec4 view covers every in-bounds element (see vec4Usable)
|
|
2938
|
+
useVecB: u32, // 1 = B's vec4 view covers every in-bounds element
|
|
2029
2939
|
}
|
|
2030
2940
|
|
|
2031
|
-
@group(0) @binding(
|
|
2941
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
2032
2942
|
|
|
2033
2943
|
var<workgroup> As: array<f32, BM * BK>;
|
|
2034
2944
|
var<workgroup> Bs: array<f32, BK * BN>;
|
|
@@ -2060,17 +2970,95 @@ fn main(
|
|
|
2060
2970
|
|
|
2061
2971
|
let numTiles = (params.k + BK - 1u) / BK;
|
|
2062
2972
|
for (var t = 0u; t < numTiles; t++) {
|
|
2063
|
-
|
|
2064
|
-
|
|
2065
|
-
|
|
2066
|
-
|
|
2067
|
-
|
|
2973
|
+
// \u2500\u2500 Load the BM\xD7BK A tile into As (vectorized along op(A)'s fast dim
|
|
2974
|
+
// when lda allows; every branch here is dispatch-uniform) \u2500\u2500
|
|
2975
|
+
if (params.useVecA == 1u && params.transA == 0u) {
|
|
2976
|
+
// No-transpose: columns contiguous. Each thread loads one vec4 of 4
|
|
2977
|
+
// columns; 64 rows \xD7 2 column-lanes = NUM_THREADS exactly, single pass.
|
|
2978
|
+
let r4 = tid / (BK / 4u);
|
|
2979
|
+
let c4 = tid % (BK / 4u);
|
|
2980
|
+
let gRow = blockRow + r4;
|
|
2981
|
+
let gCol = t * BK + c4 * 4u;
|
|
2982
|
+
var v = A4[(gRow * params.lda + gCol) / 4u];
|
|
2983
|
+
let rowOK = gRow < params.m;
|
|
2984
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.k);
|
|
2985
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.k);
|
|
2986
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.k);
|
|
2987
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.k);
|
|
2988
|
+
As[r4 * BK + c4 * 4u] = v.x;
|
|
2989
|
+
As[r4 * BK + c4 * 4u + 1u] = v.y;
|
|
2990
|
+
As[r4 * BK + c4 * 4u + 2u] = v.z;
|
|
2991
|
+
As[r4 * BK + c4 * 4u + 3u] = v.w;
|
|
2992
|
+
} else if (params.useVecA == 1u && params.transA != 0u) {
|
|
2993
|
+
// Transpose: rows contiguous within a column. Each thread loads one
|
|
2994
|
+
// vec4 of 4 rows; 16 row-lanes \xD7 8 columns = NUM_THREADS, single pass.
|
|
2995
|
+
let r4 = tid % (BM / 4u);
|
|
2996
|
+
let c = tid / (BM / 4u);
|
|
2997
|
+
let gRow = blockRow + r4 * 4u;
|
|
2998
|
+
let gCol = t * BK + c;
|
|
2999
|
+
var v = A4[(gCol * params.lda + gRow) / 4u];
|
|
3000
|
+
let colOK = gCol < params.k;
|
|
3001
|
+
v.x = select(0.0, v.x, colOK && gRow < params.m);
|
|
3002
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.m);
|
|
3003
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.m);
|
|
3004
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.m);
|
|
3005
|
+
As[(r4 * 4u) * BK + c] = v.x;
|
|
3006
|
+
As[(r4 * 4u + 1u) * BK + c] = v.y;
|
|
3007
|
+
As[(r4 * 4u + 2u) * BK + c] = v.z;
|
|
3008
|
+
As[(r4 * 4u + 3u) * BK + c] = v.w;
|
|
3009
|
+
} else {
|
|
3010
|
+
// Scalar fallback: odd stride or unhandled orientation.
|
|
3011
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
3012
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
3013
|
+
let gColA = t * BK + innerColA;
|
|
3014
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
3015
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
3016
|
+
}
|
|
2068
3017
|
}
|
|
2069
|
-
|
|
2070
|
-
|
|
2071
|
-
|
|
2072
|
-
|
|
2073
|
-
|
|
3018
|
+
|
|
3019
|
+
// \u2500\u2500 Load the BK\xD7BN B tile into Bs \u2500\u2500
|
|
3020
|
+
if (params.useVecB == 1u && params.transB == 0u) {
|
|
3021
|
+
// No-transpose: columns contiguous. 8 rows \xD7 16 column-lanes cover the
|
|
3022
|
+
// tile in one pass (BK = NUM_THREADS / (BN/4)).
|
|
3023
|
+
let r = tid / (BN / 4u);
|
|
3024
|
+
let c4 = tid % (BN / 4u);
|
|
3025
|
+
let gRow = t * BK + r;
|
|
3026
|
+
let gCol = blockCol + c4 * 4u;
|
|
3027
|
+
var v = B4[(gRow * params.ldb + gCol) / 4u];
|
|
3028
|
+
let rowOK = gRow < params.k;
|
|
3029
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.n);
|
|
3030
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.n);
|
|
3031
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.n);
|
|
3032
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.n);
|
|
3033
|
+
Bs[r * BN + c4 * 4u] = v.x;
|
|
3034
|
+
Bs[r * BN + c4 * 4u + 1u] = v.y;
|
|
3035
|
+
Bs[r * BN + c4 * 4u + 2u] = v.z;
|
|
3036
|
+
Bs[r * BN + c4 * 4u + 3u] = v.w;
|
|
3037
|
+
} else if (params.useVecB == 1u && params.transB != 0u) {
|
|
3038
|
+
// Transpose: rows contiguous within a column. 2 row-lanes \xD7 64 columns
|
|
3039
|
+
// cover the tile in one pass (BN = NUM_THREADS / (BK/4)).
|
|
3040
|
+
let r4 = tid % (BK / 4u);
|
|
3041
|
+
let c = tid / (BK / 4u);
|
|
3042
|
+
let gRow = t * BK + r4 * 4u;
|
|
3043
|
+
let gCol = blockCol + c;
|
|
3044
|
+
var v = B4[(gCol * params.ldb + gRow) / 4u];
|
|
3045
|
+
let colOK = gCol < params.n;
|
|
3046
|
+
v.x = select(0.0, v.x, colOK && gRow < params.k);
|
|
3047
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.k);
|
|
3048
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.k);
|
|
3049
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.k);
|
|
3050
|
+
Bs[(r4 * 4u) * BN + c] = v.x;
|
|
3051
|
+
Bs[(r4 * 4u + 1u) * BN + c] = v.y;
|
|
3052
|
+
Bs[(r4 * 4u + 2u) * BN + c] = v.z;
|
|
3053
|
+
Bs[(r4 * 4u + 3u) * BN + c] = v.w;
|
|
3054
|
+
} else {
|
|
3055
|
+
// Scalar fallback.
|
|
3056
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
3057
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
3058
|
+
let gColB = blockCol + innerColB;
|
|
3059
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
3060
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
3061
|
+
}
|
|
2074
3062
|
}
|
|
2075
3063
|
|
|
2076
3064
|
workgroupBarrier();
|
|
@@ -2099,13 +3087,16 @@ fn main(
|
|
|
2099
3087
|
let col = blockCol + threadCol * TN + resIdxN;
|
|
2100
3088
|
if (col < params.n) {
|
|
2101
3089
|
let cIdx = row * params.ldc + col;
|
|
2102
|
-
|
|
3090
|
+
// BLAS beta==0 semantics: C is written, not accumulated \u2014 must not
|
|
3091
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
3092
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
3093
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
2103
3094
|
}
|
|
2104
3095
|
}
|
|
2105
3096
|
}
|
|
2106
3097
|
}
|
|
2107
3098
|
}
|
|
2108
|
-
`});var
|
|
3099
|
+
`});var xe,Go=O(()=>{xe=`// sgemmtr_small: C := uplo(alpha * op(A) * op(B) + beta * C) \u2014 small-tile
|
|
2109
3100
|
// half of a two-tier dispatch, identical to sgemm_small.wgsl except the
|
|
2110
3101
|
// final output write is gated to one triangle of C by \`uplo\` \u2014 see
|
|
2111
3102
|
// sgemmtr_large.wgsl for the full rationale (shared by both tiers).
|
|
@@ -2209,13 +3200,16 @@ fn main(
|
|
|
2209
3200
|
let inTriangle = select(col >= row, col <= row, params.uplo == 0u);
|
|
2210
3201
|
if (col < params.n && inTriangle) {
|
|
2211
3202
|
let cIdx = row * params.ldc + col;
|
|
2212
|
-
|
|
3203
|
+
// BLAS beta==0 semantics: C is written, not accumulated \u2014 must not
|
|
3204
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
3205
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
3206
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
2213
3207
|
}
|
|
2214
3208
|
}
|
|
2215
3209
|
}
|
|
2216
3210
|
}
|
|
2217
3211
|
}
|
|
2218
|
-
`});var
|
|
3212
|
+
`});var ve,Eo=O(()=>{ve=`// sgemmtr_large: C := uplo(alpha * op(A) * op(B) + beta * C) \u2014 large-tile
|
|
2219
3213
|
// half of a two-tier dispatch, identical to sgemm_large.wgsl (see that file
|
|
2220
3214
|
// for the BM/BN/BK/TM/TN autotuning rationale) except the final output write
|
|
2221
3215
|
// is gated to one triangle of C by \`uplo\`, the same convention ssyr/ssyr2
|
|
@@ -2326,13 +3320,16 @@ fn main(
|
|
|
2326
3320
|
let inTriangle = select(col >= row, col <= row, params.uplo == 0u);
|
|
2327
3321
|
if (col < params.n && inTriangle) {
|
|
2328
3322
|
let cIdx = row * params.ldc + col;
|
|
2329
|
-
|
|
3323
|
+
// BLAS beta==0 semantics: C is written, not accumulated \u2014 must not
|
|
3324
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
3325
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
3326
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
2330
3327
|
}
|
|
2331
3328
|
}
|
|
2332
3329
|
}
|
|
2333
3330
|
}
|
|
2334
3331
|
}
|
|
2335
|
-
`});var
|
|
3332
|
+
`});var Do,ko=O(()=>{Do=`// symmetrize: Adense := full dense expansion of a symmetric matrix stored
|
|
2336
3333
|
// with only its \`uplo\` triangle meaningful (the other triangle is implied
|
|
2337
3334
|
// by symmetry: A[i,j] = A[j,i]). A plain element-wise pass, no tiling or
|
|
2338
3335
|
// shared memory needed \u2014 used to materialize a dense operand for routines
|
|
@@ -2363,7 +3360,7 @@ fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
|
2363
3360
|
let srcIdx = select(col * params.lda + row, row * params.lda + col, isStored);
|
|
2364
3361
|
Adense[row * params.ldd + col] = A[srcIdx];
|
|
2365
3362
|
}
|
|
2366
|
-
`});var
|
|
3363
|
+
`});var Po,No=O(()=>{Po=`// triangularize: Adense := dense expansion of op(A) (A or A^T per \`trans\`),
|
|
2367
3364
|
// zero-filling the unstored triangle (exact for a matmul) so strmm can reuse
|
|
2368
3365
|
// sgemm's kernel unchanged. \`diag=1\` substitutes 1.0 on the diagonal.
|
|
2369
3366
|
|
|
@@ -2407,7 +3404,7 @@ fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
|
2407
3404
|
|
|
2408
3405
|
Adense[row * params.ldd + col] = select(0.0, A[srcRow * params.lda + srcCol], isMeaningful);
|
|
2409
3406
|
}
|
|
2410
|
-
`});var
|
|
3407
|
+
`});var Io,Mo=O(()=>{Io=`// block_transfer: gather/scatter/scatter-subtract between a tight (blockLen
|
|
2411
3408
|
// x otherLen) block and a sub-range of a strided (any ld, row/col-major)
|
|
2412
3409
|
// buffer \u2014 needed since block offsets aren't 256-byte-aligned and block
|
|
2413
3410
|
// rows/cols aren't always one contiguous range for copyBufferToBuffer.
|
|
@@ -2449,7 +3446,8 @@ fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
|
2449
3446
|
strided[stridedIdx] = block[blockIdx];
|
|
2450
3447
|
}
|
|
2451
3448
|
}
|
|
2452
|
-
`});var Ct={};te(Ct,{shaderSources:()=>Ga});var Ga,Wt=O(()=>{pe();ge();he();ve();_e();Ee();Ge();Pe();Se();Ie();De();Te();Ce();Fe();Ue();Ve();ze();Ye();Qe();$e();rt();tt();at();st();ut();ft();ct();pt();gt();ht();vt();_t();Et();Gt();Pt();St();It();Dt();Tt();Ga={"reduction/argmax":we,"reduction/argmaxF64":be,"reduction/sum":xe,"reduction/sumF64":ye,sscal:Be,sswap:Ae,saxpy:ke,scopy:Ne,sdot:Me,sasum:Le,snrm2:Re,srot:je,srotm:We,isamax:He,sgemv_n:Oe,sgemv_t:Ke,ssymv:qe,strmv:Xe,sger:Ze,ssyr:Je,ssyr2:et,f64add:ot,"f64/dekker":it,"f64/utils/abs":nt,"f64/utils/add":lt,"f64/utils/greater":mt,"f64/utils/equal":dt,dasum:wt,idamax:bt,strsv_invert_block:xt,strsv_apply_inverse:yt,strsv_update:Bt,sgemm_small:At,sgemm_large:kt,sgemmtr_small:Nt,sgemmtr_large:Mt,symmetrize:Lt,triangularize:Rt,block_transfer:jt}});var ni={};te(ni,{GpuMatrix:()=>H,GpuVector:()=>I,cleanup:()=>le,dasum:()=>Xt,gpuName:()=>fe,idamax:()=>Jt,init:()=>ue,isamax:()=>$t,randomFloat32Array:()=>me,randomFloat64Array:()=>ce,randomTriangularFloat32Array:()=>de,sasum:()=>Yt,saxpy:()=>Ot,scopy:()=>Vt,sdot:()=>zt,sgemm:()=>mo,sgemmtr:()=>co,sgemv:()=>to,sger:()=>uo,snrm2:()=>Zt,srot:()=>ro,srotm:()=>eo,sscal:()=>Ht,sswap:()=>Ut,ssymm:()=>bo,ssymv:()=>oo,ssyr:()=>lo,ssyr2:()=>fo,ssyr2k:()=>wo,ssyrk:()=>po,strmm:()=>xo,strmv:()=>ao,strsm:()=>Ao,strsv:()=>no});function ae(a,e){return e?a.features.has("timestamp-query")?{requiredFeatures:["timestamp-query"]}:(console.warn("timestamp-query not supported on this device \u2014 benchmark mode disabled."),{}):{}}function ie(){if(!se())return{querySet:null,passDescriptor:void 0};let e=lr().createQuerySet({type:"timestamp",count:2});return{querySet:e,passDescriptor:{timestampWrites:{querySet:e,beginningOfPassWriteIndex:0,endOfPassWriteIndex:1}}}}function br(a,e){if(!e)return null;let r=lr(),o=r.createBuffer({label:"timestamp-resolve",size:16,usage:GPUBufferUsage.QUERY_RESOLVE|GPUBufferUsage.COPY_SRC});a.resolveQuerySet(e,0,2,o,0);let t=r.createBuffer({label:"timestamp-readback",size:16,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ});return a.copyBufferToBuffer(o,0,t,0,16),{tsReadBuffer:t,resolveBuffer:o,querySet:e}}async function S(a){if(!a)return;let{tsReadBuffer:e,resolveBuffer:r,querySet:o}=a;await e.mapAsync(GPUMapMode.READ);let t=new BigInt64Array(e.getMappedRange().slice());return e.unmap(),e.destroy(),r.destroy(),o.destroy(),Math.max(0,Number(t[1]-t[0]))/1e6}var Ar=null,Ir=null,ne=null,Qr=!1;async function ue({powerPreference:a="high-performance",benchmark:e=!1,dumpShaders:r=!1}={}){if(Ar)return Ar;let o;if(typeof window>"u"){let{create:s,globals:l}=await import("webgpu");Object.assign(globalThis,l),o=s(r?["enable-dawn-features=dump_shaders,disable_symbol_renaming"]:[]),ne=o}else r&&console.warn("dumpShaders has no effect in the browser \u2014 see init()'s docs."),o=navigator.gpu;if(!o)throw new Error("WebGPU not supported in this environment.");if(Ir=await o.requestAdapter({powerPreference:a})??await o.requestAdapter(),!Ir)throw new Error("No WebGPU adapter found.");Qr=e;let i=[...ae(Ir,e).requiredFeatures??[]];return Ar=await Ir.requestDevice({requiredFeatures:i}),Ar.addEventListener("uncapturederror",s=>{console.error("Uncaptured GPU error:",s.error.message)}),Ar}function le(){Ar&&(Ar.destroy(),Ar=null),Ir=null,ne=null,Qr=!1}function fe(){if(!Ir)throw new Error("WebGPU adapter not initialized \u2014 call init() first.");let{device:a,description:e}=Ir.info;return{description:e||"unknown",device:a||"unknown"}}function se(){return Qr}function lr(){if(!Ar)throw new Error("WebGPU device not initialized \u2014 call init() first.");return Ar}function d(...a){a.flat().forEach(e=>e.destroy())}function v(a,e="blas-input",r=!1){let o=lr(),t=o.limits.maxStorageBufferBindingSize,i=a.byteLength;if(i>t)throw new Error(`Buffer size ${i} bytes exceeds device limit of ${t} bytes.`);let s=r?GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC:GPUBufferUsage.STORAGE,l=o.createBuffer({label:e,size:i,usage:s,mappedAtCreation:!0}),n=a.constructor;return new n(l.getMappedRange()).set(a),l.unmap(),l}function er(a,e="blas-storage",r=0){return lr().createBuffer({label:e,size:a,usage:GPUBufferUsage.STORAGE|r})}function xr(a,e="blas-result"){return lr().createBuffer({label:e,size:a,usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC})}function N(a,e){let o=lr().createBuffer({label:"blas-readback",size:e.size,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ});return a.copyBufferToBuffer(e,0,o,0,e.size),o}function L(a,e="blas-params"){let r=lr(),o=a.length*4,t=Math.ceil(o/16)*16,i=new ArrayBuffer(t),s=new DataView(i);a.forEach(({value:n,type:u},f)=>{let m=f*4;if(u==="u32")s.setUint32(m,n,!0);else if(u==="i32")s.setInt32(m,n,!0);else if(u==="f32")s.setFloat32(m,n,!0);else throw new Error(`Unknown param type "${u}". Use "f32", "u32", or "i32".`)});let l=r.createBuffer({label:e,size:t,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST});return r.queue.writeBuffer(l,0,i),l}async function k(a,e=Float32Array){try{await a.mapAsync(GPUMapMode.READ);let r=new e(a.getMappedRange().slice());return a.unmap(),r}finally{a.destroy()}}function Pr(a){let e=a.length,r=new Float32Array(e),o=new Float32Array(e);for(let t=0;t<e;t++){let i=Math.fround(a[t]);r[t]=i,o[t]=Math.fround(a[t]-i)}return{hi:r,lo:o}}function Dr(a,e){let r=a.length,o=new Float64Array(r);for(let t=0;t<r;t++)o[t]=a[t]+e[t];return o}var I=class a{constructor(e,r,o=Float32Array,t=null){this._buf=e,this._loBuf=t,this.length=r,this.dtype=o}static from(e){if(e instanceof Float64Array){let{hi:o,lo:t}=Pr(e),i=v(o,"gpu-vector-f64-hi",!0),s=v(t,"gpu-vector-f64-lo",!0);return new a(i,e.length,Float64Array,s)}if(!(e instanceof Float32Array))throw new Error("GpuVector.from expects a Float32Array or Float64Array.");let r=v(e,"gpu-vector",!0);return new a(r,e.length,e.constructor)}async read(){let e=lr(),r=e.createCommandEncoder(),o=N(r,this._buf);if(e.queue.submit([r.finish()]),!this._loBuf)return k(o,this.dtype);let t=e.createCommandEncoder(),i=N(t,this._loBuf);e.queue.submit([t.finish()]);let[s,l]=await Promise.all([k(o,Float32Array),k(i,Float32Array)]);return Dr(s,l)}destroy(){this._buf.destroy(),this._loBuf&&this._loBuf.destroy()}};var H=class a{constructor(e,r,o,t,i=null,s="row-major"){this._buf=e,this._loBuf=i,this.rows=r,this.cols=o,this.lda=t,this.layout=s}static from(e,r,o,t,i="row-major"){if(i!=="row-major"&&i!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");let s=i==="row-major";if(t===void 0&&(t=s?o:r),!(e instanceof Float32Array)&&!(e instanceof Float64Array))throw new Error("GpuMatrix.from expects a Float32Array or Float64Array.");if(!Number.isInteger(r)||r<=0)throw new Error("rows must be a positive integer.");if(!Number.isInteger(o)||o<=0)throw new Error("cols must be a positive integer.");let l=s?o:r;if(!Number.isInteger(t)||t<l)throw new Error(`lda must be an integer >= ${s?"cols":"rows"}.`);let n=s?r:o;if(e.length<n*t)throw new Error("data does not have enough elements for the given rows, cols, and lda.");if(e instanceof Float64Array){let f=n*t,{hi:m,lo:w}=Pr(e.subarray(0,f)),c=v(m,"gpu-matrix-f64-hi",!0),p=v(w,"gpu-matrix-f64-lo",!0);return new a(c,r,o,t,p,i)}let u=v(e.subarray(0,n*t),"gpu-matrix",!0);return new a(u,r,o,t,null,i)}async read(){let e=lr(),r=e.createCommandEncoder(),o=N(r,this._buf);e.queue.submit([r.finish()]);let t=this.layout!=="column-major",i=t?this.rows:this.cols,s=t?this.cols:this.rows;if(this._loBuf){let u=e.createCommandEncoder(),f=N(u,this._loBuf);e.queue.submit([u.finish()]);let[m,w]=await Promise.all([k(o,Float32Array),k(f,Float32Array)]),c=Dr(m,w);if(this.lda===s)return c;let p=new Float64Array(i*s);for(let g=0;g<i;g++)p.set(c.subarray(g*this.lda,g*this.lda+s),g*s);return p}let l=await k(o,Float32Array);if(this.lda===s)return l;let n=new Float32Array(i*s);for(let u=0;u<i;u++)n.set(l.subarray(u*this.lda,u*this.lda+s),u*s);return n}destroy(){this._buf.destroy(),this._loBuf&&this._loBuf.destroy()}};function me(a,e=-1,r=1){let o=new Float32Array(a);for(let t=0;t<a;t++)o[t]=e+Math.random()*(r-e);return o}function ce(a,e=-1,r=1){let o=new Float64Array(a);for(let t=0;t<a;t++)o[t]=e+Math.random()*(r-e);return o}function de(a,e,r="lower",o=-1,t=1,i=5,s=15){if(r!=="lower"&&r!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(e<a)throw new Error("lda must be >= n.");let l=new Float32Array(a*e);for(let n=0;n<a;n++){for(let u=0;u<a;u++){if(n===u)continue;(r==="lower"?u<n:u>n)&&(l[n*e+u]=o+Math.random()*(t-o))}l[n*e+n]=i+Math.random()*(s-i)}return l}function B(a,e,r=0){let o=lr(),t=e.map((i,s)=>({binding:r+s,resource:i instanceof GPUBuffer?{buffer:i}:i}));return o.createBindGroup({layout:a,entries:t})}var Fo=new WeakMap;function M(a){lr().queue.submit([a.finish()])}function vr(){let a=lr(),{querySet:e,passDescriptor:r}=ie();return{commandEncoder:a.createCommandEncoder(),querySet:e,passDescriptor:r}}function ar(a,e,r,o,t){let i=a.beginComputePass(t);i.setPipeline(e),i.setBindGroup(0,r),typeof o=="number"?i.dispatchWorkgroups(o):i.dispatchWorkgroups(o.x,o.y,o.z??1),i.end(),Fo.set(a,i)}function C(a,e,r){let{commandEncoder:o,querySet:t,passDescriptor:i}=vr();ar(o,a,e,r,i);let s=br(o,t);return{commandEncoder:o,ts:s}}var Na={},Zr=new WeakMap;async function G(a,e,r="main"){Zr.has(a)||Zr.set(a,new Map);let o=Zr.get(a),t=Array.isArray(e)?e:[e],i=`${t.join("+")}::${r}`;return o.has(i)||o.set(i,await Pa(t,r)),o.get(i)}async function ka(a){if(typeof process>"u"||!process.versions?.node){let{shaderSources:e}=await Promise.resolve().then(()=>(Wt(),Ct)),r=e[a];if(!r)throw new Error(`Shader "${a}" not found in browser bundle.`);return r}else{let{readFileSync:e}=await import("fs"),{fileURLToPath:r}=await import("url"),{dirname:o,join:t}=await import("path"),i=o(r(Na.url));return e(t(i,`../shaders/${a}.wgsl`),"utf8")}}async function Pa(a,e="main"){let r=lr(),o=a.join("+"),t=(await Promise.all(a.map(ka))).join(`
|
|
2453
|
-
`),
|
|
2454
|
-
|
|
2455
|
-
`)}`);let n=e==="main"?{module:i}:{module:i,entryPoint:e},u=r.createComputePipeline({label:o,layout:"auto",compute:n});return u._shaderModule=i,u}var Sa=64,Ft=8;function mr(a,e){let r=lr().limits.maxComputeWorkgroupsPerDimension;return e===void 0?Math.min(Math.ceil(a/Sa),r):{x:Math.min(Math.ceil(e/Ft),r),y:Math.min(Math.ceil(a/Ft),r)}}async function Ht(a,e,r,o,t){let i=o instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(t))throw new Error("n and incx must be integers.");if(typeof r!="number")throw new Error("alpha must be a number.");if(Number.isNaN(r))throw new Error("alpha must not be NaN.");if(!Number.isFinite(r))throw new Error("alpha must be finite.");if(t<=0)throw new Error("incx must be positive.");if(!(o instanceof Float32Array)&&!(o instanceof I))throw new Error("x must be a Float32Array or GpuVector.");if(e<=0)return i?{}:o;if(o.length<(e-1)*t+1)throw new Error("x does not have enough elements for the given n and incx.");let s=await G(a,"sscal"),l=null,n=null,u=null;try{l=i?o._buf:v(o,"sscal-x",!0),n=L([{value:e,type:"u32"},{value:r,type:"f32"},{value:t,type:"u32"}],"sscal-params");let f=B(s.getBindGroupLayout(0),[l,n]),{commandEncoder:m,ts:w}=C(s,f,mr(e));u=i?null:N(m,l),M(m);let c=await S(w);if(i)return c!==void 0?{gpuTimeMs:c}:{};let p=await k(u,Float32Array);return u=null,c!==void 0?{x:p,gpuTimeMs:c}:p}finally{!i&&l&&d(l),n&&d(n),u&&d(u)}}async function Ut(a,e,r,o,t,i){let s=r instanceof I,l=t instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!(r instanceof Float32Array)&&!(r instanceof I))throw new Error("x must be a Float32Array or GpuVector.");if(!(t instanceof Float32Array)&&!(t instanceof I))throw new Error("y must be a Float32Array or GpuVector.");if(r.constructor!==t.constructor)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(e<=0)return s?{}:{x:r,y:t};if(r.length<(e-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(t.length<(e-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let n=await G(a,"sswap"),u=null,f=null,m=null,w=null,c=null;try{u=s?r._buf:v(r,"sswap-x",!0),f=l?t._buf:v(t,"sswap-y",!0),m=L([{value:e,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"sswap-params");let p=B(n.getBindGroupLayout(0),[u,f,m]),{commandEncoder:g,ts:h}=C(n,p,mr(e));w=s?null:N(g,u),c=l?null:N(g,f),M(g);let b=await S(h);if(s&&l)return b!==void 0?{gpuTimeMs:b}:{};let x=await k(w,Float32Array);w=null;let _=await k(c,Float32Array);return c=null,b!==void 0?{x,y:_,gpuTimeMs:b}:{x,y:_}}finally{!s&&u&&d(u),!l&&f&&d(f),m&&d(m),w&&d(w),c&&d(c)}}async function Ot(a,e,r,o,t,i,s){let l=o instanceof I,n=i instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(t)||!Number.isInteger(s))throw new Error("n, incx, and incy must be integers.");if(typeof r!="number")throw new Error("alpha must be a number.");if(Number.isNaN(r))throw new Error("alpha must not be NaN.");if(!Number.isFinite(r))throw new Error("alpha must be finite.");if(t<=0||s<=0)throw new Error("incx and incy must be positive.");if(!l&&!(o instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!n&&!(i instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(l!==n)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(e<=0)return n?{}:{y:i};if(o.length<(e-1)*t+1)throw new Error("x does not have enough elements for the given n and incx.");if(i.length<(e-1)*s+1)throw new Error("y does not have enough elements for the given n and incy.");let u=await G(a,"saxpy"),f=null,m=null,w=null,c=null;try{f=l?o._buf:v(o,"saxpy-x",!1),m=n?i._buf:v(i,"saxpy-y",!0),w=L([{value:e,type:"u32"},{value:r,type:"f32"},{value:t,type:"u32"},{value:s,type:"u32"}],"saxpy-params");let p=B(u.getBindGroupLayout(0),[f,m,w]),{commandEncoder:g,ts:h}=C(u,p,mr(e));c=n?null:N(g,m),M(g);let b=await S(h);if(n&&l)return b!==void 0?{gpuTimeMs:b}:{};let x=await k(c,Float32Array);return c=null,b!==void 0?{y:x,gpuTimeMs:b}:{y:x}}finally{!l&&f&&d(f),!n&&m&&d(m),w&&d(w),c&&d(c)}}async function Vt(a,e,r,o,t,i){let s=r instanceof I,l=t instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!s&&!(r instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!l&&!(t instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(s!==l)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(e<=0)return l?{}:{y:t};if(r.length<(e-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(t.length<(e-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let n=await G(a,"scopy"),u=null,f=null,m=null,w=null;try{u=s?r._buf:v(r,"scopy-x",!1),f=l?t._buf:v(t,"scopy-y",!0),m=L([{value:e,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"scopy-params");let c=B(n.getBindGroupLayout(0),[u,f,m]),{commandEncoder:p,ts:g}=C(n,c,mr(e));w=l?null:N(p,f),M(p);let h=await S(g);if(l&&s)return h!==void 0?{gpuTimeMs:h}:{};let b=await k(w,Float32Array);return w=null,h!==void 0?{y:b,gpuTimeMs:h}:{y:b}}finally{!s&&u&&d(u),!l&&f&&d(f),m&&d(m),w&&d(w)}}var Kt=64;async function zt(a,e,r,o,t,i){let s=r instanceof I,l=t instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!s&&!(r instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!l&&!(t instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(s!==l)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(e<=0)return{dot:0};if(r.length<(e-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(t.length<(e-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let n=await G(a,"sdot"),u=await G(a,"reduction/sum"),f=null,m=null,w=null,c=null,p=null,g=null;try{f=s?r._buf:v(r,"sdot-x",!1),m=l?t._buf:v(t,"sdot-y",!1),w=er(2*Kt*4,"sdot-partials"),c=xr(4,"sdot-result"),p=L([{value:e,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"sdot-params");let h=B(n.getBindGroupLayout(0),[f,m,w,p]),{commandEncoder:b,ts:x}=C(n,h,2*Kt);M(b);let _=B(u.getBindGroupLayout(0),[w,c]),{commandEncoder:y,ts:A}=C(u,_,1);g=N(y,c),M(y);let P=k(g,Float32Array);g=null;let[E,T,D]=await Promise.all([S(x),S(A),P]);return E!==void 0&&T!==void 0?{dot:D[0],gpuTimeMs:E+T}:{dot:D[0]}}finally{!s&&f&&d(f),!l&&m&&d(m),w&&d(w),c&&d(c),p&&d(p),g&&d(g)}}var qt=64;async function Yt(a,e,r,o){let t=r instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(o<=0)throw new Error("incx must be positive.");if(!t&&!(r instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(e<=0)return{asum:0};if(r.length<(e-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let i=await G(a,"sasum"),s=await G(a,"reduction/sum"),l=null,n=null,u=null,f=null,m=null;try{l=t?r._buf:v(r,"sasum-x",!1),n=er(2*qt*4,"sasum-partials"),u=xr(4,"sasum-result"),f=L([{value:e,type:"u32"},{value:o,type:"u32"}],"sasum-params");let w=B(i.getBindGroupLayout(0),[l,n,f]),{commandEncoder:c,ts:p}=C(i,w,2*qt);M(c);let g=B(s.getBindGroupLayout(0),[n,u]),{commandEncoder:h,ts:b}=C(s,g,1);m=N(h,u),M(h);let x=k(m,Float32Array);m=null;let[_,y,A]=await Promise.all([S(p),S(b),x]);return _!==void 0&&y!==void 0?{asum:A[0],gpuTimeMs:_+y}:{asum:A[0]}}finally{!t&&l&&d(l),n&&d(n),u&&d(u),f&&d(f),m&&d(m)}}var $r=64;async function Xt(a,e,r,o){let t=r instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(o<=0)throw new Error("incx must be positive.");if(!t&&!(r instanceof Float64Array))throw new Error("x must be a Float64Array or GpuVector.");if(t&&r.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(e<=0)return{asum:0};if(r.length<(e-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let i=["f64/dekker","f64/utils/abs","f64/utils/add"],s=await G(a,[...i,"dasum"]),l=await G(a,[...i,"reduction/sumF64"]),n=null,u=null,f=null,m=null,w=null,c=null,p=null,g=null,h=null;try{if(t)n=r._buf,u=r._loBuf;else{let{hi:K,lo:W}=Pr(r.map(Math.abs));n=v(K,"dasum-xHi",!1),u=v(W,"dasum-xLo",!1)}f=er(2*$r*4,"dasum-partialsHi"),m=er(2*$r*4,"dasum-partialsLo"),w=xr(4,"dasum-result-hi"),c=xr(4,"dasum-result-lo"),p=L([{value:e,type:"u32"},{value:o,type:"u32"}],"dasum-params");let b=B(s.getBindGroupLayout(0),[n,u,f,m,p]),{commandEncoder:x,ts:_}=C(s,b,2*$r);M(x);let y=B(l.getBindGroupLayout(0),[f,m,w,c]),{commandEncoder:A,ts:P}=C(l,y,1);g=N(A,w),h=N(A,c),M(A);let E=k(g,Float32Array),T=k(h,Float32Array);g=null,h=null;let[D,R,j,F]=await Promise.all([S(_),S(P),E,T]),V=Dr(j,F)[0];return D!==void 0&&R!==void 0?{asum:V,gpuTimeMs:D+R}:{asum:V}}finally{!t&&n&&d(n),!t&&u&&d(u),f&&d(f),m&&d(m),w&&d(w),c&&d(c),p&&d(p),g&&d(g),h&&d(h)}}var Qt=64;async function Zt(a,e,r,o){let t=r instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(o<=0)throw new Error("incx must be positive.");if(!t&&!(r instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(e<=0)return{nrm2:0};if(r.length<(e-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let i=await G(a,"snrm2"),s=await G(a,"reduction/sum"),l=null,n=null,u=null,f=null,m=null;try{l=t?r._buf:v(r,"snrm2-x",!1),n=er(2*Qt*4,"snrm2-partials"),u=xr(4,"snrm2-result"),f=L([{value:e,type:"u32"},{value:o,type:"u32"}],"snrm2-params");let w=B(i.getBindGroupLayout(0),[l,n,f]),{commandEncoder:c,ts:p}=C(i,w,2*Qt);M(c);let g=B(s.getBindGroupLayout(0),[n,u]),{commandEncoder:h,ts:b}=C(s,g,1);m=N(h,u),M(h);let x=k(m,Float32Array);m=null;let[_,y,A]=await Promise.all([S(p),S(b),x]),P=Math.sqrt(A[0]);return _!==void 0&&y!==void 0?{nrm2:P,gpuTimeMs:_+y}:{nrm2:P}}finally{!t&&l&&d(l),n&&d(n),u&&d(u),f&&d(f),m&&d(m)}}var Jr=64;async function $t(a,e,r,o){let t=r instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(o<=0)throw new Error("incx must be positive.");if(!t&&!(r instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(e<=0)return{index:0};if(r.length<(e-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let i=await G(a,"isamax"),s=await G(a,"reduction/argmax"),l=null,n=null,u=null,f=null,m=null,w=null;try{l=t?r._buf:v(r,"isamax-x",!1),n=er(2*Jr*4,"isamax-partials-val"),u=er(2*Jr*4,"isamax-partials-idx"),f=xr(4,"isamax-result"),m=L([{value:e,type:"u32"},{value:o,type:"u32"}],"isamax-params");let c=B(i.getBindGroupLayout(0),[l,n,u,m]),{commandEncoder:p,ts:g}=C(i,c,2*Jr);M(p);let h=B(s.getBindGroupLayout(0),[n,u,f]),{commandEncoder:b,ts:x}=C(s,h,1);w=N(b,f),M(b);let _=k(w,Uint32Array);w=null;let[y,A,P]=await Promise.all([S(g),S(x),_]),E=P[0];return y!==void 0&&A!==void 0?{index:E,gpuTimeMs:y+A}:{index:E}}finally{!t&&l&&d(l),n&&d(n),u&&d(u),f&&d(f),m&&d(m),w&&d(w)}}var Kr=64;async function Jt(a,e,r,o){let t=r instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(o<=0)throw new Error("incx must be positive.");if(!t&&!(r instanceof Float64Array))throw new Error("x must be a Float64Array or GpuVector.");if(t&&r.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(e<=0)return{index:0};if(r.length<(e-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let i=["f64/dekker","f64/utils/abs","f64/utils/greater","f64/utils/equal"],s=await G(a,[...i,"idamax"],"idamax_main"),l=await G(a,[...i,"reduction/argmaxF64"],"reduce_f64"),n=null,u=null,f=null,m=null,w=null,c=null,p=null,g=null;try{if(t)n=r._buf,u=r._loBuf;else{let{hi:j,lo:F}=Pr(r);n=v(j,"idamax-xHi",!1),u=v(F,"idamax-xLo",!1)}f=er(2*Kr*4,"idamax-partials-val-hi"),m=er(2*Kr*4,"idamax-partials-val-lo"),w=er(2*Kr*4,"idamax-partials-idx"),c=xr(4,"idamax-result"),p=L([{value:e,type:"u32"},{value:o,type:"u32"}],"idamax-params");let h=B(s.getBindGroupLayout(0),[n,u,f,m,w,p]),{commandEncoder:b,ts:x}=C(s,h,2*Kr);M(b);let _=B(l.getBindGroupLayout(0),[f,m,w,c]),{commandEncoder:y,ts:A}=C(l,_,1);g=N(y,c),M(y);let P=k(g,Uint32Array);g=null;let[E,T,D]=await Promise.all([S(x),S(A),P]),R=D[0];return E!==void 0&&T!==void 0?{index:R,gpuTimeMs:E+T}:{index:R}}finally{!t&&n&&d(n),!t&&u&&d(u),f&&d(f),m&&d(m),w&&d(w),c&&d(c),p&&d(p),g&&d(g)}}async function ro(a,e,r,o,t,i,s,l){let n=r instanceof I,u=t instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(typeof s!="number")throw new Error("c must be a number.");if(typeof l!="number")throw new Error("s must be a number.");if(Number.isNaN(s)||Number.isNaN(l))throw new Error("c and s must not be NaN.");if(!Number.isFinite(s))throw new Error("c must be finite.");if(!Number.isFinite(l))throw new Error("s must be finite.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!n&&!(r instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!u&&!(t instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(n!==u)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(e<=0)return n?{}:{x:r,y:t};if(r.length<(e-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(t.length<(e-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let f=await G(a,"srot"),m=null,w=null,c=null,p=null,g=null;try{m=n?r._buf:v(r,"srot-x",!0),w=u?t._buf:v(t,"srot-y",!0),c=L([{value:e,type:"u32"},{value:s,type:"f32"},{value:l,type:"f32"},{value:o,type:"u32"},{value:i,type:"u32"}],"srot-params");let h=B(f.getBindGroupLayout(0),[m,w,c]),{commandEncoder:b,ts:x}=C(f,h,mr(e));p=n?null:N(b,m),g=u?null:N(b,w),M(b);let _=await S(x);if(n&&u)return _!==void 0?{gpuTimeMs:_}:{};let y=k(p,Float32Array),A=k(g,Float32Array);p=null,g=null;let[P,E]=await Promise.all([y,A]);return _!==void 0?{x:P,y:E,gpuTimeMs:_}:{x:P,y:E}}finally{!n&&m&&d(m),!u&&w&&d(w),c&&d(c),p&&d(p),g&&d(g)}}async function eo(a,e,r,o,t,i,s){let l=r instanceof I,n=t instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(!(s instanceof Float32Array)||s.length!==5)throw new Error("param must be a Float32Array of length 5.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!l&&!(r instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!n&&!(t instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(l!==n)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(e<=0||s[0]===-2)return l?{}:{x:r,y:t};if(r.length<(e-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(t.length<(e-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let u=await G(a,"srotm"),f=null,m=null,w=null,c=null,p=null,g=null;try{f=l?r._buf:v(r,"srotm-x",!0),m=n?t._buf:v(t,"srotm-y",!0),w=v(s,"srotm-param",!1),c=L([{value:e,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"srotm-params");let h=B(u.getBindGroupLayout(0),[f,m,w,c]),{commandEncoder:b,ts:x}=C(u,h,mr(e));p=l?null:N(b,f),g=n?null:N(b,m),M(b);let _=await S(x);if(l&&n)return _!==void 0?{gpuTimeMs:_}:{};let y=k(p,Float32Array),A=k(g,Float32Array);p=null,g=null;let[P,E]=await Promise.all([y,A]);return _!==void 0?{x:P,y:E,gpuTimeMs:_}:{x:P,y:E}}finally{!l&&f&&d(f),!n&&m&&d(m),w&&d(w),c&&d(c),p&&d(p),g&&d(g)}}async function to(a,e,r,o,t,i,s,l,n,u,f,m,w="row-major"){let c=i instanceof H,p=l instanceof I,g=f instanceof I;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="no-transpose"&&e!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(w!=="row-major"&&w!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof t!="number")throw new Error("alpha must be a number.");if(Number.isNaN(t))throw new Error("alpha must not be NaN.");if(!Number.isFinite(t))throw new Error("alpha must be finite.");if(typeof u!="number")throw new Error("beta must be a number.");if(Number.isNaN(u))throw new Error("beta must not be NaN.");if(!Number.isFinite(u))throw new Error("beta must be finite.");if(!Number.isInteger(r)||!Number.isInteger(o)||!Number.isInteger(n)||!Number.isInteger(m)||!Number.isInteger(s))throw new Error("m, n, incx, incy, and lda must be integers.");if(n<=0||m<=0)throw new Error("incx and incy must be positive.");if(!c&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!p&&!(l instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!g&&!(f instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(p!==g)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(p&&!c)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(c&&!p)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(p&&l._buf===f._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(c&&s!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(c&&(i.rows<r||i.cols<o))throw new Error("A is too small for the given m and n.");if(r<0||o<0)throw new Error("m and n must be non-negative.");if(r===0||o===0)return g?{}:{y:f};(c?i.layout:w)==="column-major"&&([r,o]=[o,r],e=e==="no-transpose"?"transpose":"no-transpose");let b=e==="no-transpose",x=b?o:r,_=b?r:o;if(s<o)throw new Error("lda must be >= n.");if(!c&&i.length<(r-1)*s+o)throw new Error("A does not have enough elements for the given m, n, and lda.");if(l.length<(x-1)*n+1)throw new Error("x does not have enough elements for the given dimensions and incx.");if(f.length<(_-1)*m+1)throw new Error("y does not have enough elements for the given dimensions and incy.");let A=await G(a,b?"sgemv_n":"sgemv_t"),P=c?i._buf:v(i,"sgemv-A",!1),E=p?l._buf:v(l,"sgemv-x",!1),T=g?f._buf:v(f,"sgemv-y",!0),D=L([{value:r,type:"u32"},{value:o,type:"u32"},{value:t,type:"f32"},{value:u,type:"f32"},{value:n,type:"u32"},{value:m,type:"u32"},{value:s,type:"u32"}],"sgemv-params");try{let R=B(A.getBindGroupLayout(0),[P,E,T,D]),j=b?Math.min(r,a.limits.maxComputeWorkgroupsPerDimension):mr(_),{commandEncoder:F,ts:V}=C(A,R,j),K=g?null:N(F,T);M(F);let W=await S(V);if(g)return W!==void 0?{gpuTimeMs:W}:{};let $=await k(K,Float32Array);return W!==void 0?{y:$,gpuTimeMs:W}:{y:$}}finally{c||d(P),p||d(E),g||d(T),d(D)}}async function oo(a,e,r,o,t,i,s,l,n,u,f,m="row-major"){let w=s instanceof I,c=u instanceof I,p=t instanceof H;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(m!=="row-major"&&m!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(r)||!Number.isInteger(l)||!Number.isInteger(f)||!Number.isInteger(i))throw new Error("n, incx, incy, and lda must be integers.");if(typeof o!="number")throw new Error("alpha must be a number.");if(Number.isNaN(o))throw new Error("alpha must not be NaN.");if(!Number.isFinite(o))throw new Error("alpha must be finite.");if(typeof n!="number")throw new Error("beta must be a number.");if(Number.isNaN(n))throw new Error("beta must not be NaN.");if(!Number.isFinite(n))throw new Error("beta must be finite.");if(l<=0||f<=0)throw new Error("incx and incy must be positive.");if(i<r)throw new Error("lda must be >= n.");if(!p&&!(t instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!w&&!(s instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!c&&!(u instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(w!==c)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(w&&!p)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(p&&!w)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(w&&s._buf===u._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(p&&i!==t.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(p&&(t.rows<r||t.cols<r))throw new Error("A is too small for the given n.");if(r<0)throw new Error("n must be non-negative.");if(r===0)return c?{}:{y:u};if(!p&&t.length<(r-1)*i+r)throw new Error("A does not have enough elements for the given n and lda.");if(s.length<(r-1)*l+1)throw new Error("x does not have enough elements for the given n and incx.");if(u.length<(r-1)*f+1)throw new Error("y does not have enough elements for the given n and incy.");let h=(p?t.layout:m)==="column-major"?e==="upper":e==="lower",b=await G(a,"ssymv"),x=null,_=null,y=null,A=null;try{x=p?t._buf:v(t,"ssymv-A",!1),_=w?s._buf:v(s,"ssymv-x",!1),y=c?u._buf:v(u,"ssymv-y",!0),A=L([{value:r,type:"u32"},{value:o,type:"f32"},{value:n,type:"f32"},{value:l,type:"u32"},{value:f,type:"u32"},{value:i,type:"u32"},{value:h?0:1,type:"u32"}],"ssymv-params");let P=B(b.getBindGroupLayout(0),[x,_,y,A]),E=Math.min(r,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:T,ts:D}=C(b,P,E),R=c?null:N(T,y);M(T);let j=await S(D);if(c)return j!==void 0?{gpuTimeMs:j}:{};let F=await k(R,Float32Array);return j!==void 0?{y:F,gpuTimeMs:j}:{y:F}}finally{!p&&x&&d(x),!w&&_&&d(_),!c&&y&&d(y),A&&d(A)}}async function ao(a,e,r,o,t,i,s,l,n,u,f,m="row-major"){let w=l instanceof I,c=u instanceof I,p=i instanceof H,g=o==="unit";if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(r!=="no-transpose"&&r!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(!g&&o!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(m!=="row-major"&&m!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(t)||!Number.isInteger(n)||!Number.isInteger(f)||!Number.isInteger(s))throw new Error("n, incx, incy, and lda must be integers.");if(n<=0||f<=0)throw new Error("incx and incy must be positive.");if(s<t)throw new Error("lda must be >= n.");if(!p&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!w&&!(l instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!c&&!(u instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(w!==c)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(w&&l._buf===u._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(w&&!p)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(p&&!w)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(p&&c&&i._buf===u._buf)throw new Error("A and y must not reference the same GPU buffer.");if(p&&s!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(p&&(i.rows<t||i.cols<t))throw new Error("A is too small for the given n.");if(t<0)throw new Error("n must be non-negative.");if(t===0)return c?{}:{y:u};if(!p&&i.length<(t-1)*s+t)throw new Error("A does not have enough elements for the given n and lda.");if(l.length<(t-1)*n+1)throw new Error("x does not have enough elements for the given n and incx.");if(u.length<(t-1)*f+1)throw new Error("y does not have enough elements for the given n and incy.");let b=(p?i.layout:m)==="column-major",x=b?e==="upper":e==="lower",_=b?r==="transpose":r==="no-transpose",y=await G(a,"strmv"),A=null,P=null,E=null,T=null;try{A=p?i._buf:v(i,"strmv-A",!1),P=w?l._buf:v(l,"strmv-x",!1),E=c?u._buf:v(u,"strmv-y",!0),T=L([{value:t,type:"u32"},{value:n,type:"u32"},{value:f,type:"u32"},{value:s,type:"u32"},{value:_?0:1,type:"u32"},{value:x?0:1,type:"u32"},{value:g?1:0,type:"u32"}],"strmv-params");let D=B(y.getBindGroupLayout(0),[A,P,E,T]),R=Math.min(t,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:j,ts:F}=C(y,D,R),V=c?null:N(j,E);M(j);let K=await S(F);if(c)return K!==void 0?{gpuTimeMs:K}:{};let W=await k(V,Float32Array);return K!==void 0?{y:W,gpuTimeMs:K}:{y:W}}finally{!p&&A&&d(A),!w&&P&&d(P),!c&&E&&d(E),T&&d(T)}}var Gr=64;function io(a,e,r){let o=new ArrayBuffer(a*e),t=new DataView(o);for(let i=0;i<a;i++){let s=r(i),l=i*e;s.forEach((n,u)=>t.setUint32(l+u*4,n,!0))}return o}function so(a,e,r){let o=a.createBuffer({label:r,size:e.byteLength,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST});return a.queue.writeBuffer(o,0,e),o}async function no(a,e,r,o,t,i,s,l,n,u="row-major"){let f=l instanceof I,m=i instanceof H,w=o==="unit";if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(r!=="no-transpose"&&r!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(!w&&o!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(u!=="row-major"&&u!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(t)||!Number.isInteger(n)||!Number.isInteger(s))throw new Error("n, incx, and lda must be integers.");if(n<=0)throw new Error("incx must be positive.");if(s<t)throw new Error("lda must be >= n.");if(!m&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!f&&!(l instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(f&&!m)throw new Error("A must be a GpuMatrix when x is a GpuVector.");if(m&&!f)throw new Error("x must be a GpuVector when A is a GpuMatrix.");if(m&&s!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(m&&(i.rows<t||i.cols<t))throw new Error("A is too small for the given n.");if(t<0)throw new Error("n must be non-negative.");if(t===0)return f?{}:{x:l};if(!m&&i.length<(t-1)*s+t)throw new Error("A does not have enough elements for the given n and lda.");if(l.length<(t-1)*n+1)throw new Error("x does not have enough elements for the given n and incx.");let p=(m?i.layout:u)==="column-major",g=p?e==="upper":e==="lower",h=p?r==="transpose":r==="no-transpose",b=await G(a,"strsv_invert_block"),x=await G(a,"strsv_apply_inverse"),_=await G(a,"strsv_update"),y=h===g,A=[];for(let W=0;W<t;W+=Gr)A.push(W);y||A.reverse();let P=A.length,E=a.limits.maxComputeWorkgroupsPerDimension,T=a.limits.minUniformBufferOffsetAlignment,D=null,R=null,j=null,F=null,V=null,K=null;try{D=m?i._buf:v(i,"strsv-A",!1),R=f?l._buf:v(l,"strsv-x",!0),j=er(P*Gr*Gr*4,"strsv-Ainv");let W=io(P,T,q=>{let z=q*Gr,X=Math.min(z+Gr,t);return[n,q,z,X]});F=so(a,W,"strsv-apply-params");let $=io(P,T,q=>{let z=q*Gr,X=Math.min(z+Gr,t);return[t,n,s,h?0:1,g?0:1,z,X]});V=so(a,$,"strsv-update-params");let{commandEncoder:Y,querySet:J}=vr();K=L([{value:t,type:"u32"},{value:s,type:"u32"},{value:h?0:1,type:"u32"},{value:g?0:1,type:"u32"},{value:w?1:0,type:"u32"}],"strsv-invert-params");let nr=B(b.getBindGroupLayout(0),[D,j,K]);ar(Y,b,nr,{x:Gr,y:P},J?{timestampWrites:{querySet:J,beginningOfPassWriteIndex:0}}:void 0);for(let q=0;q<A.length;q++){let z=A[q],X=Math.min(z+Gr,t),Q=z/Gr,tr=q===A.length-1,fr=Q*T,or=B(x.getBindGroupLayout(0),[j,R,{buffer:F,offset:fr,size:16}]);ar(Y,x,or,1,tr&&J?{timestampWrites:{querySet:J,endOfPassWriteIndex:1}}:void 0);let dr=y?t-X:z;if(dr===0)continue;let Br=B(_.getBindGroupLayout(0),[D,R,{buffer:V,offset:fr,size:32}]),yr=Math.min(dr,E);ar(Y,_,Br,yr)}let sr=br(Y,J),Z=f?null:N(Y,R);M(Y);let rr=await S(sr);if(f)return rr!==void 0?{gpuTimeMs:rr}:{};let U=await k(Z,Float32Array);return rr!==void 0?{x:U,gpuTimeMs:rr}:{x:U}}finally{!m&&D&&d(D),!f&&R&&d(R),j&&d(j),F&&d(F),V&&d(V),K&&d(K)}}async function uo(a,e,r,o,t,i,s,l,n,u,f="row-major"){let m=n instanceof H;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(f!=="row-major"&&f!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof o!="number")throw new Error("alpha must be a number.");if(Number.isNaN(o))throw new Error("alpha must not be NaN.");if(!Number.isFinite(o))throw new Error("alpha must be finite.");if(!Number.isInteger(e)||!Number.isInteger(r)||!Number.isInteger(i)||!Number.isInteger(l)||!Number.isInteger(u))throw new Error("m, n, incx, incy, and lda must be integers.");if(i<=0||l<=0)throw new Error("incx and incy must be positive.");if(!m&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(m&&u!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(m&&(n.rows<e||n.cols<r))throw new Error("A is too small for the given m and n.");(m?n.layout:f)==="column-major"&&([e,r]=[r,e],[t,s]=[s,t],[i,l]=[l,i]);let c=t instanceof I,p=s instanceof I;if(u<r)throw new Error("lda must be >= n.");if(!c&&!(t instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!p&&!(s instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(c!==p)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(c&&!m)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(m&&!c)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(m&&c&&n._buf===t._buf)throw new Error("A and x must not reference the same GPU buffer.");if(m&&p&&n._buf===s._buf)throw new Error("A and y must not reference the same GPU buffer.");if(e<0||r<0)throw new Error("m and n must be non-negative.");if(e===0||r===0)return m?{}:{A:n};if(!m&&n.length<(e-1)*u+r)throw new Error("A does not have enough elements for the given m, n, and lda.");if(t.length<(e-1)*i+1)throw new Error("x does not have enough elements for the given m and incx.");if(s.length<(r-1)*l+1)throw new Error("y does not have enough elements for the given n and incy.");let g=await G(a,"sger"),h=null,b=null,x=null,_=null;try{h=c?t._buf:v(t,"sger-x",!1),b=p?s._buf:v(s,"sger-y",!1),x=m?n._buf:v(n,"sger-A",!0),_=L([{value:e,type:"u32"},{value:r,type:"u32"},{value:o,type:"f32"},{value:i,type:"u32"},{value:l,type:"u32"},{value:u,type:"u32"}],"sger-params");let y=B(g.getBindGroupLayout(0),[h,b,x,_]),A=Math.min(e,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:P,ts:E}=C(g,y,A),T=m?null:N(P,x);M(P);let D=await S(E);if(m)return D!==void 0?{gpuTimeMs:D}:{};let R=await k(T,Float32Array);return D!==void 0?{A:R,gpuTimeMs:D}:{A:R}}finally{!c&&h&&d(h),!p&&b&&d(b),!m&&x&&d(x),_&&d(_)}}async function lo(a,e,r,o,t,i,s,l,n="row-major"){let u=t instanceof I,f=s instanceof H;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(n!=="row-major"&&n!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(r)||!Number.isInteger(i)||!Number.isInteger(l))throw new Error("n, incx, and lda must be integers.");if(typeof o!="number")throw new Error("alpha must be a number.");if(Number.isNaN(o))throw new Error("alpha must not be NaN.");if(!Number.isFinite(o))throw new Error("alpha must be finite.");if(i<=0)throw new Error("incx must be positive.");if(l<r)throw new Error("lda must be >= n.");if(!f&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!u&&!(t instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(u&&!f)throw new Error("A must be a GpuMatrix when x is a GpuVector.");if(f&&!u)throw new Error("x must be a GpuVector when A is a GpuMatrix.");if(f&&u&&s._buf===t._buf)throw new Error("A and x must not reference the same GPU buffer.");if(f&&l!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(f&&(s.rows<r||s.cols<r))throw new Error("A is too small for the given n.");if(r<0)throw new Error("n must be non-negative.");if(r===0)return f?{}:{A:s};if(!f&&s.length<(r-1)*l+r)throw new Error("A does not have enough elements for the given n and lda.");if(t.length<(r-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");let w=(f?s.layout:n)==="column-major"?e==="upper":e==="lower",c=await G(a,"ssyr"),p=null,g=null,h=null;try{p=u?t._buf:v(t,"ssyr-x",!1),g=f?s._buf:v(s,"ssyr-A",!0),h=L([{value:r,type:"u32"},{value:o,type:"f32"},{value:i,type:"u32"},{value:l,type:"u32"},{value:w?0:1,type:"u32"}],"ssyr-params");let b=B(c.getBindGroupLayout(0),[p,g,h]),x=Math.min(r,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:_,ts:y}=C(c,b,x),A=f?null:N(_,g);M(_);let P=await S(y);if(f)return P!==void 0?{gpuTimeMs:P}:{};let E=await k(A,Float32Array);return P!==void 0?{A:E,gpuTimeMs:P}:{A:E}}finally{!u&&p&&d(p),!f&&g&&d(g),h&&d(h)}}async function fo(a,e,r,o,t,i,s,l,n,u,f="row-major"){let m=t instanceof I,w=s instanceof I,c=n instanceof H;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(f!=="row-major"&&f!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(r)||!Number.isInteger(i)||!Number.isInteger(l)||!Number.isInteger(u))throw new Error("n, incx, incy, and lda must be integers.");if(typeof o!="number")throw new Error("alpha must be a number.");if(Number.isNaN(o))throw new Error("alpha must not be NaN.");if(!Number.isFinite(o))throw new Error("alpha must be finite.");if(i<=0||l<=0)throw new Error("incx and incy must be positive.");if(u<r)throw new Error("lda must be >= n.");if(!c&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!m&&!(t instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!w&&!(s instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(m!==w)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(m&&!c)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(c&&!m)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(c&&m&&n._buf===t._buf)throw new Error("A and x must not reference the same GPU buffer.");if(c&&w&&n._buf===s._buf)throw new Error("A and y must not reference the same GPU buffer.");if(m&&t._buf===s._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(c&&u!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(c&&(n.rows<r||n.cols<r))throw new Error("A is too small for the given n.");if(r<0)throw new Error("n must be non-negative.");if(r===0)return c?{}:{A:n};if(!c&&n.length<(r-1)*u+r)throw new Error("A does not have enough elements for the given n and lda.");if(t.length<(r-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");if(s.length<(r-1)*l+1)throw new Error("y does not have enough elements for the given n and incy.");let g=(c?n.layout:f)==="column-major"?e==="upper":e==="lower",h=await G(a,"ssyr2"),b=null,x=null,_=null,y=null;try{b=m?t._buf:v(t,"ssyr2-x",!1),x=w?s._buf:v(s,"ssyr2-y",!1),_=c?n._buf:v(n,"ssyr2-A",!0),y=L([{value:r,type:"u32"},{value:o,type:"f32"},{value:i,type:"u32"},{value:l,type:"u32"},{value:u,type:"u32"},{value:g?0:1,type:"u32"}],"ssyr2-params");let A=B(h.getBindGroupLayout(0),[b,x,_,y]),P=Math.min(r,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:E,ts:T}=C(h,A,P),D=c?null:N(E,_);M(E);let R=await S(T);if(c)return R!==void 0?{gpuTimeMs:R}:{};let j=await k(D,Float32Array);return R!==void 0?{A:j,gpuTimeMs:R}:{A:j}}finally{!m&&b&&d(b),!w&&x&&d(x),!c&&_&&d(_),y&&d(y)}}var Ma=32,Ia=32,La=64,Da=64,Ra=36;async function mo(a,e,r,o,t,i,s,l,n,u,f,m,w,c,p="row-major"){let g=l instanceof H,h=u instanceof H,b=w instanceof H;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="no-transpose"&&e!=="transpose")throw new Error("transA must be 'no-transpose' or 'transpose'.");if(r!=="no-transpose"&&r!=="transpose")throw new Error("transB must be 'no-transpose' or 'transpose'.");if(p!=="row-major"&&p!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof s!="number")throw new Error("alpha must be a number.");if(Number.isNaN(s))throw new Error("alpha must not be NaN.");if(!Number.isFinite(s))throw new Error("alpha must be finite.");if(typeof m!="number")throw new Error("beta must be a number.");if(Number.isNaN(m))throw new Error("beta must not be NaN.");if(!Number.isFinite(m))throw new Error("beta must be finite.");if(!Number.isInteger(o)||!Number.isInteger(t)||!Number.isInteger(i)||!Number.isInteger(n)||!Number.isInteger(f)||!Number.isInteger(c))throw new Error("m, n, k, lda, ldb, and ldc must be integers.");if(!g&&!(l instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!h&&!(u instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(!b&&!(w instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if((g||h)&&!b)throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");if(b&&(!g||!h))throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");if(o<0||t<0||i<0)throw new Error("m, n, and k must be non-negative.");if(o===0||t===0)return b?{}:{C:w};let x=g?l.layout:p,_=h?u.layout:p,y=b?w.layout:p,A=x==="column-major"?i:o,P=x==="column-major"?o:i,E=e==="no-transpose"?A:P,T=e==="no-transpose"?P:A;if(n<T)throw new Error(`lda must be >= ${x==="column-major"?"rows":"cols"} of A as stored.`);if(g){if(n!==l.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[rr,U]=e==="no-transpose"?[o,i]:[i,o];if(l.rows<rr||l.cols<U)throw new Error("A is too small for the given m, k, and transA.")}else if(l.length<(E-1)*n+T)throw new Error("A does not have enough elements for the given dimensions and lda.");let D=_==="column-major"?t:i,R=_==="column-major"?i:t,j=r==="no-transpose"?D:R,F=r==="no-transpose"?R:D;if(f<F)throw new Error(`ldb must be >= ${_==="column-major"?"rows":"cols"} of B as stored.`);if(h){if(f!==u.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");let[rr,U]=r==="no-transpose"?[i,t]:[t,i];if(u.rows<rr||u.cols<U)throw new Error("B is too small for the given n, k, and transB.")}else if(u.length<(j-1)*f+F)throw new Error("B does not have enough elements for the given dimensions and ldb.");let V=y==="column-major"?t:o,K=y==="column-major"?o:t;if(c<K)throw new Error(`ldc must be >= ${y==="column-major"?"rows":"cols"} of C as stored.`);if(b){if(c!==w.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(w.rows<o||w.cols<t)throw new Error("C is too small for the given m and n.")}else if(w.length<(V-1)*c+K)throw new Error("C does not have enough elements for the given dimensions and ldc.");x==="column-major"&&(e=e==="no-transpose"?"transpose":"no-transpose"),_==="column-major"&&(r=r==="no-transpose"?"transpose":"no-transpose"),y==="column-major"&&([l,u]=[u,l],[g,h]=[h,g],[n,f]=[f,n],[e,r]=[r==="no-transpose"?"transpose":"no-transpose",e==="no-transpose"?"transpose":"no-transpose"],[o,t]=[t,o]);let W=Math.ceil(t/Da),$=Math.ceil(o/La),Y=W*$>=Ra,J=await G(a,Y?"sgemm_large":"sgemm_small"),nr=g?l._buf:v(l,"sgemm-A",!1),ur=h?u._buf:v(u,"sgemm-B",!1),sr=b?w._buf:v(w,"sgemm-C",!0),Z=L([{value:o,type:"u32"},{value:t,type:"u32"},{value:i,type:"u32"},{value:s,type:"f32"},{value:m,type:"f32"},{value:n,type:"u32"},{value:f,type:"u32"},{value:c,type:"u32"},{value:e==="transpose"?1:0,type:"u32"},{value:r==="transpose"?1:0,type:"u32"}],"sgemm-params");try{let rr=B(J.getBindGroupLayout(0),[nr,ur,sr,Z]),U=Y?{x:Math.min(W,a.limits.maxComputeWorkgroupsPerDimension),y:Math.min($,a.limits.maxComputeWorkgroupsPerDimension)}:{x:Math.min(Math.ceil(t/Ia),a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(Math.ceil(o/Ma),a.limits.maxComputeWorkgroupsPerDimension)},{commandEncoder:q,ts:z}=C(J,rr,U),X=b?null:N(q,sr);M(q);let Q=await S(z);if(b)return Q!==void 0?{gpuTimeMs:Q}:{};let tr=await k(X,Float32Array);return Q!==void 0?{C:tr,gpuTimeMs:Q}:{C:tr}}finally{g||d(nr),h||d(ur),b||d(sr),d(Z)}}var Ta=32,ja=32,Ca=64,Wa=64,Fa=36;async function co(a,e,r,o,t,i,s,l,n,u,f,m,w,c,p,g="row-major"){let h=n instanceof H,b=f instanceof H,x=c instanceof H;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(r!=="no-transpose"&&r!=="transpose")throw new Error("transA must be 'no-transpose' or 'transpose'.");if(o!=="no-transpose"&&o!=="transpose")throw new Error("transB must be 'no-transpose' or 'transpose'.");if(g!=="row-major"&&g!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof l!="number")throw new Error("alpha must be a number.");if(Number.isNaN(l))throw new Error("alpha must not be NaN.");if(!Number.isFinite(l))throw new Error("alpha must be finite.");if(typeof w!="number")throw new Error("beta must be a number.");if(Number.isNaN(w))throw new Error("beta must not be NaN.");if(!Number.isFinite(w))throw new Error("beta must be finite.");if(!Number.isInteger(t)||!Number.isInteger(i)||!Number.isInteger(s)||!Number.isInteger(u)||!Number.isInteger(m)||!Number.isInteger(p))throw new Error("m, n, k, lda, ldb, and ldc must be integers.");if(!h&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!b&&!(f instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(!x&&!(c instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if((h||b)&&!x)throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");if(x&&(!h||!b))throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");if(t<0||i<0||s<0)throw new Error("m, n, and k must be non-negative.");if(t===0||i===0)return x?{}:{C:c};let _=h?n.layout:g,y=b?f.layout:g,A=x?c.layout:g,P=_==="column-major"?s:t,E=_==="column-major"?t:s,T=r==="no-transpose"?P:E,D=r==="no-transpose"?E:P;if(u<D)throw new Error(`lda must be >= ${_==="column-major"?"rows":"cols"} of A as stored.`);if(h){if(u!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[U,q]=r==="no-transpose"?[t,s]:[s,t];if(n.rows<U||n.cols<q)throw new Error("A is too small for the given m, k, and transA.")}else if(n.length<(T-1)*u+D)throw new Error("A does not have enough elements for the given dimensions and lda.");let R=y==="column-major"?i:s,j=y==="column-major"?s:i,F=o==="no-transpose"?R:j,V=o==="no-transpose"?j:R;if(m<V)throw new Error(`ldb must be >= ${y==="column-major"?"rows":"cols"} of B as stored.`);if(b){if(m!==f.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");let[U,q]=o==="no-transpose"?[s,i]:[i,s];if(f.rows<U||f.cols<q)throw new Error("B is too small for the given n, k, and transB.")}else if(f.length<(F-1)*m+V)throw new Error("B does not have enough elements for the given dimensions and ldb.");let K=A==="column-major"?i:t,W=A==="column-major"?t:i;if(p<W)throw new Error(`ldc must be >= ${A==="column-major"?"rows":"cols"} of C as stored.`);if(x){if(p!==c.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(c.rows<t||c.cols<i)throw new Error("C is too small for the given m and n.")}else if(c.length<(K-1)*p+W)throw new Error("C does not have enough elements for the given dimensions and ldc.");_==="column-major"&&(r=r==="no-transpose"?"transpose":"no-transpose"),y==="column-major"&&(o=o==="no-transpose"?"transpose":"no-transpose"),A==="column-major"&&([n,f]=[f,n],[h,b]=[b,h],[u,m]=[m,u],[r,o]=[o==="no-transpose"?"transpose":"no-transpose",r==="no-transpose"?"transpose":"no-transpose"],[t,i]=[i,t],e=e==="lower"?"upper":"lower");let $=Math.ceil(i/Wa),Y=Math.ceil(t/Ca),J=$*Y>=Fa,nr=await G(a,J?"sgemmtr_large":"sgemmtr_small"),ur=h?n._buf:v(n,"sgemmtr-A",!1),sr=b?f._buf:v(f,"sgemmtr-B",!1),Z=x?c._buf:v(c,"sgemmtr-C",!0),rr=L([{value:t,type:"u32"},{value:i,type:"u32"},{value:s,type:"u32"},{value:l,type:"f32"},{value:w,type:"f32"},{value:u,type:"u32"},{value:m,type:"u32"},{value:p,type:"u32"},{value:r==="transpose"?1:0,type:"u32"},{value:o==="transpose"?1:0,type:"u32"},{value:e==="upper"?1:0,type:"u32"}],"sgemmtr-params");try{let U=B(nr.getBindGroupLayout(0),[ur,sr,Z,rr]),q=J?{x:Math.min($,a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(Y,a.limits.maxComputeWorkgroupsPerDimension)}:{x:Math.min(Math.ceil(i/ja),a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(Math.ceil(t/Ta),a.limits.maxComputeWorkgroupsPerDimension)},{commandEncoder:z,ts:X}=C(nr,U,q),Q=x?null:N(z,Z);M(z);let tr=await S(X);if(x)return tr!==void 0?{gpuTimeMs:tr}:{};let fr=await k(Q,Float32Array);return tr!==void 0?{C:fr,gpuTimeMs:tr}:{C:fr}}finally{h||d(ur),b||d(sr),x||d(Z),d(rr)}}var Ha=32,Ua=32,Oa=64,Va=64,Ka=36;async function po(a,e,r,o,t,i,s,l,n,u,f,m="row-major"){let w=s instanceof H,c=u instanceof H;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(r!=="no-transpose"&&r!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(m!=="row-major"&&m!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof i!="number")throw new Error("alpha must be a number.");if(Number.isNaN(i))throw new Error("alpha must not be NaN.");if(!Number.isFinite(i))throw new Error("alpha must be finite.");if(typeof n!="number")throw new Error("beta must be a number.");if(Number.isNaN(n))throw new Error("beta must not be NaN.");if(!Number.isFinite(n))throw new Error("beta must be finite.");if(!Number.isInteger(o)||!Number.isInteger(t)||!Number.isInteger(l)||!Number.isInteger(f))throw new Error("n, k, lda, and ldc must be integers.");if(!w&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!c&&!(u instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if(w&&!c)throw new Error("C must be a GpuMatrix when A is a GpuMatrix.");if(c&&!w)throw new Error("A must be a GpuMatrix when C is a GpuMatrix.");if(o<0||t<0)throw new Error("n and k must be non-negative.");if(o===0)return c?{}:{C:u};let p=w?s.layout:m,g=c?u.layout:m,h=p==="column-major"?t:o,b=p==="column-major"?o:t,x=r==="no-transpose"?h:b,_=r==="no-transpose"?b:h;if(l<_)throw new Error(`lda must be >= ${p==="column-major"?"rows":"cols"} of A as stored.`);if(w){if(l!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[W,$]=r==="no-transpose"?[o,t]:[t,o];if(s.rows<W||s.cols<$)throw new Error("A is too small for the given n, k, and trans.")}else if(s.length<(x-1)*l+_)throw new Error("A does not have enough elements for the given dimensions and lda.");if(f<o)throw new Error("ldc must be >= n.");if(c){if(f!==u.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(u.rows<o||u.cols<o)throw new Error("C is too small for the given n.")}else if(u.length<(o-1)*f+o)throw new Error("C does not have enough elements for the given dimensions and ldc.");let y=r;p==="column-major"&&(y=y==="no-transpose"?"transpose":"no-transpose");let A=y==="no-transpose"?"transpose":"no-transpose",P=e;g==="column-major"&&([y,A]=[A==="no-transpose"?"transpose":"no-transpose",y==="no-transpose"?"transpose":"no-transpose"],P=P==="lower"?"upper":"lower");let E=Math.ceil(o/Va),T=Math.ceil(o/Oa),D=E*T>=Ka,R=await G(a,D?"sgemmtr_large":"sgemmtr_small"),j=w?s._buf:v(s,"ssyrk-A",!1),F=c?u._buf:v(u,"ssyrk-C",!0),V=w?er(j.size,"ssyrk-B",GPUBufferUsage.COPY_DST):v(s,"ssyrk-B",!1),K=L([{value:o,type:"u32"},{value:o,type:"u32"},{value:t,type:"u32"},{value:i,type:"f32"},{value:n,type:"f32"},{value:l,type:"u32"},{value:l,type:"u32"},{value:f,type:"u32"},{value:y==="transpose"?1:0,type:"u32"},{value:A==="transpose"?1:0,type:"u32"},{value:P==="upper"?1:0,type:"u32"}],"ssyrk-params");try{let W=B(R.getBindGroupLayout(0),[j,V,F,K]),$=D?{x:Math.min(E,a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(T,a.limits.maxComputeWorkgroupsPerDimension)}:{x:Math.min(Math.ceil(o/Ua),a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(Math.ceil(o/Ha),a.limits.maxComputeWorkgroupsPerDimension)},{commandEncoder:Y,querySet:J,passDescriptor:nr}=vr();w&&Y.copyBufferToBuffer(j,0,V,0,j.size),ar(Y,R,W,$,nr);let ur=br(Y,J),sr=c?null:N(Y,F);M(Y);let Z=await S(ur);if(c)return Z!==void 0?{gpuTimeMs:Z}:{};let rr=await k(sr,Float32Array);return Z!==void 0?{C:rr,gpuTimeMs:Z}:{C:rr}}finally{w||d(j),d(V),c||d(F),d(K)}}var za=32,qa=32,Ya=64,Xa=64,Qa=36;async function wo(a,e,r,o,t,i,s,l,n,u,f,m,w,c="row-major"){let p=s instanceof H,g=n instanceof H,h=m instanceof H;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(r!=="no-transpose"&&r!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(c!=="row-major"&&c!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof i!="number")throw new Error("alpha must be a number.");if(Number.isNaN(i))throw new Error("alpha must not be NaN.");if(!Number.isFinite(i))throw new Error("alpha must be finite.");if(typeof f!="number")throw new Error("beta must be a number.");if(Number.isNaN(f))throw new Error("beta must not be NaN.");if(!Number.isFinite(f))throw new Error("beta must be finite.");if(!Number.isInteger(o)||!Number.isInteger(t)||!Number.isInteger(l)||!Number.isInteger(u)||!Number.isInteger(w))throw new Error("n, k, lda, ldb, and ldc must be integers.");if(!p&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!g&&!(n instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(!h&&!(m instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if((p||g)&&!h)throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");if(h&&(!p||!g))throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");if(o<0||t<0)throw new Error("n and k must be non-negative.");if(o===0)return h?{}:{C:m};let b=p?s.layout:c,x=g?n.layout:c,_=h?m.layout:c,y=b==="column-major"?t:o,A=b==="column-major"?o:t,P=r==="no-transpose"?y:A,E=r==="no-transpose"?A:y;if(l<E)throw new Error(`lda must be >= ${b==="column-major"?"rows":"cols"} of A as stored.`);if(p){if(l!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[X,Q]=r==="no-transpose"?[o,t]:[t,o];if(s.rows<X||s.cols<Q)throw new Error("A is too small for the given n, k, and trans.")}else if(s.length<(P-1)*l+E)throw new Error("A does not have enough elements for the given dimensions and lda.");let T=x==="column-major"?t:o,D=x==="column-major"?o:t,R=r==="no-transpose"?T:D,j=r==="no-transpose"?D:T;if(u<j)throw new Error(`ldb must be >= ${x==="column-major"?"rows":"cols"} of B as stored.`);if(g){if(u!==n.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");let[X,Q]=r==="no-transpose"?[o,t]:[t,o];if(n.rows<X||n.cols<Q)throw new Error("B is too small for the given n, k, and trans.")}else if(n.length<(R-1)*u+j)throw new Error("B does not have enough elements for the given dimensions and ldb.");if(w<o)throw new Error("ldc must be >= n.");if(h){if(w!==m.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(m.rows<o||m.cols<o)throw new Error("C is too small for the given n.")}else if(m.length<(o-1)*w+o)throw new Error("C does not have enough elements for the given dimensions and ldc.");let F=r;b==="column-major"&&(F=F==="no-transpose"?"transpose":"no-transpose");let V=r;x==="column-major"&&(V=V==="no-transpose"?"transpose":"no-transpose");let K=_==="column-major"?e==="lower"?"upper":"lower":e,W=X=>X==="no-transpose"?"transpose":"no-transpose";function $(X,Q,tr,fr,or,ir){let dr=X,Br=W(fr);return _!=="column-major"?{transX:dr,X:Q,ldX:tr,transY:Br,Y:or,ldY:ir}:{transX:W(Br),X:or,ldX:ir,transY:W(dr),Y:Q,ldY:tr}}let Y=Math.ceil(o/Xa),J=Math.ceil(o/Ya),nr=Y*J>=Qa,ur=await G(a,nr?"sgemmtr_large":"sgemmtr_small"),sr=nr?{x:Math.min(Y,a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(J,a.limits.maxComputeWorkgroupsPerDimension)}:{x:Math.min(Math.ceil(o/qa),a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(Math.ceil(o/za),a.limits.maxComputeWorkgroupsPerDimension)},Z=p?s._buf:v(s,"ssyr2k-A",!1),rr=g?n._buf:v(n,"ssyr2k-B",!1),U=h?m._buf:v(m,"ssyr2k-C",!0),q=null,z=null;try{let X=$(F,Z,l,V,rr,u),Q=$(V,rr,u,F,Z,l),tr=(Er,wr)=>L([{value:o,type:"u32"},{value:o,type:"u32"},{value:t,type:"u32"},{value:i,type:"f32"},{value:wr,type:"f32"},{value:Er.ldX,type:"u32"},{value:Er.ldY,type:"u32"},{value:w,type:"u32"},{value:Er.transX==="transpose"?1:0,type:"u32"},{value:Er.transY==="transpose"?1:0,type:"u32"},{value:K==="upper"?1:0,type:"u32"}],"ssyr2k-params");q=tr(X,f),z=tr(Q,1);let fr=B(ur.getBindGroupLayout(0),[X.X,X.Y,U,q]),or=B(ur.getBindGroupLayout(0),[Q.X,Q.Y,U,z]),{commandEncoder:ir,querySet:dr}=vr(),Br=dr?{timestampWrites:{querySet:dr,beginningOfPassWriteIndex:0}}:void 0,yr=dr?{timestampWrites:{querySet:dr,endOfPassWriteIndex:1}}:void 0;ar(ir,ur,fr,sr,Br),ar(ir,ur,or,sr,yr);let _r=br(ir,dr),gr=h?null:N(ir,U);M(ir);let pr=await S(_r);if(h)return pr!==void 0?{gpuTimeMs:pr}:{};let cr=await k(gr,Float32Array);return pr!==void 0?{C:cr,gpuTimeMs:pr}:{C:cr}}finally{p||d(Z),g||d(rr),h||d(U),q&&d(q),z&&d(z)}}var Za=32,$a=32,Ja=64,ri=64,ei=36,go=8;async function bo(a,e,r,o,t,i,s,l,n,u,f,m,w,c="row-major"){let p=s instanceof H,g=n instanceof H,h=m instanceof H;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="left"&&e!=="right")throw new Error("side must be 'left' or 'right'.");if(r!=="lower"&&r!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(c!=="row-major"&&c!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof i!="number")throw new Error("alpha must be a number.");if(Number.isNaN(i))throw new Error("alpha must not be NaN.");if(!Number.isFinite(i))throw new Error("alpha must be finite.");if(typeof f!="number")throw new Error("beta must be a number.");if(Number.isNaN(f))throw new Error("beta must not be NaN.");if(!Number.isFinite(f))throw new Error("beta must be finite.");if(!Number.isInteger(o)||!Number.isInteger(t)||!Number.isInteger(l)||!Number.isInteger(u)||!Number.isInteger(w))throw new Error("m, n, lda, ldb, and ldc must be integers.");if(!p&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!g&&!(n instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(!h&&!(m instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if((p||g)&&!h)throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");if(h&&(!p||!g))throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");if(o<0||t<0)throw new Error("m and n must be non-negative.");if(o===0||t===0)return h?{}:{C:m};let b=p?s.layout:c,x=g?n.layout:c,_=h?m.layout:c,y=e==="left"?o:t;if(l<y)throw new Error("lda must be >= "+(e==="left"?"m":"n")+".");if(p){if(l!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(s.rows<y||s.cols<y)throw new Error("A is too small for the given m/n and side.")}else if(s.length<(y-1)*l+y)throw new Error("A does not have enough elements for the given dimensions and lda.");let A=x==="column-major"?t:o,P=x==="column-major"?o:t;if(u<P)throw new Error(`ldb must be >= ${x==="column-major"?"rows":"cols"} of B as stored.`);if(g){if(u!==n.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");if(n.rows<o||n.cols<t)throw new Error("B is too small for the given m and n.")}else if(n.length<(A-1)*u+P)throw new Error("B does not have enough elements for the given dimensions and ldb.");let E=_==="column-major"?t:o,T=_==="column-major"?o:t;if(w<T)throw new Error(`ldc must be >= ${_==="column-major"?"rows":"cols"} of C as stored.`);if(h){if(w!==m.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(m.rows<o||m.cols<t)throw new Error("C is too small for the given m and n.")}else if(m.length<(E-1)*w+T)throw new Error("C does not have enough elements for the given dimensions and ldc.");let D=b==="column-major"?r==="lower"?"upper":"lower":r,R=x==="column-major"?"transpose":"no-transpose",j="no-transpose",F=o,V=t,K=y,W=e==="left"?j:R,$=e==="left"?R:j,Y=ir=>ir==="no-transpose"?"transpose":"no-transpose",J=e==="right";_==="column-major"&&([W,$]=[Y($),Y(W)],J=!J,[F,V]=[V,F]);let nr=y,ur=Math.ceil(V/ri),sr=Math.ceil(F/Ja),Z=ur*sr>=ei,rr=await G(a,Z?"sgemm_large":"sgemm_small"),U=await G(a,"symmetrize"),q=Z?{x:Math.min(ur,a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(sr,a.limits.maxComputeWorkgroupsPerDimension)}:{x:Math.min(Math.ceil(V/$a),a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(Math.ceil(F/Za),a.limits.maxComputeWorkgroupsPerDimension)},z=p?s._buf:v(s,"ssymm-A",!1),X=g?n._buf:v(n,"ssymm-B",!1),Q=h?m._buf:v(m,"ssymm-C",!0),tr=er(y*nr*4,"ssymm-Adense"),fr=null,or=null;try{fr=L([{value:y,type:"u32"},{value:l,type:"u32"},{value:nr,type:"u32"},{value:D==="upper"?1:0,type:"u32"}],"ssymm-sym-params");let ir=B(U.getBindGroupLayout(0),[z,tr,fr]),dr=J?X:tr,Br=J?u:nr,yr=J?tr:X;or=L([{value:F,type:"u32"},{value:V,type:"u32"},{value:K,type:"u32"},{value:i,type:"f32"},{value:f,type:"f32"},{value:Br,type:"u32"},{value:J?nr:u,type:"u32"},{value:w,type:"u32"},{value:W==="transpose"?1:0,type:"u32"},{value:$==="transpose"?1:0,type:"u32"}],"ssymm-gemm-params");let gr=B(rr.getBindGroupLayout(0),[dr,yr,Q,or]),{commandEncoder:pr,querySet:cr}=vr(),Er=cr?{timestampWrites:{querySet:cr,beginningOfPassWriteIndex:0}}:void 0,wr=cr?{timestampWrites:{querySet:cr,endOfPassWriteIndex:1}}:void 0;ar(pr,U,ir,{x:Math.ceil(y/go),y:Math.ceil(y/go)},Er),ar(pr,rr,gr,q,wr);let kr=br(pr,cr),Nr=h?null:N(pr,Q);M(pr);let Lr=await S(kr);if(h)return Lr!==void 0?{gpuTimeMs:Lr}:{};let Fr=await k(Nr,Float32Array);return Lr!==void 0?{C:Fr,gpuTimeMs:Lr}:{C:Fr}}finally{p||d(z),g||d(X),h||d(Q),d(tr),fr&&d(fr),or&&d(or)}}var ti=32,oi=32,ai=64,ii=64,si=36,ho=8;async function xo(a,e,r,o,t,i,s,l,n,u,f,m,w="row-major"){let c=n instanceof H,p=f instanceof H,g=t==="unit";if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="left"&&e!=="right")throw new Error("side must be 'left' or 'right'.");if(r!=="lower"&&r!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(o!=="no-transpose"&&o!=="transpose")throw new Error("transA must be 'no-transpose' or 'transpose'.");if(!g&&t!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(w!=="row-major"&&w!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof l!="number")throw new Error("alpha must be a number.");if(Number.isNaN(l))throw new Error("alpha must not be NaN.");if(!Number.isFinite(l))throw new Error("alpha must be finite.");if(!Number.isInteger(i)||!Number.isInteger(s)||!Number.isInteger(u)||!Number.isInteger(m))throw new Error("m, n, lda, and ldb must be integers.");if(!c&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!p&&!(f instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(c!==p)throw new Error("A and B must both be GpuMatrix or both be Float32Array.");if(i<0||s<0)throw new Error("m and n must be non-negative.");if(i===0||s===0)return p?{}:{B:f};let h=c?n.layout:w,b=p?f.layout:w,x=e==="left"?i:s;if(u<x)throw new Error("lda must be >= "+(e==="left"?"m":"n")+".");if(c){if(u!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(n.rows<x||n.cols<x)throw new Error("A is too small for the given m/n and side.")}else if(n.length<(x-1)*u+x)throw new Error("A does not have enough elements for the given dimensions and lda.");let _=b==="column-major"?s:i,y=b==="column-major"?i:s;if(m<y)throw new Error(`ldb must be >= ${b==="column-major"?"rows":"cols"} of B as stored.`);if(p){if(m!==f.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");if(f.rows<i||f.cols<s)throw new Error("B is too small for the given m and n.")}else if(f.length<(_-1)*m+y)throw new Error("B does not have enough elements for the given dimensions and ldb.");let A=h==="column-major"?r==="lower"?"upper":"lower":r,P=h==="column-major"?o==="no-transpose"?"transpose":"no-transpose":o,E=b==="column-major"?"transpose":"no-transpose",T="no-transpose",D=i,R=s,j=x,F=e==="left"?T:E,V=e==="left"?E:T,K=fr=>fr==="no-transpose"?"transpose":"no-transpose",W=e==="right";b==="column-major"&&([F,V]=[K(V),K(F)],W=!W,[D,R]=[R,D]);let $=x,Y=Math.ceil(R/ii),J=Math.ceil(D/ai),nr=Y*J>=si,ur=await G(a,nr?"sgemm_large":"sgemm_small"),sr=await G(a,"triangularize"),Z=nr?{x:Math.min(Y,a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(J,a.limits.maxComputeWorkgroupsPerDimension)}:{x:Math.min(Math.ceil(R/oi),a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(Math.ceil(D/ti),a.limits.maxComputeWorkgroupsPerDimension)},rr=c?n._buf:v(n,"strmm-A",!1),U=p?f._buf:v(f,"strmm-B",!0),q=er(x*$*4,"strmm-Adense"),z=er(_*m*4,"strmm-out",GPUBufferUsage.COPY_SRC|GPUBufferUsage.COPY_DST),X=null,Q=null,tr=!1;try{X=L([{value:x,type:"u32"},{value:u,type:"u32"},{value:$,type:"u32"},{value:A==="upper"?1:0,type:"u32"},{value:P==="transpose"?1:0,type:"u32"},{value:g?1:0,type:"u32"}],"strmm-tri-params");let fr=B(sr.getBindGroupLayout(0),[rr,q,X]),or=W?U:q,ir=W?m:$,dr=W?q:U;Q=L([{value:D,type:"u32"},{value:R,type:"u32"},{value:j,type:"u32"},{value:l,type:"f32"},{value:0,type:"f32"},{value:ir,type:"u32"},{value:W?$:m,type:"u32"},{value:m,type:"u32"},{value:F==="transpose"?1:0,type:"u32"},{value:V==="transpose"?1:0,type:"u32"}],"strmm-gemm-params");let yr=B(ur.getBindGroupLayout(0),[or,dr,z,Q]),{commandEncoder:_r,querySet:gr}=vr();_r.copyBufferToBuffer(U,0,z,0,Math.min(U.size,z.size));let pr=gr?{timestampWrites:{querySet:gr,beginningOfPassWriteIndex:0}}:void 0,cr=gr?{timestampWrites:{querySet:gr,endOfPassWriteIndex:1}}:void 0;ar(_r,sr,fr,{x:Math.ceil(x/ho),y:Math.ceil(x/ho)},pr),ar(_r,ur,yr,Z,cr);let Er=br(_r,gr),wr=p?null:N(_r,z);M(_r);let kr=await S(Er);if(p)return d(f._buf),f._buf=z,tr=!0,kr!==void 0?{gpuTimeMs:kr}:{};let Nr=await k(wr,Float32Array);return kr!==void 0?{B:Nr,gpuTimeMs:kr}:{B:Nr}}finally{c||d(rr),p||d(U),d(q),tr||d(z),X&&d(X),Q&&d(Q)}}var hr=64,vo=32,yo=32,_o=64,Bo=64,Eo=36;async function Ao(a,e,r,o,t,i,s,l,n,u,f,m,w="row-major"){let c=n instanceof H,p=f instanceof H,g=t==="unit";if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(e!=="left"&&e!=="right")throw new Error("side must be 'left' or 'right'.");if(r!=="lower"&&r!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(o!=="no-transpose"&&o!=="transpose")throw new Error("transA must be 'no-transpose' or 'transpose'.");if(!g&&t!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(w!=="row-major"&&w!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof l!="number")throw new Error("alpha must be a number.");if(Number.isNaN(l))throw new Error("alpha must not be NaN.");if(!Number.isFinite(l))throw new Error("alpha must be finite.");if(!Number.isInteger(i)||!Number.isInteger(s)||!Number.isInteger(u)||!Number.isInteger(m))throw new Error("m, n, lda, and ldb must be integers.");if(!c&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!p&&!(f instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(c!==p)throw new Error("A and B must both be GpuMatrix or both be Float32Array.");if(i<0||s<0)throw new Error("m and n must be non-negative.");if(i===0||s===0)return p?{}:{B:f};let h=c?n.layout:w,b=p?f.layout:w,x=e==="left"?i:s;if(u<x)throw new Error("lda must be >= "+(e==="left"?"m":"n")+".");if(c){if(u!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(n.rows<x||n.cols<x)throw new Error("A is too small for the given m/n and side.")}else if(n.length<(x-1)*u+x)throw new Error("A does not have enough elements for the given dimensions and lda.");let _=b==="column-major"?s:i,y=b==="column-major"?i:s;if(m<y)throw new Error(`ldb must be >= ${b==="column-major"?"rows":"cols"} of B as stored.`);if(p){if(m!==f.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");if(f.rows<i||f.cols<s)throw new Error("B is too small for the given m and n.")}else if(f.length<(_-1)*m+y)throw new Error("B does not have enough elements for the given dimensions and ldb.");let A=h==="column-major"?r==="lower"?"upper":"lower":r,P=h==="column-major"?o==="no-transpose"?"transpose":"no-transpose":o,E=e==="left"?s:i,T=e==="left",D=P==="no-transpose"==(A==="lower"),R=e==="left"?D:!D,j=[];for(let U=0;U<x;U+=hr)j.push(U);R||j.reverse();let F=j.length,V=await G(a,"strsv_invert_block"),K=await G(a,"block_transfer"),W=await G(a,"sscal"),$=c?n._buf:v(n,"strsm-A",!1),Y=p?f._buf:v(f,"strsm-B",!0),J=er(F*hr*hr*4,"strsm-Ainv"),nr=[],ur=[];function sr(U,q){let z=er(U,q);return ur.push(z),z}function Z(U,q){let z=L(U,q);return nr.push(z),z}let rr=(_-1)*m+y;try{let U=null;if(l!==1){let gr=Z([{value:rr,type:"u32"},{value:l,type:"f32"},{value:1,type:"u32"}],"strsm-scale-params");U=B(W.getBindGroupLayout(0),[Y,gr])}let q=Z([{value:x,type:"u32"},{value:u,type:"u32"},{value:P==="transpose"?1:0,type:"u32"},{value:A==="upper"?1:0,type:"u32"},{value:g?1:0,type:"u32"}],"strsm-invert-params"),z=B(V.getBindGroupLayout(0),[$,J,q]),X=sr(hr*E*4,"strsm-Bblock"),Q=sr(hr*E*4,"strsm-Xblock"),tr=sr(x*hr*4,"strsm-Aoff"),fr=sr(x*E*4,"strsm-delta"),{commandEncoder:or,querySet:ir}=vr();if(l===0){let gr=ir?{timestampWrites:{querySet:ir,beginningOfPassWriteIndex:0,endOfPassWriteIndex:1}}:void 0;ar(or,W,U,mr(rr),gr)}else{U&&ar(or,W,U,mr(rr)),ar(or,V,z,{x:hr,y:F},ir?{timestampWrites:{querySet:ir,beginningOfPassWriteIndex:0}}:void 0);for(let pr=0;pr<j.length;pr++){let cr=j[pr],Er=Math.min(cr+hr,x),wr=Er-cr,kr=cr/hr,Nr=pr===j.length-1,Lr=Z([{value:cr,type:"u32"},{value:wr,type:"u32"},{value:0,type:"u32"},{value:E,type:"u32"},{value:m,type:"u32"},{value:b==="column-major"?1:0,type:"u32"},{value:T?1:0,type:"u32"},{value:2,type:"u32"}],"strsm-gather-B-params"),Fr=B(K.getBindGroupLayout(0),[X,Y,Lr]);ar(or,K,Fr,mr(wr,E));{let Sr=wr,Mr=E,zr=wr,Tr=Math.ceil(Mr/Bo),jr=Math.ceil(Sr/_o),Cr=Tr*jr>=Eo,Wr=await G(a,Cr?"sgemm_large":"sgemm_small"),qr=Z([{value:Sr,type:"u32"},{value:Mr,type:"u32"},{value:zr,type:"u32"},{value:1,type:"f32"},{value:0,type:"f32"},{value:hr,type:"u32"},{value:E,type:"u32"},{value:E,type:"u32"},{value:e==="right"?1:0,type:"u32"},{value:0,type:"u32"}],"strsm-apply-params"),Yr=B(Wr.getBindGroupLayout(0),[{buffer:J,offset:kr*hr*hr*4,size:hr*hr*4},X,Q,qr]),Xr=Cr?{x:Math.min(Tr,a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(jr,a.limits.maxComputeWorkgroupsPerDimension)}:{x:Math.min(Math.ceil(Mr/yo),a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(Math.ceil(Sr/vo),a.limits.maxComputeWorkgroupsPerDimension)};ar(or,Wr,Yr,Xr)}let Hr=R?Er:0,re=R?x:cr,ee=Hr<re,Go=Z([{value:cr,type:"u32"},{value:wr,type:"u32"},{value:0,type:"u32"},{value:E,type:"u32"},{value:m,type:"u32"},{value:b==="column-major"?1:0,type:"u32"},{value:T?1:0,type:"u32"},{value:0,type:"u32"}],"strsm-scatter-params"),ko=B(K.getBindGroupLayout(0),[Q,Y,Go]),Po=Nr&&!ee&&ir?{timestampWrites:{querySet:ir,endOfPassWriteIndex:1}}:void 0;if(ar(or,K,ko,mr(wr,E),Po),!ee)continue;let Rr=re-Hr,No=Z([{value:Hr,type:"u32"},{value:Rr,type:"u32"},{value:cr,type:"u32"},{value:wr,type:"u32"},{value:u,type:"u32"},{value:P==="transpose"?1:0,type:"u32"},{value:T?1:0,type:"u32"},{value:2,type:"u32"}],"strsm-gather-A-params"),So=B(K.getBindGroupLayout(0),[tr,$,No]);ar(or,K,So,mr(Rr,wr));{let Sr=Rr,Mr=E,zr=wr,Tr=Math.ceil(Mr/Bo),jr=Math.ceil(Sr/_o),Cr=Tr*jr>=Eo,Wr=await G(a,Cr?"sgemm_large":"sgemm_small"),qr=Z([{value:Sr,type:"u32"},{value:Mr,type:"u32"},{value:zr,type:"u32"},{value:1,type:"f32"},{value:0,type:"f32"},{value:wr,type:"u32"},{value:E,type:"u32"},{value:E,type:"u32"},{value:0,type:"u32"},{value:0,type:"u32"}],"strsm-update-params"),Yr=B(Wr.getBindGroupLayout(0),[tr,Q,fr,qr]),Xr=Cr?{x:Math.min(Tr,a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(jr,a.limits.maxComputeWorkgroupsPerDimension)}:{x:Math.min(Math.ceil(Mr/yo),a.limits.maxComputeWorkgroupsPerDimension),y:Math.min(Math.ceil(Sr/vo),a.limits.maxComputeWorkgroupsPerDimension)};ar(or,Wr,Yr,Xr)}let Mo=Z([{value:Hr,type:"u32"},{value:Rr,type:"u32"},{value:0,type:"u32"},{value:E,type:"u32"},{value:m,type:"u32"},{value:b==="column-major"?1:0,type:"u32"},{value:T?1:0,type:"u32"},{value:1,type:"u32"}],"strsm-scatter-sub-params"),Io=B(K.getBindGroupLayout(0),[fr,Y,Mo]),Lo=Nr&&ir?{timestampWrites:{querySet:ir,endOfPassWriteIndex:1}}:void 0;ar(or,K,Io,mr(Rr,E),Lo)}}let dr=br(or,ir),Br=p?null:N(or,Y);M(or);let yr=await S(dr);if(p)return yr!==void 0?{gpuTimeMs:yr}:{};let _r=await k(Br,Float32Array);return yr!==void 0?{B:_r,gpuTimeMs:yr}:{B:_r}}finally{c||d($),p||d(Y),d(J),d(ur),d(nr)}}return Wo(ni);})();
|
|
3449
|
+
`});var Lo={};qe(Lo,{routineShaders:()=>or,shaderSources:()=>Ii});var or,Ii,Ro=O(()=>{$e();Ze();Je();et();ot();it();nt();ut();mt();dt();pt();wt();bt();xt();_t();Bt();At();St();Et();kt();Dt();Pt();It();Rt();qt();Ct();jt();Ht();Vt();zt();Yt();$t();Qt();ro();to();ao();so();lo();fo();co();po();wo();bo();xo();_o();Ao();So();Go();Eo();ko();No();Mo();or={};or.sscal={sscal:ke};or.cscal={cscal:Qe};or.sswap={sswap:rt};or.dswap={dswap:tt};or.saxpy={saxpy:at};or.scopy={scopy:st};or.dcopy={dcopy:lt};or.sdot={sdot:ft,"reduction/sum":De};or.sasum={sasum:ct,"reduction/sum":De};or.snrm2={snrm2:gt,"reduction/scaledSum":ht};or.isamax={isamax:yt,"reduction/argmax":vt};or.dasum={"f64/dekker":Kr,"f64/utils/abs":ye,"f64/utils/add":zr,dasum:Gt,"reduction/sumF64":Ne};or.ddot={"f64/dekker":Kr,"f64/utils/add":zr,"f64/utils/multiply":$r,ddot:Nt,"reduction/sumF64":Ne};or.dscal={"f64/dekker":Kr,"f64/utils/add":zr,"f64/utils/multiply":$r,dscal:Mt};or.daxpy={"f64/dekker":Kr,"f64/utils/add":zr,"f64/utils/multiply":$r,daxpy:Lt};or.idamax={"f64/dekker":Kr,"f64/utils/abs":ye,"f64/utils/greater":Pe,"f64/utils/equal":Tt,idamax:Ft,"reduction/argmaxF64":Wt};or.srot={srot:Ot};or.drot={"f64/dekker":Kr,"f64/utils/add":zr,"f64/utils/multiply":$r,drot:Kt};or.srotm={srotm:Ut};or.drotm={"f64/dekker":Kr,"f64/utils/add":zr,"f64/utils/multiply":$r,drotm:Xt};or.dnrm2={"f64/dekker":Kr,"f64/utils/abs":ye,"f64/utils/greater":Pe,"f64/utils/add":zr,"f64/utils/multiply":$r,"f64/utils/divide":Zt,"f64/utils/sqrt":Jt,dnrm2:eo,"reduction/scaledSumF64":oo};or.sgemv={sgemv_n:io,sgemv_t:no};or.ssymv={ssymv:uo};or.strmv={strmv:mo};or.strsv={strsv_invert_block:Me,strsv_apply_inverse:go,strsv_update:ho};or.sger={sger:yo};or.ssyr={ssyr:vo};or.ssyr2={ssyr2:Bo};or.sgemm={sgemm_small:ue,sgemm_large:fe};or.sgemmtr={sgemmtr_small:xe,sgemmtr_large:ve};or.ssyrk={sgemmtr_small:xe,sgemmtr_large:ve};or.ssyr2k={sgemmtr_small:xe,sgemmtr_large:ve};or.ssymm={sgemm_small:ue,sgemm_large:fe,symmetrize:Do};or.strmm={sgemm_small:ue,sgemm_large:fe,triangularize:Po};or.strsm={strsv_invert_block:Me,block_transfer:Io,sscal:ke,sgemm_small:ue,sgemm_large:fe};Ii=Object.assign({},...Object.values(or))});var Ti={};qe(Ti,{Complex32:()=>Wr,Complex32Array:()=>_r,Complex64:()=>jr,Complex64Array:()=>Gr,GpuMatrix:()=>X,GpuVector:()=>N,cleanup:()=>Oe,cscal:()=>To,dasum:()=>Uo,daxpy:()=>Ho,dcopy:()=>Vo,ddot:()=>Yo,dnrm2:()=>$o,drot:()=>ra,drotm:()=>ta,dscal:()=>Co,dswap:()=>jo,gpuName:()=>Ve,idamax:()=>Qo,init:()=>He,isamax:()=>Zo,randomFloat32Array:()=>Ue,randomFloat64Array:()=>Ye,randomTriangularFloat32Array:()=>Xe,sasum:()=>zo,saxpy:()=>Wo,scopy:()=>Oo,sdot:()=>Ko,sgemm:()=>da,sgemmtr:()=>ca,sgemv:()=>oa,sger:()=>ua,snrm2:()=>Xo,srot:()=>Jo,srotm:()=>ea,sscal:()=>qo,sswap:()=>Fo,ssymm:()=>wa,ssymv:()=>aa,ssyr:()=>fa,ssyr2:()=>ma,ssyr2k:()=>ga,ssyrk:()=>pa,strmm:()=>ha,strmv:()=>ia,strsm:()=>ba,strsv:()=>la});function Ce(r,t){return t?r.features.has("timestamp-query")?{requiredFeatures:["timestamp-query"]}:(console.warn("timestamp-query not supported on this device \u2014 benchmark mode disabled."),{}):{}}function Fe(r){if(!je(r))return{querySet:null,passDescriptor:void 0};let t=r.createQuerySet({type:"timestamp",count:2});return{querySet:t,passDescriptor:{timestampWrites:{querySet:t,beginningOfPassWriteIndex:0,endOfPassWriteIndex:1}}}}function Lr(r,t,e){if(!e)return null;let o=r.createBuffer({label:"timestamp-resolve",size:16,usage:GPUBufferUsage.QUERY_RESOLVE|GPUBufferUsage.COPY_SRC});t.resolveQuerySet(e,0,2,o,0);let a=r.createBuffer({label:"timestamp-readback",size:16,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ});return t.copyBufferToBuffer(o,0,a,0,16),{tsReadBuffer:a,resolveBuffer:o,querySet:e}}async function P(r){if(!r)return;let{tsReadBuffer:t,resolveBuffer:e,querySet:o}=r;await t.mapAsync(GPUMapMode.READ);let a=new BigInt64Array(t.getMappedRange().slice());return t.unmap(),t.destroy(),e.destroy(),o.destroy(),Math.max(0,Number(a[1]-a[0]))/1e6}var Qr=null,Se=!1,Jr=new Map,le=new WeakMap,Vr=null,We=({powerPreference:r,benchmark:t})=>`${r}::${t}`;async function He({powerPreference:r="high-performance",benchmark:t=!1,dumpShaders:e=!1}={}){let o={powerPreference:r,benchmark:t,dumpShaders:e},a=We(o),i=Jr.get(a);if(i)return i;if(Qr)e!==Se&&typeof window>"u"&&console.warn(`dumpShaders: ${e} was requested, but the WebGPU instance was already created with dumpShaders: ${Se}. The first init() call fixes this for the process.`);else if(typeof window>"u"){let{create:m,globals:p}=await import("webgpu");Object.assign(globalThis,p),Qr=m(e?["enable-dawn-features=dump_shaders,disable_symbol_renaming"]:[]),Se=e}else e&&console.warn("dumpShaders has no effect in the browser \u2014 see init()'s docs."),Qr=navigator.gpu;if(!Qr)throw new Error("WebGPU not supported in this environment.");let s=await Qr.requestAdapter({powerPreference:r})??await Qr.requestAdapter();if(!s)throw new Error("No WebGPU adapter found.");let n=[...Ce(s,t).requiredFeatures??[]],f=await s.requestDevice({requiredFeatures:n});f.addEventListener("uncapturederror",m=>{console.error("Uncaptured GPU error:",m.error.message)});let l=n.includes("timestamp-query");return le.set(f,{adapter:s,benchmark:l,options:o}),Jr.set(a,f),Vr||(Vr=f),f}function Oe(r){if(r===void 0){for(let e of Jr.values())e.destroy();Jr.clear(),Vr=null;return}let t=le.get(r);t&&(Jr.delete(We(t.options)),le.delete(r),r.destroy(),Vr===r&&(Vr=Jr.values().next().value??null))}function Ve(r=Vr){let t=r&&le.get(r);if(!t)throw new Error("WebGPU adapter not initialized \u2014 call init() first.");let{device:e,description:o}=t.adapter.info;return{description:o||"unknown",device:e||"unknown"}}function je(r=Vr){return le.get(r)?.benchmark??!1}function re(){if(!Vr)throw new Error("WebGPU device not initialized \u2014 call init() first.");return Vr}function d(...r){r.flat().forEach(t=>t.destroy())}function Ge(r,t,e){let o=r.limits.maxStorageBufferBindingSize;if(t>o)throw new Error(`Buffer "${e}" needs ${t} bytes, exceeding this device's maxStorageBufferBindingSize (${o} bytes). The operands are too large for this device.`)}function x(r,t,e="blas-input",o=!1){let a=t.byteLength;Ge(r,a,e);let i=o?GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC:GPUBufferUsage.STORAGE,s=r.createBuffer({label:e,size:a,usage:i,mappedAtCreation:!0}),u=t.constructor;return new u(s.getMappedRange()).set(t),s.unmap(),s}function tr(r,t,e="blas-storage",o=0){return Ge(r,t,e),r.createBuffer({label:e,size:t,usage:GPUBufferUsage.STORAGE|o})}function Br(r,t,e="blas-result"){return Ge(r,t,e),r.createBuffer({label:e,size:t,usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC})}function G(r,t,e){let o=r.createBuffer({label:"blas-readback",size:e.size,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ});return t.copyBufferToBuffer(e,0,o,0,e.size),o}var ee=16,Ke=new WeakMap;function La(r){let t=Ke.get(r);return t||(t=r.createBuffer({label:"blas-vec4-fallback",size:ee,usage:GPUBufferUsage.STORAGE}),Ke.set(r,t)),t}function Sr(r,t){let e=t instanceof GPUBuffer?t:t.buffer,o=t instanceof GPUBuffer?0:t.offset??0,a=t instanceof GPUBuffer?t.size:t.size??e.size-o,i=Math.floor(a/ee)*ee;return i<ee?{buffer:La(r),offset:0,size:ee}:{buffer:e,offset:o,size:i}}function Ee(r,t,e,o){if(t%4!==0)return!1;let a=r instanceof GPUBuffer?r:r.buffer,i=r instanceof GPUBuffer?0:r.offset??0,s=r instanceof GPUBuffer?a.size:r.size??a.size-i,u=Math.floor(s/ee)*4;if(u<=0)return!1;let n=(Math.max(e,1)-1)*t+(Math.max(o,1)-1);return Math.floor(n/4)*4+4<=u}function I(r,t,e="blas-params"){let o=t.length*4,a=Math.ceil(o/16)*16,i=new ArrayBuffer(a),s=new DataView(i);t.forEach(({value:n,type:f},l)=>{let m=l*4;if(f==="u32")s.setUint32(m,n,!0);else if(f==="i32")s.setInt32(m,n,!0);else if(f==="f32")s.setFloat32(m,n,!0);else throw new Error(`Unknown param type "${f}". Use "f32", "u32", or "i32".`)});let u=r.createBuffer({label:e,size:a,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST});return r.queue.writeBuffer(u,0,i),u}async function S(r,t=Float32Array){try{await r.mapAsync(GPUMapMode.READ);let e=new t(r.getMappedRange().slice());return r.unmap(),e}finally{r.destroy()}}function rr(r){let t=r.length,e=new Float32Array(t),o=new Float32Array(t);for(let a=0;a<t;a++){let i=Math.fround(r[a]);e[a]=i,o[a]=Math.fround(r[a]-i)}return{hi:e,lo:o}}function dr(r,t){let e=r.length,o=new Float64Array(e);for(let a=0;a<e;a++)o[a]=r[a]+t[a];return o}var jr=class{constructor(t,e){this.re=t,this.im=e}},Gr=class extends Array{constructor(t){if(t===void 0){super();return}if(typeof t=="number"){super(t);for(let o=0;o<t;o++)this[o]=new jr(0,0);return}let e=Array.from(t);if(super(),e.length!==0){if(e[0]instanceof jr){for(let o of e){if(!(o instanceof jr))throw new Error("Complex64Array expects every element to be a Complex64.");this.push(o)}return}if(e.length%2!==0)throw new Error("Complex64Array expects an even number of interleaved [re, im, ...] values.");for(let o=0;o<e.length;o+=2){if(typeof e[o]!="number"||typeof e[o+1]!="number")throw new Error("Complex64Array expects interleaved [re, im, ...] values to be numbers.");this.push(new jr(e[o],e[o+1]))}}}};function te(r,t=r.length){let e=new Float32Array(t*2);for(let o=0;o<t;o++)e[o*2]=r[o].re,e[o*2+1]=r[o].im;return e}function he(r,t=r.length){let e=new Float64Array(t),o=new Float64Array(t);for(let l=0;l<t;l++)e[l]=r[l].re,o[l]=r[l].im;let{hi:a,lo:i}=rr(e),{hi:s,lo:u}=rr(o),n=new Float32Array(t*2),f=new Float32Array(t*2);for(let l=0;l<t;l++)n[l*2]=a[l],n[l*2+1]=s[l],f[l*2]=i[l],f[l*2+1]=u[l];return{hi:n,lo:f}}function be(r,t){let e=r.length/2,o=new Float32Array(e),a=new Float32Array(e),i=new Float32Array(e),s=new Float32Array(e);for(let l=0;l<e;l++)o[l]=r[l*2],i[l]=r[l*2+1],a[l]=t[l*2],s[l]=t[l*2+1];let u=dr(o,a),n=dr(i,s),f=new Gr(e);for(let l=0;l<e;l++)f[l]=new jr(u[l],n[l]);return f}var Wr=class{constructor(t,e){this.re=Math.fround(t),this.im=Math.fround(e)}},_r=class extends Array{constructor(t){if(t===void 0){super();return}if(typeof t=="number"){super(t);for(let o=0;o<t;o++)this[o]=new Wr(0,0);return}let e=Array.from(t);if(super(),e.length!==0){if(e[0]instanceof Wr){for(let o of e){if(!(o instanceof Wr))throw new Error("Complex32Array expects every element to be a Complex32.");this.push(o)}return}if(e.length%2!==0)throw new Error("Complex32Array expects an even number of interleaved [re, im, ...] values.");for(let o=0;o<e.length;o+=2){if(typeof e[o]!="number"||typeof e[o+1]!="number")throw new Error("Complex32Array expects interleaved [re, im, ...] values to be numbers.");this.push(new Wr(e[o],e[o+1]))}}}};var N=class r{constructor(t,e,o=Float32Array,a=null,i=null){this._buf=t,this._loBuf=a,this.length=e,this.dtype=o,this.device=i??re()}static from(t,e){let o=t instanceof GPUDevice,a=o?t:re(),i=o?e:t;if(i instanceof Float64Array){let{hi:u,lo:n}=rr(i),f=x(a,u,"gpu-vector-f64-hi",!0),l=x(a,n,"gpu-vector-f64-lo",!0);return new r(f,i.length,Float64Array,l,a)}if(i instanceof _r){let u=x(a,te(i),"gpu-vector-complex32",!0);return new r(u,i.length,_r,null,a)}if(i instanceof Gr){let{hi:u,lo:n}=he(i),f=x(a,u,"gpu-vector-complex64-hi",!0),l=x(a,n,"gpu-vector-complex64-lo",!0);return new r(f,i.length,Gr,l,a)}if(!(i instanceof Float32Array))throw new Error("GpuVector.from expects a Float32Array, Float64Array, Complex32Array, or Complex64Array.");let s=x(a,i,"gpu-vector",!0);return new r(s,i.length,i.constructor,null,a)}async read(){let t=this.device,e=t.createCommandEncoder(),o=G(t,e,this._buf);if(t.queue.submit([e.finish()]),this.dtype===_r)return new _r(await S(o,Float32Array));if(!this._loBuf)return S(o,this.dtype);let a=t.createCommandEncoder(),i=G(t,a,this._loBuf);t.queue.submit([a.finish()]);let[s,u]=await Promise.all([S(o,Float32Array),S(i,Float32Array)]);return this.dtype===Gr?be(s,u):dr(s,u)}destroy(){this._buf.destroy(),this._loBuf&&this._loBuf.destroy()}};var X=class r{constructor(t,e,o,a,i=null,s="row-major",u=null,n=Float32Array){this._buf=t,this._loBuf=i,this.rows=e,this.cols=o,this.lda=a,this.layout=s,this.dtype=n,this.device=u??re()}static from(t,...e){let o=t instanceof GPUDevice,a=o?t:re(),i=o?e.shift():t,[s,u,n,f="row-major"]=e;if(f!=="row-major"&&f!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");let l=f==="row-major";if(n===void 0&&(n=l?u:s),!(i instanceof Float32Array)&&!(i instanceof Float64Array)&&!(i instanceof _r)&&!(i instanceof Gr))throw new Error("GpuMatrix.from expects a Float32Array, Float64Array, Complex32Array, or Complex64Array.");if(!Number.isInteger(s)||s<=0)throw new Error("rows must be a positive integer.");if(!Number.isInteger(u)||u<=0)throw new Error("cols must be a positive integer.");let m=l?u:s;if(!Number.isInteger(n)||n<m)throw new Error(`lda must be an integer >= ${l?"cols":"rows"}.`);let p=l?s:u;if(i.length<p*n)throw new Error("data does not have enough elements for the given rows, cols, and lda.");if(i instanceof Float64Array){let g=p*n,{hi:w,lo:h}=rr(i.subarray(0,g)),b=x(a,w,"gpu-matrix-f64-hi",!0),y=x(a,h,"gpu-matrix-f64-lo",!0);return new r(b,s,u,n,y,f,a,Float64Array)}if(i instanceof _r){let g=x(a,te(i,p*n),"gpu-matrix-complex32",!0);return new r(g,s,u,n,null,f,a,_r)}if(i instanceof Gr){let{hi:g,lo:w}=he(i,p*n),h=x(a,g,"gpu-matrix-complex64-hi",!0),b=x(a,w,"gpu-matrix-complex64-lo",!0);return new r(h,s,u,n,b,f,a,Gr)}let c=x(a,i.subarray(0,p*n),"gpu-matrix",!0);return new r(c,s,u,n,null,f,a)}async read(){let t=this.device,e=t.createCommandEncoder(),o=G(t,e,this._buf);t.queue.submit([e.finish()]);let a=this.layout!=="column-major",i=a?this.rows:this.cols,s=a?this.cols:this.rows;if(this.dtype===_r){let f=new _r(await S(o,Float32Array));if(this.lda===s)return f;let l=new _r(i*s);for(let m=0;m<i;m++)for(let p=0;p<s;p++)l[m*s+p]=f[m*this.lda+p];return l}if(this._loBuf){let f=t.createCommandEncoder(),l=G(t,f,this._loBuf);t.queue.submit([f.finish()]);let[m,p]=await Promise.all([S(o,Float32Array),S(l,Float32Array)]);if(this.dtype===Gr){let w=be(m,p);if(this.lda===s)return w;let h=new Gr(i*s);for(let b=0;b<i;b++)for(let y=0;y<s;y++)h[b*s+y]=w[b*this.lda+y];return h}let c=dr(m,p);if(this.lda===s)return c;let g=new Float64Array(i*s);for(let w=0;w<i;w++)g.set(c.subarray(w*this.lda,w*this.lda+s),w*s);return g}let u=await S(o,Float32Array);if(this.lda===s)return u;let n=new Float32Array(i*s);for(let f=0;f<i;f++)n.set(u.subarray(f*this.lda,f*this.lda+s),f*s);return n}destroy(){this._buf.destroy(),this._loBuf&&this._loBuf.destroy()}};function ze(r){let t=r>>>0;return function(){t=t+1831565813|0;let e=Math.imul(t^t>>>15,1|t);return e=e+Math.imul(e^e>>>7,61|e)^e,((e^e>>>14)>>>0)/4294967296}}function Ue(r,t=-1,e=1,o){let a=new Float32Array(r),i=o===void 0?Math.random:ze(o);for(let s=0;s<r;s++)a[s]=t+i()*(e-t);return a}function Ye(r,t=-1,e=1,o){let a=new Float64Array(r),i=o===void 0?Math.random:ze(o);for(let s=0;s<r;s++)a[s]=t+i()*(e-t);return a}function Xe(r,t,e="lower",o=-1,a=1,i=5,s=15,u="row-major"){if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(u!=="row-major"&&u!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(t<r)throw new Error("lda must be >= n.");let n=u==="column-major",f=(m,p)=>n?p*t+m:m*t+p,l=new Float32Array(r*t);for(let m=0;m<r;m++){for(let p=0;p<r;p++){if(m===p)continue;(e==="lower"?p<m:p>m)&&(l[f(m,p)]=o+Math.random()*(a-o))}l[f(m,m)]=i+Math.random()*(s-i)}return l}function E(r,t,e,o=0){let a=e.map((i,s)=>({binding:o+s,resource:i instanceof GPUBuffer?{buffer:i}:i}));return r.createBindGroup({layout:t,entries:a})}function M(r,t){r.queue.submit([t.finish()])}function qr(r){let{querySet:t,passDescriptor:e}=Fe(r);return{commandEncoder:r.createCommandEncoder(),querySet:t,passDescriptor:e}}function gr(r,t,e,o,a){let i=r.beginComputePass(a);i.setPipeline(t),i.setBindGroup(0,e),typeof o=="number"?i.dispatchWorkgroups(o):i.dispatchWorkgroups(o.x,o.y,o.z??1),i.end()}function j(r,t,e,o){let{commandEncoder:a,querySet:i,passDescriptor:s}=qr(r);gr(a,t,e,o,s);let u=Lr(r,a,i);return{commandEncoder:a,ts:u}}var qi={},Ie=new WeakMap;async function D(r,t,e="main"){Ie.has(r)||Ie.set(r,new Map);let o=Ie.get(r),a=Array.isArray(t)?t:[t],i=`${a.join("+")}::${e}`;if(!o.has(i)){let s=Ri(r,a,e).catch(u=>{throw o.delete(i),u});o.set(i,s)}return o.get(i)}async function Li(r){if(typeof process>"u"||!process.versions?.node){let{shaderSources:t}=await Promise.resolve().then(()=>(Ro(),Lo)),e=t[r];if(!e)throw new Error(`Shader "${r}" not found in browser bundle.`);return e}else{let{readFileSync:t}=await import("fs"),{fileURLToPath:e}=await import("url"),{dirname:o,join:a}=await import("path"),i=o(e(qi.url));return t(a(i,`../shaders/${r}.wgsl`),"utf8")}}async function Ri(r,t,e="main"){let o=t.join("+"),a=await Promise.all(t.map(Li)),i=0,s=a.map((g,w)=>{let h=g.split(`
|
|
3450
|
+
`).length,b={name:t[w],startLine:i+1,endLine:i+h};return i+=h,b}),u=g=>{let w=g&&s.find(h=>g>=h.startLine&&g<=h.endLine);return w?`${w.name}.wgsl:${g-w.startLine+1}`:`line ${g}`},n=a.join(`
|
|
3451
|
+
`),f=r.createShaderModule({label:o,code:n}),m=(await f.getCompilationInfo()).messages.filter(g=>g.type==="error");if(m.length>0)throw new Error(`Shader "${o}" compilation failed:
|
|
3452
|
+
${m.map(g=>` ${u(g.lineNum)}: ${g.message}`).join(`
|
|
3453
|
+
`)}`);let p=e==="main"?{module:f}:{module:f,entryPoint:e},c=r.createComputePipeline({label:o,layout:"auto",compute:p});return c._shaderModule=f,c}function cr(r,t,e){let o=r.limits.maxComputeWorkgroupsPerDimension;return e===void 0?Math.min(Math.ceil(t/64),o):{x:Math.min(Math.ceil(e/8),o),y:Math.min(Math.ceil(t/8),o)}}function U(r,t,e,o="x"){let a=r.limits.maxComputeWorkgroupsPerDimension;if(t>a)throw new Error(`${e}: this problem needs ${t} workgroups in ${o}, but the device allows ${a} (maxComputeWorkgroupsPerDimension). The operands are too large for this device \u2014 split the operation into smaller blocks.`);return t}function Zr(r,t,e,o){return o===void 0?U(r,Math.ceil(e/64),t):{x:U(r,Math.ceil(o/8),t,"x"),y:U(r,Math.ceil(e/8),t,"y")}}function q(r){if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.")}function T(r,t,e){for(let[o,a]of Object.entries(e))if(!(!(a instanceof N)&&!(a instanceof X))&&a.device!==r)throw new Error(`${t}: ${o} belongs to a different GPUDevice than the one passed in. GPU buffers cannot be shared across devices \u2014 recreate the operand on this device, or call the routine with the device that owns it.`)}async function qo(r,t,e,o,a){let i=o instanceof N;if(q(r),T(r,"sscal",{x:o}),!Number.isInteger(t)||!Number.isInteger(a))throw new Error("n and incx must be integers.");if(typeof e!="number")throw new Error("alpha must be a number.");if(Number.isNaN(e))throw new Error("alpha must not be NaN.");if(!Number.isFinite(e))throw new Error("alpha must be finite.");if(a<=0)throw new Error("incx must be positive.");if(!(o instanceof Float32Array)&&!(o instanceof N))throw new Error("x must be a Float32Array or GpuVector.");if(t<=0)return i?{}:{x:o};if(o.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");let s=await D(r,"sscal"),u=null,n=null,f=null;try{u=i?o._buf:x(r,o,"sscal-x",!0),n=I(r,[{value:t,type:"u32"},{value:e,type:"f32"},{value:a,type:"u32"}],"sscal-params");let l=E(r,s.getBindGroupLayout(0),[u,n]),{commandEncoder:m,ts:p}=j(r,s,l,cr(r,t));f=i?null:G(r,m,u),M(r,m);let c=await P(p);if(i)return c!==void 0?{gpuTimeMs:c}:{};let g=await S(f,Float32Array);return f=null,c!==void 0?{x:g,gpuTimeMs:c}:{x:g}}finally{!i&&u&&d(u),n&&d(n),f&&d(f)}}async function To(r,t,e,o,a){let i=o instanceof N;if(q(r),T(r,"cscal",{x:o}),!Number.isInteger(t)||!Number.isInteger(a))throw new Error("n and incx must be integers.");if(!(e instanceof Wr))throw new Error("alpha must be a Complex32.");if(Number.isNaN(e.re)||Number.isNaN(e.im))throw new Error("alpha must not be NaN.");if(!Number.isFinite(e.re)||!Number.isFinite(e.im))throw new Error("alpha must be finite.");if(a<=0)throw new Error("incx must be positive.");if(!(o instanceof _r)&&!i)throw new Error("x must be a Complex32Array or GpuVector.");if(i&&o.dtype!==_r)throw new Error("x must be a Complex32Array-backed GpuVector.");if(t<=0)return i?{}:{x:o};if(o.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");let s=await D(r,"cscal"),u=null,n=null,f=null;try{u=i?o._buf:x(r,te(o),"cscal-x",!0),n=I(r,[{value:t,type:"u32"},{value:e.re,type:"f32"},{value:e.im,type:"f32"},{value:a,type:"u32"}],"cscal-params");let l=E(r,s.getBindGroupLayout(0),[u,n]),{commandEncoder:m,ts:p}=j(r,s,l,cr(r,t));f=i?null:G(r,m,u),M(r,m);let c=await P(p);if(i)return c!==void 0?{gpuTimeMs:c}:{};let g=await S(f,Float32Array);f=null;let w=new _r(g);return c!==void 0?{x:w,gpuTimeMs:c}:{x:w}}finally{!i&&u&&d(u),n&&d(n),f&&d(f)}}async function Co(r,t,e,o,a){let i=o instanceof N;if(q(r),!Number.isInteger(t)||!Number.isInteger(a))throw new Error("n and incx must be integers.");if(typeof e!="number")throw new Error("alpha must be a number.");if(Number.isNaN(e))throw new Error("alpha must not be NaN.");if(!Number.isFinite(e))throw new Error("alpha must be finite.");if(!(o instanceof Float64Array)&&!i)throw new Error("x must be a Float64Array or GpuVector.");if(i&&o.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(a<=0)throw new Error("incx must be positive.");if(T(r,"dscal",{x:o}),t<=0)return i?{}:{x:o};if(o.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");let u=await D(r,[...["f64/dekker","f64/utils/add","f64/utils/multiply"],"dscal"]),{hi:n,lo:f}=rr(new Float64Array([e])),l=null,m=null,p=null,c=null,g=null;try{if(i)l=o._buf,m=o._loBuf;else{let{hi:k,lo:B}=rr(o);l=x(r,k,"dscal-xHi",!0),m=x(r,B,"dscal-xLo",!0)}p=I(r,[{value:t,type:"u32"},{value:n[0],type:"f32"},{value:f[0],type:"f32"},{value:a,type:"u32"}],"dscal-params");let w=E(r,u.getBindGroupLayout(0),[l,m,p]),{commandEncoder:h,ts:b}=j(r,u,w,cr(r,t));c=i?null:G(r,h,l),g=i?null:G(r,h,m),M(r,h);let y=await P(b);if(i)return y!==void 0?{gpuTimeMs:y}:{};let v=await S(c,Float32Array);c=null;let _=await S(g,Float32Array);g=null;let A=dr(v,_);return y!==void 0?{x:A,gpuTimeMs:y}:{x:A}}finally{!i&&l&&d(l),!i&&m&&d(m),p&&d(p),c&&d(c),g&&d(g)}}async function Fo(r,t,e,o,a,i){let s=e instanceof N,u=a instanceof N;if(q(r),T(r,"sswap",{x:e,y:a}),!Number.isInteger(t)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!(e instanceof Float32Array)&&!(e instanceof N))throw new Error("x must be a Float32Array or GpuVector.");if(!(a instanceof Float32Array)&&!(a instanceof N))throw new Error("y must be a Float32Array or GpuVector.");if(e.constructor!==a.constructor)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(t<=0)return s?{}:{x:e,y:a};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(a.length<(t-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let n=await D(r,"sswap"),f=null,l=null,m=null,p=null,c=null;try{f=s?e._buf:x(r,e,"sswap-x",!0),l=u?a._buf:x(r,a,"sswap-y",!0),m=I(r,[{value:t,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"sswap-params");let g=E(r,n.getBindGroupLayout(0),[f,l,m]),{commandEncoder:w,ts:h}=j(r,n,g,cr(r,t));p=s?null:G(r,w,f),c=u?null:G(r,w,l),M(r,w);let b=await P(h);if(s)return b!==void 0?{gpuTimeMs:b}:{};let y=await S(p,Float32Array);p=null;let v=await S(c,Float32Array);return c=null,b!==void 0?{x:y,y:v,gpuTimeMs:b}:{x:y,y:v}}finally{!s&&f&&d(f),!u&&l&&d(l),m&&d(m),p&&d(p),c&&d(c)}}async function jo(r,t,e,o,a,i){let s=e instanceof N,u=a instanceof N;if(q(r),T(r,"dswap",{x:e,y:a}),!Number.isInteger(t)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!(e instanceof Float64Array)&&!s)throw new Error("x must be a Float64Array or GpuVector.");if(!(a instanceof Float64Array)&&!u)throw new Error("y must be a Float64Array or GpuVector.");if(s&&e.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(u&&a.dtype!==Float64Array)throw new Error("y must be a Float64Array-backed GpuVector.");if(s!==u)throw new Error("x and y must be the same type (both Float64Array or both GpuVector).");if(t<=0)return s?{}:{x:e,y:a};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(a.length<(t-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let n=await D(r,"dswap"),f=null,l=null,m=null,p=null,c=null,g=null,w=null,h=null,b=null;try{if(s)f=e._buf,l=e._loBuf,m=a._buf,p=a._loBuf;else{let W=rr(e),V=rr(a);f=x(r,W.hi,"dswap-xHi",!0),l=x(r,W.lo,"dswap-xLo",!0),m=x(r,V.hi,"dswap-yHi",!0),p=x(r,V.lo,"dswap-yLo",!0)}c=I(r,[{value:t,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"dswap-params");let y=E(r,n.getBindGroupLayout(0),[f,l,m,p,c]),{commandEncoder:v,ts:_}=j(r,n,y,cr(r,t));g=s?null:G(r,v,f),w=s?null:G(r,v,l),h=u?null:G(r,v,m),b=u?null:G(r,v,p),M(r,v);let A=await P(_);if(s)return A!==void 0?{gpuTimeMs:A}:{};let k=await S(g,Float32Array);g=null;let B=await S(w,Float32Array);w=null;let L=await S(h,Float32Array);h=null;let C=await S(b,Float32Array);b=null;let R=dr(k,B),F=dr(L,C);return A!==void 0?{x:R,y:F,gpuTimeMs:A}:{x:R,y:F}}finally{!s&&f&&d(f),!s&&l&&d(l),!u&&m&&d(m),!u&&p&&d(p),c&&d(c),g&&d(g),w&&d(w),h&&d(h),b&&d(b)}}async function Wo(r,t,e,o,a,i,s){let u=o instanceof N,n=i instanceof N;if(q(r),T(r,"saxpy",{x:o,y:i}),!Number.isInteger(t)||!Number.isInteger(a)||!Number.isInteger(s))throw new Error("n, incx, and incy must be integers.");if(typeof e!="number")throw new Error("alpha must be a number.");if(Number.isNaN(e))throw new Error("alpha must not be NaN.");if(!Number.isFinite(e))throw new Error("alpha must be finite.");if(a<=0||s<=0)throw new Error("incx and incy must be positive.");if(!u&&!(o instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!n&&!(i instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(u!==n)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(t<=0)return n?{}:{y:i};if(o.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(i.length<(t-1)*s+1)throw new Error("y does not have enough elements for the given n and incy.");let f=await D(r,"saxpy"),l=null,m=null,p=null,c=null;try{l=u?o._buf:x(r,o,"saxpy-x",!1),m=n?i._buf:x(r,i,"saxpy-y",!0),p=I(r,[{value:t,type:"u32"},{value:e,type:"f32"},{value:a,type:"u32"},{value:s,type:"u32"}],"saxpy-params");let g=E(r,f.getBindGroupLayout(0),[l,m,p]),{commandEncoder:w,ts:h}=j(r,f,g,cr(r,t));c=n?null:G(r,w,m),M(r,w);let b=await P(h);if(n)return b!==void 0?{gpuTimeMs:b}:{};let y=await S(c,Float32Array);return c=null,b!==void 0?{y,gpuTimeMs:b}:{y}}finally{!u&&l&&d(l),!n&&m&&d(m),p&&d(p),c&&d(c)}}async function Ho(r,t,e,o,a,i,s){let u=o instanceof N,n=i instanceof N;if(q(r),!Number.isInteger(t)||!Number.isInteger(a)||!Number.isInteger(s))throw new Error("n, incx, and incy must be integers.");if(typeof e!="number")throw new Error("alpha must be a number.");if(Number.isNaN(e))throw new Error("alpha must not be NaN.");if(!Number.isFinite(e))throw new Error("alpha must be finite.");if(!(o instanceof Float64Array)&&!u)throw new Error("x must be a Float64Array or GpuVector.");if(!(i instanceof Float64Array)&&!n)throw new Error("y must be a Float64Array or GpuVector.");if(u&&o.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(n&&i.dtype!==Float64Array)throw new Error("y must be a Float64Array-backed GpuVector.");if(u!==n)throw new Error("x and y must be the same type (both Float64Array or both GpuVector).");if(a<=0||s<=0)throw new Error("incx and incy must be positive.");if(T(r,"daxpy",{x:o,y:i}),t<=0)return n?{}:{y:i};if(o.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(i.length<(t-1)*s+1)throw new Error("y does not have enough elements for the given n and incy.");let l=await D(r,[...["f64/dekker","f64/utils/add","f64/utils/multiply"],"daxpy"]),{hi:m,lo:p}=rr(new Float64Array([e])),c=null,g=null,w=null,h=null,b=null,y=null,v=null;try{if(u)c=o._buf,g=o._loBuf,w=i._buf,h=i._loBuf;else{let F=rr(o),W=rr(i);c=x(r,F.hi,"daxpy-xHi",!1),g=x(r,F.lo,"daxpy-xLo",!1),w=x(r,W.hi,"daxpy-yHi",!0),h=x(r,W.lo,"daxpy-yLo",!0)}b=I(r,[{value:t,type:"u32"},{value:m[0],type:"f32"},{value:p[0],type:"f32"},{value:a,type:"u32"},{value:s,type:"u32"}],"daxpy-params");let _=E(r,l.getBindGroupLayout(0),[c,g,w,h,b]),{commandEncoder:A,ts:k}=j(r,l,_,cr(r,t));y=n?null:G(r,A,w),v=n?null:G(r,A,h),M(r,A);let B=await P(k);if(n)return B!==void 0?{gpuTimeMs:B}:{};let L=await S(y,Float32Array);y=null;let C=await S(v,Float32Array);v=null;let R=dr(L,C);return B!==void 0?{y:R,gpuTimeMs:B}:{y:R}}finally{!u&&c&&d(c),!u&&g&&d(g),!n&&w&&d(w),!n&&h&&d(h),b&&d(b),y&&d(y),v&&d(v)}}async function Oo(r,t,e,o,a,i){let s=e instanceof N,u=a instanceof N;if(q(r),T(r,"scopy",{x:e,y:a}),!Number.isInteger(t)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!s&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!u&&!(a instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(s!==u)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(t<=0)return u?{}:{y:a};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(a.length<(t-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let n=await D(r,"scopy"),f=null,l=null,m=null,p=null;try{f=s?e._buf:x(r,e,"scopy-x",!1),l=u?a._buf:x(r,a,"scopy-y",!0),m=I(r,[{value:t,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"scopy-params");let c=E(r,n.getBindGroupLayout(0),[f,l,m]),{commandEncoder:g,ts:w}=j(r,n,c,cr(r,t));p=u?null:G(r,g,l),M(r,g);let h=await P(w);if(u)return h!==void 0?{gpuTimeMs:h}:{};let b=await S(p,Float32Array);return p=null,h!==void 0?{y:b,gpuTimeMs:h}:{y:b}}finally{!s&&f&&d(f),!u&&l&&d(l),m&&d(m),p&&d(p)}}async function Vo(r,t,e,o,a,i){let s=e instanceof N,u=a instanceof N;if(q(r),T(r,"dcopy",{x:e,y:a}),!Number.isInteger(t)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!(e instanceof Float64Array)&&!s)throw new Error("x must be a Float64Array or GpuVector.");if(!(a instanceof Float64Array)&&!u)throw new Error("y must be a Float64Array or GpuVector.");if(s&&e.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(u&&a.dtype!==Float64Array)throw new Error("y must be a Float64Array-backed GpuVector.");if(s!==u)throw new Error("x and y must be the same type (both Float64Array or both GpuVector).");if(t<=0)return u?{}:{y:a};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(a.length<(t-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let n=await D(r,"dcopy"),f=null,l=null,m=null,p=null,c=null,g=null,w=null;try{if(s)f=e._buf,l=e._loBuf,m=a._buf,p=a._loBuf;else{let B=rr(e),L=rr(a);f=x(r,B.hi,"dcopy-xHi",!1),l=x(r,B.lo,"dcopy-xLo",!1),m=x(r,L.hi,"dcopy-yHi",!0),p=x(r,L.lo,"dcopy-yLo",!0)}c=I(r,[{value:t,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"dcopy-params");let h=E(r,n.getBindGroupLayout(0),[f,l,m,p,c]),{commandEncoder:b,ts:y}=j(r,n,h,cr(r,t));g=u?null:G(r,b,m),w=u?null:G(r,b,p),M(r,b);let v=await P(y);if(u)return v!==void 0?{gpuTimeMs:v}:{};let _=await S(g,Float32Array);g=null;let A=await S(w,Float32Array);w=null;let k=dr(_,A);return v!==void 0?{y:k,gpuTimeMs:v}:{y:k}}finally{!s&&f&&d(f),!s&&l&&d(l),!u&&m&&d(m),!u&&p&&d(p),c&&d(c),g&&d(g),w&&d(w)}}async function Ko(r,t,e,o,a,i){let s=e instanceof N,u=a instanceof N;if(q(r),T(r,"sdot",{x:e,y:a}),!Number.isInteger(t)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!s&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!u&&!(a instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(s!==u)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(t<=0)return{dot:0};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(a.length<(t-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let n=await D(r,"sdot"),f=await D(r,"reduction/sum"),l=null,m=null,p=null,c=null,g=null,w=null;try{l=s?e._buf:x(r,e,"sdot-x",!1),m=u?a._buf:x(r,a,"sdot-y",!1),p=tr(r,512,"sdot-partials"),c=Br(r,4,"sdot-result"),g=I(r,[{value:t,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"sdot-params");let h=E(r,n.getBindGroupLayout(0),[l,m,p,g]),{commandEncoder:b,ts:y}=j(r,n,h,128);M(r,b);let v=E(r,f.getBindGroupLayout(0),[p,c]),{commandEncoder:_,ts:A}=j(r,f,v,1);w=G(r,_,c),M(r,_);let k=S(w,Float32Array);w=null;let[B,L,C]=await Promise.all([P(y),P(A),k]);return B!==void 0&&L!==void 0?{dot:C[0],gpuTimeMs:B+L}:{dot:C[0]}}finally{!s&&l&&d(l),!u&&m&&d(m),p&&d(p),c&&d(c),g&&d(g),w&&d(w)}}async function zo(r,t,e,o){let a=e instanceof N;if(q(r),T(r,"sasum",{x:e}),!Number.isInteger(t)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(o<=0)throw new Error("incx must be positive.");if(!a&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(t<=0)return{asum:0};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let i=await D(r,"sasum"),s=await D(r,"reduction/sum"),u=null,n=null,f=null,l=null,m=null;try{u=a?e._buf:x(r,e,"sasum-x",!1),n=tr(r,512,"sasum-partials"),f=Br(r,4,"sasum-result"),l=I(r,[{value:t,type:"u32"},{value:o,type:"u32"}],"sasum-params");let p=E(r,i.getBindGroupLayout(0),[u,n,l]),{commandEncoder:c,ts:g}=j(r,i,p,128);M(r,c);let w=E(r,s.getBindGroupLayout(0),[n,f]),{commandEncoder:h,ts:b}=j(r,s,w,1);m=G(r,h,f),M(r,h);let y=S(m,Float32Array);m=null;let[v,_,A]=await Promise.all([P(g),P(b),y]);return v!==void 0&&_!==void 0?{asum:A[0],gpuTimeMs:v+_}:{asum:A[0]}}finally{!a&&u&&d(u),n&&d(n),f&&d(f),l&&d(l),m&&d(m)}}async function Uo(r,t,e,o){let a=e instanceof N;if(q(r),T(r,"dasum",{x:e}),!Number.isInteger(t)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(o<=0)throw new Error("incx must be positive.");if(!a&&!(e instanceof Float64Array))throw new Error("x must be a Float64Array or GpuVector.");if(a&&e.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(t<=0)return{asum:0};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let i=["f64/dekker","f64/utils/abs","f64/utils/add"],s=await D(r,[...i,"dasum"]),u=await D(r,[...i,"reduction/sumF64"]),n=null,f=null,l=null,m=null,p=null,c=null,g=null,w=null,h=null;try{if(a)n=e._buf,f=e._loBuf;else{let{hi:z,lo:H}=rr(e.map(Math.abs));n=x(r,z,"dasum-xHi",!1),f=x(r,H,"dasum-xLo",!1)}l=tr(r,512,"dasum-partialsHi"),m=tr(r,512,"dasum-partialsLo"),p=Br(r,4,"dasum-result-hi"),c=Br(r,4,"dasum-result-lo"),g=I(r,[{value:t,type:"u32"},{value:o,type:"u32"}],"dasum-params");let b=E(r,s.getBindGroupLayout(0),[n,f,l,m,g]),{commandEncoder:y,ts:v}=j(r,s,b,128);M(r,y);let _=E(r,u.getBindGroupLayout(0),[l,m,p,c]),{commandEncoder:A,ts:k}=j(r,u,_,1);w=G(r,A,p),h=G(r,A,c),M(r,A);let B=S(w,Float32Array),L=S(h,Float32Array);w=null,h=null;let[C,R,F,W]=await Promise.all([P(v),P(k),B,L]),V=dr(F,W)[0];return C!==void 0&&R!==void 0?{asum:V,gpuTimeMs:C+R}:{asum:V}}finally{!a&&n&&d(n),!a&&f&&d(f),l&&d(l),m&&d(m),p&&d(p),c&&d(c),g&&d(g),w&&d(w),h&&d(h)}}async function Yo(r,t,e,o,a,i){let s=e instanceof N,u=a instanceof N;if(q(r),T(r,"ddot",{x:e,y:a}),!Number.isInteger(t)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!s&&!(e instanceof Float64Array))throw new Error("x must be a Float64Array or GpuVector.");if(!u&&!(a instanceof Float64Array))throw new Error("y must be a Float64Array or GpuVector.");if(s&&e.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(u&&a.dtype!==Float64Array)throw new Error("y must be a Float64Array-backed GpuVector.");if(s!==u)throw new Error("x and y must be the same type (both Float64Array or both GpuVector).");if(t<=0)return{dot:0};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(a.length<(t-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let n=["f64/dekker","f64/utils/add"],f=await D(r,[...n,"f64/utils/multiply","ddot"]),l=await D(r,[...n,"reduction/sumF64"]),m=null,p=null,c=null,g=null,w=null,h=null,b=null,y=null,v=null,_=null,A=null;try{if(s)m=e._buf,p=e._loBuf,c=a._buf,g=a._loBuf;else{let ir=rr(e),lr=rr(a);m=x(r,ir.hi,"ddot-xHi",!1),p=x(r,ir.lo,"ddot-xLo",!1),c=x(r,lr.hi,"ddot-yHi",!1),g=x(r,lr.lo,"ddot-yLo",!1)}w=tr(r,512,"ddot-partialsHi"),h=tr(r,512,"ddot-partialsLo"),b=Br(r,4,"ddot-result-hi"),y=Br(r,4,"ddot-result-lo"),v=I(r,[{value:t,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"ddot-params");let k=E(r,f.getBindGroupLayout(0),[m,p,c,g,w,h,v]),{commandEncoder:B,ts:L}=j(r,f,k,128);M(r,B);let C=E(r,l.getBindGroupLayout(0),[w,h,b,y]),{commandEncoder:R,ts:F}=j(r,l,C,1);_=G(r,R,b),A=G(r,R,y),M(r,R);let W=S(_,Float32Array),V=S(A,Float32Array);_=null,A=null;let[z,H,$,K]=await Promise.all([P(L),P(F),W,V]),Y=dr($,K)[0];return z!==void 0&&H!==void 0?{dot:Y,gpuTimeMs:z+H}:{dot:Y}}finally{!s&&m&&d(m),!s&&p&&d(p),!u&&c&&d(c),!u&&g&&d(g),w&&d(w),h&&d(h),b&&d(b),y&&d(y),v&&d(v),_&&d(_),A&&d(A)}}async function Xo(r,t,e,o){let a=e instanceof N;if(q(r),T(r,"snrm2",{x:e}),!Number.isInteger(t)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(o<=0)throw new Error("incx must be positive.");if(!a&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(t<=0)return{nrm2:0};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let i=await D(r,"snrm2"),s=await D(r,"reduction/scaledSum"),u=null,n=null,f=null,l=null,m=null,p=null;try{u=a?e._buf:x(r,e,"snrm2-x",!1),n=tr(r,512,"snrm2-partials-scale"),f=tr(r,512,"snrm2-partials-ssq"),l=Br(r,4,"snrm2-result"),m=I(r,[{value:t,type:"u32"},{value:o,type:"u32"}],"snrm2-params");let c=E(r,i.getBindGroupLayout(0),[u,n,f,m]),{commandEncoder:g,ts:w}=j(r,i,c,128);M(r,g);let h=E(r,s.getBindGroupLayout(0),[n,f,l]),{commandEncoder:b,ts:y}=j(r,s,h,1);p=G(r,b,l),M(r,b);let v=S(p,Float32Array);p=null;let[_,A,k]=await Promise.all([P(w),P(y),v]),B=k[0];return _!==void 0&&A!==void 0?{nrm2:B,gpuTimeMs:_+A}:{nrm2:B}}finally{!a&&u&&d(u),n&&d(n),f&&d(f),l&&d(l),m&&d(m),p&&d(p)}}async function $o(r,t,e,o){let a=e instanceof N;if(q(r),T(r,"dnrm2",{x:e}),!Number.isInteger(t)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(o<=0)throw new Error("incx must be positive.");if(!a&&!(e instanceof Float64Array))throw new Error("x must be a Float64Array or GpuVector.");if(a&&e.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(t<=0)return{nrm2:0};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let i=["f64/dekker","f64/utils/abs","f64/utils/greater","f64/utils/add","f64/utils/multiply","f64/utils/divide","f64/utils/sqrt"],s=await D(r,[...i,"dnrm2"]),u=await D(r,[...i,"reduction/scaledSumF64"]),n=null,f=null,l=null,m=null,p=null,c=null,g=null,w=null,h=null,b=null,y=null;try{if(a)n=e._buf,f=e._loBuf;else{let{hi:$,lo:K}=rr(e);n=x(r,$,"dnrm2-xHi",!1),f=x(r,K,"dnrm2-xLo",!1)}l=tr(r,512,"dnrm2-partials-scaleHi"),m=tr(r,512,"dnrm2-partials-scaleLo"),p=tr(r,512,"dnrm2-partials-ssqHi"),c=tr(r,512,"dnrm2-partials-ssqLo"),g=Br(r,4,"dnrm2-result-hi"),w=Br(r,4,"dnrm2-result-lo"),h=I(r,[{value:t,type:"u32"},{value:o,type:"u32"}],"dnrm2-params");let v=E(r,s.getBindGroupLayout(0),[n,f,l,m,p,c,h]),{commandEncoder:_,ts:A}=j(r,s,v,128);M(r,_);let k=E(r,u.getBindGroupLayout(0),[l,m,p,c,g,w]),{commandEncoder:B,ts:L}=j(r,u,k,1);b=G(r,B,g),y=G(r,B,w),M(r,B);let C=S(b,Float32Array),R=S(y,Float32Array);b=null,y=null;let[F,W,V,z]=await Promise.all([P(A),P(L),C,R]),H=dr(V,z)[0];return F!==void 0&&W!==void 0?{nrm2:H,gpuTimeMs:F+W}:{nrm2:H}}finally{!a&&n&&d(n),!a&&f&&d(f),l&&d(l),m&&d(m),p&&d(p),c&&d(c),g&&d(g),w&&d(w),h&&d(h),b&&d(b),y&&d(y)}}async function Zo(r,t,e,o){let a=e instanceof N;if(q(r),T(r,"isamax",{x:e}),!Number.isInteger(t)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(o<=0)throw new Error("incx must be positive.");if(!a&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(t<=0)return{index:0};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let i=await D(r,"isamax"),s=await D(r,"reduction/argmax"),u=null,n=null,f=null,l=null,m=null,p=null;try{u=a?e._buf:x(r,e,"isamax-x",!1),n=tr(r,512,"isamax-partials-val"),f=tr(r,512,"isamax-partials-idx"),l=Br(r,4,"isamax-result"),m=I(r,[{value:t,type:"u32"},{value:o,type:"u32"}],"isamax-params");let c=E(r,i.getBindGroupLayout(0),[u,n,f,m]),{commandEncoder:g,ts:w}=j(r,i,c,128);M(r,g);let h=E(r,s.getBindGroupLayout(0),[n,f,l]),{commandEncoder:b,ts:y}=j(r,s,h,1);p=G(r,b,l),M(r,b);let v=S(p,Uint32Array);p=null;let[_,A,k]=await Promise.all([P(w),P(y),v]),B=k[0];return _!==void 0&&A!==void 0?{index:B,gpuTimeMs:_+A}:{index:B}}finally{!a&&u&&d(u),n&&d(n),f&&d(f),l&&d(l),m&&d(m),p&&d(p)}}async function Qo(r,t,e,o){let a=e instanceof N;if(q(r),T(r,"idamax",{x:e}),!Number.isInteger(t)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(o<=0)throw new Error("incx must be positive.");if(!a&&!(e instanceof Float64Array))throw new Error("x must be a Float64Array or GpuVector.");if(a&&e.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(t<=0)return{index:0};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let i=["f64/dekker","f64/utils/abs","f64/utils/greater","f64/utils/equal"],s=await D(r,[...i,"idamax"],"idamax_main"),u=await D(r,[...i,"reduction/argmaxF64"],"reduce_f64"),n=null,f=null,l=null,m=null,p=null,c=null,g=null,w=null;try{if(a)n=e._buf,f=e._loBuf;else{let{hi:F,lo:W}=rr(e);n=x(r,F,"idamax-xHi",!1),f=x(r,W,"idamax-xLo",!1)}l=tr(r,512,"idamax-partials-val-hi"),m=tr(r,512,"idamax-partials-val-lo"),p=tr(r,512,"idamax-partials-idx"),c=Br(r,4,"idamax-result"),g=I(r,[{value:t,type:"u32"},{value:o,type:"u32"}],"idamax-params");let h=E(r,s.getBindGroupLayout(0),[n,f,l,m,p,g]),{commandEncoder:b,ts:y}=j(r,s,h,128);M(r,b);let v=E(r,u.getBindGroupLayout(0),[l,m,p,c]),{commandEncoder:_,ts:A}=j(r,u,v,1);w=G(r,_,c),M(r,_);let k=S(w,Uint32Array);w=null;let[B,L,C]=await Promise.all([P(y),P(A),k]),R=C[0];return B!==void 0&&L!==void 0?{index:R,gpuTimeMs:B+L}:{index:R}}finally{!a&&n&&d(n),!a&&f&&d(f),l&&d(l),m&&d(m),p&&d(p),c&&d(c),g&&d(g),w&&d(w)}}async function Jo(r,t,e,o,a,i,s,u){let n=e instanceof N,f=a instanceof N;if(q(r),T(r,"srot",{x:e,y:a}),!Number.isInteger(t)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(typeof s!="number")throw new Error("c must be a number.");if(typeof u!="number")throw new Error("s must be a number.");if(Number.isNaN(s)||Number.isNaN(u))throw new Error("c and s must not be NaN.");if(!Number.isFinite(s))throw new Error("c must be finite.");if(!Number.isFinite(u))throw new Error("s must be finite.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!n&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!f&&!(a instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(n!==f)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(t<=0)return n?{}:{x:e,y:a};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(a.length<(t-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let l=await D(r,"srot"),m=null,p=null,c=null,g=null,w=null;try{m=n?e._buf:x(r,e,"srot-x",!0),p=f?a._buf:x(r,a,"srot-y",!0),c=I(r,[{value:t,type:"u32"},{value:s,type:"f32"},{value:u,type:"f32"},{value:o,type:"u32"},{value:i,type:"u32"}],"srot-params");let h=E(r,l.getBindGroupLayout(0),[m,p,c]),{commandEncoder:b,ts:y}=j(r,l,h,cr(r,t));g=n?null:G(r,b,m),w=f?null:G(r,b,p),M(r,b);let v=await P(y);if(n)return v!==void 0?{gpuTimeMs:v}:{};let _=S(g,Float32Array),A=S(w,Float32Array);g=null,w=null;let[k,B]=await Promise.all([_,A]);return v!==void 0?{x:k,y:B,gpuTimeMs:v}:{x:k,y:B}}finally{!n&&m&&d(m),!f&&p&&d(p),c&&d(c),g&&d(g),w&&d(w)}}async function ra(r,t,e,o,a,i,s,u){let n=e instanceof N,f=a instanceof N;if(q(r),T(r,"drot",{x:e,y:a}),!Number.isInteger(t)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(typeof s!="number")throw new Error("c must be a number.");if(typeof u!="number")throw new Error("s must be a number.");if(Number.isNaN(s)||Number.isNaN(u))throw new Error("c and s must not be NaN.");if(!Number.isFinite(s))throw new Error("c must be finite.");if(!Number.isFinite(u))throw new Error("s must be finite.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!(e instanceof Float64Array)&&!n)throw new Error("x must be a Float64Array or GpuVector.");if(!(a instanceof Float64Array)&&!f)throw new Error("y must be a Float64Array or GpuVector.");if(n&&e.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(f&&a.dtype!==Float64Array)throw new Error("y must be a Float64Array-backed GpuVector.");if(n!==f)throw new Error("x and y must be the same type (both Float64Array or both GpuVector).");if(t<=0)return n?{}:{x:e,y:a};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(a.length<(t-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let m=await D(r,[...["f64/dekker","f64/utils/add","f64/utils/multiply"],"drot"]),{hi:p,lo:c}=rr(new Float64Array([s])),{hi:g,lo:w}=rr(new Float64Array([u])),h=null,b=null,y=null,v=null,_=null,A=null,k=null,B=null,L=null;try{if(n)h=e._buf,b=e._loBuf,y=a._buf,v=a._loBuf;else{let ir=rr(e),lr=rr(a);h=x(r,ir.hi,"drot-xHi",!0),b=x(r,ir.lo,"drot-xLo",!0),y=x(r,lr.hi,"drot-yHi",!0),v=x(r,lr.lo,"drot-yLo",!0)}_=I(r,[{value:t,type:"u32"},{value:p[0],type:"f32"},{value:c[0],type:"f32"},{value:g[0],type:"f32"},{value:w[0],type:"f32"},{value:o,type:"u32"},{value:i,type:"u32"}],"drot-params");let C=E(r,m.getBindGroupLayout(0),[h,b,y,v,_]),{commandEncoder:R,ts:F}=j(r,m,C,cr(r,t));A=n?null:G(r,R,h),k=n?null:G(r,R,b),B=f?null:G(r,R,y),L=f?null:G(r,R,v),M(r,R);let W=await P(F);if(n)return W!==void 0?{gpuTimeMs:W}:{};let V=await S(A,Float32Array);A=null;let z=await S(k,Float32Array);k=null;let H=await S(B,Float32Array);B=null;let $=await S(L,Float32Array);L=null;let K=dr(V,z),Y=dr(H,$);return W!==void 0?{x:K,y:Y,gpuTimeMs:W}:{x:K,y:Y}}finally{!n&&h&&d(h),!n&&b&&d(b),!f&&y&&d(y),!f&&v&&d(v),_&&d(_),A&&d(A),k&&d(k),B&&d(B),L&&d(L)}}async function ea(r,t,e,o,a,i,s){let u=e instanceof N,n=a instanceof N;if(q(r),T(r,"srotm",{x:e,y:a}),!Number.isInteger(t)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(!(s instanceof Float32Array)||s.length!==5)throw new Error("param must be a Float32Array of length 5.");if(s[0]!==-2&&s[0]!==-1&&s[0]!==0&&s[0]!==1)throw new Error("param[0] (flag) must be one of -2, -1, 0, or 1.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!u&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!n&&!(a instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(u!==n)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(t<=0||s[0]===-2)return u?{}:{x:e,y:a};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(a.length<(t-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let f=await D(r,"srotm"),l=null,m=null,p=null,c=null,g=null,w=null;try{l=u?e._buf:x(r,e,"srotm-x",!0),m=n?a._buf:x(r,a,"srotm-y",!0),p=x(r,s,"srotm-param",!1),c=I(r,[{value:t,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"srotm-params");let h=E(r,f.getBindGroupLayout(0),[l,m,p,c]),{commandEncoder:b,ts:y}=j(r,f,h,cr(r,t));g=u?null:G(r,b,l),w=n?null:G(r,b,m),M(r,b);let v=await P(y);if(u)return v!==void 0?{gpuTimeMs:v}:{};let _=S(g,Float32Array),A=S(w,Float32Array);g=null,w=null;let[k,B]=await Promise.all([_,A]);return v!==void 0?{x:k,y:B,gpuTimeMs:v}:{x:k,y:B}}finally{!u&&l&&d(l),!n&&m&&d(m),p&&d(p),c&&d(c),g&&d(g),w&&d(w)}}async function ta(r,t,e,o,a,i,s){let u=e instanceof N,n=a instanceof N;if(q(r),T(r,"drotm",{x:e,y:a}),!Number.isInteger(t)||!Number.isInteger(o)||!Number.isInteger(i))throw new Error("n, incx, and incy must be integers.");if(!(s instanceof Float64Array)||s.length!==5)throw new Error("param must be a Float64Array of length 5.");if(s[0]!==-2&&s[0]!==-1&&s[0]!==0&&s[0]!==1)throw new Error("param[0] (flag) must be one of -2, -1, 0, or 1.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!(e instanceof Float64Array)&&!u)throw new Error("x must be a Float64Array or GpuVector.");if(!(a instanceof Float64Array)&&!n)throw new Error("y must be a Float64Array or GpuVector.");if(u&&e.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(n&&a.dtype!==Float64Array)throw new Error("y must be a Float64Array-backed GpuVector.");if(u!==n)throw new Error("x and y must be the same type (both Float64Array or both GpuVector).");if(t<=0||s[0]===-2)return u?{}:{x:e,y:a};if(e.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(a.length<(t-1)*i+1)throw new Error("y does not have enough elements for the given n and incy.");let l=await D(r,[...["f64/dekker","f64/utils/add","f64/utils/multiply"],"drotm"]),{hi:m,lo:p}=rr(s),c=null,g=null,w=null,h=null,b=null,y=null,v=null,_=null,A=null,k=null,B=null;try{if(u)c=e._buf,g=e._loBuf,w=a._buf,h=a._loBuf;else{let Y=rr(e),ir=rr(a);c=x(r,Y.hi,"drotm-xHi",!0),g=x(r,Y.lo,"drotm-xLo",!0),w=x(r,ir.hi,"drotm-yHi",!0),h=x(r,ir.lo,"drotm-yLo",!0)}b=x(r,m,"drotm-paramHi",!1),y=x(r,p,"drotm-paramLo",!1),v=I(r,[{value:t,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"drotm-params");let L=E(r,l.getBindGroupLayout(0),[c,g,w,h,b,y,v]),{commandEncoder:C,ts:R}=j(r,l,L,cr(r,t));_=u?null:G(r,C,c),A=u?null:G(r,C,g),k=n?null:G(r,C,w),B=n?null:G(r,C,h),M(r,C);let F=await P(R);if(u)return F!==void 0?{gpuTimeMs:F}:{};let W=await S(_,Float32Array);_=null;let V=await S(A,Float32Array);A=null;let z=await S(k,Float32Array);k=null;let H=await S(B,Float32Array);B=null;let $=dr(W,V),K=dr(z,H);return F!==void 0?{x:$,y:K,gpuTimeMs:F}:{x:$,y:K}}finally{!u&&c&&d(c),!u&&g&&d(g),!n&&w&&d(w),!n&&h&&d(h),b&&d(b),y&&d(y),v&&d(v),_&&d(_),A&&d(A),k&&d(k),B&&d(B)}}async function oa(r,t,e,o,a,i,s,u,n,f,l,m,p="row-major"){let c=i instanceof X,g=u instanceof N,w=l instanceof N;if(q(r),T(r,"sgemv",{A:i,x:u,y:l}),t!=="no-transpose"&&t!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(p!=="row-major"&&p!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof a!="number")throw new Error("alpha must be a number.");if(Number.isNaN(a))throw new Error("alpha must not be NaN.");if(!Number.isFinite(a))throw new Error("alpha must be finite.");if(typeof f!="number")throw new Error("beta must be a number.");if(Number.isNaN(f))throw new Error("beta must not be NaN.");if(!Number.isFinite(f))throw new Error("beta must be finite.");if(!Number.isInteger(e)||!Number.isInteger(o)||!Number.isInteger(n)||!Number.isInteger(m)||!Number.isInteger(s))throw new Error("m, n, incx, incy, and lda must be integers.");if(n<=0||m<=0)throw new Error("incx and incy must be positive.");if(!c&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!g&&!(u instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!w&&!(l instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(g!==w)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(g&&!c)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(c&&!g)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(g&&u._buf===l._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(c&&w&&i._buf===l._buf)throw new Error("A and y must not reference the same GPU buffer.");if(c&&s!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(c&&(i.rows<e||i.cols<o))throw new Error("A is too small for the given m and n.");if(e<0||o<0)throw new Error("m and n must be non-negative.");if(e===0||o===0)return w?{}:{y:l};(c?i.layout:p)==="column-major"&&([e,o]=[o,e],t=t==="no-transpose"?"transpose":"no-transpose");let b=t==="no-transpose",y=b?o:e,v=b?e:o;if(s<o)throw new Error("lda must be >= n.");if(!c&&i.length<(e-1)*s+o)throw new Error("A does not have enough elements for the given m, n, and lda.");if(u.length<(y-1)*n+1)throw new Error("x does not have enough elements for the given dimensions and incx.");if(l.length<(v-1)*m+1)throw new Error("y does not have enough elements for the given dimensions and incy.");let A=await D(r,b?"sgemv_n":"sgemv_t"),k=null,B=null,L=null,C=null;try{k=c?i._buf:x(r,i,"sgemv-A",!1),B=g?u._buf:x(r,u,"sgemv-x",!1),L=w?l._buf:x(r,l,"sgemv-y",!0),C=I(r,[{value:e,type:"u32"},{value:o,type:"u32"},{value:a,type:"f32"},{value:f,type:"f32"},{value:n,type:"u32"},{value:m,type:"u32"},{value:s,type:"u32"}],"sgemv-params");let R=E(r,A.getBindGroupLayout(0),[k,B,L,C]),F=b?Math.min(e,r.limits.maxComputeWorkgroupsPerDimension):Zr(r,"sgemv",v),{commandEncoder:W,ts:V}=j(r,A,R,F),z=w?null:G(r,W,L);M(r,W);let H=await P(V);if(w)return H!==void 0?{gpuTimeMs:H}:{};let $=await S(z,Float32Array);return H!==void 0?{y:$,gpuTimeMs:H}:{y:$}}finally{!c&&k&&d(k),!g&&B&&d(B),!w&&L&&d(L),C&&d(C)}}async function aa(r,t,e,o,a,i,s,u,n,f,l,m="row-major"){let p=s instanceof N,c=f instanceof N,g=a instanceof X;if(q(r),T(r,"ssymv",{A:a,x:s,y:f}),t!=="lower"&&t!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(m!=="row-major"&&m!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(e)||!Number.isInteger(u)||!Number.isInteger(l)||!Number.isInteger(i))throw new Error("n, incx, incy, and lda must be integers.");if(typeof o!="number")throw new Error("alpha must be a number.");if(Number.isNaN(o))throw new Error("alpha must not be NaN.");if(!Number.isFinite(o))throw new Error("alpha must be finite.");if(typeof n!="number")throw new Error("beta must be a number.");if(Number.isNaN(n))throw new Error("beta must not be NaN.");if(!Number.isFinite(n))throw new Error("beta must be finite.");if(u<=0||l<=0)throw new Error("incx and incy must be positive.");if(i<e)throw new Error("lda must be >= n.");if(!g&&!(a instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!p&&!(s instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!c&&!(f instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(p!==c)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(p&&!g)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(g&&!p)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(p&&s._buf===f._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(g&&i!==a.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(g&&(a.rows<e||a.cols<e))throw new Error("A is too small for the given n.");if(e<0)throw new Error("n must be non-negative.");if(e===0)return c?{}:{y:f};if(!g&&a.length<(e-1)*i+e)throw new Error("A does not have enough elements for the given n and lda.");if(s.length<(e-1)*u+1)throw new Error("x does not have enough elements for the given n and incx.");if(f.length<(e-1)*l+1)throw new Error("y does not have enough elements for the given n and incy.");let h=(g?a.layout:m)==="column-major"?t==="upper":t==="lower",b=await D(r,"ssymv"),y=null,v=null,_=null,A=null;try{y=g?a._buf:x(r,a,"ssymv-A",!1),v=p?s._buf:x(r,s,"ssymv-x",!1),_=c?f._buf:x(r,f,"ssymv-y",!0),A=I(r,[{value:e,type:"u32"},{value:o,type:"f32"},{value:n,type:"f32"},{value:u,type:"u32"},{value:l,type:"u32"},{value:i,type:"u32"},{value:h?0:1,type:"u32"}],"ssymv-params");let k=E(r,b.getBindGroupLayout(0),[y,v,_,A]),B=Math.min(e,r.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:L,ts:C}=j(r,b,k,B),R=c?null:G(r,L,_);M(r,L);let F=await P(C);if(c)return F!==void 0?{gpuTimeMs:F}:{};let W=await S(R,Float32Array);return F!==void 0?{y:W,gpuTimeMs:F}:{y:W}}finally{!g&&y&&d(y),!p&&v&&d(v),!c&&_&&d(_),A&&d(A)}}async function ia(r,t,e,o,a,i,s,u,n,f,l,m="row-major"){let p=u instanceof N,c=f instanceof N,g=i instanceof X,w=o==="unit";if(q(r),T(r,"strmv",{A:i,x:u,y:f}),t!=="lower"&&t!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(e!=="no-transpose"&&e!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(!w&&o!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(m!=="row-major"&&m!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(a)||!Number.isInteger(n)||!Number.isInteger(l)||!Number.isInteger(s))throw new Error("n, incx, incy, and lda must be integers.");if(n<=0||l<=0)throw new Error("incx and incy must be positive.");if(s<a)throw new Error("lda must be >= n.");if(!g&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!p&&!(u instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!c&&!(f instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(p!==c)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(p&&u._buf===f._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(p&&!g)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(g&&!p)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(g&&c&&i._buf===f._buf)throw new Error("A and y must not reference the same GPU buffer.");if(g&&s!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(g&&(i.rows<a||i.cols<a))throw new Error("A is too small for the given n.");if(a<0)throw new Error("n must be non-negative.");if(a===0)return c?{}:{y:f};if(!g&&i.length<(a-1)*s+a)throw new Error("A does not have enough elements for the given n and lda.");if(u.length<(a-1)*n+1)throw new Error("x does not have enough elements for the given n and incx.");if(f.length<(a-1)*l+1)throw new Error("y does not have enough elements for the given n and incy.");let b=(g?i.layout:m)==="column-major",y=b?t==="upper":t==="lower",v=b?e==="transpose":e==="no-transpose",_=await D(r,"strmv"),A=null,k=null,B=null,L=null;try{A=g?i._buf:x(r,i,"strmv-A",!1),k=p?u._buf:x(r,u,"strmv-x",!1),B=c?f._buf:x(r,f,"strmv-y",!0),L=I(r,[{value:a,type:"u32"},{value:n,type:"u32"},{value:l,type:"u32"},{value:s,type:"u32"},{value:v?0:1,type:"u32"},{value:y?0:1,type:"u32"},{value:w?1:0,type:"u32"}],"strmv-params");let C=E(r,_.getBindGroupLayout(0),[A,k,B,L]),R=Math.min(a,r.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:F,ts:W}=j(r,_,C,R),V=c?null:G(r,F,B);M(r,F);let z=await P(W);if(c)return z!==void 0?{gpuTimeMs:z}:{};let H=await S(V,Float32Array);return z!==void 0?{y:H,gpuTimeMs:z}:{y:H}}finally{!g&&A&&d(A),!p&&k&&d(k),!c&&B&&d(B),L&&d(L)}}function sa(r,t,e){let o=new ArrayBuffer(r*t),a=new DataView(o);for(let i=0;i<r;i++){let s=e(i),u=i*t;s.forEach((n,f)=>a.setUint32(u+f*4,n,!0))}return o}function na(r,t,e){let o=r.createBuffer({label:e,size:t.byteLength,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST});return r.queue.writeBuffer(o,0,t),o}async function la(r,t,e,o,a,i,s,u,n,f="row-major"){let l=u instanceof N,m=i instanceof X,p=o==="unit";if(q(r),T(r,"strsv",{A:i,x:u}),t!=="lower"&&t!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(e!=="no-transpose"&&e!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(!p&&o!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(f!=="row-major"&&f!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(a)||!Number.isInteger(n)||!Number.isInteger(s))throw new Error("n, incx, and lda must be integers.");if(n<=0)throw new Error("incx must be positive.");if(s<a)throw new Error("lda must be >= n.");if(!m&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!l&&!(u instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(l&&!m)throw new Error("A must be a GpuMatrix when x is a GpuVector.");if(m&&!l)throw new Error("x must be a GpuVector when A is a GpuMatrix.");if(m&&l&&i._buf===u._buf)throw new Error("A and x must not reference the same GPU buffer.");if(m&&s!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(m&&(i.rows<a||i.cols<a))throw new Error("A is too small for the given n.");if(a<0)throw new Error("n must be non-negative.");if(a===0)return l?{}:{x:u};if(!m&&i.length<(a-1)*s+a)throw new Error("A does not have enough elements for the given n and lda.");if(u.length<(a-1)*n+1)throw new Error("x does not have enough elements for the given n and incx.");let g=(m?i.layout:f)==="column-major",w=g?t==="upper":t==="lower",h=g?e==="transpose":e==="no-transpose",b=await D(r,"strsv_invert_block"),y=await D(r,"strsv_apply_inverse"),v=await D(r,"strsv_update"),_=h===w,A=[];for(let H=0;H<a;H+=64)A.push(H);_||A.reverse();let k=A.length,B=r.limits.maxComputeWorkgroupsPerDimension,L=r.limits.minUniformBufferOffsetAlignment,C=null,R=null,F=null,W=null,V=null,z=null;try{C=m?i._buf:x(r,i,"strsv-A",!1),R=l?u._buf:x(r,u,"strsv-x",!0),F=tr(r,k*64*64*4,"strsv-Ainv");let H=sa(k,L,er=>{let Z=er*64,J=Math.min(Z+64,a);return[n,er,Z,J]});W=na(r,H,"strsv-apply-params");let $=sa(k,L,er=>{let Z=er*64,J=Math.min(Z+64,a);return[a,n,s,h?0:1,w?0:1,Z,J]});V=na(r,$,"strsv-update-params");let{commandEncoder:K,querySet:Y}=qr(r);z=I(r,[{value:a,type:"u32"},{value:s,type:"u32"},{value:h?0:1,type:"u32"},{value:w?0:1,type:"u32"},{value:p?1:0,type:"u32"}],"strsv-invert-params");let ir=E(r,b.getBindGroupLayout(0),[C,F,z]);gr(K,b,ir,{x:64,y:k},Y?{timestampWrites:{querySet:Y,beginningOfPassWriteIndex:0}}:void 0);for(let er=0;er<A.length;er++){let Z=A[er],J=Math.min(Z+64,a),sr=Z/64,pr=er===A.length-1,br=sr*L,fr=E(r,y.getBindGroupLayout(0),[F,R,{buffer:W,offset:br,size:16}]);gr(K,y,fr,1,pr&&Y?{timestampWrites:{querySet:Y,endOfPassWriteIndex:1}}:void 0);let yr=_?a-J:Z;if(yr===0)continue;let Cr=E(r,v.getBindGroupLayout(0),[C,R,{buffer:V,offset:br,size:32}]),Rr=Math.min(yr,B);gr(K,v,Cr,Rr)}let hr=Lr(r,K,Y),nr=l?null:G(r,K,R);M(r,K);let mr=await P(hr);if(l)return mr!==void 0?{gpuTimeMs:mr}:{};let Q=await S(nr,Float32Array);return mr!==void 0?{x:Q,gpuTimeMs:mr}:{x:Q}}finally{!m&&C&&d(C),!l&&R&&d(R),F&&d(F),W&&d(W),V&&d(V),z&&d(z)}}async function ua(r,t,e,o,a,i,s,u,n,f,l="row-major"){let m=n instanceof X;if(q(r),T(r,"sger",{A:n,x:a,y:s}),l!=="row-major"&&l!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof o!="number")throw new Error("alpha must be a number.");if(Number.isNaN(o))throw new Error("alpha must not be NaN.");if(!Number.isFinite(o))throw new Error("alpha must be finite.");if(!Number.isInteger(t)||!Number.isInteger(e)||!Number.isInteger(i)||!Number.isInteger(u)||!Number.isInteger(f))throw new Error("m, n, incx, incy, and lda must be integers.");if(i<=0||u<=0)throw new Error("incx and incy must be positive.");if(!m&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(m&&f!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(m&&(n.rows<t||n.cols<e))throw new Error("A is too small for the given m and n.");(m?n.layout:l)==="column-major"&&([t,e]=[e,t],[a,s]=[s,a],[i,u]=[u,i]);let c=a instanceof N,g=s instanceof N;if(f<e)throw new Error("lda must be >= n.");if(!c&&!(a instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!g&&!(s instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(c!==g)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(c&&!m)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(m&&!c)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(m&&c&&n._buf===a._buf)throw new Error("A and x must not reference the same GPU buffer.");if(m&&g&&n._buf===s._buf)throw new Error("A and y must not reference the same GPU buffer.");if(t<0||e<0)throw new Error("m and n must be non-negative.");if(t===0||e===0)return m?{}:{A:n};if(!m&&n.length<(t-1)*f+e)throw new Error("A does not have enough elements for the given m, n, and lda.");if(a.length<(t-1)*i+1)throw new Error("x does not have enough elements for the given m and incx.");if(s.length<(e-1)*u+1)throw new Error("y does not have enough elements for the given n and incy.");let w=await D(r,"sger"),h=null,b=null,y=null,v=null;try{h=c?a._buf:x(r,a,"sger-x",!1),b=g?s._buf:x(r,s,"sger-y",!1),y=m?n._buf:x(r,n,"sger-A",!0),v=I(r,[{value:t,type:"u32"},{value:e,type:"u32"},{value:o,type:"f32"},{value:i,type:"u32"},{value:u,type:"u32"},{value:f,type:"u32"}],"sger-params");let _=E(r,w.getBindGroupLayout(0),[h,b,y,v]),A=Math.min(t,r.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:k,ts:B}=j(r,w,_,A),L=m?null:G(r,k,y);M(r,k);let C=await P(B);if(m)return C!==void 0?{gpuTimeMs:C}:{};let R=await S(L,Float32Array);return C!==void 0?{A:R,gpuTimeMs:C}:{A:R}}finally{!c&&h&&d(h),!g&&b&&d(b),!m&&y&&d(y),v&&d(v)}}async function fa(r,t,e,o,a,i,s,u,n="row-major"){let f=a instanceof N,l=s instanceof X;if(q(r),T(r,"ssyr",{A:s,x:a}),t!=="lower"&&t!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(n!=="row-major"&&n!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(e)||!Number.isInteger(i)||!Number.isInteger(u))throw new Error("n, incx, and lda must be integers.");if(typeof o!="number")throw new Error("alpha must be a number.");if(Number.isNaN(o))throw new Error("alpha must not be NaN.");if(!Number.isFinite(o))throw new Error("alpha must be finite.");if(i<=0)throw new Error("incx must be positive.");if(u<e)throw new Error("lda must be >= n.");if(!l&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!f&&!(a instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(f&&!l)throw new Error("A must be a GpuMatrix when x is a GpuVector.");if(l&&!f)throw new Error("x must be a GpuVector when A is a GpuMatrix.");if(l&&f&&s._buf===a._buf)throw new Error("A and x must not reference the same GPU buffer.");if(l&&u!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(l&&(s.rows<e||s.cols<e))throw new Error("A is too small for the given n.");if(e<0)throw new Error("n must be non-negative.");if(e===0)return l?{}:{A:s};if(!l&&s.length<(e-1)*u+e)throw new Error("A does not have enough elements for the given n and lda.");if(a.length<(e-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");let p=(l?s.layout:n)==="column-major"?t==="upper":t==="lower",c=await D(r,"ssyr"),g=null,w=null,h=null;try{g=f?a._buf:x(r,a,"ssyr-x",!1),w=l?s._buf:x(r,s,"ssyr-A",!0),h=I(r,[{value:e,type:"u32"},{value:o,type:"f32"},{value:i,type:"u32"},{value:u,type:"u32"},{value:p?0:1,type:"u32"}],"ssyr-params");let b=E(r,c.getBindGroupLayout(0),[g,w,h]),y=Math.min(e,r.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:v,ts:_}=j(r,c,b,y),A=l?null:G(r,v,w);M(r,v);let k=await P(_);if(l)return k!==void 0?{gpuTimeMs:k}:{};let B=await S(A,Float32Array);return k!==void 0?{A:B,gpuTimeMs:k}:{A:B}}finally{!f&&g&&d(g),!l&&w&&d(w),h&&d(h)}}async function ma(r,t,e,o,a,i,s,u,n,f,l="row-major"){let m=a instanceof N,p=s instanceof N,c=n instanceof X;if(q(r),T(r,"ssyr2",{A:n,x:a,y:s}),t!=="lower"&&t!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(l!=="row-major"&&l!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(e)||!Number.isInteger(i)||!Number.isInteger(u)||!Number.isInteger(f))throw new Error("n, incx, incy, and lda must be integers.");if(typeof o!="number")throw new Error("alpha must be a number.");if(Number.isNaN(o))throw new Error("alpha must not be NaN.");if(!Number.isFinite(o))throw new Error("alpha must be finite.");if(i<=0||u<=0)throw new Error("incx and incy must be positive.");if(f<e)throw new Error("lda must be >= n.");if(!c&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!m&&!(a instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!p&&!(s instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(m!==p)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(m&&!c)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(c&&!m)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(c&&m&&n._buf===a._buf)throw new Error("A and x must not reference the same GPU buffer.");if(c&&p&&n._buf===s._buf)throw new Error("A and y must not reference the same GPU buffer.");if(m&&a._buf===s._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(c&&f!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(c&&(n.rows<e||n.cols<e))throw new Error("A is too small for the given n.");if(e<0)throw new Error("n must be non-negative.");if(e===0)return c?{}:{A:n};if(!c&&n.length<(e-1)*f+e)throw new Error("A does not have enough elements for the given n and lda.");if(a.length<(e-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");if(s.length<(e-1)*u+1)throw new Error("y does not have enough elements for the given n and incy.");let w=(c?n.layout:l)==="column-major"?t==="upper":t==="lower",h=await D(r,"ssyr2"),b=null,y=null,v=null,_=null;try{b=m?a._buf:x(r,a,"ssyr2-x",!1),y=p?s._buf:x(r,s,"ssyr2-y",!1),v=c?n._buf:x(r,n,"ssyr2-A",!0),_=I(r,[{value:e,type:"u32"},{value:o,type:"f32"},{value:i,type:"u32"},{value:u,type:"u32"},{value:f,type:"u32"},{value:w?0:1,type:"u32"}],"ssyr2-params");let A=E(r,h.getBindGroupLayout(0),[b,y,v,_]),k=Math.min(e,r.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:B,ts:L}=j(r,h,A,k),C=c?null:G(r,B,v);M(r,B);let R=await P(L);if(c)return R!==void 0?{gpuTimeMs:R}:{};let F=await S(C,Float32Array);return R!==void 0?{A:F,gpuTimeMs:R}:{A:F}}finally{!m&&b&&d(b),!p&&y&&d(y),!c&&v&&d(v),_&&d(_)}}async function da(r,t,e,o,a,i,s,u,n,f,l,m,p,c,g="row-major"){let w=u instanceof X,h=f instanceof X,b=p instanceof X;if(q(r),T(r,"sgemm",{A:u,B:f,C:p}),t!=="no-transpose"&&t!=="transpose")throw new Error("transA must be 'no-transpose' or 'transpose'.");if(e!=="no-transpose"&&e!=="transpose")throw new Error("transB must be 'no-transpose' or 'transpose'.");if(g!=="row-major"&&g!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof s!="number")throw new Error("alpha must be a number.");if(Number.isNaN(s))throw new Error("alpha must not be NaN.");if(!Number.isFinite(s))throw new Error("alpha must be finite.");if(typeof m!="number")throw new Error("beta must be a number.");if(Number.isNaN(m))throw new Error("beta must not be NaN.");if(!Number.isFinite(m))throw new Error("beta must be finite.");if(!Number.isInteger(o)||!Number.isInteger(a)||!Number.isInteger(i)||!Number.isInteger(n)||!Number.isInteger(l)||!Number.isInteger(c))throw new Error("m, n, k, lda, ldb, and ldc must be integers.");if(!w&&!(u instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!h&&!(f instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(!b&&!(p instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if((w||h)&&!b)throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");if(b&&(!w||!h))throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");if(o<0||a<0||i<0)throw new Error("m, n, and k must be non-negative.");if(n<=0||l<=0||c<=0)throw new Error("lda, ldb, and ldc must be positive.");if(o===0||a===0)return b?{}:{C:p};let y=w?u.layout:g,v=h?f.layout:g,_=b?p.layout:g,A=y==="column-major"?i:o,k=y==="column-major"?o:i,B=t==="no-transpose"?A:k,L=t==="no-transpose"?k:A;if(n<L)throw new Error(`lda must be >= ${y==="column-major"?"rows":"cols"} of A as stored.`);if(w){if(n!==u.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[J,sr]=t==="no-transpose"?[o,i]:[i,o];if(u.rows<J||u.cols<sr)throw new Error("A is too small for the given m, k, and transA.")}else if(u.length<(B-1)*n+L)throw new Error("A does not have enough elements for the given dimensions and lda.");let C=v==="column-major"?a:i,R=v==="column-major"?i:a,F=e==="no-transpose"?C:R,W=e==="no-transpose"?R:C;if(l<W)throw new Error(`ldb must be >= ${v==="column-major"?"rows":"cols"} of B as stored.`);if(h){if(l!==f.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");let[J,sr]=e==="no-transpose"?[i,a]:[a,i];if(f.rows<J||f.cols<sr)throw new Error("B is too small for the given n, k, and transB.")}else if(f.length<(F-1)*l+W)throw new Error("B does not have enough elements for the given dimensions and ldb.");let V=_==="column-major"?a:o,z=_==="column-major"?o:a;if(c<z)throw new Error(`ldc must be >= ${_==="column-major"?"rows":"cols"} of C as stored.`);if(b){if(c!==p.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(p.rows<o||p.cols<a)throw new Error("C is too small for the given m and n.")}else if(p.length<(V-1)*c+z)throw new Error("C does not have enough elements for the given dimensions and ldc.");y==="column-major"&&(t=t==="no-transpose"?"transpose":"no-transpose"),v==="column-major"&&(e=e==="no-transpose"?"transpose":"no-transpose"),_==="column-major"&&([u,f]=[f,u],[w,h]=[h,w],[n,l]=[l,n],[t,e]=[e==="no-transpose"?"transpose":"no-transpose",t==="no-transpose"?"transpose":"no-transpose"],[o,a]=[a,o]);let H=Math.ceil(a/64),$=Math.ceil(o/64),K=H*$>=36,Y=await D(r,K?"sgemm_large":"sgemm_small"),ir=w?u._buf:x(r,u,"sgemm-A",!1),lr=h?f._buf:x(r,f,"sgemm-B",!1),hr=b?p._buf:x(r,p,"sgemm-C",!0),nr=t==="no-transpose",mr=e==="no-transpose",Q=nr&&Ee(ir,n,o,i),er=Ee(lr,l,mr?i:a,mr?a:i),Z=I(r,[{value:o,type:"u32"},{value:a,type:"u32"},{value:i,type:"u32"},{value:s,type:"f32"},{value:m,type:"f32"},{value:n,type:"u32"},{value:l,type:"u32"},{value:c,type:"u32"},{value:t==="transpose"?1:0,type:"u32"},{value:e==="transpose"?1:0,type:"u32"},{value:Q?1:0,type:"u32"},{value:er?1:0,type:"u32"}],"sgemm-params");try{let J=E(r,Y.getBindGroupLayout(0),[ir,Sr(r,ir),lr,Sr(r,lr),hr,Z]),sr=K?{x:U(r,H,"sgemm","x"),y:U(r,$,"sgemm","y")}:{x:U(r,Math.ceil(a/32),"sgemm","x"),y:U(r,Math.ceil(o/32),"sgemm","y")},{commandEncoder:pr,ts:br}=j(r,Y,J,sr),fr=b?null:G(r,pr,hr);M(r,pr);let ur=await P(br);if(b)return ur!==void 0?{gpuTimeMs:ur}:{};let yr=await S(fr,Float32Array);return ur!==void 0?{C:yr,gpuTimeMs:ur}:{C:yr}}finally{w||d(ir),h||d(lr),b||d(hr),d(Z)}}async function ca(r,t,e,o,a,i,s,u,n,f,l,m,p,c,g,w="row-major"){let h=n instanceof X,b=l instanceof X,y=c instanceof X;if(q(r),T(r,"sgemmtr",{A:n,B:l,C:c}),t!=="lower"&&t!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(e!=="no-transpose"&&e!=="transpose")throw new Error("transA must be 'no-transpose' or 'transpose'.");if(o!=="no-transpose"&&o!=="transpose")throw new Error("transB must be 'no-transpose' or 'transpose'.");if(w!=="row-major"&&w!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof u!="number")throw new Error("alpha must be a number.");if(Number.isNaN(u))throw new Error("alpha must not be NaN.");if(!Number.isFinite(u))throw new Error("alpha must be finite.");if(typeof p!="number")throw new Error("beta must be a number.");if(Number.isNaN(p))throw new Error("beta must not be NaN.");if(!Number.isFinite(p))throw new Error("beta must be finite.");if(!Number.isInteger(a)||!Number.isInteger(i)||!Number.isInteger(s)||!Number.isInteger(f)||!Number.isInteger(m)||!Number.isInteger(g))throw new Error("m, n, k, lda, ldb, and ldc must be integers.");if(!h&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!b&&!(l instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(!y&&!(c instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if((h||b)&&!y)throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");if(y&&(!h||!b))throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");if(a<0||i<0||s<0)throw new Error("m, n, and k must be non-negative.");if(f<=0||m<=0||g<=0)throw new Error("lda, ldb, and ldc must be positive.");if(a===0||i===0)return y?{}:{C:c};let v=h?n.layout:w,_=b?l.layout:w,A=y?c.layout:w,k=v==="column-major"?s:a,B=v==="column-major"?a:s,L=e==="no-transpose"?k:B,C=e==="no-transpose"?B:k;if(f<C)throw new Error(`lda must be >= ${v==="column-major"?"rows":"cols"} of A as stored.`);if(h){if(f!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[Q,er]=e==="no-transpose"?[a,s]:[s,a];if(n.rows<Q||n.cols<er)throw new Error("A is too small for the given m, k, and transA.")}else if(n.length<(L-1)*f+C)throw new Error("A does not have enough elements for the given dimensions and lda.");let R=_==="column-major"?i:s,F=_==="column-major"?s:i,W=o==="no-transpose"?R:F,V=o==="no-transpose"?F:R;if(m<V)throw new Error(`ldb must be >= ${_==="column-major"?"rows":"cols"} of B as stored.`);if(b){if(m!==l.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");let[Q,er]=o==="no-transpose"?[s,i]:[i,s];if(l.rows<Q||l.cols<er)throw new Error("B is too small for the given n, k, and transB.")}else if(l.length<(W-1)*m+V)throw new Error("B does not have enough elements for the given dimensions and ldb.");let z=A==="column-major"?i:a,H=A==="column-major"?a:i;if(g<H)throw new Error(`ldc must be >= ${A==="column-major"?"rows":"cols"} of C as stored.`);if(y){if(g!==c.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(c.rows<a||c.cols<i)throw new Error("C is too small for the given m and n.")}else if(c.length<(z-1)*g+H)throw new Error("C does not have enough elements for the given dimensions and ldc.");v==="column-major"&&(e=e==="no-transpose"?"transpose":"no-transpose"),_==="column-major"&&(o=o==="no-transpose"?"transpose":"no-transpose"),A==="column-major"&&([n,l]=[l,n],[h,b]=[b,h],[f,m]=[m,f],[e,o]=[o==="no-transpose"?"transpose":"no-transpose",e==="no-transpose"?"transpose":"no-transpose"],[a,i]=[i,a],t=t==="lower"?"upper":"lower");let $=Math.ceil(i/64),K=Math.ceil(a/64),Y=$*K>=36,ir=await D(r,Y?"sgemmtr_large":"sgemmtr_small"),lr=h?n._buf:x(r,n,"sgemmtr-A",!1),hr=b?l._buf:x(r,l,"sgemmtr-B",!1),nr=y?c._buf:x(r,c,"sgemmtr-C",!0),mr=I(r,[{value:a,type:"u32"},{value:i,type:"u32"},{value:s,type:"u32"},{value:u,type:"f32"},{value:p,type:"f32"},{value:f,type:"u32"},{value:m,type:"u32"},{value:g,type:"u32"},{value:e==="transpose"?1:0,type:"u32"},{value:o==="transpose"?1:0,type:"u32"},{value:t==="upper"?1:0,type:"u32"}],"sgemmtr-params");try{let Q=E(r,ir.getBindGroupLayout(0),[lr,hr,nr,mr]),er=Y?{x:U(r,$,"sgemmtr","x"),y:U(r,K,"sgemmtr","y")}:{x:U(r,Math.ceil(i/32),"sgemmtr","x"),y:U(r,Math.ceil(a/32),"sgemmtr","y")},{commandEncoder:Z,ts:J}=j(r,ir,Q,er),sr=y?null:G(r,Z,nr);M(r,Z);let pr=await P(J);if(y)return pr!==void 0?{gpuTimeMs:pr}:{};let br=await S(sr,Float32Array);return pr!==void 0?{C:br,gpuTimeMs:pr}:{C:br}}finally{h||d(lr),b||d(hr),y||d(nr),d(mr)}}async function pa(r,t,e,o,a,i,s,u,n,f,l,m="row-major"){let p=s instanceof X,c=f instanceof X;if(q(r),T(r,"ssyrk",{A:s,C:f}),t!=="lower"&&t!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(e!=="no-transpose"&&e!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(m!=="row-major"&&m!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof i!="number")throw new Error("alpha must be a number.");if(Number.isNaN(i))throw new Error("alpha must not be NaN.");if(!Number.isFinite(i))throw new Error("alpha must be finite.");if(typeof n!="number")throw new Error("beta must be a number.");if(Number.isNaN(n))throw new Error("beta must not be NaN.");if(!Number.isFinite(n))throw new Error("beta must be finite.");if(!Number.isInteger(o)||!Number.isInteger(a)||!Number.isInteger(u)||!Number.isInteger(l))throw new Error("n, k, lda, and ldc must be integers.");if(!p&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!c&&!(f instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if(p&&!c)throw new Error("C must be a GpuMatrix when A is a GpuMatrix.");if(c&&!p)throw new Error("A must be a GpuMatrix when C is a GpuMatrix.");if(o<0||a<0)throw new Error("n and k must be non-negative.");if(u<=0||l<=0)throw new Error("lda and ldc must be positive.");if(o===0)return c?{}:{C:f};let g=p?s.layout:m,w=c?f.layout:m,h=g==="column-major"?a:o,b=g==="column-major"?o:a,y=e==="no-transpose"?h:b,v=e==="no-transpose"?b:h;if(u<v)throw new Error(`lda must be >= ${g==="column-major"?"rows":"cols"} of A as stored.`);if(p){if(u!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[H,$]=e==="no-transpose"?[o,a]:[a,o];if(s.rows<H||s.cols<$)throw new Error("A is too small for the given n, k, and trans.")}else if(s.length<(y-1)*u+v)throw new Error("A does not have enough elements for the given dimensions and lda.");if(l<o)throw new Error("ldc must be >= n.");if(c){if(l!==f.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(f.rows<o||f.cols<o)throw new Error("C is too small for the given n.")}else if(f.length<(o-1)*l+o)throw new Error("C does not have enough elements for the given dimensions and ldc.");let _=e;g==="column-major"&&(_=_==="no-transpose"?"transpose":"no-transpose");let A=_==="no-transpose"?"transpose":"no-transpose",k=t;w==="column-major"&&([_,A]=[A==="no-transpose"?"transpose":"no-transpose",_==="no-transpose"?"transpose":"no-transpose"],k=k==="lower"?"upper":"lower");let B=Math.ceil(o/64),L=Math.ceil(o/64),C=B*L>=36,R=await D(r,C?"sgemmtr_large":"sgemmtr_small"),F=p?s._buf:x(r,s,"ssyrk-A",!1),W=c?f._buf:x(r,f,"ssyrk-C",!0),V=p?tr(r,F.size,"ssyrk-B",GPUBufferUsage.COPY_DST):x(r,s,"ssyrk-B",!1),z=I(r,[{value:o,type:"u32"},{value:o,type:"u32"},{value:a,type:"u32"},{value:i,type:"f32"},{value:n,type:"f32"},{value:u,type:"u32"},{value:u,type:"u32"},{value:l,type:"u32"},{value:_==="transpose"?1:0,type:"u32"},{value:A==="transpose"?1:0,type:"u32"},{value:k==="upper"?1:0,type:"u32"}],"ssyrk-params");try{let H=E(r,R.getBindGroupLayout(0),[F,V,W,z]),$=C?{x:U(r,B,"ssyrk","x"),y:U(r,L,"ssyrk","y")}:{x:U(r,Math.ceil(o/32),"ssyrk","x"),y:U(r,Math.ceil(o/32),"ssyrk","y")},{commandEncoder:K,querySet:Y,passDescriptor:ir}=qr(r);p&&K.copyBufferToBuffer(F,0,V,0,F.size),gr(K,R,H,$,ir);let lr=Lr(r,K,Y),hr=c?null:G(r,K,W);M(r,K);let nr=await P(lr);if(c)return nr!==void 0?{gpuTimeMs:nr}:{};let mr=await S(hr,Float32Array);return nr!==void 0?{C:mr,gpuTimeMs:nr}:{C:mr}}finally{p||d(F),d(V),c||d(W),d(z)}}async function ga(r,t,e,o,a,i,s,u,n,f,l,m,p,c="row-major"){let g=s instanceof X,w=n instanceof X,h=m instanceof X;if(q(r),T(r,"ssyr2k",{A:s,B:n,C:m}),t!=="lower"&&t!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(e!=="no-transpose"&&e!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(c!=="row-major"&&c!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof i!="number")throw new Error("alpha must be a number.");if(Number.isNaN(i))throw new Error("alpha must not be NaN.");if(!Number.isFinite(i))throw new Error("alpha must be finite.");if(typeof l!="number")throw new Error("beta must be a number.");if(Number.isNaN(l))throw new Error("beta must not be NaN.");if(!Number.isFinite(l))throw new Error("beta must be finite.");if(!Number.isInteger(o)||!Number.isInteger(a)||!Number.isInteger(u)||!Number.isInteger(f)||!Number.isInteger(p))throw new Error("n, k, lda, ldb, and ldc must be integers.");if(!g&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!w&&!(n instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(!h&&!(m instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if((g||w)&&!h)throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");if(h&&(!g||!w))throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");if(o<0||a<0)throw new Error("n and k must be non-negative.");if(u<=0||f<=0||p<=0)throw new Error("lda, ldb, and ldc must be positive.");if(o===0)return h?{}:{C:m};let b=g?s.layout:c,y=w?n.layout:c,v=h?m.layout:c,_=b==="column-major"?a:o,A=b==="column-major"?o:a,k=e==="no-transpose"?_:A,B=e==="no-transpose"?A:_;if(u<B)throw new Error(`lda must be >= ${b==="column-major"?"rows":"cols"} of A as stored.`);if(g){if(u!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[J,sr]=e==="no-transpose"?[o,a]:[a,o];if(s.rows<J||s.cols<sr)throw new Error("A is too small for the given n, k, and trans.")}else if(s.length<(k-1)*u+B)throw new Error("A does not have enough elements for the given dimensions and lda.");let L=y==="column-major"?a:o,C=y==="column-major"?o:a,R=e==="no-transpose"?L:C,F=e==="no-transpose"?C:L;if(f<F)throw new Error(`ldb must be >= ${y==="column-major"?"rows":"cols"} of B as stored.`);if(w){if(f!==n.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");let[J,sr]=e==="no-transpose"?[o,a]:[a,o];if(n.rows<J||n.cols<sr)throw new Error("B is too small for the given n, k, and trans.")}else if(n.length<(R-1)*f+F)throw new Error("B does not have enough elements for the given dimensions and ldb.");if(p<o)throw new Error("ldc must be >= n.");if(h){if(p!==m.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(m.rows<o||m.cols<o)throw new Error("C is too small for the given n.")}else if(m.length<(o-1)*p+o)throw new Error("C does not have enough elements for the given dimensions and ldc.");let W=e;b==="column-major"&&(W=W==="no-transpose"?"transpose":"no-transpose");let V=e;y==="column-major"&&(V=V==="no-transpose"?"transpose":"no-transpose");let z=v==="column-major"?t==="lower"?"upper":"lower":t,H=J=>J==="no-transpose"?"transpose":"no-transpose";function $(J,sr,pr,br,fr,ur){let yr=J,Cr=H(br);return v!=="column-major"?{transX:yr,X:sr,ldX:pr,transY:Cr,Y:fr,ldY:ur}:{transX:H(Cr),X:fr,ldX:ur,transY:H(yr),Y:sr,ldY:pr}}let K=Math.ceil(o/64),Y=Math.ceil(o/64),ir=K*Y>=36,lr=await D(r,ir?"sgemmtr_large":"sgemmtr_small"),hr=ir?{x:U(r,K,"ssyr2k","x"),y:U(r,Y,"ssyr2k","y")}:{x:U(r,Math.ceil(o/32),"ssyr2k","x"),y:U(r,Math.ceil(o/32),"ssyr2k","y")},nr=g?s._buf:x(r,s,"ssyr2k-A",!1),mr=w?n._buf:x(r,n,"ssyr2k-B",!1),Q=h?m._buf:x(r,m,"ssyr2k-C",!0),er=null,Z=null;try{let J=$(W,nr,u,V,mr,f),sr=$(V,mr,f,W,nr,u),pr=(Ir,Ar)=>I(r,[{value:o,type:"u32"},{value:o,type:"u32"},{value:a,type:"u32"},{value:i,type:"f32"},{value:Ar,type:"f32"},{value:Ir.ldX,type:"u32"},{value:Ir.ldY,type:"u32"},{value:p,type:"u32"},{value:Ir.transX==="transpose"?1:0,type:"u32"},{value:Ir.transY==="transpose"?1:0,type:"u32"},{value:z==="upper"?1:0,type:"u32"}],"ssyr2k-params");er=pr(J,l),Z=pr(sr,1);let br=E(r,lr.getBindGroupLayout(0),[J.X,J.Y,Q,er]),fr=E(r,lr.getBindGroupLayout(0),[sr.X,sr.Y,Q,Z]),{commandEncoder:ur,querySet:yr}=qr(r),Cr=yr?{timestampWrites:{querySet:yr,beginningOfPassWriteIndex:0}}:void 0,Rr=yr?{timestampWrites:{querySet:yr,endOfPassWriteIndex:1}}:void 0;gr(ur,lr,br,hr,Cr),gr(ur,lr,fr,hr,Rr);let Tr=Lr(r,ur,yr),Er=h?null:G(r,ur,Q);M(r,ur);let vr=await P(Tr);if(h)return vr!==void 0?{gpuTimeMs:vr}:{};let xr=await S(Er,Float32Array);return vr!==void 0?{C:xr,gpuTimeMs:vr}:{C:xr}}finally{g||d(nr),w||d(mr),h||d(Q),er&&d(er),Z&&d(Z)}}async function wa(r,t,e,o,a,i,s,u,n,f,l,m,p,c="row-major"){let g=s instanceof X,w=n instanceof X,h=m instanceof X;if(q(r),T(r,"ssymm",{A:s,B:n,C:m}),t!=="left"&&t!=="right")throw new Error("side must be 'left' or 'right'.");if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(c!=="row-major"&&c!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof i!="number")throw new Error("alpha must be a number.");if(Number.isNaN(i))throw new Error("alpha must not be NaN.");if(!Number.isFinite(i))throw new Error("alpha must be finite.");if(typeof l!="number")throw new Error("beta must be a number.");if(Number.isNaN(l))throw new Error("beta must not be NaN.");if(!Number.isFinite(l))throw new Error("beta must be finite.");if(!Number.isInteger(o)||!Number.isInteger(a)||!Number.isInteger(u)||!Number.isInteger(f)||!Number.isInteger(p))throw new Error("m, n, lda, ldb, and ldc must be integers.");if(!g&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!w&&!(n instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(!h&&!(m instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if((g||w)&&!h)throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");if(h&&(!g||!w))throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");if(o<0||a<0)throw new Error("m and n must be non-negative.");if(o===0||a===0)return h?{}:{C:m};let b=g?s.layout:c,y=w?n.layout:c,v=h?m.layout:c,_=t==="left"?o:a;if(u<_)throw new Error("lda must be >= "+(t==="left"?"m":"n")+".");if(g){if(u!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(s.rows<_||s.cols<_)throw new Error("A is too small for the given m/n and side.")}else if(s.length<(_-1)*u+_)throw new Error("A does not have enough elements for the given dimensions and lda.");let A=y==="column-major"?a:o,k=y==="column-major"?o:a;if(f<k)throw new Error(`ldb must be >= ${y==="column-major"?"rows":"cols"} of B as stored.`);if(w){if(f!==n.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");if(n.rows<o||n.cols<a)throw new Error("B is too small for the given m and n.")}else if(n.length<(A-1)*f+k)throw new Error("B does not have enough elements for the given dimensions and ldb.");let B=v==="column-major"?a:o,L=v==="column-major"?o:a;if(p<L)throw new Error(`ldc must be >= ${v==="column-major"?"rows":"cols"} of C as stored.`);if(h){if(p!==m.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(m.rows<o||m.cols<a)throw new Error("C is too small for the given m and n.")}else if(m.length<(B-1)*p+L)throw new Error("C does not have enough elements for the given dimensions and ldc.");let C=b==="column-major"?e==="lower"?"upper":"lower":e,R=y==="column-major"?"transpose":"no-transpose",F="no-transpose",W=o,V=a,z=_,H=t==="left"?F:R,$=t==="left"?R:F,K=ur=>ur==="no-transpose"?"transpose":"no-transpose",Y=t==="right";v==="column-major"&&([H,$]=[K($),K(H)],Y=!Y,[W,V]=[V,W]);let ir=_,lr=Math.ceil(V/64),hr=Math.ceil(W/64),nr=lr*hr>=36,mr=await D(r,nr?"sgemm_large":"sgemm_small"),Q=await D(r,"symmetrize"),er=nr?{x:U(r,lr,"ssymm","x"),y:U(r,hr,"ssymm","y")}:{x:U(r,Math.ceil(V/32),"ssymm","x"),y:U(r,Math.ceil(W/32),"ssymm","y")},Z=g?s._buf:x(r,s,"ssymm-A",!1),J=w?n._buf:x(r,n,"ssymm-B",!1),sr=h?m._buf:x(r,m,"ssymm-C",!0),pr=tr(r,_*ir*4,"ssymm-Adense"),br=null,fr=null;try{br=I(r,[{value:_,type:"u32"},{value:u,type:"u32"},{value:ir,type:"u32"},{value:C==="upper"?1:0,type:"u32"}],"ssymm-sym-params");let ur=E(r,Q.getBindGroupLayout(0),[Z,pr,br]),yr=Y?J:pr,Cr=Y?f:ir,Rr=Y?pr:J;fr=I(r,[{value:W,type:"u32"},{value:V,type:"u32"},{value:z,type:"u32"},{value:i,type:"f32"},{value:l,type:"f32"},{value:Cr,type:"u32"},{value:Y?ir:f,type:"u32"},{value:p,type:"u32"},{value:H==="transpose"?1:0,type:"u32"},{value:$==="transpose"?1:0,type:"u32"}],"ssymm-gemm-params");let Er=E(r,mr.getBindGroupLayout(0),[yr,Sr(r,yr),Rr,Sr(r,Rr),sr,fr]),{commandEncoder:vr,querySet:xr}=qr(r),Ir=xr?{timestampWrites:{querySet:xr,beginningOfPassWriteIndex:0}}:void 0,Ar=xr?{timestampWrites:{querySet:xr,endOfPassWriteIndex:1}}:void 0;gr(vr,Q,ur,{x:Math.ceil(_/8),y:Math.ceil(_/8)},Ir),gr(vr,mr,Er,er,Ar);let Fr=Lr(r,vr,xr),Or=h?null:G(r,vr,sr);M(r,vr);let Ur=await P(Fr);if(h)return Ur!==void 0?{gpuTimeMs:Ur}:{};let me=await S(Or,Float32Array);return Ur!==void 0?{C:me,gpuTimeMs:Ur}:{C:me}}finally{g||d(Z),w||d(J),h||d(sr),d(pr),br&&d(br),fr&&d(fr)}}async function ha(r,t,e,o,a,i,s,u,n,f,l,m,p="row-major"){let c=n instanceof X,g=l instanceof X,w=a==="unit";if(q(r),T(r,"strmm",{A:n,B:l}),t!=="left"&&t!=="right")throw new Error("side must be 'left' or 'right'.");if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(o!=="no-transpose"&&o!=="transpose")throw new Error("transA must be 'no-transpose' or 'transpose'.");if(!w&&a!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(p!=="row-major"&&p!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof u!="number")throw new Error("alpha must be a number.");if(Number.isNaN(u))throw new Error("alpha must not be NaN.");if(!Number.isFinite(u))throw new Error("alpha must be finite.");if(!Number.isInteger(i)||!Number.isInteger(s)||!Number.isInteger(f)||!Number.isInteger(m))throw new Error("m, n, lda, and ldb must be integers.");if(!c&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!g&&!(l instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(c!==g)throw new Error("A and B must both be GpuMatrix or both be Float32Array.");if(i<0||s<0)throw new Error("m and n must be non-negative.");if(i===0||s===0)return g?{}:{B:l};let h=c?n.layout:p,b=g?l.layout:p,y=t==="left"?i:s;if(f<y)throw new Error("lda must be >= "+(t==="left"?"m":"n")+".");if(c){if(f!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(n.rows<y||n.cols<y)throw new Error("A is too small for the given m/n and side.")}else if(n.length<(y-1)*f+y)throw new Error("A does not have enough elements for the given dimensions and lda.");let v=b==="column-major"?s:i,_=b==="column-major"?i:s;if(m<_)throw new Error(`ldb must be >= ${b==="column-major"?"rows":"cols"} of B as stored.`);if(g){if(m!==l.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");if(l.rows<i||l.cols<s)throw new Error("B is too small for the given m and n.")}else if(l.length<(v-1)*m+_)throw new Error("B does not have enough elements for the given dimensions and ldb.");let A=h==="column-major"?e==="lower"?"upper":"lower":e,k=h==="column-major"?o==="no-transpose"?"transpose":"no-transpose":o,B=b==="column-major"?"transpose":"no-transpose",L="no-transpose",C=i,R=s,F=y,W=t==="left"?L:B,V=t==="left"?B:L,z=br=>br==="no-transpose"?"transpose":"no-transpose",H=t==="right";b==="column-major"&&([W,V]=[z(V),z(W)],H=!H,[C,R]=[R,C]);let $=y,K=Math.ceil(R/64),Y=Math.ceil(C/64),ir=K*Y>=36,lr=await D(r,ir?"sgemm_large":"sgemm_small"),hr=await D(r,"triangularize"),nr=ir?{x:U(r,K,"strmm","x"),y:U(r,Y,"strmm","y")}:{x:U(r,Math.ceil(R/32),"strmm","x"),y:U(r,Math.ceil(C/32),"strmm","y")},mr=null,Q=null,er=null,Z=null,J=null,sr=null,pr=!1;try{mr=c?n._buf:x(r,n,"strmm-A",!1),Q=g?l._buf:x(r,l,"strmm-B",!0),er=tr(r,y*$*4,"strmm-Adense"),Z=tr(r,v*m*4,"strmm-out",GPUBufferUsage.COPY_SRC|GPUBufferUsage.COPY_DST),J=I(r,[{value:y,type:"u32"},{value:f,type:"u32"},{value:$,type:"u32"},{value:A==="upper"?1:0,type:"u32"},{value:k==="transpose"?1:0,type:"u32"},{value:w?1:0,type:"u32"}],"strmm-tri-params");let br=E(r,hr.getBindGroupLayout(0),[mr,er,J]),fr=H?Q:er,ur=H?m:$,yr=H?er:Q;sr=I(r,[{value:C,type:"u32"},{value:R,type:"u32"},{value:F,type:"u32"},{value:u,type:"f32"},{value:0,type:"f32"},{value:ur,type:"u32"},{value:H?$:m,type:"u32"},{value:m,type:"u32"},{value:W==="transpose"?1:0,type:"u32"},{value:V==="transpose"?1:0,type:"u32"}],"strmm-gemm-params");let Rr=E(r,lr.getBindGroupLayout(0),[fr,Sr(r,fr),yr,Sr(r,yr),Z,sr]),{commandEncoder:Tr,querySet:Er}=qr(r);Tr.copyBufferToBuffer(Q,0,Z,0,Math.min(Q.size,Z.size));let vr=Er?{timestampWrites:{querySet:Er,beginningOfPassWriteIndex:0}}:void 0,xr=Er?{timestampWrites:{querySet:Er,endOfPassWriteIndex:1}}:void 0;gr(Tr,hr,br,{x:Math.ceil(y/8),y:Math.ceil(y/8)},vr),gr(Tr,lr,Rr,nr,xr);let Ir=Lr(r,Tr,Er),Ar=g?null:G(r,Tr,Z);M(r,Tr);let Fr=await P(Ir);if(g)return d(l._buf),l._buf=Z,pr=!0,Fr!==void 0?{gpuTimeMs:Fr}:{};let Or=await S(Ar,Float32Array);return Fr!==void 0?{B:Or,gpuTimeMs:Fr}:{B:Or}}finally{!c&&mr&&d(mr),!g&&Q&&d(Q),er&&d(er),Z&&!pr&&d(Z),J&&d(J),sr&&d(sr)}}async function ba(r,t,e,o,a,i,s,u,n,f,l,m,p="row-major"){let c=n instanceof X,g=l instanceof X,w=a==="unit";if(q(r),T(r,"strsm",{A:n,B:l}),t!=="left"&&t!=="right")throw new Error("side must be 'left' or 'right'.");if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(o!=="no-transpose"&&o!=="transpose")throw new Error("transA must be 'no-transpose' or 'transpose'.");if(!w&&a!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(p!=="row-major"&&p!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof u!="number")throw new Error("alpha must be a number.");if(Number.isNaN(u))throw new Error("alpha must not be NaN.");if(!Number.isFinite(u))throw new Error("alpha must be finite.");if(!Number.isInteger(i)||!Number.isInteger(s)||!Number.isInteger(f)||!Number.isInteger(m))throw new Error("m, n, lda, and ldb must be integers.");if(!c&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!g&&!(l instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(c!==g)throw new Error("A and B must both be GpuMatrix or both be Float32Array.");if(i<0||s<0)throw new Error("m and n must be non-negative.");if(i===0||s===0)return g?{}:{B:l};let h=c?n.layout:p,b=g?l.layout:p,y=t==="left"?i:s;if(f<y)throw new Error("lda must be >= "+(t==="left"?"m":"n")+".");if(c){if(f!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(n.rows<y||n.cols<y)throw new Error("A is too small for the given m/n and side.")}else if(n.length<(y-1)*f+y)throw new Error("A does not have enough elements for the given dimensions and lda.");let v=b==="column-major"?s:i,_=b==="column-major"?i:s;if(m<_)throw new Error(`ldb must be >= ${b==="column-major"?"rows":"cols"} of B as stored.`);if(g){if(m!==l.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");if(l.rows<i||l.cols<s)throw new Error("B is too small for the given m and n.")}else if(l.length<(v-1)*m+_)throw new Error("B does not have enough elements for the given dimensions and ldb.");let A=h==="column-major"?e==="lower"?"upper":"lower":e,k=h==="column-major"?o==="no-transpose"?"transpose":"no-transpose":o,B=t==="left"?s:i,L=t==="left",C=k==="no-transpose"==(A==="lower"),R=t==="left"?C:!C,F=[];for(let Q=0;Q<y;Q+=64)F.push(Q);R||F.reverse();let W=F.length,V=await D(r,"strsv_invert_block"),z=await D(r,"block_transfer"),H=await D(r,"sscal"),$=null,K=null,Y=null,ir=[],lr=[];function hr(Q,er){let Z=tr(r,Q,er);return lr.push(Z),Z}function nr(Q,er){let Z=I(r,Q,er);return ir.push(Z),Z}let mr=(v-1)*m+_;try{$=c?n._buf:x(r,n,"strsm-A",!1),K=g?l._buf:x(r,l,"strsm-B",!0),Y=tr(r,W*64*64*4,"strsm-Ainv");let Q=null;if(u!==1&&u!==0){let Er=nr([{value:mr,type:"u32"},{value:u,type:"f32"},{value:1,type:"u32"}],"strsm-scale-params");Q=E(r,H.getBindGroupLayout(0),[K,Er])}let er=nr([{value:y,type:"u32"},{value:f,type:"u32"},{value:k==="transpose"?1:0,type:"u32"},{value:A==="upper"?1:0,type:"u32"},{value:w?1:0,type:"u32"}],"strsm-invert-params"),Z=E(r,V.getBindGroupLayout(0),[$,Y,er]),J=hr(64*B*4,"strsm-Bblock"),sr=hr(64*B*4,"strsm-Xblock"),pr=hr(y*64*4,"strsm-Aoff"),br=hr(y*B*4,"strsm-delta"),{commandEncoder:fr,querySet:ur}=qr(r);if(u===0){let Er=Math.ceil(_/64),vr=Math.ceil(v/64),xr=Er*vr>=36,Ir=await D(r,xr?"sgemm_large":"sgemm_small"),Ar=nr([{value:v,type:"u32"},{value:_,type:"u32"},{value:0,type:"u32"},{value:0,type:"f32"},{value:0,type:"f32"},{value:1,type:"u32"},{value:1,type:"u32"},{value:m,type:"u32"},{value:0,type:"u32"},{value:0,type:"u32"}],"strsm-zero-params"),Fr=E(r,Ir.getBindGroupLayout(0),[Y,Sr(r,Y),Y,Sr(r,Y),K,Ar]),Or=xr?{x:U(r,Er,"strsm","x"),y:U(r,vr,"strsm","y")}:{x:U(r,Math.ceil(_/32),"strsm","x"),y:U(r,Math.ceil(v/32),"strsm","y")};gr(fr,Ir,Fr,Or,ur?{timestampWrites:{querySet:ur,beginningOfPassWriteIndex:0,endOfPassWriteIndex:1}}:void 0)}else{Q&&gr(fr,H,Q,cr(r,mr)),gr(fr,V,Z,{x:64,y:W},ur?{timestampWrites:{querySet:ur,beginningOfPassWriteIndex:0}}:void 0);for(let vr=0;vr<F.length;vr++){let xr=F[vr],Ir=Math.min(xr+64,y),Ar=Ir-xr,Fr=xr/64,Or=vr===F.length-1,Ur=nr([{value:xr,type:"u32"},{value:Ar,type:"u32"},{value:0,type:"u32"},{value:B,type:"u32"},{value:m,type:"u32"},{value:b==="column-major"?1:0,type:"u32"},{value:L?1:0,type:"u32"},{value:2,type:"u32"}],"strsm-gather-B-params"),me=E(r,z.getBindGroupLayout(0),[J,K,Ur]);gr(fr,z,me,Zr(r,"strsm",Ar,B));{let Yr=Ar,Xr=B,_e=Ar,ae=Math.ceil(Xr/64),ie=Math.ceil(Yr/64),se=ae*ie>=36,ne=await D(r,se?"sgemm_large":"sgemm_small"),Be=nr([{value:Yr,type:"u32"},{value:Xr,type:"u32"},{value:_e,type:"u32"},{value:1,type:"f32"},{value:0,type:"f32"},{value:64,type:"u32"},{value:B,type:"u32"},{value:B,type:"u32"},{value:t==="right"?1:0,type:"u32"},{value:0,type:"u32"}],"strsm-apply-params"),ce={buffer:Y,offset:Fr*64*64*4,size:4096*4},Ae=E(r,ne.getBindGroupLayout(0),[ce,Sr(r,ce),J,Sr(r,J),sr,Be]),Ea=se?{x:U(r,ae,"strsm","x"),y:U(r,ie,"strsm","y")}:{x:U(r,Math.ceil(Xr/32),"strsm","x"),y:U(r,Math.ceil(Yr/32),"strsm","y")};gr(fr,ne,Ae,Ea)}let de=R?Ir:0,Le=R?y:xr,Re=de<Le,ya=nr([{value:xr,type:"u32"},{value:Ar,type:"u32"},{value:0,type:"u32"},{value:B,type:"u32"},{value:m,type:"u32"},{value:b==="column-major"?1:0,type:"u32"},{value:L?1:0,type:"u32"},{value:0,type:"u32"}],"strsm-scatter-params"),xa=E(r,z.getBindGroupLayout(0),[sr,K,ya]),va=Or&&!Re&&ur?{timestampWrites:{querySet:ur,endOfPassWriteIndex:1}}:void 0;if(gr(fr,z,xa,Zr(r,"strsm",Ar,B),va),!Re)continue;let oe=Le-de,_a=nr([{value:de,type:"u32"},{value:oe,type:"u32"},{value:xr,type:"u32"},{value:Ar,type:"u32"},{value:f,type:"u32"},{value:k==="transpose"?1:0,type:"u32"},{value:L?1:0,type:"u32"},{value:2,type:"u32"}],"strsm-gather-A-params"),Ba=E(r,z.getBindGroupLayout(0),[pr,$,_a]);gr(fr,z,Ba,Zr(r,"strsm",oe,Ar));{let Yr=oe,Xr=B,_e=Ar,ae=Math.ceil(Xr/64),ie=Math.ceil(Yr/64),se=ae*ie>=36,ne=await D(r,se?"sgemm_large":"sgemm_small"),Be=nr([{value:Yr,type:"u32"},{value:Xr,type:"u32"},{value:_e,type:"u32"},{value:1,type:"f32"},{value:0,type:"f32"},{value:Ar,type:"u32"},{value:B,type:"u32"},{value:B,type:"u32"},{value:0,type:"u32"},{value:0,type:"u32"}],"strsm-update-params"),ce=E(r,ne.getBindGroupLayout(0),[pr,Sr(r,pr),sr,Sr(r,sr),br,Be]),Ae=se?{x:U(r,ae,"strsm","x"),y:U(r,ie,"strsm","y")}:{x:U(r,Math.ceil(Xr/32),"strsm","x"),y:U(r,Math.ceil(Yr/32),"strsm","y")};gr(fr,ne,ce,Ae)}let Aa=nr([{value:de,type:"u32"},{value:oe,type:"u32"},{value:0,type:"u32"},{value:B,type:"u32"},{value:m,type:"u32"},{value:b==="column-major"?1:0,type:"u32"},{value:L?1:0,type:"u32"},{value:1,type:"u32"}],"strsm-scatter-sub-params"),Sa=E(r,z.getBindGroupLayout(0),[br,K,Aa]),Ga=Or&&ur?{timestampWrites:{querySet:ur,endOfPassWriteIndex:1}}:void 0;gr(fr,z,Sa,Zr(r,"strsm",oe,B),Ga)}}let yr=Lr(r,fr,ur),Cr=g?null:G(r,fr,K);M(r,fr);let Rr=await P(yr);if(g)return Rr!==void 0?{gpuTimeMs:Rr}:{};let Tr=await S(Cr,Float32Array);return Rr!==void 0?{B:Tr,gpuTimeMs:Rr}:{B:Tr}}finally{!c&&$&&d($),!g&&K&&d(K),Y&&d(Y),d(lr),d(ir)}}return Ia(Ti);})();
|