wgblas 2.0.0 → 2.1.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 +18 -18
- package/dist/wgblas.browser.js +1273 -1239
- package/index.d.mts +38 -6
- package/package.json +2 -1
- package/src/classes/GpuMatrix.mjs +17 -10
- package/src/classes/GpuVector.mjs +28 -10
- package/src/dasum/dasum.d.mts +4 -4
- package/src/dasum/dasum.mjs +19 -17
- package/src/devdocs.mjs +13 -0
- package/src/idamax/idamax.d.mts +20 -2
- package/src/idamax/idamax.mjs +18 -16
- package/src/init.mjs +114 -56
- package/src/isamax/isamax.d.mts +20 -2
- package/src/isamax/isamax.mjs +16 -14
- package/src/random/random.d.mts +1 -0
- package/src/sasum/sasum.d.mts +2 -2
- package/src/sasum/sasum.mjs +15 -13
- package/src/saxpy/saxpy.d.mts +2 -2
- package/src/saxpy/saxpy.mjs +10 -8
- package/src/scopy/scopy.d.mts +2 -2
- package/src/scopy/scopy.mjs +10 -8
- package/src/sdot/sdot.d.mts +2 -2
- package/src/sdot/sdot.mjs +16 -14
- package/src/sgemm/sgemm.mjs +28 -15
- package/src/sgemmtr/sgemmtr.mjs +16 -15
- package/src/sgemv/sgemv.mjs +38 -26
- package/src/sger/sger.mjs +10 -8
- package/src/shaders/index.mjs +164 -14
- package/src/shaders/reduction/scaledSum.wgsl +65 -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 +33 -21
- package/src/srot/srot.d.mts +2 -4
- package/src/srot/srot.mjs +11 -9
- package/src/srotm/srotm.d.mts +2 -4
- package/src/srotm/srotm.mjs +19 -10
- package/src/sscal/sscal.d.mts +3 -3
- package/src/sscal/sscal.mjs +12 -10
- package/src/sswap/sswap.d.mts +2 -2
- package/src/sswap/sswap.mjs +11 -9
- package/src/ssymm/ssymm.mjs +31 -22
- package/src/ssymv/ssymv.mjs +10 -8
- package/src/ssyr/ssyr.mjs +9 -7
- package/src/ssyr2/ssyr2.mjs +10 -8
- package/src/ssyr2k/ssyr2k.mjs +18 -17
- package/src/ssyrk/ssyrk.mjs +18 -17
- package/src/strmm/strmm.mjs +47 -32
- package/src/strmv/strmv.mjs +10 -8
- package/src/strsm/strsm.mjs +54 -36
- package/src/strsv/strsv.mjs +16 -12
- package/src/util/benchmark.mjs +4 -6
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +113 -19
- package/src/util/compute.mjs +6 -9
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +34 -0
- package/src/util/pipeline.mjs +5 -6
- package/src/util/workgroup.mjs +55 -7
- package/src/shaders/browser-shaders.mjs +0 -81
package/dist/wgblas.browser.js
CHANGED
|
@@ -1,163 +1,4 @@
|
|
|
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();
|
|
29
|
-
|
|
30
|
-
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
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
|
-
}
|
|
41
|
-
|
|
42
|
-
if (i == 0u) { result[0] = tile_idx[0]; }
|
|
43
|
-
}
|
|
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
|
-
|
|
52
|
-
@group(0) @binding(0) var<storage, read> partialsValHi: array<f32>;
|
|
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>;
|
|
56
|
-
|
|
57
|
-
const WGS: u32 = 64;
|
|
58
|
-
|
|
59
|
-
var<workgroup> tile_val: array<DD, 64>;
|
|
60
|
-
var<workgroup> tile_idx: array<u32, 64>;
|
|
61
|
-
|
|
62
|
-
@compute @workgroup_size(64)
|
|
63
|
-
fn reduce_f64(
|
|
64
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
65
|
-
) {
|
|
66
|
-
let i = lid.x;
|
|
67
|
-
let a_val = DD(partialsValHi[i], partialsValLo[i]);
|
|
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];
|
|
76
|
-
}
|
|
77
|
-
workgroupBarrier();
|
|
78
|
-
|
|
79
|
-
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
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
|
-
}
|
|
91
|
-
|
|
92
|
-
if (i == 0u) { result[0] = tile_idx[0]; }
|
|
93
|
-
}
|
|
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
|
-
|
|
98
|
-
@group(0) @binding(0) var<storage, read> partials: array<f32>;
|
|
99
|
-
@group(0) @binding(1) var<storage, read_write> result: array<f32>;
|
|
100
|
-
|
|
101
|
-
const WGS: u32 = 64;
|
|
102
|
-
|
|
103
|
-
var<workgroup> tile: array<f32, 64>;
|
|
104
|
-
|
|
105
|
-
@compute @workgroup_size(64)
|
|
106
|
-
fn reduce(
|
|
107
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
108
|
-
) {
|
|
109
|
-
let i = lid.x;
|
|
110
|
-
tile[i] = partials[i] + partials[i + WGS];
|
|
111
|
-
workgroupBarrier();
|
|
112
|
-
|
|
113
|
-
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
114
|
-
if (i < s) { tile[i] += tile[i + s]; }
|
|
115
|
-
workgroupBarrier();
|
|
116
|
-
}
|
|
117
|
-
|
|
118
|
-
if (i == 0u) { result[0] = tile[0]; }
|
|
119
|
-
}
|
|
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
|
-
|
|
127
|
-
@group(0) @binding(0) var<storage, read> partialsHi: array<f32>;
|
|
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>;
|
|
131
|
-
|
|
132
|
-
const WGS: u32 = 64;
|
|
133
|
-
|
|
134
|
-
var<workgroup> tile: array<DD, 64>;
|
|
135
|
-
|
|
136
|
-
@compute @workgroup_size(64)
|
|
137
|
-
fn reduce_f64(
|
|
138
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
139
|
-
) {
|
|
140
|
-
let i = lid.x;
|
|
141
|
-
let a = DD(partialsHi[i], partialsLo[i]);
|
|
142
|
-
let b = DD(partialsHi[i + WGS], partialsLo[i + WGS]);
|
|
143
|
-
tile[i] = ddAddProtected(a, b, i);
|
|
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;
|
|
158
|
-
}
|
|
159
|
-
}
|
|
160
|
-
`});var Be,_e=O(()=>{Be=`// sscal: x = alpha * x
|
|
1
|
+
var wgblas=(()=>{var Lo=Object.create;var se=Object.defineProperty;var Wo=Object.getOwnPropertyDescriptor;var Fo=Object.getOwnPropertyNames;var qo=Object.getPrototypeOf,Uo=Object.prototype.hasOwnProperty;var ne=(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 V=(r,t,e)=>()=>{if(e)throw e[0];try{return r&&(t=r(r=0)),t}catch(a){throw e=[a],a}};var Ee=(r,t)=>{for(var e in t)se(r,e,{get:t[e],enumerable:!0})},Ae=(r,t,e,a)=>{if(t&&typeof t=="object"||typeof t=="function")for(let o of Fo(t))!Uo.call(r,o)&&o!==e&&se(r,o,{get:()=>t[o],enumerable:!(a=Wo(t,o))||a.enumerable});return r};var ie=(r,t,e)=>(e=r!=null?Lo(qo(r)):{},Ae(t||!r||!r.__esModule?se(e,"default",{value:r,enumerable:!0}):e,r)),Oo=r=>Ae(se({},"__esModule",{value:!0}),r);var ge,Ce=V(()=>{ge=`// sscal: x = alpha * x
|
|
161
2
|
|
|
162
3
|
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
163
4
|
|
|
@@ -180,7 +21,7 @@ fn main(
|
|
|
180
21
|
x[id * params.x_inc] = params.alpha * x[id * params.x_inc];
|
|
181
22
|
}
|
|
182
23
|
}
|
|
183
|
-
`});var
|
|
24
|
+
`});var We,Le=V(()=>{We=`// sswap: x <-> y
|
|
184
25
|
|
|
185
26
|
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
186
27
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
@@ -206,7 +47,7 @@ fn main(
|
|
|
206
47
|
y[id * params.y_inc] = temp;
|
|
207
48
|
}
|
|
208
49
|
}
|
|
209
|
-
`});var
|
|
50
|
+
`});var qe,Fe=V(()=>{qe=`// saxpy: y = alpha * x + y
|
|
210
51
|
|
|
211
52
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
212
53
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
@@ -231,7 +72,7 @@ fn main(
|
|
|
231
72
|
y[id * params.y_inc] = params.alpha * x[id * params.x_inc] + y[id * params.y_inc];
|
|
232
73
|
}
|
|
233
74
|
}
|
|
234
|
-
`});var
|
|
75
|
+
`});var Oe,Ue=V(()=>{Oe=`// scopy: y = x
|
|
235
76
|
|
|
236
77
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
237
78
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
@@ -255,7 +96,7 @@ fn main(
|
|
|
255
96
|
y[id * params.y_inc] = x[id * params.x_inc];
|
|
256
97
|
}
|
|
257
98
|
}
|
|
258
|
-
`});var
|
|
99
|
+
`});var Ve,Ke=V(()=>{Ve=`// sdot: result = sum(x[i] * y[i])
|
|
259
100
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sum.wgsl.
|
|
260
101
|
|
|
261
102
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -308,7 +149,33 @@ fn main(
|
|
|
308
149
|
|
|
309
150
|
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
310
151
|
}
|
|
311
|
-
`});var
|
|
152
|
+
`});var be,ze=V(()=>{be=`// sum reduction: collapses 2*WGS partials into one scalar.
|
|
153
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
154
|
+
// partials must have exactly 2*WGS entries.
|
|
155
|
+
|
|
156
|
+
@group(0) @binding(0) var<storage, read> partials: array<f32>;
|
|
157
|
+
@group(0) @binding(1) var<storage, read_write> result: array<f32>;
|
|
158
|
+
|
|
159
|
+
const WGS: u32 = 64;
|
|
160
|
+
|
|
161
|
+
var<workgroup> tile: array<f32, 64>;
|
|
162
|
+
|
|
163
|
+
@compute @workgroup_size(64)
|
|
164
|
+
fn reduce(
|
|
165
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
166
|
+
) {
|
|
167
|
+
let i = lid.x;
|
|
168
|
+
tile[i] = partials[i] + partials[i + WGS];
|
|
169
|
+
workgroupBarrier();
|
|
170
|
+
|
|
171
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
172
|
+
if (i < s) { tile[i] += tile[i + s]; }
|
|
173
|
+
workgroupBarrier();
|
|
174
|
+
}
|
|
175
|
+
|
|
176
|
+
if (i == 0u) { result[0] = tile[0]; }
|
|
177
|
+
}
|
|
178
|
+
`});var Ye,He=V(()=>{Ye=`// sasum: result = sum(|x[i]|)
|
|
312
179
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/abssum.wgsl.
|
|
313
180
|
|
|
314
181
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -359,12 +226,25 @@ fn main(
|
|
|
359
226
|
|
|
360
227
|
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
361
228
|
}
|
|
362
|
-
`});var
|
|
363
|
-
//
|
|
364
|
-
|
|
365
|
-
|
|
366
|
-
|
|
367
|
-
|
|
229
|
+
`});var $e,Xe=V(()=>{$e=`// snrm2: result = sqrt(sum(x[i] * x[i])), computed via scaled accumulation
|
|
230
|
+
// (Blue's algorithm / reference BLAS's SLASSQ) rather than naive squaring \u2014
|
|
231
|
+
// naive \`sum += x_i * x_i\` overflows to inf for |x_i| \u2273 1.8e19 (f32's
|
|
232
|
+
// squaring range is only sqrt(f32_max)) and loses precision on tiny
|
|
233
|
+
// magnitudes squaring into the denormal range. Running state is (scale,
|
|
234
|
+
// ssq) with true-sum-of-squares == scale\xB2 \xB7 ssq: scale tracks the largest
|
|
235
|
+
// |x_i| seen so far, and every other contribution is expressed *relative
|
|
236
|
+
// to* scale (never squared in absolute terms), so ssq stays near 1
|
|
237
|
+
// regardless of x's magnitude range. Merging two independent partials
|
|
238
|
+
// (ssqMerge) is associative, so this composes with the same 4-way-ILP +
|
|
239
|
+
// tree-reduction shape every other Level 1 reduction here uses \u2014 see
|
|
240
|
+
// reduction/scaledSum.wgsl for the pass-2 counterpart, which finishes with
|
|
241
|
+
// scale\xB7sqrt(ssq).
|
|
242
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/scaledSum.wgsl.
|
|
243
|
+
|
|
244
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
245
|
+
@group(0) @binding(1) var<storage, read_write> partialsScale: array<f32>;
|
|
246
|
+
@group(0) @binding(2) var<storage, read_write> partialsSsq: array<f32>;
|
|
247
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
368
248
|
|
|
369
249
|
struct Params {
|
|
370
250
|
n: u32,
|
|
@@ -373,7 +253,36 @@ struct Params {
|
|
|
373
253
|
|
|
374
254
|
const WGS: u32 = 64;
|
|
375
255
|
|
|
376
|
-
|
|
256
|
+
struct ScaleSsq {
|
|
257
|
+
scale: f32,
|
|
258
|
+
ssq: f32,
|
|
259
|
+
}
|
|
260
|
+
|
|
261
|
+
// Folds one more |value| into a running (scale, ssq) pair.
|
|
262
|
+
fn ssqAccum(acc: ScaleSsq, absxi: f32) -> ScaleSsq {
|
|
263
|
+
if (absxi == 0.0) { return acc; }
|
|
264
|
+
if (absxi > acc.scale) {
|
|
265
|
+
let r = acc.scale / absxi; // 0/absxi == 0 on the first nonzero value \u2014 safe
|
|
266
|
+
return ScaleSsq(absxi, 1.0 + acc.ssq * r * r);
|
|
267
|
+
}
|
|
268
|
+
let r = absxi / acc.scale; // reached only once acc.scale > 0 (absxi <= acc.scale and absxi > 0)
|
|
269
|
+
return ScaleSsq(acc.scale, acc.ssq + r * r);
|
|
270
|
+
}
|
|
271
|
+
|
|
272
|
+
// Associative merge of two independent (scale, ssq) partials \u2014 lets this
|
|
273
|
+
// compose with a tree reduction exactly like a plain sum would.
|
|
274
|
+
fn ssqMerge(a: ScaleSsq, b: ScaleSsq) -> ScaleSsq {
|
|
275
|
+
if (a.scale == 0.0 && b.scale == 0.0) { return ScaleSsq(0.0, 1.0); }
|
|
276
|
+
if (a.scale >= b.scale) {
|
|
277
|
+
let r = b.scale / a.scale;
|
|
278
|
+
return ScaleSsq(a.scale, a.ssq + b.ssq * r * r);
|
|
279
|
+
}
|
|
280
|
+
let r = a.scale / b.scale;
|
|
281
|
+
return ScaleSsq(b.scale, b.ssq + a.ssq * r * r);
|
|
282
|
+
}
|
|
283
|
+
|
|
284
|
+
var<workgroup> tileScale: array<f32, 64>;
|
|
285
|
+
var<workgroup> tileSsq: array<f32, 64>;
|
|
377
286
|
|
|
378
287
|
@compute @workgroup_size(64)
|
|
379
288
|
fn main(
|
|
@@ -382,136 +291,129 @@ fn main(
|
|
|
382
291
|
@builtin(workgroup_id) wgid: vec3u,
|
|
383
292
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
384
293
|
) {
|
|
385
|
-
var acc0
|
|
386
|
-
var acc1
|
|
387
|
-
var acc2
|
|
388
|
-
var acc3
|
|
294
|
+
var acc0 = ScaleSsq(0.0, 1.0);
|
|
295
|
+
var acc1 = ScaleSsq(0.0, 1.0);
|
|
296
|
+
var acc2 = ScaleSsq(0.0, 1.0);
|
|
297
|
+
var acc3 = ScaleSsq(0.0, 1.0);
|
|
389
298
|
|
|
390
299
|
let stride = num_wg.x * WGS;
|
|
391
300
|
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
392
301
|
|
|
393
302
|
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;
|
|
303
|
+
acc0 = ssqAccum(acc0, abs(x[ id * params.x_inc]));
|
|
304
|
+
acc1 = ssqAccum(acc1, abs(x[(id + stride) * params.x_inc]));
|
|
305
|
+
acc2 = ssqAccum(acc2, abs(x[(id + 2u * stride) * params.x_inc]));
|
|
306
|
+
acc3 = ssqAccum(acc3, abs(x[(id + 3u * stride) * params.x_inc]));
|
|
402
307
|
}
|
|
403
308
|
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
404
|
-
|
|
405
|
-
acc0 += v * v;
|
|
309
|
+
acc0 = ssqAccum(acc0, abs(x[id * params.x_inc]));
|
|
406
310
|
}
|
|
407
311
|
|
|
408
|
-
|
|
312
|
+
let combined = ssqMerge(ssqMerge(acc0, acc1), ssqMerge(acc2, acc3));
|
|
313
|
+
tileScale[lid.x] = combined.scale;
|
|
314
|
+
tileSsq[lid.x] = combined.ssq;
|
|
409
315
|
workgroupBarrier();
|
|
410
316
|
|
|
411
317
|
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
412
|
-
if (lid.x < s) {
|
|
318
|
+
if (lid.x < s) {
|
|
319
|
+
let merged = ssqMerge(
|
|
320
|
+
ScaleSsq(tileScale[lid.x], tileSsq[lid.x]),
|
|
321
|
+
ScaleSsq(tileScale[lid.x + s], tileSsq[lid.x + s]),
|
|
322
|
+
);
|
|
323
|
+
tileScale[lid.x] = merged.scale;
|
|
324
|
+
tileSsq[lid.x] = merged.ssq;
|
|
325
|
+
}
|
|
413
326
|
workgroupBarrier();
|
|
414
327
|
}
|
|
415
328
|
|
|
416
|
-
if (lid.x == 0u) {
|
|
329
|
+
if (lid.x == 0u) {
|
|
330
|
+
partialsScale[wgid.x] = tileScale[0];
|
|
331
|
+
partialsSsq[wgid.x] = tileSsq[0];
|
|
332
|
+
}
|
|
417
333
|
}
|
|
418
|
-
`});var
|
|
334
|
+
`});var Qe,Ze=V(()=>{Qe=`// scaledSum reduction: collapses 2*WGS (scale, ssq) partials from
|
|
335
|
+
// snrm2.wgsl into the final norm \u2014 sqrt(scale\xB2 \xB7 ssq) == scale \xB7 sqrt(ssq).
|
|
336
|
+
// Mirrors reduction/sum.wgsl's shape exactly, merging via ssqMerge (see
|
|
337
|
+
// snrm2.wgsl for the derivation) instead of plain \`+\`, and taking the final
|
|
338
|
+
// sqrt here rather than on the CPU \u2014 unlike sasum/sdot's plain sum, "sum of
|
|
339
|
+
// squares" isn't a meaningful standalone value to hand back, only
|
|
340
|
+
// scale\xB7sqrt(ssq) is.
|
|
341
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
342
|
+
// partialsScale/partialsSsq must have exactly 2*WGS entries each.
|
|
419
343
|
|
|
420
|
-
@group(0) @binding(0) var<storage,
|
|
421
|
-
@group(0) @binding(1) var<storage,
|
|
344
|
+
@group(0) @binding(0) var<storage, read> partialsScale: array<f32>;
|
|
345
|
+
@group(0) @binding(1) var<storage, read> partialsSsq: array<f32>;
|
|
346
|
+
@group(0) @binding(2) var<storage, read_write> result: array<f32>;
|
|
422
347
|
|
|
423
|
-
|
|
424
|
-
|
|
425
|
-
|
|
426
|
-
|
|
427
|
-
|
|
428
|
-
|
|
348
|
+
const WGS: u32 = 64;
|
|
349
|
+
|
|
350
|
+
// True sum-of-squares represented so far == scale\xB2 \xB7 ssq \u2014 see snrm2.wgsl.
|
|
351
|
+
struct ScaleSsq {
|
|
352
|
+
scale: f32,
|
|
353
|
+
ssq: f32,
|
|
429
354
|
}
|
|
430
355
|
|
|
431
|
-
|
|
356
|
+
// Associative merge of two independent (scale, ssq) partials.
|
|
357
|
+
fn ssqMerge(a: ScaleSsq, b: ScaleSsq) -> ScaleSsq {
|
|
358
|
+
if (a.scale == 0.0 && b.scale == 0.0) { return ScaleSsq(0.0, 1.0); }
|
|
359
|
+
if (a.scale >= b.scale) {
|
|
360
|
+
let r = b.scale / a.scale;
|
|
361
|
+
return ScaleSsq(a.scale, a.ssq + b.ssq * r * r);
|
|
362
|
+
}
|
|
363
|
+
let r = a.scale / b.scale;
|
|
364
|
+
return ScaleSsq(b.scale, b.ssq + a.ssq * r * r);
|
|
365
|
+
}
|
|
432
366
|
|
|
433
|
-
|
|
367
|
+
var<workgroup> tileScale: array<f32, 64>;
|
|
368
|
+
var<workgroup> tileSsq: array<f32, 64>;
|
|
434
369
|
|
|
435
370
|
@compute @workgroup_size(64)
|
|
436
|
-
fn
|
|
437
|
-
@builtin(
|
|
438
|
-
@builtin(num_workgroups) num_wg: vec3u,
|
|
371
|
+
fn reduce_scaled(
|
|
372
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
439
373
|
) {
|
|
440
|
-
|
|
441
|
-
|
|
442
|
-
|
|
443
|
-
|
|
444
|
-
|
|
374
|
+
let i = lid.x;
|
|
375
|
+
let merged0 = ssqMerge(
|
|
376
|
+
ScaleSsq(partialsScale[i], partialsSsq[i]),
|
|
377
|
+
ScaleSsq(partialsScale[i + WGS], partialsSsq[i + WGS]),
|
|
378
|
+
);
|
|
379
|
+
tileScale[i] = merged0.scale;
|
|
380
|
+
tileSsq[i] = merged0.ssq;
|
|
381
|
+
workgroupBarrier();
|
|
382
|
+
|
|
383
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
384
|
+
if (i < s) {
|
|
385
|
+
let merged = ssqMerge(
|
|
386
|
+
ScaleSsq(tileScale[i], tileSsq[i]),
|
|
387
|
+
ScaleSsq(tileScale[i + s], tileSsq[i + s]),
|
|
388
|
+
);
|
|
389
|
+
tileScale[i] = merged.scale;
|
|
390
|
+
tileSsq[i] = merged.ssq;
|
|
391
|
+
}
|
|
392
|
+
workgroupBarrier();
|
|
393
|
+
}
|
|
394
|
+
|
|
395
|
+
if (i == 0u) {
|
|
396
|
+
result[0] = tileScale[0] * sqrt(tileSsq[0]);
|
|
445
397
|
}
|
|
446
398
|
}
|
|
447
|
-
`});var
|
|
448
|
-
//
|
|
449
|
-
// param = [ flag, h11, h21, h12, h22 ]
|
|
450
|
-
// flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
|
|
399
|
+
`});var rt,Je=V(()=>{rt=`// isamax: returns index of element with largest absolute value
|
|
400
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/argmax.wgsl.
|
|
451
401
|
|
|
452
|
-
@group(0) @binding(0) var<storage,
|
|
453
|
-
@group(0) @binding(1) var<storage, read_write>
|
|
454
|
-
@group(0) @binding(2) var<storage,
|
|
402
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
403
|
+
@group(0) @binding(1) var<storage, read_write> partials_val: array<f32>;
|
|
404
|
+
@group(0) @binding(2) var<storage, read_write> partials_idx: array<u32>;
|
|
405
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
455
406
|
|
|
456
407
|
struct Params {
|
|
457
408
|
n: u32,
|
|
458
409
|
x_inc: u32,
|
|
459
|
-
y_inc: u32,
|
|
460
410
|
}
|
|
461
411
|
|
|
462
|
-
@group(0) @binding(3) var<uniform> params: Params;
|
|
463
|
-
|
|
464
412
|
const WGS: u32 = 64;
|
|
465
413
|
|
|
466
|
-
|
|
467
|
-
|
|
468
|
-
|
|
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;
|
|
475
|
-
|
|
476
|
-
if (flag == -1.0) {
|
|
477
|
-
// full 2x2 matrix
|
|
478
|
-
h11 = param[1]; h21 = param[2];
|
|
479
|
-
h12 = param[3]; h22 = param[4];
|
|
480
|
-
} else if (flag == 0.0) {
|
|
481
|
-
// diagonal fixed at 1
|
|
482
|
-
h11 = 1.0; h21 = param[2];
|
|
483
|
-
h12 = param[3]; h22 = 1.0;
|
|
484
|
-
} else if (flag == 1.0) {
|
|
485
|
-
// flag == 1.0: off-diagonal fixed at +1 / -1
|
|
486
|
-
h11 = param[1]; h21 = -1.0;
|
|
487
|
-
h12 = 1.0; h22 = param[4];
|
|
488
|
-
}
|
|
489
|
-
|
|
490
|
-
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
491
|
-
let xi = x[id * params.x_inc];
|
|
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;
|
|
495
|
-
}
|
|
496
|
-
}
|
|
497
|
-
`});var He,Fe=O(()=>{He=`// isamax: returns index of element with largest absolute value
|
|
498
|
-
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/argmax.wgsl.
|
|
499
|
-
|
|
500
|
-
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
501
|
-
@group(0) @binding(1) var<storage, read_write> partials_val: array<f32>;
|
|
502
|
-
@group(0) @binding(2) var<storage, read_write> partials_idx: array<u32>;
|
|
503
|
-
@group(0) @binding(3) var<uniform> params: Params;
|
|
504
|
-
|
|
505
|
-
struct Params {
|
|
506
|
-
n: u32,
|
|
507
|
-
x_inc: u32,
|
|
508
|
-
}
|
|
509
|
-
|
|
510
|
-
const WGS: u32 = 64;
|
|
511
|
-
|
|
512
|
-
var<workgroup> tile_val: array<f32, 64>;
|
|
513
|
-
var<workgroup> tile_idx: array<u32, 64>;
|
|
514
|
-
|
|
414
|
+
var<workgroup> tile_val: array<f32, 64>;
|
|
415
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
416
|
+
|
|
515
417
|
@compute @workgroup_size(64)
|
|
516
418
|
fn main(
|
|
517
419
|
@builtin(global_invocation_id) gid: vec3u,
|
|
@@ -576,1076 +478,842 @@ fn main(
|
|
|
576
478
|
partials_idx[wgid.x] = tile_idx[0];
|
|
577
479
|
}
|
|
578
480
|
}
|
|
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>;
|
|
481
|
+
`});var tt,et=V(()=>{tt=`// amax reduction: collapses 2*WGS (value, index) pairs into one index.
|
|
482
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
483
|
+
// partials_val and partials_idx must have exactly 2*WGS entries.
|
|
591
484
|
|
|
592
|
-
|
|
593
|
-
|
|
594
|
-
|
|
595
|
-
alpha: f32,
|
|
596
|
-
beta: f32,
|
|
597
|
-
incx: u32,
|
|
598
|
-
incy: u32,
|
|
599
|
-
lda: u32,
|
|
600
|
-
}
|
|
485
|
+
@group(0) @binding(0) var<storage, read> partials_val: array<f32>;
|
|
486
|
+
@group(0) @binding(1) var<storage, read> partials_idx: array<u32>;
|
|
487
|
+
@group(0) @binding(2) var<storage, read_write> result: array<u32>;
|
|
601
488
|
|
|
602
|
-
|
|
489
|
+
const WGS: u32 = 64;
|
|
603
490
|
|
|
604
|
-
|
|
605
|
-
var<workgroup>
|
|
491
|
+
var<workgroup> tile_val: array<f32, 64>;
|
|
492
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
606
493
|
|
|
607
494
|
@compute @workgroup_size(64)
|
|
608
|
-
fn
|
|
609
|
-
@builtin(
|
|
610
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
611
|
-
@builtin(num_workgroups) nwg: vec3u,
|
|
495
|
+
fn reduce(
|
|
496
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
612
497
|
) {
|
|
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
|
-
}
|
|
498
|
+
let i = lid.x;
|
|
499
|
+
let a_val = partials_val[i];
|
|
500
|
+
let b_val = partials_val[i + WGS];
|
|
501
|
+
if (b_val > a_val || (b_val == a_val && partials_idx[i + WGS] < partials_idx[i])) {
|
|
502
|
+
tile_val[i] = b_val;
|
|
503
|
+
tile_idx[i] = partials_idx[i + WGS];
|
|
504
|
+
} else {
|
|
505
|
+
tile_val[i] = a_val;
|
|
506
|
+
tile_idx[i] = partials_idx[i];
|
|
507
|
+
}
|
|
508
|
+
workgroupBarrier();
|
|
635
509
|
|
|
636
|
-
|
|
637
|
-
|
|
638
|
-
|
|
639
|
-
|
|
640
|
-
if
|
|
641
|
-
|
|
510
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
511
|
+
if (i < s) {
|
|
512
|
+
let c_val = tile_val[i];
|
|
513
|
+
let d_val = tile_val[i + s];
|
|
514
|
+
if (d_val > c_val || (d_val == c_val && tile_idx[i + s] < tile_idx[i])) {
|
|
515
|
+
tile_val[i] = d_val;
|
|
516
|
+
tile_idx[i] = tile_idx[i + s];
|
|
642
517
|
}
|
|
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
518
|
}
|
|
650
|
-
// All 64 threads must agree before the next row reuses scratch[].
|
|
651
519
|
workgroupBarrier();
|
|
652
520
|
}
|
|
653
|
-
}
|
|
654
|
-
`});var Ke,Ve=O(()=>{Ke=`// sgemv_t: y = alpha * A^T * x + beta * y (A is m\xD7n row-major, transposed)
|
|
655
|
-
// each thread owns one column of A \u2192 one element of y (length n)
|
|
656
|
-
// tiles over x (length m) using shared memory; four independent accumulators
|
|
657
|
-
// let the GPU pipeline A reads across j within each tile (ILP=4)
|
|
658
|
-
|
|
659
|
-
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
660
|
-
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
661
|
-
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
662
521
|
|
|
663
|
-
|
|
664
|
-
m: u32,
|
|
665
|
-
n: u32,
|
|
666
|
-
alpha: f32,
|
|
667
|
-
beta: f32,
|
|
668
|
-
incx: u32,
|
|
669
|
-
incy: u32,
|
|
670
|
-
lda: u32,
|
|
522
|
+
if (i == 0u) { result[0] = tile_idx[0]; }
|
|
671
523
|
}
|
|
524
|
+
`});var he,ot=V(()=>{he=`// Double-double arithmetic via Dekker's algorithm \u2014 an alternative to
|
|
525
|
+
// f64add.wgsl's bit-exact IEEE-754 emulation. Doesn't touch that path.
|
|
526
|
+
//
|
|
527
|
+
// A double-double number is a pair (hi, lo) of f32 with hi+lo approximating
|
|
528
|
+
// a higher-precision value, hi holding the leading bits and lo the rounding
|
|
529
|
+
// error hi lost. ~48 bits of mantissa vs f32's 24, less than real f64's 52.
|
|
530
|
+
//
|
|
531
|
+
// No bindings, no entry point \u2014 a helper library, concatenated with a
|
|
532
|
+
// consumer's own bindings/entry point by getPipeline (WGSL has no #include).
|
|
533
|
+
// The DD struct lives here \u2014 abs.wgsl/add.wgsl/greater.wgsl/equal.wgsl all
|
|
534
|
+
// use it but don't redefine it (WGSL errors on duplicate struct definitions
|
|
535
|
+
// once concatenated), so any consumer using those must concatenate this
|
|
536
|
+
// file too, first.
|
|
672
537
|
|
|
673
|
-
|
|
674
|
-
|
|
675
|
-
|
|
676
|
-
|
|
677
|
-
|
|
678
|
-
@compute @workgroup_size(64)
|
|
679
|
-
fn main(
|
|
680
|
-
@builtin(global_invocation_id) gid: vec3u,
|
|
681
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
682
|
-
) {
|
|
683
|
-
// each thread owns column col of A \u2192 output y[col]
|
|
684
|
-
let col = gid.x;
|
|
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;
|
|
691
|
-
|
|
692
|
-
for (var base = 0u; base < m_floor; base += WGS) {
|
|
693
|
-
// cooperative load: all 64 threads fill x_tile with x[base..base+WGS]
|
|
694
|
-
x_tile[lid.x] = x[(base + lid.x) * params.incx];
|
|
695
|
-
workgroupBarrier();
|
|
696
|
-
|
|
697
|
-
// 4-unrolled inner loop: 4 independent A reads let the GPU pipeline
|
|
698
|
-
// global-memory requests within each tile. WGS=64 divides by 4 exactly.
|
|
699
|
-
if (col < params.n) {
|
|
700
|
-
for (var j = 0u; j < WGS; j += 4u) {
|
|
701
|
-
acc0 += A[(base + j ) * params.lda + col] * x_tile[j ];
|
|
702
|
-
acc1 += A[(base + j + 1) * params.lda + col] * x_tile[j + 1];
|
|
703
|
-
acc2 += A[(base + j + 2) * params.lda + col] * x_tile[j + 2];
|
|
704
|
-
acc3 += A[(base + j + 3) * params.lda + col] * x_tile[j + 3];
|
|
705
|
-
}
|
|
706
|
-
}
|
|
707
|
-
workgroupBarrier();
|
|
708
|
-
}
|
|
538
|
+
struct DD {
|
|
539
|
+
hi: f32,
|
|
540
|
+
lo: f32,
|
|
541
|
+
}
|
|
542
|
+
`});var ye,at=V(()=>{ye=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
709
543
|
|
|
710
|
-
|
|
711
|
-
|
|
712
|
-
|
|
713
|
-
|
|
714
|
-
|
|
715
|
-
let yi = col * params.incy;
|
|
716
|
-
y[yi] = params.alpha * (acc0 + acc1 + acc2 + acc3) + params.beta * y[yi];
|
|
544
|
+
// |a| for a double-double pair. Negation is exact (no rounding), so this is
|
|
545
|
+
// just a sign flip on both components \u2014 hi alone determines the pair's sign.
|
|
546
|
+
fn ddAbs(a: DD) -> DD {
|
|
547
|
+
if (a.hi < 0.0) {
|
|
548
|
+
return DD(-a.hi, -a.lo);
|
|
717
549
|
}
|
|
550
|
+
return a;
|
|
718
551
|
}
|
|
719
|
-
`});var
|
|
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.
|
|
552
|
+
`});var nt,st=V(()=>{nt=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
725
553
|
|
|
726
|
-
|
|
727
|
-
|
|
728
|
-
|
|
554
|
+
// \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
|
|
555
|
+
//
|
|
556
|
+
// twoSum/fastTwoSum's error term \`e\` should be nonzero (that's the point \u2014
|
|
557
|
+
// \`s\` is rounded). A front-end-optimizer bug zeros it anyway, on both NVIDIA
|
|
558
|
+
// and Mesa (ANV + llvmpipe), via two different mechanisms needing two fixes:
|
|
559
|
+
// bitcast-based subtraction (\`fsub\`/\`negf\`, fixes NVIDIA) and materializing
|
|
560
|
+
// the sum through workgroup memory + workgroupBarrier() (fixes Mesa). Only
|
|
561
|
+
// both together (ddAddProtected) is verified correct everywhere \u2014 the plain
|
|
562
|
+
// twoSum/fastTwoSum/ddAdd below are reference-only, not safe to use.
|
|
563
|
+
fn negf(x: f32) -> f32 {
|
|
564
|
+
return bitcast<f32>(bitcast<u32>(x) ^ 0x80000000u);
|
|
565
|
+
}
|
|
566
|
+
fn fsub(a: f32, b: f32) -> f32 {
|
|
567
|
+
return a + negf(b);
|
|
568
|
+
}
|
|
729
569
|
|
|
730
|
-
|
|
731
|
-
|
|
732
|
-
|
|
733
|
-
|
|
734
|
-
|
|
735
|
-
|
|
736
|
-
|
|
737
|
-
uplo: u32, // 0 = lower, 1 = upper
|
|
570
|
+
// Knuth/M\xF8ller's TwoSum: s = fl(a+b), e = exact rounding error, a+b == s+e.
|
|
571
|
+
// Works for any a, b. UNPROTECTED \u2014 see header above.
|
|
572
|
+
fn twoSum(a: f32, b: f32) -> DD {
|
|
573
|
+
let s = a + b;
|
|
574
|
+
let v = s - a;
|
|
575
|
+
let e = (a - (s - v)) + (b - v);
|
|
576
|
+
return DD(s, e);
|
|
738
577
|
}
|
|
739
578
|
|
|
740
|
-
|
|
741
|
-
|
|
742
|
-
|
|
743
|
-
|
|
744
|
-
|
|
745
|
-
|
|
746
|
-
fn main(
|
|
747
|
-
@builtin(workgroup_id) wgid: vec3u,
|
|
748
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
749
|
-
@builtin(num_workgroups) nwg: vec3u,
|
|
750
|
-
) {
|
|
751
|
-
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
752
|
-
var acc = 0.0f;
|
|
753
|
-
|
|
754
|
-
// y[i] = \u03A3_j A[i,j] * x[j]
|
|
755
|
-
for (var j = lid.x; j < params.n; j += WGS) {
|
|
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
|
-
}
|
|
774
|
-
|
|
775
|
-
// Parallel reduction: 64 \u2192 1
|
|
776
|
-
scratch[lid.x] = acc;
|
|
777
|
-
workgroupBarrier();
|
|
778
|
-
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
779
|
-
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
780
|
-
workgroupBarrier();
|
|
781
|
-
}
|
|
782
|
-
|
|
783
|
-
if lid.x == 0u {
|
|
784
|
-
y[i * params.incy] = params.alpha * scratch[0] + params.beta * y[i * params.incy];
|
|
785
|
-
}
|
|
786
|
-
}
|
|
579
|
+
// Dekker's Fast-Two-Sum: same contract, but only correct when |a| >= |b|.
|
|
580
|
+
// UNPROTECTED \u2014 see header above.
|
|
581
|
+
fn fastTwoSum(a: f32, b: f32) -> DD {
|
|
582
|
+
let s = a + b;
|
|
583
|
+
let e = b - (s - a);
|
|
584
|
+
return DD(s, e);
|
|
787
585
|
}
|
|
788
|
-
`});var Xe,Ye=O(()=>{Xe=`// strmv: y = op(A) * x
|
|
789
|
-
// A is n\xD7n triangular, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
790
|
-
// op(A) is A (trans=0) or A^T (trans=1).
|
|
791
|
-
// diag=1 (unit) treats the diagonal as 1 without reading A's diagonal values.
|
|
792
|
-
// One workgroup per row, grid-stride outer loop.
|
|
793
|
-
|
|
794
|
-
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
795
|
-
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
796
|
-
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
797
586
|
|
|
798
|
-
|
|
799
|
-
|
|
800
|
-
|
|
801
|
-
|
|
802
|
-
|
|
803
|
-
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
804
|
-
uplo: u32, // 0 = lower, 1 = upper
|
|
805
|
-
diag: u32, // 0 = non-unit, 1 = unit
|
|
587
|
+
// Double-double addition (Dekker's Add2). UNPROTECTED \u2014 see header above.
|
|
588
|
+
fn ddAdd(a: DD, b: DD) -> DD {
|
|
589
|
+
let s = twoSum(a.hi, b.hi);
|
|
590
|
+
let loSum = a.lo + b.lo;
|
|
591
|
+
return fastTwoSum(s.hi, s.lo + loSum);
|
|
806
592
|
}
|
|
807
593
|
|
|
808
|
-
|
|
809
|
-
|
|
810
|
-
|
|
811
|
-
|
|
812
|
-
|
|
813
|
-
|
|
814
|
-
|
|
815
|
-
|
|
816
|
-
|
|
817
|
-
|
|
818
|
-
|
|
819
|
-
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
820
|
-
var acc = 0.0f;
|
|
594
|
+
// \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
|
|
595
|
+
//
|
|
596
|
+
// Bitcast subtraction + workgroup-barrier materialization, verified correct
|
|
597
|
+
// on all three backends tested. Costs a real barrier: fine for O(1)-per-
|
|
598
|
+
// thread or O(log n) reduction use, not a long per-element loop. A
|
|
599
|
+
// workgroupBarrier() requires uniform control flow, so:
|
|
600
|
+
// - \`threadSlot\` must be unique per concurrent caller (e.g. local_invocation_index).
|
|
601
|
+
// - Every thread in the workgroup must call this the same number of times
|
|
602
|
+
// \u2014 including ones whose result gets discarded. Compute unconditionally;
|
|
603
|
+
// only the write-back should be conditional.
|
|
604
|
+
var<workgroup> dekkerScratch: array<f32, 64>;
|
|
821
605
|
|
|
822
|
-
|
|
823
|
-
|
|
824
|
-
|
|
825
|
-
|
|
826
|
-
|
|
827
|
-
|
|
828
|
-
|
|
829
|
-
|
|
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
|
-
}
|
|
606
|
+
fn twoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
|
|
607
|
+
dekkerScratch[threadSlot] = a + b;
|
|
608
|
+
workgroupBarrier();
|
|
609
|
+
let s = dekkerScratch[threadSlot];
|
|
610
|
+
let v = fsub(s, a);
|
|
611
|
+
let e = fsub(a, fsub(s, v)) + fsub(b, v);
|
|
612
|
+
return DD(s, e);
|
|
613
|
+
}
|
|
877
614
|
|
|
878
|
-
|
|
879
|
-
|
|
880
|
-
|
|
881
|
-
|
|
882
|
-
|
|
883
|
-
|
|
884
|
-
|
|
615
|
+
fn fastTwoSumProtected(a: f32, b: f32, threadSlot: u32) -> DD {
|
|
616
|
+
dekkerScratch[threadSlot] = a + b;
|
|
617
|
+
workgroupBarrier();
|
|
618
|
+
let s = dekkerScratch[threadSlot];
|
|
619
|
+
let e = fsub(b, fsub(s, a));
|
|
620
|
+
return DD(s, e);
|
|
621
|
+
}
|
|
885
622
|
|
|
886
|
-
|
|
887
|
-
|
|
888
|
-
|
|
889
|
-
|
|
623
|
+
// Protected double-double addition \u2014 same contract as ddAdd, but exact.
|
|
624
|
+
fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
625
|
+
let s = twoSumProtected(a.hi, b.hi, threadSlot);
|
|
626
|
+
let loSum = a.lo + b.lo;
|
|
627
|
+
return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
|
|
890
628
|
}
|
|
891
|
-
`});var
|
|
629
|
+
`});var ut,it=V(()=>{ut=`// dasum: sum(|x[i]|), double-double (Dekker). Same ILP=4 shape as sasum.wgsl;
|
|
630
|
+
// see f64/utils/add.wgsl for ddAddProtected and why plain ddAdd isn't safe.
|
|
631
|
+
// GpuVector input isn't pre-abs'd, so ddAbs() (f64/utils/abs.wgsl) applies
|
|
632
|
+
// unconditionally below.
|
|
892
633
|
|
|
893
|
-
@group(0) @binding(0) var<storage, read>
|
|
894
|
-
@group(0) @binding(1) var<storage, read>
|
|
895
|
-
@group(0) @binding(2) var<storage, read_write>
|
|
634
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
635
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
636
|
+
@group(0) @binding(2) var<storage, read_write> partialsHi: array<f32>;
|
|
637
|
+
@group(0) @binding(3) var<storage, read_write> partialsLo: array<f32>;
|
|
638
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
896
639
|
|
|
897
640
|
struct Params {
|
|
898
|
-
m: u32,
|
|
899
641
|
n: u32,
|
|
900
|
-
|
|
901
|
-
incx: u32,
|
|
902
|
-
incy: u32,
|
|
903
|
-
lda: u32,
|
|
642
|
+
x_inc: u32,
|
|
904
643
|
}
|
|
905
644
|
|
|
906
|
-
|
|
645
|
+
const WGS: u32 = 64;
|
|
907
646
|
|
|
908
|
-
|
|
647
|
+
var<workgroup> tile: array<DD, 64>;
|
|
909
648
|
|
|
910
649
|
@compute @workgroup_size(64)
|
|
911
|
-
fn
|
|
912
|
-
@builtin(
|
|
913
|
-
@builtin(local_invocation_id)
|
|
914
|
-
@builtin(
|
|
650
|
+
fn dasum_main(
|
|
651
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
652
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
653
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
654
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
915
655
|
) {
|
|
916
|
-
|
|
917
|
-
|
|
918
|
-
|
|
919
|
-
|
|
920
|
-
// 4-unrolled loop: each iteration issues 4 independent A/y accesses.
|
|
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
|
-
}
|
|
937
|
-
}
|
|
938
|
-
}
|
|
939
|
-
`});var Je,$e=O(()=>{Je=`// ssyr: A := alpha * x * x^T + A (symmetric rank-1 update)
|
|
940
|
-
// A is n\xD7n symmetric; only the triangle specified by uplo is referenced/updated,
|
|
941
|
-
// the other triangle is implied by symmetry (not touched).
|
|
942
|
-
|
|
943
|
-
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
944
|
-
@group(0) @binding(1) var<storage, read_write> A: array<f32>;
|
|
656
|
+
var acc0 = DD(0.0, 0.0);
|
|
657
|
+
var acc1 = DD(0.0, 0.0);
|
|
658
|
+
var acc2 = DD(0.0, 0.0);
|
|
659
|
+
var acc3 = DD(0.0, 0.0);
|
|
945
660
|
|
|
946
|
-
|
|
947
|
-
n
|
|
948
|
-
alpha: f32,
|
|
949
|
-
incx: u32,
|
|
950
|
-
lda: u32,
|
|
951
|
-
uplo: u32, // 0 = lower, 1 = upper
|
|
952
|
-
}
|
|
661
|
+
let stride = num_wg.x * WGS;
|
|
662
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
953
663
|
|
|
954
|
-
|
|
664
|
+
// Same trip count for every thread, but driven by a counter, not \`id\`
|
|
665
|
+
// itself (ddAddProtected's barrier needs a provably-uniform loop bound).
|
|
666
|
+
let mainIters = n4_floor / (4u * stride);
|
|
667
|
+
for (var iter = 0u; iter < mainIters; iter++) {
|
|
668
|
+
let id = gid.x + iter * 4u * stride;
|
|
669
|
+
let i0 = id * params.x_inc;
|
|
670
|
+
let i1 = (id + stride) * params.x_inc;
|
|
671
|
+
let i2 = (id + 2u * stride) * params.x_inc;
|
|
672
|
+
let i3 = (id + 3u * stride) * params.x_inc;
|
|
673
|
+
acc0 = ddAddProtected(acc0, ddAbs(DD(xHi[i0], xLo[i0])), lid.x);
|
|
674
|
+
acc1 = ddAddProtected(acc1, ddAbs(DD(xHi[i1], xLo[i1])), lid.x);
|
|
675
|
+
acc2 = ddAddProtected(acc2, ddAbs(DD(xHi[i2], xLo[i2])), lid.x);
|
|
676
|
+
acc3 = ddAddProtected(acc3, ddAbs(DD(xHi[i3], xLo[i3])), lid.x);
|
|
677
|
+
}
|
|
955
678
|
|
|
956
|
-
|
|
679
|
+
// Tail is ragged (0-3 extra per thread) \u2014 pad to this workgroup's worst case.
|
|
680
|
+
let wgBaseGid = wgid.x * WGS;
|
|
681
|
+
var tailIters = 0u;
|
|
682
|
+
if (n4_floor + wgBaseGid < params.n) {
|
|
683
|
+
tailIters = (params.n - 1u - n4_floor - wgBaseGid) / stride + 1u;
|
|
684
|
+
}
|
|
685
|
+
for (var iter = 0u; iter < tailIters; iter++) {
|
|
686
|
+
let id = n4_floor + gid.x + iter * stride;
|
|
687
|
+
let valid = id < params.n;
|
|
688
|
+
let i = select(0u, id * params.x_inc, valid);
|
|
689
|
+
let loaded = ddAbs(DD(xHi[i], xLo[i])); // select() has no DD overload
|
|
690
|
+
let contribution = DD(select(0.0, loaded.hi, valid), select(0.0, loaded.lo, valid));
|
|
691
|
+
acc0 = ddAddProtected(acc0, contribution, lid.x);
|
|
692
|
+
}
|
|
957
693
|
|
|
958
|
-
|
|
959
|
-
|
|
960
|
-
|
|
961
|
-
|
|
962
|
-
@builtin(num_workgroups) nwg: vec3u,
|
|
963
|
-
) {
|
|
964
|
-
for (var row = wgid.x; row < params.n; row += nwg.x) {
|
|
965
|
-
let xi = params.alpha * x[row * params.incx];
|
|
966
|
-
let row_base = row * params.lda;
|
|
694
|
+
let combined01 = ddAddProtected(acc0, acc1, lid.x);
|
|
695
|
+
let combined23 = ddAddProtected(acc2, acc3, lid.x);
|
|
696
|
+
tile[lid.x] = ddAddProtected(combined01, combined23, lid.x);
|
|
697
|
+
workgroupBarrier();
|
|
967
698
|
|
|
968
|
-
|
|
969
|
-
|
|
970
|
-
|
|
971
|
-
|
|
972
|
-
|
|
973
|
-
|
|
974
|
-
|
|
975
|
-
|
|
976
|
-
|
|
977
|
-
}
|
|
699
|
+
// Inactive threads combine against a throwaway partner and discard it
|
|
700
|
+
// (ddAddProtected must be called unconditionally by every thread).
|
|
701
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
702
|
+
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
703
|
+
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
704
|
+
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
705
|
+
if (lid.x < s) { tile[lid.x] = combined; }
|
|
706
|
+
workgroupBarrier();
|
|
707
|
+
}
|
|
978
708
|
|
|
979
|
-
|
|
980
|
-
|
|
981
|
-
|
|
982
|
-
for (var col: u32 = colStart + lid.x; col < n4_floor; col += 4u * WGS) {
|
|
983
|
-
let idx0 = row_base + col;
|
|
984
|
-
let idx1 = row_base + col + WGS;
|
|
985
|
-
let idx2 = row_base + col + 2u * WGS;
|
|
986
|
-
let idx3 = row_base + col + 3u * WGS;
|
|
987
|
-
A[idx0] = xi * x[ col * params.incx] + A[idx0];
|
|
988
|
-
A[idx1] = xi * x[(col + WGS) * params.incx] + A[idx1];
|
|
989
|
-
A[idx2] = xi * x[(col + 2u * WGS) * params.incx] + A[idx2];
|
|
990
|
-
A[idx3] = xi * x[(col + 3u * WGS) * params.incx] + A[idx3];
|
|
991
|
-
}
|
|
992
|
-
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
993
|
-
for (var col: u32 = n4_floor + lid.x; col < colEnd; col += WGS) {
|
|
994
|
-
let idx = row_base + col;
|
|
995
|
-
A[idx] = xi * x[col * params.incx] + A[idx];
|
|
996
|
-
}
|
|
709
|
+
if (lid.x == 0u) {
|
|
710
|
+
partialsHi[wgid.x] = tile[0].hi;
|
|
711
|
+
partialsLo[wgid.x] = tile[0].lo;
|
|
997
712
|
}
|
|
998
713
|
}
|
|
999
|
-
`});var
|
|
1000
|
-
//
|
|
1001
|
-
//
|
|
1002
|
-
|
|
1003
|
-
|
|
1004
|
-
|
|
1005
|
-
@group(0) @binding(2) var<storage, read_write> A: array<f32>;
|
|
714
|
+
`});var mt,lt=V(()=>{mt=`// sum reduction (f64, double-double): collapses 2*WGS partial (hi, lo) pairs
|
|
715
|
+
// into one, using ddAddProtected instead of plain f32 \`+\` (see
|
|
716
|
+
// reduction/sum.wgsl for the f32 original this mirrors).
|
|
717
|
+
// dispatch: 1 workgroup of WGS threads. partialsHi/partialsLo must have
|
|
718
|
+
// exactly 2*WGS entries each. Concatenated after f64/dekker.wgsl (DD struct)
|
|
719
|
+
// and f64/utils/add.wgsl (ddAddProtected \u2014 see it for why plain ddAdd isn't safe).
|
|
1006
720
|
|
|
1007
|
-
|
|
1008
|
-
|
|
1009
|
-
|
|
1010
|
-
|
|
1011
|
-
incy: u32,
|
|
1012
|
-
lda: u32,
|
|
1013
|
-
uplo: u32, // 0 = lower, 1 = upper
|
|
1014
|
-
}
|
|
721
|
+
@group(0) @binding(0) var<storage, read> partialsHi: array<f32>;
|
|
722
|
+
@group(0) @binding(1) var<storage, read> partialsLo: array<f32>;
|
|
723
|
+
@group(0) @binding(2) var<storage, read_write> resultHi: array<f32, 1>;
|
|
724
|
+
@group(0) @binding(3) var<storage, read_write> resultLo: array<f32, 1>;
|
|
1015
725
|
|
|
1016
|
-
|
|
726
|
+
const WGS: u32 = 64;
|
|
1017
727
|
|
|
1018
|
-
|
|
728
|
+
var<workgroup> tile: array<DD, 64>;
|
|
1019
729
|
|
|
1020
730
|
@compute @workgroup_size(64)
|
|
1021
|
-
fn
|
|
1022
|
-
@builtin(
|
|
1023
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
1024
|
-
@builtin(num_workgroups) nwg: vec3u,
|
|
731
|
+
fn reduce_f64(
|
|
732
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1025
733
|
) {
|
|
1026
|
-
|
|
1027
|
-
|
|
1028
|
-
|
|
1029
|
-
|
|
734
|
+
let i = lid.x;
|
|
735
|
+
let a = DD(partialsHi[i], partialsLo[i]);
|
|
736
|
+
let b = DD(partialsHi[i + WGS], partialsLo[i + WGS]);
|
|
737
|
+
tile[i] = ddAddProtected(a, b, i);
|
|
738
|
+
workgroupBarrier();
|
|
1030
739
|
|
|
1031
|
-
|
|
1032
|
-
|
|
1033
|
-
|
|
1034
|
-
|
|
1035
|
-
|
|
1036
|
-
|
|
1037
|
-
|
|
1038
|
-
|
|
1039
|
-
colEnd = row + 1u;
|
|
1040
|
-
}
|
|
740
|
+
// ddAddProtected must be called unconditionally by every thread.
|
|
741
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
742
|
+
let partner = select(i, i + s, i < s);
|
|
743
|
+
let combined = ddAddProtected(tile[i], tile[partner], i);
|
|
744
|
+
workgroupBarrier();
|
|
745
|
+
if (i < s) { tile[i] = combined; }
|
|
746
|
+
workgroupBarrier();
|
|
747
|
+
}
|
|
1041
748
|
|
|
1042
|
-
|
|
1043
|
-
|
|
1044
|
-
|
|
1045
|
-
for (var col: u32 = colStart + lid.x; col < n4_floor; col += 4u * WGS) {
|
|
1046
|
-
let idx0 = row_base + col;
|
|
1047
|
-
let idx1 = row_base + col + WGS;
|
|
1048
|
-
let idx2 = row_base + col + 2u * WGS;
|
|
1049
|
-
let idx3 = row_base + col + 3u * WGS;
|
|
1050
|
-
A[idx0] = xi * y[ col * params.incy] + yi * x[ col * params.incx] + A[idx0];
|
|
1051
|
-
A[idx1] = xi * y[(col + WGS) * params.incy] + yi * x[(col + WGS) * params.incx] + A[idx1];
|
|
1052
|
-
A[idx2] = xi * y[(col + 2u * WGS) * params.incy] + yi * x[(col + 2u * WGS) * params.incx] + A[idx2];
|
|
1053
|
-
A[idx3] = xi * y[(col + 3u * WGS) * params.incy] + yi * x[(col + 3u * WGS) * params.incx] + A[idx3];
|
|
1054
|
-
}
|
|
1055
|
-
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
1056
|
-
for (var col: u32 = n4_floor + lid.x; col < colEnd; col += WGS) {
|
|
1057
|
-
let idx = row_base + col;
|
|
1058
|
-
A[idx] = xi * y[col * params.incy] + yi * x[col * params.incx] + A[idx];
|
|
1059
|
-
}
|
|
749
|
+
if (i == 0u) {
|
|
750
|
+
resultHi[0] = tile[0].hi;
|
|
751
|
+
resultLo[0] = tile[0].lo;
|
|
1060
752
|
}
|
|
1061
753
|
}
|
|
1062
|
-
`});var
|
|
1063
|
-
// value, aux: raw u32 bits \u2014 see src/util/f64pack.mjs; decode()/encode()
|
|
1064
|
-
// below are the WGSL mirror of that file's packedToFields()/fieldsToPacked()),
|
|
1065
|
-
// producing the sum as another [main, aux] pair.
|
|
1066
|
-
//
|
|
1067
|
-
// Implements IEEE-754 binary64 addition (align, add/subtract significands,
|
|
1068
|
-
// normalize, round-to-nearest-even) using only u32 bitwise/integer
|
|
1069
|
-
// arithmetic \u2014 WGSL has no 64-bit integer type or arbitrary-precision
|
|
1070
|
-
// integers, so each operand's 53-bit significand is carried as a two-word
|
|
1071
|
-
// (hi, lo) pair, widened by 3 bits at the bottom to hold guard/round/sticky
|
|
1072
|
-
// information while aligning exponents.
|
|
1073
|
-
|
|
1074
|
-
const EXP_ALL_ONES: u32 = 0x7ffu;
|
|
1075
|
-
const BIAS: i32 = 1023;
|
|
1076
|
-
const QUIET_NAN_MANTISSA_HI: u32 = 1u << 19u; // bit51 of the 52-bit mantissa -> canonical quiet NaN
|
|
754
|
+
`});var ct,ft=V(()=>{ct=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
1077
755
|
|
|
1078
|
-
|
|
1079
|
-
|
|
1080
|
-
|
|
1081
|
-
|
|
1082
|
-
|
|
756
|
+
// a > b for double-double pairs. hi dominates (|lo| <= ulp(hi)/2 always), so
|
|
757
|
+
// comparing hi alone is correct except on an exact hi tie, when lo breaks it.
|
|
758
|
+
// A plain comparison, not a rounding-identity subtraction \u2014 no reassociation
|
|
759
|
+
// risk, so unlike twoSum/fastTwoSum this needs no protection.
|
|
760
|
+
fn ddGreater(a: DD, b: DD) -> bool {
|
|
761
|
+
if (a.hi != b.hi) {
|
|
762
|
+
return a.hi > b.hi;
|
|
763
|
+
}
|
|
764
|
+
return a.lo > b.lo;
|
|
1083
765
|
}
|
|
766
|
+
`});var dt,pt=V(()=>{dt=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
1084
767
|
|
|
1085
|
-
//
|
|
1086
|
-
//
|
|
1087
|
-
|
|
1088
|
-
|
|
1089
|
-
// comment above fieldsToPacked, and dasum.wgsl's xAux/partialsAux bindings.
|
|
1090
|
-
struct Packed {
|
|
1091
|
-
main: f32,
|
|
1092
|
-
aux: u32,
|
|
768
|
+
// a == b for double-double pairs \u2014 exact field equality, no rounding
|
|
769
|
+
// involved, so (like ddGreater) this needs no protection.
|
|
770
|
+
fn ddEqual(a: DD, b: DD) -> bool {
|
|
771
|
+
return a.hi == b.hi && a.lo == b.lo;
|
|
1093
772
|
}
|
|
773
|
+
`});var gt,wt=V(()=>{gt=`// idamax: returns index of element with largest absolute value (f64, double-double)
|
|
774
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/argmaxF64.wgsl.
|
|
775
|
+
// Concatenated after f64/dekker.wgsl (DD struct), f64/utils/abs.wgsl (ddAbs),
|
|
776
|
+
// f64/utils/greater.wgsl (ddGreater), and f64/utils/equal.wgsl (ddEqual).
|
|
1094
777
|
|
|
1095
|
-
|
|
1096
|
-
|
|
1097
|
-
|
|
1098
|
-
|
|
1099
|
-
|
|
778
|
+
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
779
|
+
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
780
|
+
@group(0) @binding(2) var<storage, read_write> partialsValHi: array<f32>;
|
|
781
|
+
@group(0) @binding(3) var<storage, read_write> partialsValLo: array<f32>;
|
|
782
|
+
@group(0) @binding(4) var<storage, read_write> partialsIdx: array<u32>;
|
|
783
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
1100
784
|
|
|
1101
|
-
|
|
1102
|
-
|
|
1103
|
-
|
|
785
|
+
struct Params {
|
|
786
|
+
n: u32,
|
|
787
|
+
x_inc: u32,
|
|
788
|
+
}
|
|
1104
789
|
|
|
1105
|
-
|
|
1106
|
-
let mantExtra29 = ((auxExp8 & 0x3fu) << 23u) | auxMant23;
|
|
790
|
+
const WGS: u32 = 64;
|
|
1107
791
|
|
|
1108
|
-
|
|
1109
|
-
|
|
1110
|
-
let mantTop3 = mantMain & 0x7u;
|
|
1111
|
-
let lo = (mantTop3 << 29u) | mantExtra29;
|
|
792
|
+
var<workgroup> tile_val: array<DD, 64>;
|
|
793
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
1112
794
|
|
|
1113
|
-
|
|
1114
|
-
|
|
795
|
+
@compute @workgroup_size(64)
|
|
796
|
+
fn idamax_main(
|
|
797
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
798
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
799
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
800
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
801
|
+
) {
|
|
802
|
+
// DD(-1.0, 0.0) is a safe sentinel: any |x[i]| >= 0 beats it,
|
|
803
|
+
// so workgroups with no elements lose gracefully in the epilogue.
|
|
804
|
+
var best_val0 = DD(-1.0, 0.0); var best_idx0: u32 = 0u;
|
|
805
|
+
var best_val1 = DD(-1.0, 0.0); var best_idx1: u32 = 0u;
|
|
806
|
+
var best_val2 = DD(-1.0, 0.0); var best_idx2: u32 = 0u;
|
|
807
|
+
var best_val3 = DD(-1.0, 0.0); var best_idx3: u32 = 0u;
|
|
1115
808
|
|
|
1116
|
-
|
|
1117
|
-
|
|
1118
|
-
let expMain = rawExp >> 3u;
|
|
1119
|
-
let expExtra = rawExp & 0x7u;
|
|
809
|
+
let stride = num_wg.x * WGS;
|
|
810
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
1120
811
|
|
|
1121
|
-
|
|
1122
|
-
|
|
1123
|
-
|
|
812
|
+
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
813
|
+
let i0 = id * params.x_inc;
|
|
814
|
+
let i1 = (id + stride) * params.x_inc;
|
|
815
|
+
let i2 = (id + 2u * stride) * params.x_inc;
|
|
816
|
+
let i3 = (id + 3u * stride) * params.x_inc;
|
|
817
|
+
let v0 = ddAbs(DD(xHi[i0], xLo[i0]));
|
|
818
|
+
let v1 = ddAbs(DD(xHi[i1], xLo[i1]));
|
|
819
|
+
let v2 = ddAbs(DD(xHi[i2], xLo[i2]));
|
|
820
|
+
let v3 = ddAbs(DD(xHi[i3], xLo[i3]));
|
|
821
|
+
if (ddGreater(v0, best_val0)) { best_val0 = v0; best_idx0 = id; }
|
|
822
|
+
if (ddGreater(v1, best_val1)) { best_val1 = v1; best_idx1 = id + stride; }
|
|
823
|
+
if (ddGreater(v2, best_val2)) { best_val2 = v2; best_idx2 = id + 2u * stride; }
|
|
824
|
+
if (ddGreater(v3, best_val3)) { best_val3 = v3; best_idx3 = id + 3u * stride; }
|
|
825
|
+
}
|
|
826
|
+
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
827
|
+
let i = id * params.x_inc;
|
|
828
|
+
let v = ddAbs(DD(xHi[i], xLo[i]));
|
|
829
|
+
if (ddGreater(v, best_val0)) { best_val0 = v; best_idx0 = id; }
|
|
830
|
+
}
|
|
1124
831
|
|
|
1125
|
-
|
|
832
|
+
// merge 4 independent lanes; prefer lower index on tie (first occurrence wins)
|
|
833
|
+
if (ddGreater(best_val1, best_val0) ||
|
|
834
|
+
(ddEqual(best_val1, best_val0) && best_idx1 < best_idx0)) {
|
|
835
|
+
best_val0 = best_val1; best_idx0 = best_idx1;
|
|
836
|
+
}
|
|
837
|
+
if (ddGreater(best_val2, best_val0) ||
|
|
838
|
+
(ddEqual(best_val2, best_val0) && best_idx2 < best_idx0)) {
|
|
839
|
+
best_val0 = best_val2; best_idx0 = best_idx2;
|
|
840
|
+
}
|
|
841
|
+
if (ddGreater(best_val3, best_val0) ||
|
|
842
|
+
(ddEqual(best_val3, best_val0) && best_idx3 < best_idx0)) {
|
|
843
|
+
best_val0 = best_val3; best_idx0 = best_idx3;
|
|
844
|
+
}
|
|
1126
845
|
|
|
1127
|
-
|
|
1128
|
-
|
|
1129
|
-
|
|
1130
|
-
let auxMant23 = mantExtra29 & 0x7fffffu;
|
|
1131
|
-
let auxExp8 = (auxExpTop2 << 6u) | auxExpBot6;
|
|
846
|
+
tile_val[lid.x] = best_val0;
|
|
847
|
+
tile_idx[lid.x] = best_idx0;
|
|
848
|
+
workgroupBarrier();
|
|
1132
849
|
|
|
1133
|
-
|
|
850
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
851
|
+
if (lid.x < s) {
|
|
852
|
+
let a_val = tile_val[lid.x];
|
|
853
|
+
let b_val = tile_val[lid.x + s];
|
|
854
|
+
if (ddGreater(b_val, a_val) ||
|
|
855
|
+
(ddEqual(b_val, a_val) && tile_idx[lid.x + s] < tile_idx[lid.x])) {
|
|
856
|
+
tile_val[lid.x] = b_val;
|
|
857
|
+
tile_idx[lid.x] = tile_idx[lid.x + s];
|
|
858
|
+
}
|
|
859
|
+
}
|
|
860
|
+
workgroupBarrier();
|
|
861
|
+
}
|
|
1134
862
|
|
|
1135
|
-
|
|
863
|
+
if (lid.x == 0u) {
|
|
864
|
+
partialsValHi[wgid.x] = tile_val[0].hi;
|
|
865
|
+
partialsValLo[wgid.x] = tile_val[0].lo;
|
|
866
|
+
partialsIdx[wgid.x] = tile_idx[0];
|
|
867
|
+
}
|
|
1136
868
|
}
|
|
869
|
+
`});var ht,bt=V(()=>{ht=`// amax reduction (f64, double-double): collapses 2*WGS (value, index) pairs
|
|
870
|
+
// into one index, using ddGreater/ddEqual instead of plain f32 \`>\`/\`==\` (see
|
|
871
|
+
// reduction/argmax.wgsl for the f32 original this mirrors).
|
|
872
|
+
// dispatch: 1 workgroup of WGS threads. partialsValHi/partialsValLo/
|
|
873
|
+
// partialsIdx must have exactly 2*WGS entries each. Concatenated after
|
|
874
|
+
// f64/dekker.wgsl (DD struct), f64/utils/greater.wgsl (ddGreater), and
|
|
875
|
+
// f64/utils/equal.wgsl (ddEqual).
|
|
876
|
+
|
|
877
|
+
@group(0) @binding(0) var<storage, read> partialsValHi: array<f32>;
|
|
878
|
+
@group(0) @binding(1) var<storage, read> partialsValLo: array<f32>;
|
|
879
|
+
@group(0) @binding(2) var<storage, read> partialsIdx: array<u32>;
|
|
880
|
+
@group(0) @binding(3) var<storage, read_write> result: array<u32>;
|
|
1137
881
|
|
|
1138
|
-
|
|
1139
|
-
struct Shifted { hi: u32, lo: u32, sticky: u32 }
|
|
882
|
+
const WGS: u32 = 64;
|
|
1140
883
|
|
|
1141
|
-
|
|
1142
|
-
|
|
1143
|
-
|
|
1144
|
-
|
|
1145
|
-
|
|
1146
|
-
|
|
1147
|
-
|
|
1148
|
-
|
|
1149
|
-
|
|
1150
|
-
|
|
1151
|
-
if (
|
|
1152
|
-
|
|
1153
|
-
|
|
1154
|
-
|
|
1155
|
-
|
|
884
|
+
var<workgroup> tile_val: array<DD, 64>;
|
|
885
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
886
|
+
|
|
887
|
+
@compute @workgroup_size(64)
|
|
888
|
+
fn reduce_f64(
|
|
889
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
890
|
+
) {
|
|
891
|
+
let i = lid.x;
|
|
892
|
+
let a_val = DD(partialsValHi[i], partialsValLo[i]);
|
|
893
|
+
let b_val = DD(partialsValHi[i + WGS], partialsValLo[i + WGS]);
|
|
894
|
+
if (ddGreater(b_val, a_val) ||
|
|
895
|
+
(ddEqual(b_val, a_val) && partialsIdx[i + WGS] < partialsIdx[i])) {
|
|
896
|
+
tile_val[i] = b_val;
|
|
897
|
+
tile_idx[i] = partialsIdx[i + WGS];
|
|
898
|
+
} else {
|
|
899
|
+
tile_val[i] = a_val;
|
|
900
|
+
tile_idx[i] = partialsIdx[i];
|
|
1156
901
|
}
|
|
1157
|
-
|
|
1158
|
-
|
|
902
|
+
workgroupBarrier();
|
|
903
|
+
|
|
904
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
905
|
+
if (i < s) {
|
|
906
|
+
let c_val = tile_val[i];
|
|
907
|
+
let d_val = tile_val[i + s];
|
|
908
|
+
if (ddGreater(d_val, c_val) ||
|
|
909
|
+
(ddEqual(d_val, c_val) && tile_idx[i + s] < tile_idx[i])) {
|
|
910
|
+
tile_val[i] = d_val;
|
|
911
|
+
tile_idx[i] = tile_idx[i + s];
|
|
912
|
+
}
|
|
913
|
+
}
|
|
914
|
+
workgroupBarrier();
|
|
1159
915
|
}
|
|
1160
|
-
|
|
1161
|
-
|
|
1162
|
-
let newLo = hi >> m;
|
|
1163
|
-
return Shifted(0u, newLo, select(0u, 1u, stickyBits != 0u));
|
|
916
|
+
|
|
917
|
+
if (i == 0u) { result[0] = tile_idx[0]; }
|
|
1164
918
|
}
|
|
919
|
+
`});var xt,yt=V(()=>{xt=`// srot: x = c*x + s*y, y = -s*x + c*y
|
|
1165
920
|
|
|
1166
|
-
|
|
1167
|
-
|
|
1168
|
-
|
|
1169
|
-
|
|
1170
|
-
|
|
1171
|
-
|
|
1172
|
-
|
|
1173
|
-
|
|
1174
|
-
|
|
1175
|
-
|
|
1176
|
-
|
|
1177
|
-
|
|
1178
|
-
|
|
1179
|
-
|
|
921
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
922
|
+
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
923
|
+
|
|
924
|
+
struct Params {
|
|
925
|
+
n: u32,
|
|
926
|
+
c: f32,
|
|
927
|
+
s: f32,
|
|
928
|
+
x_inc: u32,
|
|
929
|
+
y_inc: u32,
|
|
930
|
+
}
|
|
931
|
+
|
|
932
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
933
|
+
|
|
934
|
+
const WGS: u32 = 64;
|
|
935
|
+
|
|
936
|
+
@compute @workgroup_size(64)
|
|
937
|
+
fn main(
|
|
938
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
939
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
940
|
+
) {
|
|
941
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
942
|
+
let xi = x[id * params.x_inc];
|
|
943
|
+
let yi = y[id * params.y_inc];
|
|
944
|
+
x[id * params.x_inc] = params.c * xi + params.s * yi;
|
|
945
|
+
y[id * params.y_inc] = -params.s * xi + params.c * yi;
|
|
1180
946
|
}
|
|
1181
|
-
let m = n - 32u;
|
|
1182
|
-
return Pair(lo << m, 0u);
|
|
1183
947
|
}
|
|
948
|
+
`});var _t,vt=V(()=>{_t=`// srotm: applies modified Givens rotation H to vectors x and y.
|
|
949
|
+
// param[0] = flag: -1 (full H), 0 (unit diagonal), 1 (unit off-diagonal)
|
|
950
|
+
// param = [ flag, h11, h21, h12, h22 ]
|
|
951
|
+
// flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
|
|
1184
952
|
|
|
1185
|
-
|
|
1186
|
-
|
|
1187
|
-
|
|
1188
|
-
let sumHi = aHi + bHi + carry;
|
|
1189
|
-
return Pair(sumHi, sumLo);
|
|
1190
|
-
}
|
|
953
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
954
|
+
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
955
|
+
@group(0) @binding(2) var<storage, read> param: array<f32>;
|
|
1191
956
|
|
|
1192
|
-
|
|
1193
|
-
|
|
1194
|
-
|
|
1195
|
-
|
|
1196
|
-
let diffHi = aHi - bHi - borrow;
|
|
1197
|
-
return Pair(diffHi, diffLo);
|
|
957
|
+
struct Params {
|
|
958
|
+
n: u32,
|
|
959
|
+
x_inc: u32,
|
|
960
|
+
y_inc: u32,
|
|
1198
961
|
}
|
|
1199
962
|
|
|
1200
|
-
|
|
1201
|
-
return aHi > bHi || (aHi == bHi && aLo >= bLo);
|
|
1202
|
-
}
|
|
963
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1203
964
|
|
|
1204
|
-
|
|
1205
|
-
// encoded Packed pair \u2014 lets a caller that's accumulating many values in a
|
|
1206
|
-
// row (e.g. dasum.wgsl's per-thread reduction loop) keep the running total
|
|
1207
|
-
// in Fields form the whole time, only encoding once at the very end, instead
|
|
1208
|
-
// of paying a decode+encode round-trip on every single addition. computeSum
|
|
1209
|
-
// (below) is the Packed-in/Packed-out convenience wrapper around this.
|
|
1210
|
-
fn addFields(a: Fields, b: Fields) -> Fields {
|
|
1211
|
-
let aIsNaN = a.rawExp == EXP_ALL_ONES && (a.mantissaHi != 0u || a.lo != 0u);
|
|
1212
|
-
let bIsNaN = b.rawExp == EXP_ALL_ONES && (b.mantissaHi != 0u || b.lo != 0u);
|
|
1213
|
-
if (aIsNaN || bIsNaN) {
|
|
1214
|
-
return Fields(0u, EXP_ALL_ONES, QUIET_NAN_MANTISSA_HI, 0u);
|
|
1215
|
-
}
|
|
965
|
+
const WGS: u32 = 64;
|
|
1216
966
|
|
|
1217
|
-
|
|
1218
|
-
|
|
1219
|
-
|
|
1220
|
-
|
|
1221
|
-
|
|
1222
|
-
|
|
1223
|
-
return Fields(a.sign, EXP_ALL_ONES, 0u, 0u);
|
|
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;
|
|
1264
|
-
}
|
|
967
|
+
@compute @workgroup_size(64)
|
|
968
|
+
fn main(
|
|
969
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
970
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
971
|
+
) {
|
|
972
|
+
let flag = param[0];
|
|
1265
973
|
|
|
1266
|
-
|
|
1267
|
-
|
|
1268
|
-
let alignedHiQ = shiftedQ.hi;
|
|
1269
|
-
let alignedLoQ = shiftedQ.lo | shiftedQ.sticky; // fold sticky into bit0
|
|
974
|
+
var h11: f32; var h12: f32;
|
|
975
|
+
var h21: f32; var h22: f32;
|
|
1270
976
|
|
|
1271
|
-
|
|
1272
|
-
|
|
1273
|
-
|
|
1274
|
-
|
|
1275
|
-
} else {
|
|
1276
|
-
|
|
1277
|
-
|
|
977
|
+
if (flag == -1.0) {
|
|
978
|
+
// full 2x2 matrix
|
|
979
|
+
h11 = param[1]; h21 = param[2];
|
|
980
|
+
h12 = param[3]; h22 = param[4];
|
|
981
|
+
} else if (flag == 0.0) {
|
|
982
|
+
// diagonal fixed at 1
|
|
983
|
+
h11 = 1.0; h21 = param[2];
|
|
984
|
+
h12 = param[3]; h22 = 1.0;
|
|
985
|
+
} else if (flag == 1.0) {
|
|
986
|
+
// flag == 1.0: off-diagonal fixed at +1 / -1
|
|
987
|
+
h11 = param[1]; h21 = -1.0;
|
|
988
|
+
h12 = 1.0; h22 = param[4];
|
|
1278
989
|
}
|
|
1279
990
|
|
|
1280
|
-
|
|
1281
|
-
|
|
991
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
992
|
+
let xi = x[id * params.x_inc];
|
|
993
|
+
let yi = y[id * params.y_inc];
|
|
994
|
+
x[id * params.x_inc] = h11 * xi + h12 * yi;
|
|
995
|
+
y[id * params.y_inc] = h21 * xi + h22 * yi;
|
|
1282
996
|
}
|
|
997
|
+
}
|
|
998
|
+
`});var Et,Bt=V(()=>{Et=`// sgemv_n: y = alpha * A * x + beta * y (A is m\xD7n row-major, no-transpose)
|
|
999
|
+
//
|
|
1000
|
+
// One workgroup per output row, with a grid-stride outer loop so the shader
|
|
1001
|
+
// still covers all rows when m exceeds maxComputeWorkgroupsPerDimension.
|
|
1002
|
+
// Threads stride through A[row, :] and x with coalesced reads (consecutive
|
|
1003
|
+
// threads \u2192 consecutive addresses). Four independent accumulators let the GPU
|
|
1004
|
+
// pipeline memory requests across iterations (ILP=4), hiding the
|
|
1005
|
+
// global-memory latency.
|
|
1283
1006
|
|
|
1284
|
-
|
|
1285
|
-
|
|
1007
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1008
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
1009
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
1286
1010
|
|
|
1287
|
-
|
|
1288
|
-
|
|
1289
|
-
|
|
1290
|
-
|
|
1291
|
-
|
|
1292
|
-
|
|
1293
|
-
|
|
1294
|
-
|
|
1295
|
-
|
|
1296
|
-
let shiftAmt = targetLSBScale - commonExp2;
|
|
1011
|
+
struct Params {
|
|
1012
|
+
m: u32,
|
|
1013
|
+
n: u32,
|
|
1014
|
+
alpha: f32,
|
|
1015
|
+
beta: f32,
|
|
1016
|
+
incx: u32,
|
|
1017
|
+
incy: u32,
|
|
1018
|
+
lda: u32,
|
|
1019
|
+
}
|
|
1297
1020
|
|
|
1298
|
-
|
|
1299
|
-
if (shiftAmt <= 0) {
|
|
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
|
-
}
|
|
1021
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1317
1022
|
|
|
1318
|
-
|
|
1319
|
-
|
|
1320
|
-
let sh = shr_sticky(keepHi, keepLo, 1u); // dropped bit is guaranteed 0 here
|
|
1321
|
-
keepHi = sh.hi; keepLo = sh.lo;
|
|
1322
|
-
resultExpBase = resultExpBase + 1;
|
|
1323
|
-
}
|
|
1023
|
+
const WGS: u32 = 64u;
|
|
1024
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
1324
1025
|
|
|
1325
|
-
|
|
1326
|
-
|
|
1327
|
-
|
|
1328
|
-
|
|
1329
|
-
|
|
1330
|
-
|
|
1331
|
-
|
|
1332
|
-
|
|
1333
|
-
|
|
1334
|
-
|
|
1335
|
-
|
|
1026
|
+
@compute @workgroup_size(64)
|
|
1027
|
+
fn main(
|
|
1028
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1029
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1030
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
1031
|
+
) {
|
|
1032
|
+
// Grid-stride loop: each workgroup handles ceil(m / nwg.x) rows.
|
|
1033
|
+
for (var row = wgid.x; row < params.m; row += nwg.x) {
|
|
1034
|
+
let row_base = row * params.lda;
|
|
1035
|
+
var acc0: f32 = 0.0;
|
|
1036
|
+
var acc1: f32 = 0.0;
|
|
1037
|
+
var acc2: f32 = 0.0;
|
|
1038
|
+
var acc3: f32 = 0.0;
|
|
1336
1039
|
|
|
1337
|
-
//
|
|
1338
|
-
//
|
|
1339
|
-
|
|
1340
|
-
|
|
1341
|
-
|
|
1342
|
-
|
|
1343
|
-
|
|
1344
|
-
|
|
1345
|
-
|
|
1346
|
-
|
|
1347
|
-
//
|
|
1348
|
-
|
|
1349
|
-
|
|
1350
|
-
|
|
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.
|
|
1040
|
+
// 4-unrolled loop: each iteration issues 4 independent loads for A and x.
|
|
1041
|
+
// The accumulators are independent so the GPU can overlap the memory
|
|
1042
|
+
// requests rather than serialising them behind a dependency chain.
|
|
1043
|
+
let n4_floor = (params.n / (4u * WGS)) * (4u * WGS);
|
|
1044
|
+
for (var j: u32 = lid.x; j < n4_floor; j += 4u * WGS) {
|
|
1045
|
+
acc0 += A[row_base + j ] * x[ j * params.incx];
|
|
1046
|
+
acc1 += A[row_base + j + WGS ] * x[(j + WGS) * params.incx];
|
|
1047
|
+
acc2 += A[row_base + j + 2u * WGS ] * x[(j + 2u * WGS) * params.incx];
|
|
1048
|
+
acc3 += A[row_base + j + 3u * WGS ] * x[(j + 3u * WGS) * params.incx];
|
|
1049
|
+
}
|
|
1050
|
+
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
1051
|
+
for (var j: u32 = n4_floor + lid.x; j < params.n; j += WGS) {
|
|
1052
|
+
acc0 += A[row_base + j] * x[j * params.incx];
|
|
1053
|
+
}
|
|
1356
1054
|
|
|
1357
|
-
|
|
1358
|
-
|
|
1359
|
-
|
|
1360
|
-
|
|
1361
|
-
|
|
1055
|
+
// Parallel reduction: 64 \u2192 32 \u2192 16 \u2192 8 \u2192 4 \u2192 2 \u2192 1
|
|
1056
|
+
scratch[lid.x] = acc0 + acc1 + acc2 + acc3;
|
|
1057
|
+
workgroupBarrier();
|
|
1058
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
1059
|
+
if lid.x < stride {
|
|
1060
|
+
scratch[lid.x] += scratch[lid.x + stride];
|
|
1061
|
+
}
|
|
1062
|
+
workgroupBarrier();
|
|
1063
|
+
}
|
|
1362
1064
|
|
|
1363
|
-
|
|
1364
|
-
|
|
1365
|
-
|
|
1366
|
-
|
|
1367
|
-
|
|
1065
|
+
if lid.x == 0u {
|
|
1066
|
+
let yi = row * params.incy;
|
|
1067
|
+
// BLAS beta==0 semantics: y is written, not accumulated \u2014 must not read y.
|
|
1068
|
+
let acc = params.alpha * scratch[0];
|
|
1069
|
+
y[yi] = select(acc, acc + params.beta * y[yi], params.beta != 0.0);
|
|
1070
|
+
}
|
|
1071
|
+
// All 64 threads must agree before the next row reuses scratch[].
|
|
1072
|
+
workgroupBarrier();
|
|
1368
1073
|
}
|
|
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
1074
|
}
|
|
1075
|
+
`});var Gt,At=V(()=>{Gt=`// sgemv_t: y = alpha * A^T * x + beta * y (A is m\xD7n row-major, transposed)
|
|
1076
|
+
// each thread owns one column of A \u2192 one element of y (length n)
|
|
1077
|
+
// tiles over x (length m) using shared memory; four independent accumulators
|
|
1078
|
+
// let the GPU pipeline A reads across j within each tile (ILP=4)
|
|
1397
1079
|
|
|
1398
|
-
|
|
1399
|
-
|
|
1400
|
-
|
|
1401
|
-
let s = a + b;
|
|
1402
|
-
let e = b - (s - a);
|
|
1403
|
-
return DD(s, e);
|
|
1404
|
-
}
|
|
1080
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1081
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
1082
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
1405
1083
|
|
|
1406
|
-
|
|
1407
|
-
|
|
1408
|
-
|
|
1409
|
-
|
|
1410
|
-
|
|
1084
|
+
struct Params {
|
|
1085
|
+
m: u32,
|
|
1086
|
+
n: u32,
|
|
1087
|
+
alpha: f32,
|
|
1088
|
+
beta: f32,
|
|
1089
|
+
incx: u32,
|
|
1090
|
+
incy: u32,
|
|
1091
|
+
lda: u32,
|
|
1411
1092
|
}
|
|
1412
1093
|
|
|
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>;
|
|
1094
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1424
1095
|
|
|
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
|
-
}
|
|
1096
|
+
const WGS: u32 = 64u;
|
|
1097
|
+
var<workgroup> x_tile: array<f32, 64>;
|
|
1433
1098
|
|
|
1434
|
-
|
|
1435
|
-
|
|
1436
|
-
|
|
1437
|
-
|
|
1438
|
-
|
|
1439
|
-
|
|
1440
|
-
|
|
1099
|
+
@compute @workgroup_size(64)
|
|
1100
|
+
fn main(
|
|
1101
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
1102
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1103
|
+
) {
|
|
1104
|
+
// each thread owns column col of A \u2192 output y[col]
|
|
1105
|
+
let col = gid.x;
|
|
1106
|
+
// tile over x (length m, the rows of A)
|
|
1107
|
+
let m_floor = (params.m / WGS) * WGS;
|
|
1108
|
+
var acc0: f32 = 0.0;
|
|
1109
|
+
var acc1: f32 = 0.0;
|
|
1110
|
+
var acc2: f32 = 0.0;
|
|
1111
|
+
var acc3: f32 = 0.0;
|
|
1441
1112
|
|
|
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.
|
|
1113
|
+
for (var base = 0u; base < m_floor; base += WGS) {
|
|
1114
|
+
// cooperative load: all 64 threads fill x_tile with x[base..base+WGS]
|
|
1115
|
+
x_tile[lid.x] = x[(base + lid.x) * params.incx];
|
|
1116
|
+
workgroupBarrier();
|
|
1449
1117
|
|
|
1450
|
-
//
|
|
1451
|
-
//
|
|
1452
|
-
|
|
1453
|
-
|
|
1454
|
-
|
|
1455
|
-
|
|
1456
|
-
|
|
1118
|
+
// 4-unrolled inner loop: 4 independent A reads let the GPU pipeline
|
|
1119
|
+
// global-memory requests within each tile. WGS=64 divides by 4 exactly.
|
|
1120
|
+
if (col < params.n) {
|
|
1121
|
+
for (var j = 0u; j < WGS; j += 4u) {
|
|
1122
|
+
acc0 += A[(base + j ) * params.lda + col] * x_tile[j ];
|
|
1123
|
+
acc1 += A[(base + j + 1) * params.lda + col] * x_tile[j + 1];
|
|
1124
|
+
acc2 += A[(base + j + 2) * params.lda + col] * x_tile[j + 2];
|
|
1125
|
+
acc3 += A[(base + j + 3) * params.lda + col] * x_tile[j + 3];
|
|
1126
|
+
}
|
|
1127
|
+
}
|
|
1128
|
+
workgroupBarrier();
|
|
1457
1129
|
}
|
|
1458
|
-
return a.lo > b.lo;
|
|
1459
|
-
}
|
|
1460
|
-
`});var dt,ct=O(()=>{dt=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
1461
1130
|
|
|
1462
|
-
|
|
1463
|
-
//
|
|
1464
|
-
|
|
1465
|
-
|
|
1131
|
+
if (col < params.n) {
|
|
1132
|
+
// remainder: m not divisible by WGS \u2014 short loop, single accumulator fine
|
|
1133
|
+
for (var k = m_floor; k < params.m; k++) {
|
|
1134
|
+
acc0 += A[k * params.lda + col] * x[k * params.incx];
|
|
1135
|
+
}
|
|
1136
|
+
let yi = col * params.incy;
|
|
1137
|
+
// BLAS beta==0 semantics: y is written, not accumulated \u2014 must not read y.
|
|
1138
|
+
let acc = params.alpha * (acc0 + acc1 + acc2 + acc3);
|
|
1139
|
+
y[yi] = select(acc, acc + params.beta * y[yi], params.beta != 0.0);
|
|
1140
|
+
}
|
|
1466
1141
|
}
|
|
1467
|
-
`});var
|
|
1468
|
-
//
|
|
1469
|
-
//
|
|
1470
|
-
//
|
|
1142
|
+
`});var kt,St=V(()=>{kt=`// ssymv: y = alpha * A * x + beta * y
|
|
1143
|
+
// A is n\xD7n symmetric, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
1144
|
+
// The logical matrix is fully dense (symmetric), so each row's dot product
|
|
1145
|
+
// sums over all n columns; entries on the unstored side of the diagonal are
|
|
1146
|
+
// fetched from their mirror position (A[i,j] == A[j,i]).
|
|
1147
|
+
// One workgroup per row, grid-stride outer loop.
|
|
1471
1148
|
|
|
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;
|
|
1149
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1150
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
1151
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
1477
1152
|
|
|
1478
1153
|
struct Params {
|
|
1479
1154
|
n: u32,
|
|
1480
|
-
|
|
1155
|
+
alpha: f32,
|
|
1156
|
+
beta: f32,
|
|
1157
|
+
incx: u32,
|
|
1158
|
+
incy: u32,
|
|
1159
|
+
lda: u32,
|
|
1160
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
1481
1161
|
}
|
|
1482
1162
|
|
|
1483
|
-
|
|
1163
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1484
1164
|
|
|
1485
|
-
|
|
1165
|
+
const WGS: u32 = 64u;
|
|
1166
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
1486
1167
|
|
|
1487
1168
|
@compute @workgroup_size(64)
|
|
1488
|
-
fn
|
|
1489
|
-
@builtin(
|
|
1490
|
-
@builtin(local_invocation_id)
|
|
1491
|
-
@builtin(
|
|
1492
|
-
@builtin(num_workgroups) num_wg: vec3u,
|
|
1169
|
+
fn main(
|
|
1170
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1171
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1172
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
1493
1173
|
) {
|
|
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
|
-
}
|
|
1174
|
+
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
1175
|
+
var acc = 0.0f;
|
|
1531
1176
|
|
|
1532
|
-
|
|
1533
|
-
|
|
1534
|
-
|
|
1535
|
-
|
|
1177
|
+
// y[i] = \u03A3_j A[i,j] * x[j]
|
|
1178
|
+
for (var j = lid.x; j < params.n; j += WGS) {
|
|
1179
|
+
var aVal: f32;
|
|
1180
|
+
if params.uplo == 0u {
|
|
1181
|
+
// Lower: A[i,j] stored at A[i*lda+j] for j \u2264 i, mirrored from A[j*lda+i] otherwise
|
|
1182
|
+
if j <= i {
|
|
1183
|
+
aVal = A[i * params.lda + j];
|
|
1184
|
+
} else {
|
|
1185
|
+
aVal = A[j * params.lda + i];
|
|
1186
|
+
}
|
|
1187
|
+
} else {
|
|
1188
|
+
// Upper: A[i,j] stored at A[i*lda+j] for j \u2265 i, mirrored from A[j*lda+i] otherwise
|
|
1189
|
+
if j >= i {
|
|
1190
|
+
aVal = A[i * params.lda + j];
|
|
1191
|
+
} else {
|
|
1192
|
+
aVal = A[j * params.lda + i];
|
|
1193
|
+
}
|
|
1194
|
+
}
|
|
1195
|
+
acc += aVal * x[j * params.incx];
|
|
1196
|
+
}
|
|
1536
1197
|
|
|
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; }
|
|
1198
|
+
// Parallel reduction: 64 \u2192 1
|
|
1199
|
+
scratch[lid.x] = acc;
|
|
1544
1200
|
workgroupBarrier();
|
|
1545
|
-
|
|
1201
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
1202
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
1203
|
+
workgroupBarrier();
|
|
1204
|
+
}
|
|
1546
1205
|
|
|
1547
|
-
|
|
1548
|
-
|
|
1549
|
-
|
|
1206
|
+
if lid.x == 0u {
|
|
1207
|
+
// BLAS beta==0 semantics: y is written, not accumulated \u2014 must not read y.
|
|
1208
|
+
let acc = params.alpha * scratch[0];
|
|
1209
|
+
y[i * params.incy] = select(acc, acc + params.beta * y[i * params.incy], params.beta != 0.0);
|
|
1210
|
+
}
|
|
1550
1211
|
}
|
|
1551
1212
|
}
|
|
1552
|
-
`});var
|
|
1553
|
-
//
|
|
1554
|
-
//
|
|
1555
|
-
//
|
|
1213
|
+
`});var Mt,Nt=V(()=>{Mt=`// strmv: y = op(A) * x
|
|
1214
|
+
// A is n\xD7n triangular, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
1215
|
+
// op(A) is A (trans=0) or A^T (trans=1).
|
|
1216
|
+
// diag=1 (unit) treats the diagonal as 1 without reading A's diagonal values.
|
|
1217
|
+
// One workgroup per row, grid-stride outer loop.
|
|
1556
1218
|
|
|
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;
|
|
1219
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1220
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
1221
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
1563
1222
|
|
|
1564
1223
|
struct Params {
|
|
1565
1224
|
n: u32,
|
|
1566
|
-
|
|
1225
|
+
incx: u32,
|
|
1226
|
+
incy: u32,
|
|
1227
|
+
lda: u32,
|
|
1228
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
1229
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
1230
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
1567
1231
|
}
|
|
1568
1232
|
|
|
1569
|
-
|
|
1233
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1570
1234
|
|
|
1571
|
-
|
|
1572
|
-
var<workgroup>
|
|
1235
|
+
const WGS: u32 = 64u;
|
|
1236
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
1573
1237
|
|
|
1574
1238
|
@compute @workgroup_size(64)
|
|
1575
|
-
fn
|
|
1576
|
-
@builtin(
|
|
1577
|
-
@builtin(local_invocation_id)
|
|
1578
|
-
@builtin(
|
|
1579
|
-
@builtin(num_workgroups) num_wg: vec3u,
|
|
1239
|
+
fn main(
|
|
1240
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1241
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1242
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
1580
1243
|
) {
|
|
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();
|
|
1244
|
+
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
1245
|
+
var acc = 0.0f;
|
|
1628
1246
|
|
|
1629
|
-
|
|
1630
|
-
|
|
1631
|
-
|
|
1632
|
-
|
|
1633
|
-
|
|
1634
|
-
|
|
1635
|
-
|
|
1636
|
-
|
|
1247
|
+
if params.trans == 0u {
|
|
1248
|
+
// No-transpose: y[i] = \u03A3_j A[i,j] * x[j]
|
|
1249
|
+
if params.uplo == 0u {
|
|
1250
|
+
// Lower: A[i,j] stored at A[i*lda+j] for j \u2264 i
|
|
1251
|
+
for (var j = lid.x; j <= i; j += WGS) {
|
|
1252
|
+
var aVal: f32;
|
|
1253
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
1254
|
+
if params.diag == 1u && j == i {
|
|
1255
|
+
aVal = 1.0;
|
|
1256
|
+
} else if ( j <= i ) {
|
|
1257
|
+
aVal = A[i * params.lda + j];
|
|
1258
|
+
}
|
|
1259
|
+
acc += aVal * x[j * params.incx];
|
|
1260
|
+
}
|
|
1261
|
+
} else {
|
|
1262
|
+
// Upper: A[i,j] stored at A[i*lda+j] for j \u2265 i
|
|
1263
|
+
for (var j = i + lid.x; j < params.n; j += WGS) {
|
|
1264
|
+
var aVal: f32;
|
|
1265
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
1266
|
+
if params.diag == 1u && j == i {
|
|
1267
|
+
aVal = 1.0;
|
|
1268
|
+
} else if ( j >= i ) {
|
|
1269
|
+
aVal = A[i * params.lda + j];
|
|
1270
|
+
}
|
|
1271
|
+
acc += aVal * x[j * params.incx];
|
|
1272
|
+
}
|
|
1273
|
+
}
|
|
1274
|
+
} else {
|
|
1275
|
+
// Transpose: y[i] = \u03A3_j A[j,i] * x[j]
|
|
1276
|
+
if params.uplo == 0u {
|
|
1277
|
+
// Lower: A[j,i] stored at A[j*lda+i] for j \u2265 i
|
|
1278
|
+
for (var j = i + lid.x; j < params.n; j += WGS) {
|
|
1279
|
+
var aVal: f32;
|
|
1280
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
1281
|
+
if params.diag == 1u && j == i {
|
|
1282
|
+
aVal = 1.0;
|
|
1283
|
+
} else if ( j >= i ) {
|
|
1284
|
+
aVal = A[j * params.lda + i];
|
|
1285
|
+
}
|
|
1286
|
+
acc += aVal * x[j * params.incx];
|
|
1287
|
+
}
|
|
1288
|
+
} else {
|
|
1289
|
+
// Upper: A[j,i] stored at A[j*lda+i] for j \u2264 i
|
|
1290
|
+
for (var j = lid.x; j <= i; j += WGS) {
|
|
1291
|
+
var aVal: f32;
|
|
1292
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
1293
|
+
if params.diag == 1u && j == i {
|
|
1294
|
+
aVal = 1.0;
|
|
1295
|
+
} else if ( j <= i ) {
|
|
1296
|
+
aVal = A[j * params.lda + i];
|
|
1297
|
+
}
|
|
1298
|
+
acc += aVal * x[j * params.incx];
|
|
1299
|
+
}
|
|
1637
1300
|
}
|
|
1638
1301
|
}
|
|
1302
|
+
|
|
1303
|
+
// Parallel reduction: 64 \u2192 1
|
|
1304
|
+
scratch[lid.x] = acc;
|
|
1639
1305
|
workgroupBarrier();
|
|
1640
|
-
|
|
1306
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
1307
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
1308
|
+
workgroupBarrier();
|
|
1309
|
+
}
|
|
1641
1310
|
|
|
1642
|
-
|
|
1643
|
-
|
|
1644
|
-
|
|
1645
|
-
partialsIdx[wgid.x] = tile_idx[0];
|
|
1311
|
+
if lid.x == 0u {
|
|
1312
|
+
y[ i * params.incy ] = scratch[0];
|
|
1313
|
+
}
|
|
1646
1314
|
}
|
|
1647
1315
|
}
|
|
1648
|
-
`});var
|
|
1316
|
+
`});var xe,It=V(()=>{xe=`// strsv_invert_block: computes ONE column (workgroup_id.x) of ONE block's
|
|
1649
1317
|
// (workgroup_id.y) explicit inverse, via the same one-row-at-a-time
|
|
1650
1318
|
// substitution as strsv_block.wgsl, but solving against a unit basis vector
|
|
1651
1319
|
// e_col instead of the real right-hand side, and writing to a dense
|
|
@@ -1754,7 +1422,7 @@ fn strsv_invert_block_main(
|
|
|
1754
1422
|
workgroupBarrier();
|
|
1755
1423
|
}
|
|
1756
1424
|
}
|
|
1757
|
-
`});var
|
|
1425
|
+
`});var Pt,Rt=V(()=>{Pt=`// strsv_apply_inverse: given a precomputed block inverse (from
|
|
1758
1426
|
// strsv_invert_block.wgsl), computes this block's solution as a dense
|
|
1759
1427
|
// matrix-vector multiply against the block's current remainder in x \u2014
|
|
1760
1428
|
// replacing what the old strsv_block.wgsl did via a genuinely sequential,
|
|
@@ -1801,7 +1469,7 @@ fn strsv_apply_inverse_main(@builtin(local_invocation_id) lid: vec3u) {
|
|
|
1801
1469
|
}
|
|
1802
1470
|
x[(params.blockStart + lid.x) * params.incx] = acc;
|
|
1803
1471
|
}
|
|
1804
|
-
`});var
|
|
1472
|
+
`});var Tt,Dt=V(()=>{Tt=`// strsv_update: subtracts a solved block's contribution from every
|
|
1805
1473
|
// remaining row in parallel (one workgroup per row, like strmv.wgsl) \u2014
|
|
1806
1474
|
// this is what turns strsv's O(n) sequential stages into O(n/blockSize).
|
|
1807
1475
|
// No diag/masking needed: this region never touches the diagonal.
|
|
@@ -1876,13 +1544,191 @@ fn strsv_update_main(
|
|
|
1876
1544
|
workgroupBarrier();
|
|
1877
1545
|
}
|
|
1878
1546
|
}
|
|
1879
|
-
`});var
|
|
1547
|
+
`});var Ct,jt=V(()=>{Ct=`// sger: A := alpha * x * y^T + A (rank-1 update, A is m\xD7n general/dense)
|
|
1548
|
+
|
|
1549
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
1550
|
+
@group(0) @binding(1) var<storage, read> y: array<f32>;
|
|
1551
|
+
@group(0) @binding(2) var<storage, read_write> A: array<f32>;
|
|
1552
|
+
|
|
1553
|
+
struct Params {
|
|
1554
|
+
m: u32,
|
|
1555
|
+
n: u32,
|
|
1556
|
+
alpha: f32,
|
|
1557
|
+
incx: u32,
|
|
1558
|
+
incy: u32,
|
|
1559
|
+
lda: u32,
|
|
1560
|
+
}
|
|
1561
|
+
|
|
1562
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1563
|
+
|
|
1564
|
+
const WGS: u32 = 64u;
|
|
1565
|
+
|
|
1566
|
+
@compute @workgroup_size(64)
|
|
1567
|
+
fn main(
|
|
1568
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1569
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1570
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
1571
|
+
) {
|
|
1572
|
+
for (var row = wgid.x; row < params.m; row += nwg.x) {
|
|
1573
|
+
let xi = params.alpha * x[row * params.incx];
|
|
1574
|
+
let row_base = row * params.lda;
|
|
1575
|
+
|
|
1576
|
+
// 4-unrolled loop: each iteration issues 4 independent A/y accesses.
|
|
1577
|
+
let n4_floor = (params.n / (4u * WGS)) * (4u * WGS);
|
|
1578
|
+
for (var col: u32 = lid.x; col < n4_floor; col += 4u * WGS) {
|
|
1579
|
+
let idx0 = row_base + col;
|
|
1580
|
+
let idx1 = row_base + col + WGS;
|
|
1581
|
+
let idx2 = row_base + col + 2u * WGS;
|
|
1582
|
+
let idx3 = row_base + col + 3u * WGS;
|
|
1583
|
+
A[idx0] = xi * y[ col * params.incy] + A[idx0];
|
|
1584
|
+
A[idx1] = xi * y[(col + WGS) * params.incy] + A[idx1];
|
|
1585
|
+
A[idx2] = xi * y[(col + 2u * WGS) * params.incy] + A[idx2];
|
|
1586
|
+
A[idx3] = xi * y[(col + 3u * WGS) * params.incy] + A[idx3];
|
|
1587
|
+
}
|
|
1588
|
+
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
1589
|
+
for (var col: u32 = n4_floor + lid.x; col < params.n; col += WGS) {
|
|
1590
|
+
let idx = row_base + col;
|
|
1591
|
+
A[idx] = xi * y[col * params.incy] + A[idx];
|
|
1592
|
+
}
|
|
1593
|
+
}
|
|
1594
|
+
}
|
|
1595
|
+
`});var Wt,Lt=V(()=>{Wt=`// ssyr: A := alpha * x * x^T + A (symmetric rank-1 update)
|
|
1596
|
+
// A is n\xD7n symmetric; only the triangle specified by uplo is referenced/updated,
|
|
1597
|
+
// the other triangle is implied by symmetry (not touched).
|
|
1598
|
+
|
|
1599
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
1600
|
+
@group(0) @binding(1) var<storage, read_write> A: array<f32>;
|
|
1601
|
+
|
|
1602
|
+
struct Params {
|
|
1603
|
+
n: u32,
|
|
1604
|
+
alpha: f32,
|
|
1605
|
+
incx: u32,
|
|
1606
|
+
lda: u32,
|
|
1607
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
1608
|
+
}
|
|
1609
|
+
|
|
1610
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
1611
|
+
|
|
1612
|
+
const WGS: u32 = 64u;
|
|
1613
|
+
|
|
1614
|
+
@compute @workgroup_size(64)
|
|
1615
|
+
fn main(
|
|
1616
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1617
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1618
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
1619
|
+
) {
|
|
1620
|
+
for (var row = wgid.x; row < params.n; row += nwg.x) {
|
|
1621
|
+
let xi = params.alpha * x[row * params.incx];
|
|
1622
|
+
let row_base = row * params.lda;
|
|
1623
|
+
|
|
1624
|
+
// Stored-triangle column range for this row: lower [0,row], upper [row,n).
|
|
1625
|
+
var colStart: u32;
|
|
1626
|
+
var colEnd: u32;
|
|
1627
|
+
if params.uplo == 1u {
|
|
1628
|
+
colStart = row;
|
|
1629
|
+
colEnd = params.n;
|
|
1630
|
+
} else {
|
|
1631
|
+
colStart = 0u;
|
|
1632
|
+
colEnd = row + 1u;
|
|
1633
|
+
}
|
|
1634
|
+
|
|
1635
|
+
// 4-unrolled loop over the stored range.
|
|
1636
|
+
let rangeLen = colEnd - colStart;
|
|
1637
|
+
let n4_floor = colStart + (rangeLen / (4u * WGS)) * (4u * WGS);
|
|
1638
|
+
for (var col: u32 = colStart + lid.x; col < n4_floor; col += 4u * WGS) {
|
|
1639
|
+
let idx0 = row_base + col;
|
|
1640
|
+
let idx1 = row_base + col + WGS;
|
|
1641
|
+
let idx2 = row_base + col + 2u * WGS;
|
|
1642
|
+
let idx3 = row_base + col + 3u * WGS;
|
|
1643
|
+
A[idx0] = xi * x[ col * params.incx] + A[idx0];
|
|
1644
|
+
A[idx1] = xi * x[(col + WGS) * params.incx] + A[idx1];
|
|
1645
|
+
A[idx2] = xi * x[(col + 2u * WGS) * params.incx] + A[idx2];
|
|
1646
|
+
A[idx3] = xi * x[(col + 3u * WGS) * params.incx] + A[idx3];
|
|
1647
|
+
}
|
|
1648
|
+
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
1649
|
+
for (var col: u32 = n4_floor + lid.x; col < colEnd; col += WGS) {
|
|
1650
|
+
let idx = row_base + col;
|
|
1651
|
+
A[idx] = xi * x[col * params.incx] + A[idx];
|
|
1652
|
+
}
|
|
1653
|
+
}
|
|
1654
|
+
}
|
|
1655
|
+
`});var qt,Ft=V(()=>{qt=`// ssyr2: A := alpha * x * y^T + alpha * y * x^T + A (symmetric rank-2 update)
|
|
1656
|
+
// A is n\xD7n symmetric; only the triangle specified by uplo is referenced/updated,
|
|
1657
|
+
// the other triangle is implied by symmetry (not touched).
|
|
1658
|
+
|
|
1659
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
1660
|
+
@group(0) @binding(1) var<storage, read> y: array<f32>;
|
|
1661
|
+
@group(0) @binding(2) var<storage, read_write> A: array<f32>;
|
|
1662
|
+
|
|
1663
|
+
struct Params {
|
|
1664
|
+
n: u32,
|
|
1665
|
+
alpha: f32,
|
|
1666
|
+
incx: u32,
|
|
1667
|
+
incy: u32,
|
|
1668
|
+
lda: u32,
|
|
1669
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
1670
|
+
}
|
|
1671
|
+
|
|
1672
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1673
|
+
|
|
1674
|
+
const WGS: u32 = 64u;
|
|
1675
|
+
|
|
1676
|
+
@compute @workgroup_size(64)
|
|
1677
|
+
fn main(
|
|
1678
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1679
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1680
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
1681
|
+
) {
|
|
1682
|
+
for (var row = wgid.x; row < params.n; row += nwg.x) {
|
|
1683
|
+
let xi = params.alpha * x[row * params.incx];
|
|
1684
|
+
let yi = params.alpha * y[row * params.incy];
|
|
1685
|
+
let row_base = row * params.lda;
|
|
1686
|
+
|
|
1687
|
+
// Stored-triangle column range for this row: lower [0,row], upper [row,n).
|
|
1688
|
+
var colStart: u32;
|
|
1689
|
+
var colEnd: u32;
|
|
1690
|
+
if params.uplo == 1u {
|
|
1691
|
+
colStart = row;
|
|
1692
|
+
colEnd = params.n;
|
|
1693
|
+
} else {
|
|
1694
|
+
colStart = 0u;
|
|
1695
|
+
colEnd = row + 1u;
|
|
1696
|
+
}
|
|
1697
|
+
|
|
1698
|
+
// 4-unrolled loop over the stored range.
|
|
1699
|
+
let rangeLen = colEnd - colStart;
|
|
1700
|
+
let n4_floor = colStart + (rangeLen / (4u * WGS)) * (4u * WGS);
|
|
1701
|
+
for (var col: u32 = colStart + lid.x; col < n4_floor; col += 4u * WGS) {
|
|
1702
|
+
let idx0 = row_base + col;
|
|
1703
|
+
let idx1 = row_base + col + WGS;
|
|
1704
|
+
let idx2 = row_base + col + 2u * WGS;
|
|
1705
|
+
let idx3 = row_base + col + 3u * WGS;
|
|
1706
|
+
A[idx0] = xi * y[ col * params.incy] + yi * x[ col * params.incx] + A[idx0];
|
|
1707
|
+
A[idx1] = xi * y[(col + WGS) * params.incy] + yi * x[(col + WGS) * params.incx] + A[idx1];
|
|
1708
|
+
A[idx2] = xi * y[(col + 2u * WGS) * params.incy] + yi * x[(col + 2u * WGS) * params.incx] + A[idx2];
|
|
1709
|
+
A[idx3] = xi * y[(col + 3u * WGS) * params.incy] + yi * x[(col + 3u * WGS) * params.incx] + A[idx3];
|
|
1710
|
+
}
|
|
1711
|
+
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
1712
|
+
for (var col: u32 = n4_floor + lid.x; col < colEnd; col += WGS) {
|
|
1713
|
+
let idx = row_base + col;
|
|
1714
|
+
A[idx] = xi * y[col * params.incy] + yi * x[col * params.incx] + A[idx];
|
|
1715
|
+
}
|
|
1716
|
+
}
|
|
1717
|
+
}
|
|
1718
|
+
`});var re,Ut=V(()=>{re=`// sgemm_small: C = alpha * op(A) * op(B) + beta * C \u2014 small-tile half of
|
|
1880
1719
|
// the two-tier autotuned dispatch (see sgemm.mjs and sgemm_large.wgsl).
|
|
1881
1720
|
// BM=BN=32, BK=8, TM=TN=2 \u2014 wins over the large tile below a 6x6=36
|
|
1882
1721
|
// workgroup grid of 64-tiles, where the large tile doesn't have enough
|
|
1883
1722
|
// workgroups to fill the GPU. Same structure as sgemm_large.wgsl (2D
|
|
1884
1723
|
// register-blocked, shared-memory-tiled), just smaller.
|
|
1885
1724
|
//
|
|
1725
|
+
// A and B are bound twice \u2014 scalar array<f32> and array<vec4<f32>> views of
|
|
1726
|
+
// the same GPUBuffer (see vec4ViewBinding) \u2014 so each tile load can issue
|
|
1727
|
+
// 16-byte vector reads along op(A)/op(B)'s contiguous dimension when the
|
|
1728
|
+
// stride allows it (stride % 4 == 0 keeps every row base 16-byte aligned).
|
|
1729
|
+
// NUM_THREADS (256) exceeds some small-tile load shapes, so the vectorized
|
|
1730
|
+
// paths whose lane count doesn't tile exactly guard their As/Bs stores.
|
|
1731
|
+
//
|
|
1886
1732
|
// col mapped to gid.x for coalesced B/C access (row-major: col contiguous).
|
|
1887
1733
|
|
|
1888
1734
|
const BM: u32 = 32u;
|
|
@@ -1896,9 +1742,11 @@ const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 256
|
|
|
1896
1742
|
const STRIDE_A: u32 = NUM_THREADS / BK;
|
|
1897
1743
|
const STRIDE_B: u32 = NUM_THREADS / BN;
|
|
1898
1744
|
|
|
1899
|
-
@group(0) @binding(0) var<storage, read> A:
|
|
1900
|
-
@group(0) @binding(1) var<storage, read>
|
|
1901
|
-
@group(0) @binding(2) var<storage,
|
|
1745
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1746
|
+
@group(0) @binding(1) var<storage, read> A4: array<vec4<f32>>;
|
|
1747
|
+
@group(0) @binding(2) var<storage, read> B: array<f32>;
|
|
1748
|
+
@group(0) @binding(3) var<storage, read> B4: array<vec4<f32>>;
|
|
1749
|
+
@group(0) @binding(4) var<storage, read_write> C: array<f32>;
|
|
1902
1750
|
|
|
1903
1751
|
struct Params {
|
|
1904
1752
|
m: u32,
|
|
@@ -1911,9 +1759,11 @@ struct Params {
|
|
|
1911
1759
|
ldc: u32,
|
|
1912
1760
|
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
1913
1761
|
transB: u32,
|
|
1762
|
+
useVecA: u32, // 1 = A's vec4 view covers every in-bounds element (see vec4Usable)
|
|
1763
|
+
useVecB: u32, // 1 = B's vec4 view covers every in-bounds element
|
|
1914
1764
|
}
|
|
1915
1765
|
|
|
1916
|
-
@group(0) @binding(
|
|
1766
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
1917
1767
|
|
|
1918
1768
|
var<workgroup> As: array<f32, BM * BK>;
|
|
1919
1769
|
var<workgroup> Bs: array<f32, BK * BN>;
|
|
@@ -1943,17 +1793,103 @@ fn main(
|
|
|
1943
1793
|
|
|
1944
1794
|
let numTiles = (params.k + BK - 1u) / BK;
|
|
1945
1795
|
for (var t = 0u; t < numTiles; t++) {
|
|
1946
|
-
|
|
1947
|
-
|
|
1948
|
-
|
|
1949
|
-
|
|
1950
|
-
|
|
1796
|
+
// \u2500\u2500 Load the BM\xD7BK A tile into As (vectorized along op(A)'s fast dim
|
|
1797
|
+
// when lda allows; every branch here is dispatch-uniform) \u2500\u2500
|
|
1798
|
+
if (params.useVecA == 1u && params.transA == 0u) {
|
|
1799
|
+
// No-transpose: columns contiguous, one vec4 per thread. NUM_THREADS
|
|
1800
|
+
// spans BM/4\xD7(BK/4) several times over \u2014 guard the store.
|
|
1801
|
+
let r4 = tid / (BK / 4u);
|
|
1802
|
+
let c4 = tid % (BK / 4u);
|
|
1803
|
+
if (r4 < BM) {
|
|
1804
|
+
let gRow = blockRow + r4;
|
|
1805
|
+
let gCol = t * BK + c4 * 4u;
|
|
1806
|
+
var v = A4[(gRow * params.lda + gCol) / 4u];
|
|
1807
|
+
let rowOK = gRow < params.m;
|
|
1808
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.k);
|
|
1809
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.k);
|
|
1810
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.k);
|
|
1811
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.k);
|
|
1812
|
+
As[r4 * BK + c4 * 4u] = v.x;
|
|
1813
|
+
As[r4 * BK + c4 * 4u + 1u] = v.y;
|
|
1814
|
+
As[r4 * BK + c4 * 4u + 2u] = v.z;
|
|
1815
|
+
As[r4 * BK + c4 * 4u + 3u] = v.w;
|
|
1816
|
+
}
|
|
1817
|
+
} else if (params.useVecA == 1u && params.transA != 0u) {
|
|
1818
|
+
// Transpose: rows contiguous within a column. NUM_THREADS over-spans
|
|
1819
|
+
// the BK-column tile \u2014 guard the store.
|
|
1820
|
+
let r4 = tid % (BM / 4u);
|
|
1821
|
+
let c = tid / (BM / 4u);
|
|
1822
|
+
if (c < BK) {
|
|
1823
|
+
let gRow = blockRow + r4 * 4u;
|
|
1824
|
+
let gCol = t * BK + c;
|
|
1825
|
+
var v = A4[(gCol * params.lda + gRow) / 4u];
|
|
1826
|
+
let colOK = gCol < params.k;
|
|
1827
|
+
v.x = select(0.0, v.x, colOK && gRow < params.m);
|
|
1828
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.m);
|
|
1829
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.m);
|
|
1830
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.m);
|
|
1831
|
+
As[(r4 * 4u) * BK + c] = v.x;
|
|
1832
|
+
As[(r4 * 4u + 1u) * BK + c] = v.y;
|
|
1833
|
+
As[(r4 * 4u + 2u) * BK + c] = v.z;
|
|
1834
|
+
As[(r4 * 4u + 3u) * BK + c] = v.w;
|
|
1835
|
+
}
|
|
1836
|
+
} else {
|
|
1837
|
+
// Scalar fallback: odd stride or unhandled orientation.
|
|
1838
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
1839
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
1840
|
+
let gColA = t * BK + innerColA;
|
|
1841
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
1842
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
1843
|
+
}
|
|
1951
1844
|
}
|
|
1952
|
-
|
|
1953
|
-
|
|
1954
|
-
|
|
1955
|
-
|
|
1956
|
-
|
|
1845
|
+
|
|
1846
|
+
// \u2500\u2500 Load the BK\xD7BN B tile into Bs \u2500\u2500
|
|
1847
|
+
if (params.useVecB == 1u && params.transB == 0u) {
|
|
1848
|
+
// No-transpose: columns contiguous, one vec4 per thread. NUM_THREADS
|
|
1849
|
+
// over-spans the BK-row tile \u2014 guard the store.
|
|
1850
|
+
let r = tid / (BN / 4u);
|
|
1851
|
+
let c4 = tid % (BN / 4u);
|
|
1852
|
+
if (r < BK) {
|
|
1853
|
+
let gRow = t * BK + r;
|
|
1854
|
+
let gCol = blockCol + c4 * 4u;
|
|
1855
|
+
var v = B4[(gRow * params.ldb + gCol) / 4u];
|
|
1856
|
+
let rowOK = gRow < params.k;
|
|
1857
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.n);
|
|
1858
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.n);
|
|
1859
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.n);
|
|
1860
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.n);
|
|
1861
|
+
Bs[r * BN + c4 * 4u] = v.x;
|
|
1862
|
+
Bs[r * BN + c4 * 4u + 1u] = v.y;
|
|
1863
|
+
Bs[r * BN + c4 * 4u + 2u] = v.z;
|
|
1864
|
+
Bs[r * BN + c4 * 4u + 3u] = v.w;
|
|
1865
|
+
}
|
|
1866
|
+
} else if (params.useVecB == 1u && params.transB != 0u) {
|
|
1867
|
+
// Transpose: rows contiguous within a column, one vec4 per thread \u2014
|
|
1868
|
+
// NUM_THREADS over-spans the 32-column tile, so guard the store.
|
|
1869
|
+
let r4 = tid % (BK / 4u);
|
|
1870
|
+
let c = tid / (BK / 4u);
|
|
1871
|
+
if (c < BN) {
|
|
1872
|
+
let gRow = t * BK + r4 * 4u;
|
|
1873
|
+
let gCol = blockCol + c;
|
|
1874
|
+
var v = B4[(gCol * params.ldb + gRow) / 4u];
|
|
1875
|
+
let colOK = gCol < params.n;
|
|
1876
|
+
v.x = select(0.0, v.x, colOK && gRow < params.k);
|
|
1877
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.k);
|
|
1878
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.k);
|
|
1879
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.k);
|
|
1880
|
+
Bs[(r4 * 4u) * BN + c] = v.x;
|
|
1881
|
+
Bs[(r4 * 4u + 1u) * BN + c] = v.y;
|
|
1882
|
+
Bs[(r4 * 4u + 2u) * BN + c] = v.z;
|
|
1883
|
+
Bs[(r4 * 4u + 3u) * BN + c] = v.w;
|
|
1884
|
+
}
|
|
1885
|
+
} else {
|
|
1886
|
+
// Scalar fallback.
|
|
1887
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
1888
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
1889
|
+
let gColB = blockCol + innerColB;
|
|
1890
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
1891
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
1892
|
+
}
|
|
1957
1893
|
}
|
|
1958
1894
|
|
|
1959
1895
|
workgroupBarrier();
|
|
@@ -1982,13 +1918,16 @@ fn main(
|
|
|
1982
1918
|
let col = blockCol + threadCol * TN + resIdxN;
|
|
1983
1919
|
if (col < params.n) {
|
|
1984
1920
|
let cIdx = row * params.ldc + col;
|
|
1985
|
-
|
|
1921
|
+
// BLAS beta==0 semantics: C is written, not accumulated \u2014 must not
|
|
1922
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
1923
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
1924
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
1986
1925
|
}
|
|
1987
1926
|
}
|
|
1988
1927
|
}
|
|
1989
1928
|
}
|
|
1990
1929
|
}
|
|
1991
|
-
`});var
|
|
1930
|
+
`});var ee,Ot=V(()=>{ee=`// sgemm_large: C = alpha * op(A) * op(B) + beta * C \u2014 large-tile half of
|
|
1992
1931
|
// the two-tier autotuned dispatch (see sgemm.mjs and sgemm_small.wgsl).
|
|
1993
1932
|
// BM=BN=64, BK=8, TM=8, TN=4 (128 threads/workgroup) \u2014 the kernel 9
|
|
1994
1933
|
// autotuning winner (temp/autotune_sweep.mjs, temp/gen_sweep_kernel.mjs,
|
|
@@ -1996,9 +1935,13 @@ fn main(
|
|
|
1996
1935
|
// single-tier baseline at n=512, +84% at n=1024. But BM=64 loses to BM=32
|
|
1997
1936
|
// below a 6x6=36 workgroup grid (not enough workgroups to fill the GPU at
|
|
1998
1937
|
// that tile size), hence the two-tier split rather than one global config.
|
|
1999
|
-
//
|
|
2000
|
-
//
|
|
2001
|
-
//
|
|
1938
|
+
//
|
|
1939
|
+
// A and B are bound twice \u2014 scalar array<f32> and array<vec4<f32>> views of
|
|
1940
|
+
// the same GPUBuffer (see vec4ViewBinding) \u2014 so each tile load can issue
|
|
1941
|
+
// 16-byte vector reads along op(A)/op(B)'s contiguous dimension when the
|
|
1942
|
+
// stride allows it (stride % 4 == 0 keeps every row base 16-byte aligned).
|
|
1943
|
+
// Transposed or odd-stride operands take the scalar path; both paths
|
|
1944
|
+
// zero-fill out-of-bounds components identically.
|
|
2002
1945
|
|
|
2003
1946
|
const BM: u32 = 64u;
|
|
2004
1947
|
const BN: u32 = 64u;
|
|
@@ -2011,9 +1954,11 @@ const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 128
|
|
|
2011
1954
|
const STRIDE_A: u32 = NUM_THREADS / BK; // rows of As covered per load-loop step
|
|
2012
1955
|
const STRIDE_B: u32 = NUM_THREADS / BN; // rows of Bs covered per load-loop step
|
|
2013
1956
|
|
|
2014
|
-
@group(0) @binding(0) var<storage, read> A:
|
|
2015
|
-
@group(0) @binding(1) var<storage, read>
|
|
2016
|
-
@group(0) @binding(2) var<storage,
|
|
1957
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1958
|
+
@group(0) @binding(1) var<storage, read> A4: array<vec4<f32>>;
|
|
1959
|
+
@group(0) @binding(2) var<storage, read> B: array<f32>;
|
|
1960
|
+
@group(0) @binding(3) var<storage, read> B4: array<vec4<f32>>;
|
|
1961
|
+
@group(0) @binding(4) var<storage, read_write> C: array<f32>;
|
|
2017
1962
|
|
|
2018
1963
|
struct Params {
|
|
2019
1964
|
m: u32,
|
|
@@ -2026,9 +1971,11 @@ struct Params {
|
|
|
2026
1971
|
ldc: u32,
|
|
2027
1972
|
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
2028
1973
|
transB: u32,
|
|
1974
|
+
useVecA: u32, // 1 = A's vec4 view covers every in-bounds element (see vec4Usable)
|
|
1975
|
+
useVecB: u32, // 1 = B's vec4 view covers every in-bounds element
|
|
2029
1976
|
}
|
|
2030
1977
|
|
|
2031
|
-
@group(0) @binding(
|
|
1978
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
2032
1979
|
|
|
2033
1980
|
var<workgroup> As: array<f32, BM * BK>;
|
|
2034
1981
|
var<workgroup> Bs: array<f32, BK * BN>;
|
|
@@ -2060,17 +2007,95 @@ fn main(
|
|
|
2060
2007
|
|
|
2061
2008
|
let numTiles = (params.k + BK - 1u) / BK;
|
|
2062
2009
|
for (var t = 0u; t < numTiles; t++) {
|
|
2063
|
-
|
|
2064
|
-
|
|
2065
|
-
|
|
2066
|
-
|
|
2067
|
-
|
|
2010
|
+
// \u2500\u2500 Load the BM\xD7BK A tile into As (vectorized along op(A)'s fast dim
|
|
2011
|
+
// when lda allows; every branch here is dispatch-uniform) \u2500\u2500
|
|
2012
|
+
if (params.useVecA == 1u && params.transA == 0u) {
|
|
2013
|
+
// No-transpose: columns contiguous. Each thread loads one vec4 of 4
|
|
2014
|
+
// columns; 64 rows \xD7 2 column-lanes = NUM_THREADS exactly, single pass.
|
|
2015
|
+
let r4 = tid / (BK / 4u);
|
|
2016
|
+
let c4 = tid % (BK / 4u);
|
|
2017
|
+
let gRow = blockRow + r4;
|
|
2018
|
+
let gCol = t * BK + c4 * 4u;
|
|
2019
|
+
var v = A4[(gRow * params.lda + gCol) / 4u];
|
|
2020
|
+
let rowOK = gRow < params.m;
|
|
2021
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.k);
|
|
2022
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.k);
|
|
2023
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.k);
|
|
2024
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.k);
|
|
2025
|
+
As[r4 * BK + c4 * 4u] = v.x;
|
|
2026
|
+
As[r4 * BK + c4 * 4u + 1u] = v.y;
|
|
2027
|
+
As[r4 * BK + c4 * 4u + 2u] = v.z;
|
|
2028
|
+
As[r4 * BK + c4 * 4u + 3u] = v.w;
|
|
2029
|
+
} else if (params.useVecA == 1u && params.transA != 0u) {
|
|
2030
|
+
// Transpose: rows contiguous within a column. Each thread loads one
|
|
2031
|
+
// vec4 of 4 rows; 16 row-lanes \xD7 8 columns = NUM_THREADS, single pass.
|
|
2032
|
+
let r4 = tid % (BM / 4u);
|
|
2033
|
+
let c = tid / (BM / 4u);
|
|
2034
|
+
let gRow = blockRow + r4 * 4u;
|
|
2035
|
+
let gCol = t * BK + c;
|
|
2036
|
+
var v = A4[(gCol * params.lda + gRow) / 4u];
|
|
2037
|
+
let colOK = gCol < params.k;
|
|
2038
|
+
v.x = select(0.0, v.x, colOK && gRow < params.m);
|
|
2039
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.m);
|
|
2040
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.m);
|
|
2041
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.m);
|
|
2042
|
+
As[(r4 * 4u) * BK + c] = v.x;
|
|
2043
|
+
As[(r4 * 4u + 1u) * BK + c] = v.y;
|
|
2044
|
+
As[(r4 * 4u + 2u) * BK + c] = v.z;
|
|
2045
|
+
As[(r4 * 4u + 3u) * BK + c] = v.w;
|
|
2046
|
+
} else {
|
|
2047
|
+
// Scalar fallback: odd stride or unhandled orientation.
|
|
2048
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
2049
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
2050
|
+
let gColA = t * BK + innerColA;
|
|
2051
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
2052
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
2053
|
+
}
|
|
2068
2054
|
}
|
|
2069
|
-
|
|
2070
|
-
|
|
2071
|
-
|
|
2072
|
-
|
|
2073
|
-
|
|
2055
|
+
|
|
2056
|
+
// \u2500\u2500 Load the BK\xD7BN B tile into Bs \u2500\u2500
|
|
2057
|
+
if (params.useVecB == 1u && params.transB == 0u) {
|
|
2058
|
+
// No-transpose: columns contiguous. 8 rows \xD7 16 column-lanes cover the
|
|
2059
|
+
// tile in one pass (BK = NUM_THREADS / (BN/4)).
|
|
2060
|
+
let r = tid / (BN / 4u);
|
|
2061
|
+
let c4 = tid % (BN / 4u);
|
|
2062
|
+
let gRow = t * BK + r;
|
|
2063
|
+
let gCol = blockCol + c4 * 4u;
|
|
2064
|
+
var v = B4[(gRow * params.ldb + gCol) / 4u];
|
|
2065
|
+
let rowOK = gRow < params.k;
|
|
2066
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.n);
|
|
2067
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.n);
|
|
2068
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.n);
|
|
2069
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.n);
|
|
2070
|
+
Bs[r * BN + c4 * 4u] = v.x;
|
|
2071
|
+
Bs[r * BN + c4 * 4u + 1u] = v.y;
|
|
2072
|
+
Bs[r * BN + c4 * 4u + 2u] = v.z;
|
|
2073
|
+
Bs[r * BN + c4 * 4u + 3u] = v.w;
|
|
2074
|
+
} else if (params.useVecB == 1u && params.transB != 0u) {
|
|
2075
|
+
// Transpose: rows contiguous within a column. 2 row-lanes \xD7 64 columns
|
|
2076
|
+
// cover the tile in one pass (BN = NUM_THREADS / (BK/4)).
|
|
2077
|
+
let r4 = tid % (BK / 4u);
|
|
2078
|
+
let c = tid / (BK / 4u);
|
|
2079
|
+
let gRow = t * BK + r4 * 4u;
|
|
2080
|
+
let gCol = blockCol + c;
|
|
2081
|
+
var v = B4[(gCol * params.ldb + gRow) / 4u];
|
|
2082
|
+
let colOK = gCol < params.n;
|
|
2083
|
+
v.x = select(0.0, v.x, colOK && gRow < params.k);
|
|
2084
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.k);
|
|
2085
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.k);
|
|
2086
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.k);
|
|
2087
|
+
Bs[(r4 * 4u) * BN + c] = v.x;
|
|
2088
|
+
Bs[(r4 * 4u + 1u) * BN + c] = v.y;
|
|
2089
|
+
Bs[(r4 * 4u + 2u) * BN + c] = v.z;
|
|
2090
|
+
Bs[(r4 * 4u + 3u) * BN + c] = v.w;
|
|
2091
|
+
} else {
|
|
2092
|
+
// Scalar fallback.
|
|
2093
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
2094
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
2095
|
+
let gColB = blockCol + innerColB;
|
|
2096
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
2097
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
2098
|
+
}
|
|
2074
2099
|
}
|
|
2075
2100
|
|
|
2076
2101
|
workgroupBarrier();
|
|
@@ -2099,13 +2124,16 @@ fn main(
|
|
|
2099
2124
|
let col = blockCol + threadCol * TN + resIdxN;
|
|
2100
2125
|
if (col < params.n) {
|
|
2101
2126
|
let cIdx = row * params.ldc + col;
|
|
2102
|
-
|
|
2127
|
+
// BLAS beta==0 semantics: C is written, not accumulated \u2014 must not
|
|
2128
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
2129
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
2130
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
2103
2131
|
}
|
|
2104
2132
|
}
|
|
2105
2133
|
}
|
|
2106
2134
|
}
|
|
2107
2135
|
}
|
|
2108
|
-
`});var
|
|
2136
|
+
`});var ue,Kt=V(()=>{ue=`// sgemmtr_small: C := uplo(alpha * op(A) * op(B) + beta * C) \u2014 small-tile
|
|
2109
2137
|
// half of a two-tier dispatch, identical to sgemm_small.wgsl except the
|
|
2110
2138
|
// final output write is gated to one triangle of C by \`uplo\` \u2014 see
|
|
2111
2139
|
// sgemmtr_large.wgsl for the full rationale (shared by both tiers).
|
|
@@ -2209,13 +2237,16 @@ fn main(
|
|
|
2209
2237
|
let inTriangle = select(col >= row, col <= row, params.uplo == 0u);
|
|
2210
2238
|
if (col < params.n && inTriangle) {
|
|
2211
2239
|
let cIdx = row * params.ldc + col;
|
|
2212
|
-
|
|
2240
|
+
// BLAS beta==0 semantics: C is written, not accumulated \u2014 must not
|
|
2241
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
2242
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
2243
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
2213
2244
|
}
|
|
2214
2245
|
}
|
|
2215
2246
|
}
|
|
2216
2247
|
}
|
|
2217
2248
|
}
|
|
2218
|
-
`});var
|
|
2249
|
+
`});var le,Vt=V(()=>{le=`// sgemmtr_large: C := uplo(alpha * op(A) * op(B) + beta * C) \u2014 large-tile
|
|
2219
2250
|
// half of a two-tier dispatch, identical to sgemm_large.wgsl (see that file
|
|
2220
2251
|
// for the BM/BN/BK/TM/TN autotuning rationale) except the final output write
|
|
2221
2252
|
// is gated to one triangle of C by \`uplo\`, the same convention ssyr/ssyr2
|
|
@@ -2326,13 +2357,16 @@ fn main(
|
|
|
2326
2357
|
let inTriangle = select(col >= row, col <= row, params.uplo == 0u);
|
|
2327
2358
|
if (col < params.n && inTriangle) {
|
|
2328
2359
|
let cIdx = row * params.ldc + col;
|
|
2329
|
-
|
|
2360
|
+
// BLAS beta==0 semantics: C is written, not accumulated \u2014 must not
|
|
2361
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
2362
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
2363
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
2330
2364
|
}
|
|
2331
2365
|
}
|
|
2332
2366
|
}
|
|
2333
2367
|
}
|
|
2334
2368
|
}
|
|
2335
|
-
`});var
|
|
2369
|
+
`});var Ht,zt=V(()=>{Ht=`// symmetrize: Adense := full dense expansion of a symmetric matrix stored
|
|
2336
2370
|
// with only its \`uplo\` triangle meaningful (the other triangle is implied
|
|
2337
2371
|
// by symmetry: A[i,j] = A[j,i]). A plain element-wise pass, no tiling or
|
|
2338
2372
|
// shared memory needed \u2014 used to materialize a dense operand for routines
|
|
@@ -2363,7 +2397,7 @@ fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
|
2363
2397
|
let srcIdx = select(col * params.lda + row, row * params.lda + col, isStored);
|
|
2364
2398
|
Adense[row * params.ldd + col] = A[srcIdx];
|
|
2365
2399
|
}
|
|
2366
|
-
`});var
|
|
2400
|
+
`});var Xt,Yt=V(()=>{Xt=`// triangularize: Adense := dense expansion of op(A) (A or A^T per \`trans\`),
|
|
2367
2401
|
// zero-filling the unstored triangle (exact for a matmul) so strmm can reuse
|
|
2368
2402
|
// sgemm's kernel unchanged. \`diag=1\` substitutes 1.0 on the diagonal.
|
|
2369
2403
|
|
|
@@ -2407,7 +2441,7 @@ fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
|
2407
2441
|
|
|
2408
2442
|
Adense[row * params.ldd + col] = select(0.0, A[srcRow * params.lda + srcCol], isMeaningful);
|
|
2409
2443
|
}
|
|
2410
|
-
`});var
|
|
2444
|
+
`});var Zt,$t=V(()=>{Zt=`// block_transfer: gather/scatter/scatter-subtract between a tight (blockLen
|
|
2411
2445
|
// x otherLen) block and a sub-range of a strided (any ld, row/col-major)
|
|
2412
2446
|
// buffer \u2014 needed since block offsets aren't 256-byte-aligned and block
|
|
2413
2447
|
// rows/cols aren't always one contiguous range for copyBufferToBuffer.
|
|
@@ -2449,7 +2483,7 @@ fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
|
2449
2483
|
strided[stridedIdx] = block[blockIdx];
|
|
2450
2484
|
}
|
|
2451
2485
|
}
|
|
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
|
-
${l.map(
|
|
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);})();
|
|
2486
|
+
`});var Qt={};Ee(Qt,{routineShaders:()=>nr,shaderSources:()=>Ia});var nr,Ia,Jt=V(()=>{Ce();Le();Fe();Ue();Ke();ze();He();Xe();Ze();Je();et();ot();at();st();it();lt();ft();pt();wt();bt();yt();vt();Bt();At();St();Nt();It();Rt();Dt();jt();Lt();Ft();Ut();Ot();Kt();Vt();zt();Yt();$t();nr={};nr.sscal={sscal:ge};nr.sswap={sswap:We};nr.saxpy={saxpy:qe};nr.scopy={scopy:Oe};nr.sdot={sdot:Ve,"reduction/sum":be};nr.sasum={sasum:Ye,"reduction/sum":be};nr.snrm2={snrm2:$e,"reduction/scaledSum":Qe};nr.isamax={isamax:rt,"reduction/argmax":tt};nr.dasum={"f64/dekker":he,"f64/utils/abs":ye,"f64/utils/add":nt,dasum:ut,"reduction/sumF64":mt};nr.idamax={"f64/dekker":he,"f64/utils/abs":ye,"f64/utils/greater":ct,"f64/utils/equal":dt,idamax:gt,"reduction/argmaxF64":ht};nr.srot={srot:xt};nr.srotm={srotm:_t};nr.sgemv={sgemv_n:Et,sgemv_t:Gt};nr.ssymv={ssymv:kt};nr.strmv={strmv:Mt};nr.strsv={strsv_invert_block:xe,strsv_apply_inverse:Pt,strsv_update:Tt};nr.sger={sger:Ct};nr.ssyr={ssyr:Wt};nr.ssyr2={ssyr2:qt};nr.sgemm={sgemm_small:re,sgemm_large:ee};nr.sgemmtr={sgemmtr_small:ue,sgemmtr_large:le};nr.ssyrk={sgemmtr_small:ue,sgemmtr_large:le};nr.ssyr2k={sgemmtr_small:ue,sgemmtr_large:le};nr.ssymm={sgemm_small:re,sgemm_large:ee,symmetrize:Ht};nr.strmm={sgemm_small:re,sgemm_large:ee,triangularize:Xt};nr.strsm={strsv_invert_block:xe,block_transfer:Zt,sscal:ge,sgemm_small:re,sgemm_large:ee};Ia=Object.assign({},...Object.values(nr))});var Ta={};Ee(Ta,{GpuMatrix:()=>F,GpuVector:()=>I,cleanup:()=>Ie,dasum:()=>no,gpuName:()=>Re,idamax:()=>lo,init:()=>Me,isamax:()=>uo,randomFloat32Array:()=>De,randomFloat64Array:()=>Te,randomTriangularFloat32Array:()=>je,sasum:()=>so,saxpy:()=>to,scopy:()=>oo,sdot:()=>ao,sgemm:()=>_o,sgemmtr:()=>Bo,sgemv:()=>co,sger:()=>yo,snrm2:()=>io,srot:()=>mo,srotm:()=>fo,sscal:()=>ro,sswap:()=>eo,ssymm:()=>Go,ssymv:()=>po,ssyr:()=>xo,ssyr2:()=>vo,ssyr2k:()=>Ao,ssyrk:()=>Eo,strmm:()=>So,strmv:()=>wo,strsm:()=>ko,strsv:()=>ho});function Ge(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 Se(r){if(!ke(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 vr(r,t,e){if(!e)return null;let a=r.createBuffer({label:"timestamp-resolve",size:16,usage:GPUBufferUsage.QUERY_RESOLVE|GPUBufferUsage.COPY_SRC});t.resolveQuerySet(e,0,2,a,0);let o=r.createBuffer({label:"timestamp-readback",size:16,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ});return t.copyBufferToBuffer(a,0,o,0,16),{tsReadBuffer:o,resolveBuffer:a,querySet:e}}async function M(r){if(!r)return;let{tsReadBuffer:t,resolveBuffer:e,querySet:a}=r;await t.mapAsync(GPUMapMode.READ);let o=new BigInt64Array(t.getMappedRange().slice());return t.unmap(),t.destroy(),e.destroy(),a.destroy(),Math.max(0,Number(o[1]-o[0]))/1e6}var Or=null,pe=!1,Kr=new Map,Jr=new WeakMap,Tr=null,Ne=({powerPreference:r,benchmark:t})=>`${r}::${t}`;async function Me({powerPreference:r="high-performance",benchmark:t=!1,dumpShaders:e=!1}={}){let a={powerPreference:r,benchmark:t,dumpShaders:e},o=Ne(a),s=Kr.get(o);if(s)return s;if(Or)e!==pe&&typeof window>"u"&&console.warn(`dumpShaders: ${e} was requested, but the WebGPU instance was already created with dumpShaders: ${pe}. The first init() call fixes this for the process.`);else if(typeof window>"u"){let{create:f,globals:d}=await import("webgpu");Object.assign(globalThis,d),Or=f(e?["enable-dawn-features=dump_shaders,disable_symbol_renaming"]:[]),pe=e}else e&&console.warn("dumpShaders has no effect in the browser \u2014 see init()'s docs."),Or=navigator.gpu;if(!Or)throw new Error("WebGPU not supported in this environment.");let n=await Or.requestAdapter({powerPreference:r})??await Or.requestAdapter();if(!n)throw new Error("No WebGPU adapter found.");let i=[...Ge(n,t).requiredFeatures??[]],u=await n.requestDevice({requiredFeatures:i});u.addEventListener("uncapturederror",f=>{console.error("Uncaptured GPU error:",f.error.message)});let m=i.includes("timestamp-query");return Jr.set(u,{adapter:n,benchmark:m,options:a}),Kr.set(o,u),Tr||(Tr=u),u}function Ie(r){if(r===void 0){for(let e of Kr.values())e.destroy();Kr.clear(),Tr=null;return}let t=Jr.get(r);t&&(Kr.delete(Ne(t.options)),Jr.delete(r),r.destroy(),Tr===r&&(Tr=Kr.values().next().value??null))}function Re(r=Tr){let t=r&&Jr.get(r);if(!t)throw new Error("WebGPU adapter not initialized \u2014 call init() first.");let{device:e,description:a}=t.adapter.info;return{description:a||"unknown",device:e||"unknown"}}function ke(r=Tr){return Jr.get(r)?.benchmark??!1}function Vr(){if(!Tr)throw new Error("WebGPU device not initialized \u2014 call init() first.");return Tr}function p(...r){r.flat().forEach(t=>t.destroy())}function de(r,t,e){let a=r.limits.maxStorageBufferBindingSize;if(t>a)throw new Error(`Buffer "${e}" needs ${t} bytes, exceeding this device's maxStorageBufferBindingSize (${a} bytes). The operands are too large for this device.`)}function x(r,t,e="blas-input",a=!1){let o=t.byteLength;de(r,o,e);let s=a?GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC:GPUBufferUsage.STORAGE,n=r.createBuffer({label:e,size:o,usage:s,mappedAtCreation:!0}),l=t.constructor;return new l(n.getMappedRange()).set(t),n.unmap(),n}function sr(r,t,e="blas-storage",a=0){return de(r,t,e),r.createBuffer({label:e,size:t,usage:GPUBufferUsage.STORAGE|a})}function Nr(r,t,e="blas-result"){return de(r,t,e),r.createBuffer({label:e,size:t,usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC})}function N(r,t,e){let a=r.createBuffer({label:"blas-readback",size:e.size,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ});return t.copyBufferToBuffer(e,0,a,0,e.size),a}var zr=16,Pe=new WeakMap;function Ko(r){let t=Pe.get(r);return t||(t=r.createBuffer({label:"blas-vec4-fallback",size:zr,usage:GPUBufferUsage.STORAGE}),Pe.set(r,t)),t}function _r(r,t){let e=t instanceof GPUBuffer?t:t.buffer,a=t instanceof GPUBuffer?0:t.offset??0,o=t instanceof GPUBuffer?t.size:t.size??e.size-a,s=Math.floor(o/zr)*zr;return s<zr?{buffer:Ko(r),offset:0,size:zr}:{buffer:e,offset:a,size:s}}function we(r,t,e,a){if(t%4!==0)return!1;let o=r instanceof GPUBuffer?r:r.buffer,s=r instanceof GPUBuffer?0:r.offset??0,n=r instanceof GPUBuffer?o.size:r.size??o.size-s,l=Math.floor(n/zr)*4;if(l<=0)return!1;let i=(Math.max(e,1)-1)*t+(Math.max(a,1)-1);return Math.floor(i/4)*4+4<=l}function P(r,t,e="blas-params"){let a=t.length*4,o=Math.ceil(a/16)*16,s=new ArrayBuffer(o),n=new DataView(s);t.forEach(({value:i,type:u},m)=>{let f=m*4;if(u==="u32")n.setUint32(f,i,!0);else if(u==="i32")n.setInt32(f,i,!0);else if(u==="f32")n.setFloat32(f,i,!0);else throw new Error(`Unknown param type "${u}". Use "f32", "u32", or "i32".`)});let l=r.createBuffer({label:e,size:o,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST});return r.queue.writeBuffer(l,0,s),l}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 Cr(r){let t=r.length,e=new Float32Array(t),a=new Float32Array(t);for(let o=0;o<t;o++){let s=Math.fround(r[o]);e[o]=s,a[o]=Math.fround(r[o]-s)}return{hi:e,lo:a}}function Hr(r,t){let e=r.length,a=new Float64Array(e);for(let o=0;o<e;o++)a[o]=r[o]+t[o];return a}var I=class r{constructor(t,e,a=Float32Array,o=null,s=null){this._buf=t,this._loBuf=o,this.length=e,this.dtype=a,this.device=s??Vr()}static from(t,e){let a=t instanceof GPUDevice,o=a?t:Vr(),s=a?e:t;if(s instanceof Float64Array){let{hi:l,lo:i}=Cr(s),u=x(o,l,"gpu-vector-f64-hi",!0),m=x(o,i,"gpu-vector-f64-lo",!0);return new r(u,s.length,Float64Array,m,o)}if(!(s instanceof Float32Array))throw new Error("GpuVector.from expects a Float32Array or Float64Array.");let n=x(o,s,"gpu-vector",!0);return new r(n,s.length,s.constructor,null,o)}async read(){let t=this.device,e=t.createCommandEncoder(),a=N(t,e,this._buf);if(t.queue.submit([e.finish()]),!this._loBuf)return S(a,this.dtype);let o=t.createCommandEncoder(),s=N(t,o,this._loBuf);t.queue.submit([o.finish()]);let[n,l]=await Promise.all([S(a,Float32Array),S(s,Float32Array)]);return Hr(n,l)}destroy(){this._buf.destroy(),this._loBuf&&this._loBuf.destroy()}};var F=class r{constructor(t,e,a,o,s=null,n="row-major",l=null){this._buf=t,this._loBuf=s,this.rows=e,this.cols=a,this.lda=o,this.layout=n,this.device=l??Vr()}static from(t,...e){let a=t instanceof GPUDevice,o=a?t:Vr(),s=a?e.shift():t,[n,l,i,u="row-major"]=e;if(u!=="row-major"&&u!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");let m=u==="row-major";if(i===void 0&&(i=m?l:n),!(s instanceof Float32Array)&&!(s instanceof Float64Array))throw new Error("GpuMatrix.from expects a Float32Array or Float64Array.");if(!Number.isInteger(n)||n<=0)throw new Error("rows must be a positive integer.");if(!Number.isInteger(l)||l<=0)throw new Error("cols must be a positive integer.");let f=m?l:n;if(!Number.isInteger(i)||i<f)throw new Error(`lda must be an integer >= ${m?"cols":"rows"}.`);let d=m?n:l;if(s.length<d*i)throw new Error("data does not have enough elements for the given rows, cols, and lda.");if(s instanceof Float64Array){let w=d*i,{hi:g,lo:h}=Cr(s.subarray(0,w)),b=x(o,g,"gpu-matrix-f64-hi",!0),y=x(o,h,"gpu-matrix-f64-lo",!0);return new r(b,n,l,i,y,u,o)}let c=x(o,s.subarray(0,d*i),"gpu-matrix",!0);return new r(c,n,l,i,null,u,o)}async read(){let t=this.device,e=t.createCommandEncoder(),a=N(t,e,this._buf);t.queue.submit([e.finish()]);let o=this.layout!=="column-major",s=o?this.rows:this.cols,n=o?this.cols:this.rows;if(this._loBuf){let u=t.createCommandEncoder(),m=N(t,u,this._loBuf);t.queue.submit([u.finish()]);let[f,d]=await Promise.all([S(a,Float32Array),S(m,Float32Array)]),c=Hr(f,d);if(this.lda===n)return c;let w=new Float64Array(s*n);for(let g=0;g<s;g++)w.set(c.subarray(g*this.lda,g*this.lda+n),g*n);return w}let l=await S(a,Float32Array);if(this.lda===n)return l;let i=new Float32Array(s*n);for(let u=0;u<s;u++)i.set(l.subarray(u*this.lda,u*this.lda+n),u*n);return i}destroy(){this._buf.destroy(),this._loBuf&&this._loBuf.destroy()}};function De(r,t=-1,e=1){let a=new Float32Array(r);for(let o=0;o<r;o++)a[o]=t+Math.random()*(e-t);return a}function Te(r,t=-1,e=1){let a=new Float64Array(r);for(let o=0;o<r;o++)a[o]=t+Math.random()*(e-t);return a}function je(r,t,e="lower",a=-1,o=1,s=5,n=15){if(e!=="lower"&&e!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(t<r)throw new Error("lda must be >= n.");let l=new Float32Array(r*t);for(let i=0;i<r;i++){for(let u=0;u<r;u++){if(i===u)continue;(e==="lower"?u<i:u>i)&&(l[i*t+u]=a+Math.random()*(o-a))}l[i*t+i]=s+Math.random()*(n-s)}return l}function E(r,t,e,a=0){let o=e.map((s,n)=>({binding:a+n,resource:s instanceof GPUBuffer?{buffer:s}:s}));return r.createBindGroup({layout:t,entries:o})}var Vo=new WeakMap;function R(r,t){r.queue.submit([t.finish()])}function Mr(r){let{querySet:t,passDescriptor:e}=Se(r);return{commandEncoder:r.createCommandEncoder(),querySet:t,passDescriptor:e}}function ur(r,t,e,a,o){let s=r.beginComputePass(o);s.setPipeline(t),s.setBindGroup(0,e),typeof a=="number"?s.dispatchWorkgroups(a):s.dispatchWorkgroups(a.x,a.y,a.z??1),s.end(),Vo.set(r,s)}function W(r,t,e,a){let{commandEncoder:o,querySet:s,passDescriptor:n}=Mr(r);ur(o,t,e,a,n);let l=vr(r,o,s);return{commandEncoder:o,ts:l}}var Da={},ve=new WeakMap;async function G(r,t,e="main"){ve.has(r)||ve.set(r,new Map);let a=ve.get(r),o=Array.isArray(t)?t:[t],s=`${o.join("+")}::${e}`;return a.has(s)||a.set(s,await Pa(r,o,e)),a.get(s)}async function Ra(r){if(typeof process>"u"||!process.versions?.node){let{shaderSources:t}=await Promise.resolve().then(()=>(Jt(),Qt)),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:a,join:o}=await import("path"),s=a(e(Da.url));return t(o(s,`../shaders/${r}.wgsl`),"utf8")}}async function Pa(r,t,e="main"){let a=t.join("+"),o=(await Promise.all(t.map(Ra))).join(`
|
|
2487
|
+
`),s=r.createShaderModule({label:a,code:o}),l=(await s.getCompilationInfo()).messages.filter(m=>m.type==="error");if(l.length>0)throw new Error(`Shader "${a}" compilation failed:
|
|
2488
|
+
${l.map(m=>` line ${m.lineNum}: ${m.message}`).join(`
|
|
2489
|
+
`)}`);let i=e==="main"?{module:s}:{module:s,entryPoint:e},u=r.createComputePipeline({label:a,layout:"auto",compute:i});return u._shaderModule=s,u}function yr(r,t,e){let a=r.limits.maxComputeWorkgroupsPerDimension;return e===void 0?Math.min(Math.ceil(t/64),a):{x:Math.min(Math.ceil(e/8),a),y:Math.min(Math.ceil(t/8),a)}}function O(r,t,e,a="x"){let o=r.limits.maxComputeWorkgroupsPerDimension;if(t>o)throw new Error(`${e}: this problem needs ${t} workgroups in ${a}, but the device allows ${o} (maxComputeWorkgroupsPerDimension). The operands are too large for this device \u2014 split the operation into smaller blocks.`);return t}function qr(r,t,e,a){return a===void 0?O(r,Math.ceil(e/64),t):{x:O(r,Math.ceil(a/8),t,"x"),y:O(r,Math.ceil(e/8),t,"y")}}function T(r,t,e){for(let[a,o]of Object.entries(e))if(!(!(o instanceof I)&&!(o instanceof F))&&o.device!==r)throw new Error(`${t}: ${a} 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 ro(r,t,e,a,o){let s=a instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"sscal",{x:a}),!Number.isInteger(t)||!Number.isInteger(o))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<=0)throw new Error("incx must be positive.");if(!(a instanceof Float32Array)&&!(a instanceof I))throw new Error("x must be a Float32Array or GpuVector.");if(t<=0)return s?{}:{x:a};if(a.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let n=await G(r,"sscal"),l=null,i=null,u=null;try{l=s?a._buf:x(r,a,"sscal-x",!0),i=P(r,[{value:t,type:"u32"},{value:e,type:"f32"},{value:o,type:"u32"}],"sscal-params");let m=E(r,n.getBindGroupLayout(0),[l,i]),{commandEncoder:f,ts:d}=W(r,n,m,yr(r,t));u=s?null:N(r,f,l),R(r,f);let c=await M(d);if(s)return c!==void 0?{gpuTimeMs:c}:{};let w=await S(u,Float32Array);return u=null,c!==void 0?{x:w,gpuTimeMs:c}:{x:w}}finally{!s&&l&&p(l),i&&p(i),u&&p(u)}}async function eo(r,t,e,a,o,s){let n=e instanceof I,l=o instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"sswap",{x:e,y:o}),!Number.isInteger(t)||!Number.isInteger(a)||!Number.isInteger(s))throw new Error("n, incx, and incy must be integers.");if(a<=0||s<=0)throw new Error("incx and incy must be positive.");if(!(e instanceof Float32Array)&&!(e instanceof I))throw new Error("x must be a Float32Array or GpuVector.");if(!(o instanceof Float32Array)&&!(o instanceof I))throw new Error("y must be a Float32Array or GpuVector.");if(e.constructor!==o.constructor)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(t<=0)return n?{}:{x:e,y:o};if(e.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(o.length<(t-1)*s+1)throw new Error("y does not have enough elements for the given n and incy.");let i=await G(r,"sswap"),u=null,m=null,f=null,d=null,c=null;try{u=n?e._buf:x(r,e,"sswap-x",!0),m=l?o._buf:x(r,o,"sswap-y",!0),f=P(r,[{value:t,type:"u32"},{value:a,type:"u32"},{value:s,type:"u32"}],"sswap-params");let w=E(r,i.getBindGroupLayout(0),[u,m,f]),{commandEncoder:g,ts:h}=W(r,i,w,yr(r,t));d=n?null:N(r,g,u),c=l?null:N(r,g,m),R(r,g);let b=await M(h);if(n&&l)return b!==void 0?{gpuTimeMs:b}:{};let y=await S(d,Float32Array);d=null;let _=await S(c,Float32Array);return c=null,b!==void 0?{x:y,y:_,gpuTimeMs:b}:{x:y,y:_}}finally{!n&&u&&p(u),!l&&m&&p(m),f&&p(f),d&&p(d),c&&p(c)}}async function to(r,t,e,a,o,s,n){let l=a instanceof I,i=s instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"saxpy",{x:a,y:s}),!Number.isInteger(t)||!Number.isInteger(o)||!Number.isInteger(n))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<=0||n<=0)throw new Error("incx and incy must be positive.");if(!l&&!(a instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!i&&!(s instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(l!==i)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(t<=0)return i?{}:{y:s};if(a.length<(t-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(s.length<(t-1)*n+1)throw new Error("y does not have enough elements for the given n and incy.");let u=await G(r,"saxpy"),m=null,f=null,d=null,c=null;try{m=l?a._buf:x(r,a,"saxpy-x",!1),f=i?s._buf:x(r,s,"saxpy-y",!0),d=P(r,[{value:t,type:"u32"},{value:e,type:"f32"},{value:o,type:"u32"},{value:n,type:"u32"}],"saxpy-params");let w=E(r,u.getBindGroupLayout(0),[m,f,d]),{commandEncoder:g,ts:h}=W(r,u,w,yr(r,t));c=i?null:N(r,g,f),R(r,g);let b=await M(h);if(i&&l)return b!==void 0?{gpuTimeMs:b}:{};let y=await S(c,Float32Array);return c=null,b!==void 0?{y,gpuTimeMs:b}:{y}}finally{!l&&m&&p(m),!i&&f&&p(f),d&&p(d),c&&p(c)}}async function oo(r,t,e,a,o,s){let n=e instanceof I,l=o instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"scopy",{x:e,y:o}),!Number.isInteger(t)||!Number.isInteger(a)||!Number.isInteger(s))throw new Error("n, incx, and incy must be integers.");if(a<=0||s<=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(!l&&!(o instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(n!==l)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(t<=0)return l?{}:{y:o};if(e.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(o.length<(t-1)*s+1)throw new Error("y does not have enough elements for the given n and incy.");let i=await G(r,"scopy"),u=null,m=null,f=null,d=null;try{u=n?e._buf:x(r,e,"scopy-x",!1),m=l?o._buf:x(r,o,"scopy-y",!0),f=P(r,[{value:t,type:"u32"},{value:a,type:"u32"},{value:s,type:"u32"}],"scopy-params");let c=E(r,i.getBindGroupLayout(0),[u,m,f]),{commandEncoder:w,ts:g}=W(r,i,c,yr(r,t));d=l?null:N(r,w,m),R(r,w);let h=await M(g);if(l&&n)return h!==void 0?{gpuTimeMs:h}:{};let b=await S(d,Float32Array);return d=null,h!==void 0?{y:b,gpuTimeMs:h}:{y:b}}finally{!n&&u&&p(u),!l&&m&&p(m),f&&p(f),d&&p(d)}}async function ao(r,t,e,a,o,s){let n=e instanceof I,l=o instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"sdot",{x:e,y:o}),!Number.isInteger(t)||!Number.isInteger(a)||!Number.isInteger(s))throw new Error("n, incx, and incy must be integers.");if(a<=0||s<=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(!l&&!(o instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(n!==l)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)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(o.length<(t-1)*s+1)throw new Error("y does not have enough elements for the given n and incy.");let i=await G(r,"sdot"),u=await G(r,"reduction/sum"),m=null,f=null,d=null,c=null,w=null,g=null;try{m=n?e._buf:x(r,e,"sdot-x",!1),f=l?o._buf:x(r,o,"sdot-y",!1),d=sr(r,512,"sdot-partials"),c=Nr(r,4,"sdot-result"),w=P(r,[{value:t,type:"u32"},{value:a,type:"u32"},{value:s,type:"u32"}],"sdot-params");let h=E(r,i.getBindGroupLayout(0),[m,f,d,w]),{commandEncoder:b,ts:y}=W(r,i,h,128);R(r,b);let _=E(r,u.getBindGroupLayout(0),[d,c]),{commandEncoder:v,ts:A}=W(r,u,_,1);g=N(r,v,c),R(r,v);let k=S(g,Float32Array);g=null;let[B,j,D]=await Promise.all([M(y),M(A),k]);return B!==void 0&&j!==void 0?{dot:D[0],gpuTimeMs:B+j}:{dot:D[0]}}finally{!n&&m&&p(m),!l&&f&&p(f),d&&p(d),c&&p(c),w&&p(w),g&&p(g)}}async function so(r,t,e,a){let o=e instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"sasum",{x:e}),!Number.isInteger(t)||!Number.isInteger(a))throw new Error("n and incx must be integers.");if(a<=0)throw new Error("incx must be positive.");if(!o&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(t<=0)return{asum:0};if(e.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");let s=await G(r,"sasum"),n=await G(r,"reduction/sum"),l=null,i=null,u=null,m=null,f=null;try{l=o?e._buf:x(r,e,"sasum-x",!1),i=sr(r,512,"sasum-partials"),u=Nr(r,4,"sasum-result"),m=P(r,[{value:t,type:"u32"},{value:a,type:"u32"}],"sasum-params");let d=E(r,s.getBindGroupLayout(0),[l,i,m]),{commandEncoder:c,ts:w}=W(r,s,d,128);R(r,c);let g=E(r,n.getBindGroupLayout(0),[i,u]),{commandEncoder:h,ts:b}=W(r,n,g,1);f=N(r,h,u),R(r,h);let y=S(f,Float32Array);f=null;let[_,v,A]=await Promise.all([M(w),M(b),y]);return _!==void 0&&v!==void 0?{asum:A[0],gpuTimeMs:_+v}:{asum:A[0]}}finally{!o&&l&&p(l),i&&p(i),u&&p(u),m&&p(m),f&&p(f)}}async function no(r,t,e,a){let o=e instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"dasum",{x:e}),!Number.isInteger(t)||!Number.isInteger(a))throw new Error("n and incx must be integers.");if(a<=0)throw new Error("incx must be positive.");if(!o&&!(e instanceof Float64Array))throw new Error("x must be a Float64Array or GpuVector.");if(o&&e.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(t<=0)return{asum:0};if(e.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");let s=["f64/dekker","f64/utils/abs","f64/utils/add"],n=await G(r,[...s,"dasum"]),l=await G(r,[...s,"reduction/sumF64"]),i=null,u=null,m=null,f=null,d=null,c=null,w=null,g=null,h=null;try{if(o)i=e._buf,u=e._loBuf;else{let{hi:X,lo:q}=Cr(e.map(Math.abs));i=x(r,X,"dasum-xHi",!1),u=x(r,q,"dasum-xLo",!1)}m=sr(r,512,"dasum-partialsHi"),f=sr(r,512,"dasum-partialsLo"),d=Nr(r,4,"dasum-result-hi"),c=Nr(r,4,"dasum-result-lo"),w=P(r,[{value:t,type:"u32"},{value:a,type:"u32"}],"dasum-params");let b=E(r,n.getBindGroupLayout(0),[i,u,m,f,w]),{commandEncoder:y,ts:_}=W(r,n,b,128);R(r,y);let v=E(r,l.getBindGroupLayout(0),[m,f,d,c]),{commandEncoder:A,ts:k}=W(r,l,v,1);g=N(r,A,d),h=N(r,A,c),R(r,A);let B=S(g,Float32Array),j=S(h,Float32Array);g=null,h=null;let[D,C,L,U]=await Promise.all([M(_),M(k),B,j]),H=Hr(L,U)[0];return D!==void 0&&C!==void 0?{asum:H,gpuTimeMs:D+C}:{asum:H}}finally{!o&&i&&p(i),!o&&u&&p(u),m&&p(m),f&&p(f),d&&p(d),c&&p(c),w&&p(w),g&&p(g),h&&p(h)}}async function io(r,t,e,a){let o=e instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"snrm2",{x:e}),!Number.isInteger(t)||!Number.isInteger(a))throw new Error("n and incx must be integers.");if(a<=0)throw new Error("incx must be positive.");if(!o&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(t<=0)return{nrm2:0};if(e.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");let s=await G(r,"snrm2"),n=await G(r,"reduction/scaledSum"),l=null,i=null,u=null,m=null,f=null,d=null;try{l=o?e._buf:x(r,e,"snrm2-x",!1),i=sr(r,512,"snrm2-partials-scale"),u=sr(r,512,"snrm2-partials-ssq"),m=Nr(r,4,"snrm2-result"),f=P(r,[{value:t,type:"u32"},{value:a,type:"u32"}],"snrm2-params");let c=E(r,s.getBindGroupLayout(0),[l,i,u,f]),{commandEncoder:w,ts:g}=W(r,s,c,128);R(r,w);let h=E(r,n.getBindGroupLayout(0),[i,u,m]),{commandEncoder:b,ts:y}=W(r,n,h,1);d=N(r,b,m),R(r,b);let _=S(d,Float32Array);d=null;let[v,A,k]=await Promise.all([M(g),M(y),_]),B=k[0];return v!==void 0&&A!==void 0?{nrm2:B,gpuTimeMs:v+A}:{nrm2:B}}finally{!o&&l&&p(l),i&&p(i),u&&p(u),m&&p(m),f&&p(f),d&&p(d)}}async function uo(r,t,e,a){let o=e instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"isamax",{x:e}),!Number.isInteger(t)||!Number.isInteger(a))throw new Error("n and incx must be integers.");if(a<=0)throw new Error("incx must be positive.");if(!o&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(t<=0)return{index:0};if(e.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");let s=await G(r,"isamax"),n=await G(r,"reduction/argmax"),l=null,i=null,u=null,m=null,f=null,d=null;try{l=o?e._buf:x(r,e,"isamax-x",!1),i=sr(r,512,"isamax-partials-val"),u=sr(r,512,"isamax-partials-idx"),m=Nr(r,4,"isamax-result"),f=P(r,[{value:t,type:"u32"},{value:a,type:"u32"}],"isamax-params");let c=E(r,s.getBindGroupLayout(0),[l,i,u,f]),{commandEncoder:w,ts:g}=W(r,s,c,128);R(r,w);let h=E(r,n.getBindGroupLayout(0),[i,u,m]),{commandEncoder:b,ts:y}=W(r,n,h,1);d=N(r,b,m),R(r,b);let _=S(d,Uint32Array);d=null;let[v,A,k]=await Promise.all([M(g),M(y),_]),B=k[0];return v!==void 0&&A!==void 0?{index:B,gpuTimeMs:v+A}:{index:B}}finally{!o&&l&&p(l),i&&p(i),u&&p(u),m&&p(m),f&&p(f),d&&p(d)}}async function lo(r,t,e,a){let o=e instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"idamax",{x:e}),!Number.isInteger(t)||!Number.isInteger(a))throw new Error("n and incx must be integers.");if(a<=0)throw new Error("incx must be positive.");if(!o&&!(e instanceof Float64Array))throw new Error("x must be a Float64Array or GpuVector.");if(o&&e.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(t<=0)return{index:0};if(e.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");let s=["f64/dekker","f64/utils/abs","f64/utils/greater","f64/utils/equal"],n=await G(r,[...s,"idamax"],"idamax_main"),l=await G(r,[...s,"reduction/argmaxF64"],"reduce_f64"),i=null,u=null,m=null,f=null,d=null,c=null,w=null,g=null;try{if(o)i=e._buf,u=e._loBuf;else{let{hi:L,lo:U}=Cr(e);i=x(r,L,"idamax-xHi",!1),u=x(r,U,"idamax-xLo",!1)}m=sr(r,512,"idamax-partials-val-hi"),f=sr(r,512,"idamax-partials-val-lo"),d=sr(r,512,"idamax-partials-idx"),c=Nr(r,4,"idamax-result"),w=P(r,[{value:t,type:"u32"},{value:a,type:"u32"}],"idamax-params");let h=E(r,n.getBindGroupLayout(0),[i,u,m,f,d,w]),{commandEncoder:b,ts:y}=W(r,n,h,128);R(r,b);let _=E(r,l.getBindGroupLayout(0),[m,f,d,c]),{commandEncoder:v,ts:A}=W(r,l,_,1);g=N(r,v,c),R(r,v);let k=S(g,Uint32Array);g=null;let[B,j,D]=await Promise.all([M(y),M(A),k]),C=D[0];return B!==void 0&&j!==void 0?{index:C,gpuTimeMs:B+j}:{index:C}}finally{!o&&i&&p(i),!o&&u&&p(u),m&&p(m),f&&p(f),d&&p(d),c&&p(c),w&&p(w),g&&p(g)}}async function mo(r,t,e,a,o,s,n,l){let i=e instanceof I,u=o instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"srot",{x:e,y:o}),!Number.isInteger(t)||!Number.isInteger(a)||!Number.isInteger(s))throw new Error("n, incx, and incy must be integers.");if(typeof n!="number")throw new Error("c must be a number.");if(typeof l!="number")throw new Error("s must be a number.");if(Number.isNaN(n)||Number.isNaN(l))throw new Error("c and s must not be NaN.");if(!Number.isFinite(n))throw new Error("c must be finite.");if(!Number.isFinite(l))throw new Error("s must be finite.");if(a<=0||s<=0)throw new Error("incx and incy must be positive.");if(!i&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!u&&!(o instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(i!==u)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(t<=0)return i?{}:{x:e,y:o};if(e.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(o.length<(t-1)*s+1)throw new Error("y does not have enough elements for the given n and incy.");let m=await G(r,"srot"),f=null,d=null,c=null,w=null,g=null;try{f=i?e._buf:x(r,e,"srot-x",!0),d=u?o._buf:x(r,o,"srot-y",!0),c=P(r,[{value:t,type:"u32"},{value:n,type:"f32"},{value:l,type:"f32"},{value:a,type:"u32"},{value:s,type:"u32"}],"srot-params");let h=E(r,m.getBindGroupLayout(0),[f,d,c]),{commandEncoder:b,ts:y}=W(r,m,h,yr(r,t));w=i?null:N(r,b,f),g=u?null:N(r,b,d),R(r,b);let _=await M(y);if(i&&u)return _!==void 0?{gpuTimeMs:_}:{};let v=S(w,Float32Array),A=S(g,Float32Array);w=null,g=null;let[k,B]=await Promise.all([v,A]);return _!==void 0?{x:k,y:B,gpuTimeMs:_}:{x:k,y:B}}finally{!i&&f&&p(f),!u&&d&&p(d),c&&p(c),w&&p(w),g&&p(g)}}async function fo(r,t,e,a,o,s,n){let l=e instanceof I,i=o instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"srotm",{x:e,y:o}),!Number.isInteger(t)||!Number.isInteger(a)||!Number.isInteger(s))throw new Error("n, incx, and incy must be integers.");if(!(n instanceof Float32Array)||n.length!==5)throw new Error("param must be a Float32Array of length 5.");if(n[0]!==-2&&n[0]!==-1&&n[0]!==0&&n[0]!==1)throw new Error("param[0] (flag) must be one of -2, -1, 0, or 1.");if(a<=0||s<=0)throw new Error("incx and incy must be positive.");if(!l&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!i&&!(o instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(l!==i)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(t<=0||n[0]===-2)return l?{}:{x:e,y:o};if(e.length<(t-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(o.length<(t-1)*s+1)throw new Error("y does not have enough elements for the given n and incy.");let u=await G(r,"srotm"),m=null,f=null,d=null,c=null,w=null,g=null;try{m=l?e._buf:x(r,e,"srotm-x",!0),f=i?o._buf:x(r,o,"srotm-y",!0),d=x(r,n,"srotm-param",!1),c=P(r,[{value:t,type:"u32"},{value:a,type:"u32"},{value:s,type:"u32"}],"srotm-params");let h=E(r,u.getBindGroupLayout(0),[m,f,d,c]),{commandEncoder:b,ts:y}=W(r,u,h,yr(r,t));w=l?null:N(r,b,m),g=i?null:N(r,b,f),R(r,b);let _=await M(y);if(l&&i)return _!==void 0?{gpuTimeMs:_}:{};let v=S(w,Float32Array),A=S(g,Float32Array);w=null,g=null;let[k,B]=await Promise.all([v,A]);return _!==void 0?{x:k,y:B,gpuTimeMs:_}:{x:k,y:B}}finally{!l&&m&&p(m),!i&&f&&p(f),d&&p(d),c&&p(c),w&&p(w),g&&p(g)}}async function co(r,t,e,a,o,s,n,l,i,u,m,f,d="row-major"){let c=s instanceof F,w=l instanceof I,g=m instanceof I;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"sgemv",{A:s,x:l,y:m}),t!=="no-transpose"&&t!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(d!=="row-major"&&d!=="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(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(e)||!Number.isInteger(a)||!Number.isInteger(i)||!Number.isInteger(f)||!Number.isInteger(n))throw new Error("m, n, incx, incy, and lda must be integers.");if(i<=0||f<=0)throw new Error("incx and incy must be positive.");if(!c&&!(s 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(!g&&!(m instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(w!==g)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(w&&!c)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(c&&!w)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(w&&l._buf===m._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(c&&g&&s._buf===m._buf)throw new Error("A and y must not reference the same GPU buffer.");if(c&&n!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(c&&(s.rows<e||s.cols<a))throw new Error("A is too small for the given m and n.");if(e<0||a<0)throw new Error("m and n must be non-negative.");if(e===0||a===0)return g?{}:{y:m};(c?s.layout:d)==="column-major"&&([e,a]=[a,e],t=t==="no-transpose"?"transpose":"no-transpose");let b=t==="no-transpose",y=b?a:e,_=b?e:a;if(n<a)throw new Error("lda must be >= n.");if(!c&&s.length<(e-1)*n+a)throw new Error("A does not have enough elements for the given m, n, and lda.");if(l.length<(y-1)*i+1)throw new Error("x does not have enough elements for the given dimensions and incx.");if(m.length<(_-1)*f+1)throw new Error("y does not have enough elements for the given dimensions and incy.");let A=await G(r,b?"sgemv_n":"sgemv_t"),k=null,B=null,j=null,D=null;try{k=c?s._buf:x(r,s,"sgemv-A",!1),B=w?l._buf:x(r,l,"sgemv-x",!1),j=g?m._buf:x(r,m,"sgemv-y",!0),D=P(r,[{value:e,type:"u32"},{value:a,type:"u32"},{value:o,type:"f32"},{value:u,type:"f32"},{value:i,type:"u32"},{value:f,type:"u32"},{value:n,type:"u32"}],"sgemv-params");let C=E(r,A.getBindGroupLayout(0),[k,B,j,D]),L=b?Math.min(e,r.limits.maxComputeWorkgroupsPerDimension):qr(r,"sgemv",_),{commandEncoder:U,ts:H}=W(r,A,C,L),X=g?null:N(r,U,j);R(r,U);let q=await M(H);if(g)return q!==void 0?{gpuTimeMs:q}:{};let J=await S(X,Float32Array);return q!==void 0?{y:J,gpuTimeMs:q}:{y:J}}finally{!c&&k&&p(k),!w&&B&&p(B),!g&&j&&p(j),D&&p(D)}}async function po(r,t,e,a,o,s,n,l,i,u,m,f="row-major"){let d=n instanceof I,c=u instanceof I,w=o instanceof F;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"ssymv",{A:o,x:n,y:u}),t!=="lower"&&t!=="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(e)||!Number.isInteger(l)||!Number.isInteger(m)||!Number.isInteger(s))throw new Error("n, incx, incy, and lda must be integers.");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 i!="number")throw new Error("beta must be a number.");if(Number.isNaN(i))throw new Error("beta must not be NaN.");if(!Number.isFinite(i))throw new Error("beta must be finite.");if(l<=0||m<=0)throw new Error("incx and incy must be positive.");if(s<e)throw new Error("lda must be >= n.");if(!w&&!(o instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!d&&!(n 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(d!==c)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(d&&!w)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(w&&!d)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(d&&n._buf===u._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(w&&s!==o.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(w&&(o.rows<e||o.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:u};if(!w&&o.length<(e-1)*s+e)throw new Error("A does not have enough elements for the given n and lda.");if(n.length<(e-1)*l+1)throw new Error("x does not have enough elements for the given n and incx.");if(u.length<(e-1)*m+1)throw new Error("y does not have enough elements for the given n and incy.");let h=(w?o.layout:f)==="column-major"?t==="upper":t==="lower",b=await G(r,"ssymv"),y=null,_=null,v=null,A=null;try{y=w?o._buf:x(r,o,"ssymv-A",!1),_=d?n._buf:x(r,n,"ssymv-x",!1),v=c?u._buf:x(r,u,"ssymv-y",!0),A=P(r,[{value:e,type:"u32"},{value:a,type:"f32"},{value:i,type:"f32"},{value:l,type:"u32"},{value:m,type:"u32"},{value:s,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:j,ts:D}=W(r,b,k,B),C=c?null:N(r,j,v);R(r,j);let L=await M(D);if(c)return L!==void 0?{gpuTimeMs:L}:{};let U=await S(C,Float32Array);return L!==void 0?{y:U,gpuTimeMs:L}:{y:U}}finally{!w&&y&&p(y),!d&&_&&p(_),!c&&v&&p(v),A&&p(A)}}async function wo(r,t,e,a,o,s,n,l,i,u,m,f="row-major"){let d=l instanceof I,c=u instanceof I,w=s instanceof F,g=a==="unit";if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"strmv",{A:s,x:l,y: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(!g&&a!=="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(o)||!Number.isInteger(i)||!Number.isInteger(m)||!Number.isInteger(n))throw new Error("n, incx, incy, and lda must be integers.");if(i<=0||m<=0)throw new Error("incx and incy must be positive.");if(n<o)throw new Error("lda must be >= n.");if(!w&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!d&&!(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(d!==c)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(d&&l._buf===u._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(d&&!w)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(w&&!d)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(w&&c&&s._buf===u._buf)throw new Error("A and y must not reference the same GPU buffer.");if(w&&n!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(w&&(s.rows<o||s.cols<o))throw new Error("A is too small for the given n.");if(o<0)throw new Error("n must be non-negative.");if(o===0)return c?{}:{y:u};if(!w&&s.length<(o-1)*n+o)throw new Error("A does not have enough elements for the given n and lda.");if(l.length<(o-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");if(u.length<(o-1)*m+1)throw new Error("y does not have enough elements for the given n and incy.");let b=(w?s.layout:f)==="column-major",y=b?t==="upper":t==="lower",_=b?e==="transpose":e==="no-transpose",v=await G(r,"strmv"),A=null,k=null,B=null,j=null;try{A=w?s._buf:x(r,s,"strmv-A",!1),k=d?l._buf:x(r,l,"strmv-x",!1),B=c?u._buf:x(r,u,"strmv-y",!0),j=P(r,[{value:o,type:"u32"},{value:i,type:"u32"},{value:m,type:"u32"},{value:n,type:"u32"},{value:_?0:1,type:"u32"},{value:y?0:1,type:"u32"},{value:g?1:0,type:"u32"}],"strmv-params");let D=E(r,v.getBindGroupLayout(0),[A,k,B,j]),C=Math.min(o,r.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:L,ts:U}=W(r,v,D,C),H=c?null:N(r,L,B);R(r,L);let X=await M(U);if(c)return X!==void 0?{gpuTimeMs:X}:{};let q=await S(H,Float32Array);return X!==void 0?{y:q,gpuTimeMs:X}:{y:q}}finally{!w&&A&&p(A),!d&&k&&p(k),!c&&B&&p(B),j&&p(j)}}function go(r,t,e){let a=new ArrayBuffer(r*t),o=new DataView(a);for(let s=0;s<r;s++){let n=e(s),l=s*t;n.forEach((i,u)=>o.setUint32(l+u*4,i,!0))}return a}function bo(r,t,e){let a=r.createBuffer({label:e,size:t.byteLength,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST});return r.queue.writeBuffer(a,0,t),a}async function ho(r,t,e,a,o,s,n,l,i,u="row-major"){let m=l instanceof I,f=s instanceof F,d=a==="unit";if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"strsv",{A:s,x:l}),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(!d&&a!=="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(o)||!Number.isInteger(i)||!Number.isInteger(n))throw new Error("n, incx, and lda must be integers.");if(i<=0)throw new Error("incx must be positive.");if(n<o)throw new Error("lda must be >= n.");if(!f&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!m&&!(l instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(m&&!f)throw new Error("A must be a GpuMatrix when x is a GpuVector.");if(f&&!m)throw new Error("x must be a GpuVector when A is a GpuMatrix.");if(f&&m&&s._buf===l._buf)throw new Error("A and x must not reference the same GPU buffer.");if(f&&n!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(f&&(s.rows<o||s.cols<o))throw new Error("A is too small for the given n.");if(o<0)throw new Error("n must be non-negative.");if(o===0)return m?{}:{x:l};if(!f&&s.length<(o-1)*n+o)throw new Error("A does not have enough elements for the given n and lda.");if(l.length<(o-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");let w=(f?s.layout:u)==="column-major",g=w?t==="upper":t==="lower",h=w?e==="transpose":e==="no-transpose",b=await G(r,"strsv_invert_block"),y=await G(r,"strsv_apply_inverse"),_=await G(r,"strsv_update"),v=h===g,A=[];for(let q=0;q<o;q+=64)A.push(q);v||A.reverse();let k=A.length,B=r.limits.maxComputeWorkgroupsPerDimension,j=r.limits.minUniformBufferOffsetAlignment,D=null,C=null,L=null,U=null,H=null,X=null;try{D=f?s._buf:x(r,s,"strsv-A",!1),C=m?l._buf:x(r,l,"strsv-x",!0),L=sr(r,k*64*64*4,"strsv-Ainv");let q=go(k,j,$=>{let K=$*64,Y=Math.min(K+64,o);return[i,$,K,Y]});U=bo(r,q,"strsv-apply-params");let J=go(k,j,$=>{let K=$*64,Y=Math.min(K+64,o);return[o,i,n,h?0:1,g?0:1,K,Y]});H=bo(r,J,"strsv-update-params");let{commandEncoder:Z,querySet:rr}=Mr(r);X=P(r,[{value:o,type:"u32"},{value:n,type:"u32"},{value:h?0:1,type:"u32"},{value:g?0:1,type:"u32"},{value:d?1:0,type:"u32"}],"strsv-invert-params");let lr=E(r,b.getBindGroupLayout(0),[D,L,X]);ur(Z,b,lr,{x:64,y:k},rr?{timestampWrites:{querySet:rr,beginningOfPassWriteIndex:0}}:void 0);for(let $=0;$<A.length;$++){let K=A[$],Y=Math.min(K+64,o),Q=K/64,ir=$===A.length-1,dr=Q*j,ar=E(r,y.getBindGroupLayout(0),[L,C,{buffer:U,offset:dr,size:16}]);ur(Z,y,ar,1,ir&&rr?{timestampWrites:{querySet:rr,endOfPassWriteIndex:1}}:void 0);let wr=v?o-Y:K;if(wr===0)continue;let Rr=E(r,_.getBindGroupLayout(0),[D,C,{buffer:H,offset:dr,size:32}]),kr=Math.min(wr,B);ur(Z,_,Rr,kr)}let pr=vr(r,Z,rr),er=m?null:N(r,Z,C);R(r,Z);let or=await M(pr);if(m)return or!==void 0?{gpuTimeMs:or}:{};let z=await S(er,Float32Array);return or!==void 0?{x:z,gpuTimeMs:or}:{x:z}}finally{!f&&D&&p(D),!m&&C&&p(C),L&&p(L),U&&p(U),H&&p(H),X&&p(X)}}async function yo(r,t,e,a,o,s,n,l,i,u,m="row-major"){let f=i instanceof F;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"sger",{A:i,x:o,y:n}),m!=="row-major"&&m!=="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(!Number.isInteger(t)||!Number.isInteger(e)||!Number.isInteger(s)||!Number.isInteger(l)||!Number.isInteger(u))throw new Error("m, n, incx, incy, and lda must be integers.");if(s<=0||l<=0)throw new Error("incx and incy must be positive.");if(!f&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(f&&u!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(f&&(i.rows<t||i.cols<e))throw new Error("A is too small for the given m and n.");(f?i.layout:m)==="column-major"&&([t,e]=[e,t],[o,n]=[n,o],[s,l]=[l,s]);let c=o instanceof I,w=n instanceof I;if(u<e)throw new Error("lda must be >= n.");if(!c&&!(o instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!w&&!(n instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(c!==w)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(c&&!f)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(f&&!c)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(f&&c&&i._buf===o._buf)throw new Error("A and x must not reference the same GPU buffer.");if(f&&w&&i._buf===n._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 f?{}:{A:i};if(!f&&i.length<(t-1)*u+e)throw new Error("A does not have enough elements for the given m, n, and lda.");if(o.length<(t-1)*s+1)throw new Error("x does not have enough elements for the given m and incx.");if(n.length<(e-1)*l+1)throw new Error("y does not have enough elements for the given n and incy.");let g=await G(r,"sger"),h=null,b=null,y=null,_=null;try{h=c?o._buf:x(r,o,"sger-x",!1),b=w?n._buf:x(r,n,"sger-y",!1),y=f?i._buf:x(r,i,"sger-A",!0),_=P(r,[{value:t,type:"u32"},{value:e,type:"u32"},{value:a,type:"f32"},{value:s,type:"u32"},{value:l,type:"u32"},{value:u,type:"u32"}],"sger-params");let v=E(r,g.getBindGroupLayout(0),[h,b,y,_]),A=Math.min(t,r.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:k,ts:B}=W(r,g,v,A),j=f?null:N(r,k,y);R(r,k);let D=await M(B);if(f)return D!==void 0?{gpuTimeMs:D}:{};let C=await S(j,Float32Array);return D!==void 0?{A:C,gpuTimeMs:D}:{A:C}}finally{!c&&h&&p(h),!w&&b&&p(b),!f&&y&&p(y),_&&p(_)}}async function xo(r,t,e,a,o,s,n,l,i="row-major"){let u=o instanceof I,m=n instanceof F;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"ssyr",{A:n,x:o}),t!=="lower"&&t!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(i!=="row-major"&&i!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(e)||!Number.isInteger(s)||!Number.isInteger(l))throw new Error("n, incx, and lda must be integers.");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(s<=0)throw new Error("incx must be positive.");if(l<e)throw new Error("lda must be >= n.");if(!m&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!u&&!(o instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(u&&!m)throw new Error("A must be a GpuMatrix when x is a GpuVector.");if(m&&!u)throw new Error("x must be a GpuVector when A is a GpuMatrix.");if(m&&u&&n._buf===o._buf)throw new Error("A and x must not reference the same GPU buffer.");if(m&&l!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(m&&(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 m?{}:{A:n};if(!m&&n.length<(e-1)*l+e)throw new Error("A does not have enough elements for the given n and lda.");if(o.length<(e-1)*s+1)throw new Error("x does not have enough elements for the given n and incx.");let d=(m?n.layout:i)==="column-major"?t==="upper":t==="lower",c=await G(r,"ssyr"),w=null,g=null,h=null;try{w=u?o._buf:x(r,o,"ssyr-x",!1),g=m?n._buf:x(r,n,"ssyr-A",!0),h=P(r,[{value:e,type:"u32"},{value:a,type:"f32"},{value:s,type:"u32"},{value:l,type:"u32"},{value:d?0:1,type:"u32"}],"ssyr-params");let b=E(r,c.getBindGroupLayout(0),[w,g,h]),y=Math.min(e,r.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:_,ts:v}=W(r,c,b,y),A=m?null:N(r,_,g);R(r,_);let k=await M(v);if(m)return k!==void 0?{gpuTimeMs:k}:{};let B=await S(A,Float32Array);return k!==void 0?{A:B,gpuTimeMs:k}:{A:B}}finally{!u&&w&&p(w),!m&&g&&p(g),h&&p(h)}}async function vo(r,t,e,a,o,s,n,l,i,u,m="row-major"){let f=o instanceof I,d=n instanceof I,c=i instanceof F;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"ssyr2",{A:i,x:o,y:n}),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(s)||!Number.isInteger(l)||!Number.isInteger(u))throw new Error("n, incx, incy, and lda must be integers.");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(s<=0||l<=0)throw new Error("incx and incy must be positive.");if(u<e)throw new Error("lda must be >= n.");if(!c&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!f&&!(o instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!d&&!(n instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(f!==d)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(f&&!c)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(c&&!f)throw new Error("x and y must be GpuVectors when A is a GpuMatrix.");if(c&&f&&i._buf===o._buf)throw new Error("A and x must not reference the same GPU buffer.");if(c&&d&&i._buf===n._buf)throw new Error("A and y must not reference the same GPU buffer.");if(f&&o._buf===n._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(c&&u!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(c&&(i.rows<e||i.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:i};if(!c&&i.length<(e-1)*u+e)throw new Error("A does not have enough elements for the given n and lda.");if(o.length<(e-1)*s+1)throw new Error("x does not have enough elements for the given n and incx.");if(n.length<(e-1)*l+1)throw new Error("y does not have enough elements for the given n and incy.");let g=(c?i.layout:m)==="column-major"?t==="upper":t==="lower",h=await G(r,"ssyr2"),b=null,y=null,_=null,v=null;try{b=f?o._buf:x(r,o,"ssyr2-x",!1),y=d?n._buf:x(r,n,"ssyr2-y",!1),_=c?i._buf:x(r,i,"ssyr2-A",!0),v=P(r,[{value:e,type:"u32"},{value:a,type:"f32"},{value:s,type:"u32"},{value:l,type:"u32"},{value:u,type:"u32"},{value:g?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:j}=W(r,h,A,k),D=c?null:N(r,B,_);R(r,B);let C=await M(j);if(c)return C!==void 0?{gpuTimeMs:C}:{};let L=await S(D,Float32Array);return C!==void 0?{A:L,gpuTimeMs:C}:{A:L}}finally{!f&&b&&p(b),!d&&y&&p(y),!c&&_&&p(_),v&&p(v)}}async function _o(r,t,e,a,o,s,n,l,i,u,m,f,d,c,w="row-major"){let g=l instanceof F,h=u instanceof F,b=d instanceof F;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"sgemm",{A:l,B:u,C:d}),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(w!=="row-major"&&w!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(typeof n!="number")throw new Error("alpha must be a number.");if(Number.isNaN(n))throw new Error("alpha must not be NaN.");if(!Number.isFinite(n))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(a)||!Number.isInteger(o)||!Number.isInteger(s)||!Number.isInteger(i)||!Number.isInteger(m)||!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&&!(d 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(a<0||o<0||s<0)throw new Error("m, n, and k must be non-negative.");if(a===0||o===0)return b?{}:{C:d};let y=g?l.layout:w,_=h?u.layout:w,v=b?d.layout:w,A=y==="column-major"?s:a,k=y==="column-major"?a:s,B=t==="no-transpose"?A:k,j=t==="no-transpose"?k:A;if(i<j)throw new Error(`lda must be >= ${y==="column-major"?"rows":"cols"} of A as stored.`);if(g){if(i!==l.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[Y,Q]=t==="no-transpose"?[a,s]:[s,a];if(l.rows<Y||l.cols<Q)throw new Error("A is too small for the given m, k, and transA.")}else if(l.length<(B-1)*i+j)throw new Error("A does not have enough elements for the given dimensions and lda.");let D=_==="column-major"?o:s,C=_==="column-major"?s:o,L=e==="no-transpose"?D:C,U=e==="no-transpose"?C:D;if(m<U)throw new Error(`ldb must be >= ${_==="column-major"?"rows":"cols"} of B as stored.`);if(h){if(m!==u.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");let[Y,Q]=e==="no-transpose"?[s,o]:[o,s];if(u.rows<Y||u.cols<Q)throw new Error("B is too small for the given n, k, and transB.")}else if(u.length<(L-1)*m+U)throw new Error("B does not have enough elements for the given dimensions and ldb.");let H=v==="column-major"?o:a,X=v==="column-major"?a:o;if(c<X)throw new Error(`ldc must be >= ${v==="column-major"?"rows":"cols"} of C as stored.`);if(b){if(c!==d.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(d.rows<a||d.cols<o)throw new Error("C is too small for the given m and n.")}else if(d.length<(H-1)*c+X)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"),_==="column-major"&&(e=e==="no-transpose"?"transpose":"no-transpose"),v==="column-major"&&([l,u]=[u,l],[g,h]=[h,g],[i,m]=[m,i],[t,e]=[e==="no-transpose"?"transpose":"no-transpose",t==="no-transpose"?"transpose":"no-transpose"],[a,o]=[o,a]);let q=Math.ceil(o/64),J=Math.ceil(a/64),Z=q*J>=36,rr=await G(r,Z?"sgemm_large":"sgemm_small"),lr=g?l._buf:x(r,l,"sgemm-A",!1),cr=h?u._buf:x(r,u,"sgemm-B",!1),pr=b?d._buf:x(r,d,"sgemm-C",!0),er=t==="no-transpose",or=e==="no-transpose",z=er&&we(lr,i,a,s),$=we(cr,m,or?s:o,or?o:s),K=P(r,[{value:a,type:"u32"},{value:o,type:"u32"},{value:s,type:"u32"},{value:n,type:"f32"},{value:f,type:"f32"},{value:i,type:"u32"},{value:m,type:"u32"},{value:c,type:"u32"},{value:t==="transpose"?1:0,type:"u32"},{value:e==="transpose"?1:0,type:"u32"},{value:z?1:0,type:"u32"},{value:$?1:0,type:"u32"}],"sgemm-params");try{let Y=E(r,rr.getBindGroupLayout(0),[lr,_r(r,lr),cr,_r(r,cr),pr,K]),Q=Z?{x:O(r,q,"sgemm","x"),y:O(r,J,"sgemm","y")}:{x:O(r,Math.ceil(o/32),"sgemm","x"),y:O(r,Math.ceil(a/32),"sgemm","y")},{commandEncoder:ir,ts:dr}=W(r,rr,Y,Q),ar=b?null:N(r,ir,pr);R(r,ir);let tr=await M(dr);if(b)return tr!==void 0?{gpuTimeMs:tr}:{};let wr=await S(ar,Float32Array);return tr!==void 0?{C:wr,gpuTimeMs:tr}:{C:wr}}finally{g||p(lr),h||p(cr),b||p(pr),p(K)}}async function Bo(r,t,e,a,o,s,n,l,i,u,m,f,d,c,w,g="row-major"){let h=i instanceof F,b=m instanceof F,y=c instanceof F;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"sgemmtr",{A:i,B:m,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(a!=="no-transpose"&&a!=="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 d!="number")throw new Error("beta must be a number.");if(Number.isNaN(d))throw new Error("beta must not be NaN.");if(!Number.isFinite(d))throw new Error("beta must be finite.");if(!Number.isInteger(o)||!Number.isInteger(s)||!Number.isInteger(n)||!Number.isInteger(u)||!Number.isInteger(f)||!Number.isInteger(w))throw new Error("m, n, k, lda, ldb, and ldc must be integers.");if(!h&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!b&&!(m 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(o<0||s<0||n<0)throw new Error("m, n, and k must be non-negative.");if(o===0||s===0)return y?{}:{C:c};let _=h?i.layout:g,v=b?m.layout:g,A=y?c.layout:g,k=_==="column-major"?n:o,B=_==="column-major"?o:n,j=e==="no-transpose"?k:B,D=e==="no-transpose"?B:k;if(u<D)throw new Error(`lda must be >= ${_==="column-major"?"rows":"cols"} of A as stored.`);if(h){if(u!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[z,$]=e==="no-transpose"?[o,n]:[n,o];if(i.rows<z||i.cols<$)throw new Error("A is too small for the given m, k, and transA.")}else if(i.length<(j-1)*u+D)throw new Error("A does not have enough elements for the given dimensions and lda.");let C=v==="column-major"?s:n,L=v==="column-major"?n:s,U=a==="no-transpose"?C:L,H=a==="no-transpose"?L:C;if(f<H)throw new Error(`ldb must be >= ${v==="column-major"?"rows":"cols"} of B as stored.`);if(b){if(f!==m.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");let[z,$]=a==="no-transpose"?[n,s]:[s,n];if(m.rows<z||m.cols<$)throw new Error("B is too small for the given n, k, and transB.")}else if(m.length<(U-1)*f+H)throw new Error("B does not have enough elements for the given dimensions and ldb.");let X=A==="column-major"?s:o,q=A==="column-major"?o:s;if(w<q)throw new Error(`ldc must be >= ${A==="column-major"?"rows":"cols"} of C as stored.`);if(y){if(w!==c.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(c.rows<o||c.cols<s)throw new Error("C is too small for the given m and n.")}else if(c.length<(X-1)*w+q)throw new Error("C does not have enough elements for the given dimensions and ldc.");_==="column-major"&&(e=e==="no-transpose"?"transpose":"no-transpose"),v==="column-major"&&(a=a==="no-transpose"?"transpose":"no-transpose"),A==="column-major"&&([i,m]=[m,i],[h,b]=[b,h],[u,f]=[f,u],[e,a]=[a==="no-transpose"?"transpose":"no-transpose",e==="no-transpose"?"transpose":"no-transpose"],[o,s]=[s,o],t=t==="lower"?"upper":"lower");let J=Math.ceil(s/64),Z=Math.ceil(o/64),rr=J*Z>=36,lr=await G(r,rr?"sgemmtr_large":"sgemmtr_small"),cr=h?i._buf:x(r,i,"sgemmtr-A",!1),pr=b?m._buf:x(r,m,"sgemmtr-B",!1),er=y?c._buf:x(r,c,"sgemmtr-C",!0),or=P(r,[{value:o,type:"u32"},{value:s,type:"u32"},{value:n,type:"u32"},{value:l,type:"f32"},{value:d,type:"f32"},{value:u,type:"u32"},{value:f,type:"u32"},{value:w,type:"u32"},{value:e==="transpose"?1:0,type:"u32"},{value:a==="transpose"?1:0,type:"u32"},{value:t==="upper"?1:0,type:"u32"}],"sgemmtr-params");try{let z=E(r,lr.getBindGroupLayout(0),[cr,pr,er,or]),$=rr?{x:O(r,J,"sgemmtr","x"),y:O(r,Z,"sgemmtr","y")}:{x:O(r,Math.ceil(s/32),"sgemmtr","x"),y:O(r,Math.ceil(o/32),"sgemmtr","y")},{commandEncoder:K,ts:Y}=W(r,lr,z,$),Q=y?null:N(r,K,er);R(r,K);let ir=await M(Y);if(y)return ir!==void 0?{gpuTimeMs:ir}:{};let dr=await S(Q,Float32Array);return ir!==void 0?{C:dr,gpuTimeMs:ir}:{C:dr}}finally{h||p(cr),b||p(pr),y||p(er),p(or)}}async function Eo(r,t,e,a,o,s,n,l,i,u,m,f="row-major"){let d=n instanceof F,c=u instanceof F;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"ssyrk",{A:n,C: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(f!=="row-major"&&f!=="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 i!="number")throw new Error("beta must be a number.");if(Number.isNaN(i))throw new Error("beta must not be NaN.");if(!Number.isFinite(i))throw new Error("beta must be finite.");if(!Number.isInteger(a)||!Number.isInteger(o)||!Number.isInteger(l)||!Number.isInteger(m))throw new Error("n, k, lda, and ldc must be integers.");if(!d&&!(n 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(d&&!c)throw new Error("C must be a GpuMatrix when A is a GpuMatrix.");if(c&&!d)throw new Error("A must be a GpuMatrix when C is a GpuMatrix.");if(a<0||o<0)throw new Error("n and k must be non-negative.");if(a===0)return c?{}:{C:u};let w=d?n.layout:f,g=c?u.layout:f,h=w==="column-major"?o:a,b=w==="column-major"?a:o,y=e==="no-transpose"?h:b,_=e==="no-transpose"?b:h;if(l<_)throw new Error(`lda must be >= ${w==="column-major"?"rows":"cols"} of A as stored.`);if(d){if(l!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[q,J]=e==="no-transpose"?[a,o]:[o,a];if(n.rows<q||n.cols<J)throw new Error("A is too small for the given n, k, and trans.")}else if(n.length<(y-1)*l+_)throw new Error("A does not have enough elements for the given dimensions and lda.");if(m<a)throw new Error("ldc must be >= n.");if(c){if(m!==u.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(u.rows<a||u.cols<a)throw new Error("C is too small for the given n.")}else if(u.length<(a-1)*m+a)throw new Error("C does not have enough elements for the given dimensions and ldc.");let v=e;w==="column-major"&&(v=v==="no-transpose"?"transpose":"no-transpose");let A=v==="no-transpose"?"transpose":"no-transpose",k=t;g==="column-major"&&([v,A]=[A==="no-transpose"?"transpose":"no-transpose",v==="no-transpose"?"transpose":"no-transpose"],k=k==="lower"?"upper":"lower");let B=Math.ceil(a/64),j=Math.ceil(a/64),D=B*j>=36,C=await G(r,D?"sgemmtr_large":"sgemmtr_small"),L=d?n._buf:x(r,n,"ssyrk-A",!1),U=c?u._buf:x(r,u,"ssyrk-C",!0),H=d?sr(r,L.size,"ssyrk-B",GPUBufferUsage.COPY_DST):x(r,n,"ssyrk-B",!1),X=P(r,[{value:a,type:"u32"},{value:a,type:"u32"},{value:o,type:"u32"},{value:s,type:"f32"},{value:i,type:"f32"},{value:l,type:"u32"},{value:l,type:"u32"},{value:m,type:"u32"},{value:v==="transpose"?1:0,type:"u32"},{value:A==="transpose"?1:0,type:"u32"},{value:k==="upper"?1:0,type:"u32"}],"ssyrk-params");try{let q=E(r,C.getBindGroupLayout(0),[L,H,U,X]),J=D?{x:O(r,B,"ssyrk","x"),y:O(r,j,"ssyrk","y")}:{x:O(r,Math.ceil(a/32),"ssyrk","x"),y:O(r,Math.ceil(a/32),"ssyrk","y")},{commandEncoder:Z,querySet:rr,passDescriptor:lr}=Mr(r);d&&Z.copyBufferToBuffer(L,0,H,0,L.size),ur(Z,C,q,J,lr);let cr=vr(r,Z,rr),pr=c?null:N(r,Z,U);R(r,Z);let er=await M(cr);if(c)return er!==void 0?{gpuTimeMs:er}:{};let or=await S(pr,Float32Array);return er!==void 0?{C:or,gpuTimeMs:er}:{C:or}}finally{d||p(L),p(H),c||p(U),p(X)}}async function Ao(r,t,e,a,o,s,n,l,i,u,m,f,d,c="row-major"){let w=n instanceof F,g=i instanceof F,h=f instanceof F;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"ssyr2k",{A:n,B:i,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(c!=="row-major"&&c!=="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(a)||!Number.isInteger(o)||!Number.isInteger(l)||!Number.isInteger(u)||!Number.isInteger(d))throw new Error("n, k, lda, ldb, and ldc must be integers.");if(!w&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!g&&!(i instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(!h&&!(f instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if((w||g)&&!h)throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");if(h&&(!w||!g))throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");if(a<0||o<0)throw new Error("n and k must be non-negative.");if(a===0)return h?{}:{C:f};let b=w?n.layout:c,y=g?i.layout:c,_=h?f.layout:c,v=b==="column-major"?o:a,A=b==="column-major"?a:o,k=e==="no-transpose"?v:A,B=e==="no-transpose"?A:v;if(l<B)throw new Error(`lda must be >= ${b==="column-major"?"rows":"cols"} of A as stored.`);if(w){if(l!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");let[Y,Q]=e==="no-transpose"?[a,o]:[o,a];if(n.rows<Y||n.cols<Q)throw new Error("A is too small for the given n, k, and trans.")}else if(n.length<(k-1)*l+B)throw new Error("A does not have enough elements for the given dimensions and lda.");let j=y==="column-major"?o:a,D=y==="column-major"?a:o,C=e==="no-transpose"?j:D,L=e==="no-transpose"?D:j;if(u<L)throw new Error(`ldb must be >= ${y==="column-major"?"rows":"cols"} of B as stored.`);if(g){if(u!==i.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");let[Y,Q]=e==="no-transpose"?[a,o]:[o,a];if(i.rows<Y||i.cols<Q)throw new Error("B is too small for the given n, k, and trans.")}else if(i.length<(C-1)*u+L)throw new Error("B does not have enough elements for the given dimensions and ldb.");if(d<a)throw new Error("ldc must be >= n.");if(h){if(d!==f.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(f.rows<a||f.cols<a)throw new Error("C is too small for the given n.")}else if(f.length<(a-1)*d+a)throw new Error("C does not have enough elements for the given dimensions and ldc.");let U=e;b==="column-major"&&(U=U==="no-transpose"?"transpose":"no-transpose");let H=e;y==="column-major"&&(H=H==="no-transpose"?"transpose":"no-transpose");let X=_==="column-major"?t==="lower"?"upper":"lower":t,q=Y=>Y==="no-transpose"?"transpose":"no-transpose";function J(Y,Q,ir,dr,ar,tr){let wr=Y,Rr=q(dr);return _!=="column-major"?{transX:wr,X:Q,ldX:ir,transY:Rr,Y:ar,ldY:tr}:{transX:q(Rr),X:ar,ldX:tr,transY:q(wr),Y:Q,ldY:ir}}let Z=Math.ceil(a/64),rr=Math.ceil(a/64),lr=Z*rr>=36,cr=await G(r,lr?"sgemmtr_large":"sgemmtr_small"),pr=lr?{x:O(r,Z,"ssyr2k","x"),y:O(r,rr,"ssyr2k","y")}:{x:O(r,Math.ceil(a/32),"ssyr2k","x"),y:O(r,Math.ceil(a/32),"ssyr2k","y")},er=w?n._buf:x(r,n,"ssyr2k-A",!1),or=g?i._buf:x(r,i,"ssyr2k-B",!1),z=h?f._buf:x(r,f,"ssyr2k-C",!0),$=null,K=null;try{let Y=J(U,er,l,H,or,u),Q=J(H,or,u,U,er,l),ir=(Pr,hr)=>P(r,[{value:a,type:"u32"},{value:a,type:"u32"},{value:o,type:"u32"},{value:s,type:"f32"},{value:hr,type:"f32"},{value:Pr.ldX,type:"u32"},{value:Pr.ldY,type:"u32"},{value:d,type:"u32"},{value:Pr.transX==="transpose"?1:0,type:"u32"},{value:Pr.transY==="transpose"?1:0,type:"u32"},{value:X==="upper"?1:0,type:"u32"}],"ssyr2k-params");$=ir(Y,m),K=ir(Q,1);let dr=E(r,cr.getBindGroupLayout(0),[Y.X,Y.Y,z,$]),ar=E(r,cr.getBindGroupLayout(0),[Q.X,Q.Y,z,K]),{commandEncoder:tr,querySet:wr}=Mr(r),Rr=wr?{timestampWrites:{querySet:wr,beginningOfPassWriteIndex:0}}:void 0,kr=wr?{timestampWrites:{querySet:wr,endOfPassWriteIndex:1}}:void 0;ur(tr,cr,dr,pr,Rr),ur(tr,cr,ar,pr,kr);let Ir=vr(r,tr,wr),xr=h?null:N(r,tr,z);R(r,tr);let br=await M(Ir);if(h)return br!==void 0?{gpuTimeMs:br}:{};let gr=await S(xr,Float32Array);return br!==void 0?{C:gr,gpuTimeMs:br}:{C:gr}}finally{w||p(er),g||p(or),h||p(z),$&&p($),K&&p(K)}}async function Go(r,t,e,a,o,s,n,l,i,u,m,f,d,c="row-major"){let w=n instanceof F,g=i instanceof F,h=f instanceof F;if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"ssymm",{A:n,B:i,C:f}),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 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(a)||!Number.isInteger(o)||!Number.isInteger(l)||!Number.isInteger(u)||!Number.isInteger(d))throw new Error("m, n, lda, ldb, and ldc must be integers.");if(!w&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!g&&!(i instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(!h&&!(f instanceof Float32Array))throw new Error("C must be a Float32Array or GpuMatrix.");if((w||g)&&!h)throw new Error("C must be a GpuMatrix when A or B is a GpuMatrix.");if(h&&(!w||!g))throw new Error("A and B must be GpuMatrix when C is a GpuMatrix.");if(a<0||o<0)throw new Error("m and n must be non-negative.");if(a===0||o===0)return h?{}:{C:f};let b=w?n.layout:c,y=g?i.layout:c,_=h?f.layout:c,v=t==="left"?a:o;if(l<v)throw new Error("lda must be >= "+(t==="left"?"m":"n")+".");if(w){if(l!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(n.rows<v||n.cols<v)throw new Error("A is too small for the given m/n and side.")}else if(n.length<(v-1)*l+v)throw new Error("A does not have enough elements for the given dimensions and lda.");let A=y==="column-major"?o:a,k=y==="column-major"?a:o;if(u<k)throw new Error(`ldb must be >= ${y==="column-major"?"rows":"cols"} of B as stored.`);if(g){if(u!==i.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");if(i.rows<a||i.cols<o)throw new Error("B is too small for the given m and n.")}else if(i.length<(A-1)*u+k)throw new Error("B does not have enough elements for the given dimensions and ldb.");let B=_==="column-major"?o:a,j=_==="column-major"?a:o;if(d<j)throw new Error(`ldc must be >= ${_==="column-major"?"rows":"cols"} of C as stored.`);if(h){if(d!==f.lda)throw new Error("ldc must match C.lda when C is a GpuMatrix.");if(f.rows<a||f.cols<o)throw new Error("C is too small for the given m and n.")}else if(f.length<(B-1)*d+j)throw new Error("C does not have enough elements for the given dimensions and ldc.");let D=b==="column-major"?e==="lower"?"upper":"lower":e,C=y==="column-major"?"transpose":"no-transpose",L="no-transpose",U=a,H=o,X=v,q=t==="left"?L:C,J=t==="left"?C:L,Z=tr=>tr==="no-transpose"?"transpose":"no-transpose",rr=t==="right";_==="column-major"&&([q,J]=[Z(J),Z(q)],rr=!rr,[U,H]=[H,U]);let lr=v,cr=Math.ceil(H/64),pr=Math.ceil(U/64),er=cr*pr>=36,or=await G(r,er?"sgemm_large":"sgemm_small"),z=await G(r,"symmetrize"),$=er?{x:O(r,cr,"ssymm","x"),y:O(r,pr,"ssymm","y")}:{x:O(r,Math.ceil(H/32),"ssymm","x"),y:O(r,Math.ceil(U/32),"ssymm","y")},K=w?n._buf:x(r,n,"ssymm-A",!1),Y=g?i._buf:x(r,i,"ssymm-B",!1),Q=h?f._buf:x(r,f,"ssymm-C",!0),ir=sr(r,v*lr*4,"ssymm-Adense"),dr=null,ar=null;try{dr=P(r,[{value:v,type:"u32"},{value:l,type:"u32"},{value:lr,type:"u32"},{value:D==="upper"?1:0,type:"u32"}],"ssymm-sym-params");let tr=E(r,z.getBindGroupLayout(0),[K,ir,dr]),wr=rr?Y:ir,Rr=rr?u:lr,kr=rr?ir:Y;ar=P(r,[{value:U,type:"u32"},{value:H,type:"u32"},{value:X,type:"u32"},{value:s,type:"f32"},{value:m,type:"f32"},{value:Rr,type:"u32"},{value:rr?lr:u,type:"u32"},{value:d,type:"u32"},{value:q==="transpose"?1:0,type:"u32"},{value:J==="transpose"?1:0,type:"u32"}],"ssymm-gemm-params");let xr=E(r,or.getBindGroupLayout(0),[wr,_r(r,wr),kr,_r(r,kr),Q,ar]),{commandEncoder:br,querySet:gr}=Mr(r),Pr=gr?{timestampWrites:{querySet:gr,beginningOfPassWriteIndex:0}}:void 0,hr=gr?{timestampWrites:{querySet:gr,endOfPassWriteIndex:1}}:void 0;ur(br,z,tr,{x:Math.ceil(v/8),y:Math.ceil(v/8)},Pr),ur(br,or,xr,$,hr);let jr=vr(r,br,gr),Lr=h?null:N(r,br,Q);R(r,br);let Ur=await M(jr);if(h)return Ur!==void 0?{gpuTimeMs:Ur}:{};let te=await S(Lr,Float32Array);return Ur!==void 0?{C:te,gpuTimeMs:Ur}:{C:te}}finally{w||p(K),g||p(Y),h||p(Q),p(ir),dr&&p(dr),ar&&p(ar)}}async function So(r,t,e,a,o,s,n,l,i,u,m,f,d="row-major"){let c=i instanceof F,w=m instanceof F,g=o==="unit";if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"strmm",{A:i,B: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(a!=="no-transpose"&&a!=="transpose")throw new Error("transA must be 'no-transpose' or 'transpose'.");if(!g&&o!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(d!=="row-major"&&d!=="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(s)||!Number.isInteger(n)||!Number.isInteger(u)||!Number.isInteger(f))throw new Error("m, n, lda, and ldb must be integers.");if(!c&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!w&&!(m instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(c!==w)throw new Error("A and B must both be GpuMatrix or both be Float32Array.");if(s<0||n<0)throw new Error("m and n must be non-negative.");if(s===0||n===0)return w?{}:{B:m};let h=c?i.layout:d,b=w?m.layout:d,y=t==="left"?s:n;if(u<y)throw new Error("lda must be >= "+(t==="left"?"m":"n")+".");if(c){if(u!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(i.rows<y||i.cols<y)throw new Error("A is too small for the given m/n and side.")}else if(i.length<(y-1)*u+y)throw new Error("A does not have enough elements for the given dimensions and lda.");let _=b==="column-major"?n:s,v=b==="column-major"?s:n;if(f<v)throw new Error(`ldb must be >= ${b==="column-major"?"rows":"cols"} of B as stored.`);if(w){if(f!==m.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");if(m.rows<s||m.cols<n)throw new Error("B is too small for the given m and n.")}else if(m.length<(_-1)*f+v)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"?a==="no-transpose"?"transpose":"no-transpose":a,B=b==="column-major"?"transpose":"no-transpose",j="no-transpose",D=s,C=n,L=y,U=t==="left"?j:B,H=t==="left"?B:j,X=dr=>dr==="no-transpose"?"transpose":"no-transpose",q=t==="right";b==="column-major"&&([U,H]=[X(H),X(U)],q=!q,[D,C]=[C,D]);let J=y,Z=Math.ceil(C/64),rr=Math.ceil(D/64),lr=Z*rr>=36,cr=await G(r,lr?"sgemm_large":"sgemm_small"),pr=await G(r,"triangularize"),er=lr?{x:O(r,Z,"strmm","x"),y:O(r,rr,"strmm","y")}:{x:O(r,Math.ceil(C/32),"strmm","x"),y:O(r,Math.ceil(D/32),"strmm","y")},or=null,z=null,$=null,K=null,Y=null,Q=null,ir=!1;try{or=c?i._buf:x(r,i,"strmm-A",!1),z=w?m._buf:x(r,m,"strmm-B",!0),$=sr(r,y*J*4,"strmm-Adense"),K=sr(r,_*f*4,"strmm-out",GPUBufferUsage.COPY_SRC|GPUBufferUsage.COPY_DST),Y=P(r,[{value:y,type:"u32"},{value:u,type:"u32"},{value:J,type:"u32"},{value:A==="upper"?1:0,type:"u32"},{value:k==="transpose"?1:0,type:"u32"},{value:g?1:0,type:"u32"}],"strmm-tri-params");let dr=E(r,pr.getBindGroupLayout(0),[or,$,Y]),ar=q?z:$,tr=q?f:J,wr=q?$:z;Q=P(r,[{value:D,type:"u32"},{value:C,type:"u32"},{value:L,type:"u32"},{value:l,type:"f32"},{value:0,type:"f32"},{value:tr,type:"u32"},{value:q?J:f,type:"u32"},{value:f,type:"u32"},{value:U==="transpose"?1:0,type:"u32"},{value:H==="transpose"?1:0,type:"u32"}],"strmm-gemm-params");let kr=E(r,cr.getBindGroupLayout(0),[ar,_r(r,ar),wr,_r(r,wr),K,Q]),{commandEncoder:Ir,querySet:xr}=Mr(r);Ir.copyBufferToBuffer(z,0,K,0,Math.min(z.size,K.size));let br=xr?{timestampWrites:{querySet:xr,beginningOfPassWriteIndex:0}}:void 0,gr=xr?{timestampWrites:{querySet:xr,endOfPassWriteIndex:1}}:void 0;ur(Ir,pr,dr,{x:Math.ceil(y/8),y:Math.ceil(y/8)},br),ur(Ir,cr,kr,er,gr);let Pr=vr(r,Ir,xr),hr=w?null:N(r,Ir,K);R(r,Ir);let jr=await M(Pr);if(w)return p(m._buf),m._buf=K,ir=!0,jr!==void 0?{gpuTimeMs:jr}:{};let Lr=await S(hr,Float32Array);return jr!==void 0?{B:Lr,gpuTimeMs:jr}:{B:Lr}}finally{!c&&or&&p(or),!w&&z&&p(z),$&&p($),K&&!ir&&p(K),Y&&p(Y),Q&&p(Q)}}async function ko(r,t,e,a,o,s,n,l,i,u,m,f,d="row-major"){let c=i instanceof F,w=m instanceof F,g=o==="unit";if(!(r instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(T(r,"strsm",{A:i,B: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(a!=="no-transpose"&&a!=="transpose")throw new Error("transA must be 'no-transpose' or 'transpose'.");if(!g&&o!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(d!=="row-major"&&d!=="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(s)||!Number.isInteger(n)||!Number.isInteger(u)||!Number.isInteger(f))throw new Error("m, n, lda, and ldb must be integers.");if(!c&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!w&&!(m instanceof Float32Array))throw new Error("B must be a Float32Array or GpuMatrix.");if(c!==w)throw new Error("A and B must both be GpuMatrix or both be Float32Array.");if(s<0||n<0)throw new Error("m and n must be non-negative.");if(s===0||n===0)return w?{}:{B:m};let h=c?i.layout:d,b=w?m.layout:d,y=t==="left"?s:n;if(u<y)throw new Error("lda must be >= "+(t==="left"?"m":"n")+".");if(c){if(u!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(i.rows<y||i.cols<y)throw new Error("A is too small for the given m/n and side.")}else if(i.length<(y-1)*u+y)throw new Error("A does not have enough elements for the given dimensions and lda.");let _=b==="column-major"?n:s,v=b==="column-major"?s:n;if(f<v)throw new Error(`ldb must be >= ${b==="column-major"?"rows":"cols"} of B as stored.`);if(w){if(f!==m.lda)throw new Error("ldb must match B.lda when B is a GpuMatrix.");if(m.rows<s||m.cols<n)throw new Error("B is too small for the given m and n.")}else if(m.length<(_-1)*f+v)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"?a==="no-transpose"?"transpose":"no-transpose":a,B=t==="left"?n:s,j=t==="left",D=k==="no-transpose"==(A==="lower"),C=t==="left"?D:!D,L=[];for(let z=0;z<y;z+=64)L.push(z);C||L.reverse();let U=L.length,H=await G(r,"strsv_invert_block"),X=await G(r,"block_transfer"),q=await G(r,"sscal"),J=null,Z=null,rr=null,lr=[],cr=[];function pr(z,$){let K=sr(r,z,$);return cr.push(K),K}function er(z,$){let K=P(r,z,$);return lr.push(K),K}let or=(_-1)*f+v;try{J=c?i._buf:x(r,i,"strsm-A",!1),Z=w?m._buf:x(r,m,"strsm-B",!0),rr=sr(r,U*64*64*4,"strsm-Ainv");let z=null;if(l!==1){let xr=er([{value:or,type:"u32"},{value:l,type:"f32"},{value:1,type:"u32"}],"strsm-scale-params");z=E(r,q.getBindGroupLayout(0),[Z,xr])}let $=er([{value:y,type:"u32"},{value:u,type:"u32"},{value:k==="transpose"?1:0,type:"u32"},{value:A==="upper"?1:0,type:"u32"},{value:g?1:0,type:"u32"}],"strsm-invert-params"),K=E(r,H.getBindGroupLayout(0),[J,rr,$]),Y=pr(64*B*4,"strsm-Bblock"),Q=pr(64*B*4,"strsm-Xblock"),ir=pr(y*64*4,"strsm-Aoff"),dr=pr(y*B*4,"strsm-delta"),{commandEncoder:ar,querySet:tr}=Mr(r);if(l===0){let xr=tr?{timestampWrites:{querySet:tr,beginningOfPassWriteIndex:0,endOfPassWriteIndex:1}}:void 0;ur(ar,q,z,yr(r,or),xr)}else{z&&ur(ar,q,z,yr(r,or)),ur(ar,H,K,{x:64,y:U},tr?{timestampWrites:{querySet:tr,beginningOfPassWriteIndex:0}}:void 0);for(let br=0;br<L.length;br++){let gr=L[br],Pr=Math.min(gr+64,y),hr=Pr-gr,jr=gr/64,Lr=br===L.length-1,Ur=er([{value:gr,type:"u32"},{value:hr,type:"u32"},{value:0,type:"u32"},{value:B,type:"u32"},{value:f,type:"u32"},{value:b==="column-major"?1:0,type:"u32"},{value:j?1:0,type:"u32"},{value:2,type:"u32"}],"strsm-gather-B-params"),te=E(r,X.getBindGroupLayout(0),[Y,Z,Ur]);ur(ar,X,te,qr(r,"strsm",hr,B));{let Wr=hr,Fr=B,me=hr,Xr=Math.ceil(Fr/64),$r=Math.ceil(Wr/64),Zr=Xr*$r>=36,Qr=await G(r,Zr?"sgemm_large":"sgemm_small"),fe=er([{value:Wr,type:"u32"},{value:Fr,type:"u32"},{value:me,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"),ae={buffer:rr,offset:jr*64*64*4,size:4096*4},ce=E(r,Qr.getBindGroupLayout(0),[ae,_r(r,ae),Y,_r(r,Y),Q,fe]),Co=Zr?{x:O(r,Xr,"strsm","x"),y:O(r,$r,"strsm","y")}:{x:O(r,Math.ceil(Fr/32),"strsm","x"),y:O(r,Math.ceil(Wr/32),"strsm","y")};ur(ar,Qr,ce,Co)}let oe=C?Pr:0,_e=C?y:gr,Be=oe<_e,No=er([{value:gr,type:"u32"},{value:hr,type:"u32"},{value:0,type:"u32"},{value:B,type:"u32"},{value:f,type:"u32"},{value:b==="column-major"?1:0,type:"u32"},{value:j?1:0,type:"u32"},{value:0,type:"u32"}],"strsm-scatter-params"),Mo=E(r,X.getBindGroupLayout(0),[Q,Z,No]),Io=Lr&&!Be&&tr?{timestampWrites:{querySet:tr,endOfPassWriteIndex:1}}:void 0;if(ur(ar,X,Mo,qr(r,"strsm",hr,B),Io),!Be)continue;let Yr=_e-oe,Ro=er([{value:oe,type:"u32"},{value:Yr,type:"u32"},{value:gr,type:"u32"},{value:hr,type:"u32"},{value:u,type:"u32"},{value:k==="transpose"?1:0,type:"u32"},{value:j?1:0,type:"u32"},{value:2,type:"u32"}],"strsm-gather-A-params"),Po=E(r,X.getBindGroupLayout(0),[ir,J,Ro]);ur(ar,X,Po,qr(r,"strsm",Yr,hr));{let Wr=Yr,Fr=B,me=hr,Xr=Math.ceil(Fr/64),$r=Math.ceil(Wr/64),Zr=Xr*$r>=36,Qr=await G(r,Zr?"sgemm_large":"sgemm_small"),fe=er([{value:Wr,type:"u32"},{value:Fr,type:"u32"},{value:me,type:"u32"},{value:1,type:"f32"},{value:0,type:"f32"},{value:hr,type:"u32"},{value:B,type:"u32"},{value:B,type:"u32"},{value:0,type:"u32"},{value:0,type:"u32"}],"strsm-update-params"),ae=E(r,Qr.getBindGroupLayout(0),[ir,_r(r,ir),Q,_r(r,Q),dr,fe]),ce=Zr?{x:O(r,Xr,"strsm","x"),y:O(r,$r,"strsm","y")}:{x:O(r,Math.ceil(Fr/32),"strsm","x"),y:O(r,Math.ceil(Wr/32),"strsm","y")};ur(ar,Qr,ae,ce)}let Do=er([{value:oe,type:"u32"},{value:Yr,type:"u32"},{value:0,type:"u32"},{value:B,type:"u32"},{value:f,type:"u32"},{value:b==="column-major"?1:0,type:"u32"},{value:j?1:0,type:"u32"},{value:1,type:"u32"}],"strsm-scatter-sub-params"),To=E(r,X.getBindGroupLayout(0),[dr,Z,Do]),jo=Lr&&tr?{timestampWrites:{querySet:tr,endOfPassWriteIndex:1}}:void 0;ur(ar,X,To,qr(r,"strsm",Yr,B),jo)}}let wr=vr(r,ar,tr),Rr=w?null:N(r,ar,Z);R(r,ar);let kr=await M(wr);if(w)return kr!==void 0?{gpuTimeMs:kr}:{};let Ir=await S(Rr,Float32Array);return kr!==void 0?{B:Ir,gpuTimeMs:kr}:{B:Ir}}finally{!c&&J&&p(J),!w&&Z&&p(Z),rr&&p(rr),p(cr),p(lr)}}return Oo(Ta);})();
|