wgblas 0.1.2 → 1.1.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +3 -0
- package/dist/wgblas.browser.js +1249 -37
- package/index.d.mts +9 -0
- package/index.mjs +9 -0
- package/package.json +47 -1
- package/src/classes/GpuMatrix.d.mts +98 -0
- package/src/classes/GpuMatrix.mjs +109 -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 +127 -0
- package/src/sgemv/sgemv.mjs +148 -0
- package/src/sger/sger.d.mts +111 -0
- package/src/sger/sger.mjs +136 -0
- package/src/shaders/browser-shaders.mjs +26 -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/sger.wgsl +48 -0
- package/src/shaders/snrm2.wgsl +22 -4
- package/src/shaders/ssymv.wgsl +69 -0
- package/src/shaders/ssyr.wgsl +60 -0
- package/src/shaders/ssyr2.wgsl +63 -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 +117 -0
- package/src/ssymv/ssymv.mjs +135 -0
- package/src/ssyr/ssyr.d.mts +100 -0
- package/src/ssyr/ssyr.mjs +106 -0
- package/src/ssyr2/ssyr2.d.mts +112 -0
- package/src/ssyr2/ssyr2.mjs +130 -0
- package/src/strmv/strmv.d.mts +117 -0
- package/src/strmv/strmv.mjs +138 -0
- package/src/strsv/strsv.d.mts +106 -0
- package/src/strsv/strsv.mjs +207 -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
|
@@ -0,0 +1,69 @@
|
|
|
1
|
+
// ssymv: y = alpha * A * x + beta * y
|
|
2
|
+
// A is n×n symmetric, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
3
|
+
// The logical matrix is fully dense (symmetric), so each row's dot product
|
|
4
|
+
// sums over all n columns; entries on the unstored side of the diagonal are
|
|
5
|
+
// fetched from their mirror position (A[i,j] == A[j,i]).
|
|
6
|
+
// One workgroup per row, grid-stride outer loop.
|
|
7
|
+
|
|
8
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
9
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
10
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
11
|
+
|
|
12
|
+
struct Params {
|
|
13
|
+
n: u32,
|
|
14
|
+
alpha: f32,
|
|
15
|
+
beta: f32,
|
|
16
|
+
incx: u32,
|
|
17
|
+
incy: u32,
|
|
18
|
+
lda: u32,
|
|
19
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
20
|
+
}
|
|
21
|
+
|
|
22
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
23
|
+
|
|
24
|
+
const WGS: u32 = 64u;
|
|
25
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
26
|
+
|
|
27
|
+
@compute @workgroup_size(64)
|
|
28
|
+
fn main(
|
|
29
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
30
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
31
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
32
|
+
) {
|
|
33
|
+
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
34
|
+
var acc = 0.0f;
|
|
35
|
+
|
|
36
|
+
// y[i] = Σ_j A[i,j] * x[j]
|
|
37
|
+
for (var j = lid.x; j < params.n; j += WGS) {
|
|
38
|
+
var aVal: f32;
|
|
39
|
+
if params.uplo == 0u {
|
|
40
|
+
// Lower: A[i,j] stored at A[i*lda+j] for j ≤ i, mirrored from A[j*lda+i] otherwise
|
|
41
|
+
if j <= i {
|
|
42
|
+
aVal = A[i * params.lda + j];
|
|
43
|
+
} else {
|
|
44
|
+
aVal = A[j * params.lda + i];
|
|
45
|
+
}
|
|
46
|
+
} else {
|
|
47
|
+
// Upper: A[i,j] stored at A[i*lda+j] for j ≥ i, mirrored from A[j*lda+i] otherwise
|
|
48
|
+
if j >= i {
|
|
49
|
+
aVal = A[i * params.lda + j];
|
|
50
|
+
} else {
|
|
51
|
+
aVal = A[j * params.lda + i];
|
|
52
|
+
}
|
|
53
|
+
}
|
|
54
|
+
acc += aVal * x[j * params.incx];
|
|
55
|
+
}
|
|
56
|
+
|
|
57
|
+
// Parallel reduction: 64 → 1
|
|
58
|
+
scratch[lid.x] = acc;
|
|
59
|
+
workgroupBarrier();
|
|
60
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
61
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
62
|
+
workgroupBarrier();
|
|
63
|
+
}
|
|
64
|
+
|
|
65
|
+
if lid.x == 0u {
|
|
66
|
+
y[i * params.incy] = params.alpha * scratch[0] + params.beta * y[i * params.incy];
|
|
67
|
+
}
|
|
68
|
+
}
|
|
69
|
+
}
|
|
@@ -0,0 +1,60 @@
|
|
|
1
|
+
// ssyr: A := alpha * x * x^T + A (symmetric rank-1 update)
|
|
2
|
+
// A is n×n symmetric; only the triangle specified by uplo is referenced/updated,
|
|
3
|
+
// the other triangle is implied by symmetry (not touched).
|
|
4
|
+
|
|
5
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
6
|
+
@group(0) @binding(1) var<storage, read_write> A: array<f32>;
|
|
7
|
+
|
|
8
|
+
struct Params {
|
|
9
|
+
n: u32,
|
|
10
|
+
alpha: f32,
|
|
11
|
+
incx: u32,
|
|
12
|
+
lda: u32,
|
|
13
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
14
|
+
}
|
|
15
|
+
|
|
16
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
17
|
+
|
|
18
|
+
const WGS: u32 = 64u;
|
|
19
|
+
|
|
20
|
+
@compute @workgroup_size(64)
|
|
21
|
+
fn main(
|
|
22
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
23
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
24
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
25
|
+
) {
|
|
26
|
+
for (var row = wgid.x; row < params.n; row += nwg.x) {
|
|
27
|
+
let xi = params.alpha * x[row * params.incx];
|
|
28
|
+
let row_base = row * params.lda;
|
|
29
|
+
|
|
30
|
+
// Stored-triangle column range for this row: lower [0,row], upper [row,n).
|
|
31
|
+
var colStart: u32;
|
|
32
|
+
var colEnd: u32;
|
|
33
|
+
if params.uplo == 1u {
|
|
34
|
+
colStart = row;
|
|
35
|
+
colEnd = params.n;
|
|
36
|
+
} else {
|
|
37
|
+
colStart = 0u;
|
|
38
|
+
colEnd = row + 1u;
|
|
39
|
+
}
|
|
40
|
+
|
|
41
|
+
// 4-unrolled loop over the stored range.
|
|
42
|
+
let rangeLen = colEnd - colStart;
|
|
43
|
+
let n4_floor = colStart + (rangeLen / (4u * WGS)) * (4u * WGS);
|
|
44
|
+
for (var col: u32 = colStart + lid.x; col < n4_floor; col += 4u * WGS) {
|
|
45
|
+
let idx0 = row_base + col;
|
|
46
|
+
let idx1 = row_base + col + WGS;
|
|
47
|
+
let idx2 = row_base + col + 2u * WGS;
|
|
48
|
+
let idx3 = row_base + col + 3u * WGS;
|
|
49
|
+
A[idx0] = xi * x[ col * params.incx] + A[idx0];
|
|
50
|
+
A[idx1] = xi * x[(col + WGS) * params.incx] + A[idx1];
|
|
51
|
+
A[idx2] = xi * x[(col + 2u * WGS) * params.incx] + A[idx2];
|
|
52
|
+
A[idx3] = xi * x[(col + 3u * WGS) * params.incx] + A[idx3];
|
|
53
|
+
}
|
|
54
|
+
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
55
|
+
for (var col: u32 = n4_floor + lid.x; col < colEnd; col += WGS) {
|
|
56
|
+
let idx = row_base + col;
|
|
57
|
+
A[idx] = xi * x[col * params.incx] + A[idx];
|
|
58
|
+
}
|
|
59
|
+
}
|
|
60
|
+
}
|
|
@@ -0,0 +1,63 @@
|
|
|
1
|
+
// ssyr2: A := alpha * x * y^T + alpha * y * x^T + A (symmetric rank-2 update)
|
|
2
|
+
// A is n×n symmetric; only the triangle specified by uplo is referenced/updated,
|
|
3
|
+
// the other triangle is implied by symmetry (not touched).
|
|
4
|
+
|
|
5
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
6
|
+
@group(0) @binding(1) var<storage, read> y: array<f32>;
|
|
7
|
+
@group(0) @binding(2) var<storage, read_write> A: array<f32>;
|
|
8
|
+
|
|
9
|
+
struct Params {
|
|
10
|
+
n: u32,
|
|
11
|
+
alpha: f32,
|
|
12
|
+
incx: u32,
|
|
13
|
+
incy: u32,
|
|
14
|
+
lda: u32,
|
|
15
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
16
|
+
}
|
|
17
|
+
|
|
18
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
19
|
+
|
|
20
|
+
const WGS: u32 = 64u;
|
|
21
|
+
|
|
22
|
+
@compute @workgroup_size(64)
|
|
23
|
+
fn main(
|
|
24
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
25
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
26
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
27
|
+
) {
|
|
28
|
+
for (var row = wgid.x; row < params.n; row += nwg.x) {
|
|
29
|
+
let xi = params.alpha * x[row * params.incx];
|
|
30
|
+
let yi = params.alpha * y[row * params.incy];
|
|
31
|
+
let row_base = row * params.lda;
|
|
32
|
+
|
|
33
|
+
// Stored-triangle column range for this row: lower [0,row], upper [row,n).
|
|
34
|
+
var colStart: u32;
|
|
35
|
+
var colEnd: u32;
|
|
36
|
+
if params.uplo == 1u {
|
|
37
|
+
colStart = row;
|
|
38
|
+
colEnd = params.n;
|
|
39
|
+
} else {
|
|
40
|
+
colStart = 0u;
|
|
41
|
+
colEnd = row + 1u;
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
// 4-unrolled loop over the stored range.
|
|
45
|
+
let rangeLen = colEnd - colStart;
|
|
46
|
+
let n4_floor = colStart + (rangeLen / (4u * WGS)) * (4u * WGS);
|
|
47
|
+
for (var col: u32 = colStart + lid.x; col < n4_floor; col += 4u * WGS) {
|
|
48
|
+
let idx0 = row_base + col;
|
|
49
|
+
let idx1 = row_base + col + WGS;
|
|
50
|
+
let idx2 = row_base + col + 2u * WGS;
|
|
51
|
+
let idx3 = row_base + col + 3u * WGS;
|
|
52
|
+
A[idx0] = xi * y[ col * params.incy] + yi * x[ col * params.incx] + A[idx0];
|
|
53
|
+
A[idx1] = xi * y[(col + WGS) * params.incy] + yi * x[(col + WGS) * params.incx] + A[idx1];
|
|
54
|
+
A[idx2] = xi * y[(col + 2u * WGS) * params.incy] + yi * x[(col + 2u * WGS) * params.incx] + A[idx2];
|
|
55
|
+
A[idx3] = xi * y[(col + 3u * WGS) * params.incy] + yi * x[(col + 3u * WGS) * params.incx] + A[idx3];
|
|
56
|
+
}
|
|
57
|
+
// Scalar tail: at most 3*WGS elements left after the unrolled block.
|
|
58
|
+
for (var col: u32 = n4_floor + lid.x; col < colEnd; col += WGS) {
|
|
59
|
+
let idx = row_base + col;
|
|
60
|
+
A[idx] = xi * y[col * params.incy] + yi * x[col * params.incx] + A[idx];
|
|
61
|
+
}
|
|
62
|
+
}
|
|
63
|
+
}
|
|
@@ -0,0 +1,103 @@
|
|
|
1
|
+
// strmv: y = op(A) * x
|
|
2
|
+
// A is n×n triangular, lower (uplo=0) or upper (uplo=1) triangle stored.
|
|
3
|
+
// op(A) is A (trans=0) or A^T (trans=1).
|
|
4
|
+
// diag=1 (unit) treats the diagonal as 1 without reading A's diagonal values.
|
|
5
|
+
// One workgroup per row, grid-stride outer loop.
|
|
6
|
+
|
|
7
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
8
|
+
@group(0) @binding(1) var<storage, read> x: array<f32>;
|
|
9
|
+
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
|
|
10
|
+
|
|
11
|
+
struct Params {
|
|
12
|
+
n: u32,
|
|
13
|
+
incx: u32,
|
|
14
|
+
incy: u32,
|
|
15
|
+
lda: u32,
|
|
16
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
17
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
18
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
19
|
+
}
|
|
20
|
+
|
|
21
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
22
|
+
|
|
23
|
+
const WGS: u32 = 64u;
|
|
24
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
25
|
+
|
|
26
|
+
@compute @workgroup_size(64)
|
|
27
|
+
fn main(
|
|
28
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
29
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
30
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
31
|
+
) {
|
|
32
|
+
for (var i = wgid.x; i < params.n; i += nwg.x) {
|
|
33
|
+
var acc = 0.0f;
|
|
34
|
+
|
|
35
|
+
if params.trans == 0u {
|
|
36
|
+
// No-transpose: y[i] = Σ_j A[i,j] * x[j]
|
|
37
|
+
if params.uplo == 0u {
|
|
38
|
+
// Lower: A[i,j] stored at A[i*lda+j] for j ≤ i
|
|
39
|
+
for (var j = lid.x; j <= i; j += WGS) {
|
|
40
|
+
var aVal: f32;
|
|
41
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
42
|
+
if params.diag == 1u && j == i {
|
|
43
|
+
aVal = 1.0;
|
|
44
|
+
} else if ( j <= i ) {
|
|
45
|
+
aVal = A[i * params.lda + j];
|
|
46
|
+
}
|
|
47
|
+
acc += aVal * x[j * params.incx];
|
|
48
|
+
}
|
|
49
|
+
} else {
|
|
50
|
+
// Upper: A[i,j] stored at A[i*lda+j] for j ≥ i
|
|
51
|
+
for (var j = i + lid.x; j < params.n; j += WGS) {
|
|
52
|
+
var aVal: f32;
|
|
53
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
54
|
+
if params.diag == 1u && j == i {
|
|
55
|
+
aVal = 1.0;
|
|
56
|
+
} else if ( j >= i ) {
|
|
57
|
+
aVal = A[i * params.lda + j];
|
|
58
|
+
}
|
|
59
|
+
acc += aVal * x[j * params.incx];
|
|
60
|
+
}
|
|
61
|
+
}
|
|
62
|
+
} else {
|
|
63
|
+
// Transpose: y[i] = Σ_j A[j,i] * x[j]
|
|
64
|
+
if params.uplo == 0u {
|
|
65
|
+
// Lower: A[j,i] stored at A[j*lda+i] for j ≥ i
|
|
66
|
+
for (var j = i + lid.x; j < params.n; j += WGS) {
|
|
67
|
+
var aVal: f32;
|
|
68
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
69
|
+
if params.diag == 1u && j == i {
|
|
70
|
+
aVal = 1.0;
|
|
71
|
+
} else if ( j >= i ) {
|
|
72
|
+
aVal = A[j * params.lda + i];
|
|
73
|
+
}
|
|
74
|
+
acc += aVal * x[j * params.incx];
|
|
75
|
+
}
|
|
76
|
+
} else {
|
|
77
|
+
// Upper: A[j,i] stored at A[j*lda+i] for j ≤ i
|
|
78
|
+
for (var j = lid.x; j <= i; j += WGS) {
|
|
79
|
+
var aVal: f32;
|
|
80
|
+
// unit diagonal: use 1 instead of A's actual diagonal value
|
|
81
|
+
if params.diag == 1u && j == i {
|
|
82
|
+
aVal = 1.0;
|
|
83
|
+
} else if ( j <= i ) {
|
|
84
|
+
aVal = A[j * params.lda + i];
|
|
85
|
+
}
|
|
86
|
+
acc += aVal * x[j * params.incx];
|
|
87
|
+
}
|
|
88
|
+
}
|
|
89
|
+
}
|
|
90
|
+
|
|
91
|
+
// Parallel reduction: 64 → 1
|
|
92
|
+
scratch[lid.x] = acc;
|
|
93
|
+
workgroupBarrier();
|
|
94
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
95
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
96
|
+
workgroupBarrier();
|
|
97
|
+
}
|
|
98
|
+
|
|
99
|
+
if lid.x == 0u {
|
|
100
|
+
y[ i * params.incy ] = scratch[0];
|
|
101
|
+
}
|
|
102
|
+
}
|
|
103
|
+
}
|
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
// strsv_apply_inverse: given a precomputed block inverse (from
|
|
2
|
+
// strsv_invert_block.wgsl), computes this block's solution as a dense
|
|
3
|
+
// matrix-vector multiply against the block's current remainder in x —
|
|
4
|
+
// replacing what the old strsv_block.wgsl did via a genuinely sequential,
|
|
5
|
+
// barrier-per-row substitution.
|
|
6
|
+
//
|
|
7
|
+
// All blockLen rows are computed in parallel within a single workgroup: the
|
|
8
|
+
// remainder is loaded into workgroup-shared memory once, then each thread
|
|
9
|
+
// independently computes one full row's dot product from that shared copy.
|
|
10
|
+
// No further synchronization is needed after the load — every thread only
|
|
11
|
+
// reads shared memory from then on (never written again within this call)
|
|
12
|
+
// and writes a distinct element of x, so there's no cross-thread hazard to
|
|
13
|
+
// guard against.
|
|
14
|
+
|
|
15
|
+
@group(0) @binding(0) var<storage, read> Ainv: array<f32>;
|
|
16
|
+
@group(0) @binding(1) var<storage, read_write> x: array<f32>;
|
|
17
|
+
|
|
18
|
+
struct Params {
|
|
19
|
+
incx: u32,
|
|
20
|
+
blockIndex: u32,
|
|
21
|
+
blockStart: u32,
|
|
22
|
+
blockEnd: u32,
|
|
23
|
+
}
|
|
24
|
+
|
|
25
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
26
|
+
|
|
27
|
+
const BLOCK_SIZE: u32 = 64u;
|
|
28
|
+
var<workgroup> xLocal: array<f32, 64>;
|
|
29
|
+
|
|
30
|
+
@compute @workgroup_size(64)
|
|
31
|
+
fn strsv_apply_inverse_main(@builtin(local_invocation_id) lid: vec3u) {
|
|
32
|
+
let blockLen = params.blockEnd - params.blockStart;
|
|
33
|
+
|
|
34
|
+
if (lid.x < blockLen) {
|
|
35
|
+
xLocal[lid.x] = x[(params.blockStart + lid.x) * params.incx];
|
|
36
|
+
}
|
|
37
|
+
workgroupBarrier();
|
|
38
|
+
|
|
39
|
+
if (lid.x >= blockLen) { return; }
|
|
40
|
+
|
|
41
|
+
let ainvBase = params.blockIndex * BLOCK_SIZE * BLOCK_SIZE;
|
|
42
|
+
var acc = 0.0f;
|
|
43
|
+
for (var j = 0u; j < blockLen; j++) {
|
|
44
|
+
acc += Ainv[ainvBase + lid.x * BLOCK_SIZE + j] * xLocal[j];
|
|
45
|
+
}
|
|
46
|
+
x[(params.blockStart + lid.x) * params.incx] = acc;
|
|
47
|
+
}
|
|
@@ -0,0 +1,109 @@
|
|
|
1
|
+
// strsv_invert_block: computes ONE column (workgroup_id.x) of ONE block's
|
|
2
|
+
// (workgroup_id.y) explicit inverse, via the same one-row-at-a-time
|
|
3
|
+
// substitution as strsv_block.wgsl, but solving against a unit basis vector
|
|
4
|
+
// e_col instead of the real right-hand side, and writing to a dense
|
|
5
|
+
// (BLOCK_SIZE x BLOCK_SIZE, row-major) scratch buffer per block instead of
|
|
6
|
+
// mutating x. Dispatched once for the whole matrix (2D: BLOCK_SIZE columns x
|
|
7
|
+
// numBlocks), fully in parallel -- unlike the sequential per-block main
|
|
8
|
+
// loop in strsv.mjs, no block's inverse depends on any other block or on x.
|
|
9
|
+
//
|
|
10
|
+
// A triangular block's inverse is itself triangular: forward (effectively-
|
|
11
|
+
// lower, e.g. no-trans+lower) blocks have inverse column col nonzero only
|
|
12
|
+
// for row>=col, solved in increasing row order; backward (effectively-
|
|
13
|
+
// upper) blocks have it nonzero only for row<=col, solved in decreasing
|
|
14
|
+
// order. Rows outside a column's nonzero range are written as literal 0 —
|
|
15
|
+
// strsv_apply_inverse.wgsl's dense matvec depends on that, not just on
|
|
16
|
+
// those entries being mathematically implied zero.
|
|
17
|
+
|
|
18
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
19
|
+
@group(0) @binding(1) var<storage, read_write> Ainv: array<f32>;
|
|
20
|
+
|
|
21
|
+
struct Params {
|
|
22
|
+
n: u32,
|
|
23
|
+
lda: u32,
|
|
24
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
25
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
26
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
27
|
+
}
|
|
28
|
+
|
|
29
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
30
|
+
|
|
31
|
+
const WGS: u32 = 64u;
|
|
32
|
+
const BLOCK_SIZE: u32 = 64u;
|
|
33
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
34
|
+
|
|
35
|
+
fn readA(i: u32, j: u32) -> f32 {
|
|
36
|
+
if params.trans == 0u {
|
|
37
|
+
return A[i * params.lda + j];
|
|
38
|
+
} else {
|
|
39
|
+
return A[j * params.lda + i];
|
|
40
|
+
}
|
|
41
|
+
}
|
|
42
|
+
|
|
43
|
+
@compute @workgroup_size(64)
|
|
44
|
+
fn strsv_invert_block_main(
|
|
45
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
46
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
47
|
+
) {
|
|
48
|
+
let col = wgid.x;
|
|
49
|
+
let blockIndex = wgid.y;
|
|
50
|
+
let blockStart = blockIndex * BLOCK_SIZE;
|
|
51
|
+
var blockEnd = blockStart + BLOCK_SIZE;
|
|
52
|
+
if (blockEnd > params.n) { blockEnd = params.n; }
|
|
53
|
+
let blockLen = blockEnd - blockStart;
|
|
54
|
+
|
|
55
|
+
if (col >= blockLen) { return; }
|
|
56
|
+
|
|
57
|
+
let ainvBase = blockIndex * BLOCK_SIZE * BLOCK_SIZE;
|
|
58
|
+
let forward = (params.trans == 0u) == (params.uplo == 0u);
|
|
59
|
+
|
|
60
|
+
if forward {
|
|
61
|
+
for (var r = lid.x; r < col; r += WGS) {
|
|
62
|
+
Ainv[ainvBase + r * BLOCK_SIZE + col] = 0.0;
|
|
63
|
+
}
|
|
64
|
+
} else {
|
|
65
|
+
for (var r = col + 1u + lid.x; r < blockLen; r += WGS) {
|
|
66
|
+
Ainv[ainvBase + r * BLOCK_SIZE + col] = 0.0;
|
|
67
|
+
}
|
|
68
|
+
}
|
|
69
|
+
storageBarrier();
|
|
70
|
+
workgroupBarrier();
|
|
71
|
+
|
|
72
|
+
let numSteps = select(col + 1u, blockLen - col, forward);
|
|
73
|
+
for (var step = 0u; step < numSteps; step++) {
|
|
74
|
+
let localRow = select(col - step, col + step, forward);
|
|
75
|
+
let i = blockStart + localRow;
|
|
76
|
+
|
|
77
|
+
var acc = 0.0f;
|
|
78
|
+
if forward {
|
|
79
|
+
for (var lj = col + lid.x; lj < localRow; lj += WGS) {
|
|
80
|
+
acc += readA(i, blockStart + lj) * Ainv[ainvBase + lj * BLOCK_SIZE + col];
|
|
81
|
+
}
|
|
82
|
+
} else {
|
|
83
|
+
for (var lj = localRow + 1u + lid.x; lj <= col; lj += WGS) {
|
|
84
|
+
acc += readA(i, blockStart + lj) * Ainv[ainvBase + lj * BLOCK_SIZE + col];
|
|
85
|
+
}
|
|
86
|
+
}
|
|
87
|
+
|
|
88
|
+
scratch[lid.x] = acc;
|
|
89
|
+
workgroupBarrier();
|
|
90
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
91
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
92
|
+
workgroupBarrier();
|
|
93
|
+
}
|
|
94
|
+
|
|
95
|
+
if lid.x == 0u {
|
|
96
|
+
let e = select(0.0, 1.0, localRow == col);
|
|
97
|
+
let rhs = e - scratch[0];
|
|
98
|
+
var val: f32;
|
|
99
|
+
if params.diag == 1u {
|
|
100
|
+
val = rhs;
|
|
101
|
+
} else {
|
|
102
|
+
val = rhs / A[i * params.lda + i];
|
|
103
|
+
}
|
|
104
|
+
Ainv[ainvBase + localRow * BLOCK_SIZE + col] = val;
|
|
105
|
+
}
|
|
106
|
+
storageBarrier();
|
|
107
|
+
workgroupBarrier();
|
|
108
|
+
}
|
|
109
|
+
}
|
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
// strsv_update: subtracts a solved block's contribution from every
|
|
2
|
+
// remaining row in parallel (one workgroup per row, like strmv.wgsl) —
|
|
3
|
+
// this is what turns strsv's O(n) sequential stages into O(n/blockSize).
|
|
4
|
+
// No diag/masking needed: this region never touches the diagonal.
|
|
5
|
+
|
|
6
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
7
|
+
@group(0) @binding(1) var<storage, read_write> x: array<f32>;
|
|
8
|
+
|
|
9
|
+
struct Params {
|
|
10
|
+
n: u32,
|
|
11
|
+
incx: u32,
|
|
12
|
+
lda: u32,
|
|
13
|
+
trans: u32, // 0 = no-transpose, 1 = transpose
|
|
14
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
15
|
+
blockStart: u32,
|
|
16
|
+
blockEnd: u32, // exclusive
|
|
17
|
+
}
|
|
18
|
+
|
|
19
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
20
|
+
|
|
21
|
+
const WGS: u32 = 64u;
|
|
22
|
+
var<workgroup> scratch: array<f32, 64>;
|
|
23
|
+
|
|
24
|
+
@compute @workgroup_size(64)
|
|
25
|
+
fn strsv_update_main(
|
|
26
|
+
@builtin(workgroup_id) wgid: vec3u,
|
|
27
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
28
|
+
@builtin(num_workgroups) nwg: vec3u,
|
|
29
|
+
) {
|
|
30
|
+
// forward: remaining rows are [blockEnd,n); backward: [0,blockStart).
|
|
31
|
+
let forward = (params.trans == 0u) == (params.uplo == 0u);
|
|
32
|
+
|
|
33
|
+
var rangeStart: u32;
|
|
34
|
+
var rangeEnd: u32;
|
|
35
|
+
|
|
36
|
+
if forward {
|
|
37
|
+
rangeStart = params.blockEnd;
|
|
38
|
+
rangeEnd = params.n;
|
|
39
|
+
} else {
|
|
40
|
+
rangeStart = 0u;
|
|
41
|
+
rangeEnd = params.blockStart;
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
if (rangeStart >= rangeEnd) { return; }
|
|
45
|
+
let count = rangeEnd - rangeStart;
|
|
46
|
+
|
|
47
|
+
for (var idx = wgid.x; idx < count; idx += nwg.x) {
|
|
48
|
+
let i = rangeStart + idx;
|
|
49
|
+
|
|
50
|
+
// No-trans reads A[i,j]; transpose reads A[j,i] — uplo only sets the range above.
|
|
51
|
+
var acc = 0.0f;
|
|
52
|
+
if params.trans == 0u {
|
|
53
|
+
for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
|
|
54
|
+
acc += A[i * params.lda + j] * x[j * params.incx];
|
|
55
|
+
}
|
|
56
|
+
} else {
|
|
57
|
+
for (var j = params.blockStart + lid.x; j < params.blockEnd; j += WGS) {
|
|
58
|
+
acc += A[j * params.lda + i] * x[j * params.incx];
|
|
59
|
+
}
|
|
60
|
+
}
|
|
61
|
+
|
|
62
|
+
// Parallel reduction: 64 → 1
|
|
63
|
+
scratch[lid.x] = acc;
|
|
64
|
+
workgroupBarrier();
|
|
65
|
+
for (var stride = WGS >> 1u; stride > 0u; stride >>= 1u) {
|
|
66
|
+
if lid.x < stride { scratch[lid.x] += scratch[lid.x + stride]; }
|
|
67
|
+
workgroupBarrier();
|
|
68
|
+
}
|
|
69
|
+
|
|
70
|
+
if lid.x == 0u {
|
|
71
|
+
x[i * params.incx] -= scratch[0];
|
|
72
|
+
}
|
|
73
|
+
workgroupBarrier();
|
|
74
|
+
}
|
|
75
|
+
}
|
package/src/snrm2/snrm2.mjs
CHANGED
|
@@ -34,67 +34,71 @@ export async function snrm2(device, n, x, incx) {
|
|
|
34
34
|
const pipelineMain = await getPipeline(device, "snrm2");
|
|
35
35
|
const pipelineReduce = await getPipeline(device, "reduction/sum");
|
|
36
36
|
|
|
37
|
-
|
|
38
|
-
|
|
39
|
-
|
|
40
|
-
|
|
41
|
-
|
|
42
|
-
{ value: n, type: "u32" },
|
|
43
|
-
{ value: incx, type: "u32" },
|
|
44
|
-
],
|
|
45
|
-
"snrm2-params",
|
|
46
|
-
);
|
|
37
|
+
let xBuffer = null;
|
|
38
|
+
let partialsBuffer = null;
|
|
39
|
+
let resultBuffer = null;
|
|
40
|
+
let paramsBuffer = null;
|
|
41
|
+
let readBuffer = null;
|
|
47
42
|
|
|
48
|
-
|
|
49
|
-
xBuffer,
|
|
50
|
-
partialsBuffer,
|
|
51
|
-
|
|
52
|
-
|
|
53
|
-
|
|
54
|
-
|
|
55
|
-
|
|
56
|
-
|
|
57
|
-
|
|
43
|
+
try {
|
|
44
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "snrm2-x", false);
|
|
45
|
+
partialsBuffer = createStorageBuffer(2 * WGS * 4, "snrm2-partials"); // 2*WGS partial sums of f32
|
|
46
|
+
resultBuffer = createResultBuffer(4, "snrm2-result"); // final f32 scalar
|
|
47
|
+
paramsBuffer = createParamsBuffer(
|
|
48
|
+
[
|
|
49
|
+
{ value: n, type: "u32" },
|
|
50
|
+
{ value: incx, type: "u32" },
|
|
51
|
+
],
|
|
52
|
+
"snrm2-params",
|
|
53
|
+
);
|
|
54
|
+
|
|
55
|
+
const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
|
|
56
|
+
xBuffer,
|
|
57
|
+
partialsBuffer,
|
|
58
|
+
paramsBuffer,
|
|
59
|
+
]);
|
|
60
|
+
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
61
|
+
pipelineMain,
|
|
62
|
+
bgMain,
|
|
63
|
+
2 * WGS,
|
|
64
|
+
); // dispatch 2*WGS workgroups
|
|
58
65
|
|
|
59
|
-
|
|
66
|
+
submit(enc1);
|
|
60
67
|
|
|
61
|
-
|
|
62
|
-
|
|
63
|
-
|
|
64
|
-
|
|
65
|
-
|
|
66
|
-
|
|
67
|
-
|
|
68
|
-
|
|
69
|
-
|
|
70
|
-
|
|
68
|
+
const bgReduce = createBindGroup(pipelineReduce.getBindGroupLayout(0), [
|
|
69
|
+
partialsBuffer,
|
|
70
|
+
resultBuffer,
|
|
71
|
+
]);
|
|
72
|
+
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
73
|
+
pipelineReduce,
|
|
74
|
+
bgReduce,
|
|
75
|
+
1,
|
|
76
|
+
); // reduce partials to single result
|
|
77
|
+
readBuffer = stageReadback(enc2, resultBuffer);
|
|
71
78
|
|
|
72
|
-
|
|
79
|
+
submit(enc2);
|
|
73
80
|
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
extractTimestamp(ts2),
|
|
77
|
-
extractResult(readBuffer, Float32Array),
|
|
78
|
-
]);
|
|
81
|
+
const resultPromise = extractResult(readBuffer, Float32Array);
|
|
82
|
+
readBuffer = null; // ownership transferred — extractResult's own finally destroys it
|
|
79
83
|
|
|
80
|
-
|
|
81
|
-
|
|
84
|
+
const [gpuTime1, gpuTime2, sqsumArr] = await Promise.all([
|
|
85
|
+
extractTimestamp(ts1),
|
|
86
|
+
extractTimestamp(ts2),
|
|
87
|
+
resultPromise,
|
|
88
|
+
]);
|
|
89
|
+
|
|
90
|
+
// sqrt is taken on CPU after the GPU sum-of-squares reduction
|
|
91
|
+
const nrm2 = Math.sqrt(sqsumArr[0]);
|
|
82
92
|
|
|
83
|
-
if (xIsGpu) {
|
|
84
|
-
destroyBuffers(partialsBuffer, resultBuffer, paramsBuffer, readBuffer);
|
|
85
93
|
if (gpuTime1 !== undefined && gpuTime2 !== undefined)
|
|
86
94
|
return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
|
|
87
95
|
return { nrm2 };
|
|
96
|
+
} finally {
|
|
97
|
+
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
98
|
+
if (partialsBuffer) destroyBuffers(partialsBuffer);
|
|
99
|
+
if (resultBuffer) destroyBuffers(resultBuffer);
|
|
100
|
+
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
101
|
+
// Only reached if submit(enc2) threw before ownership was transferred above.
|
|
102
|
+
if (readBuffer) destroyBuffers(readBuffer);
|
|
88
103
|
}
|
|
89
|
-
|
|
90
|
-
destroyBuffers(
|
|
91
|
-
xBuffer,
|
|
92
|
-
partialsBuffer,
|
|
93
|
-
resultBuffer,
|
|
94
|
-
paramsBuffer,
|
|
95
|
-
readBuffer,
|
|
96
|
-
);
|
|
97
|
-
if (gpuTime1 !== undefined && gpuTime2 !== undefined)
|
|
98
|
-
return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
|
|
99
|
-
return { nrm2 };
|
|
100
104
|
}
|