wgblas 0.1.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/LICENSE +201 -0
- package/README.md +161 -0
- package/dist/wgblas.browser.js +422 -0
- package/index.d.mts +94 -0
- package/index.mjs +16 -0
- package/package.json +127 -0
- package/src/classes/GpuVector.d.mts +69 -0
- package/src/classes/GpuVector.mjs +33 -0
- package/src/devdocs.mjs +81 -0
- package/src/index.mjs +4 -0
- package/src/init.mjs +105 -0
- package/src/isamax/isamax.d.mts +46 -0
- package/src/isamax/isamax.mjs +115 -0
- package/src/random/random.d.mts +61 -0
- package/src/random/random.mjs +11 -0
- package/src/sasum/sasum.d.mts +44 -0
- package/src/sasum/sasum.mjs +98 -0
- package/src/saxpy/saxpy.d.mts +54 -0
- package/src/saxpy/saxpy.mjs +90 -0
- package/src/scopy/scopy.d.mts +50 -0
- package/src/scopy/scopy.mjs +87 -0
- package/src/sdot/sdot.d.mts +52 -0
- package/src/sdot/sdot.mjs +118 -0
- package/src/shaders/browser-shaders.mjs +27 -0
- package/src/shaders/index.mjs +27 -0
- package/src/snrm2/snrm2.d.mts +44 -0
- package/src/snrm2/snrm2.mjs +100 -0
- package/src/srot/srot.d.mts +62 -0
- package/src/srot/srot.mjs +95 -0
- package/src/srotm/srotm.d.mts +60 -0
- package/src/srotm/srotm.mjs +94 -0
- package/src/sscal/sscal.d.mts +46 -0
- package/src/sscal/sscal.mjs +71 -0
- package/src/sswap/sswap.d.mts +50 -0
- package/src/sswap/sswap.mjs +90 -0
- package/src/util/benchmark.mjs +103 -0
- package/src/util/bindgroup.mjs +22 -0
- package/src/util/buffer.mjs +160 -0
- package/src/util/compute.mjs +52 -0
- package/src/util/index.mjs +12 -0
- package/src/util/pipeline.mjs +82 -0
- package/src/util/result.mjs +19 -0
- package/src/util/workgroup.mjs +31 -0
|
@@ -0,0 +1,422 @@
|
|
|
1
|
+
var wgblas=(()=>{var jr=Object.create;var L=Object.defineProperty;var Qr=Object.getOwnPropertyDescriptor;var Xr=Object.getOwnPropertyNames;var Hr=Object.getPrototypeOf,Kr=Object.prototype.hasOwnProperty;var q=(t=>typeof require<"u"?require:typeof Proxy<"u"?new Proxy(t,{get:(r,e)=>(typeof require<"u"?require:r)[e]}):t)(function(t){if(typeof require<"u")return require.apply(this,arguments);throw Error('Dynamic require of "'+t+'" is not supported')});var I=(t,r,e)=>()=>{if(e)throw e[0];try{return t&&(r=t(t=0)),r}catch(a){throw e=[a],a}};var X=(t,r)=>{for(var e in r)L(t,e,{get:r[e],enumerable:!0})},H=(t,r,e,a)=>{if(r&&typeof r=="object"||typeof r=="function")for(let o of Xr(r))!Kr.call(t,o)&&o!==e&&L(t,o,{get:()=>r[o],enumerable:!(a=Qr(r,o))||a.enumerable});return t};var Y=(t,r,e)=>(e=t!=null?jr(Hr(t)):{},H(r||!t||!t.__esModule?L(e,"default",{value:t,enumerable:!0}):e,t)),Zr=t=>H(L({},"__esModule",{value:!0}),t);var ur,sr=I(()=>{ur=`// amax reduction: collapses 2*WGS (value, index) pairs into one index.
|
|
2
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
3
|
+
// partials_val and partials_idx must have exactly 2*WGS entries.
|
|
4
|
+
|
|
5
|
+
@group(0) @binding(0) var<storage, read> partials_val: array<f32>;
|
|
6
|
+
@group(0) @binding(1) var<storage, read> partials_idx: array<u32>;
|
|
7
|
+
@group(0) @binding(2) var<storage, read_write> result: array<u32>;
|
|
8
|
+
|
|
9
|
+
const WGS: u32 = 64;
|
|
10
|
+
|
|
11
|
+
var<workgroup> tile_val: array<f32, 64>;
|
|
12
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
13
|
+
|
|
14
|
+
@compute @workgroup_size(64)
|
|
15
|
+
fn reduce(
|
|
16
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
17
|
+
) {
|
|
18
|
+
let i = lid.x;
|
|
19
|
+
let a_val = partials_val[i];
|
|
20
|
+
let b_val = partials_val[i + WGS];
|
|
21
|
+
if (b_val > a_val || (b_val == a_val && partials_idx[i + WGS] < partials_idx[i])) {
|
|
22
|
+
tile_val[i] = b_val;
|
|
23
|
+
tile_idx[i] = partials_idx[i + WGS];
|
|
24
|
+
} else {
|
|
25
|
+
tile_val[i] = a_val;
|
|
26
|
+
tile_idx[i] = partials_idx[i];
|
|
27
|
+
}
|
|
28
|
+
workgroupBarrier();
|
|
29
|
+
|
|
30
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
31
|
+
if (i < s) {
|
|
32
|
+
let c_val = tile_val[i];
|
|
33
|
+
let d_val = tile_val[i + s];
|
|
34
|
+
if (d_val > c_val || (d_val == c_val && tile_idx[i + s] < tile_idx[i])) {
|
|
35
|
+
tile_val[i] = d_val;
|
|
36
|
+
tile_idx[i] = tile_idx[i + s];
|
|
37
|
+
}
|
|
38
|
+
}
|
|
39
|
+
workgroupBarrier();
|
|
40
|
+
}
|
|
41
|
+
|
|
42
|
+
if (i == 0u) { result[0] = tile_idx[0]; }
|
|
43
|
+
}
|
|
44
|
+
`});var cr,fr=I(()=>{cr=`// sum reduction: collapses 2*WGS partials into one scalar.
|
|
45
|
+
// dispatch: 1 workgroup of WGS threads.
|
|
46
|
+
// partials must have exactly 2*WGS entries.
|
|
47
|
+
|
|
48
|
+
@group(0) @binding(0) var<storage, read> partials: array<f32>;
|
|
49
|
+
@group(0) @binding(1) var<storage, read_write> result: array<f32>;
|
|
50
|
+
|
|
51
|
+
const WGS: u32 = 64;
|
|
52
|
+
|
|
53
|
+
var<workgroup> tile: array<f32, 64>;
|
|
54
|
+
|
|
55
|
+
@compute @workgroup_size(64)
|
|
56
|
+
fn reduce(
|
|
57
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
58
|
+
) {
|
|
59
|
+
let i = lid.x;
|
|
60
|
+
tile[i] = partials[i] + partials[i + WGS];
|
|
61
|
+
workgroupBarrier();
|
|
62
|
+
|
|
63
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
64
|
+
if (i < s) { tile[i] += tile[i + s]; }
|
|
65
|
+
workgroupBarrier();
|
|
66
|
+
}
|
|
67
|
+
|
|
68
|
+
if (i == 0u) { result[0] = tile[0]; }
|
|
69
|
+
}
|
|
70
|
+
`});var mr,pr=I(()=>{mr=`// sscal: x = alpha * x
|
|
71
|
+
|
|
72
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
73
|
+
|
|
74
|
+
struct Params {
|
|
75
|
+
n: u32,
|
|
76
|
+
alpha: f32,
|
|
77
|
+
x_inc: u32,
|
|
78
|
+
}
|
|
79
|
+
|
|
80
|
+
@group(0) @binding(1) var<uniform> params: Params;
|
|
81
|
+
|
|
82
|
+
const WGS: u32 = 64;
|
|
83
|
+
|
|
84
|
+
@compute @workgroup_size(64)
|
|
85
|
+
fn main(
|
|
86
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
87
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
88
|
+
) {
|
|
89
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
90
|
+
x[id * params.x_inc] = params.alpha * x[id * params.x_inc];
|
|
91
|
+
}
|
|
92
|
+
}
|
|
93
|
+
`});var dr,lr=I(()=>{dr=`// sswap: x <-> y
|
|
94
|
+
|
|
95
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
96
|
+
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
97
|
+
|
|
98
|
+
struct Params {
|
|
99
|
+
n: u32,
|
|
100
|
+
x_inc: u32,
|
|
101
|
+
y_inc: u32,
|
|
102
|
+
}
|
|
103
|
+
|
|
104
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
105
|
+
|
|
106
|
+
const WGS: u32 = 64;
|
|
107
|
+
|
|
108
|
+
@compute @workgroup_size(64)
|
|
109
|
+
fn main(
|
|
110
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
111
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
112
|
+
) {
|
|
113
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
114
|
+
let temp = x[id * params.x_inc];
|
|
115
|
+
x[id * params.x_inc] = y[id * params.y_inc];
|
|
116
|
+
y[id * params.y_inc] = temp;
|
|
117
|
+
}
|
|
118
|
+
}
|
|
119
|
+
`});var wr,gr=I(()=>{wr=`// saxpy: y = alpha * x + y
|
|
120
|
+
|
|
121
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
122
|
+
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
123
|
+
|
|
124
|
+
struct Params {
|
|
125
|
+
n: u32,
|
|
126
|
+
alpha: f32,
|
|
127
|
+
x_inc: u32,
|
|
128
|
+
y_inc: u32,
|
|
129
|
+
}
|
|
130
|
+
|
|
131
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
132
|
+
|
|
133
|
+
const WGS: u32 = 64;
|
|
134
|
+
|
|
135
|
+
@compute @workgroup_size(64)
|
|
136
|
+
fn main(
|
|
137
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
138
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
139
|
+
) {
|
|
140
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
141
|
+
y[id * params.y_inc] = params.alpha * x[id * params.x_inc] + y[id * params.y_inc];
|
|
142
|
+
}
|
|
143
|
+
}
|
|
144
|
+
`});var vr,yr=I(()=>{vr=`// scopy: y = x
|
|
145
|
+
|
|
146
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
147
|
+
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
148
|
+
|
|
149
|
+
struct Params {
|
|
150
|
+
n: u32,
|
|
151
|
+
x_inc: u32,
|
|
152
|
+
y_inc: u32,
|
|
153
|
+
}
|
|
154
|
+
|
|
155
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
156
|
+
|
|
157
|
+
const WGS: u32 = 64;
|
|
158
|
+
|
|
159
|
+
@compute @workgroup_size(64)
|
|
160
|
+
fn main(
|
|
161
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
162
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
163
|
+
) {
|
|
164
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
165
|
+
y[id * params.y_inc] = x[id * params.x_inc];
|
|
166
|
+
}
|
|
167
|
+
}
|
|
168
|
+
`});var xr,br=I(()=>{xr=`// sdot: result = sum(x[i] * y[i])
|
|
169
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sum.wgsl.
|
|
170
|
+
|
|
171
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
172
|
+
@group(0) @binding(1) var<storage, read> y: array<f32>;
|
|
173
|
+
@group(0) @binding(2) var<storage, read_write> partials: array<f32>;
|
|
174
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
175
|
+
|
|
176
|
+
struct Params {
|
|
177
|
+
n: u32,
|
|
178
|
+
x_inc: u32,
|
|
179
|
+
y_inc: u32,
|
|
180
|
+
}
|
|
181
|
+
|
|
182
|
+
const WGS: u32 = 64;
|
|
183
|
+
|
|
184
|
+
var<workgroup> tile: array<f32, 64>;
|
|
185
|
+
|
|
186
|
+
@compute @workgroup_size(64)
|
|
187
|
+
fn main(
|
|
188
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
189
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
190
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
191
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
192
|
+
) {
|
|
193
|
+
var acc: f32 = 0.0;
|
|
194
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
195
|
+
acc += x[id * params.x_inc] * y[id * params.y_inc];
|
|
196
|
+
}
|
|
197
|
+
tile[lid.x] = acc;
|
|
198
|
+
workgroupBarrier();
|
|
199
|
+
|
|
200
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
201
|
+
if (lid.x < s) { tile[lid.x] += tile[lid.x + s]; }
|
|
202
|
+
workgroupBarrier();
|
|
203
|
+
}
|
|
204
|
+
|
|
205
|
+
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
206
|
+
}
|
|
207
|
+
`});var _r,hr=I(()=>{_r=`// sasum: result = sum(|x[i]|)
|
|
208
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/abssum.wgsl.
|
|
209
|
+
|
|
210
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
211
|
+
@group(0) @binding(1) var<storage, read_write> partials: array<f32>;
|
|
212
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
213
|
+
|
|
214
|
+
struct Params {
|
|
215
|
+
n: u32,
|
|
216
|
+
x_inc: u32,
|
|
217
|
+
}
|
|
218
|
+
|
|
219
|
+
const WGS: u32 = 64;
|
|
220
|
+
|
|
221
|
+
var<workgroup> tile: array<f32, 64>;
|
|
222
|
+
|
|
223
|
+
@compute @workgroup_size(64)
|
|
224
|
+
fn main(
|
|
225
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
226
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
227
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
228
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
229
|
+
) {
|
|
230
|
+
var acc: f32 = 0.0;
|
|
231
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
232
|
+
acc += abs(x[id * params.x_inc]);
|
|
233
|
+
}
|
|
234
|
+
tile[lid.x] = acc;
|
|
235
|
+
workgroupBarrier();
|
|
236
|
+
|
|
237
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
238
|
+
if (lid.x < s) { tile[lid.x] += tile[lid.x + s]; }
|
|
239
|
+
workgroupBarrier();
|
|
240
|
+
}
|
|
241
|
+
|
|
242
|
+
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
243
|
+
}
|
|
244
|
+
`});var Br,Gr=I(()=>{Br=`// snrm2: result = sqrt(sum(x[i] * x[i]))
|
|
245
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sqsum.wgsl.
|
|
246
|
+
|
|
247
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
248
|
+
@group(0) @binding(1) var<storage, read_write> partials: array<f32>;
|
|
249
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
250
|
+
|
|
251
|
+
struct Params {
|
|
252
|
+
n: u32,
|
|
253
|
+
x_inc: u32,
|
|
254
|
+
}
|
|
255
|
+
|
|
256
|
+
const WGS: u32 = 64;
|
|
257
|
+
|
|
258
|
+
var<workgroup> tile: array<f32, 64>;
|
|
259
|
+
|
|
260
|
+
@compute @workgroup_size(64)
|
|
261
|
+
fn main(
|
|
262
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
263
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
264
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
265
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
266
|
+
) {
|
|
267
|
+
var acc: f32 = 0.0;
|
|
268
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
269
|
+
let v = x[id * params.x_inc];
|
|
270
|
+
acc += v * v;
|
|
271
|
+
}
|
|
272
|
+
tile[lid.x] = acc;
|
|
273
|
+
workgroupBarrier();
|
|
274
|
+
|
|
275
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
276
|
+
if (lid.x < s) { tile[lid.x] += tile[lid.x + s]; }
|
|
277
|
+
workgroupBarrier();
|
|
278
|
+
}
|
|
279
|
+
|
|
280
|
+
if (lid.x == 0u) { partials[wgid.x] = tile[0]; }
|
|
281
|
+
}
|
|
282
|
+
`});var Pr,Er=I(()=>{Pr=`// srot: x = c*x + s*y, y = -s*x + c*y
|
|
283
|
+
|
|
284
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
285
|
+
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
286
|
+
|
|
287
|
+
struct Params {
|
|
288
|
+
n: u32,
|
|
289
|
+
c: f32,
|
|
290
|
+
s: f32,
|
|
291
|
+
x_inc: u32,
|
|
292
|
+
y_inc: u32,
|
|
293
|
+
}
|
|
294
|
+
|
|
295
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
296
|
+
|
|
297
|
+
const WGS: u32 = 64;
|
|
298
|
+
|
|
299
|
+
@compute @workgroup_size(64)
|
|
300
|
+
fn main(
|
|
301
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
302
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
303
|
+
) {
|
|
304
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
305
|
+
let xi = x[id * params.x_inc];
|
|
306
|
+
let yi = y[id * params.y_inc];
|
|
307
|
+
x[id * params.x_inc] = params.c * xi + params.s * yi;
|
|
308
|
+
y[id * params.y_inc] = -params.s * xi + params.c * yi;
|
|
309
|
+
}
|
|
310
|
+
}
|
|
311
|
+
`});var kr,Ar=I(()=>{kr=`// srotm: applies modified Givens rotation H to vectors x and y.
|
|
312
|
+
// param[0] = flag: -1 (full H), 0 (unit diagonal), 1 (unit off-diagonal)
|
|
313
|
+
// param = [ flag, h11, h21, h12, h22 ]
|
|
314
|
+
// flag == -2 (identity/no-op) is handled in JS before dispatch reaches here.
|
|
315
|
+
|
|
316
|
+
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
|
|
317
|
+
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
|
|
318
|
+
@group(0) @binding(2) var<storage, read> param: array<f32>;
|
|
319
|
+
|
|
320
|
+
struct Params {
|
|
321
|
+
n: u32,
|
|
322
|
+
x_inc: u32,
|
|
323
|
+
y_inc: u32,
|
|
324
|
+
}
|
|
325
|
+
|
|
326
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
327
|
+
|
|
328
|
+
const WGS: u32 = 64;
|
|
329
|
+
|
|
330
|
+
@compute @workgroup_size(64)
|
|
331
|
+
fn main(
|
|
332
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
333
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
334
|
+
) {
|
|
335
|
+
let flag = param[0];
|
|
336
|
+
|
|
337
|
+
var h11: f32; var h12: f32;
|
|
338
|
+
var h21: f32; var h22: f32;
|
|
339
|
+
|
|
340
|
+
if (flag == -1.0) {
|
|
341
|
+
// full 2x2 matrix
|
|
342
|
+
h11 = param[1]; h21 = param[2];
|
|
343
|
+
h12 = param[3]; h22 = param[4];
|
|
344
|
+
} else if (flag == 0.0) {
|
|
345
|
+
// diagonal fixed at 1
|
|
346
|
+
h11 = 1.0; h21 = param[2];
|
|
347
|
+
h12 = param[3]; h22 = 1.0;
|
|
348
|
+
} else if (flag == 1.0) {
|
|
349
|
+
// flag == 1.0: off-diagonal fixed at +1 / -1
|
|
350
|
+
h11 = param[1]; h21 = -1.0;
|
|
351
|
+
h12 = 1.0; h22 = param[4];
|
|
352
|
+
}
|
|
353
|
+
|
|
354
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
355
|
+
let xi = x[id * params.x_inc];
|
|
356
|
+
let yi = y[id * params.y_inc];
|
|
357
|
+
x[id * params.x_inc] = h11 * xi + h12 * yi;
|
|
358
|
+
y[id * params.y_inc] = h21 * xi + h22 * yi;
|
|
359
|
+
}
|
|
360
|
+
}
|
|
361
|
+
`});var Fr,Sr=I(()=>{Fr=`// isamax: returns index of element with largest absolute value
|
|
362
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/argmax.wgsl.
|
|
363
|
+
|
|
364
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
365
|
+
@group(0) @binding(1) var<storage, read_write> partials_val: array<f32>;
|
|
366
|
+
@group(0) @binding(2) var<storage, read_write> partials_idx: array<u32>;
|
|
367
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
368
|
+
|
|
369
|
+
struct Params {
|
|
370
|
+
n: u32,
|
|
371
|
+
x_inc: u32,
|
|
372
|
+
}
|
|
373
|
+
|
|
374
|
+
const WGS: u32 = 64;
|
|
375
|
+
|
|
376
|
+
var<workgroup> tile_val: array<f32, 64>;
|
|
377
|
+
var<workgroup> tile_idx: array<u32, 64>;
|
|
378
|
+
|
|
379
|
+
@compute @workgroup_size(64)
|
|
380
|
+
fn main(
|
|
381
|
+
@builtin(global_invocation_id) gid: vec3u,
|
|
382
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
383
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
384
|
+
@builtin(num_workgroups) num_wg: vec3u,
|
|
385
|
+
) {
|
|
386
|
+
// -1.0 is a safe sentinel: any |x[i]| >= 0 beats it,
|
|
387
|
+
// so workgroups with no elements lose gracefully in the epilogue.
|
|
388
|
+
var best_val: f32 = -1.0;
|
|
389
|
+
var best_idx: u32 = 0u;
|
|
390
|
+
|
|
391
|
+
for (var id = gid.x; id < params.n; id += num_wg.x * WGS) {
|
|
392
|
+
let v = abs(x[id * params.x_inc]);
|
|
393
|
+
if (v > best_val) {
|
|
394
|
+
best_val = v;
|
|
395
|
+
best_idx = id;
|
|
396
|
+
}
|
|
397
|
+
}
|
|
398
|
+
|
|
399
|
+
tile_val[lid.x] = best_val;
|
|
400
|
+
tile_idx[lid.x] = best_idx;
|
|
401
|
+
workgroupBarrier();
|
|
402
|
+
|
|
403
|
+
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
404
|
+
if (lid.x < s) {
|
|
405
|
+
let a_val = tile_val[lid.x];
|
|
406
|
+
let b_val = tile_val[lid.x + s];
|
|
407
|
+
if (b_val > a_val || (b_val == a_val && tile_idx[lid.x + s] < tile_idx[lid.x])) {
|
|
408
|
+
tile_val[lid.x] = b_val;
|
|
409
|
+
tile_idx[lid.x] = tile_idx[lid.x + s];
|
|
410
|
+
}
|
|
411
|
+
}
|
|
412
|
+
workgroupBarrier();
|
|
413
|
+
}
|
|
414
|
+
|
|
415
|
+
if (lid.x == 0u) {
|
|
416
|
+
partials_val[wgid.x] = tile_val[0];
|
|
417
|
+
partials_idx[wgid.x] = tile_idx[0];
|
|
418
|
+
}
|
|
419
|
+
}
|
|
420
|
+
`});var Rr={};X(Rr,{shaderSources:()=>pe});var pe,Wr=I(()=>{sr();fr();pr();lr();gr();yr();br();hr();Gr();Er();Ar();Sr();pe={"reduction/argmax":ur,"reduction/sum":cr,sscal:mr,sswap:dr,saxpy:wr,scopy:vr,sdot:xr,sasum:_r,snrm2:Br,srot:Pr,srotm:kr,isamax:Fr}});var we={};X(we,{GpuVector:()=>m,cleanup:()=>ar,gpuName:()=>or,init:()=>tr,isamax:()=>qr,randomFloat32Array:()=>ir,randomFloat64Array:()=>nr,sasum:()=>zr,saxpy:()=>Tr,scopy:()=>Dr,sdot:()=>Vr,snrm2:()=>Lr,srot:()=>Yr,srotm:()=>$r,sscal:()=>Ir,sswap:()=>Mr});function K(t,r){return r?t.features.has("timestamp-query")?{requiredFeatures:["timestamp-query"]}:(console.warn("timestamp-query not supported on this device \u2014 benchmark mode disabled."),{}):{}}function Z(){if(!rr())return{querySet:null,passDescriptor:void 0};let r=F().createQuerySet({type:"timestamp",count:2});return{querySet:r,passDescriptor:{timestampWrites:{querySet:r,beginningOfPassWriteIndex:0,endOfPassWriteIndex:1}}}}function J(t,r){if(!r)return null;let e=F(),a=e.createBuffer({label:"timestamp-resolve",size:16,usage:GPUBufferUsage.QUERY_RESOLVE|GPUBufferUsage.COPY_SRC});t.resolveQuerySet(r,0,2,a,0);let o=e.createBuffer({label:"timestamp-readback",size:16,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ});return t.copyBufferToBuffer(a,0,o,0,16),{tsReadBuffer:o,resolveBuffer:a,querySet:r}}async function G(t){if(!t)return;let{tsReadBuffer:r,resolveBuffer:e,querySet:a}=t;await r.mapAsync(GPUMapMode.READ);let o=new BigInt64Array(r.getMappedRange().slice());return r.unmap(),r.destroy(),e.destroy(),a.destroy(),Number(o[1]-o[0])/1e6}var T=null,C=null,er=null,$=!1;async function tr({powerPreference:t="high-performance",benchmark:r=!1}={}){if(T)return T;let e;if(typeof window>"u"){let{create:a,globals:o}=await import("webgpu");Object.assign(globalThis,o),e=a([]),er=e}else e=navigator.gpu;if(!e)throw new Error("WebGPU not supported in this environment.");if(C=await e.requestAdapter({powerPreference:t})??await e.requestAdapter(),!C)throw new Error("No WebGPU adapter found.");return $=r,T=await C.requestDevice(K(C,r)),T.addEventListener("uncapturederror",a=>{console.error("Uncaptured GPU error:",a.error.message)}),T}function ar(){T&&(T.destroy(),T=null),C=null,er=null,$=!1}function or(){if(!C)throw new Error("WebGPU adapter not initialized \u2014 call init() first.");let{device:t,description:r}=C.info;return{description:r||"unknown",device:t||"unknown"}}function rr(){return $}function F(){if(!T)throw new Error("WebGPU device not initialized \u2014 call init() first.");return T}function w(...t){t.flat().forEach(r=>r.destroy())}function v(t,r="blas-input",e=!1){let a=F(),o=a.limits.maxStorageBufferBindingSize,n=t.byteLength;if(n>o)throw new Error(`Buffer size ${n} bytes exceeds device limit of ${o} bytes.`);let s=e?GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC:GPUBufferUsage.STORAGE,i=a.createBuffer({label:r,size:n,usage:s,mappedAtCreation:!0});return new Float32Array(i.getMappedRange()).set(t),i.unmap(),i}function N(t,r="blas-storage"){return F().createBuffer({label:r,size:t,usage:GPUBufferUsage.STORAGE})}function V(t,r="blas-result"){return F().createBuffer({label:r,size:t,usage:GPUBufferUsage.STORAGE|GPUBufferUsage.COPY_SRC})}function h(t,r){let a=F().createBuffer({label:"blas-readback",size:r.size,usage:GPUBufferUsage.COPY_DST|GPUBufferUsage.MAP_READ});return t.copyBufferToBuffer(r,0,a,0,r.size),a}function R(t,r="blas-params"){let e=F(),a=t.length*4,o=Math.ceil(a/16)*16,n=new ArrayBuffer(o),s=new DataView(n);t.forEach(({value:u,type:f},c)=>{let p=c*4;if(f==="u32")s.setUint32(p,u,!0);else if(f==="i32")s.setInt32(p,u,!0);else if(f==="f32")s.setFloat32(p,u,!0);else throw new Error(`Unknown param type "${f}". Use "f32", "u32", or "i32".`)});let i=e.createBuffer({label:r,size:o,usage:GPUBufferUsage.UNIFORM|GPUBufferUsage.COPY_DST});return e.queue.writeBuffer(i,0,n),i}async function _(t,r=Float32Array){await t.mapAsync(GPUMapMode.READ);let e=new r(t.getMappedRange().slice());return t.unmap(),e}var m=class t{constructor(r,e,a=Float32Array){this._buf=r,this.length=e,this.dtype=a}static from(r){if(!(r instanceof Float32Array))throw new Error("GpuVector.from expects a Float32Array.");let e=v(r,"gpu-vector",!0);return new t(e,r.length,r.constructor)}async read(){let r=F(),e=r.createCommandEncoder(),a=h(e,this._buf);return r.queue.submit([e.finish()]),_(a,this.dtype)}destroy(){this._buf.destroy()}};function ir(t,r=-1,e=1){let a=new Float32Array(t);for(let o=0;o<t;o++)a[o]=r+Math.random()*(e-r);return a}function nr(t,r=-1,e=1){let a=new Float64Array(t);for(let o=0;o<t;o++)a[o]=r+Math.random()*(e-r);return a}function B(t,r,e=null){let a=F(),n=(e?[...r,e]:[...r]).map((s,i)=>({binding:i,resource:{buffer:s}}));return a.createBindGroup({layout:t,entries:n})}function E(t){F().queue.submit([t.finish()])}function P(t,r,e){let a=F(),{querySet:o,passDescriptor:n}=Z(),s=a.createCommandEncoder(),i=s.beginComputePass(n);i.setPipeline(t),i.setBindGroup(0,r),typeof e=="number"?i.dispatchWorkgroups(e):i.dispatchWorkgroups(e.x,e.y),i.end();let u=J(s,o);return s._passEncoder=i,{commandEncoder:s,ts:u}}var de={},j=new WeakMap;async function A(t,r){j.has(t)||j.set(t,new Map);let e=j.get(t);return e.has(r)||e.set(r,await le(r)),e.get(r)}async function me(t){if(typeof window<"u"){let{shaderSources:r}=await Promise.resolve().then(()=>(Wr(),Rr)),e=r[t];if(!e)throw new Error(`Shader "${t}" not found in browser bundle.`);return e}else{let{readFileSync:r}=await import("fs"),{fileURLToPath:e}=await import("url"),{dirname:a,join:o}=await import("path"),n=a(e(de.url));return r(o(n,`../shaders/${t}.wgsl`),"utf8")}}async function le(t){let r=F(),e=await me(t),a=r.createShaderModule({label:t,code:e}),n=(await a.getCompilationInfo()).messages.filter(i=>i.type==="error");if(n.length>0)throw new Error(`Shader "${t}" compilation failed:
|
|
421
|
+
${n.map(i=>` line ${i.lineNum}: ${i.message}`).join(`
|
|
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);})();
|
package/index.d.mts
ADDED
|
@@ -0,0 +1,94 @@
|
|
|
1
|
+
/**
|
|
2
|
+
* @module docs
|
|
3
|
+
*/
|
|
4
|
+
export { GpuVector } from "./src/classes/GpuVector.mjs";
|
|
5
|
+
export {
|
|
6
|
+
randomFloat32Array,
|
|
7
|
+
randomFloat64Array,
|
|
8
|
+
} from "./src/random/random.mjs";
|
|
9
|
+
export { sscal } from "./src/sscal/sscal.mjs";
|
|
10
|
+
export { sswap } from "./src/sswap/sswap.mjs";
|
|
11
|
+
export { saxpy } from "./src/saxpy/saxpy.mjs";
|
|
12
|
+
export { scopy } from "./src/scopy/scopy.mjs";
|
|
13
|
+
export { sdot } from "./src/sdot/sdot.mjs";
|
|
14
|
+
export { sasum } from "./src/sasum/sasum.mjs";
|
|
15
|
+
export { snrm2 } from "./src/snrm2/snrm2.mjs";
|
|
16
|
+
export { isamax } from "./src/isamax/isamax.mjs";
|
|
17
|
+
export { srot } from "./src/srot/srot.mjs";
|
|
18
|
+
export { srotm } from "./src/srotm/srotm.mjs";
|
|
19
|
+
|
|
20
|
+
/**
|
|
21
|
+
* Initializes the WebGPU device.
|
|
22
|
+
*
|
|
23
|
+
* @param options.powerPreference - GPU power preference (default: `"high-performance"`).
|
|
24
|
+
* This is a hint to the browser: on dual-GPU systems, `"high-performance"` typically favors the discrete GPU
|
|
25
|
+
* and `"low-power"` favors the integrated one.
|
|
26
|
+
* See [MDN: GPU.requestAdapter()](https://developer.mozilla.org/en-US/docs/Web/API/GPU/requestAdapter).
|
|
27
|
+
* @param options.benchmark - enable GPU timestamp queries; BLAS functions return `{ result, gpuTimeMs }` (default: `false`)
|
|
28
|
+
*
|
|
29
|
+
* @example Default (high-performance GPU)
|
|
30
|
+
* ```js
|
|
31
|
+
* import { init, gpuName } from "wgblas";
|
|
32
|
+
* await init();
|
|
33
|
+
* const { description, device } = gpuName();
|
|
34
|
+
* console.log("description:", description, "device:", device);
|
|
35
|
+
* ```
|
|
36
|
+
*
|
|
37
|
+
* @example Low-power (integrated GPU)
|
|
38
|
+
* ```js
|
|
39
|
+
* import { init, gpuName } from "wgblas";
|
|
40
|
+
* await init({ powerPreference: "low-power" });
|
|
41
|
+
* const { description, device } = gpuName();
|
|
42
|
+
* console.log("description:", description, "device:", device);
|
|
43
|
+
* ```
|
|
44
|
+
*
|
|
45
|
+
* @example Benchmark mode
|
|
46
|
+
* ```js
|
|
47
|
+
* import { init, sscal } from "wgblas";
|
|
48
|
+
* const device = await init({ benchmark: true });
|
|
49
|
+
* const n = 5;
|
|
50
|
+
* const alpha = 2.0;
|
|
51
|
+
* const x = new Float32Array([1, 2, 3, 4, 5]);
|
|
52
|
+
* const { result, gpuTimeMs } = await sscal(device, n, alpha, x, 1);
|
|
53
|
+
* console.log(`Result: [${Array.from(result).join(", ")}]`);
|
|
54
|
+
* console.log(`GPU time: ${gpuTimeMs.toFixed(3)} ms`);
|
|
55
|
+
* ```
|
|
56
|
+
* @see [Source code: init.mjs](https://github.com/manit2004/wgblas/blob/main/src/init.mjs#L18-L54)
|
|
57
|
+
* @category Core
|
|
58
|
+
*/
|
|
59
|
+
export declare function init(options?: {
|
|
60
|
+
powerPreference?: GPUPowerPreference;
|
|
61
|
+
benchmark?: boolean;
|
|
62
|
+
}): Promise<GPUDevice>;
|
|
63
|
+
|
|
64
|
+
/**
|
|
65
|
+
* Destroys the WebGPU device, releases the adapter, resets benchmark state, and fires all internal
|
|
66
|
+
* cleanup callbacks (e.g. releasing cached GPU pipelines and buffers). Call when done (required in Node.js to prevent crash on exit).
|
|
67
|
+
*
|
|
68
|
+
* @example
|
|
69
|
+
* ```js
|
|
70
|
+
* import { init, cleanup, gpuName } from "wgblas";
|
|
71
|
+
*
|
|
72
|
+
* await init();
|
|
73
|
+
* console.log("GPU:", gpuName().description);
|
|
74
|
+
* if (typeof process !== "undefined") cleanup(); // Node.js only — browser cleanup is automatic
|
|
75
|
+
* ```
|
|
76
|
+
* @see [Source code: init.mjs](https://github.com/manit2004/wgblas/blob/main/src/init.mjs#L56-L65)
|
|
77
|
+
* @category Core
|
|
78
|
+
*/
|
|
79
|
+
export declare function cleanup(): void;
|
|
80
|
+
|
|
81
|
+
/**
|
|
82
|
+
* Returns the GPU device name from the WebGPU adapter info. Must be called after `init()`.
|
|
83
|
+
*
|
|
84
|
+
* @example
|
|
85
|
+
* ```js
|
|
86
|
+
* import { init, gpuName } from "wgblas";
|
|
87
|
+
* await init();
|
|
88
|
+
* const { description, device } = gpuName();
|
|
89
|
+
* console.log("description:", description, "device:", device);
|
|
90
|
+
* ```
|
|
91
|
+
* @see [Source code: init.mjs](https://github.com/manit2004/wgblas/blob/main/src/init.mjs#L81-L87)
|
|
92
|
+
* @category Core
|
|
93
|
+
*/
|
|
94
|
+
export declare function gpuName(): { description: string; device: string };
|
package/index.mjs
ADDED
|
@@ -0,0 +1,16 @@
|
|
|
1
|
+
export { init, cleanup, gpuName } from "./src/init.mjs";
|
|
2
|
+
export { GpuVector } from "./src/classes/GpuVector.mjs";
|
|
3
|
+
export {
|
|
4
|
+
randomFloat32Array,
|
|
5
|
+
randomFloat64Array,
|
|
6
|
+
} from "./src/random/random.mjs";
|
|
7
|
+
export { sscal } from "./src/sscal/sscal.mjs";
|
|
8
|
+
export { sswap } from "./src/sswap/sswap.mjs";
|
|
9
|
+
export { saxpy } from "./src/saxpy/saxpy.mjs";
|
|
10
|
+
export { scopy } from "./src/scopy/scopy.mjs";
|
|
11
|
+
export { sdot } from "./src/sdot/sdot.mjs";
|
|
12
|
+
export { sasum } from "./src/sasum/sasum.mjs";
|
|
13
|
+
export { snrm2 } from "./src/snrm2/snrm2.mjs";
|
|
14
|
+
export { isamax } from "./src/isamax/isamax.mjs";
|
|
15
|
+
export { srot } from "./src/srot/srot.mjs";
|
|
16
|
+
export { srotm } from "./src/srotm/srotm.mjs";
|