wgblas 0.1.2 → 1.0.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +3 -0
- package/dist/wgblas.browser.js +1078 -37
- package/index.d.mts +6 -0
- package/index.mjs +6 -0
- package/package.json +32 -1
- package/src/classes/GpuMatrix.d.mts +85 -0
- package/src/classes/GpuMatrix.mjs +91 -0
- package/src/classes/GpuVector.d.mts +11 -7
- package/src/classes/GpuVector.mjs +31 -5
- package/src/dasum/dasum.d.mts +48 -0
- package/src/dasum/dasum.mjs +121 -0
- package/src/init.mjs +3 -1
- package/src/isamax/isamax.mjs +66 -67
- package/src/random/random.d.mts +37 -0
- package/src/random/random.mjs +17 -0
- package/src/sasum/sasum.mjs +55 -51
- package/src/saxpy/saxpy.mjs +48 -35
- package/src/scopy/scopy.mjs +43 -32
- package/src/sdot/sdot.mjs +60 -55
- package/src/sgemv/sgemv.d.mts +118 -0
- package/src/sgemv/sgemv.mjs +141 -0
- package/src/shaders/browser-shaders.mjs +20 -0
- package/src/shaders/dasum.wgsl +98 -0
- package/src/shaders/f64add.wgsl +281 -0
- package/src/shaders/isamax.wgsl +32 -9
- package/src/shaders/reduction/sumF64.wgsl +49 -0
- package/src/shaders/sasum.wgsl +18 -4
- package/src/shaders/sdot.wgsl +18 -4
- package/src/shaders/sgemv_n.wgsl +75 -0
- package/src/shaders/sgemv_t.wgsl +65 -0
- package/src/shaders/snrm2.wgsl +22 -4
- package/src/shaders/ssymv.wgsl +69 -0
- package/src/shaders/strmv.wgsl +103 -0
- package/src/shaders/strsv_apply_inverse.wgsl +47 -0
- package/src/shaders/strsv_invert_block.wgsl +109 -0
- package/src/shaders/strsv_update.wgsl +75 -0
- package/src/snrm2/snrm2.mjs +56 -52
- package/src/srot/srot.mjs +57 -41
- package/src/srotm/srotm.mjs +54 -38
- package/src/sscal/sscal.mjs +43 -32
- package/src/sswap/sswap.mjs +49 -34
- package/src/ssymv/ssymv.d.mts +109 -0
- package/src/ssymv/ssymv.mjs +130 -0
- package/src/strmv/strmv.d.mts +109 -0
- package/src/strmv/strmv.mjs +132 -0
- package/src/strsv/strsv.d.mts +98 -0
- package/src/strsv/strsv.mjs +212 -0
- package/src/util/benchmark.mjs +1 -1
- package/src/util/bindgroup.mjs +14 -10
- package/src/util/buffer.mjs +7 -2
- package/src/util/compute.mjs +41 -15
- package/src/util/f64pack.mjs +152 -0
- package/src/util/pipeline.mjs +32 -17
- package/src/util/result.mjs +8 -4
- package/src/util/workgroup.mjs +10 -10
package/dist/wgblas.browser.js
CHANGED
|
@@ -1,4 +1,4 @@
|
|
|
1
|
-
var wgblas=(()=>{var
|
|
1
|
+
var wgblas=(()=>{var Je=Object.create;var or=Object.defineProperty;var rt=Object.getOwnPropertyDescriptor;var et=Object.getOwnPropertyNames;var tt=Object.getPrototypeOf,at=Object.prototype.hasOwnProperty;var nr=(a=>typeof require<"u"?require:typeof Proxy<"u"?new Proxy(a,{get:(r,e)=>(typeof require<"u"?require:r)[e]}):a)(function(a){if(typeof require<"u")return require.apply(this,arguments);throw Error('Dynamic require of "'+a+'" is not supported')});var W=(a,r,e)=>()=>{if(e)throw e[0];try{return a&&(r=a(a=0)),r}catch(i){throw e=[i],i}};var hr=(a,r)=>{for(var e in r)or(a,e,{get:r[e],enumerable:!0})},xr=(a,r,e,i)=>{if(r&&typeof r=="object"||typeof r=="function")for(let t of et(r))!at.call(a,t)&&t!==e&&or(a,t,{get:()=>r[t],enumerable:!(i=rt(r,t))||i.enumerable});return a};var sr=(a,r,e)=>(e=a!=null?Je(tt(a)):{},xr(r||!a||!a.__esModule?or(e,"default",{value:a,enumerable:!0}):e,a)),it=a=>xr(or({},"__esModule",{value:!0}),a);var jr,Ir=W(()=>{jr=`// amax reduction: collapses 2*WGS (value, index) pairs into one index.
|
|
2
2
|
// dispatch: 1 workgroup of WGS threads.
|
|
3
3
|
// partials_val and partials_idx must have exactly 2*WGS entries.
|
|
4
4
|
|
|
@@ -41,7 +41,7 @@ fn reduce(
|
|
|
41
41
|
|
|
42
42
|
if (i == 0u) { result[0] = tile_idx[0]; }
|
|
43
43
|
}
|
|
44
|
-
`});var
|
|
44
|
+
`});var Wr,Nr=W(()=>{Wr=`// sum reduction: collapses 2*WGS partials into one scalar.
|
|
45
45
|
// dispatch: 1 workgroup of WGS threads.
|
|
46
46
|
// partials must have exactly 2*WGS entries.
|
|
47
47
|
|
|
@@ -67,7 +67,56 @@ fn reduce(
|
|
|
67
67
|
|
|
68
68
|
if (i == 0u) { result[0] = tile[0]; }
|
|
69
69
|
}
|
|
70
|
-
`});var
|
|
70
|
+
`});var Ur,Mr=W(()=>{Ur=`// sum reduction (f64): collapses 2*WGS partial [main, aux] pairs into one,
|
|
71
|
+
// using computeSum instead of plain f32 \`+\` (see reduction/sum.wgsl for the
|
|
72
|
+
// f32 original this mirrors).
|
|
73
|
+
// dispatch: 1 workgroup of WGS threads. partialsMain/partialsAux must have
|
|
74
|
+
// exactly 2*WGS entries each.
|
|
75
|
+
//
|
|
76
|
+
// Concatenated after f64add.wgsl by getPipeline (WGSL has no #include),
|
|
77
|
+
// reusing its decode/encode/computeSum and Packed struct \u2014 f64add.wgsl
|
|
78
|
+
// declares no bindings and no entry point of its own (just helper functions),
|
|
79
|
+
// so bindings here start at 0 and the entry point is simply \`reduce_f64\`.
|
|
80
|
+
//
|
|
81
|
+
// partialsAux/result's aux slot are array<u32>, not array<f32> \u2014 aux's bits
|
|
82
|
+
// must never pass through an f32-typed storage slot (NaN-bit-pattern
|
|
83
|
+
// corruption risk, see f64pack.mjs and the Packed struct comment above
|
|
84
|
+
// decode()/encode() in f64add.wgsl).
|
|
85
|
+
|
|
86
|
+
@group(0) @binding(0) var<storage, read> partialsMain: array<f32>;
|
|
87
|
+
@group(0) @binding(1) var<storage, read> partialsAux: array<u32>;
|
|
88
|
+
@group(0) @binding(2) var<storage, read_write> resultMain: array<f32, 1>;
|
|
89
|
+
@group(0) @binding(3) var<storage, read_write> resultAux: array<u32, 1>;
|
|
90
|
+
|
|
91
|
+
const WGS: u32 = 64;
|
|
92
|
+
|
|
93
|
+
var<workgroup> tile: array<Packed, 64>;
|
|
94
|
+
|
|
95
|
+
fn addPair(a: Packed, b: Packed) -> Packed {
|
|
96
|
+
return computeSum(decode(bitcast<u32>(a.main), a.aux), decode(bitcast<u32>(b.main), b.aux));
|
|
97
|
+
}
|
|
98
|
+
|
|
99
|
+
@compute @workgroup_size(64)
|
|
100
|
+
fn reduce_f64(
|
|
101
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
102
|
+
) {
|
|
103
|
+
let i = lid.x;
|
|
104
|
+
let a = Packed(partialsMain[i], partialsAux[i]);
|
|
105
|
+
let b = Packed(partialsMain[i + WGS], partialsAux[i + WGS]);
|
|
106
|
+
tile[i] = addPair(a, b);
|
|
107
|
+
workgroupBarrier();
|
|
108
|
+
|
|
109
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
110
|
+
if (i < s) { tile[i] = addPair(tile[i], tile[i + s]); }
|
|
111
|
+
workgroupBarrier();
|
|
112
|
+
}
|
|
113
|
+
|
|
114
|
+
if (i == 0u) {
|
|
115
|
+
resultMain[0] = tile[0].main;
|
|
116
|
+
resultAux[0] = tile[0].aux;
|
|
117
|
+
}
|
|
118
|
+
}
|
|
119
|
+
`});var Rr,Tr=W(()=>{Rr=`// sscal: x = alpha * x
|
|
71
120
|
|
|
72
121
|
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
73
122
|
|
|
@@ -90,7 +139,7 @@ fn main(
|
|
|
90
139
|
x[id * params.x_inc] = params.alpha * x[id * params.x_inc];
|
|
91
140
|
}
|
|
92
141
|
}
|
|
93
|
-
`});var
|
|
142
|
+
`});var Hr,Vr=W(()=>{Hr=`// sswap: x <-> y
|
|
94
143
|
|
|
95
144
|
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
96
145
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
@@ -116,7 +165,7 @@ fn main(
|
|
|
116
165
|
y[id * params.y_inc] = temp;
|
|
117
166
|
}
|
|
118
167
|
}
|
|
119
|
-
`});var
|
|
168
|
+
`});var Cr,Dr=W(()=>{Cr=`// saxpy: y = alpha * x + y
|
|
120
169
|
|
|
121
170
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
122
171
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
@@ -141,7 +190,7 @@ fn main(
|
|
|
141
190
|
y[id * params.y_inc] = params.alpha * x[id * params.x_inc] + y[id * params.y_inc];
|
|
142
191
|
}
|
|
143
192
|
}
|
|
144
|
-
`});var
|
|
193
|
+
`});var zr,Or=W(()=>{zr=`// scopy: y = x
|
|
145
194
|
|
|
146
195
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
147
196
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
@@ -165,7 +214,7 @@ fn main(
|
|
|
165
214
|
y[id * params.y_inc] = x[id * params.x_inc];
|
|
166
215
|
}
|
|
167
216
|
}
|
|
168
|
-
`});var
|
|
217
|
+
`});var qr,Qr=W(()=>{qr=`// sdot: result = sum(x[i] * y[i])
|
|
169
218
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sum.wgsl.
|
|
170
219
|
|
|
171
220
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -190,11 +239,25 @@ fn main(
|
|
|
190
239
|
@builtin(workgroup_id) wgid: vec3u,
|
|
191
240
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
192
241
|
) {
|
|
193
|
-
var
|
|
194
|
-
|
|
195
|
-
|
|
242
|
+
var acc0: f32 = 0.0;
|
|
243
|
+
var acc1: f32 = 0.0;
|
|
244
|
+
var acc2: f32 = 0.0;
|
|
245
|
+
var acc3: f32 = 0.0;
|
|
246
|
+
|
|
247
|
+
let stride = num_wg.x * WGS;
|
|
248
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
249
|
+
|
|
250
|
+
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
251
|
+
acc0 += x[ id * params.x_inc] * y[ id * params.y_inc];
|
|
252
|
+
acc1 += x[(id + stride) * params.x_inc] * y[(id + stride) * params.y_inc];
|
|
253
|
+
acc2 += x[(id + 2u * stride) * params.x_inc] * y[(id + 2u * stride) * params.y_inc];
|
|
254
|
+
acc3 += x[(id + 3u * stride) * params.x_inc] * y[(id + 3u * stride) * params.y_inc];
|
|
196
255
|
}
|
|
197
|
-
|
|
256
|
+
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
257
|
+
acc0 += x[id * params.x_inc] * y[id * params.y_inc];
|
|
258
|
+
}
|
|
259
|
+
|
|
260
|
+
tile[lid.x] = acc0 + acc1 + acc2 + acc3;
|
|
198
261
|
workgroupBarrier();
|
|
199
262
|
|
|
200
263
|
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
@@ -204,7 +267,7 @@ fn main(
|
|
|
204
267
|
|
|
205
268
|
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
206
269
|
}
|
|
207
|
-
`});var
|
|
270
|
+
`});var Kr,Zr=W(()=>{Kr=`// sasum: result = sum(|x[i]|)
|
|
208
271
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/abssum.wgsl.
|
|
209
272
|
|
|
210
273
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -227,11 +290,25 @@ fn main(
|
|
|
227
290
|
@builtin(workgroup_id) wgid: vec3u,
|
|
228
291
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
229
292
|
) {
|
|
230
|
-
var
|
|
231
|
-
|
|
232
|
-
|
|
293
|
+
var acc0: f32 = 0.0;
|
|
294
|
+
var acc1: f32 = 0.0;
|
|
295
|
+
var acc2: f32 = 0.0;
|
|
296
|
+
var acc3: f32 = 0.0;
|
|
297
|
+
|
|
298
|
+
let stride = num_wg.x * WGS;
|
|
299
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
300
|
+
|
|
301
|
+
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
302
|
+
acc0 += abs(x[ id * params.x_inc]);
|
|
303
|
+
acc1 += abs(x[(id + stride) * params.x_inc]);
|
|
304
|
+
acc2 += abs(x[(id + 2u * stride) * params.x_inc]);
|
|
305
|
+
acc3 += abs(x[(id + 3u * stride) * params.x_inc]);
|
|
233
306
|
}
|
|
234
|
-
|
|
307
|
+
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
308
|
+
acc0 += abs(x[id * params.x_inc]);
|
|
309
|
+
}
|
|
310
|
+
|
|
311
|
+
tile[lid.x] = acc0 + acc1 + acc2 + acc3;
|
|
235
312
|
workgroupBarrier();
|
|
236
313
|
|
|
237
314
|
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
@@ -241,7 +318,7 @@ fn main(
|
|
|
241
318
|
|
|
242
319
|
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
243
320
|
}
|
|
244
|
-
`});var
|
|
321
|
+
`});var $r,Xr=W(()=>{$r=`// snrm2: result = sqrt(sum(x[i] * x[i]))
|
|
245
322
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sqsum.wgsl.
|
|
246
323
|
|
|
247
324
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -264,12 +341,30 @@ fn main(
|
|
|
264
341
|
@builtin(workgroup_id) wgid: vec3u,
|
|
265
342
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
266
343
|
) {
|
|
267
|
-
var
|
|
268
|
-
|
|
344
|
+
var acc0: f32 = 0.0;
|
|
345
|
+
var acc1: f32 = 0.0;
|
|
346
|
+
var acc2: f32 = 0.0;
|
|
347
|
+
var acc3: f32 = 0.0;
|
|
348
|
+
|
|
349
|
+
let stride = num_wg.x * WGS;
|
|
350
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
351
|
+
|
|
352
|
+
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
353
|
+
let v0 = x[ id * params.x_inc];
|
|
354
|
+
let v1 = x[(id + stride) * params.x_inc];
|
|
355
|
+
let v2 = x[(id + 2u * stride) * params.x_inc];
|
|
356
|
+
let v3 = x[(id + 3u * stride) * params.x_inc];
|
|
357
|
+
acc0 += v0 * v0;
|
|
358
|
+
acc1 += v1 * v1;
|
|
359
|
+
acc2 += v2 * v2;
|
|
360
|
+
acc3 += v3 * v3;
|
|
361
|
+
}
|
|
362
|
+
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
269
363
|
let v = x[id * params.x_inc];
|
|
270
|
-
|
|
364
|
+
acc0 += v * v;
|
|
271
365
|
}
|
|
272
|
-
|
|
366
|
+
|
|
367
|
+
tile[lid.x] = acc0 + acc1 + acc2 + acc3;
|
|
273
368
|
workgroupBarrier();
|
|
274
369
|
|
|
275
370
|
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
@@ -279,7 +374,7 @@ fn main(
|
|
|
279
374
|
|
|
280
375
|
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
281
376
|
}
|
|
282
|
-
`});var
|
|
377
|
+
`});var Jr,Yr=W(()=>{Jr=`// srot: x = c*x + s*y, y = -s*x + c*y
|
|
283
378
|
|
|
284
379
|
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
285
380
|
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
@@ -308,7 +403,7 @@ fn main(
|
|
|
308
403
|
y[id * params.y_inc] = -params.s * xi + params.c * yi;
|
|
309
404
|
}
|
|
310
405
|
}
|
|
311
|
-
`});var
|
|
406
|
+
`});var ee,re=W(()=>{ee=`// srotm: applies modified Givens rotation H to vectors x and y.
|
|
312
407
|
// param[0] = flag: -1 (full H), 0 (unit diagonal), 1 (unit off-diagonal)
|
|
313
408
|
// param = [ flag, h11, h21, h12, h22 ]
|
|
314
409
|
// flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
|
|
@@ -358,7 +453,7 @@ fn main(
|
|
|
358
453
|
y[id * params.y_inc] = h21 * xi + h22 * yi;
|
|
359
454
|
}
|
|
360
455
|
}
|
|
361
|
-
`});var
|
|
456
|
+
`});var ae,te=W(()=>{ae=`// isamax: returns index of element with largest absolute value
|
|
362
457
|
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/argmax.wgsl.
|
|
363
458
|
|
|
364
459
|
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
@@ -385,19 +480,42 @@ fn main(
|
|
|
385
480
|
) {
|
|
386
481
|
// -1.0 is a safe sentinel: any |x[i]| >= 0 beats it,
|
|
387
482
|
// so workgroups with no elements lose gracefully in the epilogue.
|
|
388
|
-
var
|
|
389
|
-
var
|
|
390
|
-
|
|
391
|
-
|
|
483
|
+
var best_val0: f32 = -1.0; var best_idx0: u32 = 0u;
|
|
484
|
+
var best_val1: f32 = -1.0; var best_idx1: u32 = 0u;
|
|
485
|
+
var best_val2: f32 = -1.0; var best_idx2: u32 = 0u;
|
|
486
|
+
var best_val3: f32 = -1.0; var best_idx3: u32 = 0u;
|
|
487
|
+
|
|
488
|
+
let stride = num_wg.x * WGS;
|
|
489
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
490
|
+
|
|
491
|
+
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
492
|
+
let v0 = abs(x[ id * params.x_inc]);
|
|
493
|
+
let v1 = abs(x[(id + stride) * params.x_inc]);
|
|
494
|
+
let v2 = abs(x[(id + 2u * stride) * params.x_inc]);
|
|
495
|
+
let v3 = abs(x[(id + 3u * stride) * params.x_inc]);
|
|
496
|
+
if (v0 > best_val0) { best_val0 = v0; best_idx0 = id; }
|
|
497
|
+
if (v1 > best_val1) { best_val1 = v1; best_idx1 = id + stride; }
|
|
498
|
+
if (v2 > best_val2) { best_val2 = v2; best_idx2 = id + 2u * stride; }
|
|
499
|
+
if (v3 > best_val3) { best_val3 = v3; best_idx3 = id + 3u * stride; }
|
|
500
|
+
}
|
|
501
|
+
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
392
502
|
let v = abs(x[id * params.x_inc]);
|
|
393
|
-
if (v >
|
|
394
|
-
best_val = v;
|
|
395
|
-
best_idx = id;
|
|
396
|
-
}
|
|
503
|
+
if (v > best_val0) { best_val0 = v; best_idx0 = id; }
|
|
397
504
|
}
|
|
398
505
|
|
|
399
|
-
|
|
400
|
-
|
|
506
|
+
// merge 4 independent lanes; prefer lower index on tie (first occurrence wins)
|
|
507
|
+
if (best_val1 > best_val0 || (best_val1 == best_val0 && best_idx1 < best_idx0)) {
|
|
508
|
+
best_val0 = best_val1; best_idx0 = best_idx1;
|
|
509
|
+
}
|
|
510
|
+
if (best_val2 > best_val0 || (best_val2 == best_val0 && best_idx2 < best_idx0)) {
|
|
511
|
+
best_val0 = best_val2; best_idx0 = best_idx2;
|
|
512
|
+
}
|
|
513
|
+
if (best_val3 > best_val0 || (best_val3 == best_val0 && best_idx3 < best_idx0)) {
|
|
514
|
+
best_val0 = best_val3; best_idx0 = best_idx3;
|
|
515
|
+
}
|
|
516
|
+
|
|
517
|
+
tile_val[lid.x] = best_val0;
|
|
518
|
+
tile_idx[lid.x] = best_idx0;
|
|
401
519
|
workgroupBarrier();
|
|
402
520
|
|
|
403
521
|
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
@@ -417,6 +535,929 @@ fn main(
|
|
|
417
535
|
partials_idx[wgid.x] = tile_idx[0];
|
|
418
536
|
}
|
|
419
537
|
}
|
|
420
|
-
`});var
|
|
421
|
-
|
|
422
|
-
`)}`);let s=r.createComputePipeline({label:t,layout:"auto",compute:{module:a}});return s._shaderModule=a,s}var ge=64,Ur=8;function M(t,r){let e=F().limits.maxComputeWorkgroupsPerDimension;return r===void 0?Math.min(Math.ceil(t/ge),e):{x:Math.min(Math.ceil(r/Ur),e),y:Math.min(Math.ceil(t/Ur),e)}}async function Ir(t,r,e,a,o){let n=a instanceof m;if(!(t instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(o))throw new Error("n and incx must be integers.");if(isNaN(e))throw new Error("alpha must not be NaN.");if(!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 m))throw new Error("x must be a Float32Array or GpuVector.");if(r<=0)return n?{}:a;if(a.length<(r-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");let s=await A(t,"sscal"),i=n?a._buf:v(a,"sscal-x",!0),u=R([{value:r,type:"u32"},{value:e,type:"f32"},{value:o,type:"u32"}],"sscal-params"),f=B(s.getBindGroupLayout(0),[i,u]),{commandEncoder:c,ts:p}=P(s,f,M(r)),y=n?null:h(c,i);E(c);let l=await G(p);if(n)return w(u),l!==void 0?{gpuTimeMs:l}:{};let k=await _(y,Float32Array);return w(i,u,y),l!==void 0?{result:k,gpuTimeMs:l}:k}async function Mr(t,r,e,a,o,n){let s=e instanceof m,i=o instanceof m;if(!(t instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(a)||!Number.isInteger(n))throw new Error("n, incx, and incy must be integers.");if(a<=0||n<=0)throw new Error("incx and incy must be positive.");if(!(e instanceof Float32Array)&&!(e instanceof m))throw new Error("x must be a Float32Array or GpuVector.");if(!(o instanceof Float32Array)&&!(o instanceof m))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(r<=0)return s?{}:{x:e,y:o};if(e.length<(r-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(o.length<(r-1)*n+1)throw new Error("y does not have enough elements for the given n and incy.");let u=await A(t,"sswap"),f=s?e._buf:v(e,"sswap-x",!0),c=i?o._buf:v(o,"sswap-y",!0),p=R([{value:r,type:"u32"},{value:a,type:"u32"},{value:n,type:"u32"}],"sswap-params"),y=B(u.getBindGroupLayout(0),[f,c,p]),{commandEncoder:l,ts:k}=P(u,y,M(r)),x=s?null:h(l,f),S=i?null:h(l,c);E(l);let d=await G(k);if(s&&i)return w(p),d!==void 0?{gpuTimeMs:d}:{};let b=await _(x,Float32Array),g=await _(S,Float32Array);return w(f,x,c,S,p),d!==void 0?{x:b,y:g,gpuTimeMs:d}:{x:b,y:g}}async function Tr(t,r,e,a,o,n,s){let i=a instanceof m,u=n instanceof m;if(!(t instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(o)||!Number.isInteger(s))throw new Error("n, incx, and incy must be integers.");if(isNaN(e))throw new Error("alpha must not be NaN.");if(!isFinite(e))throw new Error("alpha must be finite.");if(o<=0||s<=0)throw new Error("incx and incy must be positive.");if(!i&&!(a instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!u&&!(n 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(r<=0)return u?{}:{y:n};if(a.length<(r-1)*o+1)throw new Error("x does not have enough elements for the given n and incx.");if(n.length<(r-1)*s+1)throw new Error("y does not have enough elements for the given n and incy.");let f=await A(t,"saxpy"),c=i?a._buf:v(a,"saxpy-x",!1),p=u?n._buf:v(n,"saxpy-y",!0),y=R([{value:r,type:"u32"},{value:e,type:"f32"},{value:o,type:"u32"},{value:s,type:"u32"}],"saxpy-params"),l=B(f.getBindGroupLayout(0),[c,p,y]),{commandEncoder:k,ts:x}=P(f,l,M(r)),S=u?null:h(k,p);E(k);let d=await G(x);if(u&&i)return w(y),d!==void 0?{gpuTimeMs:d}:{};let b=await _(S,Float32Array);return w(c,p,y,S),d!==void 0?{y:b,gpuTimeMs:d}:{y:b}}async function Dr(t,r,e,a,o,n){let s=e instanceof m,i=o instanceof m;if(!(t instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(a)||!Number.isInteger(n))throw new Error("n, incx, and incy must be integers.");if(a<=0||n<=0)throw new Error("incx and incy must be positive.");if(!s&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!i&&!(o instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(s!==i)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(r<=0)return i?{}:{y:o};if(e.length<(r-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(o.length<(r-1)*n+1)throw new Error("y does not have enough elements for the given n and incy.");let u=await A(t,"scopy"),f=s?e._buf:v(e,"scopy-x",!1),c=i?o._buf:v(o,"scopy-y",!0),p=R([{value:r,type:"u32"},{value:a,type:"u32"},{value:n,type:"u32"}],"scopy-params"),y=B(u.getBindGroupLayout(0),[f,c,p]),{commandEncoder:l,ts:k}=P(u,y,M(r)),x=i?null:h(l,c);E(l);let S=await G(k);if(i&&s)return w(p),S!==void 0?{gpuTimeMs:S}:{};let d=await _(x,Float32Array);return w(f,c,p,x),S!==void 0?{y:d,gpuTimeMs:S}:{y:d}}var Nr=64;async function Vr(t,r,e,a,o,n){let s=e instanceof m,i=o instanceof m;if(!(t instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(a)||!Number.isInteger(n))throw new Error("n, incx, and incy must be integers.");if(a<=0||n<=0)throw new Error("incx and incy must be positive.");if(!s&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!i&&!(o instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(s!==i)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(r<=0)return{dot:0};if(e.length<(r-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(o.length<(r-1)*n+1)throw new Error("y does not have enough elements for the given n and incy.");let u=await A(t,"sdot"),f=await A(t,"reduction/sum"),c=s?e._buf:v(e,"sdot-x",!1),p=i?o._buf:v(o,"sdot-y",!1),y=N(2*Nr*4,"sdot-partials"),l=V(4,"sdot-result"),k=R([{value:r,type:"u32"},{value:a,type:"u32"},{value:n,type:"u32"}],"sdot-params"),x=B(u.getBindGroupLayout(0),[c,p,y,k]),{commandEncoder:S,ts:d}=P(u,x,2*Nr);E(S);let b=B(f.getBindGroupLayout(0),[y,l]),{commandEncoder:g,ts:W}=P(f,b,1),U=h(g,l);E(g);let[D,z,O]=await Promise.all([G(d),G(W),_(U,Float32Array)]);return s&&i?(w(y,l,k,U),D!==void 0&&z!==void 0?{dot:O[0],gpuTimeMs:D+z}:{dot:O[0]}):(w(c,p,y,l,k,U),D!==void 0&&z!==void 0?{dot:O[0],gpuTimeMs:D+z}:{dot:O[0]})}var Cr=64;async function zr(t,r,e,a){let o=e instanceof m;if(!(t instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!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(r<=0)return{asum:0};if(e.length<(r-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");let n=await A(t,"sasum"),s=await A(t,"reduction/sum"),i=o?e._buf:v(e,"sasum-x",!1),u=N(2*Cr*4,"sasum-partials"),f=V(4,"sasum-result"),c=R([{value:r,type:"u32"},{value:a,type:"u32"}],"sasum-params"),p=B(n.getBindGroupLayout(0),[i,u,c]),{commandEncoder:y,ts:l}=P(n,p,2*Cr);E(y);let k=B(s.getBindGroupLayout(0),[u,f]),{commandEncoder:x,ts:S}=P(s,k,1),d=h(x,f);E(x);let[b,g,W]=await Promise.all([G(l),G(S),_(d,Float32Array)]);return o?(w(u,f,c,d),b!==void 0&&g!==void 0?{asum:W[0],gpuTimeMs:b+g}:{asum:W[0]}):(w(i,u,f,c,d),b!==void 0&&g!==void 0?{asum:W[0],gpuTimeMs:b+g}:{asum:W[0]})}var Or=64;async function Lr(t,r,e,a){let o=e instanceof m;if(!(t instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!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(r<=0)return{nrm2:0};if(e.length<(r-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");let n=await A(t,"snrm2"),s=await A(t,"reduction/sum"),i=o?e._buf:v(e,"snrm2-x",!1),u=N(2*Or*4,"snrm2-partials"),f=V(4,"snrm2-result"),c=R([{value:r,type:"u32"},{value:a,type:"u32"}],"snrm2-params"),p=B(n.getBindGroupLayout(0),[i,u,c]),{commandEncoder:y,ts:l}=P(n,p,2*Or);E(y);let k=B(s.getBindGroupLayout(0),[u,f]),{commandEncoder:x,ts:S}=P(s,k,1),d=h(x,f);E(x);let[b,g,W]=await Promise.all([G(l),G(S),_(d,Float32Array)]),U=Math.sqrt(W[0]);return o?(w(u,f,c,d),b!==void 0&&g!==void 0?{nrm2:U,gpuTimeMs:b+g}:{nrm2:U}):(w(i,u,f,c,d),b!==void 0&&g!==void 0?{nrm2:U,gpuTimeMs:b+g}:{nrm2:U})}var Q=64;async function qr(t,r,e,a){let o=e instanceof m;if(!(t instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!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(r<=0)return{index:0};if(e.length<(r-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");let n=await A(t,"isamax"),s=await A(t,"reduction/argmax"),i=o?e._buf:v(e,"isamax-x",!1),u=N(2*Q*4,"isamax-partials-val"),f=N(2*Q*4,"isamax-partials-idx"),c=V(4,"isamax-result"),p=R([{value:r,type:"u32"},{value:a,type:"u32"}],"isamax-params"),y=B(n.getBindGroupLayout(0),[i,u,f,p]),{commandEncoder:l,ts:k}=P(n,y,2*Q);E(l);let x=B(s.getBindGroupLayout(0),[u,f,c]),{commandEncoder:S,ts:d}=P(s,x,1),b=h(S,c);E(S);let[g,W,U]=await Promise.all([G(k),G(d),_(b,Uint32Array)]),D=U[0];return o?(w(u,f,c,p,b),g!==void 0&&W!==void 0?{index:D,gpuTimeMs:g+W}:{index:D}):(w(i,u,f,c,p,b),g!==void 0&&W!==void 0?{index:D,gpuTimeMs:g+W}:{index:D})}async function Yr(t,r,e,a,o,n,s,i){let u=e instanceof m,f=o instanceof m;if(!(t instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(a)||!Number.isInteger(n))throw new Error("n, incx, and incy must be integers.");if(isNaN(s)||isNaN(i))throw new Error("c and s must not be NaN.");if(!isFinite(s))throw new Error("c must be finite.");if(!isFinite(i))throw new Error("s must be finite.");if(a<=0||n<=0)throw new Error("incx and incy must be positive.");if(!u&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!f&&!(o instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(u!==f)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(r<=0)return u?{}:{x:e,y:o};if(e.length<(r-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(o.length<(r-1)*n+1)throw new Error("y does not have enough elements for the given n and incy.");let c=await A(t,"srot"),p=u?e._buf:v(e,"srot-x",!0),y=f?o._buf:v(o,"srot-y",!0),l=R([{value:r,type:"u32"},{value:s,type:"f32"},{value:i,type:"f32"},{value:a,type:"u32"},{value:n,type:"u32"}],"srot-params"),k=B(c.getBindGroupLayout(0),[p,y,l]),{commandEncoder:x,ts:S}=P(c,k,M(r)),d=u?null:h(x,p),b=f?null:h(x,y);E(x);let g=await G(S);if(u&&f)return w(l),g!==void 0?{gpuTimeMs:g}:{};let[W,U]=await Promise.all([_(d,Float32Array),_(b,Float32Array)]);return w(p,y,l,d,b),g!==void 0?{x:W,y:U,gpuTimeMs:g}:{x:W,y:U}}async function $r(t,r,e,a,o,n,s){let i=e instanceof m,u=o instanceof m;if(!(t instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(a)||!Number.isInteger(n))throw new Error("n, incx, and incy must be integers.");if(!(s instanceof Float32Array)||s.length!==5)throw new Error("param must be a Float32Array of length 5.");if(a<=0||n<=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(r<=0||s[0]===-2)return i?{}:{x:e,y:o};if(e.length<(r-1)*a+1)throw new Error("x does not have enough elements for the given n and incx.");if(o.length<(r-1)*n+1)throw new Error("y does not have enough elements for the given n and incy.");let f=await A(t,"srotm"),c=i?e._buf:v(e,"srotm-x",!0),p=u?o._buf:v(o,"srotm-y",!0),y=v(s,"srotm-param",!1),l=R([{value:r,type:"u32"},{value:a,type:"u32"},{value:n,type:"u32"}],"srotm-params"),k=B(f.getBindGroupLayout(0),[c,p,y,l]),{commandEncoder:x,ts:S}=P(f,k,M(r)),d=i?null:h(x,c),b=u?null:h(x,p);E(x);let g=await G(S);if(i&&u)return w(y,l),g!==void 0?{gpuTimeMs:g}:{};let[W,U]=await Promise.all([_(d,Float32Array),_(b,Float32Array)]);return w(c,p,y,l,d,b),g!==void 0?{x:W,y:U,gpuTimeMs:g}:{x:W,y:U}}return Zr(we);})();
|
|
538
|
+
`});var oe,ie=W(()=>{oe=`// sgemv_n: y = alpha * A * x + beta * y (A is m\xD7n row-major, no-transpose)
|
|
539
|
+
//
|
|
540
|
+
// One workgroup per output row, with a grid-stride outer loop so the shader
|
|
541
|
+
// still covers all rows when m exceeds maxComputeWorkgroupsPerDimension.
|
|
542
|
+
// Threads stride through A[row, :] and x with coalesced reads (consecutive
|
|
543
|
+
// threads \u2192 consecutive addresses). Four independent accumulators let the GPU
|
|
544
|
+
// pipeline memory requests across iterations (ILP=4), hiding the
|
|
545
|
+
// global-memory latency.
|
|
546
|
+
|
|
547
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
548
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
549
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
550
|
+
|
|
551
|
+
struct Params {
|
|
552
|
+
m: u32,
|
|
553
|
+
n: u32,
|
|
554
|
+
alpha: f32,
|
|
555
|
+
beta: f32,
|
|
556
|
+
incx: u32,
|
|
557
|
+
incy: u32,
|
|
558
|
+
lda: u32,
|
|
559
|
+
}
|
|
560
|
+
|
|
561
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
562
|
+
|
|
563
|
+
const WGS: u32 = 64u;
|
|
564
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
565
|
+
|
|
566
|
+
@compute @workgroup_size(64)
|
|
567
|
+
fn main(
|
|
568
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
569
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
570
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
571
|
+
) {
|
|
572
|
+
// Grid-stride loop: each workgroup handles ceil(m / nwg.x) rows.
|
|
573
|
+
for (var row = wgid.x; row < params.m; row += nwg.x) {
|
|
574
|
+
let row_base = row * params.lda;
|
|
575
|
+
var acc0: f32 = 0.0;
|
|
576
|
+
var acc1: f32 = 0.0;
|
|
577
|
+
var acc2: f32 = 0.0;
|
|
578
|
+
var acc3: f32 = 0.0;
|
|
579
|
+
|
|
580
|
+
// 4-unrolled loop: each iteration issues 4 independent loads for A and x.
|
|
581
|
+
// The accumulators are independent so the GPU can overlap the memory
|
|
582
|
+
// requests rather than serialising them behind a dependency chain.
|
|
583
|
+
let n4_floor = (params.n / (4u * WGS)) * (4u * WGS);
|
|
584
|
+
for (var j: u32 = lid.x; j < n4_floor; j += 4u * WGS) {
|
|
585
|
+
acc0 += A[row_base + j ] * x[ j * params.incx];
|
|
586
|
+
acc1 += A[row_base + j + WGS ] * x[(j + WGS) * params.incx];
|
|
587
|
+
acc2 += A[row_base + j + 2u * WGS ] * x[(j + 2u * WGS) * params.incx];
|
|
588
|
+
acc3 += A[row_base + j + 3u * WGS ] * x[(j + 3u * WGS) * params.incx];
|
|
589
|
+
}
|
|
590
|
+
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
591
|
+
for (var j: u32 = n4_floor + lid.x; j < params.n; j += WGS) {
|
|
592
|
+
acc0 += A[row_base + j] * x[j * params.incx];
|
|
593
|
+
}
|
|
594
|
+
|
|
595
|
+
// Parallel reduction: 64 \u2192 32 \u2192 16 \u2192 8 \u2192 4 \u2192 2 \u2192 1
|
|
596
|
+
scratch[lid.x] = acc0 + acc1 + acc2 + acc3;
|
|
597
|
+
workgroupBarrier();
|
|
598
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
599
|
+
if lid.x < stride {
|
|
600
|
+
scratch[lid.x] += scratch[lid.x + stride];
|
|
601
|
+
}
|
|
602
|
+
workgroupBarrier();
|
|
603
|
+
}
|
|
604
|
+
|
|
605
|
+
if lid.x == 0u {
|
|
606
|
+
let yi = row * params.incy;
|
|
607
|
+
y[yi] = params.alpha * scratch[0] + params.beta * y[yi];
|
|
608
|
+
}
|
|
609
|
+
// All 64 threads must agree before the next row reuses scratch[].
|
|
610
|
+
workgroupBarrier();
|
|
611
|
+
}
|
|
612
|
+
}
|
|
613
|
+
`});var se,ne=W(()=>{se=`// sgemv_t: y = alpha * A^T * x + beta * y (A is m\xD7n row-major, transposed)
|
|
614
|
+
// each thread owns one column of A \u2192 one element of y (length n)
|
|
615
|
+
// tiles over x (length m) using shared memory; four independent accumulators
|
|
616
|
+
// let the GPU pipeline A reads across j within each tile (ILP=4)
|
|
617
|
+
|
|
618
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
619
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
620
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
621
|
+
|
|
622
|
+
struct Params {
|
|
623
|
+
m: u32,
|
|
624
|
+
n: u32,
|
|
625
|
+
alpha: f32,
|
|
626
|
+
beta: f32,
|
|
627
|
+
incx: u32,
|
|
628
|
+
incy: u32,
|
|
629
|
+
lda: u32,
|
|
630
|
+
}
|
|
631
|
+
|
|
632
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
633
|
+
|
|
634
|
+
const WGS: u32 = 64u;
|
|
635
|
+
var<workgroup> x_tile: array<f32, 64>;
|
|
636
|
+
|
|
637
|
+
@compute @workgroup_size(64)
|
|
638
|
+
fn main(
|
|
639
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
640
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
641
|
+
) {
|
|
642
|
+
// each thread owns column col of A \u2192 output y[col]
|
|
643
|
+
let col = gid.x;
|
|
644
|
+
// tile over x (length m, the rows of A)
|
|
645
|
+
let m_floor = (params.m / WGS) * WGS;
|
|
646
|
+
var acc0: f32 = 0.0;
|
|
647
|
+
var acc1: f32 = 0.0;
|
|
648
|
+
var acc2: f32 = 0.0;
|
|
649
|
+
var acc3: f32 = 0.0;
|
|
650
|
+
|
|
651
|
+
for (var base = 0u; base < m_floor; base += WGS) {
|
|
652
|
+
// cooperative load: all 64 threads fill x_tile with x[base..base+WGS]
|
|
653
|
+
x_tile[lid.x] = x[(base + lid.x) * params.incx];
|
|
654
|
+
workgroupBarrier();
|
|
655
|
+
|
|
656
|
+
// 4-unrolled inner loop: 4 independent A reads let the GPU pipeline
|
|
657
|
+
// global-memory requests within each tile. WGS=64 divides by 4 exactly.
|
|
658
|
+
if (col < params.n) {
|
|
659
|
+
for (var j = 0u; j < WGS; j += 4u) {
|
|
660
|
+
acc0 += A[(base + j ) * params.lda + col] * x_tile[j ];
|
|
661
|
+
acc1 += A[(base + j + 1) * params.lda + col] * x_tile[j + 1];
|
|
662
|
+
acc2 += A[(base + j + 2) * params.lda + col] * x_tile[j + 2];
|
|
663
|
+
acc3 += A[(base + j + 3) * params.lda + col] * x_tile[j + 3];
|
|
664
|
+
}
|
|
665
|
+
}
|
|
666
|
+
workgroupBarrier();
|
|
667
|
+
}
|
|
668
|
+
|
|
669
|
+
if (col < params.n) {
|
|
670
|
+
// remainder: m not divisible by WGS \u2014 short loop, single accumulator fine
|
|
671
|
+
for (var k = m_floor; k < params.m; k++) {
|
|
672
|
+
acc0 += A[k * params.lda + col] * x[k * params.incx];
|
|
673
|
+
}
|
|
674
|
+
let yi = col * params.incy;
|
|
675
|
+
y[yi] = params.alpha * (acc0 + acc1 + acc2 + acc3) + params.beta * y[yi];
|
|
676
|
+
}
|
|
677
|
+
}
|
|
678
|
+
`});var le,ue=W(()=>{le=`// ssymv: y = alpha * A * x + beta * y
|
|
679
|
+
// A is n\xD7n symmetric, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
680
|
+
// The logical matrix is fully dense (symmetric), so each row's dot product
|
|
681
|
+
// sums over all n columns; entries on the unstored side of the diagonal are
|
|
682
|
+
// fetched from their mirror position (A[i,j] == A[j,i]).
|
|
683
|
+
// One workgroup per row, grid-stride outer loop.
|
|
684
|
+
|
|
685
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
686
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
687
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
688
|
+
|
|
689
|
+
struct Params {
|
|
690
|
+
n: u32,
|
|
691
|
+
alpha: f32,
|
|
692
|
+
beta: f32,
|
|
693
|
+
incx: u32,
|
|
694
|
+
incy: u32,
|
|
695
|
+
lda: u32,
|
|
696
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
697
|
+
}
|
|
698
|
+
|
|
699
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
700
|
+
|
|
701
|
+
const WGS: u32 = 64u;
|
|
702
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
703
|
+
|
|
704
|
+
@compute @workgroup_size(64)
|
|
705
|
+
fn main(
|
|
706
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
707
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
708
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
709
|
+
) {
|
|
710
|
+
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
711
|
+
var acc = 0.0f;
|
|
712
|
+
|
|
713
|
+
// y[i] = \u03A3_j A[i,j] * x[j]
|
|
714
|
+
for (var j = lid.x; j < params.n; j += WGS) {
|
|
715
|
+
var aVal: f32;
|
|
716
|
+
if params.uplo == 0u {
|
|
717
|
+
// Lower: A[i,j] stored at A[i*lda+j] for j \u2264 i, mirrored from A[j*lda+i] otherwise
|
|
718
|
+
if j <= i {
|
|
719
|
+
aVal = A[i * params.lda + j];
|
|
720
|
+
} else {
|
|
721
|
+
aVal = A[j * params.lda + i];
|
|
722
|
+
}
|
|
723
|
+
} else {
|
|
724
|
+
// Upper: A[i,j] stored at A[i*lda+j] for j \u2265 i, mirrored from A[j*lda+i] otherwise
|
|
725
|
+
if j >= i {
|
|
726
|
+
aVal = A[i * params.lda + j];
|
|
727
|
+
} else {
|
|
728
|
+
aVal = A[j * params.lda + i];
|
|
729
|
+
}
|
|
730
|
+
}
|
|
731
|
+
acc += aVal * x[j * params.incx];
|
|
732
|
+
}
|
|
733
|
+
|
|
734
|
+
// Parallel reduction: 64 \u2192 1
|
|
735
|
+
scratch[lid.x] = acc;
|
|
736
|
+
workgroupBarrier();
|
|
737
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
738
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
739
|
+
workgroupBarrier();
|
|
740
|
+
}
|
|
741
|
+
|
|
742
|
+
if lid.x == 0u {
|
|
743
|
+
y[i * params.incy] = params.alpha * scratch[0] + params.beta * y[i * params.incy];
|
|
744
|
+
}
|
|
745
|
+
}
|
|
746
|
+
}
|
|
747
|
+
`});var fe,ce=W(()=>{fe=`// strmv: y = op(A) * x
|
|
748
|
+
// A is n\xD7n triangular, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
749
|
+
// op(A) is A (trans=0) or A^T (trans=1).
|
|
750
|
+
// diag=1 (unit) treats the diagonal as 1 without reading A's diagonal values.
|
|
751
|
+
// One workgroup per row, grid-stride outer loop.
|
|
752
|
+
|
|
753
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
754
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
755
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
756
|
+
|
|
757
|
+
struct Params {
|
|
758
|
+
n: u32,
|
|
759
|
+
incx: u32,
|
|
760
|
+
incy: u32,
|
|
761
|
+
lda: u32,
|
|
762
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
763
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
764
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
765
|
+
}
|
|
766
|
+
|
|
767
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
768
|
+
|
|
769
|
+
const WGS: u32 = 64u;
|
|
770
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
771
|
+
|
|
772
|
+
@compute @workgroup_size(64)
|
|
773
|
+
fn main(
|
|
774
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
775
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
776
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
777
|
+
) {
|
|
778
|
+
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
779
|
+
var acc = 0.0f;
|
|
780
|
+
|
|
781
|
+
if params.trans == 0u {
|
|
782
|
+
// No-transpose: y[i] = \u03A3_j A[i,j] * x[j]
|
|
783
|
+
if params.uplo == 0u {
|
|
784
|
+
// Lower: A[i,j] stored at A[i*lda+j] for j \u2264 i
|
|
785
|
+
for (var j = lid.x; j <= i; j += WGS) {
|
|
786
|
+
var aVal: f32;
|
|
787
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
788
|
+
if params.diag == 1u && j == i {
|
|
789
|
+
aVal = 1.0;
|
|
790
|
+
} else if ( j <= i ) {
|
|
791
|
+
aVal = A[i * params.lda + j];
|
|
792
|
+
}
|
|
793
|
+
acc += aVal * x[j * params.incx];
|
|
794
|
+
}
|
|
795
|
+
} else {
|
|
796
|
+
// Upper: A[i,j] stored at A[i*lda+j] for j \u2265 i
|
|
797
|
+
for (var j = i + lid.x; j < params.n; j += WGS) {
|
|
798
|
+
var aVal: f32;
|
|
799
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
800
|
+
if params.diag == 1u && j == i {
|
|
801
|
+
aVal = 1.0;
|
|
802
|
+
} else if ( j >= i ) {
|
|
803
|
+
aVal = A[i * params.lda + j];
|
|
804
|
+
}
|
|
805
|
+
acc += aVal * x[j * params.incx];
|
|
806
|
+
}
|
|
807
|
+
}
|
|
808
|
+
} else {
|
|
809
|
+
// Transpose: y[i] = \u03A3_j A[j,i] * x[j]
|
|
810
|
+
if params.uplo == 0u {
|
|
811
|
+
// Lower: A[j,i] stored at A[j*lda+i] for j \u2265 i
|
|
812
|
+
for (var j = i + lid.x; j < params.n; j += WGS) {
|
|
813
|
+
var aVal: f32;
|
|
814
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
815
|
+
if params.diag == 1u && j == i {
|
|
816
|
+
aVal = 1.0;
|
|
817
|
+
} else if ( j >= i ) {
|
|
818
|
+
aVal = A[j * params.lda + i];
|
|
819
|
+
}
|
|
820
|
+
acc += aVal * x[j * params.incx];
|
|
821
|
+
}
|
|
822
|
+
} else {
|
|
823
|
+
// Upper: A[j,i] stored at A[j*lda+i] for j \u2264 i
|
|
824
|
+
for (var j = lid.x; j <= i; j += WGS) {
|
|
825
|
+
var aVal: f32;
|
|
826
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
827
|
+
if params.diag == 1u && j == i {
|
|
828
|
+
aVal = 1.0;
|
|
829
|
+
} else if ( j <= i ) {
|
|
830
|
+
aVal = A[j * params.lda + i];
|
|
831
|
+
}
|
|
832
|
+
acc += aVal * x[j * params.incx];
|
|
833
|
+
}
|
|
834
|
+
}
|
|
835
|
+
}
|
|
836
|
+
|
|
837
|
+
// Parallel reduction: 64 \u2192 1
|
|
838
|
+
scratch[lid.x] = acc;
|
|
839
|
+
workgroupBarrier();
|
|
840
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
841
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
842
|
+
workgroupBarrier();
|
|
843
|
+
}
|
|
844
|
+
|
|
845
|
+
if lid.x == 0u {
|
|
846
|
+
y[ i * params.incy ] = scratch[0];
|
|
847
|
+
}
|
|
848
|
+
}
|
|
849
|
+
}
|
|
850
|
+
`});var me,de=W(()=>{me=`// f64add: adds two doubles, each packed as a [main, aux] pair (main: f32
|
|
851
|
+
// value, aux: raw u32 bits \u2014 see src/util/f64pack.mjs; decode()/encode()
|
|
852
|
+
// below are the WGSL mirror of that file's packedToFields()/fieldsToPacked()),
|
|
853
|
+
// producing the sum as another [main, aux] pair.
|
|
854
|
+
//
|
|
855
|
+
// Implements IEEE-754 binary64 addition (align, add/subtract significands,
|
|
856
|
+
// normalize, round-to-nearest-even) using only u32 bitwise/integer
|
|
857
|
+
// arithmetic \u2014 WGSL has no 64-bit integer type or arbitrary-precision
|
|
858
|
+
// integers, so each operand's 53-bit significand is carried as a two-word
|
|
859
|
+
// (hi, lo) pair, widened by 3 bits at the bottom to hold guard/round/sticky
|
|
860
|
+
// information while aligning exponents.
|
|
861
|
+
|
|
862
|
+
const EXP_ALL_ONES: u32 = 0x7ffu;
|
|
863
|
+
const BIAS: i32 = 1023;
|
|
864
|
+
const QUIET_NAN_MANTISSA_HI: u32 = 1u << 19u; // bit51 of the 52-bit mantissa -> canonical quiet NaN
|
|
865
|
+
|
|
866
|
+
struct Fields {
|
|
867
|
+
sign: u32,
|
|
868
|
+
rawExp: u32,
|
|
869
|
+
mantissaHi: u32, // 20 bits
|
|
870
|
+
lo: u32, // 32 bits
|
|
871
|
+
}
|
|
872
|
+
|
|
873
|
+
// A packed [main, aux] result \u2014 aux stays a raw u32; it must never be stored
|
|
874
|
+
// as an array<f32>/treated as a real float (bit pattern can land on a NaN/
|
|
875
|
+
// Infinity exponent for perfectly ordinary doubles \u2014 an f32-typed storage
|
|
876
|
+
// slot canonicalizes/corrupts that on any round-trip). See f64pack.mjs's
|
|
877
|
+
// comment above fieldsToPacked, and dasum.wgsl's xAux/partialsAux bindings.
|
|
878
|
+
struct Packed {
|
|
879
|
+
main: f32,
|
|
880
|
+
aux: u32,
|
|
881
|
+
}
|
|
882
|
+
|
|
883
|
+
// Mirrors packedToFields() in f64pack.mjs.
|
|
884
|
+
fn decode(mainBits: u32, auxBits: u32) -> Fields {
|
|
885
|
+
let sign = mainBits >> 31u;
|
|
886
|
+
let expMain = (mainBits >> 23u) & 0xffu;
|
|
887
|
+
let mantMain = mainBits & 0x7fffffu;
|
|
888
|
+
|
|
889
|
+
let auxSign = auxBits >> 31u;
|
|
890
|
+
let auxExp8 = (auxBits >> 23u) & 0xffu;
|
|
891
|
+
let auxMant23 = auxBits & 0x7fffffu;
|
|
892
|
+
|
|
893
|
+
let expExtra = (auxSign << 2u) | (auxExp8 >> 6u);
|
|
894
|
+
let mantExtra29 = ((auxExp8 & 0x3fu) << 23u) | auxMant23;
|
|
895
|
+
|
|
896
|
+
let rawExp = (expMain << 3u) | expExtra;
|
|
897
|
+
let mantissaHi = mantMain >> 3u;
|
|
898
|
+
let mantTop3 = mantMain & 0x7u;
|
|
899
|
+
let lo = (mantTop3 << 29u) | mantExtra29;
|
|
900
|
+
|
|
901
|
+
return Fields(sign, rawExp, mantissaHi, lo);
|
|
902
|
+
}
|
|
903
|
+
|
|
904
|
+
// Mirrors fieldsToPacked() in f64pack.mjs.
|
|
905
|
+
fn encode(sign: u32, rawExp: u32, mantissaHi: u32, lo: u32) -> Packed {
|
|
906
|
+
let expMain = rawExp >> 3u;
|
|
907
|
+
let expExtra = rawExp & 0x7u;
|
|
908
|
+
|
|
909
|
+
let mantTop3 = lo >> 29u;
|
|
910
|
+
let mantMain = (mantissaHi << 3u) | mantTop3;
|
|
911
|
+
let mantExtra29 = lo & 0x1fffffffu;
|
|
912
|
+
|
|
913
|
+
let mainBits = (sign << 31u) | (expMain << 23u) | mantMain;
|
|
914
|
+
|
|
915
|
+
let auxSign = (expExtra >> 2u) & 0x1u;
|
|
916
|
+
let auxExpTop2 = expExtra & 0x3u;
|
|
917
|
+
let auxExpBot6 = mantExtra29 >> 23u;
|
|
918
|
+
let auxMant23 = mantExtra29 & 0x7fffffu;
|
|
919
|
+
let auxExp8 = (auxExpTop2 << 6u) | auxExpBot6;
|
|
920
|
+
|
|
921
|
+
let auxBits = (auxSign << 31u) | (auxExp8 << 23u) | auxMant23;
|
|
922
|
+
|
|
923
|
+
return Packed(bitcast<f32>(mainBits), auxBits);
|
|
924
|
+
}
|
|
925
|
+
|
|
926
|
+
struct Pair { hi: u32, lo: u32 }
|
|
927
|
+
struct Shifted { hi: u32, lo: u32, sticky: u32 }
|
|
928
|
+
|
|
929
|
+
// Two-word right shift by 0..64+ bits, folding every shifted-out 1 bit into
|
|
930
|
+
// a returned sticky flag \u2014 used only for the (potentially huge) exponent
|
|
931
|
+
// alignment shift, where exact bits can't all be kept.
|
|
932
|
+
fn shr_sticky(hi: u32, lo: u32, n: u32) -> Shifted {
|
|
933
|
+
if (n == 0u) {
|
|
934
|
+
return Shifted(hi, lo, 0u);
|
|
935
|
+
}
|
|
936
|
+
if (n >= 64u) {
|
|
937
|
+
return Shifted(0u, 0u, select(0u, 1u, hi != 0u || lo != 0u));
|
|
938
|
+
}
|
|
939
|
+
if (n < 32u) {
|
|
940
|
+
let stickyBits = lo & ((1u << n) - 1u);
|
|
941
|
+
let newLo = (lo >> n) | (hi << (32u - n));
|
|
942
|
+
let newHi = hi >> n;
|
|
943
|
+
return Shifted(newHi, newLo, select(0u, 1u, stickyBits != 0u));
|
|
944
|
+
}
|
|
945
|
+
if (n == 32u) {
|
|
946
|
+
return Shifted(0u, hi, select(0u, 1u, lo != 0u));
|
|
947
|
+
}
|
|
948
|
+
let m = n - 32u;
|
|
949
|
+
let stickyBits = lo | (hi & ((1u << m) - 1u));
|
|
950
|
+
let newLo = hi >> m;
|
|
951
|
+
return Shifted(0u, newLo, select(0u, 1u, stickyBits != 0u));
|
|
952
|
+
}
|
|
953
|
+
|
|
954
|
+
// Two-word left shift by 0..63 bits \u2014 used only to renormalize after
|
|
955
|
+
// cancellation, by an amount that exactly matches the leading-zero count,
|
|
956
|
+
// so nothing meaningful is ever lost off the top.
|
|
957
|
+
fn shl(hi: u32, lo: u32, n: u32) -> Pair {
|
|
958
|
+
if (n == 0u) {
|
|
959
|
+
return Pair(hi, lo);
|
|
960
|
+
}
|
|
961
|
+
if (n < 32u) {
|
|
962
|
+
let newHi = (hi << n) | (lo >> (32u - n));
|
|
963
|
+
let newLo = lo << n;
|
|
964
|
+
return Pair(newHi, newLo);
|
|
965
|
+
}
|
|
966
|
+
if (n == 32u) {
|
|
967
|
+
return Pair(lo, 0u);
|
|
968
|
+
}
|
|
969
|
+
let m = n - 32u;
|
|
970
|
+
return Pair(lo << m, 0u);
|
|
971
|
+
}
|
|
972
|
+
|
|
973
|
+
fn add64(aHi: u32, aLo: u32, bHi: u32, bLo: u32) -> Pair {
|
|
974
|
+
let sumLo = aLo + bLo;
|
|
975
|
+
let carry = select(0u, 1u, sumLo < aLo); // wrapped around -> there was a carry
|
|
976
|
+
let sumHi = aHi + bHi + carry;
|
|
977
|
+
return Pair(sumHi, sumLo);
|
|
978
|
+
}
|
|
979
|
+
|
|
980
|
+
// Assumes (aHi:aLo) >= (bHi:bLo) \u2014 callers guarantee this so no sign handling is needed.
|
|
981
|
+
fn sub64(aHi: u32, aLo: u32, bHi: u32, bLo: u32) -> Pair {
|
|
982
|
+
let borrow = select(0u, 1u, aLo < bLo);
|
|
983
|
+
let diffLo = aLo - bLo;
|
|
984
|
+
let diffHi = aHi - bHi - borrow;
|
|
985
|
+
return Pair(diffHi, diffLo);
|
|
986
|
+
}
|
|
987
|
+
|
|
988
|
+
fn ge64(aHi: u32, aLo: u32, bHi: u32, bLo: u32) -> bool {
|
|
989
|
+
return aHi > bHi || (aHi == bHi && aLo >= bLo);
|
|
990
|
+
}
|
|
991
|
+
|
|
992
|
+
// The actual IEEE-754 addition, returning decoded Fields rather than an
|
|
993
|
+
// encoded Packed pair \u2014 lets a caller that's accumulating many values in a
|
|
994
|
+
// row (e.g. dasum.wgsl's per-thread reduction loop) keep the running total
|
|
995
|
+
// in Fields form the whole time, only encoding once at the very end, instead
|
|
996
|
+
// of paying a decode+encode round-trip on every single addition. computeSum
|
|
997
|
+
// (below) is the Packed-in/Packed-out convenience wrapper around this.
|
|
998
|
+
fn addFields(a: Fields, b: Fields) -> Fields {
|
|
999
|
+
let aIsNaN = a.rawExp == EXP_ALL_ONES && (a.mantissaHi != 0u || a.lo != 0u);
|
|
1000
|
+
let bIsNaN = b.rawExp == EXP_ALL_ONES && (b.mantissaHi != 0u || b.lo != 0u);
|
|
1001
|
+
if (aIsNaN || bIsNaN) {
|
|
1002
|
+
return Fields(0u, EXP_ALL_ONES, QUIET_NAN_MANTISSA_HI, 0u);
|
|
1003
|
+
}
|
|
1004
|
+
|
|
1005
|
+
let aIsInf = a.rawExp == EXP_ALL_ONES; // mantissa==0 here since NaN is excluded above
|
|
1006
|
+
let bIsInf = b.rawExp == EXP_ALL_ONES;
|
|
1007
|
+
if (aIsInf && bIsInf) {
|
|
1008
|
+
if (a.sign != b.sign) {
|
|
1009
|
+
return Fields(0u, EXP_ALL_ONES, QUIET_NAN_MANTISSA_HI, 0u);
|
|
1010
|
+
}
|
|
1011
|
+
return Fields(a.sign, EXP_ALL_ONES, 0u, 0u);
|
|
1012
|
+
}
|
|
1013
|
+
if (aIsInf) { return Fields(a.sign, EXP_ALL_ONES, 0u, 0u); }
|
|
1014
|
+
if (bIsInf) { return Fields(b.sign, EXP_ALL_ONES, 0u, 0u); }
|
|
1015
|
+
|
|
1016
|
+
let aIsZero = a.rawExp == 0u && a.mantissaHi == 0u && a.lo == 0u;
|
|
1017
|
+
let bIsZero = b.rawExp == 0u && b.mantissaHi == 0u && b.lo == 0u;
|
|
1018
|
+
if (aIsZero && bIsZero) {
|
|
1019
|
+
return Fields(a.sign & b.sign, 0u, 0u, 0u);
|
|
1020
|
+
}
|
|
1021
|
+
if (aIsZero) { return Fields(b.sign, b.rawExp, b.mantissaHi, b.lo); }
|
|
1022
|
+
if (bIsZero) { return Fields(a.sign, a.rawExp, a.mantissaHi, a.lo); }
|
|
1023
|
+
|
|
1024
|
+
// Effective (unbiased) exponent \u2014 subnormals share the smallest normal
|
|
1025
|
+
// exponent for alignment purposes and have no implicit leading 1.
|
|
1026
|
+
var expA = i32(a.rawExp) - BIAS;
|
|
1027
|
+
if (a.rawExp == 0u) { expA = 1 - BIAS; }
|
|
1028
|
+
var expB = i32(b.rawExp) - BIAS;
|
|
1029
|
+
if (b.rawExp == 0u) { expB = 1 - BIAS; }
|
|
1030
|
+
|
|
1031
|
+
let implicitA = select(0u, 1u, a.rawExp != 0u);
|
|
1032
|
+
let implicitB = select(0u, 1u, b.rawExp != 0u);
|
|
1033
|
+
|
|
1034
|
+
// Widen each 53-bit significand (implicit + 52-bit mantissa) by 3 zero
|
|
1035
|
+
// bits at the bottom \u2014 room for guard/round/sticky once alignment shifts happen.
|
|
1036
|
+
let sigHiA = (implicitA << 23u) | (a.mantissaHi << 3u) | (a.lo >> 29u);
|
|
1037
|
+
let sigLoA = a.lo << 3u;
|
|
1038
|
+
let sigHiB = (implicitB << 23u) | (b.mantissaHi << 3u) | (b.lo >> 29u);
|
|
1039
|
+
let sigLoB = b.lo << 3u;
|
|
1040
|
+
|
|
1041
|
+
// P = the operand with the larger exponent (Q = the other); on a tie, P =
|
|
1042
|
+
// whichever has the larger significand \u2014 keeps subtraction below always
|
|
1043
|
+
// non-negative without needing signed magnitudes.
|
|
1044
|
+
var signP: u32; var expP: i32; var sigHiP: u32; var sigLoP: u32;
|
|
1045
|
+
var signQ: u32; var expQ: i32; var sigHiQ: u32; var sigLoQ: u32;
|
|
1046
|
+
if (expA > expB || (expA == expB && ge64(sigHiA, sigLoA, sigHiB, sigLoB))) {
|
|
1047
|
+
signP = a.sign; expP = expA; sigHiP = sigHiA; sigLoP = sigLoA;
|
|
1048
|
+
signQ = b.sign; expQ = expB; sigHiQ = sigHiB; sigLoQ = sigLoB;
|
|
1049
|
+
} else {
|
|
1050
|
+
signP = b.sign; expP = expB; sigHiP = sigHiB; sigLoP = sigLoB;
|
|
1051
|
+
signQ = a.sign; expQ = expA; sigHiQ = sigHiA; sigLoQ = sigLoA;
|
|
1052
|
+
}
|
|
1053
|
+
|
|
1054
|
+
let diff = u32(expP - expQ);
|
|
1055
|
+
let shiftedQ = shr_sticky(sigHiQ, sigLoQ, diff);
|
|
1056
|
+
let alignedHiQ = shiftedQ.hi;
|
|
1057
|
+
let alignedLoQ = shiftedQ.lo | shiftedQ.sticky; // fold sticky into bit0
|
|
1058
|
+
|
|
1059
|
+
var sumHi: u32; var sumLo: u32;
|
|
1060
|
+
if (signP == signQ) {
|
|
1061
|
+
let s = add64(sigHiP, sigLoP, alignedHiQ, alignedLoQ);
|
|
1062
|
+
sumHi = s.hi; sumLo = s.lo;
|
|
1063
|
+
} else {
|
|
1064
|
+
let s = sub64(sigHiP, sigLoP, alignedHiQ, alignedLoQ); // P >= Q(aligned) by construction
|
|
1065
|
+
sumHi = s.hi; sumLo = s.lo;
|
|
1066
|
+
}
|
|
1067
|
+
|
|
1068
|
+
if (sumHi == 0u && sumLo == 0u) {
|
|
1069
|
+
return Fields(0u, 0u, 0u, 0u); // exact cancellation -> +0
|
|
1070
|
+
}
|
|
1071
|
+
|
|
1072
|
+
// commonExp2: per-bit scale of sumLo's bit0 in the widened representation.
|
|
1073
|
+
let commonExp2 = expP - 55;
|
|
1074
|
+
|
|
1075
|
+
var leadPos: i32;
|
|
1076
|
+
if (sumHi != 0u) {
|
|
1077
|
+
leadPos = 32 + i32(31u - countLeadingZeros(sumHi));
|
|
1078
|
+
} else {
|
|
1079
|
+
leadPos = i32(31u - countLeadingZeros(sumLo));
|
|
1080
|
+
}
|
|
1081
|
+
let tentativeExp = leadPos + commonExp2;
|
|
1082
|
+
var targetLSBScale = tentativeExp - 52;
|
|
1083
|
+
if (tentativeExp < -1022) { targetLSBScale = -1074; }
|
|
1084
|
+
let shiftAmt = targetLSBScale - commonExp2;
|
|
1085
|
+
|
|
1086
|
+
var keepHi: u32; var keepLo: u32;
|
|
1087
|
+
if (shiftAmt <= 0) {
|
|
1088
|
+
let sh = shl(sumHi, sumLo, u32(-shiftAmt)); // exact \u2014 cancellation only, never loses bits
|
|
1089
|
+
keepHi = sh.hi; keepLo = sh.lo;
|
|
1090
|
+
} else {
|
|
1091
|
+
// Only reached without cancellation (same-sign add, or a tied-exponent
|
|
1092
|
+
// subtract with no shrinkage) \u2014 shiftAmt here is always exactly 3 or 4,
|
|
1093
|
+
// so the dropped bits are fully known from sumLo directly (no sticky
|
|
1094
|
+
// approximation needed, unlike the Q-alignment shift above).
|
|
1095
|
+
let n = u32(shiftAmt);
|
|
1096
|
+
let remainder = sumLo & ((1u << n) - 1u);
|
|
1097
|
+
let halfway = 1u << (n - 1u);
|
|
1098
|
+
let sh = shr_sticky(sumHi, sumLo, n);
|
|
1099
|
+
keepHi = sh.hi; keepLo = sh.lo;
|
|
1100
|
+
if (remainder > halfway || (remainder == halfway && (keepLo & 1u) != 0u)) {
|
|
1101
|
+
let inc = add64(keepHi, keepLo, 0u, 1u);
|
|
1102
|
+
keepHi = inc.hi; keepLo = inc.lo;
|
|
1103
|
+
}
|
|
1104
|
+
}
|
|
1105
|
+
|
|
1106
|
+
var resultExpBase = targetLSBScale;
|
|
1107
|
+
if ((keepHi & (1u << 21u)) != 0u) { // rounding carried past the 53-bit budget
|
|
1108
|
+
let sh = shr_sticky(keepHi, keepLo, 1u); // dropped bit is guaranteed 0 here
|
|
1109
|
+
keepHi = sh.hi; keepLo = sh.lo;
|
|
1110
|
+
resultExpBase = resultExpBase + 1;
|
|
1111
|
+
}
|
|
1112
|
+
|
|
1113
|
+
let resultSign = signP;
|
|
1114
|
+
if ((keepHi & (1u << 20u)) != 0u) { // normal-shaped result
|
|
1115
|
+
let unbiasedExp = 52 + resultExpBase;
|
|
1116
|
+
let rawExpFinal = unbiasedExp + BIAS;
|
|
1117
|
+
if (rawExpFinal >= 2047) {
|
|
1118
|
+
return Fields(resultSign, EXP_ALL_ONES, 0u, 0u); // overflow -> Infinity
|
|
1119
|
+
}
|
|
1120
|
+
return Fields(resultSign, u32(rawExpFinal), keepHi & 0xfffffu, keepLo);
|
|
1121
|
+
}
|
|
1122
|
+
return Fields(resultSign, 0u, keepHi & 0xfffffu, keepLo); // subnormal result
|
|
1123
|
+
}
|
|
1124
|
+
|
|
1125
|
+
// Packed-in/Packed-out convenience wrapper around addFields \u2014 encodes once,
|
|
1126
|
+
// after the math, rather than addFields itself needing to know about Packed.
|
|
1127
|
+
fn computeSum(a: Fields, b: Fields) -> Packed {
|
|
1128
|
+
let f = addFields(a, b);
|
|
1129
|
+
return encode(f.sign, f.rawExp, f.mantissaHi, f.lo);
|
|
1130
|
+
}
|
|
1131
|
+
`});var ge,pe=W(()=>{ge=`// dasum: result = sum(|x[i]|)
|
|
1132
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sumF64.wgsl.
|
|
1133
|
+
// Same structure as sasum.wgsl \u2014 every value is now a [main, aux] pair
|
|
1134
|
+
// (see src/util/f64pack.mjs) and every \`+\`/\`+=\` is computeSum via addPair
|
|
1135
|
+
// instead of plain f32 addition. Concatenated after f64add.wgsl by
|
|
1136
|
+
// getPipeline (WGSL has no #include), reusing its decode/encode/computeSum/
|
|
1137
|
+
// addFields and Packed struct \u2014 f64add.wgsl declares no bindings and no entry
|
|
1138
|
+
// point of its own (just helper functions), so bindings here start at 0 and
|
|
1139
|
+
// the entry point is simply \`dasum_main\`.
|
|
1140
|
+
//
|
|
1141
|
+
// xAux/partialsAux are array<u32>, not array<f32> \u2014 aux's bits must never
|
|
1142
|
+
// pass through an f32-typed storage slot (NaN-bit-pattern corruption risk,
|
|
1143
|
+
// see f64pack.mjs and the Packed struct comment above decode()/encode() in
|
|
1144
|
+
// f64add.wgsl); Packed (from f64add.wgsl) keeps aux as u32 in registers/
|
|
1145
|
+
// workgroup memory too.
|
|
1146
|
+
//
|
|
1147
|
+
// Per-thread accumulation (acc0..acc3) stays in DECODED Fields form for the
|
|
1148
|
+
// entire strided loop below, via addFields \u2014 not re-encoded to Packed and
|
|
1149
|
+
// re-decoded on every single element like a naive version would. Only the
|
|
1150
|
+
// freshly-loaded x[idx] needs decoding each iteration (unavoidable, it's new
|
|
1151
|
+
// data every time); the running total never leaves Fields form until the
|
|
1152
|
+
// four accumulators are combined and encoded exactly once, right before
|
|
1153
|
+
// writing into workgroup-shared \`tile\`. The cross-thread reduction tree
|
|
1154
|
+
// after that still goes through Packed per level (unavoidable \u2014 each level
|
|
1155
|
+
// combines values that live in different threads' registers via shared
|
|
1156
|
+
// memory), but that's a fixed 6 levels regardless of n, unlike the strided
|
|
1157
|
+
// loop above whose iteration count scales with n.
|
|
1158
|
+
|
|
1159
|
+
@group(0) @binding(0) var<storage, read> xMain: array<f32>;
|
|
1160
|
+
@group(0) @binding(1) var<storage, read> xAux: array<u32>;
|
|
1161
|
+
@group(0) @binding(2) var<storage, read_write> partialsMain: array<f32>;
|
|
1162
|
+
@group(0) @binding(3) var<storage, read_write> partialsAux: array<u32>;
|
|
1163
|
+
@group(0) @binding(4) var<uniform> params: Params;
|
|
1164
|
+
|
|
1165
|
+
struct Params {
|
|
1166
|
+
n: u32,
|
|
1167
|
+
x_inc: u32,
|
|
1168
|
+
}
|
|
1169
|
+
|
|
1170
|
+
const WGS: u32 = 64;
|
|
1171
|
+
|
|
1172
|
+
var<workgroup> tile: array<Packed, 64>;
|
|
1173
|
+
|
|
1174
|
+
// a + b, where a/b are [main, aux] pairs \u2014 computeSum takes decoded Fields.
|
|
1175
|
+
// Only used for the cross-thread reduction tree below; the per-thread
|
|
1176
|
+
// strided loop uses addFields directly instead (see module comment).
|
|
1177
|
+
fn addPair(a: Packed, b: Packed) -> Packed {
|
|
1178
|
+
return computeSum(decode(bitcast<u32>(a.main), a.aux), decode(bitcast<u32>(b.main), b.aux));
|
|
1179
|
+
}
|
|
1180
|
+
|
|
1181
|
+
// |x| for a packed double is abs(main) with aux untouched \u2014 only main's
|
|
1182
|
+
// sign bit carries the double's sign (see fieldsToPacked() in f64pack.mjs).
|
|
1183
|
+
// Returns decoded Fields directly (not Packed) for the per-thread loop.
|
|
1184
|
+
fn absFields(idx: u32) -> Fields {
|
|
1185
|
+
return decode(bitcast<u32>(abs(xMain[idx])), xAux[idx]);
|
|
1186
|
+
}
|
|
1187
|
+
|
|
1188
|
+
@compute @workgroup_size(64)
|
|
1189
|
+
fn dasum_main(
|
|
1190
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
1191
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1192
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1193
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
1194
|
+
) {
|
|
1195
|
+
var acc0: Fields = Fields(0u, 0u, 0u, 0u);
|
|
1196
|
+
var acc1: Fields = Fields(0u, 0u, 0u, 0u);
|
|
1197
|
+
var acc2: Fields = Fields(0u, 0u, 0u, 0u);
|
|
1198
|
+
var acc3: Fields = Fields(0u, 0u, 0u, 0u);
|
|
1199
|
+
|
|
1200
|
+
let stride = num_wg.x * WGS;
|
|
1201
|
+
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
1202
|
+
|
|
1203
|
+
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
1204
|
+
acc0 = addFields(acc0, absFields( id * params.x_inc));
|
|
1205
|
+
acc1 = addFields(acc1, absFields((id + stride) * params.x_inc));
|
|
1206
|
+
acc2 = addFields(acc2, absFields((id + 2u * stride) * params.x_inc));
|
|
1207
|
+
acc3 = addFields(acc3, absFields((id + 3u * stride) * params.x_inc));
|
|
1208
|
+
}
|
|
1209
|
+
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
1210
|
+
acc0 = addFields(acc0, absFields(id * params.x_inc));
|
|
1211
|
+
}
|
|
1212
|
+
|
|
1213
|
+
// Combine the 4 per-thread accumulators in Fields form too \u2014 still no
|
|
1214
|
+
// encode/decode needed, since none of them have touched Packed yet.
|
|
1215
|
+
let combined = addFields(addFields(acc0, acc1), addFields(acc2, acc3));
|
|
1216
|
+
tile[lid.x] = encode(combined.sign, combined.rawExp, combined.mantissaHi, combined.lo);
|
|
1217
|
+
workgroupBarrier();
|
|
1218
|
+
|
|
1219
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
1220
|
+
if (lid.x < s) { tile[lid.x] = addPair(tile[lid.x], tile[lid.x + s]); }
|
|
1221
|
+
workgroupBarrier();
|
|
1222
|
+
}
|
|
1223
|
+
|
|
1224
|
+
if (lid.x == 0u) {
|
|
1225
|
+
partialsMain[wgid.x] = tile[0].main;
|
|
1226
|
+
partialsAux[wgid.x] = tile[0].aux;
|
|
1227
|
+
}
|
|
1228
|
+
}
|
|
1229
|
+
`});var be,we=W(()=>{be=`// strsv_invert_block: computes ONE column (workgroup_id.x) of ONE block's
|
|
1230
|
+
// (workgroup_id.y) explicit inverse, via the same one-row-at-a-time
|
|
1231
|
+
// substitution as strsv_block.wgsl, but solving against a unit basis vector
|
|
1232
|
+
// e_col instead of the real right-hand side, and writing to a dense
|
|
1233
|
+
// (BLOCK_SIZE x BLOCK_SIZE, row-major) scratch buffer per block instead of
|
|
1234
|
+
// mutating x. Dispatched once for the whole matrix (2D: BLOCK_SIZE columns x
|
|
1235
|
+
// numBlocks), fully in parallel -- unlike the sequential per-block main
|
|
1236
|
+
// loop in strsv.mjs, no block's inverse depends on any other block or on x.
|
|
1237
|
+
//
|
|
1238
|
+
// A triangular block's inverse is itself triangular: forward (effectively-
|
|
1239
|
+
// lower, e.g. no-trans+lower) blocks have inverse column col nonzero only
|
|
1240
|
+
// for row>=col, solved in increasing row order; backward (effectively-
|
|
1241
|
+
// upper) blocks have it nonzero only for row<=col, solved in decreasing
|
|
1242
|
+
// order. Rows outside a column's nonzero range are written as literal 0 \u2014
|
|
1243
|
+
// strsv_apply_inverse.wgsl's dense matvec depends on that, not just on
|
|
1244
|
+
// those entries being mathematically implied zero.
|
|
1245
|
+
|
|
1246
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1247
|
+
@group(0) @binding(1) var<storage, read_write> Ainv: array<f32>;
|
|
1248
|
+
|
|
1249
|
+
struct Params {
|
|
1250
|
+
n: u32,
|
|
1251
|
+
lda: u32,
|
|
1252
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
1253
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
1254
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
1255
|
+
}
|
|
1256
|
+
|
|
1257
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
1258
|
+
|
|
1259
|
+
const WGS: u32 = 64u;
|
|
1260
|
+
const BLOCK_SIZE: u32 = 64u;
|
|
1261
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
1262
|
+
|
|
1263
|
+
fn readA(i: u32, j: u32) -> f32 {
|
|
1264
|
+
if params.trans == 0u {
|
|
1265
|
+
return A[i * params.lda + j];
|
|
1266
|
+
} else {
|
|
1267
|
+
return A[j * params.lda + i];
|
|
1268
|
+
}
|
|
1269
|
+
}
|
|
1270
|
+
|
|
1271
|
+
@compute @workgroup_size(64)
|
|
1272
|
+
fn strsv_invert_block_main(
|
|
1273
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1274
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1275
|
+
) {
|
|
1276
|
+
let col = wgid.x;
|
|
1277
|
+
let blockIndex = wgid.y;
|
|
1278
|
+
let blockStart = blockIndex * BLOCK_SIZE;
|
|
1279
|
+
var blockEnd = blockStart + BLOCK_SIZE;
|
|
1280
|
+
if (blockEnd > params.n) { blockEnd = params.n; }
|
|
1281
|
+
let blockLen = blockEnd - blockStart;
|
|
1282
|
+
|
|
1283
|
+
if (col >= blockLen) { return; }
|
|
1284
|
+
|
|
1285
|
+
let ainvBase = blockIndex * BLOCK_SIZE * BLOCK_SIZE;
|
|
1286
|
+
let forward = (params.trans == 0u) == (params.uplo == 0u);
|
|
1287
|
+
|
|
1288
|
+
if forward {
|
|
1289
|
+
for (var r = lid.x; r < col; r += WGS) {
|
|
1290
|
+
Ainv[ainvBase + r * BLOCK_SIZE + col] = 0.0;
|
|
1291
|
+
}
|
|
1292
|
+
} else {
|
|
1293
|
+
for (var r = col + 1u + lid.x; r < blockLen; r += WGS) {
|
|
1294
|
+
Ainv[ainvBase + r * BLOCK_SIZE + col] = 0.0;
|
|
1295
|
+
}
|
|
1296
|
+
}
|
|
1297
|
+
storageBarrier();
|
|
1298
|
+
workgroupBarrier();
|
|
1299
|
+
|
|
1300
|
+
let numSteps = select(col + 1u, blockLen - col, forward);
|
|
1301
|
+
for (var step = 0u; step < numSteps; step++) {
|
|
1302
|
+
let localRow = select(col - step, col + step, forward);
|
|
1303
|
+
let i = blockStart + localRow;
|
|
1304
|
+
|
|
1305
|
+
var acc = 0.0f;
|
|
1306
|
+
if forward {
|
|
1307
|
+
for (var lj = col + lid.x; lj < localRow; lj += WGS) {
|
|
1308
|
+
acc += readA(i, blockStart + lj) * Ainv[ainvBase + lj * BLOCK_SIZE + col];
|
|
1309
|
+
}
|
|
1310
|
+
} else {
|
|
1311
|
+
for (var lj = localRow + 1u + lid.x; lj <= col; lj += WGS) {
|
|
1312
|
+
acc += readA(i, blockStart + lj) * Ainv[ainvBase + lj * BLOCK_SIZE + col];
|
|
1313
|
+
}
|
|
1314
|
+
}
|
|
1315
|
+
|
|
1316
|
+
scratch[lid.x] = acc;
|
|
1317
|
+
workgroupBarrier();
|
|
1318
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
1319
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
1320
|
+
workgroupBarrier();
|
|
1321
|
+
}
|
|
1322
|
+
|
|
1323
|
+
if lid.x == 0u {
|
|
1324
|
+
let e = select(0.0, 1.0, localRow == col);
|
|
1325
|
+
let rhs = e - scratch[0];
|
|
1326
|
+
var val: f32;
|
|
1327
|
+
if params.diag == 1u {
|
|
1328
|
+
val = rhs;
|
|
1329
|
+
} else {
|
|
1330
|
+
val = rhs / A[i * params.lda + i];
|
|
1331
|
+
}
|
|
1332
|
+
Ainv[ainvBase + localRow * BLOCK_SIZE + col] = val;
|
|
1333
|
+
}
|
|
1334
|
+
storageBarrier();
|
|
1335
|
+
workgroupBarrier();
|
|
1336
|
+
}
|
|
1337
|
+
}
|
|
1338
|
+
`});var xe,he=W(()=>{xe=`// strsv_apply_inverse: given a precomputed block inverse (from
|
|
1339
|
+
// strsv_invert_block.wgsl), computes this block's solution as a dense
|
|
1340
|
+
// matrix-vector multiply against the block's current remainder in x \u2014
|
|
1341
|
+
// replacing what the old strsv_block.wgsl did via a genuinely sequential,
|
|
1342
|
+
// barrier-per-row substitution.
|
|
1343
|
+
//
|
|
1344
|
+
// All blockLen rows are computed in parallel within a single workgroup: the
|
|
1345
|
+
// remainder is loaded into workgroup-shared memory once, then each thread
|
|
1346
|
+
// independently computes one full row's dot product from that shared copy.
|
|
1347
|
+
// No further synchronization is needed after the load \u2014 every thread only
|
|
1348
|
+
// reads shared memory from then on (never written again within this call)
|
|
1349
|
+
// and writes a distinct element of x, so there's no cross-thread hazard to
|
|
1350
|
+
// guard against.
|
|
1351
|
+
|
|
1352
|
+
@group(0) @binding(0) var<storage, read> Ainv: array<f32>;
|
|
1353
|
+
@group(0) @binding(1) var<storage, read_write> x: array<f32>;
|
|
1354
|
+
|
|
1355
|
+
struct Params {
|
|
1356
|
+
incx: u32,
|
|
1357
|
+
blockIndex: u32,
|
|
1358
|
+
blockStart: u32,
|
|
1359
|
+
blockEnd: u32,
|
|
1360
|
+
}
|
|
1361
|
+
|
|
1362
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
1363
|
+
|
|
1364
|
+
const BLOCK_SIZE: u32 = 64u;
|
|
1365
|
+
var<workgroup> xLocal: array<f32, 64>;
|
|
1366
|
+
|
|
1367
|
+
@compute @workgroup_size(64)
|
|
1368
|
+
fn strsv_apply_inverse_main(@builtin(local_invocation_id) lid: vec3u) {
|
|
1369
|
+
let blockLen = params.blockEnd - params.blockStart;
|
|
1370
|
+
|
|
1371
|
+
if (lid.x < blockLen) {
|
|
1372
|
+
xLocal[lid.x] = x[(params.blockStart + lid.x) * params.incx];
|
|
1373
|
+
}
|
|
1374
|
+
workgroupBarrier();
|
|
1375
|
+
|
|
1376
|
+
if (lid.x >= blockLen) { return; }
|
|
1377
|
+
|
|
1378
|
+
let ainvBase = params.blockIndex * BLOCK_SIZE * BLOCK_SIZE;
|
|
1379
|
+
var acc = 0.0f;
|
|
1380
|
+
for (var j = 0u; j < blockLen; j++) {
|
|
1381
|
+
acc += Ainv[ainvBase + lid.x * BLOCK_SIZE + j] * xLocal[j];
|
|
1382
|
+
}
|
|
1383
|
+
x[(params.blockStart + lid.x) * params.incx] = acc;
|
|
1384
|
+
}
|
|
1385
|
+
`});var ye,ve=W(()=>{ye=`// strsv_update: subtracts a solved block's contribution from every
|
|
1386
|
+
// remaining row in parallel (one workgroup per row, like strmv.wgsl) \u2014
|
|
1387
|
+
// this is what turns strsv's O(n) sequential stages into O(n/blockSize).
|
|
1388
|
+
// No diag/masking needed: this region never touches the diagonal.
|
|
1389
|
+
|
|
1390
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
1391
|
+
@group(0) @binding(1) var<storage, read_write> x: array<f32>;
|
|
1392
|
+
|
|
1393
|
+
struct Params {
|
|
1394
|
+
n: u32,
|
|
1395
|
+
incx: u32,
|
|
1396
|
+
lda: u32,
|
|
1397
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
1398
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
1399
|
+
blockStart: u32,
|
|
1400
|
+
blockEnd: u32, // exclusive
|
|
1401
|
+
}
|
|
1402
|
+
|
|
1403
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
1404
|
+
|
|
1405
|
+
const WGS: u32 = 64u;
|
|
1406
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
1407
|
+
|
|
1408
|
+
@compute @workgroup_size(64)
|
|
1409
|
+
fn strsv_update_main(
|
|
1410
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
1411
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
1412
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
1413
|
+
) {
|
|
1414
|
+
// forward: remaining rows are [blockEnd,n); backward: [0,blockStart).
|
|
1415
|
+
let forward = (params.trans == 0u) == (params.uplo == 0u);
|
|
1416
|
+
|
|
1417
|
+
var rangeStart: u32;
|
|
1418
|
+
var rangeEnd: u32;
|
|
1419
|
+
|
|
1420
|
+
if forward {
|
|
1421
|
+
rangeStart = params.blockEnd;
|
|
1422
|
+
rangeEnd = params.n;
|
|
1423
|
+
} else {
|
|
1424
|
+
rangeStart = 0u;
|
|
1425
|
+
rangeEnd = params.blockStart;
|
|
1426
|
+
}
|
|
1427
|
+
|
|
1428
|
+
if (rangeStart >= rangeEnd) { return; }
|
|
1429
|
+
let count = rangeEnd - rangeStart;
|
|
1430
|
+
|
|
1431
|
+
for (var idx = wgid.x; idx < count; idx += nwg.x) {
|
|
1432
|
+
let i = rangeStart + idx;
|
|
1433
|
+
|
|
1434
|
+
// No-trans reads A[i,j]; transpose reads A[j,i] \u2014 uplo only sets the range above.
|
|
1435
|
+
var acc = 0.0f;
|
|
1436
|
+
if params.trans == 0u {
|
|
1437
|
+
for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
|
|
1438
|
+
acc += A[i * params.lda + j] * x[j * params.incx];
|
|
1439
|
+
}
|
|
1440
|
+
} else {
|
|
1441
|
+
for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
|
|
1442
|
+
acc += A[j * params.lda + i] * x[j * params.incx];
|
|
1443
|
+
}
|
|
1444
|
+
}
|
|
1445
|
+
|
|
1446
|
+
// Parallel reduction: 64 \u2192 1
|
|
1447
|
+
scratch[lid.x] = acc;
|
|
1448
|
+
workgroupBarrier();
|
|
1449
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
1450
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
1451
|
+
workgroupBarrier();
|
|
1452
|
+
}
|
|
1453
|
+
|
|
1454
|
+
if lid.x == 0u {
|
|
1455
|
+
x[i * params.incx] -= scratch[0];
|
|
1456
|
+
}
|
|
1457
|
+
workgroupBarrier();
|
|
1458
|
+
}
|
|
1459
|
+
}
|
|
1460
|
+
`});var _e={};hr(_e,{shaderSources:()=>Nt});var Nt,Ee=W(()=>{Ir();Nr();Mr();Tr();Vr();Dr();Or();Qr();Zr();Xr();Yr();re();te();ie();ne();ue();ce();de();pe();we();he();ve();Nt={"reduction/argmax":jr,"reduction/sum":Wr,"reduction/sumF64":Ur,sscal:Rr,sswap:Hr,saxpy:Cr,scopy:zr,sdot:qr,sasum:Kr,snrm2:$r,srot:Jr,srotm:ee,isamax:ae,sgemv_n:oe,sgemv_t:se,ssymv:le,strmv:fe,f64add:me,dasum:ge,strsv_invert_block:be,strsv_apply_inverse:xe,strsv_update:ye}});var Rt={};hr(Rt,{GpuMatrix:()=>C,GpuVector:()=>h,cleanup:()=>Br,dasum:()=>je,gpuName:()=>Ar,init:()=>kr,isamax:()=>Me,randomFloat32Array:()=>Fr,randomFloat64Array:()=>Lr,sasum:()=>Ie,saxpy:()=>Ge,scopy:()=>Pe,sdot:()=>Fe,sgemv:()=>Re,snrm2:()=>We,srot:()=>Ue,srotm:()=>Te,sscal:()=>Be,sswap:()=>Ae,ssymv:()=>Ve,strmv:()=>He,strsv:()=>Oe});function vr(a,r){return r?a.features.has("timestamp-query")?{requiredFeatures:["timestamp-query"]}:(console.warn("timestamp-query not supported on this device \u2014 benchmark mode disabled."),{}):{}}function yr(){if(!_r())return{querySet:null,passDescriptor:void 0};let r=M().createQuerySet({type:"timestamp",count:2});return{querySet:r,passDescriptor:{timestampWrites:{querySet:r,beginningOfPassWriteIndex:0,endOfPassWriteIndex:1}}}}function ur(a,r){if(!r)return null;let e=M(),i=e.createBuffer({label:"timestamp-resolve",size:16,usage:GPUBufferUsage.QUERY_RESOLVE|GPUBufferUsage.COPY_SRC});a.resolveQuerySet(r,0,2,i,0);let t=e.createBuffer({label:"timestamp-readback",size:16,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ});return a.copyBufferToBuffer(i,0,t,0,16),{tsReadBuffer:t,resolveBuffer:i,querySet:r}}async function G(a){if(!a)return;let{tsReadBuffer:r,resolveBuffer:e,querySet:i}=a;await r.mapAsync(GPUMapMode.READ);let t=new BigInt64Array(r.getMappedRange().slice());return r.unmap(),r.destroy(),e.destroy(),i.destroy(),Math.max(0,Number(t[1]-t[0]))/1e6}var Q=null,$=null,Er=null,cr=!1;async function kr({powerPreference:a="high-performance",benchmark:r=!1}={}){if(Q)return Q;let e;if(typeof window>"u"){let{create:o,globals:u}=await import("webgpu");Object.assign(globalThis,u),e=o([]),Er=e}else e=navigator.gpu;if(!e)throw new Error("WebGPU not supported in this environment.");if($=await e.requestAdapter({powerPreference:a})??await e.requestAdapter(),!$)throw new Error("No WebGPU adapter found.");cr=r;let t=[...vr($,r).requiredFeatures??[]];return Q=await $.requestDevice({requiredFeatures:t}),Q.addEventListener("uncapturederror",o=>{console.error("Uncaptured GPU error:",o.error.message)}),Q}function Br(){Q&&(Q.destroy(),Q=null),$=null,Er=null,cr=!1}function Ar(){if(!$)throw new Error("WebGPU adapter not initialized \u2014 call init() first.");let{device:a,description:r}=$.info;return{description:r||"unknown",device:a||"unknown"}}function _r(){return cr}function M(){if(!Q)throw new Error("WebGPU device not initialized \u2014 call init() first.");return Q}function m(...a){a.flat().forEach(r=>r.destroy())}function b(a,r="blas-input",e=!1){let i=M(),t=i.limits.maxStorageBufferBindingSize,o=a.byteLength;if(o>t)throw new Error(`Buffer size ${o} bytes exceeds device limit of ${t} bytes.`);let u=e?GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC:GPUBufferUsage.STORAGE,n=i.createBuffer({label:r,size:o,usage:u,mappedAtCreation:!0}),l=a.constructor;return new l(n.getMappedRange()).set(a),n.unmap(),n}function H(a,r="blas-storage"){return M().createBuffer({label:r,size:a,usage:GPUBufferUsage.STORAGE})}function O(a,r="blas-result"){return M().createBuffer({label:r,size:a,usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC})}function _(a,r){let i=M().createBuffer({label:"blas-readback",size:r.size,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ});return a.copyBufferToBuffer(r,0,i,0,r.size),i}function L(a,r="blas-params"){let e=M(),i=a.length*4,t=Math.ceil(i/16)*16,o=new ArrayBuffer(t),u=new DataView(o);a.forEach(({value:l,type:s},f)=>{let c=f*4;if(s==="u32")u.setUint32(c,l,!0);else if(s==="i32")u.setInt32(c,l,!0);else if(s==="f32")u.setFloat32(c,l,!0);else throw new Error(`Unknown param type "${s}". Use "f32", "u32", or "i32".`)});let n=e.createBuffer({label:r,size:t,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST});return e.queue.writeBuffer(n,0,o),n}async function y(a,r=Float32Array){try{await a.mapAsync(GPUMapMode.READ);let e=new r(a.getMappedRange().slice());return a.unmap(),e}finally{a.destroy()}}var ot=new ArrayBuffer(8),J=new DataView(ot),Gr=new ArrayBuffer(4),Pr=new Uint32Array(Gr),Sr=new Float32Array(Gr);function nt(a){return Pr[0]=a>>>0,Sr[0]}function st(a){return Sr[0]=a,Pr[0]}function ut(a,r,e,i){let t=r>>>3,o=r&7,u=i>>>29,n=e<<3|u,l=i&536870911,s=a<<31|t<<23|n,f=o>>>2&1,c=o&3,d=l>>>23,p=l&8388607,g=c<<6|d,w=(f<<31|g<<23|p)>>>0;return[nt(s),w]}function lt(a,r){let e=st(a);r=r>>>0;let i=e>>>31,t=e>>>23&255,o=e&8388607,u=r>>>31,n=r>>>23&255,l=r&8388607,s=u<<2|n>>>6,f=(n&63)<<23|l,c=t<<3|s,d=o>>>3,g=((o&7)<<29|f)>>>0;return{sign:i,rawExp:c,mantissaHi:d,lo:g}}var ct=2040;function lr(a){J.setFloat64(0,a,!1);let r=J.getUint32(0,!1),e=J.getUint32(4,!1),i=r>>>31,t=r>>>20&2047,o=r&1048575;if(t>=ct)throw new RangeError(`packF64: |${a}| is too large to pack safely (must be finite with magnitude below ~1.4e306); main's bit pattern would itself be NaN/Infinity-shaped and get silently corrupted by any real float32 round-trip`);return ut(i,t,o,e)}function rr(a,r){let{sign:e,rawExp:i,mantissaHi:t,lo:o}=lt(a,r),u=(e<<31|i<<20|t)>>>0;return J.setUint32(0,u,!1),J.setUint32(4,o,!1),J.getFloat64(0,!1)}var h=class a{constructor(r,e,i=Float32Array,t=null){this._buf=r,this._auxBuf=t,this.length=e,this.dtype=i}static from(r){if(r instanceof Float64Array){let i=new Float32Array(r.length),t=new Uint32Array(r.length);for(let n=0;n<r.length;n++){let l=lr(r[n]);i[n]=l[0],t[n]=l[1]}let o=b(i,"gpu-vector-f64-main",!0),u=b(t,"gpu-vector-f64-aux",!0);return new a(o,r.length,Float64Array,u)}if(!(r instanceof Float32Array))throw new Error("GpuVector.from expects a Float32Array or Float64Array.");let e=b(r,"gpu-vector",!0);return new a(e,r.length,r.constructor)}async read(){let r=M(),e=r.createCommandEncoder(),i=_(e,this._buf);if(r.queue.submit([e.finish()]),!this._auxBuf)return y(i,this.dtype);let t=r.createCommandEncoder(),o=_(t,this._auxBuf);r.queue.submit([t.finish()]);let[u,n]=await Promise.all([y(i,Float32Array),y(o,Uint32Array)]),l=new Float64Array(this.length);for(let s=0;s<this.length;s++)l[s]=rr(u[s],n[s]);return l}destroy(){this._buf.destroy(),this._auxBuf&&this._auxBuf.destroy()}};var C=class a{constructor(r,e,i,t,o=null){this._buf=r,this._auxBuf=o,this.rows=e,this.cols=i,this.lda=t}static from(r,e,i,t=i){if(!(r instanceof Float32Array)&&!(r instanceof Float64Array))throw new Error("GpuMatrix.from expects a Float32Array or Float64Array.");if(!Number.isInteger(e)||e<=0)throw new Error("rows must be a positive integer.");if(!Number.isInteger(i)||i<=0)throw new Error("cols must be a positive integer.");if(!Number.isInteger(t)||t<i)throw new Error("lda must be an integer >= cols.");if(r.length<e*t)throw new Error("data does not have enough elements for the given rows and lda.");if(r instanceof Float64Array){let u=e*t,n=new Float32Array(u),l=new Uint32Array(u);for(let c=0;c<u;c++){let d=lr(r[c]);n[c]=d[0],l[c]=d[1]}let s=b(n,"gpu-matrix-f64-main",!0),f=b(l,"gpu-matrix-f64-aux",!0);return new a(s,e,i,t,f)}let o=b(r.subarray(0,e*t),"gpu-matrix",!0);return new a(o,e,i,t)}async read(){let r=M(),e=r.createCommandEncoder(),i=_(e,this._buf);if(r.queue.submit([e.finish()]),this._auxBuf){let u=r.createCommandEncoder(),n=_(u,this._auxBuf);r.queue.submit([u.finish()]);let[l,s]=await Promise.all([y(i,Float32Array),y(n,Uint32Array)]),f=new Float64Array(this.rows*this.lda);for(let d=0;d<f.length;d++)f[d]=rr(l[d],s[d]);if(this.lda===this.cols)return f;let c=new Float64Array(this.rows*this.cols);for(let d=0;d<this.rows;d++)c.set(f.subarray(d*this.lda,d*this.lda+this.cols),d*this.cols);return c}let t=await y(i,Float32Array);if(this.lda===this.cols)return t;let o=new Float32Array(this.rows*this.cols);for(let u=0;u<this.rows;u++)o.set(t.subarray(u*this.lda,u*this.lda+this.cols),u*this.cols);return o}destroy(){this._buf.destroy(),this._auxBuf&&this._auxBuf.destroy()}};function Fr(a,r=-1,e=1){let i=new Float32Array(a);for(let t=0;t<a;t++)i[t]=r+Math.random()*(e-r);return i}function Lr(a,r=-1,e=1){let i=new Float64Array(a);for(let t=0;t<a;t++)i[t]=r+Math.random()*(e-r);return i}function B(a,r,e=0){let i=M(),t=r.map((o,u)=>({binding:e+u,resource:o instanceof GPUBuffer?{buffer:o}:o}));return i.createBindGroup({layout:a,entries:t})}var ft=new WeakMap;function P(a){M().queue.submit([a.finish()])}function fr(){let a=M(),{querySet:r,passDescriptor:e}=yr();return{commandEncoder:a.createCommandEncoder(),querySet:r,passDescriptor:e}}function ar(a,r,e,i,t){let o=a.beginComputePass(t);o.setPipeline(r),o.setBindGroup(0,e),typeof i=="number"?o.dispatchWorkgroups(i):o.dispatchWorkgroups(i.x,i.y),o.end(),ft.set(a,o)}function S(a,r,e){let{commandEncoder:i,querySet:t,passDescriptor:o}=fr();ar(i,a,r,e,o);let u=ur(i,t);return{commandEncoder:i,ts:u}}var Ut={},dr=new WeakMap;async function A(a,r,e="main"){dr.has(a)||dr.set(a,new Map);let i=dr.get(a),t=Array.isArray(r)?r:[r],o=`${t.join("+")}::${e}`;return i.has(o)||i.set(o,await Mt(t,e)),i.get(o)}async function Wt(a){if(typeof process>"u"||!process.versions?.node){let{shaderSources:r}=await Promise.resolve().then(()=>(Ee(),_e)),e=r[a];if(!e)throw new Error(`Shader "${a}" not found in browser bundle.`);return e}else{let{readFileSync:r}=await import("fs"),{fileURLToPath:e}=await import("url"),{dirname:i,join:t}=await import("path"),o=i(e(Ut.url));return r(t(o,`../shaders/${a}.wgsl`),"utf8")}}async function Mt(a,r="main"){let e=M(),i=a.join("+"),t=(await Promise.all(a.map(Wt))).join(`
|
|
1461
|
+
`),o=e.createShaderModule({label:i,code:t}),n=(await o.getCompilationInfo()).messages.filter(f=>f.type==="error");if(n.length>0)throw new Error(`Shader "${i}" compilation failed:
|
|
1462
|
+
${n.map(f=>` line ${f.lineNum}: ${f.message}`).join(`
|
|
1463
|
+
`)}`);let l=r==="main"?{module:o}:{module:o,entryPoint:r},s=e.createComputePipeline({label:i,layout:"auto",compute:l});return s._shaderModule=o,s}var Tt=64,ke=8;function D(a,r){let e=M().limits.maxComputeWorkgroupsPerDimension;return r===void 0?Math.min(Math.ceil(a/Tt),e):{x:Math.min(Math.ceil(r/ke),e),y:Math.min(Math.ceil(a/ke),e)}}async function Be(a,r,e,i,t){let o=i instanceof h;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(t))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(t<=0)throw new Error("incx must be positive.");if(!(i instanceof Float32Array)&&!(i instanceof h))throw new Error("x must be a Float32Array or GpuVector.");if(r<=0)return o?{}:i;if(i.length<(r-1)*t+1)throw new Error("x does not have enough elements for the given n and incx.");let u=await A(a,"sscal"),n=null,l=null,s=null;try{n=o?i._buf:b(i,"sscal-x",!0),l=L([{value:r,type:"u32"},{value:e,type:"f32"},{value:t,type:"u32"}],"sscal-params");let f=B(u.getBindGroupLayout(0),[n,l]),{commandEncoder:c,ts:d}=S(u,f,D(r));s=o?null:_(c,n),P(c);let p=await G(d);if(o)return p!==void 0?{gpuTimeMs:p}:{};let g=await y(s,Float32Array);return s=null,p!==void 0?{x:g,gpuTimeMs:p}:g}finally{!o&&n&&m(n),l&&m(l),s&&m(s)}}async function Ae(a,r,e,i,t,o){let u=e instanceof h,n=t instanceof h;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(i)||!Number.isInteger(o))throw new Error("n, incx, and incy must be integers.");if(i<=0||o<=0)throw new Error("incx and incy must be positive.");if(!(e instanceof Float32Array)&&!(e instanceof h))throw new Error("x must be a Float32Array or GpuVector.");if(!(t instanceof Float32Array)&&!(t instanceof h))throw new Error("y must be a Float32Array or GpuVector.");if(e.constructor!==t.constructor)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(r<=0)return u?{}:{x:e,y:t};if(e.length<(r-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");if(t.length<(r-1)*o+1)throw new Error("y does not have enough elements for the given n and incy.");let l=await A(a,"sswap"),s=null,f=null,c=null,d=null,p=null;try{s=u?e._buf:b(e,"sswap-x",!0),f=n?t._buf:b(t,"sswap-y",!0),c=L([{value:r,type:"u32"},{value:i,type:"u32"},{value:o,type:"u32"}],"sswap-params");let g=B(l.getBindGroupLayout(0),[s,f,c]),{commandEncoder:w,ts:k}=S(l,g,D(r));d=u?null:_(w,s),p=n?null:_(w,f),P(w);let x=await G(k);if(u&&n)return x!==void 0?{gpuTimeMs:x}:{};let E=await y(d,Float32Array);d=null;let v=await y(p,Float32Array);return p=null,x!==void 0?{x:E,y:v,gpuTimeMs:x}:{x:E,y:v}}finally{!u&&s&&m(s),!n&&f&&m(f),c&&m(c),d&&m(d),p&&m(p)}}async function Ge(a,r,e,i,t,o,u){let n=i instanceof h,l=o instanceof h;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(t)||!Number.isInteger(u))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(t<=0||u<=0)throw new Error("incx and incy must be positive.");if(!n&&!(i 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(r<=0)return l?{}:{y:o};if(i.length<(r-1)*t+1)throw new Error("x does not have enough elements for the given n and incx.");if(o.length<(r-1)*u+1)throw new Error("y does not have enough elements for the given n and incy.");let s=await A(a,"saxpy"),f=null,c=null,d=null,p=null;try{f=n?i._buf:b(i,"saxpy-x",!1),c=l?o._buf:b(o,"saxpy-y",!0),d=L([{value:r,type:"u32"},{value:e,type:"f32"},{value:t,type:"u32"},{value:u,type:"u32"}],"saxpy-params");let g=B(s.getBindGroupLayout(0),[f,c,d]),{commandEncoder:w,ts:k}=S(s,g,D(r));p=l?null:_(w,c),P(w);let x=await G(k);if(l&&n)return x!==void 0?{gpuTimeMs:x}:{};let E=await y(p,Float32Array);return p=null,x!==void 0?{y:E,gpuTimeMs:x}:{y:E}}finally{!n&&f&&m(f),!l&&c&&m(c),d&&m(d),p&&m(p)}}async function Pe(a,r,e,i,t,o){let u=e instanceof h,n=t instanceof h;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(i)||!Number.isInteger(o))throw new Error("n, incx, and incy must be integers.");if(i<=0||o<=0)throw new Error("incx and incy must be positive.");if(!u&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!n&&!(t instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(u!==n)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(r<=0)return n?{}:{y:t};if(e.length<(r-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");if(t.length<(r-1)*o+1)throw new Error("y does not have enough elements for the given n and incy.");let l=await A(a,"scopy"),s=null,f=null,c=null,d=null;try{s=u?e._buf:b(e,"scopy-x",!1),f=n?t._buf:b(t,"scopy-y",!0),c=L([{value:r,type:"u32"},{value:i,type:"u32"},{value:o,type:"u32"}],"scopy-params");let p=B(l.getBindGroupLayout(0),[s,f,c]),{commandEncoder:g,ts:w}=S(l,p,D(r));d=n?null:_(g,f),P(g);let k=await G(w);if(n&&u)return k!==void 0?{gpuTimeMs:k}:{};let x=await y(d,Float32Array);return d=null,k!==void 0?{y:x,gpuTimeMs:k}:{y:x}}finally{!u&&s&&m(s),!n&&f&&m(f),c&&m(c),d&&m(d)}}var Se=64;async function Fe(a,r,e,i,t,o){let u=e instanceof h,n=t instanceof h;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(i)||!Number.isInteger(o))throw new Error("n, incx, and incy must be integers.");if(i<=0||o<=0)throw new Error("incx and incy must be positive.");if(!u&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!n&&!(t instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(u!==n)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(r<=0)return{dot:0};if(e.length<(r-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");if(t.length<(r-1)*o+1)throw new Error("y does not have enough elements for the given n and incy.");let l=await A(a,"sdot"),s=await A(a,"reduction/sum"),f=null,c=null,d=null,p=null,g=null,w=null;try{f=u?e._buf:b(e,"sdot-x",!1),c=n?t._buf:b(t,"sdot-y",!1),d=H(2*Se*4,"sdot-partials"),p=O(4,"sdot-result"),g=L([{value:r,type:"u32"},{value:i,type:"u32"},{value:o,type:"u32"}],"sdot-params");let k=B(l.getBindGroupLayout(0),[f,c,d,g]),{commandEncoder:x,ts:E}=S(l,k,2*Se);P(x);let v=B(s.getBindGroupLayout(0),[d,p]),{commandEncoder:F,ts:I}=S(s,v,1);w=_(F,p),P(F);let j=y(w,Float32Array);w=null;let[N,U,T]=await Promise.all([G(E),G(I),j]);return N!==void 0&&U!==void 0?{dot:T[0],gpuTimeMs:N+U}:{dot:T[0]}}finally{!u&&f&&m(f),!n&&c&&m(c),d&&m(d),p&&m(p),g&&m(g),w&&m(w)}}var Le=64;async function Ie(a,r,e,i){let t=e instanceof h;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(i))throw new Error("n and incx must be integers.");if(i<=0)throw new Error("incx must be positive.");if(!t&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(r<=0)return{asum:0};if(e.length<(r-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");let o=await A(a,"sasum"),u=await A(a,"reduction/sum"),n=null,l=null,s=null,f=null,c=null;try{n=t?e._buf:b(e,"sasum-x",!1),l=H(2*Le*4,"sasum-partials"),s=O(4,"sasum-result"),f=L([{value:r,type:"u32"},{value:i,type:"u32"}],"sasum-params");let d=B(o.getBindGroupLayout(0),[n,l,f]),{commandEncoder:p,ts:g}=S(o,d,2*Le);P(p);let w=B(u.getBindGroupLayout(0),[l,s]),{commandEncoder:k,ts:x}=S(u,w,1);c=_(k,s),P(k);let E=y(c,Float32Array);c=null;let[v,F,I]=await Promise.all([G(g),G(x),E]);return v!==void 0&&F!==void 0?{asum:I[0],gpuTimeMs:v+F}:{asum:I[0]}}finally{!t&&n&&m(n),l&&m(l),s&&m(s),f&&m(f),c&&m(c)}}var mr=64;async function je(a,r,e,i){let t=e instanceof h;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(i))throw new Error("n and incx must be integers.");if(i<=0)throw new Error("incx must be positive.");if(!t&&!(e instanceof Float64Array))throw new Error("x must be a Float64Array or GpuVector.");if(t&&e.dtype!==Float64Array)throw new Error("x must be a Float64Array-backed GpuVector.");if(r<=0)return{asum:0};if(e.length<(r-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");let o=await A(a,["f64add","dasum"]),u=await A(a,["f64add","reduction/sumF64"]),n=null,l=null,s=null,f=null,c=null,d=null,p=null,g=null;try{n=t?e:h.from(e),l=H(2*mr*4,"dasum-partialsMain"),s=H(2*mr*4,"dasum-partialsAux"),f=O(4,"dasum-result-main"),c=O(4,"dasum-result-aux"),d=L([{value:r,type:"u32"},{value:i,type:"u32"}],"dasum-params");let w=B(o.getBindGroupLayout(0),[n._buf,n._auxBuf,l,s,d]),{commandEncoder:k,ts:x}=S(o,w,2*mr);P(k);let E=B(u.getBindGroupLayout(0),[l,s,f,c]),{commandEncoder:v,ts:F}=S(u,E,1);p=_(v,f),g=_(v,c),P(v);let I=y(p,Float32Array),j=y(g,Uint32Array);p=null,g=null;let[N,U,T,R]=await Promise.all([G(x),G(F),I,j]),V=rr(T[0],R[0]);return N!==void 0&&U!==void 0?{asum:V,gpuTimeMs:N+U}:{asum:V}}finally{!t&&n&&n.destroy(),l&&m(l),s&&m(s),f&&m(f),c&&m(c),d&&m(d),p&&m(p),g&&m(g)}}var Ne=64;async function We(a,r,e,i){let t=e instanceof h;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(i))throw new Error("n and incx must be integers.");if(i<=0)throw new Error("incx must be positive.");if(!t&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(r<=0)return{nrm2:0};if(e.length<(r-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");let o=await A(a,"snrm2"),u=await A(a,"reduction/sum"),n=null,l=null,s=null,f=null,c=null;try{n=t?e._buf:b(e,"snrm2-x",!1),l=H(2*Ne*4,"snrm2-partials"),s=O(4,"snrm2-result"),f=L([{value:r,type:"u32"},{value:i,type:"u32"}],"snrm2-params");let d=B(o.getBindGroupLayout(0),[n,l,f]),{commandEncoder:p,ts:g}=S(o,d,2*Ne);P(p);let w=B(u.getBindGroupLayout(0),[l,s]),{commandEncoder:k,ts:x}=S(u,w,1);c=_(k,s),P(k);let E=y(c,Float32Array);c=null;let[v,F,I]=await Promise.all([G(g),G(x),E]),j=Math.sqrt(I[0]);return v!==void 0&&F!==void 0?{nrm2:j,gpuTimeMs:v+F}:{nrm2:j}}finally{!t&&n&&m(n),l&&m(l),s&&m(s),f&&m(f),c&&m(c)}}var pr=64;async function Me(a,r,e,i){let t=e instanceof h;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(i))throw new Error("n and incx must be integers.");if(i<=0)throw new Error("incx must be positive.");if(!t&&!(e instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(r<=0)return{index:0};if(e.length<(r-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");let o=await A(a,"isamax"),u=await A(a,"reduction/argmax"),n=null,l=null,s=null,f=null,c=null,d=null;try{n=t?e._buf:b(e,"isamax-x",!1),l=H(2*pr*4,"isamax-partials-val"),s=H(2*pr*4,"isamax-partials-idx"),f=O(4,"isamax-result"),c=L([{value:r,type:"u32"},{value:i,type:"u32"}],"isamax-params");let p=B(o.getBindGroupLayout(0),[n,l,s,c]),{commandEncoder:g,ts:w}=S(o,p,2*pr);P(g);let k=B(u.getBindGroupLayout(0),[l,s,f]),{commandEncoder:x,ts:E}=S(u,k,1);d=_(x,f),P(x);let v=y(d,Uint32Array);d=null;let[F,I,j]=await Promise.all([G(w),G(E),v]),N=j[0];return F!==void 0&&I!==void 0?{index:N,gpuTimeMs:F+I}:{index:N}}finally{!t&&n&&m(n),l&&m(l),s&&m(s),f&&m(f),c&&m(c),d&&m(d)}}async function Ue(a,r,e,i,t,o,u,n){let l=e instanceof h,s=t instanceof h;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(i)||!Number.isInteger(o))throw new Error("n, incx, and incy must be integers.");if(typeof u!="number")throw new Error("c must be a number.");if(typeof n!="number")throw new Error("s must be a number.");if(Number.isNaN(u)||Number.isNaN(n))throw new Error("c and s must not be NaN.");if(!Number.isFinite(u))throw new Error("c must be finite.");if(!Number.isFinite(n))throw new Error("s must be finite.");if(i<=0||o<=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(!s&&!(t instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(l!==s)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(r<=0)return l?{}:{x:e,y:t};if(e.length<(r-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");if(t.length<(r-1)*o+1)throw new Error("y does not have enough elements for the given n and incy.");let f=await A(a,"srot"),c=null,d=null,p=null,g=null,w=null;try{c=l?e._buf:b(e,"srot-x",!0),d=s?t._buf:b(t,"srot-y",!0),p=L([{value:r,type:"u32"},{value:u,type:"f32"},{value:n,type:"f32"},{value:i,type:"u32"},{value:o,type:"u32"}],"srot-params");let k=B(f.getBindGroupLayout(0),[c,d,p]),{commandEncoder:x,ts:E}=S(f,k,D(r));g=l?null:_(x,c),w=s?null:_(x,d),P(x);let v=await G(E);if(l&&s)return v!==void 0?{gpuTimeMs:v}:{};let F=y(g,Float32Array),I=y(w,Float32Array);g=null,w=null;let[j,N]=await Promise.all([F,I]);return v!==void 0?{x:j,y:N,gpuTimeMs:v}:{x:j,y:N}}finally{!l&&c&&m(c),!s&&d&&m(d),p&&m(p),g&&m(g),w&&m(w)}}async function Te(a,r,e,i,t,o,u){let n=e instanceof h,l=t instanceof h;if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!Number.isInteger(r)||!Number.isInteger(i)||!Number.isInteger(o))throw new Error("n, incx, and incy must be integers.");if(!(u instanceof Float32Array)||u.length!==5)throw new Error("param must be a Float32Array of length 5.");if(i<=0||o<=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&&!(t 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(r<=0||u[0]===-2)return n?{}:{x:e,y:t};if(e.length<(r-1)*i+1)throw new Error("x does not have enough elements for the given n and incx.");if(t.length<(r-1)*o+1)throw new Error("y does not have enough elements for the given n and incy.");let s=await A(a,"srotm"),f=null,c=null,d=null,p=null,g=null,w=null;try{f=n?e._buf:b(e,"srotm-x",!0),c=l?t._buf:b(t,"srotm-y",!0),d=b(u,"srotm-param",!1),p=L([{value:r,type:"u32"},{value:i,type:"u32"},{value:o,type:"u32"}],"srotm-params");let k=B(s.getBindGroupLayout(0),[f,c,d,p]),{commandEncoder:x,ts:E}=S(s,k,D(r));g=n?null:_(x,f),w=l?null:_(x,c),P(x);let v=await G(E);if(n&&l)return v!==void 0?{gpuTimeMs:v}:{};let F=y(g,Float32Array),I=y(w,Float32Array);g=null,w=null;let[j,N]=await Promise.all([F,I]);return v!==void 0?{x:j,y:N,gpuTimeMs:v}:{x:j,y:N}}finally{!n&&f&&m(f),!l&&c&&m(c),d&&m(d),p&&m(p),g&&m(g),w&&m(w)}}async function Re(a,r,e,i,t,o,u,n,l,s,f,c){let d=n instanceof h,p=f instanceof h,g=o instanceof C,w=r==="no-transpose";if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!w&&r!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(!Number.isInteger(e)||!Number.isInteger(i)||!Number.isInteger(l)||!Number.isInteger(c)||!Number.isInteger(u))throw new Error("m, n, incx, incy, and lda must be integers.");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 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(l<=0||c<=0)throw new Error("incx and incy must be positive.");if(u<i)throw new Error("lda must be >= n.");if(!g&&!(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(!p&&!(f instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(d!==p)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(d&&!g)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(d&&n._buf===f._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(g&&u!==o.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(g&&(o.rows<e||o.cols<i))throw new Error("A is too small for the given m and n.");if(e<0||i<0)throw new Error("m and n must be non-negative.");if(e===0||i===0)return p?{}:{y:f};let k=w?i:e,x=w?e:i;if(!g&&o.length<(e-1)*u+i)throw new Error("A does not have enough elements for the given m, n, and lda.");if(n.length<(k-1)*l+1)throw new Error("x does not have enough elements for the given dimensions and incx.");if(f.length<(x-1)*c+1)throw new Error("y does not have enough elements for the given dimensions and incy.");let v=await A(a,w?"sgemv_n":"sgemv_t"),F=g?o._buf:b(o,"sgemv-A",!1),I=d?n._buf:b(n,"sgemv-x",!1),j=p?f._buf:b(f,"sgemv-y",!0),N=L([{value:e,type:"u32"},{value:i,type:"u32"},{value:t,type:"f32"},{value:s,type:"f32"},{value:l,type:"u32"},{value:c,type:"u32"},{value:u,type:"u32"}],"sgemv-params");try{let U=B(v.getBindGroupLayout(0),[F,I,j,N]),T=w?Math.min(e,a.limits.maxComputeWorkgroupsPerDimension):D(x),{commandEncoder:R,ts:V}=S(v,U,T),z=p?null:_(R,j);P(R);let Y=await G(V);if(p)return Y!==void 0?{gpuTimeMs:Y}:{};let Z=await y(z,Float32Array);return Y!==void 0?{y:Z,gpuTimeMs:Y}:{y:Z}}finally{g||m(F),d||m(I),p||m(j),m(N)}}async function Ve(a,r,e,i,t,o,u,n,l,s,f){let c=u instanceof h,d=s instanceof h,p=t instanceof C,g=r==="lower";if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!g&&r!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(!Number.isInteger(e)||!Number.isInteger(n)||!Number.isInteger(f)||!Number.isInteger(o))throw new Error("n, incx, incy, and lda must be integers.");if(typeof i!="number")throw new Error("alpha must be a number.");if(Number.isNaN(i))throw new Error("alpha must not be NaN.");if(!Number.isFinite(i))throw new Error("alpha must be finite.");if(typeof l!="number")throw new Error("beta must be a number.");if(Number.isNaN(l))throw new Error("beta must not be NaN.");if(!Number.isFinite(l))throw new Error("beta must be finite.");if(n<=0||f<=0)throw new Error("incx and incy must be positive.");if(o<e)throw new Error("lda must be >= n.");if(!p&&!(t instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!c&&!(u instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!d&&!(s instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(c!==d)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(c&&!p)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(c&&u._buf===s._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(p&&o!==t.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(p&&(t.rows<e||t.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 d?{}:{y:s};if(!p&&t.length<(e-1)*o+e)throw new Error("A does not have enough elements for the given n and lda.");if(u.length<(e-1)*n+1)throw new Error("x does not have enough elements for the given n and incx.");if(s.length<(e-1)*f+1)throw new Error("y does not have enough elements for the given n and incy.");let w=await A(a,"ssymv"),k=null,x=null,E=null,v=null;try{k=p?t._buf:b(t,"ssymv-A",!1),x=c?u._buf:b(u,"ssymv-x",!1),E=d?s._buf:b(s,"ssymv-y",!0),v=L([{value:e,type:"u32"},{value:i,type:"f32"},{value:l,type:"f32"},{value:n,type:"u32"},{value:f,type:"u32"},{value:o,type:"u32"},{value:g?0:1,type:"u32"}],"ssymv-params");let F=B(w.getBindGroupLayout(0),[k,x,E,v]),I=Math.min(e,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:j,ts:N}=S(w,F,I),U=d?null:_(j,E);P(j);let T=await G(N);if(d)return T!==void 0?{gpuTimeMs:T}:{};let R=await y(U,Float32Array);return T!==void 0?{y:R,gpuTimeMs:T}:{y:R}}finally{!p&&k&&m(k),!c&&x&&m(x),!d&&E&&m(E),v&&m(v)}}async function He(a,r,e,i,t,o,u,n,l,s,f){let c=n instanceof h,d=s instanceof h,p=o instanceof C,g=r==="lower",w=e==="no-transpose",k=i==="unit";if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!g&&r!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(!w&&e!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(!k&&i!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(!Number.isInteger(t)||!Number.isInteger(l)||!Number.isInteger(f)||!Number.isInteger(u))throw new Error("n, incx, incy, and lda must be integers.");if(l<=0||f<=0)throw new Error("incx and incy must be positive.");if(u<t)throw new Error("lda must be >= n.");if(!p&&!(o instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!c&&!(n instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(!d&&!(s instanceof Float32Array))throw new Error("y must be a Float32Array or GpuVector.");if(c!==d)throw new Error("x and y must be the same type (both Float32Array or both GpuVector).");if(c&&n._buf===s._buf)throw new Error("x and y must not reference the same GPU buffer when both are GpuVectors.");if(c&&!p)throw new Error("A must be a GpuMatrix when x and y are GpuVectors.");if(p&&d&&o._buf===s._buf)throw new Error("A and y must not reference the same GPU buffer.");if(p&&u!==o.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(p&&(o.rows<t||o.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:s};if(!p&&o.length<(t-1)*u+t)throw new Error("A does not have enough elements for the given n and lda.");if(n.length<(t-1)*l+1)throw new Error("x does not have enough elements for the given n and incx.");if(s.length<(t-1)*f+1)throw new Error("y does not have enough elements for the given n and incy.");let x=await A(a,"strmv"),E=null,v=null,F=null,I=null;try{E=p?o._buf:b(o,"strmv-A",!1),v=c?n._buf:b(n,"strmv-x",!1),F=d?s._buf:b(s,"strmv-y",!0),I=L([{value:t,type:"u32"},{value:l,type:"u32"},{value:f,type:"u32"},{value:u,type:"u32"},{value:w?0:1,type:"u32"},{value:g?0:1,type:"u32"},{value:k?1:0,type:"u32"}],"strmv-params");let j=B(x.getBindGroupLayout(0),[E,v,F,I]),N=Math.min(t,a.limits.maxComputeWorkgroupsPerDimension),{commandEncoder:U,ts:T}=S(x,j,N),R=d?null:_(U,F);P(U);let V=await G(T);if(d)return V!==void 0?{gpuTimeMs:V}:{};let z=await y(R,Float32Array);return V!==void 0?{y:z,gpuTimeMs:V}:{y:z}}finally{!p&&E&&m(E),!c&&v&&m(v),!d&&F&&m(F),I&&m(I)}}var q=64;function De(a,r,e){let i=new ArrayBuffer(a*r),t=new DataView(i);for(let o=0;o<a;o++){let u=e(o),n=o*r;u.forEach((l,s)=>t.setUint32(n+s*4,l,!0))}return i}function Ce(a,r,e){let i=a.createBuffer({label:e,size:r.byteLength,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST});return a.queue.writeBuffer(i,0,r),i}async function Oe(a,r,e,i,t,o,u,n,l){let s=n instanceof h,f=o instanceof C,c=r==="lower",d=e==="no-transpose",p=i==="unit";if(!(a instanceof GPUDevice))throw new Error("device must be a GPUDevice.");if(!c&&r!=="upper")throw new Error("uplo must be 'lower' or 'upper'.");if(!d&&e!=="transpose")throw new Error("trans must be 'no-transpose' or 'transpose'.");if(!p&&i!=="non-unit")throw new Error("diag must be 'unit' or 'non-unit'.");if(!Number.isInteger(t)||!Number.isInteger(l)||!Number.isInteger(u))throw new Error("n, incx, and lda must be integers.");if(l<=0)throw new Error("incx must be positive.");if(u<t)throw new Error("lda must be >= n.");if(!f&&!(o instanceof Float32Array))throw new Error("A must be a Float32Array or GpuMatrix.");if(!s&&!(n instanceof Float32Array))throw new Error("x must be a Float32Array or GpuVector.");if(s&&!f)throw new Error("A must be a GpuMatrix when x is a GpuVector.");if(f&&u!==o.lda)throw new Error("lda must match A.lda when A is a GpuMatrix.");if(f&&(o.rows<t||o.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 s?{}:{x:n};if(!f&&o.length<(t-1)*u+t)throw new Error("A does not have enough elements for the given n and lda.");if(n.length<(t-1)*l+1)throw new Error("x does not have enough elements for the given n and incx.");let g=await A(a,"strsv_invert_block"),w=await A(a,"strsv_apply_inverse"),k=await A(a,"strsv_update"),x=d===c,E=[];for(let z=0;z<t;z+=q)E.push(z);x||E.reverse();let v=E.length,F=a.limits.maxComputeWorkgroupsPerDimension,I=a.limits.minUniformBufferOffsetAlignment,j=null,N=null,U=null,T=null,R=null,V=null;try{j=f?o._buf:b(o,"strsv-A",!1),N=s?n._buf:b(n,"strsv-x",!0),U=H(v*q*q*4,"strsv-Ainv");let z=De(v,I,K=>{let X=K*q,tr=Math.min(X+q,t);return[l,K,X,tr]});T=Ce(a,z,"strsv-apply-params");let Y=De(v,I,K=>{let X=K*q,tr=Math.min(X+q,t);return[t,l,u,d?0:1,c?0:1,X,tr]});R=Ce(a,Y,"strsv-update-params");let{commandEncoder:Z,querySet:er}=fr();V=L([{value:t,type:"u32"},{value:u,type:"u32"},{value:d?0:1,type:"u32"},{value:c?0:1,type:"u32"},{value:p?1:0,type:"u32"}],"strsv-invert-params");let ze=B(g.getBindGroupLayout(0),[j,U,V]);ar(Z,g,ze,{x:q,y:v},er?{timestampWrites:{querySet:er,beginningOfPassWriteIndex:0}}:void 0);for(let K=0;K<E.length;K++){let X=E[K],tr=Math.min(X+q,t),Ze=X/q,Ke=K===E.length-1,wr=Ze*I,Xe=B(w.getBindGroupLayout(0),[U,N,{buffer:T,offset:wr,size:16}]);ar(Z,w,Xe,1,Ke&&er?{timestampWrites:{querySet:er,endOfPassWriteIndex:1}}:void 0);let br=x?t-tr:X;if(br===0)continue;let $e=B(k.getBindGroupLayout(0),[j,N,{buffer:R,offset:wr,size:32}]),Ye=Math.min(br,F);ar(Z,k,$e,Ye)}let Qe=ur(Z,er),qe=s?null:_(Z,N);P(Z);let ir=await G(Qe);if(s)return ir!==void 0?{gpuTimeMs:ir}:{};let gr=await y(qe,Float32Array);return ir!==void 0?{x:gr,gpuTimeMs:ir}:{x:gr}}finally{!f&&j&&m(j),!s&&N&&m(N),U&&m(U),T&&m(T),R&&m(R),V&&m(V)}}return it(Rt);})();
|