wgblas 1.2.1 → 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/LICENSE +1 -1
- package/README.md +54 -72
- package/dist/wgblas.browser.js +1637 -858
- package/index.d.mts +51 -6
- package/index.mjs +8 -0
- package/package.json +56 -2
- package/src/classes/GpuMatrix.mjs +17 -10
- package/src/classes/GpuVector.mjs +28 -10
- package/src/dasum/dasum.d.mts +6 -6
- package/src/dasum/dasum.mjs +25 -21
- package/src/devdocs.mjs +13 -0
- package/src/idamax/idamax.d.mts +69 -0
- package/src/idamax/idamax.mjs +130 -0
- package/src/init.mjs +115 -49
- package/src/isamax/isamax.d.mts +21 -3
- package/src/isamax/isamax.mjs +16 -14
- package/src/random/random.d.mts +1 -0
- package/src/sasum/sasum.d.mts +3 -3
- package/src/sasum/sasum.mjs +15 -13
- package/src/saxpy/saxpy.d.mts +3 -3
- package/src/saxpy/saxpy.mjs +10 -8
- package/src/scopy/scopy.d.mts +3 -3
- package/src/scopy/scopy.mjs +10 -8
- package/src/sdot/sdot.d.mts +3 -3
- package/src/sdot/sdot.mjs +16 -14
- package/src/sgemm/sgemm.d.mts +102 -0
- package/src/sgemm/sgemm.mjs +208 -0
- package/src/sgemmtr/sgemmtr.d.mts +104 -0
- package/src/sgemmtr/sgemmtr.mjs +204 -0
- package/src/sgemv/sgemv.d.mts +1 -38
- package/src/sgemv/sgemv.mjs +42 -26
- package/src/sger/sger.d.mts +1 -34
- package/src/sger/sger.mjs +12 -8
- package/src/shaders/block_transfer.wgsl +42 -0
- package/src/shaders/dasum.wgsl +3 -2
- package/src/shaders/f64/dekker.wgsl +4 -85
- package/src/shaders/f64/utils/abs.wgsl +10 -0
- package/src/shaders/f64/utils/add.wgsl +77 -0
- package/src/shaders/f64/utils/equal.wgsl +7 -0
- package/src/shaders/f64/utils/greater.wgsl +12 -0
- package/src/shaders/f64/utils/multiply.wgsl +81 -0
- package/src/shaders/idamax.wgsl +96 -0
- package/src/shaders/index.mjs +164 -14
- package/src/shaders/reduction/argmaxF64.wgsl +50 -0
- package/src/shaders/reduction/scaledSum.wgsl +65 -0
- package/src/shaders/reduction/sumF64.wgsl +2 -2
- package/src/shaders/sgemm_large.wgsl +206 -0
- package/src/shaders/sgemm_small.wgsl +212 -0
- package/src/shaders/sgemmtr_large.wgsl +120 -0
- package/src/shaders/sgemmtr_small.wgsl +113 -0
- 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/shaders/symmetrize.wgsl +31 -0
- package/src/shaders/triangularize.wgsl +44 -0
- package/src/snrm2/snrm2.d.mts +3 -3
- package/src/snrm2/snrm2.mjs +33 -21
- package/src/srot/srot.d.mts +3 -5
- package/src/srot/srot.mjs +11 -9
- package/src/srotm/srotm.d.mts +3 -5
- package/src/srotm/srotm.mjs +19 -10
- package/src/sscal/sscal.d.mts +4 -4
- package/src/sscal/sscal.mjs +12 -10
- package/src/sswap/sswap.d.mts +3 -3
- package/src/sswap/sswap.mjs +11 -9
- package/src/ssymm/ssymm.d.mts +103 -0
- package/src/ssymm/ssymm.mjs +218 -0
- package/src/ssymv/ssymv.d.mts +1 -36
- package/src/ssymv/ssymv.mjs +12 -8
- package/src/ssyr/ssyr.d.mts +1 -30
- package/src/ssyr/ssyr.mjs +11 -7
- package/src/ssyr2/ssyr2.d.mts +1 -34
- package/src/ssyr2/ssyr2.mjs +12 -8
- package/src/ssyr2k/ssyr2k.d.mts +100 -0
- package/src/ssyr2k/ssyr2k.mjs +202 -0
- package/src/ssyrk/ssyrk.d.mts +90 -0
- package/src/ssyrk/ssyrk.mjs +177 -0
- package/src/strmm/strmm.d.mts +100 -0
- package/src/strmm/strmm.mjs +226 -0
- package/src/strmv/strmv.d.mts +1 -36
- package/src/strmv/strmv.mjs +12 -8
- package/src/strsm/strsm.d.mts +99 -0
- package/src/strsm/strsm.mjs +360 -0
- package/src/strsv/strsv.d.mts +1 -32
- package/src/strsv/strsv.mjs +18 -12
- package/src/util/benchmark.mjs +4 -6
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +116 -20
- package/src/util/compute.mjs +12 -12
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +34 -0
- package/src/util/f64.mjs +3 -3
- package/src/util/pipeline.mjs +5 -6
- package/src/util/workgroup.mjs +55 -7
- package/src/shaders/browser-shaders.mjs +0 -55
package/dist/wgblas.browser.js
CHANGED
|
@@ -1,113 +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 Wr,Nr=W(()=>{Wr=`// sum reduction: collapses 2*WGS partials into one scalar.
|
|
45
|
-
// dispatch: 1 workgroup of WGS threads.
|
|
46
|
-
// partials must have exactly 2*WGS entries.
|
|
47
|
-
|
|
48
|
-
@group(0) @binding(0) var<storage, read> partials: array<f32>;
|
|
49
|
-
@group(0) @binding(1) var<storage, read_write> result: array<f32>;
|
|
50
|
-
|
|
51
|
-
const WGS: u32 = 64;
|
|
52
|
-
|
|
53
|
-
var<workgroup> tile: array<f32, 64>;
|
|
54
|
-
|
|
55
|
-
@compute @workgroup_size(64)
|
|
56
|
-
fn reduce(
|
|
57
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
58
|
-
) {
|
|
59
|
-
let i = lid.x;
|
|
60
|
-
tile[i] = partials[i] + partials[i + WGS];
|
|
61
|
-
workgroupBarrier();
|
|
62
|
-
|
|
63
|
-
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
64
|
-
if (i < s) { tile[i] += tile[i + s]; }
|
|
65
|
-
workgroupBarrier();
|
|
66
|
-
}
|
|
67
|
-
|
|
68
|
-
if (i == 0u) { result[0] = tile[0]; }
|
|
69
|
-
}
|
|
70
|
-
`});var Mr,Dr=W(()=>{Mr=`// sum reduction (f64, double-double): collapses 2*WGS partial (hi, lo) pairs
|
|
71
|
-
// into one, using ddAddProtected instead of plain f32 \`+\` (see
|
|
72
|
-
// reduction/sum.wgsl for the f32 original this mirrors).
|
|
73
|
-
// dispatch: 1 workgroup of WGS threads. partialsHi/partialsLo must have
|
|
74
|
-
// exactly 2*WGS entries each. Concatenated after f64/dekker.wgsl for
|
|
75
|
-
// DD/ddAddProtected (see it for why plain ddAdd isn't safe).
|
|
76
|
-
|
|
77
|
-
@group(0) @binding(0) var<storage, read> partialsHi: array<f32>;
|
|
78
|
-
@group(0) @binding(1) var<storage, read> partialsLo: array<f32>;
|
|
79
|
-
@group(0) @binding(2) var<storage, read_write> resultHi: array<f32, 1>;
|
|
80
|
-
@group(0) @binding(3) var<storage, read_write> resultLo: array<f32, 1>;
|
|
81
|
-
|
|
82
|
-
const WGS: u32 = 64;
|
|
83
|
-
|
|
84
|
-
var<workgroup> tile: array<DD, 64>;
|
|
85
|
-
|
|
86
|
-
@compute @workgroup_size(64)
|
|
87
|
-
fn reduce_f64(
|
|
88
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
89
|
-
) {
|
|
90
|
-
let i = lid.x;
|
|
91
|
-
let a = DD(partialsHi[i], partialsLo[i]);
|
|
92
|
-
let b = DD(partialsHi[i + WGS], partialsLo[i + WGS]);
|
|
93
|
-
tile[i] = ddAddProtected(a, b, i);
|
|
94
|
-
workgroupBarrier();
|
|
95
|
-
|
|
96
|
-
// ddAddProtected must be called unconditionally by every thread.
|
|
97
|
-
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
98
|
-
let partner = select(i, i + s, i < s);
|
|
99
|
-
let combined = ddAddProtected(tile[i], tile[partner], i);
|
|
100
|
-
workgroupBarrier();
|
|
101
|
-
if (i < s) { tile[i] = combined; }
|
|
102
|
-
workgroupBarrier();
|
|
103
|
-
}
|
|
104
|
-
|
|
105
|
-
if (i == 0u) {
|
|
106
|
-
resultHi[0] = tile[0].hi;
|
|
107
|
-
resultLo[0] = tile[0].lo;
|
|
108
|
-
}
|
|
109
|
-
}
|
|
110
|
-
`});var Ur,Tr=W(()=>{Ur=`// 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
|
|
111
2
|
|
|
112
3
|
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
113
4
|
|
|
@@ -130,7 +21,7 @@ fn main(
|
|
|
130
21
|
x[id * params.x_inc] = params.alpha * x[id * params.x_inc];
|
|
131
22
|
}
|
|
132
23
|
}
|
|
133
|
-
`});var
|
|
24
|
+
`});var We,Le=V(()=>{We=`// sswap: x <-> y
|
|
134
25
|
|
|
135
26
|
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
136
27
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
@@ -156,7 +47,7 @@ fn main(
|
|
|
156
47
|
y[id * params.y_inc] = temp;
|
|
157
48
|
}
|
|
158
49
|
}
|
|
159
|
-
`});var
|
|
50
|
+
`});var qe,Fe=V(()=>{qe=`// saxpy: y = alpha * x + y
|
|
160
51
|
|
|
161
52
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
162
53
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
@@ -181,7 +72,7 @@ fn main(
|
|
|
181
72
|
y[id * params.y_inc] = params.alpha * x[id * params.x_inc] + y[id * params.y_inc];
|
|
182
73
|
}
|
|
183
74
|
}
|
|
184
|
-
`});var
|
|
75
|
+
`});var Oe,Ue=V(()=>{Oe=`// scopy: y = x
|
|
185
76
|
|
|
186
77
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
187
78
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
@@ -205,7 +96,7 @@ fn main(
|
|
|
205
96
|
y[id * params.y_inc] = x[id * params.x_inc];
|
|
206
97
|
}
|
|
207
98
|
}
|
|
208
|
-
`});var
|
|
99
|
+
`});var Ve,Ke=V(()=>{Ve=`// sdot: result = sum(x[i] * y[i])
|
|
209
100
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sum.wgsl.
|
|
210
101
|
|
|
211
102
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -258,7 +149,33 @@ fn main(
|
|
|
258
149
|
|
|
259
150
|
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
260
151
|
}
|
|
261
|
-
`});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]|)
|
|
262
179
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/abssum.wgsl.
|
|
263
180
|
|
|
264
181
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -309,12 +226,25 @@ fn main(
|
|
|
309
226
|
|
|
310
227
|
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
311
228
|
}
|
|
312
|
-
`});var $
|
|
313
|
-
//
|
|
314
|
-
|
|
315
|
-
|
|
316
|
-
|
|
317
|
-
|
|
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;
|
|
318
248
|
|
|
319
249
|
struct Params {
|
|
320
250
|
n: u32,
|
|
@@ -323,7 +253,36 @@ struct Params {
|
|
|
323
253
|
|
|
324
254
|
const WGS: u32 = 64;
|
|
325
255
|
|
|
326
|
-
|
|
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>;
|
|
327
286
|
|
|
328
287
|
@compute @workgroup_size(64)
|
|
329
288
|
fn main(
|
|
@@ -332,119 +291,112 @@ fn main(
|
|
|
332
291
|
@builtin(workgroup_id) wgid: vec3u,
|
|
333
292
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
334
293
|
) {
|
|
335
|
-
var acc0
|
|
336
|
-
var acc1
|
|
337
|
-
var acc2
|
|
338
|
-
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);
|
|
339
298
|
|
|
340
299
|
let stride = num_wg.x * WGS;
|
|
341
300
|
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
342
301
|
|
|
343
302
|
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
344
|
-
|
|
345
|
-
|
|
346
|
-
|
|
347
|
-
|
|
348
|
-
acc0 += v0 * v0;
|
|
349
|
-
acc1 += v1 * v1;
|
|
350
|
-
acc2 += v2 * v2;
|
|
351
|
-
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]));
|
|
352
307
|
}
|
|
353
308
|
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
354
|
-
|
|
355
|
-
acc0 += v * v;
|
|
309
|
+
acc0 = ssqAccum(acc0, abs(x[id * params.x_inc]));
|
|
356
310
|
}
|
|
357
311
|
|
|
358
|
-
|
|
312
|
+
let combined = ssqMerge(ssqMerge(acc0, acc1), ssqMerge(acc2, acc3));
|
|
313
|
+
tileScale[lid.x] = combined.scale;
|
|
314
|
+
tileSsq[lid.x] = combined.ssq;
|
|
359
315
|
workgroupBarrier();
|
|
360
316
|
|
|
361
317
|
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
362
|
-
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
|
+
}
|
|
363
326
|
workgroupBarrier();
|
|
364
327
|
}
|
|
365
328
|
|
|
366
|
-
if (lid.x == 0u) {
|
|
367
|
-
|
|
368
|
-
|
|
369
|
-
|
|
370
|
-
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
371
|
-
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
372
|
-
|
|
373
|
-
struct Params {
|
|
374
|
-
n: u32,
|
|
375
|
-
c: f32,
|
|
376
|
-
s: f32,
|
|
377
|
-
x_inc: u32,
|
|
378
|
-
y_inc: u32,
|
|
329
|
+
if (lid.x == 0u) {
|
|
330
|
+
partialsScale[wgid.x] = tileScale[0];
|
|
331
|
+
partialsSsq[wgid.x] = tileSsq[0];
|
|
332
|
+
}
|
|
379
333
|
}
|
|
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.
|
|
380
343
|
|
|
381
|
-
@group(0) @binding(
|
|
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>;
|
|
382
347
|
|
|
383
348
|
const WGS: u32 = 64;
|
|
384
349
|
|
|
385
|
-
|
|
386
|
-
|
|
387
|
-
|
|
388
|
-
|
|
389
|
-
) {
|
|
390
|
-
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
391
|
-
let xi = x[id * params.x_inc];
|
|
392
|
-
let yi = y[id * params.y_inc];
|
|
393
|
-
x[id * params.x_inc] = params.c * xi + params.s * yi;
|
|
394
|
-
y[id * params.y_inc] = -params.s * xi + params.c * yi;
|
|
395
|
-
}
|
|
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,
|
|
396
354
|
}
|
|
397
|
-
`});var ee,re=W(()=>{ee=`// srotm: applies modified Givens rotation H to vectors x and y.
|
|
398
|
-
// param[0] = flag: -1 (full H), 0 (unit diagonal), 1 (unit off-diagonal)
|
|
399
|
-
// param = [ flag, h11, h21, h12, h22 ]
|
|
400
|
-
// flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
|
|
401
|
-
|
|
402
|
-
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
403
|
-
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
404
|
-
@group(0) @binding(2) var<storage, read> param: array<f32>;
|
|
405
355
|
|
|
406
|
-
|
|
407
|
-
|
|
408
|
-
|
|
409
|
-
|
|
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);
|
|
410
365
|
}
|
|
411
366
|
|
|
412
|
-
|
|
413
|
-
|
|
414
|
-
const WGS: u32 = 64;
|
|
367
|
+
var<workgroup> tileScale: array<f32, 64>;
|
|
368
|
+
var<workgroup> tileSsq: array<f32, 64>;
|
|
415
369
|
|
|
416
370
|
@compute @workgroup_size(64)
|
|
417
|
-
fn
|
|
418
|
-
@builtin(
|
|
419
|
-
@builtin(num_workgroups) num_wg: vec3u,
|
|
371
|
+
fn reduce_scaled(
|
|
372
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
420
373
|
) {
|
|
421
|
-
let
|
|
422
|
-
|
|
423
|
-
|
|
424
|
-
|
|
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();
|
|
425
382
|
|
|
426
|
-
|
|
427
|
-
|
|
428
|
-
|
|
429
|
-
|
|
430
|
-
|
|
431
|
-
|
|
432
|
-
|
|
433
|
-
|
|
434
|
-
|
|
435
|
-
|
|
436
|
-
h11 = param[1]; h21 = -1.0;
|
|
437
|
-
h12 = 1.0; h22 = param[4];
|
|
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();
|
|
438
393
|
}
|
|
439
394
|
|
|
440
|
-
|
|
441
|
-
|
|
442
|
-
let yi = y[id * params.y_inc];
|
|
443
|
-
x[id * params.x_inc] = h11 * xi + h12 * yi;
|
|
444
|
-
y[id * params.y_inc] = h21 * xi + h22 * yi;
|
|
395
|
+
if (i == 0u) {
|
|
396
|
+
result[0] = tileScale[0] * sqrt(tileSsq[0]);
|
|
445
397
|
}
|
|
446
398
|
}
|
|
447
|
-
`});var
|
|
399
|
+
`});var rt,Je=V(()=>{rt=`// isamax: returns index of element with largest absolute value
|
|
448
400
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/argmax.wgsl.
|
|
449
401
|
|
|
450
402
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -526,36 +478,553 @@ fn main(
|
|
|
526
478
|
partials_idx[wgid.x] = tile_idx[0];
|
|
527
479
|
}
|
|
528
480
|
}
|
|
529
|
-
`});var
|
|
530
|
-
//
|
|
531
|
-
//
|
|
532
|
-
// still covers all rows when m exceeds maxComputeWorkgroupsPerDimension.
|
|
533
|
-
// Threads stride through A[row, :] and x with coalesced reads (consecutive
|
|
534
|
-
// threads \u2192 consecutive addresses). Four independent accumulators let the GPU
|
|
535
|
-
// pipeline memory requests across iterations (ILP=4), hiding the
|
|
536
|
-
// global-memory latency.
|
|
537
|
-
|
|
538
|
-
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
539
|
-
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
540
|
-
@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.
|
|
541
484
|
|
|
542
|
-
|
|
543
|
-
|
|
544
|
-
|
|
545
|
-
alpha: f32,
|
|
546
|
-
beta: f32,
|
|
547
|
-
incx: u32,
|
|
548
|
-
incy: u32,
|
|
549
|
-
lda: u32,
|
|
550
|
-
}
|
|
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>;
|
|
551
488
|
|
|
552
|
-
|
|
489
|
+
const WGS: u32 = 64;
|
|
553
490
|
|
|
554
|
-
|
|
555
|
-
var<workgroup>
|
|
491
|
+
var<workgroup> tile_val: array<f32, 64>;
|
|
492
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
556
493
|
|
|
557
494
|
@compute @workgroup_size(64)
|
|
558
|
-
fn
|
|
495
|
+
fn reduce(
|
|
496
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
497
|
+
) {
|
|
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();
|
|
509
|
+
|
|
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];
|
|
517
|
+
}
|
|
518
|
+
}
|
|
519
|
+
workgroupBarrier();
|
|
520
|
+
}
|
|
521
|
+
|
|
522
|
+
if (i == 0u) { result[0] = tile_idx[0]; }
|
|
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.
|
|
537
|
+
|
|
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.
|
|
543
|
+
|
|
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);
|
|
549
|
+
}
|
|
550
|
+
return a;
|
|
551
|
+
}
|
|
552
|
+
`});var nt,st=V(()=>{nt=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
553
|
+
|
|
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
|
+
}
|
|
569
|
+
|
|
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);
|
|
577
|
+
}
|
|
578
|
+
|
|
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);
|
|
585
|
+
}
|
|
586
|
+
|
|
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);
|
|
592
|
+
}
|
|
593
|
+
|
|
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>;
|
|
605
|
+
|
|
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
|
+
}
|
|
614
|
+
|
|
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
|
+
}
|
|
622
|
+
|
|
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);
|
|
628
|
+
}
|
|
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.
|
|
633
|
+
|
|
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;
|
|
639
|
+
|
|
640
|
+
struct Params {
|
|
641
|
+
n: u32,
|
|
642
|
+
x_inc: u32,
|
|
643
|
+
}
|
|
644
|
+
|
|
645
|
+
const WGS: u32 = 64;
|
|
646
|
+
|
|
647
|
+
var<workgroup> tile: array<DD, 64>;
|
|
648
|
+
|
|
649
|
+
@compute @workgroup_size(64)
|
|
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,
|
|
655
|
+
) {
|
|
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);
|
|
660
|
+
|
|
661
|
+
let stride = num_wg.x * WGS;
|
|
662
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
663
|
+
|
|
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
|
+
}
|
|
678
|
+
|
|
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
|
+
}
|
|
693
|
+
|
|
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();
|
|
698
|
+
|
|
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
|
+
}
|
|
708
|
+
|
|
709
|
+
if (lid.x == 0u) {
|
|
710
|
+
partialsHi[wgid.x] = tile[0].hi;
|
|
711
|
+
partialsLo[wgid.x] = tile[0].lo;
|
|
712
|
+
}
|
|
713
|
+
}
|
|
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).
|
|
720
|
+
|
|
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>;
|
|
725
|
+
|
|
726
|
+
const WGS: u32 = 64;
|
|
727
|
+
|
|
728
|
+
var<workgroup> tile: array<DD, 64>;
|
|
729
|
+
|
|
730
|
+
@compute @workgroup_size(64)
|
|
731
|
+
fn reduce_f64(
|
|
732
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
733
|
+
) {
|
|
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();
|
|
739
|
+
|
|
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
|
+
}
|
|
748
|
+
|
|
749
|
+
if (i == 0u) {
|
|
750
|
+
resultHi[0] = tile[0].hi;
|
|
751
|
+
resultLo[0] = tile[0].lo;
|
|
752
|
+
}
|
|
753
|
+
}
|
|
754
|
+
`});var ct,ft=V(()=>{ct=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
755
|
+
|
|
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;
|
|
765
|
+
}
|
|
766
|
+
`});var dt,pt=V(()=>{dt=`// Requires f64/dekker.wgsl concatenated first for the DD struct.
|
|
767
|
+
|
|
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;
|
|
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).
|
|
777
|
+
|
|
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;
|
|
784
|
+
|
|
785
|
+
struct Params {
|
|
786
|
+
n: u32,
|
|
787
|
+
x_inc: u32,
|
|
788
|
+
}
|
|
789
|
+
|
|
790
|
+
const WGS: u32 = 64;
|
|
791
|
+
|
|
792
|
+
var<workgroup> tile_val: array<DD, 64>;
|
|
793
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
794
|
+
|
|
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;
|
|
808
|
+
|
|
809
|
+
let stride = num_wg.x * WGS;
|
|
810
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
811
|
+
|
|
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
|
+
}
|
|
831
|
+
|
|
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
|
+
}
|
|
845
|
+
|
|
846
|
+
tile_val[lid.x] = best_val0;
|
|
847
|
+
tile_idx[lid.x] = best_idx0;
|
|
848
|
+
workgroupBarrier();
|
|
849
|
+
|
|
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
|
+
}
|
|
862
|
+
|
|
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
|
+
}
|
|
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>;
|
|
881
|
+
|
|
882
|
+
const WGS: u32 = 64;
|
|
883
|
+
|
|
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];
|
|
901
|
+
}
|
|
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();
|
|
915
|
+
}
|
|
916
|
+
|
|
917
|
+
if (i == 0u) { result[0] = tile_idx[0]; }
|
|
918
|
+
}
|
|
919
|
+
`});var xt,yt=V(()=>{xt=`// srot: x = c*x + s*y, y = -s*x + c*y
|
|
920
|
+
|
|
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;
|
|
946
|
+
}
|
|
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.
|
|
952
|
+
|
|
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>;
|
|
956
|
+
|
|
957
|
+
struct Params {
|
|
958
|
+
n: u32,
|
|
959
|
+
x_inc: u32,
|
|
960
|
+
y_inc: u32,
|
|
961
|
+
}
|
|
962
|
+
|
|
963
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
964
|
+
|
|
965
|
+
const WGS: u32 = 64;
|
|
966
|
+
|
|
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];
|
|
973
|
+
|
|
974
|
+
var h11: f32; var h12: f32;
|
|
975
|
+
var h21: f32; var h22: f32;
|
|
976
|
+
|
|
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];
|
|
989
|
+
}
|
|
990
|
+
|
|
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;
|
|
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.
|
|
1006
|
+
|
|
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>;
|
|
1010
|
+
|
|
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
|
+
}
|
|
1020
|
+
|
|
1021
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1022
|
+
|
|
1023
|
+
const WGS: u32 = 64u;
|
|
1024
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
1025
|
+
|
|
1026
|
+
@compute @workgroup_size(64)
|
|
1027
|
+
fn main(
|
|
559
1028
|
@builtin(workgroup_id) wgid: vec3u,
|
|
560
1029
|
@builtin(local_invocation_id) lid: vec3u,
|
|
561
1030
|
@builtin(num_workgroups) nwg: vec3u,
|
|
@@ -595,13 +1064,15 @@ fn main(
|
|
|
595
1064
|
|
|
596
1065
|
if lid.x == 0u {
|
|
597
1066
|
let yi = row * params.incy;
|
|
598
|
-
y
|
|
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);
|
|
599
1070
|
}
|
|
600
1071
|
// All 64 threads must agree before the next row reuses scratch[].
|
|
601
1072
|
workgroupBarrier();
|
|
602
1073
|
}
|
|
603
1074
|
}
|
|
604
|
-
`});var
|
|
1075
|
+
`});var Gt,At=V(()=>{Gt=`// sgemv_t: y = alpha * A^T * x + beta * y (A is m\xD7n row-major, transposed)
|
|
605
1076
|
// each thread owns one column of A \u2192 one element of y (length n)
|
|
606
1077
|
// tiles over x (length m) using shared memory; four independent accumulators
|
|
607
1078
|
// let the GPU pipeline A reads across j within each tile (ILP=4)
|
|
@@ -663,10 +1134,12 @@ fn main(
|
|
|
663
1134
|
acc0 += A[k * params.lda + col] * x[k * params.incx];
|
|
664
1135
|
}
|
|
665
1136
|
let yi = col * params.incy;
|
|
666
|
-
|
|
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);
|
|
667
1140
|
}
|
|
668
1141
|
}
|
|
669
|
-
`});var
|
|
1142
|
+
`});var kt,St=V(()=>{kt=`// ssymv: y = alpha * A * x + beta * y
|
|
670
1143
|
// A is n\xD7n symmetric, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
671
1144
|
// The logical matrix is fully dense (symmetric), so each row's dot product
|
|
672
1145
|
// sums over all n columns; entries on the unstored side of the diagonal are
|
|
@@ -731,11 +1204,13 @@ fn main(
|
|
|
731
1204
|
}
|
|
732
1205
|
|
|
733
1206
|
if lid.x == 0u {
|
|
734
|
-
|
|
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);
|
|
735
1210
|
}
|
|
736
1211
|
}
|
|
737
1212
|
}
|
|
738
|
-
`});var
|
|
1213
|
+
`});var Mt,Nt=V(()=>{Mt=`// strmv: y = op(A) * x
|
|
739
1214
|
// A is n\xD7n triangular, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
740
1215
|
// op(A) is A (trans=0) or A^T (trans=1).
|
|
741
1216
|
// diag=1 (unit) treats the diagonal as 1 without reading A's diagonal values.
|
|
@@ -834,11 +1309,242 @@ fn main(
|
|
|
834
1309
|
}
|
|
835
1310
|
|
|
836
1311
|
if lid.x == 0u {
|
|
837
|
-
y[ i * params.incy ] = scratch[0];
|
|
1312
|
+
y[ i * params.incy ] = scratch[0];
|
|
1313
|
+
}
|
|
1314
|
+
}
|
|
1315
|
+
}
|
|
1316
|
+
`});var xe,It=V(()=>{xe=`// strsv_invert_block: computes ONE column (workgroup_id.x) of ONE block's
|
|
1317
|
+
// (workgroup_id.y) explicit inverse, via the same one-row-at-a-time
|
|
1318
|
+
// substitution as strsv_block.wgsl, but solving against a unit basis vector
|
|
1319
|
+
// e_col instead of the real right-hand side, and writing to a dense
|
|
1320
|
+
// (BLOCK_SIZE x BLOCK_SIZE, row-major) scratch buffer per block instead of
|
|
1321
|
+
// mutating x. Dispatched once for the whole matrix (2D: BLOCK_SIZE columns x
|
|
1322
|
+
// numBlocks), fully in parallel -- unlike the sequential per-block main
|
|
1323
|
+
// loop in strsv.mjs, no block's inverse depends on any other block or on x.
|
|
1324
|
+
//
|
|
1325
|
+
// A triangular block's inverse is itself triangular: forward (effectively-
|
|
1326
|
+
// lower, e.g. no-trans+lower) blocks have inverse column col nonzero only
|
|
1327
|
+
// for row>=col, solved in increasing row order; backward (effectively-
|
|
1328
|
+
// upper) blocks have it nonzero only for row<=col, solved in decreasing
|
|
1329
|
+
// order. Rows outside a column's nonzero range are written as literal 0 \u2014
|
|
1330
|
+
// strsv_apply_inverse.wgsl's dense matvec depends on that, not just on
|
|
1331
|
+
// those entries being mathematically implied zero.
|
|
1332
|
+
|
|
1333
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1334
|
+
@group(0) @binding(1) var<storage, read_write> Ainv: array<f32>;
|
|
1335
|
+
|
|
1336
|
+
struct Params {
|
|
1337
|
+
n: u32,
|
|
1338
|
+
lda: u32,
|
|
1339
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
1340
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
1341
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
1342
|
+
}
|
|
1343
|
+
|
|
1344
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
1345
|
+
|
|
1346
|
+
const WGS: u32 = 64u;
|
|
1347
|
+
const BLOCK_SIZE: u32 = 64u;
|
|
1348
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
1349
|
+
|
|
1350
|
+
fn readA(i: u32, j: u32) -> f32 {
|
|
1351
|
+
if params.trans == 0u {
|
|
1352
|
+
return A[i * params.lda + j];
|
|
1353
|
+
} else {
|
|
1354
|
+
return A[j * params.lda + i];
|
|
1355
|
+
}
|
|
1356
|
+
}
|
|
1357
|
+
|
|
1358
|
+
@compute @workgroup_size(64)
|
|
1359
|
+
fn strsv_invert_block_main(
|
|
1360
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1361
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1362
|
+
) {
|
|
1363
|
+
let col = wgid.x;
|
|
1364
|
+
let blockIndex = wgid.y;
|
|
1365
|
+
let blockStart = blockIndex * BLOCK_SIZE;
|
|
1366
|
+
var blockEnd = blockStart + BLOCK_SIZE;
|
|
1367
|
+
if (blockEnd > params.n) { blockEnd = params.n; }
|
|
1368
|
+
let blockLen = blockEnd - blockStart;
|
|
1369
|
+
|
|
1370
|
+
if (col >= blockLen) { return; }
|
|
1371
|
+
|
|
1372
|
+
let ainvBase = blockIndex * BLOCK_SIZE * BLOCK_SIZE;
|
|
1373
|
+
let forward = (params.trans == 0u) == (params.uplo == 0u);
|
|
1374
|
+
|
|
1375
|
+
if forward {
|
|
1376
|
+
for (var r = lid.x; r < col; r += WGS) {
|
|
1377
|
+
Ainv[ainvBase + r * BLOCK_SIZE + col] = 0.0;
|
|
1378
|
+
}
|
|
1379
|
+
} else {
|
|
1380
|
+
for (var r = col + 1u + lid.x; r < blockLen; r += WGS) {
|
|
1381
|
+
Ainv[ainvBase + r * BLOCK_SIZE + col] = 0.0;
|
|
1382
|
+
}
|
|
1383
|
+
}
|
|
1384
|
+
storageBarrier();
|
|
1385
|
+
workgroupBarrier();
|
|
1386
|
+
|
|
1387
|
+
let numSteps = select(col + 1u, blockLen - col, forward);
|
|
1388
|
+
for (var step = 0u; step < numSteps; step++) {
|
|
1389
|
+
let localRow = select(col - step, col + step, forward);
|
|
1390
|
+
let i = blockStart + localRow;
|
|
1391
|
+
|
|
1392
|
+
var acc = 0.0f;
|
|
1393
|
+
if forward {
|
|
1394
|
+
for (var lj = col + lid.x; lj < localRow; lj += WGS) {
|
|
1395
|
+
acc += readA(i, blockStart + lj) * Ainv[ainvBase + lj * BLOCK_SIZE + col];
|
|
1396
|
+
}
|
|
1397
|
+
} else {
|
|
1398
|
+
for (var lj = localRow + 1u + lid.x; lj <= col; lj += WGS) {
|
|
1399
|
+
acc += readA(i, blockStart + lj) * Ainv[ainvBase + lj * BLOCK_SIZE + col];
|
|
1400
|
+
}
|
|
1401
|
+
}
|
|
1402
|
+
|
|
1403
|
+
scratch[lid.x] = acc;
|
|
1404
|
+
workgroupBarrier();
|
|
1405
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
1406
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
1407
|
+
workgroupBarrier();
|
|
1408
|
+
}
|
|
1409
|
+
|
|
1410
|
+
if lid.x == 0u {
|
|
1411
|
+
let e = select(0.0, 1.0, localRow == col);
|
|
1412
|
+
let rhs = e - scratch[0];
|
|
1413
|
+
var val: f32;
|
|
1414
|
+
if params.diag == 1u {
|
|
1415
|
+
val = rhs;
|
|
1416
|
+
} else {
|
|
1417
|
+
val = rhs / A[i * params.lda + i];
|
|
1418
|
+
}
|
|
1419
|
+
Ainv[ainvBase + localRow * BLOCK_SIZE + col] = val;
|
|
1420
|
+
}
|
|
1421
|
+
storageBarrier();
|
|
1422
|
+
workgroupBarrier();
|
|
1423
|
+
}
|
|
1424
|
+
}
|
|
1425
|
+
`});var Pt,Rt=V(()=>{Pt=`// strsv_apply_inverse: given a precomputed block inverse (from
|
|
1426
|
+
// strsv_invert_block.wgsl), computes this block's solution as a dense
|
|
1427
|
+
// matrix-vector multiply against the block's current remainder in x \u2014
|
|
1428
|
+
// replacing what the old strsv_block.wgsl did via a genuinely sequential,
|
|
1429
|
+
// barrier-per-row substitution.
|
|
1430
|
+
//
|
|
1431
|
+
// All blockLen rows are computed in parallel within a single workgroup: the
|
|
1432
|
+
// remainder is loaded into workgroup-shared memory once, then each thread
|
|
1433
|
+
// independently computes one full row's dot product from that shared copy.
|
|
1434
|
+
// No further synchronization is needed after the load \u2014 every thread only
|
|
1435
|
+
// reads shared memory from then on (never written again within this call)
|
|
1436
|
+
// and writes a distinct element of x, so there's no cross-thread hazard to
|
|
1437
|
+
// guard against.
|
|
1438
|
+
|
|
1439
|
+
@group(0) @binding(0) var<storage, read> Ainv: array<f32>;
|
|
1440
|
+
@group(0) @binding(1) var<storage, read_write> x: array<f32>;
|
|
1441
|
+
|
|
1442
|
+
struct Params {
|
|
1443
|
+
incx: u32,
|
|
1444
|
+
blockIndex: u32,
|
|
1445
|
+
blockStart: u32,
|
|
1446
|
+
blockEnd: u32,
|
|
1447
|
+
}
|
|
1448
|
+
|
|
1449
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
1450
|
+
|
|
1451
|
+
const BLOCK_SIZE: u32 = 64u;
|
|
1452
|
+
var<workgroup> xLocal: array<f32, 64>;
|
|
1453
|
+
|
|
1454
|
+
@compute @workgroup_size(64)
|
|
1455
|
+
fn strsv_apply_inverse_main(@builtin(local_invocation_id) lid: vec3u) {
|
|
1456
|
+
let blockLen = params.blockEnd - params.blockStart;
|
|
1457
|
+
|
|
1458
|
+
if (lid.x < blockLen) {
|
|
1459
|
+
xLocal[lid.x] = x[(params.blockStart + lid.x) * params.incx];
|
|
1460
|
+
}
|
|
1461
|
+
workgroupBarrier();
|
|
1462
|
+
|
|
1463
|
+
if (lid.x >= blockLen) { return; }
|
|
1464
|
+
|
|
1465
|
+
let ainvBase = params.blockIndex * BLOCK_SIZE * BLOCK_SIZE;
|
|
1466
|
+
var acc = 0.0f;
|
|
1467
|
+
for (var j = 0u; j < blockLen; j++) {
|
|
1468
|
+
acc += Ainv[ainvBase + lid.x * BLOCK_SIZE + j] * xLocal[j];
|
|
1469
|
+
}
|
|
1470
|
+
x[(params.blockStart + lid.x) * params.incx] = acc;
|
|
1471
|
+
}
|
|
1472
|
+
`});var Tt,Dt=V(()=>{Tt=`// strsv_update: subtracts a solved block's contribution from every
|
|
1473
|
+
// remaining row in parallel (one workgroup per row, like strmv.wgsl) \u2014
|
|
1474
|
+
// this is what turns strsv's O(n) sequential stages into O(n/blockSize).
|
|
1475
|
+
// No diag/masking needed: this region never touches the diagonal.
|
|
1476
|
+
|
|
1477
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1478
|
+
@group(0) @binding(1) var<storage, read_write> x: array<f32>;
|
|
1479
|
+
|
|
1480
|
+
struct Params {
|
|
1481
|
+
n: u32,
|
|
1482
|
+
incx: u32,
|
|
1483
|
+
lda: u32,
|
|
1484
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
1485
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
1486
|
+
blockStart: u32,
|
|
1487
|
+
blockEnd: u32, // exclusive
|
|
1488
|
+
}
|
|
1489
|
+
|
|
1490
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
1491
|
+
|
|
1492
|
+
const WGS: u32 = 64u;
|
|
1493
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
1494
|
+
|
|
1495
|
+
@compute @workgroup_size(64)
|
|
1496
|
+
fn strsv_update_main(
|
|
1497
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1498
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1499
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
1500
|
+
) {
|
|
1501
|
+
// forward: remaining rows are [blockEnd,n); backward: [0,blockStart).
|
|
1502
|
+
let forward = (params.trans == 0u) == (params.uplo == 0u);
|
|
1503
|
+
|
|
1504
|
+
var rangeStart: u32;
|
|
1505
|
+
var rangeEnd: u32;
|
|
1506
|
+
|
|
1507
|
+
if forward {
|
|
1508
|
+
rangeStart = params.blockEnd;
|
|
1509
|
+
rangeEnd = params.n;
|
|
1510
|
+
} else {
|
|
1511
|
+
rangeStart = 0u;
|
|
1512
|
+
rangeEnd = params.blockStart;
|
|
1513
|
+
}
|
|
1514
|
+
|
|
1515
|
+
if (rangeStart >= rangeEnd) { return; }
|
|
1516
|
+
let count = rangeEnd - rangeStart;
|
|
1517
|
+
|
|
1518
|
+
for (var idx = wgid.x; idx < count; idx += nwg.x) {
|
|
1519
|
+
let i = rangeStart + idx;
|
|
1520
|
+
|
|
1521
|
+
// No-trans reads A[i,j]; transpose reads A[j,i] \u2014 uplo only sets the range above.
|
|
1522
|
+
var acc = 0.0f;
|
|
1523
|
+
if params.trans == 0u {
|
|
1524
|
+
for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
|
|
1525
|
+
acc += A[i * params.lda + j] * x[j * params.incx];
|
|
1526
|
+
}
|
|
1527
|
+
} else {
|
|
1528
|
+
for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
|
|
1529
|
+
acc += A[j * params.lda + i] * x[j * params.incx];
|
|
1530
|
+
}
|
|
1531
|
+
}
|
|
1532
|
+
|
|
1533
|
+
// Parallel reduction: 64 \u2192 1
|
|
1534
|
+
scratch[lid.x] = acc;
|
|
1535
|
+
workgroupBarrier();
|
|
1536
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
1537
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
1538
|
+
workgroupBarrier();
|
|
1539
|
+
}
|
|
1540
|
+
|
|
1541
|
+
if lid.x == 0u {
|
|
1542
|
+
x[i * params.incx] -= scratch[0];
|
|
838
1543
|
}
|
|
1544
|
+
workgroupBarrier();
|
|
839
1545
|
}
|
|
840
1546
|
}
|
|
841
|
-
`});var
|
|
1547
|
+
`});var Ct,jt=V(()=>{Ct=`// sger: A := alpha * x * y^T + A (rank-1 update, A is m\xD7n general/dense)
|
|
842
1548
|
|
|
843
1549
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
844
1550
|
@group(0) @binding(1) var<storage, read> y: array<f32>;
|
|
@@ -886,7 +1592,7 @@ fn main(
|
|
|
886
1592
|
}
|
|
887
1593
|
}
|
|
888
1594
|
}
|
|
889
|
-
`});var
|
|
1595
|
+
`});var Wt,Lt=V(()=>{Wt=`// ssyr: A := alpha * x * x^T + A (symmetric rank-1 update)
|
|
890
1596
|
// A is n\xD7n symmetric; only the triangle specified by uplo is referenced/updated,
|
|
891
1597
|
// the other triangle is implied by symmetry (not touched).
|
|
892
1598
|
|
|
@@ -946,7 +1652,7 @@ fn main(
|
|
|
946
1652
|
}
|
|
947
1653
|
}
|
|
948
1654
|
}
|
|
949
|
-
`});var
|
|
1655
|
+
`});var qt,Ft=V(()=>{qt=`// ssyr2: A := alpha * x * y^T + alpha * y * x^T + A (symmetric rank-2 update)
|
|
950
1656
|
// A is n\xD7n symmetric; only the triangle specified by uplo is referenced/updated,
|
|
951
1657
|
// the other triangle is implied by symmetry (not touched).
|
|
952
1658
|
|
|
@@ -1009,702 +1715,775 @@ fn main(
|
|
|
1009
1715
|
}
|
|
1010
1716
|
}
|
|
1011
1717
|
}
|
|
1012
|
-
`});var
|
|
1013
|
-
//
|
|
1014
|
-
//
|
|
1015
|
-
//
|
|
1718
|
+
`});var re,Ut=V(()=>{re=`// sgemm_small: C = alpha * op(A) * op(B) + beta * C \u2014 small-tile half of
|
|
1719
|
+
// the two-tier autotuned dispatch (see sgemm.mjs and sgemm_large.wgsl).
|
|
1720
|
+
// BM=BN=32, BK=8, TM=TN=2 \u2014 wins over the large tile below a 6x6=36
|
|
1721
|
+
// workgroup grid of 64-tiles, where the large tile doesn't have enough
|
|
1722
|
+
// workgroups to fill the GPU. Same structure as sgemm_large.wgsl (2D
|
|
1723
|
+
// register-blocked, shared-memory-tiled), just smaller.
|
|
1016
1724
|
//
|
|
1017
|
-
//
|
|
1018
|
-
//
|
|
1019
|
-
//
|
|
1020
|
-
//
|
|
1021
|
-
// (
|
|
1022
|
-
//
|
|
1023
|
-
|
|
1024
|
-
|
|
1025
|
-
|
|
1026
|
-
const
|
|
1027
|
-
|
|
1028
|
-
|
|
1029
|
-
|
|
1030
|
-
|
|
1031
|
-
|
|
1032
|
-
|
|
1033
|
-
|
|
1034
|
-
|
|
1035
|
-
|
|
1036
|
-
|
|
1037
|
-
|
|
1038
|
-
|
|
1039
|
-
|
|
1040
|
-
|
|
1041
|
-
|
|
1042
|
-
aux: u32,
|
|
1043
|
-
}
|
|
1044
|
-
|
|
1045
|
-
// Mirrors packedToFields() in f64pack.mjs.
|
|
1046
|
-
fn decode(mainBits: u32, auxBits: u32) -> Fields {
|
|
1047
|
-
let sign = mainBits >> 31u;
|
|
1048
|
-
let expMain = (mainBits >> 23u) & 0xffu;
|
|
1049
|
-
let mantMain = mainBits & 0x7fffffu;
|
|
1050
|
-
|
|
1051
|
-
let auxSign = auxBits >> 31u;
|
|
1052
|
-
let auxExp8 = (auxBits >> 23u) & 0xffu;
|
|
1053
|
-
let auxMant23 = auxBits & 0x7fffffu;
|
|
1054
|
-
|
|
1055
|
-
let expExtra = (auxSign << 2u) | (auxExp8 >> 6u);
|
|
1056
|
-
let mantExtra29 = ((auxExp8 & 0x3fu) << 23u) | auxMant23;
|
|
1057
|
-
|
|
1058
|
-
let rawExp = (expMain << 3u) | expExtra;
|
|
1059
|
-
let mantissaHi = mantMain >> 3u;
|
|
1060
|
-
let mantTop3 = mantMain & 0x7u;
|
|
1061
|
-
let lo = (mantTop3 << 29u) | mantExtra29;
|
|
1062
|
-
|
|
1063
|
-
return Fields(sign, rawExp, mantissaHi, lo);
|
|
1064
|
-
}
|
|
1065
|
-
|
|
1066
|
-
// Mirrors fieldsToPacked() in f64pack.mjs.
|
|
1067
|
-
fn encode(sign: u32, rawExp: u32, mantissaHi: u32, lo: u32) -> Packed {
|
|
1068
|
-
let expMain = rawExp >> 3u;
|
|
1069
|
-
let expExtra = rawExp & 0x7u;
|
|
1070
|
-
|
|
1071
|
-
let mantTop3 = lo >> 29u;
|
|
1072
|
-
let mantMain = (mantissaHi << 3u) | mantTop3;
|
|
1073
|
-
let mantExtra29 = lo & 0x1fffffffu;
|
|
1074
|
-
|
|
1075
|
-
let mainBits = (sign << 31u) | (expMain << 23u) | mantMain;
|
|
1076
|
-
|
|
1077
|
-
let auxSign = (expExtra >> 2u) & 0x1u;
|
|
1078
|
-
let auxExpTop2 = expExtra & 0x3u;
|
|
1079
|
-
let auxExpBot6 = mantExtra29 >> 23u;
|
|
1080
|
-
let auxMant23 = mantExtra29 & 0x7fffffu;
|
|
1081
|
-
let auxExp8 = (auxExpTop2 << 6u) | auxExpBot6;
|
|
1082
|
-
|
|
1083
|
-
let auxBits = (auxSign << 31u) | (auxExp8 << 23u) | auxMant23;
|
|
1084
|
-
|
|
1085
|
-
return Packed(bitcast<f32>(mainBits), auxBits);
|
|
1086
|
-
}
|
|
1087
|
-
|
|
1088
|
-
struct Pair { hi: u32, lo: u32 }
|
|
1089
|
-
struct Shifted { hi: u32, lo: u32, sticky: u32 }
|
|
1090
|
-
|
|
1091
|
-
// Two-word right shift by 0..64+ bits, folding every shifted-out 1 bit into
|
|
1092
|
-
// a returned sticky flag \u2014 used only for the (potentially huge) exponent
|
|
1093
|
-
// alignment shift, where exact bits can't all be kept.
|
|
1094
|
-
fn shr_sticky(hi: u32, lo: u32, n: u32) -> Shifted {
|
|
1095
|
-
if (n == 0u) {
|
|
1096
|
-
return Shifted(hi, lo, 0u);
|
|
1097
|
-
}
|
|
1098
|
-
if (n >= 64u) {
|
|
1099
|
-
return Shifted(0u, 0u, select(0u, 1u, hi != 0u || lo != 0u));
|
|
1100
|
-
}
|
|
1101
|
-
if (n < 32u) {
|
|
1102
|
-
let stickyBits = lo & ((1u << n) - 1u);
|
|
1103
|
-
let newLo = (lo >> n) | (hi << (32u - n));
|
|
1104
|
-
let newHi = hi >> n;
|
|
1105
|
-
return Shifted(newHi, newLo, select(0u, 1u, stickyBits != 0u));
|
|
1106
|
-
}
|
|
1107
|
-
if (n == 32u) {
|
|
1108
|
-
return Shifted(0u, hi, select(0u, 1u, lo != 0u));
|
|
1109
|
-
}
|
|
1110
|
-
let m = n - 32u;
|
|
1111
|
-
let stickyBits = lo | (hi & ((1u << m) - 1u));
|
|
1112
|
-
let newLo = hi >> m;
|
|
1113
|
-
return Shifted(0u, newLo, select(0u, 1u, stickyBits != 0u));
|
|
1114
|
-
}
|
|
1115
|
-
|
|
1116
|
-
// Two-word left shift by 0..63 bits \u2014 used only to renormalize after
|
|
1117
|
-
// cancellation, by an amount that exactly matches the leading-zero count,
|
|
1118
|
-
// so nothing meaningful is ever lost off the top.
|
|
1119
|
-
fn shl(hi: u32, lo: u32, n: u32) -> Pair {
|
|
1120
|
-
if (n == 0u) {
|
|
1121
|
-
return Pair(hi, lo);
|
|
1122
|
-
}
|
|
1123
|
-
if (n < 32u) {
|
|
1124
|
-
let newHi = (hi << n) | (lo >> (32u - n));
|
|
1125
|
-
let newLo = lo << n;
|
|
1126
|
-
return Pair(newHi, newLo);
|
|
1127
|
-
}
|
|
1128
|
-
if (n == 32u) {
|
|
1129
|
-
return Pair(lo, 0u);
|
|
1130
|
-
}
|
|
1131
|
-
let m = n - 32u;
|
|
1132
|
-
return Pair(lo << m, 0u);
|
|
1133
|
-
}
|
|
1134
|
-
|
|
1135
|
-
fn add64(aHi: u32, aLo: u32, bHi: u32, bLo: u32) -> Pair {
|
|
1136
|
-
let sumLo = aLo + bLo;
|
|
1137
|
-
let carry = select(0u, 1u, sumLo < aLo); // wrapped around -> there was a carry
|
|
1138
|
-
let sumHi = aHi + bHi + carry;
|
|
1139
|
-
return Pair(sumHi, sumLo);
|
|
1140
|
-
}
|
|
1141
|
-
|
|
1142
|
-
// Assumes (aHi:aLo) >= (bHi:bLo) \u2014 callers guarantee this so no sign handling is needed.
|
|
1143
|
-
fn sub64(aHi: u32, aLo: u32, bHi: u32, bLo: u32) -> Pair {
|
|
1144
|
-
let borrow = select(0u, 1u, aLo < bLo);
|
|
1145
|
-
let diffLo = aLo - bLo;
|
|
1146
|
-
let diffHi = aHi - bHi - borrow;
|
|
1147
|
-
return Pair(diffHi, diffLo);
|
|
1148
|
-
}
|
|
1149
|
-
|
|
1150
|
-
fn ge64(aHi: u32, aLo: u32, bHi: u32, bLo: u32) -> bool {
|
|
1151
|
-
return aHi > bHi || (aHi == bHi && aLo >= bLo);
|
|
1152
|
-
}
|
|
1153
|
-
|
|
1154
|
-
// The actual IEEE-754 addition, returning decoded Fields rather than an
|
|
1155
|
-
// encoded Packed pair \u2014 lets a caller that's accumulating many values in a
|
|
1156
|
-
// row (e.g. dasum.wgsl's per-thread reduction loop) keep the running total
|
|
1157
|
-
// in Fields form the whole time, only encoding once at the very end, instead
|
|
1158
|
-
// of paying a decode+encode round-trip on every single addition. computeSum
|
|
1159
|
-
// (below) is the Packed-in/Packed-out convenience wrapper around this.
|
|
1160
|
-
fn addFields(a: Fields, b: Fields) -> Fields {
|
|
1161
|
-
let aIsNaN = a.rawExp == EXP_ALL_ONES && (a.mantissaHi != 0u || a.lo != 0u);
|
|
1162
|
-
let bIsNaN = b.rawExp == EXP_ALL_ONES && (b.mantissaHi != 0u || b.lo != 0u);
|
|
1163
|
-
if (aIsNaN || bIsNaN) {
|
|
1164
|
-
return Fields(0u, EXP_ALL_ONES, QUIET_NAN_MANTISSA_HI, 0u);
|
|
1165
|
-
}
|
|
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
|
+
//
|
|
1732
|
+
// col mapped to gid.x for coalesced B/C access (row-major: col contiguous).
|
|
1733
|
+
|
|
1734
|
+
const BM: u32 = 32u;
|
|
1735
|
+
const BN: u32 = 32u;
|
|
1736
|
+
const BK: u32 = 8u;
|
|
1737
|
+
const TM: u32 = 2u;
|
|
1738
|
+
const TN: u32 = 2u;
|
|
1739
|
+
const THREADS_X: u32 = BN / TN;
|
|
1740
|
+
const THREADS_Y: u32 = BM / TM;
|
|
1741
|
+
const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 256
|
|
1742
|
+
const STRIDE_A: u32 = NUM_THREADS / BK;
|
|
1743
|
+
const STRIDE_B: u32 = NUM_THREADS / BN;
|
|
1744
|
+
|
|
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>;
|
|
1166
1750
|
|
|
1167
|
-
|
|
1168
|
-
|
|
1169
|
-
|
|
1170
|
-
|
|
1171
|
-
|
|
1751
|
+
struct Params {
|
|
1752
|
+
m: u32,
|
|
1753
|
+
n: u32,
|
|
1754
|
+
k: u32,
|
|
1755
|
+
alpha: f32,
|
|
1756
|
+
beta: f32,
|
|
1757
|
+
lda: u32,
|
|
1758
|
+
ldb: u32,
|
|
1759
|
+
ldc: u32,
|
|
1760
|
+
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
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
|
|
1764
|
+
}
|
|
1765
|
+
|
|
1766
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
1767
|
+
|
|
1768
|
+
var<workgroup> As: array<f32, BM * BK>;
|
|
1769
|
+
var<workgroup> Bs: array<f32, BK * BN>;
|
|
1770
|
+
|
|
1771
|
+
@compute @workgroup_size(THREADS_X, THREADS_Y)
|
|
1772
|
+
fn main(
|
|
1773
|
+
@builtin(workgroup_id) wid: vec3u,
|
|
1774
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1775
|
+
@builtin(local_invocation_index) tid: u32,
|
|
1776
|
+
) {
|
|
1777
|
+
let blockRow = wid.y * BM;
|
|
1778
|
+
let blockCol = wid.x * BN;
|
|
1779
|
+
let threadCol = lid.x;
|
|
1780
|
+
let threadRow = lid.y;
|
|
1781
|
+
|
|
1782
|
+
let innerRowA = tid / BK;
|
|
1783
|
+
let innerColA = tid % BK;
|
|
1784
|
+
let innerRowB = tid / BN;
|
|
1785
|
+
let innerColB = tid % BN;
|
|
1786
|
+
|
|
1787
|
+
var threadResults: array<f32, TM * TN>;
|
|
1788
|
+
for (var i = 0u; i < TM * TN; i++) {
|
|
1789
|
+
threadResults[i] = 0.0;
|
|
1790
|
+
}
|
|
1791
|
+
var regM: array<f32, TM>;
|
|
1792
|
+
var regN: array<f32, TN>;
|
|
1793
|
+
|
|
1794
|
+
let numTiles = (params.k + BK - 1u) / BK;
|
|
1795
|
+
for (var t = 0u; t < numTiles; t++) {
|
|
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
|
+
}
|
|
1172
1844
|
}
|
|
1173
|
-
return Fields(a.sign, EXP_ALL_ONES, 0u, 0u);
|
|
1174
|
-
}
|
|
1175
|
-
if (aIsInf) { return Fields(a.sign, EXP_ALL_ONES, 0u, 0u); }
|
|
1176
|
-
if (bIsInf) { return Fields(b.sign, EXP_ALL_ONES, 0u, 0u); }
|
|
1177
|
-
|
|
1178
|
-
let aIsZero = a.rawExp == 0u && a.mantissaHi == 0u && a.lo == 0u;
|
|
1179
|
-
let bIsZero = b.rawExp == 0u && b.mantissaHi == 0u && b.lo == 0u;
|
|
1180
|
-
if (aIsZero && bIsZero) {
|
|
1181
|
-
return Fields(a.sign & b.sign, 0u, 0u, 0u);
|
|
1182
|
-
}
|
|
1183
|
-
if (aIsZero) { return Fields(b.sign, b.rawExp, b.mantissaHi, b.lo); }
|
|
1184
|
-
if (bIsZero) { return Fields(a.sign, a.rawExp, a.mantissaHi, a.lo); }
|
|
1185
|
-
|
|
1186
|
-
// Effective (unbiased) exponent \u2014 subnormals share the smallest normal
|
|
1187
|
-
// exponent for alignment purposes and have no implicit leading 1.
|
|
1188
|
-
var expA = i32(a.rawExp) - BIAS;
|
|
1189
|
-
if (a.rawExp == 0u) { expA = 1 - BIAS; }
|
|
1190
|
-
var expB = i32(b.rawExp) - BIAS;
|
|
1191
|
-
if (b.rawExp == 0u) { expB = 1 - BIAS; }
|
|
1192
|
-
|
|
1193
|
-
let implicitA = select(0u, 1u, a.rawExp != 0u);
|
|
1194
|
-
let implicitB = select(0u, 1u, b.rawExp != 0u);
|
|
1195
|
-
|
|
1196
|
-
// Widen each 53-bit significand (implicit + 52-bit mantissa) by 3 zero
|
|
1197
|
-
// bits at the bottom \u2014 room for guard/round/sticky once alignment shifts happen.
|
|
1198
|
-
let sigHiA = (implicitA << 23u) | (a.mantissaHi << 3u) | (a.lo >> 29u);
|
|
1199
|
-
let sigLoA = a.lo << 3u;
|
|
1200
|
-
let sigHiB = (implicitB << 23u) | (b.mantissaHi << 3u) | (b.lo >> 29u);
|
|
1201
|
-
let sigLoB = b.lo << 3u;
|
|
1202
|
-
|
|
1203
|
-
// P = the operand with the larger exponent (Q = the other); on a tie, P =
|
|
1204
|
-
// whichever has the larger significand \u2014 keeps subtraction below always
|
|
1205
|
-
// non-negative without needing signed magnitudes.
|
|
1206
|
-
var signP: u32; var expP: i32; var sigHiP: u32; var sigLoP: u32;
|
|
1207
|
-
var signQ: u32; var expQ: i32; var sigHiQ: u32; var sigLoQ: u32;
|
|
1208
|
-
if (expA > expB || (expA == expB && ge64(sigHiA, sigLoA, sigHiB, sigLoB))) {
|
|
1209
|
-
signP = a.sign; expP = expA; sigHiP = sigHiA; sigLoP = sigLoA;
|
|
1210
|
-
signQ = b.sign; expQ = expB; sigHiQ = sigHiB; sigLoQ = sigLoB;
|
|
1211
|
-
} else {
|
|
1212
|
-
signP = b.sign; expP = expB; sigHiP = sigHiB; sigLoP = sigLoB;
|
|
1213
|
-
signQ = a.sign; expQ = expA; sigHiQ = sigHiA; sigLoQ = sigLoA;
|
|
1214
|
-
}
|
|
1215
|
-
|
|
1216
|
-
let diff = u32(expP - expQ);
|
|
1217
|
-
let shiftedQ = shr_sticky(sigHiQ, sigLoQ, diff);
|
|
1218
|
-
let alignedHiQ = shiftedQ.hi;
|
|
1219
|
-
let alignedLoQ = shiftedQ.lo | shiftedQ.sticky; // fold sticky into bit0
|
|
1220
|
-
|
|
1221
|
-
var sumHi: u32; var sumLo: u32;
|
|
1222
|
-
if (signP == signQ) {
|
|
1223
|
-
let s = add64(sigHiP, sigLoP, alignedHiQ, alignedLoQ);
|
|
1224
|
-
sumHi = s.hi; sumLo = s.lo;
|
|
1225
|
-
} else {
|
|
1226
|
-
let s = sub64(sigHiP, sigLoP, alignedHiQ, alignedLoQ); // P >= Q(aligned) by construction
|
|
1227
|
-
sumHi = s.hi; sumLo = s.lo;
|
|
1228
|
-
}
|
|
1229
1845
|
|
|
1230
|
-
|
|
1231
|
-
|
|
1232
|
-
|
|
1233
|
-
|
|
1234
|
-
|
|
1235
|
-
|
|
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
|
+
}
|
|
1893
|
+
}
|
|
1236
1894
|
|
|
1237
|
-
|
|
1238
|
-
if (sumHi != 0u) {
|
|
1239
|
-
leadPos = 32 + i32(31u - countLeadingZeros(sumHi));
|
|
1240
|
-
} else {
|
|
1241
|
-
leadPos = i32(31u - countLeadingZeros(sumLo));
|
|
1242
|
-
}
|
|
1243
|
-
let tentativeExp = leadPos + commonExp2;
|
|
1244
|
-
var targetLSBScale = tentativeExp - 52;
|
|
1245
|
-
if (tentativeExp < -1022) { targetLSBScale = -1074; }
|
|
1246
|
-
let shiftAmt = targetLSBScale - commonExp2;
|
|
1895
|
+
workgroupBarrier();
|
|
1247
1896
|
|
|
1248
|
-
|
|
1249
|
-
|
|
1250
|
-
|
|
1251
|
-
|
|
1252
|
-
|
|
1253
|
-
|
|
1254
|
-
|
|
1255
|
-
|
|
1256
|
-
|
|
1257
|
-
|
|
1258
|
-
|
|
1259
|
-
|
|
1260
|
-
let sh = shr_sticky(sumHi, sumLo, n);
|
|
1261
|
-
keepHi = sh.hi; keepLo = sh.lo;
|
|
1262
|
-
if (remainder > halfway || (remainder == halfway && (keepLo & 1u) != 0u)) {
|
|
1263
|
-
let inc = add64(keepHi, keepLo, 0u, 1u);
|
|
1264
|
-
keepHi = inc.hi; keepLo = inc.lo;
|
|
1897
|
+
for (var dotIdx = 0u; dotIdx < BK; dotIdx++) {
|
|
1898
|
+
for (var i = 0u; i < TM; i++) {
|
|
1899
|
+
regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
|
|
1900
|
+
}
|
|
1901
|
+
for (var i = 0u; i < TN; i++) {
|
|
1902
|
+
regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
|
|
1903
|
+
}
|
|
1904
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
1905
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
1906
|
+
threadResults[resIdxM * TN + resIdxN] += regM[resIdxM] * regN[resIdxN];
|
|
1907
|
+
}
|
|
1908
|
+
}
|
|
1265
1909
|
}
|
|
1266
|
-
}
|
|
1267
1910
|
|
|
1268
|
-
|
|
1269
|
-
if ((keepHi & (1u << 21u)) != 0u) { // rounding carried past the 53-bit budget
|
|
1270
|
-
let sh = shr_sticky(keepHi, keepLo, 1u); // dropped bit is guaranteed 0 here
|
|
1271
|
-
keepHi = sh.hi; keepLo = sh.lo;
|
|
1272
|
-
resultExpBase = resultExpBase + 1;
|
|
1911
|
+
workgroupBarrier();
|
|
1273
1912
|
}
|
|
1274
1913
|
|
|
1275
|
-
|
|
1276
|
-
|
|
1277
|
-
|
|
1278
|
-
|
|
1279
|
-
|
|
1280
|
-
|
|
1914
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
1915
|
+
let row = blockRow + threadRow * TM + resIdxM;
|
|
1916
|
+
if (row < params.m) {
|
|
1917
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
1918
|
+
let col = blockCol + threadCol * TN + resIdxN;
|
|
1919
|
+
if (col < params.n) {
|
|
1920
|
+
let cIdx = row * params.ldc + col;
|
|
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);
|
|
1925
|
+
}
|
|
1926
|
+
}
|
|
1281
1927
|
}
|
|
1282
|
-
return Fields(resultSign, u32(rawExpFinal), keepHi & 0xfffffu, keepLo);
|
|
1283
|
-
}
|
|
1284
|
-
return Fields(resultSign, 0u, keepHi & 0xfffffu, keepLo); // subnormal result
|
|
1285
|
-
}
|
|
1286
|
-
|
|
1287
|
-
// Packed-in/Packed-out convenience wrapper around addFields \u2014 encodes once,
|
|
1288
|
-
// after the math, rather than addFields itself needing to know about Packed.
|
|
1289
|
-
fn computeSum(a: Fields, b: Fields) -> Packed {
|
|
1290
|
-
let f = addFields(a, b);
|
|
1291
|
-
return encode(f.sign, f.rawExp, f.mantissaHi, f.lo);
|
|
1292
|
-
}
|
|
1293
|
-
`});var ye,ve=W(()=>{ye=`// Double-double arithmetic via Dekker's algorithm \u2014 an alternative to
|
|
1294
|
-
// f64add.wgsl's bit-exact IEEE-754 emulation. Doesn't touch that path.
|
|
1295
|
-
//
|
|
1296
|
-
// A double-double number is a pair (hi, lo) of f32 with hi+lo approximating
|
|
1297
|
-
// a higher-precision value, hi holding the leading bits and lo the rounding
|
|
1298
|
-
// error hi lost. ~48 bits of mantissa vs f32's 24, less than real f64's 52.
|
|
1299
|
-
//
|
|
1300
|
-
// No bindings, no entry point \u2014 a helper library, concatenated with a
|
|
1301
|
-
// consumer's own bindings/entry point by getPipeline (WGSL has no #include).
|
|
1302
|
-
|
|
1303
|
-
struct DD {
|
|
1304
|
-
hi: f32,
|
|
1305
|
-
lo: f32,
|
|
1306
|
-
}
|
|
1307
|
-
|
|
1308
|
-
// |a| for a double-double pair. Negation is exact (no rounding), so this is
|
|
1309
|
-
// just a sign flip on both components \u2014 hi alone determines the pair's sign.
|
|
1310
|
-
fn ddAbs(a: DD) -> DD {
|
|
1311
|
-
if (a.hi < 0.0) {
|
|
1312
|
-
return DD(-a.hi, -a.lo);
|
|
1313
1928
|
}
|
|
1314
|
-
return a;
|
|
1315
|
-
}
|
|
1316
|
-
|
|
1317
|
-
// \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
|
|
1318
|
-
//
|
|
1319
|
-
// twoSum/fastTwoSum's error term \`e\` should be nonzero (that's the point \u2014
|
|
1320
|
-
// \`s\` is rounded). A front-end-optimizer bug zeros it anyway, on both NVIDIA
|
|
1321
|
-
// and Mesa (ANV + llvmpipe), via two different mechanisms needing two fixes:
|
|
1322
|
-
// bitcast-based subtraction (\`fsub\`/\`negf\`, fixes NVIDIA) and materializing
|
|
1323
|
-
// the sum through workgroup memory + workgroupBarrier() (fixes Mesa). Only
|
|
1324
|
-
// both together (ddAddProtected) is verified correct everywhere \u2014 the plain
|
|
1325
|
-
// twoSum/fastTwoSum/ddAdd below are reference-only, not safe to use.
|
|
1326
|
-
fn negf(x: f32) -> f32 {
|
|
1327
|
-
return bitcast<f32>(bitcast<u32>(x) ^ 0x80000000u);
|
|
1328
|
-
}
|
|
1329
|
-
fn fsub(a: f32, b: f32) -> f32 {
|
|
1330
|
-
return a + negf(b);
|
|
1331
|
-
}
|
|
1332
|
-
|
|
1333
|
-
// Knuth/M\xF8ller's TwoSum: s = fl(a+b), e = exact rounding error, a+b == s+e.
|
|
1334
|
-
// Works for any a, b. UNPROTECTED \u2014 see header above.
|
|
1335
|
-
fn twoSum(a: f32, b: f32) -> DD {
|
|
1336
|
-
let s = a + b;
|
|
1337
|
-
let v = s - a;
|
|
1338
|
-
let e = (a - (s - v)) + (b - v);
|
|
1339
|
-
return DD(s, e);
|
|
1340
|
-
}
|
|
1341
|
-
|
|
1342
|
-
// Dekker's Fast-Two-Sum: same contract, but only correct when |a| >= |b|.
|
|
1343
|
-
// UNPROTECTED \u2014 see header above.
|
|
1344
|
-
fn fastTwoSum(a: f32, b: f32) -> DD {
|
|
1345
|
-
let s = a + b;
|
|
1346
|
-
let e = b - (s - a);
|
|
1347
|
-
return DD(s, e);
|
|
1348
|
-
}
|
|
1349
|
-
|
|
1350
|
-
// Double-double addition (Dekker's Add2). UNPROTECTED \u2014 see header above.
|
|
1351
|
-
fn ddAdd(a: DD, b: DD) -> DD {
|
|
1352
|
-
let s = twoSum(a.hi, b.hi);
|
|
1353
|
-
let loSum = a.lo + b.lo;
|
|
1354
|
-
return fastTwoSum(s.hi, s.lo + loSum);
|
|
1355
1929
|
}
|
|
1356
|
-
|
|
1357
|
-
//
|
|
1930
|
+
`});var ee,Ot=V(()=>{ee=`// sgemm_large: C = alpha * op(A) * op(B) + beta * C \u2014 large-tile half of
|
|
1931
|
+
// the two-tier autotuned dispatch (see sgemm.mjs and sgemm_small.wgsl).
|
|
1932
|
+
// BM=BN=64, BK=8, TM=8, TN=4 (128 threads/workgroup) \u2014 the kernel 9
|
|
1933
|
+
// autotuning winner (temp/autotune_sweep.mjs, temp/gen_sweep_kernel.mjs,
|
|
1934
|
+
// swept BM/BN/BK/TM/TN and warp-tiled variants), +69% over the old BM=32
|
|
1935
|
+
// single-tier baseline at n=512, +84% at n=1024. But BM=64 loses to BM=32
|
|
1936
|
+
// below a 6x6=36 workgroup grid (not enough workgroups to fill the GPU at
|
|
1937
|
+
// that tile size), hence the two-tier split rather than one global config.
|
|
1358
1938
|
//
|
|
1359
|
-
//
|
|
1360
|
-
//
|
|
1361
|
-
//
|
|
1362
|
-
//
|
|
1363
|
-
//
|
|
1364
|
-
//
|
|
1365
|
-
|
|
1366
|
-
|
|
1367
|
-
|
|
1368
|
-
|
|
1369
|
-
|
|
1370
|
-
|
|
1371
|
-
|
|
1372
|
-
|
|
1373
|
-
|
|
1374
|
-
|
|
1375
|
-
|
|
1376
|
-
|
|
1377
|
-
|
|
1378
|
-
|
|
1379
|
-
|
|
1380
|
-
|
|
1381
|
-
|
|
1382
|
-
let e = fsub(b, fsub(s, a));
|
|
1383
|
-
return DD(s, e);
|
|
1384
|
-
}
|
|
1385
|
-
|
|
1386
|
-
// Protected double-double addition \u2014 same contract as ddAdd, but exact.
|
|
1387
|
-
fn ddAddProtected(a: DD, b: DD, threadSlot: u32) -> DD {
|
|
1388
|
-
let s = twoSumProtected(a.hi, b.hi, threadSlot);
|
|
1389
|
-
let loSum = a.lo + b.lo;
|
|
1390
|
-
return fastTwoSumProtected(s.hi, s.lo + loSum, threadSlot);
|
|
1391
|
-
}
|
|
1392
|
-
`});var Ee,_e=W(()=>{Ee=`// dasum: sum(|x[i]|), double-double (Dekker). Same ILP=4 shape as sasum.wgsl;
|
|
1393
|
-
// see f64/dekker.wgsl for ddAddProtected and why plain ddAdd isn't safe.
|
|
1394
|
-
// GpuVector input isn't pre-abs'd, so ddAbs() applies unconditionally below.
|
|
1395
|
-
|
|
1396
|
-
@group(0) @binding(0) var<storage, read> xHi: array<f32>;
|
|
1397
|
-
@group(0) @binding(1) var<storage, read> xLo: array<f32>;
|
|
1398
|
-
@group(0) @binding(2) var<storage, read_write> partialsHi: array<f32>;
|
|
1399
|
-
@group(0) @binding(3) var<storage, read_write> partialsLo: array<f32>;
|
|
1400
|
-
@group(0) @binding(4) var<uniform> params: Params;
|
|
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.
|
|
1945
|
+
|
|
1946
|
+
const BM: u32 = 64u;
|
|
1947
|
+
const BN: u32 = 64u;
|
|
1948
|
+
const BK: u32 = 8u;
|
|
1949
|
+
const TM: u32 = 8u;
|
|
1950
|
+
const TN: u32 = 4u;
|
|
1951
|
+
const THREADS_X: u32 = BN / TN;
|
|
1952
|
+
const THREADS_Y: u32 = BM / TM;
|
|
1953
|
+
const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 128
|
|
1954
|
+
const STRIDE_A: u32 = NUM_THREADS / BK; // rows of As covered per load-loop step
|
|
1955
|
+
const STRIDE_B: u32 = NUM_THREADS / BN; // rows of Bs covered per load-loop step
|
|
1956
|
+
|
|
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>;
|
|
1401
1962
|
|
|
1402
1963
|
struct Params {
|
|
1403
|
-
|
|
1404
|
-
|
|
1405
|
-
|
|
1406
|
-
|
|
1407
|
-
|
|
1408
|
-
|
|
1409
|
-
|
|
1410
|
-
|
|
1411
|
-
|
|
1412
|
-
|
|
1413
|
-
|
|
1414
|
-
|
|
1415
|
-
|
|
1416
|
-
|
|
1417
|
-
)
|
|
1418
|
-
|
|
1419
|
-
|
|
1420
|
-
|
|
1421
|
-
|
|
1422
|
-
|
|
1423
|
-
|
|
1424
|
-
|
|
1425
|
-
|
|
1426
|
-
|
|
1427
|
-
|
|
1428
|
-
let
|
|
1429
|
-
|
|
1430
|
-
|
|
1431
|
-
|
|
1432
|
-
|
|
1433
|
-
|
|
1434
|
-
|
|
1435
|
-
|
|
1436
|
-
|
|
1437
|
-
|
|
1438
|
-
|
|
1439
|
-
|
|
1440
|
-
|
|
1441
|
-
|
|
1442
|
-
|
|
1443
|
-
|
|
1444
|
-
|
|
1445
|
-
|
|
1446
|
-
|
|
1447
|
-
|
|
1448
|
-
|
|
1449
|
-
|
|
1450
|
-
|
|
1451
|
-
|
|
1452
|
-
|
|
1453
|
-
|
|
1454
|
-
|
|
1964
|
+
m: u32,
|
|
1965
|
+
n: u32,
|
|
1966
|
+
k: u32,
|
|
1967
|
+
alpha: f32,
|
|
1968
|
+
beta: f32,
|
|
1969
|
+
lda: u32,
|
|
1970
|
+
ldb: u32,
|
|
1971
|
+
ldc: u32,
|
|
1972
|
+
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
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
|
|
1976
|
+
}
|
|
1977
|
+
|
|
1978
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
1979
|
+
|
|
1980
|
+
var<workgroup> As: array<f32, BM * BK>;
|
|
1981
|
+
var<workgroup> Bs: array<f32, BK * BN>;
|
|
1982
|
+
|
|
1983
|
+
@compute @workgroup_size(THREADS_X, THREADS_Y)
|
|
1984
|
+
fn main(
|
|
1985
|
+
@builtin(workgroup_id) wid: vec3u,
|
|
1986
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1987
|
+
@builtin(local_invocation_index) tid: u32,
|
|
1988
|
+
) {
|
|
1989
|
+
let blockRow = wid.y * BM;
|
|
1990
|
+
let blockCol = wid.x * BN;
|
|
1991
|
+
let threadCol = lid.x;
|
|
1992
|
+
let threadRow = lid.y;
|
|
1993
|
+
|
|
1994
|
+
// Load indices, independent of the compute thread shape \u2014 a loop since
|
|
1995
|
+
// NUM_THREADS doesn't match the tile size 1:1 at this config.
|
|
1996
|
+
let innerRowA = tid / BK;
|
|
1997
|
+
let innerColA = tid % BK;
|
|
1998
|
+
let innerRowB = tid / BN;
|
|
1999
|
+
let innerColB = tid % BN;
|
|
2000
|
+
|
|
2001
|
+
var threadResults: array<f32, TM * TN>;
|
|
2002
|
+
for (var i = 0u; i < TM * TN; i++) {
|
|
2003
|
+
threadResults[i] = 0.0;
|
|
2004
|
+
}
|
|
2005
|
+
var regM: array<f32, TM>;
|
|
2006
|
+
var regN: array<f32, TN>;
|
|
2007
|
+
|
|
2008
|
+
let numTiles = (params.k + BK - 1u) / BK;
|
|
2009
|
+
for (var t = 0u; t < numTiles; t++) {
|
|
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
|
+
}
|
|
2054
|
+
}
|
|
1455
2055
|
|
|
1456
|
-
|
|
1457
|
-
|
|
1458
|
-
|
|
1459
|
-
|
|
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
|
+
}
|
|
2099
|
+
}
|
|
2100
|
+
|
|
2101
|
+
workgroupBarrier();
|
|
2102
|
+
|
|
2103
|
+
for (var dotIdx = 0u; dotIdx < BK; dotIdx++) {
|
|
2104
|
+
for (var i = 0u; i < TM; i++) {
|
|
2105
|
+
regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
|
|
2106
|
+
}
|
|
2107
|
+
for (var i = 0u; i < TN; i++) {
|
|
2108
|
+
regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
|
|
2109
|
+
}
|
|
2110
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
2111
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
2112
|
+
threadResults[resIdxM * TN + resIdxN] += regM[resIdxM] * regN[resIdxN];
|
|
2113
|
+
}
|
|
2114
|
+
}
|
|
2115
|
+
}
|
|
1460
2116
|
|
|
1461
|
-
// Inactive threads combine against a throwaway partner and discard it
|
|
1462
|
-
// (ddAddProtected must be called unconditionally by every thread).
|
|
1463
|
-
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
1464
|
-
let partner = select(lid.x, lid.x + s, lid.x < s);
|
|
1465
|
-
let combined = ddAddProtected(tile[lid.x], tile[partner], lid.x);
|
|
1466
|
-
workgroupBarrier(); // all threads must read tile[] above before any write below
|
|
1467
|
-
if (lid.x < s) { tile[lid.x] = combined; }
|
|
1468
2117
|
workgroupBarrier();
|
|
1469
2118
|
}
|
|
1470
2119
|
|
|
1471
|
-
|
|
1472
|
-
|
|
1473
|
-
|
|
2120
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
2121
|
+
let row = blockRow + threadRow * TM + resIdxM;
|
|
2122
|
+
if (row < params.m) {
|
|
2123
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
2124
|
+
let col = blockCol + threadCol * TN + resIdxN;
|
|
2125
|
+
if (col < params.n) {
|
|
2126
|
+
let cIdx = row * params.ldc + col;
|
|
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);
|
|
2131
|
+
}
|
|
2132
|
+
}
|
|
2133
|
+
}
|
|
1474
2134
|
}
|
|
1475
2135
|
}
|
|
1476
|
-
`});var
|
|
1477
|
-
//
|
|
1478
|
-
//
|
|
1479
|
-
//
|
|
1480
|
-
|
|
1481
|
-
|
|
1482
|
-
|
|
1483
|
-
|
|
1484
|
-
|
|
1485
|
-
|
|
1486
|
-
|
|
1487
|
-
|
|
1488
|
-
|
|
1489
|
-
|
|
1490
|
-
|
|
1491
|
-
// those entries being mathematically implied zero.
|
|
2136
|
+
`});var ue,Kt=V(()=>{ue=`// sgemmtr_small: C := uplo(alpha * op(A) * op(B) + beta * C) \u2014 small-tile
|
|
2137
|
+
// half of a two-tier dispatch, identical to sgemm_small.wgsl except the
|
|
2138
|
+
// final output write is gated to one triangle of C by \`uplo\` \u2014 see
|
|
2139
|
+
// sgemmtr_large.wgsl for the full rationale (shared by both tiers).
|
|
2140
|
+
|
|
2141
|
+
const BM: u32 = 32u;
|
|
2142
|
+
const BN: u32 = 32u;
|
|
2143
|
+
const BK: u32 = 8u;
|
|
2144
|
+
const TM: u32 = 2u;
|
|
2145
|
+
const TN: u32 = 2u;
|
|
2146
|
+
const THREADS_X: u32 = BN / TN;
|
|
2147
|
+
const THREADS_Y: u32 = BM / TM;
|
|
2148
|
+
const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 256
|
|
2149
|
+
const STRIDE_A: u32 = NUM_THREADS / BK;
|
|
2150
|
+
const STRIDE_B: u32 = NUM_THREADS / BN;
|
|
1492
2151
|
|
|
1493
2152
|
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1494
|
-
@group(0) @binding(1) var<storage,
|
|
2153
|
+
@group(0) @binding(1) var<storage, read> B: array<f32>;
|
|
2154
|
+
@group(0) @binding(2) var<storage, read_write> C: array<f32>;
|
|
1495
2155
|
|
|
1496
2156
|
struct Params {
|
|
1497
|
-
|
|
1498
|
-
|
|
1499
|
-
|
|
1500
|
-
|
|
1501
|
-
|
|
2157
|
+
m: u32,
|
|
2158
|
+
n: u32,
|
|
2159
|
+
k: u32,
|
|
2160
|
+
alpha: f32,
|
|
2161
|
+
beta: f32,
|
|
2162
|
+
lda: u32,
|
|
2163
|
+
ldb: u32,
|
|
2164
|
+
ldc: u32,
|
|
2165
|
+
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
2166
|
+
transB: u32,
|
|
2167
|
+
uplo: u32, // 0 = lower (col <= row), 1 = upper (col >= row)
|
|
1502
2168
|
}
|
|
1503
2169
|
|
|
1504
|
-
@group(0) @binding(
|
|
2170
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1505
2171
|
|
|
1506
|
-
|
|
1507
|
-
|
|
1508
|
-
var<workgroup> scratch: array<f32, 64>;
|
|
2172
|
+
var<workgroup> As: array<f32, BM * BK>;
|
|
2173
|
+
var<workgroup> Bs: array<f32, BK * BN>;
|
|
1509
2174
|
|
|
1510
|
-
|
|
1511
|
-
|
|
1512
|
-
|
|
1513
|
-
|
|
1514
|
-
|
|
2175
|
+
@compute @workgroup_size(THREADS_X, THREADS_Y)
|
|
2176
|
+
fn main(
|
|
2177
|
+
@builtin(workgroup_id) wid: vec3u,
|
|
2178
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
2179
|
+
@builtin(local_invocation_index) tid: u32,
|
|
2180
|
+
) {
|
|
2181
|
+
let blockRow = wid.y * BM;
|
|
2182
|
+
let blockCol = wid.x * BN;
|
|
2183
|
+
let threadCol = lid.x;
|
|
2184
|
+
let threadRow = lid.y;
|
|
2185
|
+
|
|
2186
|
+
let innerRowA = tid / BK;
|
|
2187
|
+
let innerColA = tid % BK;
|
|
2188
|
+
let innerRowB = tid / BN;
|
|
2189
|
+
let innerColB = tid % BN;
|
|
2190
|
+
|
|
2191
|
+
var threadResults: array<f32, TM * TN>;
|
|
2192
|
+
for (var i = 0u; i < TM * TN; i++) {
|
|
2193
|
+
threadResults[i] = 0.0;
|
|
2194
|
+
}
|
|
2195
|
+
var regM: array<f32, TM>;
|
|
2196
|
+
var regN: array<f32, TN>;
|
|
2197
|
+
|
|
2198
|
+
let numTiles = (params.k + BK - 1u) / BK;
|
|
2199
|
+
for (var t = 0u; t < numTiles; t++) {
|
|
2200
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
2201
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
2202
|
+
let gColA = t * BK + innerColA;
|
|
2203
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
2204
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
2205
|
+
}
|
|
2206
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
2207
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
2208
|
+
let gColB = blockCol + innerColB;
|
|
2209
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
2210
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
2211
|
+
}
|
|
2212
|
+
|
|
2213
|
+
workgroupBarrier();
|
|
2214
|
+
|
|
2215
|
+
for (var dotIdx = 0u; dotIdx < BK; dotIdx++) {
|
|
2216
|
+
for (var i = 0u; i < TM; i++) {
|
|
2217
|
+
regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
|
|
2218
|
+
}
|
|
2219
|
+
for (var i = 0u; i < TN; i++) {
|
|
2220
|
+
regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
|
|
2221
|
+
}
|
|
2222
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
2223
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
2224
|
+
threadResults[resIdxM * TN + resIdxN] += regM[resIdxM] * regN[resIdxN];
|
|
2225
|
+
}
|
|
2226
|
+
}
|
|
2227
|
+
}
|
|
2228
|
+
|
|
2229
|
+
workgroupBarrier();
|
|
2230
|
+
}
|
|
2231
|
+
|
|
2232
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
2233
|
+
let row = blockRow + threadRow * TM + resIdxM;
|
|
2234
|
+
if (row < params.m) {
|
|
2235
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
2236
|
+
let col = blockCol + threadCol * TN + resIdxN;
|
|
2237
|
+
let inTriangle = select(col >= row, col <= row, params.uplo == 0u);
|
|
2238
|
+
if (col < params.n && inTriangle) {
|
|
2239
|
+
let cIdx = row * params.ldc + col;
|
|
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);
|
|
2244
|
+
}
|
|
2245
|
+
}
|
|
2246
|
+
}
|
|
1515
2247
|
}
|
|
1516
2248
|
}
|
|
2249
|
+
`});var le,Vt=V(()=>{le=`// sgemmtr_large: C := uplo(alpha * op(A) * op(B) + beta * C) \u2014 large-tile
|
|
2250
|
+
// half of a two-tier dispatch, identical to sgemm_large.wgsl (see that file
|
|
2251
|
+
// for the BM/BN/BK/TM/TN autotuning rationale) except the final output write
|
|
2252
|
+
// is gated to one triangle of C by \`uplo\`, the same convention ssyr/ssyr2
|
|
2253
|
+
// use (0 = lower: col <= row, 1 = upper: col >= row). Every other element of
|
|
2254
|
+
// C \u2014 including inside the compute loop, where the full tile is still
|
|
2255
|
+
// computed regardless of uplo, only the write is masked \u2014 is left untouched.
|
|
2256
|
+
// gemmtr's uplo(C) test is a plain row/col comparison over the full m\xD7n
|
|
2257
|
+
// grid, well-defined even when m != n (not restricted to square C).
|
|
2258
|
+
|
|
2259
|
+
const BM: u32 = 64u;
|
|
2260
|
+
const BN: u32 = 64u;
|
|
2261
|
+
const BK: u32 = 8u;
|
|
2262
|
+
const TM: u32 = 8u;
|
|
2263
|
+
const TN: u32 = 4u;
|
|
2264
|
+
const THREADS_X: u32 = BN / TN;
|
|
2265
|
+
const THREADS_Y: u32 = BM / TM;
|
|
2266
|
+
const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 128
|
|
2267
|
+
const STRIDE_A: u32 = NUM_THREADS / BK; // rows of As covered per load-loop step
|
|
2268
|
+
const STRIDE_B: u32 = NUM_THREADS / BN; // rows of Bs covered per load-loop step
|
|
1517
2269
|
|
|
1518
|
-
@
|
|
1519
|
-
|
|
1520
|
-
|
|
1521
|
-
@builtin(local_invocation_id) lid: vec3u,
|
|
1522
|
-
) {
|
|
1523
|
-
let col = wgid.x;
|
|
1524
|
-
let blockIndex = wgid.y;
|
|
1525
|
-
let blockStart = blockIndex * BLOCK_SIZE;
|
|
1526
|
-
var blockEnd = blockStart + BLOCK_SIZE;
|
|
1527
|
-
if (blockEnd > params.n) { blockEnd = params.n; }
|
|
1528
|
-
let blockLen = blockEnd - blockStart;
|
|
2270
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
2271
|
+
@group(0) @binding(1) var<storage, read> B: array<f32>;
|
|
2272
|
+
@group(0) @binding(2) var<storage, read_write> C: array<f32>;
|
|
1529
2273
|
|
|
1530
|
-
|
|
2274
|
+
struct Params {
|
|
2275
|
+
m: u32,
|
|
2276
|
+
n: u32,
|
|
2277
|
+
k: u32,
|
|
2278
|
+
alpha: f32,
|
|
2279
|
+
beta: f32,
|
|
2280
|
+
lda: u32,
|
|
2281
|
+
ldb: u32,
|
|
2282
|
+
ldc: u32,
|
|
2283
|
+
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
2284
|
+
transB: u32,
|
|
2285
|
+
uplo: u32, // 0 = lower (col <= row), 1 = upper (col >= row)
|
|
2286
|
+
}
|
|
1531
2287
|
|
|
1532
|
-
|
|
1533
|
-
let forward = (params.trans == 0u) == (params.uplo == 0u);
|
|
2288
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
1534
2289
|
|
|
1535
|
-
|
|
1536
|
-
|
|
1537
|
-
|
|
2290
|
+
var<workgroup> As: array<f32, BM * BK>;
|
|
2291
|
+
var<workgroup> Bs: array<f32, BK * BN>;
|
|
2292
|
+
|
|
2293
|
+
@compute @workgroup_size(THREADS_X, THREADS_Y)
|
|
2294
|
+
fn main(
|
|
2295
|
+
@builtin(workgroup_id) wid: vec3u,
|
|
2296
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
2297
|
+
@builtin(local_invocation_index) tid: u32,
|
|
2298
|
+
) {
|
|
2299
|
+
let blockRow = wid.y * BM;
|
|
2300
|
+
let blockCol = wid.x * BN;
|
|
2301
|
+
let threadCol = lid.x;
|
|
2302
|
+
let threadRow = lid.y;
|
|
2303
|
+
|
|
2304
|
+
// Load indices, independent of the compute thread shape \u2014 a loop since
|
|
2305
|
+
// NUM_THREADS doesn't match the tile size 1:1 at this config.
|
|
2306
|
+
let innerRowA = tid / BK;
|
|
2307
|
+
let innerColA = tid % BK;
|
|
2308
|
+
let innerRowB = tid / BN;
|
|
2309
|
+
let innerColB = tid % BN;
|
|
2310
|
+
|
|
2311
|
+
var threadResults: array<f32, TM * TN>;
|
|
2312
|
+
for (var i = 0u; i < TM * TN; i++) {
|
|
2313
|
+
threadResults[i] = 0.0;
|
|
2314
|
+
}
|
|
2315
|
+
var regM: array<f32, TM>;
|
|
2316
|
+
var regN: array<f32, TN>;
|
|
2317
|
+
|
|
2318
|
+
let numTiles = (params.k + BK - 1u) / BK;
|
|
2319
|
+
for (var t = 0u; t < numTiles; t++) {
|
|
2320
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
2321
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
2322
|
+
let gColA = t * BK + innerColA;
|
|
2323
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
2324
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
1538
2325
|
}
|
|
1539
|
-
|
|
1540
|
-
|
|
1541
|
-
|
|
2326
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
2327
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
2328
|
+
let gColB = blockCol + innerColB;
|
|
2329
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
2330
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
1542
2331
|
}
|
|
1543
|
-
}
|
|
1544
|
-
storageBarrier();
|
|
1545
|
-
workgroupBarrier();
|
|
1546
2332
|
|
|
1547
|
-
|
|
1548
|
-
for (var step = 0u; step < numSteps; step++) {
|
|
1549
|
-
let localRow = select(col - step, col + step, forward);
|
|
1550
|
-
let i = blockStart + localRow;
|
|
2333
|
+
workgroupBarrier();
|
|
1551
2334
|
|
|
1552
|
-
var
|
|
1553
|
-
|
|
1554
|
-
|
|
1555
|
-
acc += readA(i, blockStart + lj) * Ainv[ainvBase + lj * BLOCK_SIZE + col];
|
|
2335
|
+
for (var dotIdx = 0u; dotIdx < BK; dotIdx++) {
|
|
2336
|
+
for (var i = 0u; i < TM; i++) {
|
|
2337
|
+
regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
|
|
1556
2338
|
}
|
|
1557
|
-
|
|
1558
|
-
|
|
1559
|
-
|
|
2339
|
+
for (var i = 0u; i < TN; i++) {
|
|
2340
|
+
regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
|
|
2341
|
+
}
|
|
2342
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
2343
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
2344
|
+
threadResults[resIdxM * TN + resIdxN] += regM[resIdxM] * regN[resIdxN];
|
|
2345
|
+
}
|
|
1560
2346
|
}
|
|
1561
2347
|
}
|
|
1562
2348
|
|
|
1563
|
-
scratch[lid.x] = acc;
|
|
1564
2349
|
workgroupBarrier();
|
|
1565
|
-
|
|
1566
|
-
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
1567
|
-
workgroupBarrier();
|
|
1568
|
-
}
|
|
2350
|
+
}
|
|
1569
2351
|
|
|
1570
|
-
|
|
1571
|
-
|
|
1572
|
-
|
|
1573
|
-
var
|
|
1574
|
-
|
|
1575
|
-
|
|
1576
|
-
|
|
1577
|
-
|
|
2352
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
2353
|
+
let row = blockRow + threadRow * TM + resIdxM;
|
|
2354
|
+
if (row < params.m) {
|
|
2355
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
2356
|
+
let col = blockCol + threadCol * TN + resIdxN;
|
|
2357
|
+
let inTriangle = select(col >= row, col <= row, params.uplo == 0u);
|
|
2358
|
+
if (col < params.n && inTriangle) {
|
|
2359
|
+
let cIdx = row * params.ldc + col;
|
|
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);
|
|
2364
|
+
}
|
|
1578
2365
|
}
|
|
1579
|
-
Ainv[ainvBase + localRow * BLOCK_SIZE + col] = val;
|
|
1580
2366
|
}
|
|
1581
|
-
storageBarrier();
|
|
1582
|
-
workgroupBarrier();
|
|
1583
2367
|
}
|
|
1584
2368
|
}
|
|
1585
|
-
`});var
|
|
1586
|
-
//
|
|
1587
|
-
//
|
|
1588
|
-
//
|
|
1589
|
-
//
|
|
1590
|
-
//
|
|
1591
|
-
// All blockLen rows are computed in parallel within a single workgroup: the
|
|
1592
|
-
// remainder is loaded into workgroup-shared memory once, then each thread
|
|
1593
|
-
// independently computes one full row's dot product from that shared copy.
|
|
1594
|
-
// No further synchronization is needed after the load \u2014 every thread only
|
|
1595
|
-
// reads shared memory from then on (never written again within this call)
|
|
1596
|
-
// and writes a distinct element of x, so there's no cross-thread hazard to
|
|
1597
|
-
// guard against.
|
|
2369
|
+
`});var Ht,zt=V(()=>{Ht=`// symmetrize: Adense := full dense expansion of a symmetric matrix stored
|
|
2370
|
+
// with only its \`uplo\` triangle meaningful (the other triangle is implied
|
|
2371
|
+
// by symmetry: A[i,j] = A[j,i]). A plain element-wise pass, no tiling or
|
|
2372
|
+
// shared memory needed \u2014 used to materialize a dense operand for routines
|
|
2373
|
+
// that read a symmetric matrix as a normal dense gemm input (e.g. ssymm),
|
|
2374
|
+
// rather than teaching the tiled gemm kernel itself to mirror-read.
|
|
1598
2375
|
|
|
1599
|
-
@group(0) @binding(0) var<storage, read>
|
|
1600
|
-
@group(0) @binding(1) var<storage, read_write>
|
|
2376
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
2377
|
+
@group(0) @binding(1) var<storage, read_write> Adense: array<f32>;
|
|
1601
2378
|
|
|
1602
2379
|
struct Params {
|
|
1603
|
-
|
|
1604
|
-
|
|
1605
|
-
|
|
1606
|
-
|
|
2380
|
+
n: u32,
|
|
2381
|
+
lda: u32,
|
|
2382
|
+
ldd: u32, // leading dimension of Adense
|
|
2383
|
+
uplo: u32, // 0 = lower (stored where col <= row), 1 = upper (col >= row)
|
|
1607
2384
|
}
|
|
1608
2385
|
|
|
1609
2386
|
@group(0) @binding(2) var<uniform> params: Params;
|
|
1610
2387
|
|
|
1611
|
-
|
|
1612
|
-
|
|
1613
|
-
|
|
1614
|
-
|
|
1615
|
-
|
|
1616
|
-
|
|
1617
|
-
|
|
1618
|
-
if (lid.x < blockLen) {
|
|
1619
|
-
xLocal[lid.x] = x[(params.blockStart + lid.x) * params.incx];
|
|
2388
|
+
@compute @workgroup_size(8, 8)
|
|
2389
|
+
fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
2390
|
+
let row = gid.y;
|
|
2391
|
+
let col = gid.x;
|
|
2392
|
+
if (row >= params.n || col >= params.n) {
|
|
2393
|
+
return;
|
|
1620
2394
|
}
|
|
1621
|
-
workgroupBarrier();
|
|
1622
|
-
|
|
1623
|
-
if (lid.x >= blockLen) { return; }
|
|
1624
2395
|
|
|
1625
|
-
let
|
|
1626
|
-
|
|
1627
|
-
|
|
1628
|
-
acc += Ainv[ainvBase + lid.x * BLOCK_SIZE + j] * xLocal[j];
|
|
1629
|
-
}
|
|
1630
|
-
x[(params.blockStart + lid.x) * params.incx] = acc;
|
|
2396
|
+
let isStored = select(col >= row, col <= row, params.uplo == 0u);
|
|
2397
|
+
let srcIdx = select(col * params.lda + row, row * params.lda + col, isStored);
|
|
2398
|
+
Adense[row * params.ldd + col] = A[srcIdx];
|
|
1631
2399
|
}
|
|
1632
|
-
`});var
|
|
1633
|
-
//
|
|
1634
|
-
//
|
|
1635
|
-
// No diag/masking needed: this region never touches the diagonal.
|
|
2400
|
+
`});var Xt,Yt=V(()=>{Xt=`// triangularize: Adense := dense expansion of op(A) (A or A^T per \`trans\`),
|
|
2401
|
+
// zero-filling the unstored triangle (exact for a matmul) so strmm can reuse
|
|
2402
|
+
// sgemm's kernel unchanged. \`diag=1\` substitutes 1.0 on the diagonal.
|
|
1636
2403
|
|
|
1637
2404
|
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1638
|
-
@group(0) @binding(1) var<storage, read_write>
|
|
2405
|
+
@group(0) @binding(1) var<storage, read_write> Adense: array<f32>;
|
|
1639
2406
|
|
|
1640
2407
|
struct Params {
|
|
1641
|
-
n:
|
|
1642
|
-
|
|
1643
|
-
|
|
1644
|
-
|
|
1645
|
-
|
|
1646
|
-
|
|
1647
|
-
blockEnd: u32, // exclusive
|
|
2408
|
+
n: u32,
|
|
2409
|
+
lda: u32,
|
|
2410
|
+
ldd: u32, // leading dimension of Adense
|
|
2411
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
2412
|
+
trans: u32, // 0 = no-transpose (op(A) = A), 1 = transpose (op(A) = A^T)
|
|
2413
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
1648
2414
|
}
|
|
1649
2415
|
|
|
1650
2416
|
@group(0) @binding(2) var<uniform> params: Params;
|
|
1651
2417
|
|
|
1652
|
-
|
|
1653
|
-
|
|
1654
|
-
|
|
1655
|
-
|
|
1656
|
-
|
|
1657
|
-
|
|
1658
|
-
|
|
1659
|
-
@builtin(num_workgroups) nwg: vec3u,
|
|
1660
|
-
) {
|
|
1661
|
-
// forward: remaining rows are [blockEnd,n); backward: [0,blockStart).
|
|
1662
|
-
let forward = (params.trans == 0u) == (params.uplo == 0u);
|
|
2418
|
+
@compute @workgroup_size(8, 8)
|
|
2419
|
+
fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
2420
|
+
let row = gid.y;
|
|
2421
|
+
let col = gid.x;
|
|
2422
|
+
if (row >= params.n || col >= params.n) {
|
|
2423
|
+
return;
|
|
2424
|
+
}
|
|
1663
2425
|
|
|
1664
|
-
|
|
1665
|
-
|
|
2426
|
+
if (row == col) {
|
|
2427
|
+
Adense[row * params.ldd + col] = select(A[row * params.lda + row], 1.0, params.diag == 1u);
|
|
2428
|
+
return;
|
|
2429
|
+
}
|
|
1666
2430
|
|
|
1667
|
-
|
|
1668
|
-
|
|
1669
|
-
|
|
2431
|
+
var isMeaningful: bool;
|
|
2432
|
+
var srcRow: u32;
|
|
2433
|
+
var srcCol: u32;
|
|
2434
|
+
if (params.trans == 0u) {
|
|
2435
|
+
isMeaningful = select(col >= row, col <= row, params.uplo == 0u);
|
|
2436
|
+
srcRow = row; srcCol = col;
|
|
1670
2437
|
} else {
|
|
1671
|
-
|
|
1672
|
-
|
|
2438
|
+
isMeaningful = select(col <= row, col >= row, params.uplo == 0u);
|
|
2439
|
+
srcRow = col; srcCol = row;
|
|
1673
2440
|
}
|
|
1674
|
-
|
|
1675
|
-
if (rangeStart >= rangeEnd) { return; }
|
|
1676
|
-
let count = rangeEnd - rangeStart;
|
|
1677
2441
|
|
|
1678
|
-
|
|
1679
|
-
|
|
2442
|
+
Adense[row * params.ldd + col] = select(0.0, A[srcRow * params.lda + srcCol], isMeaningful);
|
|
2443
|
+
}
|
|
2444
|
+
`});var Zt,$t=V(()=>{Zt=`// block_transfer: gather/scatter/scatter-subtract between a tight (blockLen
|
|
2445
|
+
// x otherLen) block and a sub-range of a strided (any ld, row/col-major)
|
|
2446
|
+
// buffer \u2014 needed since block offsets aren't 256-byte-aligned and block
|
|
2447
|
+
// rows/cols aren't always one contiguous range for copyBufferToBuffer.
|
|
1680
2448
|
|
|
1681
|
-
|
|
1682
|
-
|
|
1683
|
-
if params.trans == 0u {
|
|
1684
|
-
for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
|
|
1685
|
-
acc += A[i * params.lda + j] * x[j * params.incx];
|
|
1686
|
-
}
|
|
1687
|
-
} else {
|
|
1688
|
-
for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
|
|
1689
|
-
acc += A[j * params.lda + i] * x[j * params.incx];
|
|
1690
|
-
}
|
|
1691
|
-
}
|
|
2449
|
+
@group(0) @binding(0) var<storage, read_write> block: array<f32>; // blockLen x otherLen, block[i*otherLen+j]
|
|
2450
|
+
@group(0) @binding(1) var<storage, read_write> strided: array<f32>; // B's or A's own buffer
|
|
1692
2451
|
|
|
1693
|
-
|
|
1694
|
-
|
|
1695
|
-
|
|
1696
|
-
|
|
1697
|
-
|
|
1698
|
-
|
|
1699
|
-
|
|
2452
|
+
struct Params {
|
|
2453
|
+
blockStart: u32,
|
|
2454
|
+
blockLen: u32,
|
|
2455
|
+
otherStart: u32,
|
|
2456
|
+
otherLen: u32,
|
|
2457
|
+
ld: u32,
|
|
2458
|
+
isColMajor: u32, // 0 = row-major addressing, 1 = column-major (row/col swapped)
|
|
2459
|
+
blockIsRow: u32, // 1 = blockStart indexes strided's rows, 0 = its columns
|
|
2460
|
+
mode: u32, // 0 = scatter (strided := block), 1 = scatter_sub (strided -= block), 2 = gather (block := strided)
|
|
2461
|
+
}
|
|
1700
2462
|
|
|
1701
|
-
|
|
1702
|
-
|
|
1703
|
-
|
|
1704
|
-
|
|
2463
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
2464
|
+
|
|
2465
|
+
@compute @workgroup_size(8, 8)
|
|
2466
|
+
fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
2467
|
+
let i = gid.y; // index along the blocked axis, within the block
|
|
2468
|
+
let j = gid.x; // index along the other axis, within the block
|
|
2469
|
+
if (i >= params.blockLen || j >= params.otherLen) {
|
|
2470
|
+
return;
|
|
2471
|
+
}
|
|
2472
|
+
|
|
2473
|
+
let row = select(params.otherStart + j, params.blockStart + i, params.blockIsRow == 1u);
|
|
2474
|
+
let col = select(params.blockStart + i, params.otherStart + j, params.blockIsRow == 1u);
|
|
2475
|
+
let stridedIdx = select(row * params.ld + col, col * params.ld + row, params.isColMajor == 1u);
|
|
2476
|
+
let blockIdx = i * params.otherLen + j;
|
|
2477
|
+
|
|
2478
|
+
if (params.mode == 2u) {
|
|
2479
|
+
block[blockIdx] = strided[stridedIdx];
|
|
2480
|
+
} else if (params.mode == 1u) {
|
|
2481
|
+
strided[stridedIdx] -= block[blockIdx];
|
|
2482
|
+
} else {
|
|
2483
|
+
strided[stridedIdx] = block[blockIdx];
|
|
1705
2484
|
}
|
|
1706
2485
|
}
|
|
1707
|
-
`});var Le={};vr(Le,{shaderSources:()=>Ct});var Ct,je=W(()=>{Fr();Nr();Dr();Tr();Hr();Rr();Or();Qr();Zr();Xr();Yr();re();te();oe();ne();ue();fe();me();pe();we();he();ve();_e();Ge();Be();Se();Ct={"reduction/argmax":Ir,"reduction/sum":Wr,"reduction/sumF64":Mr,sscal:Ur,sswap:Vr,saxpy:Cr,scopy:zr,sdot:qr,sasum:Kr,snrm2:$r,srot:Jr,srotm:ee,isamax:ae,sgemv_n:ie,sgemv_t:se,ssymv:le,strmv:ce,sger:de,ssyr:ge,ssyr2:be,f64add:xe,"f64/dekker":ye,dasum:Ee,strsv_invert_block:Ae,strsv_apply_inverse:ke,strsv_update:Pe}});var Zt={};vr(Zt,{GpuMatrix:()=>V,GpuVector:()=>x,cleanup:()=>kr,dasum:()=>Ve,gpuName:()=>Sr,init:()=>Br,isamax:()=>Oe,randomFloat32Array:()=>Pr,randomFloat64Array:()=>Lr,randomTriangularFloat32Array:()=>jr,sasum:()=>He,saxpy:()=>We,scopy:()=>De,sdot:()=>Te,sgemv:()=>qe,sger:()=>Je,snrm2:()=>Ce,srot:()=>ze,srotm:()=>Qe,sscal:()=>Ie,sswap:()=>Ne,ssymv:()=>Ze,ssyr:()=>rt,ssyr2:()=>et,strmv:()=>Ke,strsv:()=>Ye});function _r(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 Er(){if(!Gr())return{querySet:null,passDescriptor:void 0};let e=U().createQuerySet({type:"timestamp",count:2});return{querySet:e,passDescriptor:{timestampWrites:{querySet:e,beginningOfPassWriteIndex:0,endOfPassWriteIndex:1}}}}function cr(a,e){if(!e)return null;let r=U(),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 P(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 K=null,J=null,Ar=null,mr=!1;async function Br({powerPreference:a="high-performance",benchmark:e=!1}={}){if(K)return K;let r;if(typeof window>"u"){let{create:i,globals:n}=await import("webgpu");Object.assign(globalThis,n),r=i([]),Ar=r}else r=navigator.gpu;if(!r)throw new Error("WebGPU not supported in this environment.");if(J=await r.requestAdapter({powerPreference:a})??await r.requestAdapter(),!J)throw new Error("No WebGPU adapter found.");mr=e;let t=[..._r(J,e).requiredFeatures??[]];return K=await J.requestDevice({requiredFeatures:t}),K.addEventListener("uncapturederror",i=>{console.error("Uncaptured GPU error:",i.error.message)}),K}function kr(){K&&(K.destroy(),K=null),J=null,Ar=null,mr=!1}function Sr(){if(!J)throw new Error("WebGPU adapter not initialized \u2014 call init() first.");let{device:a,description:e}=J.info;return{description:e||"unknown",device:a||"unknown"}}function Gr(){return mr}function U(){if(!K)throw new Error("WebGPU device not initialized \u2014 call init() first.");return K}function m(...a){a.flat().forEach(e=>e.destroy())}function b(a,e="blas-input",r=!1){let o=U(),t=o.limits.maxStorageBufferBindingSize,i=a.byteLength;if(i>t)throw new Error(`Buffer size ${i} bytes exceeds device limit of ${t} bytes.`);let n=r?GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC:GPUBufferUsage.STORAGE,u=o.createBuffer({label:e,size:i,usage:n,mappedAtCreation:!0}),s=a.constructor;return new s(u.getMappedRange()).set(a),u.unmap(),u}function C(a,e="blas-storage"){return U().createBuffer({label:e,size:a,usage:GPUBufferUsage.STORAGE})}function q(a,e="blas-result"){return U().createBuffer({label:e,size:a,usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC})}function _(a,e){let o=U().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 I(a,e="blas-params"){let r=U(),o=a.length*4,t=Math.ceil(o/16)*16,i=new ArrayBuffer(t),n=new DataView(i);a.forEach(({value:s,type:l},f)=>{let c=f*4;if(l==="u32")n.setUint32(c,s,!0);else if(l==="i32")n.setInt32(c,s,!0);else if(l==="f32")n.setFloat32(c,s,!0);else throw new Error(`Unknown param type "${l}". Use "f32", "u32", or "i32".`)});let u=r.createBuffer({label:e,size:t,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST});return r.queue.writeBuffer(u,0,i),u}async function y(a,e=Float32Array){try{await a.mapAsync(GPUMapMode.READ);let r=new e(a.getMappedRange().slice());return a.unmap(),r}finally{a.destroy()}}function er(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 tr(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 x=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}=er(e),i=b(o,"gpu-vector-f64-hi",!0),n=b(t,"gpu-vector-f64-lo",!0);return new a(i,e.length,Float64Array,n)}if(!(e instanceof Float32Array))throw new Error("GpuVector.from expects a Float32Array or Float64Array.");let r=b(e,"gpu-vector",!0);return new a(r,e.length,e.constructor)}async read(){let e=U(),r=e.createCommandEncoder(),o=_(r,this._buf);if(e.queue.submit([r.finish()]),!this._loBuf)return y(o,this.dtype);let t=e.createCommandEncoder(),i=_(t,this._loBuf);e.queue.submit([t.finish()]);let[n,u]=await Promise.all([y(o,Float32Array),y(i,Float32Array)]);return tr(n,u)}destroy(){this._buf.destroy(),this._loBuf&&this._loBuf.destroy()}};var V=class a{constructor(e,r,o,t,i=null,n="row-major"){this._buf=e,this._loBuf=i,this.rows=r,this.cols=o,this.lda=t,this.layout=n}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 n=i==="row-major";if(t===void 0&&(t=n?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 u=n?o:r;if(!Number.isInteger(t)||t<u)throw new Error(`lda must be an integer >= ${n?"cols":"rows"}.`);let s=n?r:o;if(e.length<s*t)throw new Error("data does not have enough elements for the given rows, cols, and lda.");if(e instanceof Float64Array){let f=s*t,{hi:c,lo:p}=er(e.subarray(0,f)),d=b(c,"gpu-matrix-f64-hi",!0),g=b(p,"gpu-matrix-f64-lo",!0);return new a(d,r,o,t,g,i)}let l=b(e.subarray(0,s*t),"gpu-matrix",!0);return new a(l,r,o,t,null,i)}async read(){let e=U(),r=e.createCommandEncoder(),o=_(r,this._buf);e.queue.submit([r.finish()]);let t=this.layout!=="column-major",i=t?this.rows:this.cols,n=t?this.cols:this.rows;if(this._loBuf){let l=e.createCommandEncoder(),f=_(l,this._loBuf);e.queue.submit([l.finish()]);let[c,p]=await Promise.all([y(o,Float32Array),y(f,Float32Array)]),d=tr(c,p);if(this.lda===n)return d;let g=new Float64Array(i*n);for(let w=0;w<i;w++)g.set(d.subarray(w*this.lda,w*this.lda+n),w*n);return g}let u=await y(o,Float32Array);if(this.lda===n)return u;let s=new Float32Array(i*n);for(let l=0;l<i;l++)s.set(u.subarray(l*this.lda,l*this.lda+n),l*n);return s}destroy(){this._buf.destroy(),this._loBuf&&this._loBuf.destroy()}};function Pr(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 Lr(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 jr(a,e,r="lower",o=-1,t=1,i=5,n=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 u=new Float32Array(a*e);for(let s=0;s<a;s++){for(let l=0;l<a;l++){if(s===l)continue;(r==="lower"?l<s:l>s)&&(u[s*e+l]=o+Math.random()*(t-o))}u[s*e+s]=i+Math.random()*(n-i)}return u}function A(a,e,r=0){let o=U(),t=e.map((i,n)=>({binding:r+n,resource:i instanceof GPUBuffer?{buffer:i}:i}));return o.createBindGroup({layout:a,entries:t})}var wt=new WeakMap;function L(a){U().queue.submit([a.finish()])}function dr(){let a=U(),{querySet:e,passDescriptor:r}=Er();return{commandEncoder:a.createCommandEncoder(),querySet:e,passDescriptor:r}}function ir(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),i.end(),wt.set(a,i)}function F(a,e,r){let{commandEncoder:o,querySet:t,passDescriptor:i}=dr();ir(o,a,e,r,i);let n=cr(o,t);return{commandEncoder:o,ts:n}}var Qt={},pr=new WeakMap;async function B(a,e,r="main"){pr.has(a)||pr.set(a,new Map);let o=pr.get(a),t=Array.isArray(e)?e:[e],i=`${t.join("+")}::${r}`;return o.has(i)||o.set(i,await zt(t,r)),o.get(i)}async function Ot(a){if(typeof process>"u"||!process.versions?.node){let{shaderSources:e}=await Promise.resolve().then(()=>(je(),Le)),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(Qt.url));return e(t(i,`../shaders/${a}.wgsl`),"utf8")}}async function zt(a,e="main"){let r=U(),o=a.join("+"),t=(await Promise.all(a.map(Ot))).join(`
|
|
1708
|
-
`),
|
|
1709
|
-
${
|
|
1710
|
-
`)}`);let s=e==="main"?{module:i}:{module:i,entryPoint:e},l=r.createComputePipeline({label:o,layout:"auto",compute:s});return l._shaderModule=i,l}var qt=64,Fe=8;function O(a,e){let r=U().limits.maxComputeWorkgroupsPerDimension;return e===void 0?Math.min(Math.ceil(a/qt),r):{x:Math.min(Math.ceil(e/Fe),r),y:Math.min(Math.ceil(a/Fe),r)}}async function Ie(a,e,r,o,t){let i=o instanceof x;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 x))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 n=await B(a,"sscal"),u=null,s=null,l=null;try{u=i?o._buf:b(o,"sscal-x",!0),s=I([{value:e,type:"u32"},{value:r,type:"f32"},{value:t,type:"u32"}],"sscal-params");let f=A(n.getBindGroupLayout(0),[u,s]),{commandEncoder:c,ts:p}=F(n,f,O(e));l=i?null:_(c,u),L(c);let d=await P(p);if(i)return d!==void 0?{gpuTimeMs:d}:{};let g=await y(l,Float32Array);return l=null,d!==void 0?{x:g,gpuTimeMs:d}:g}finally{!i&&u&&m(u),s&&m(s),l&&m(l)}}async function Ne(a,e,r,o,t,i){let n=r instanceof x,u=t instanceof x;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 x))throw new Error("x must be a Float32Array or GpuVector.");if(!(t instanceof Float32Array)&&!(t instanceof x))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 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 s=await B(a,"sswap"),l=null,f=null,c=null,p=null,d=null;try{l=n?r._buf:b(r,"sswap-x",!0),f=u?t._buf:b(t,"sswap-y",!0),c=I([{value:e,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"sswap-params");let g=A(s.getBindGroupLayout(0),[l,f,c]),{commandEncoder:w,ts:E}=F(s,g,O(e));p=n?null:_(w,l),d=u?null:_(w,f),L(w);let h=await P(E);if(n&&u)return h!==void 0?{gpuTimeMs:h}:{};let G=await y(p,Float32Array);p=null;let v=await y(d,Float32Array);return d=null,h!==void 0?{x:G,y:v,gpuTimeMs:h}:{x:G,y:v}}finally{!n&&l&&m(l),!u&&f&&m(f),c&&m(c),p&&m(p),d&&m(d)}}async function We(a,e,r,o,t,i,n){let u=o instanceof x,s=i instanceof x;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(e)||!Number.isInteger(t)||!Number.isInteger(n))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||n<=0)throw new Error("incx and incy must be positive.");if(!u&&!(o instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!s&&!(i instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(u!==s)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(e<=0)return s?{}:{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)*n+1)throw new Error("y does not have enough elements for the given n and incy.");let l=await B(a,"saxpy"),f=null,c=null,p=null,d=null;try{f=u?o._buf:b(o,"saxpy-x",!1),c=s?i._buf:b(i,"saxpy-y",!0),p=I([{value:e,type:"u32"},{value:r,type:"f32"},{value:t,type:"u32"},{value:n,type:"u32"}],"saxpy-params");let g=A(l.getBindGroupLayout(0),[f,c,p]),{commandEncoder:w,ts:E}=F(l,g,O(e));d=s?null:_(w,c),L(w);let h=await P(E);if(s&&u)return h!==void 0?{gpuTimeMs:h}:{};let G=await y(d,Float32Array);return d=null,h!==void 0?{y:G,gpuTimeMs:h}:{y:G}}finally{!u&&f&&m(f),!s&&c&&m(c),p&&m(p),d&&m(d)}}async function De(a,e,r,o,t,i){let n=r instanceof x,u=t instanceof x;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(!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 u?{}:{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 s=await B(a,"scopy"),l=null,f=null,c=null,p=null;try{l=n?r._buf:b(r,"scopy-x",!1),f=u?t._buf:b(t,"scopy-y",!0),c=I([{value:e,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"scopy-params");let d=A(s.getBindGroupLayout(0),[l,f,c]),{commandEncoder:g,ts:w}=F(s,d,O(e));p=u?null:_(g,f),L(g);let E=await P(w);if(u&&n)return E!==void 0?{gpuTimeMs:E}:{};let h=await y(p,Float32Array);return p=null,E!==void 0?{y:h,gpuTimeMs:E}:{y:h}}finally{!n&&l&&m(l),!u&&f&&m(f),c&&m(c),p&&m(p)}}var Me=64;async function Te(a,e,r,o,t,i){let n=r instanceof x,u=t instanceof x;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(!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{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 s=await B(a,"sdot"),l=await B(a,"reduction/sum"),f=null,c=null,p=null,d=null,g=null,w=null;try{f=n?r._buf:b(r,"sdot-x",!1),c=u?t._buf:b(t,"sdot-y",!1),p=C(2*Me*4,"sdot-partials"),d=q(4,"sdot-result"),g=I([{value:e,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"sdot-params");let E=A(s.getBindGroupLayout(0),[f,c,p,g]),{commandEncoder:h,ts:G}=F(s,E,2*Me);L(h);let v=A(l.getBindGroupLayout(0),[p,d]),{commandEncoder:k,ts:S}=F(l,v,1);w=_(k,d),L(k);let j=y(w,Float32Array);w=null;let[N,D,M]=await Promise.all([P(G),P(S),j]);return N!==void 0&&D!==void 0?{dot:M[0],gpuTimeMs:N+D}:{dot:M[0]}}finally{!n&&f&&m(f),!u&&c&&m(c),p&&m(p),d&&m(d),g&&m(g),w&&m(w)}}var Ue=64;async function He(a,e,r,o){let t=r instanceof x;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 B(a,"sasum"),n=await B(a,"reduction/sum"),u=null,s=null,l=null,f=null,c=null;try{u=t?r._buf:b(r,"sasum-x",!1),s=C(2*Ue*4,"sasum-partials"),l=q(4,"sasum-result"),f=I([{value:e,type:"u32"},{value:o,type:"u32"}],"sasum-params");let p=A(i.getBindGroupLayout(0),[u,s,f]),{commandEncoder:d,ts:g}=F(i,p,2*Ue);L(d);let w=A(n.getBindGroupLayout(0),[s,l]),{commandEncoder:E,ts:h}=F(n,w,1);c=_(E,l),L(E);let G=y(c,Float32Array);c=null;let[v,k,S]=await Promise.all([P(g),P(h),G]);return v!==void 0&&k!==void 0?{asum:S[0],gpuTimeMs:v+k}:{asum:S[0]}}finally{!t&&u&&m(u),s&&m(s),l&&m(l),f&&m(f),c&&m(c)}}var gr=64;async function Ve(a,e,r,o){let t=r instanceof x;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=await B(a,["f64/dekker","dasum"]),n=await B(a,["f64/dekker","reduction/sumF64"]),u=null,s=null,l=null,f=null,c=null,p=null,d=null,g=null,w=null;try{if(t)u=r._buf,s=r._loBuf;else{let{hi:Z,lo:z}=er(r.map(Math.abs));u=b(Z,"dasum-xHi",!1),s=b(z,"dasum-xLo",!1)}l=C(2*gr*4,"dasum-partialsHi"),f=C(2*gr*4,"dasum-partialsLo"),c=q(4,"dasum-result-hi"),p=q(4,"dasum-result-lo"),d=I([{value:e,type:"u32"},{value:o,type:"u32"}],"dasum-params");let E=A(i.getBindGroupLayout(0),[u,s,l,f,d]),{commandEncoder:h,ts:G}=F(i,E,2*gr);L(h);let v=A(n.getBindGroupLayout(0),[l,f,c,p]),{commandEncoder:k,ts:S}=F(n,v,1);g=_(k,c),w=_(k,p),L(k);let j=y(g,Float32Array),N=y(w,Float32Array);g=null,w=null;let[D,M,T,H]=await Promise.all([P(G),P(S),j,N]),R=tr(T,H)[0];return D!==void 0&&M!==void 0?{asum:R,gpuTimeMs:D+M}:{asum:R}}finally{!t&&u&&m(u),!t&&s&&m(s),l&&m(l),f&&m(f),c&&m(c),p&&m(p),d&&m(d),g&&m(g),w&&m(w)}}var Re=64;async function Ce(a,e,r,o){let t=r instanceof x;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 B(a,"snrm2"),n=await B(a,"reduction/sum"),u=null,s=null,l=null,f=null,c=null;try{u=t?r._buf:b(r,"snrm2-x",!1),s=C(2*Re*4,"snrm2-partials"),l=q(4,"snrm2-result"),f=I([{value:e,type:"u32"},{value:o,type:"u32"}],"snrm2-params");let p=A(i.getBindGroupLayout(0),[u,s,f]),{commandEncoder:d,ts:g}=F(i,p,2*Re);L(d);let w=A(n.getBindGroupLayout(0),[s,l]),{commandEncoder:E,ts:h}=F(n,w,1);c=_(E,l),L(E);let G=y(c,Float32Array);c=null;let[v,k,S]=await Promise.all([P(g),P(h),G]),j=Math.sqrt(S[0]);return v!==void 0&&k!==void 0?{nrm2:j,gpuTimeMs:v+k}:{nrm2:j}}finally{!t&&u&&m(u),s&&m(s),l&&m(l),f&&m(f),c&&m(c)}}var wr=64;async function Oe(a,e,r,o){let t=r instanceof x;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 B(a,"isamax"),n=await B(a,"reduction/argmax"),u=null,s=null,l=null,f=null,c=null,p=null;try{u=t?r._buf:b(r,"isamax-x",!1),s=C(2*wr*4,"isamax-partials-val"),l=C(2*wr*4,"isamax-partials-idx"),f=q(4,"isamax-result"),c=I([{value:e,type:"u32"},{value:o,type:"u32"}],"isamax-params");let d=A(i.getBindGroupLayout(0),[u,s,l,c]),{commandEncoder:g,ts:w}=F(i,d,2*wr);L(g);let E=A(n.getBindGroupLayout(0),[s,l,f]),{commandEncoder:h,ts:G}=F(n,E,1);p=_(h,f),L(h);let v=y(p,Uint32Array);p=null;let[k,S,j]=await Promise.all([P(w),P(G),v]),N=j[0];return k!==void 0&&S!==void 0?{index:N,gpuTimeMs:k+S}:{index:N}}finally{!t&&u&&m(u),s&&m(s),l&&m(l),f&&m(f),c&&m(c),p&&m(p)}}async function ze(a,e,r,o,t,i,n,u){let s=r instanceof x,l=t instanceof x;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 n!="number")throw new Error("c must be a number.");if(typeof u!="number")throw new Error("s must be a number.");if(Number.isNaN(n)||Number.isNaN(u))throw new Error("c and s must not be NaN.");if(!Number.isFinite(n))throw new Error("c must be finite.");if(!Number.isFinite(u))throw new Error("s must be finite.");if(o<=0||i<=0)throw new Error("incx and incy must be positive.");if(!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 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 f=await B(a,"srot"),c=null,p=null,d=null,g=null,w=null;try{c=s?r._buf:b(r,"srot-x",!0),p=l?t._buf:b(t,"srot-y",!0),d=I([{value:e,type:"u32"},{value:n,type:"f32"},{value:u,type:"f32"},{value:o,type:"u32"},{value:i,type:"u32"}],"srot-params");let E=A(f.getBindGroupLayout(0),[c,p,d]),{commandEncoder:h,ts:G}=F(f,E,O(e));g=s?null:_(h,c),w=l?null:_(h,p),L(h);let v=await P(G);if(s&&l)return v!==void 0?{gpuTimeMs:v}:{};let k=y(g,Float32Array),S=y(w,Float32Array);g=null,w=null;let[j,N]=await Promise.all([k,S]);return v!==void 0?{x:j,y:N,gpuTimeMs:v}:{x:j,y:N}}finally{!s&&c&&m(c),!l&&p&&m(p),d&&m(d),g&&m(g),w&&m(w)}}async function Qe(a,e,r,o,t,i,n){let u=r instanceof x,s=t instanceof x;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(!(n instanceof Float32Array)||n.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(!u&&!(r instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!s&&!(t instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(u!==s)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(e<=0||n[0]===-2)return u?{}:{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 l=await B(a,"srotm"),f=null,c=null,p=null,d=null,g=null,w=null;try{f=u?r._buf:b(r,"srotm-x",!0),c=s?t._buf:b(t,"srotm-y",!0),p=b(n,"srotm-param",!1),d=I([{value:e,type:"u32"},{value:o,type:"u32"},{value:i,type:"u32"}],"srotm-params");let E=A(l.getBindGroupLayout(0),[f,c,p,d]),{commandEncoder:h,ts:G}=F(l,E,O(e));g=u?null:_(h,f),w=s?null:_(h,c),L(h);let v=await P(G);if(u&&s)return v!==void 0?{gpuTimeMs:v}:{};let k=y(g,Float32Array),S=y(w,Float32Array);g=null,w=null;let[j,N]=await Promise.all([k,S]);return v!==void 0?{x:j,y:N,gpuTimeMs:v}:{x:j,y:N}}finally{!u&&f&&m(f),!s&&c&&m(c),p&&m(p),d&&m(d),g&&m(g),w&&m(w)}}async function qe(a,e,r,o,t,i,n,u,s,l,f,c,p="row-major"){let d=i instanceof V,g=u instanceof x,w=f instanceof x;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(p!=="row-major"&&p!=="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 l!="number")throw new Error("beta must be a number.");if(Number.isNaN(l))throw new Error("beta must not be NaN.");if(!Number.isFinite(l))throw new Error("beta must be finite.");if(!Number.isInteger(r)||!Number.isInteger(o)||!Number.isInteger(s)||!Number.isInteger(c)||!Number.isInteger(n))throw new Error("m, n, incx, incy, and lda must be integers.");if(s<=0||c<=0)throw new Error("incx and incy must be positive.");if(!d&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!g&&!(u instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!w&&!(f instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(g!==w)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(g&&!d)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(g&&u._buf===f._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(d&&n!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(d&&(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 w?{}:{y:f};(d?i.layout:p)==="column-major"&&([r,o]=[o,r],e=e==="no-transpose"?"transpose":"no-transpose");let h=e==="no-transpose",G=h?o:r,v=h?r:o;if(n<o)throw new Error("lda must be >= n.");if(!d&&i.length<(r-1)*n+o)throw new Error("A does not have enough elements for the given m, n, and lda.");if(u.length<(G-1)*s+1)throw new Error("x does not have enough elements for the given dimensions and incx.");if(f.length<(v-1)*c+1)throw new Error("y does not have enough elements for the given dimensions and incy.");let S=await B(a,h?"sgemv_n":"sgemv_t"),j=d?i._buf:b(i,"sgemv-A",!1),N=g?u._buf:b(u,"sgemv-x",!1),D=w?f._buf:b(f,"sgemv-y",!0),M=I([{value:r,type:"u32"},{value:o,type:"u32"},{value:t,type:"f32"},{value:l,type:"f32"},{value:s,type:"u32"},{value:c,type:"u32"},{value:n,type:"u32"}],"sgemv-params");try{let T=A(S.getBindGroupLayout(0),[j,N,D,M]),H=h?Math.min(r,a.limits.maxComputeWorkgroupsPerDimension):O(v),{commandEncoder:R,ts:Z}=F(S,T,H),z=w?null:_(R,D);L(R);let Q=await P(Z);if(w)return Q!==void 0?{gpuTimeMs:Q}:{};let nr=await y(z,Float32Array);return Q!==void 0?{y:nr,gpuTimeMs:Q}:{y:nr}}finally{d||m(j),g||m(N),w||m(D),m(M)}}async function Ze(a,e,r,o,t,i,n,u,s,l,f,c="row-major"){let p=n instanceof x,d=l instanceof x,g=t instanceof V;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(c!=="row-major"&&c!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(r)||!Number.isInteger(u)||!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 s!="number")throw new Error("beta must be a number.");if(Number.isNaN(s))throw new Error("beta must not be NaN.");if(!Number.isFinite(s))throw new Error("beta must be finite.");if(u<=0||f<=0)throw new Error("incx and incy must be positive.");if(i<r)throw new Error("lda must be >= n.");if(!g&&!(t instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!p&&!(n instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!d&&!(l instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(p!==d)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(p&&!g)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(p&&n._buf===l._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(g&&i!==t.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(g&&(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 d?{}:{y:l};if(!g&&t.length<(r-1)*i+r)throw new Error("A does not have enough elements for the given n and lda.");if(n.length<(r-1)*u+1)throw new Error("x does not have enough elements for the given n and incx.");if(l.length<(r-1)*f+1)throw new Error("y does not have enough elements for the given n and incy.");let E=(g?t.layout:c)==="column-major"?e==="upper":e==="lower",h=await B(a,"ssymv"),G=null,v=null,k=null,S=null;try{G=g?t._buf:b(t,"ssymv-A",!1),v=p?n._buf:b(n,"ssymv-x",!1),k=d?l._buf:b(l,"ssymv-y",!0),S=I([{value:r,type:"u32"},{value:o,type:"f32"},{value:s,type:"f32"},{value:u,type:"u32"},{value:f,type:"u32"},{value:i,type:"u32"},{value:E?0:1,type:"u32"}],"ssymv-params");let j=A(h.getBindGroupLayout(0),[G,v,k,S]),N=Math.min(r,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:D,ts:M}=F(h,j,N),T=d?null:_(D,k);L(D);let H=await P(M);if(d)return H!==void 0?{gpuTimeMs:H}:{};let R=await y(T,Float32Array);return H!==void 0?{y:R,gpuTimeMs:H}:{y:R}}finally{!g&&G&&m(G),!p&&v&&m(v),!d&&k&&m(k),S&&m(S)}}async function Ke(a,e,r,o,t,i,n,u,s,l,f,c="row-major"){let p=u instanceof x,d=l instanceof x,g=i instanceof V,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(c!=="row-major"&&c!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(t)||!Number.isInteger(s)||!Number.isInteger(f)||!Number.isInteger(n))throw new Error("n, incx, incy, and lda must be integers.");if(s<=0||f<=0)throw new Error("incx and incy must be positive.");if(n<t)throw new Error("lda must be >= n.");if(!g&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!p&&!(u instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!d&&!(l instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(p!==d)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(p&&u._buf===l._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(p&&!g)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(g&&d&&i._buf===l._buf)throw new Error("A and y must not reference the same GPU buffer.");if(g&&n!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(g&&(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 d?{}:{y:l};if(!g&&i.length<(t-1)*n+t)throw new Error("A does not have enough elements for the given n and lda.");if(u.length<(t-1)*s+1)throw new Error("x does not have enough elements for the given n and incx.");if(l.length<(t-1)*f+1)throw new Error("y does not have enough elements for the given n and incy.");let h=(g?i.layout:c)==="column-major",G=h?e==="upper":e==="lower",v=h?r==="transpose":r==="no-transpose",k=await B(a,"strmv"),S=null,j=null,N=null,D=null;try{S=g?i._buf:b(i,"strmv-A",!1),j=p?u._buf:b(u,"strmv-x",!1),N=d?l._buf:b(l,"strmv-y",!0),D=I([{value:t,type:"u32"},{value:s,type:"u32"},{value:f,type:"u32"},{value:n,type:"u32"},{value:v?0:1,type:"u32"},{value:G?0:1,type:"u32"},{value:w?1:0,type:"u32"}],"strmv-params");let M=A(k.getBindGroupLayout(0),[S,j,N,D]),T=Math.min(t,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:H,ts:R}=F(k,M,T),Z=d?null:_(H,N);L(H);let z=await P(R);if(d)return z!==void 0?{gpuTimeMs:z}:{};let Q=await y(Z,Float32Array);return z!==void 0?{y:Q,gpuTimeMs:z}:{y:Q}}finally{!g&&S&&m(S),!p&&j&&m(j),!d&&N&&m(N),D&&m(D)}}var X=64;function Xe(a,e,r){let o=new ArrayBuffer(a*e),t=new DataView(o);for(let i=0;i<a;i++){let n=r(i),u=i*e;n.forEach((s,l)=>t.setUint32(u+l*4,s,!0))}return o}function $e(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 Ye(a,e,r,o,t,i,n,u,s,l="row-major"){let f=u instanceof x,c=i instanceof V,p=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(!p&&o!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(l!=="row-major"&&l!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(t)||!Number.isInteger(s)||!Number.isInteger(n))throw new Error("n, incx, and lda must be integers.");if(s<=0)throw new Error("incx must be positive.");if(n<t)throw new Error("lda must be >= n.");if(!c&&!(i instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!f&&!(u instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(f&&!c)throw new Error("A must be a GpuMatrix when x is a GpuVector.");if(c&&n!==i.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(c&&(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:u};if(!c&&i.length<(t-1)*n+t)throw new Error("A does not have enough elements for the given n and lda.");if(u.length<(t-1)*s+1)throw new Error("x does not have enough elements for the given n and incx.");let g=(c?i.layout:l)==="column-major",w=g?e==="upper":e==="lower",E=g?r==="transpose":r==="no-transpose",h=await B(a,"strsv_invert_block"),G=await B(a,"strsv_apply_inverse"),v=await B(a,"strsv_update"),k=E===w,S=[];for(let Q=0;Q<t;Q+=X)S.push(Q);k||S.reverse();let j=S.length,N=a.limits.maxComputeWorkgroupsPerDimension,D=a.limits.minUniformBufferOffsetAlignment,M=null,T=null,H=null,R=null,Z=null,z=null;try{M=c?i._buf:b(i,"strsv-A",!1),T=f?u._buf:b(u,"strsv-x",!0),H=C(j*X*X*4,"strsv-Ainv");let Q=Xe(j,D,$=>{let Y=$*X,or=Math.min(Y+X,t);return[s,$,Y,or]});R=$e(a,Q,"strsv-apply-params");let nr=Xe(j,D,$=>{let Y=$*X,or=Math.min(Y+X,t);return[t,s,n,E?0:1,w?0:1,Y,or]});Z=$e(a,nr,"strsv-update-params");let{commandEncoder:rr,querySet:ar}=dr();z=I([{value:t,type:"u32"},{value:n,type:"u32"},{value:E?0:1,type:"u32"},{value:w?0:1,type:"u32"},{value:p?1:0,type:"u32"}],"strsv-invert-params");let tt=A(h.getBindGroupLayout(0),[M,H,z]);ir(rr,h,tt,{x:X,y:j},ar?{timestampWrites:{querySet:ar,beginningOfPassWriteIndex:0}}:void 0);for(let $=0;$<S.length;$++){let Y=S[$],or=Math.min(Y+X,t),it=Y/X,nt=$===S.length-1,hr=it*D,st=A(G.getBindGroupLayout(0),[H,T,{buffer:R,offset:hr,size:16}]);ir(rr,G,st,1,nt&&ar?{timestampWrites:{querySet:ar,endOfPassWriteIndex:1}}:void 0);let xr=k?t-or:Y;if(xr===0)continue;let ut=A(v.getBindGroupLayout(0),[M,T,{buffer:Z,offset:hr,size:32}]),lt=Math.min(xr,N);ir(rr,v,ut,lt)}let at=cr(rr,ar),ot=f?null:_(rr,T);L(rr);let sr=await P(at);if(f)return sr!==void 0?{gpuTimeMs:sr}:{};let br=await y(ot,Float32Array);return sr!==void 0?{x:br,gpuTimeMs:sr}:{x:br}}finally{!c&&M&&m(M),!f&&T&&m(T),H&&m(H),R&&m(R),Z&&m(Z),z&&m(z)}}async function Je(a,e,r,o,t,i,n,u,s,l,f="row-major"){let c=s instanceof V;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(u)||!Number.isInteger(l))throw new Error("m, n, incx, incy, and lda must be integers.");if(i<=0||u<=0)throw new Error("incx and incy must be positive.");if(!c&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(c&&l!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(c&&(s.rows<e||s.cols<r))throw new Error("A is too small for the given m and n.");(c?s.layout:f)==="column-major"&&([e,r]=[r,e],[t,n]=[n,t],[i,u]=[u,i]);let d=t instanceof x,g=n instanceof x;if(l<r)throw new Error("lda must be >= n.");if(!d&&!(t instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!g&&!(n instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(d!==g)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(d&&!c)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(c&&d&&s._buf===t._buf)throw new Error("A and x must not reference the same GPU buffer.");if(c&&g&&s._buf===n._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 c?{}:{A:s};if(!c&&s.length<(e-1)*l+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(n.length<(r-1)*u+1)throw new Error("y does not have enough elements for the given n and incy.");let w=await B(a,"sger"),E=null,h=null,G=null,v=null;try{E=d?t._buf:b(t,"sger-x",!1),h=g?n._buf:b(n,"sger-y",!1),G=c?s._buf:b(s,"sger-A",!0),v=I([{value:e,type:"u32"},{value:r,type:"u32"},{value:o,type:"f32"},{value:i,type:"u32"},{value:u,type:"u32"},{value:l,type:"u32"}],"sger-params");let k=A(w.getBindGroupLayout(0),[E,h,G,v]),S=Math.min(e,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:j,ts:N}=F(w,k,S),D=c?null:_(j,G);L(j);let M=await P(N);if(c)return M!==void 0?{gpuTimeMs:M}:{};let T=await y(D,Float32Array);return M!==void 0?{A:T,gpuTimeMs:M}:{A:T}}finally{!d&&E&&m(E),!g&&h&&m(h),!c&&G&&m(G),v&&m(v)}}async function rt(a,e,r,o,t,i,n,u,s="row-major"){let l=t instanceof x,f=n instanceof V;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(s!=="row-major"&&s!=="column-major")throw new Error("layout must be 'row-major' or 'column-major'.");if(!Number.isInteger(r)||!Number.isInteger(i)||!Number.isInteger(u))throw new Error("n, incx, and lda must be integers.");if(typeof o!="number")throw new Error("alpha must be a number.");if(Number.isNaN(o))throw new Error("alpha must not be NaN.");if(!Number.isFinite(o))throw new Error("alpha must be finite.");if(i<=0)throw new Error("incx must be positive.");if(u<r)throw new Error("lda must be >= n.");if(!f&&!(n instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!l&&!(t instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(l&&!f)throw new Error("A must be a GpuMatrix when x is a GpuVector.");if(f&&l&&n._buf===t._buf)throw new Error("A and x must not reference the same GPU buffer.");if(f&&u!==n.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(f&&(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 f?{}:{A:n};if(!f&&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.");let p=(f?n.layout:s)==="column-major"?e==="upper":e==="lower",d=await B(a,"ssyr"),g=null,w=null,E=null;try{g=l?t._buf:b(t,"ssyr-x",!1),w=f?n._buf:b(n,"ssyr-A",!0),E=I([{value:r,type:"u32"},{value:o,type:"f32"},{value:i,type:"u32"},{value:u,type:"u32"},{value:p?0:1,type:"u32"}],"ssyr-params");let h=A(d.getBindGroupLayout(0),[g,w,E]),G=Math.min(r,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:v,ts:k}=F(d,h,G),S=f?null:_(v,w);L(v);let j=await P(k);if(f)return j!==void 0?{gpuTimeMs:j}:{};let N=await y(S,Float32Array);return j!==void 0?{A:N,gpuTimeMs:j}:{A:N}}finally{!l&&g&&m(g),!f&&w&&m(w),E&&m(E)}}async function et(a,e,r,o,t,i,n,u,s,l,f="row-major"){let c=t instanceof x,p=n instanceof x,d=s instanceof V;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(u)||!Number.isInteger(l))throw new Error("n, incx, incy, and lda must be integers.");if(typeof o!="number")throw new Error("alpha must be a number.");if(Number.isNaN(o))throw new Error("alpha must not be NaN.");if(!Number.isFinite(o))throw new Error("alpha must be finite.");if(i<=0||u<=0)throw new Error("incx and incy must be positive.");if(l<r)throw new Error("lda must be >= n.");if(!d&&!(s instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!c&&!(t instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!p&&!(n 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&&!d)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(d&&c&&s._buf===t._buf)throw new Error("A and x must not reference the same GPU buffer.");if(d&&p&&s._buf===n._buf)throw new Error("A and y must not reference the same GPU buffer.");if(c&&t._buf===n._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(d&&l!==s.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(d&&(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 d?{}:{A:s};if(!d&&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.");if(n.length<(r-1)*u+1)throw new Error("y does not have enough elements for the given n and incy.");let w=(d?s.layout:f)==="column-major"?e==="upper":e==="lower",E=await B(a,"ssyr2"),h=null,G=null,v=null,k=null;try{h=c?t._buf:b(t,"ssyr2-x",!1),G=p?n._buf:b(n,"ssyr2-y",!1),v=d?s._buf:b(s,"ssyr2-A",!0),k=I([{value:r,type:"u32"},{value:o,type:"f32"},{value:i,type:"u32"},{value:u,type:"u32"},{value:l,type:"u32"},{value:w?0:1,type:"u32"}],"ssyr2-params");let S=A(E.getBindGroupLayout(0),[h,G,v,k]),j=Math.min(r,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:N,ts:D}=F(E,S,j),M=d?null:_(N,v);L(N);let T=await P(D);if(d)return T!==void 0?{gpuTimeMs:T}:{};let H=await y(M,Float32Array);return T!==void 0?{A:H,gpuTimeMs:T}:{A:H}}finally{!c&&h&&m(h),!p&&G&&m(G),!d&&v&&m(v),k&&m(k)}}return gt(Zt);})();
|
|
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);})();
|