wgblas 1.2.1 → 2.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/LICENSE +1 -1
- package/README.md +36 -54
- package/dist/wgblas.browser.js +779 -34
- package/index.d.mts +13 -0
- package/index.mjs +8 -0
- package/package.json +55 -2
- package/src/dasum/dasum.d.mts +2 -2
- package/src/dasum/dasum.mjs +6 -4
- package/src/idamax/idamax.d.mts +51 -0
- package/src/idamax/idamax.mjs +128 -0
- package/src/init.mjs +9 -1
- package/src/isamax/isamax.d.mts +1 -1
- package/src/sasum/sasum.d.mts +1 -1
- package/src/saxpy/saxpy.d.mts +1 -1
- package/src/scopy/scopy.d.mts +1 -1
- package/src/sdot/sdot.d.mts +1 -1
- package/src/sgemm/sgemm.d.mts +102 -0
- package/src/sgemm/sgemm.mjs +195 -0
- package/src/sgemmtr/sgemmtr.d.mts +104 -0
- package/src/sgemmtr/sgemmtr.mjs +203 -0
- package/src/sgemv/sgemv.d.mts +1 -38
- package/src/sgemv/sgemv.mjs +4 -0
- package/src/sger/sger.d.mts +1 -34
- package/src/sger/sger.mjs +2 -0
- package/src/shaders/block_transfer.wgsl +42 -0
- package/src/shaders/browser-shaders.mjs +26 -0
- package/src/shaders/dasum.wgsl +3 -2
- package/src/shaders/f64/dekker.wgsl +4 -85
- package/src/shaders/f64/utils/abs.wgsl +10 -0
- package/src/shaders/f64/utils/add.wgsl +77 -0
- package/src/shaders/f64/utils/equal.wgsl +7 -0
- package/src/shaders/f64/utils/greater.wgsl +12 -0
- package/src/shaders/f64/utils/multiply.wgsl +81 -0
- package/src/shaders/idamax.wgsl +96 -0
- package/src/shaders/reduction/argmaxF64.wgsl +50 -0
- package/src/shaders/reduction/sumF64.wgsl +2 -2
- package/src/shaders/sgemm_large.wgsl +117 -0
- package/src/shaders/sgemm_small.wgsl +112 -0
- package/src/shaders/sgemmtr_large.wgsl +117 -0
- package/src/shaders/sgemmtr_small.wgsl +110 -0
- package/src/shaders/symmetrize.wgsl +31 -0
- package/src/shaders/triangularize.wgsl +44 -0
- package/src/snrm2/snrm2.d.mts +1 -1
- package/src/srot/srot.d.mts +1 -1
- package/src/srotm/srotm.d.mts +1 -1
- package/src/sscal/sscal.d.mts +1 -1
- package/src/sswap/sswap.d.mts +1 -1
- package/src/ssymm/ssymm.d.mts +103 -0
- package/src/ssymm/ssymm.mjs +209 -0
- package/src/ssymv/ssymv.d.mts +1 -36
- package/src/ssymv/ssymv.mjs +2 -0
- package/src/ssyr/ssyr.d.mts +1 -30
- package/src/ssyr/ssyr.mjs +2 -0
- package/src/ssyr2/ssyr2.d.mts +1 -34
- package/src/ssyr2/ssyr2.mjs +2 -0
- package/src/ssyr2k/ssyr2k.d.mts +100 -0
- package/src/ssyr2k/ssyr2k.mjs +201 -0
- package/src/ssyrk/ssyrk.d.mts +90 -0
- package/src/ssyrk/ssyrk.mjs +176 -0
- package/src/strmm/strmm.d.mts +100 -0
- package/src/strmm/strmm.mjs +211 -0
- package/src/strmv/strmv.d.mts +1 -36
- package/src/strmv/strmv.mjs +2 -0
- package/src/strsm/strsm.d.mts +99 -0
- package/src/strsm/strsm.mjs +342 -0
- package/src/strsv/strsv.d.mts +1 -32
- package/src/strsv/strsv.mjs +2 -0
- package/src/util/buffer.mjs +4 -2
- package/src/util/compute.mjs +6 -3
- package/src/util/f64.mjs +3 -3
|
@@ -0,0 +1,112 @@
|
|
|
1
|
+
// sgemm_small: C = alpha * op(A) * op(B) + beta * C — small-tile half of
|
|
2
|
+
// the two-tier autotuned dispatch (see sgemm.mjs and sgemm_large.wgsl).
|
|
3
|
+
// BM=BN=32, BK=8, TM=TN=2 — wins over the large tile below a 6x6=36
|
|
4
|
+
// workgroup grid of 64-tiles, where the large tile doesn't have enough
|
|
5
|
+
// workgroups to fill the GPU. Same structure as sgemm_large.wgsl (2D
|
|
6
|
+
// register-blocked, shared-memory-tiled), just smaller.
|
|
7
|
+
//
|
|
8
|
+
// col mapped to gid.x for coalesced B/C access (row-major: col contiguous).
|
|
9
|
+
|
|
10
|
+
const BM: u32 = 32u;
|
|
11
|
+
const BN: u32 = 32u;
|
|
12
|
+
const BK: u32 = 8u;
|
|
13
|
+
const TM: u32 = 2u;
|
|
14
|
+
const TN: u32 = 2u;
|
|
15
|
+
const THREADS_X: u32 = BN / TN;
|
|
16
|
+
const THREADS_Y: u32 = BM / TM;
|
|
17
|
+
const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 256
|
|
18
|
+
const STRIDE_A: u32 = NUM_THREADS / BK;
|
|
19
|
+
const STRIDE_B: u32 = NUM_THREADS / BN;
|
|
20
|
+
|
|
21
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
22
|
+
@group(0) @binding(1) var<storage, read> B: array<f32>;
|
|
23
|
+
@group(0) @binding(2) var<storage, read_write> C: array<f32>;
|
|
24
|
+
|
|
25
|
+
struct Params {
|
|
26
|
+
m: u32,
|
|
27
|
+
n: u32,
|
|
28
|
+
k: u32,
|
|
29
|
+
alpha: f32,
|
|
30
|
+
beta: f32,
|
|
31
|
+
lda: u32,
|
|
32
|
+
ldb: u32,
|
|
33
|
+
ldc: u32,
|
|
34
|
+
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
35
|
+
transB: u32,
|
|
36
|
+
}
|
|
37
|
+
|
|
38
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
39
|
+
|
|
40
|
+
var<workgroup> As: array<f32, BM * BK>;
|
|
41
|
+
var<workgroup> Bs: array<f32, BK * BN>;
|
|
42
|
+
|
|
43
|
+
@compute @workgroup_size(THREADS_X, THREADS_Y)
|
|
44
|
+
fn main(
|
|
45
|
+
@builtin(workgroup_id) wid: vec3u,
|
|
46
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
47
|
+
@builtin(local_invocation_index) tid: u32,
|
|
48
|
+
) {
|
|
49
|
+
let blockRow = wid.y * BM;
|
|
50
|
+
let blockCol = wid.x * BN;
|
|
51
|
+
let threadCol = lid.x;
|
|
52
|
+
let threadRow = lid.y;
|
|
53
|
+
|
|
54
|
+
let innerRowA = tid / BK;
|
|
55
|
+
let innerColA = tid % BK;
|
|
56
|
+
let innerRowB = tid / BN;
|
|
57
|
+
let innerColB = tid % BN;
|
|
58
|
+
|
|
59
|
+
var threadResults: array<f32, TM * TN>;
|
|
60
|
+
for (var i = 0u; i < TM * TN; i++) {
|
|
61
|
+
threadResults[i] = 0.0;
|
|
62
|
+
}
|
|
63
|
+
var regM: array<f32, TM>;
|
|
64
|
+
var regN: array<f32, TN>;
|
|
65
|
+
|
|
66
|
+
let numTiles = (params.k + BK - 1u) / BK;
|
|
67
|
+
for (var t = 0u; t < numTiles; t++) {
|
|
68
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
69
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
70
|
+
let gColA = t * BK + innerColA;
|
|
71
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
72
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
73
|
+
}
|
|
74
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
75
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
76
|
+
let gColB = blockCol + innerColB;
|
|
77
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
78
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
79
|
+
}
|
|
80
|
+
|
|
81
|
+
workgroupBarrier();
|
|
82
|
+
|
|
83
|
+
for (var dotIdx = 0u; dotIdx < BK; dotIdx++) {
|
|
84
|
+
for (var i = 0u; i < TM; i++) {
|
|
85
|
+
regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
|
|
86
|
+
}
|
|
87
|
+
for (var i = 0u; i < TN; i++) {
|
|
88
|
+
regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
|
|
89
|
+
}
|
|
90
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
91
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
92
|
+
threadResults[resIdxM * TN + resIdxN] += regM[resIdxM] * regN[resIdxN];
|
|
93
|
+
}
|
|
94
|
+
}
|
|
95
|
+
}
|
|
96
|
+
|
|
97
|
+
workgroupBarrier();
|
|
98
|
+
}
|
|
99
|
+
|
|
100
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
101
|
+
let row = blockRow + threadRow * TM + resIdxM;
|
|
102
|
+
if (row < params.m) {
|
|
103
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
104
|
+
let col = blockCol + threadCol * TN + resIdxN;
|
|
105
|
+
if (col < params.n) {
|
|
106
|
+
let cIdx = row * params.ldc + col;
|
|
107
|
+
C[cIdx] = params.alpha * threadResults[resIdxM * TN + resIdxN] + params.beta * C[cIdx];
|
|
108
|
+
}
|
|
109
|
+
}
|
|
110
|
+
}
|
|
111
|
+
}
|
|
112
|
+
}
|
|
@@ -0,0 +1,117 @@
|
|
|
1
|
+
// sgemmtr_large: C := uplo(alpha * op(A) * op(B) + beta * C) — large-tile
|
|
2
|
+
// half of a two-tier dispatch, identical to sgemm_large.wgsl (see that file
|
|
3
|
+
// for the BM/BN/BK/TM/TN autotuning rationale) except the final output write
|
|
4
|
+
// is gated to one triangle of C by `uplo`, the same convention ssyr/ssyr2
|
|
5
|
+
// use (0 = lower: col <= row, 1 = upper: col >= row). Every other element of
|
|
6
|
+
// C — including inside the compute loop, where the full tile is still
|
|
7
|
+
// computed regardless of uplo, only the write is masked — is left untouched.
|
|
8
|
+
// gemmtr's uplo(C) test is a plain row/col comparison over the full m×n
|
|
9
|
+
// grid, well-defined even when m != n (not restricted to square C).
|
|
10
|
+
|
|
11
|
+
const BM: u32 = 64u;
|
|
12
|
+
const BN: u32 = 64u;
|
|
13
|
+
const BK: u32 = 8u;
|
|
14
|
+
const TM: u32 = 8u;
|
|
15
|
+
const TN: u32 = 4u;
|
|
16
|
+
const THREADS_X: u32 = BN / TN;
|
|
17
|
+
const THREADS_Y: u32 = BM / TM;
|
|
18
|
+
const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 128
|
|
19
|
+
const STRIDE_A: u32 = NUM_THREADS / BK; // rows of As covered per load-loop step
|
|
20
|
+
const STRIDE_B: u32 = NUM_THREADS / BN; // rows of Bs covered per load-loop step
|
|
21
|
+
|
|
22
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
23
|
+
@group(0) @binding(1) var<storage, read> B: array<f32>;
|
|
24
|
+
@group(0) @binding(2) var<storage, read_write> C: array<f32>;
|
|
25
|
+
|
|
26
|
+
struct Params {
|
|
27
|
+
m: u32,
|
|
28
|
+
n: u32,
|
|
29
|
+
k: u32,
|
|
30
|
+
alpha: f32,
|
|
31
|
+
beta: f32,
|
|
32
|
+
lda: u32,
|
|
33
|
+
ldb: u32,
|
|
34
|
+
ldc: u32,
|
|
35
|
+
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
36
|
+
transB: u32,
|
|
37
|
+
uplo: u32, // 0 = lower (col <= row), 1 = upper (col >= row)
|
|
38
|
+
}
|
|
39
|
+
|
|
40
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
41
|
+
|
|
42
|
+
var<workgroup> As: array<f32, BM * BK>;
|
|
43
|
+
var<workgroup> Bs: array<f32, BK * BN>;
|
|
44
|
+
|
|
45
|
+
@compute @workgroup_size(THREADS_X, THREADS_Y)
|
|
46
|
+
fn main(
|
|
47
|
+
@builtin(workgroup_id) wid: vec3u,
|
|
48
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
49
|
+
@builtin(local_invocation_index) tid: u32,
|
|
50
|
+
) {
|
|
51
|
+
let blockRow = wid.y * BM;
|
|
52
|
+
let blockCol = wid.x * BN;
|
|
53
|
+
let threadCol = lid.x;
|
|
54
|
+
let threadRow = lid.y;
|
|
55
|
+
|
|
56
|
+
// Load indices, independent of the compute thread shape — a loop since
|
|
57
|
+
// NUM_THREADS doesn't match the tile size 1:1 at this config.
|
|
58
|
+
let innerRowA = tid / BK;
|
|
59
|
+
let innerColA = tid % BK;
|
|
60
|
+
let innerRowB = tid / BN;
|
|
61
|
+
let innerColB = tid % BN;
|
|
62
|
+
|
|
63
|
+
var threadResults: array<f32, TM * TN>;
|
|
64
|
+
for (var i = 0u; i < TM * TN; i++) {
|
|
65
|
+
threadResults[i] = 0.0;
|
|
66
|
+
}
|
|
67
|
+
var regM: array<f32, TM>;
|
|
68
|
+
var regN: array<f32, TN>;
|
|
69
|
+
|
|
70
|
+
let numTiles = (params.k + BK - 1u) / BK;
|
|
71
|
+
for (var t = 0u; t < numTiles; t++) {
|
|
72
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
73
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
74
|
+
let gColA = t * BK + innerColA;
|
|
75
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
76
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
77
|
+
}
|
|
78
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
79
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
80
|
+
let gColB = blockCol + innerColB;
|
|
81
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
82
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
83
|
+
}
|
|
84
|
+
|
|
85
|
+
workgroupBarrier();
|
|
86
|
+
|
|
87
|
+
for (var dotIdx = 0u; dotIdx < BK; dotIdx++) {
|
|
88
|
+
for (var i = 0u; i < TM; i++) {
|
|
89
|
+
regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
|
|
90
|
+
}
|
|
91
|
+
for (var i = 0u; i < TN; i++) {
|
|
92
|
+
regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
|
|
93
|
+
}
|
|
94
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
95
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
96
|
+
threadResults[resIdxM * TN + resIdxN] += regM[resIdxM] * regN[resIdxN];
|
|
97
|
+
}
|
|
98
|
+
}
|
|
99
|
+
}
|
|
100
|
+
|
|
101
|
+
workgroupBarrier();
|
|
102
|
+
}
|
|
103
|
+
|
|
104
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
105
|
+
let row = blockRow + threadRow * TM + resIdxM;
|
|
106
|
+
if (row < params.m) {
|
|
107
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
108
|
+
let col = blockCol + threadCol * TN + resIdxN;
|
|
109
|
+
let inTriangle = select(col >= row, col <= row, params.uplo == 0u);
|
|
110
|
+
if (col < params.n && inTriangle) {
|
|
111
|
+
let cIdx = row * params.ldc + col;
|
|
112
|
+
C[cIdx] = params.alpha * threadResults[resIdxM * TN + resIdxN] + params.beta * C[cIdx];
|
|
113
|
+
}
|
|
114
|
+
}
|
|
115
|
+
}
|
|
116
|
+
}
|
|
117
|
+
}
|
|
@@ -0,0 +1,110 @@
|
|
|
1
|
+
// sgemmtr_small: C := uplo(alpha * op(A) * op(B) + beta * C) — small-tile
|
|
2
|
+
// half of a two-tier dispatch, identical to sgemm_small.wgsl except the
|
|
3
|
+
// final output write is gated to one triangle of C by `uplo` — see
|
|
4
|
+
// sgemmtr_large.wgsl for the full rationale (shared by both tiers).
|
|
5
|
+
|
|
6
|
+
const BM: u32 = 32u;
|
|
7
|
+
const BN: u32 = 32u;
|
|
8
|
+
const BK: u32 = 8u;
|
|
9
|
+
const TM: u32 = 2u;
|
|
10
|
+
const TN: u32 = 2u;
|
|
11
|
+
const THREADS_X: u32 = BN / TN;
|
|
12
|
+
const THREADS_Y: u32 = BM / TM;
|
|
13
|
+
const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 256
|
|
14
|
+
const STRIDE_A: u32 = NUM_THREADS / BK;
|
|
15
|
+
const STRIDE_B: u32 = NUM_THREADS / BN;
|
|
16
|
+
|
|
17
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
18
|
+
@group(0) @binding(1) var<storage, read> B: array<f32>;
|
|
19
|
+
@group(0) @binding(2) var<storage, read_write> C: array<f32>;
|
|
20
|
+
|
|
21
|
+
struct Params {
|
|
22
|
+
m: u32,
|
|
23
|
+
n: u32,
|
|
24
|
+
k: u32,
|
|
25
|
+
alpha: f32,
|
|
26
|
+
beta: f32,
|
|
27
|
+
lda: u32,
|
|
28
|
+
ldb: u32,
|
|
29
|
+
ldc: u32,
|
|
30
|
+
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
31
|
+
transB: u32,
|
|
32
|
+
uplo: u32, // 0 = lower (col <= row), 1 = upper (col >= row)
|
|
33
|
+
}
|
|
34
|
+
|
|
35
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
36
|
+
|
|
37
|
+
var<workgroup> As: array<f32, BM * BK>;
|
|
38
|
+
var<workgroup> Bs: array<f32, BK * BN>;
|
|
39
|
+
|
|
40
|
+
@compute @workgroup_size(THREADS_X, THREADS_Y)
|
|
41
|
+
fn main(
|
|
42
|
+
@builtin(workgroup_id) wid: vec3u,
|
|
43
|
+
@builtin(local_invocation_id) lid: vec3u,
|
|
44
|
+
@builtin(local_invocation_index) tid: u32,
|
|
45
|
+
) {
|
|
46
|
+
let blockRow = wid.y * BM;
|
|
47
|
+
let blockCol = wid.x * BN;
|
|
48
|
+
let threadCol = lid.x;
|
|
49
|
+
let threadRow = lid.y;
|
|
50
|
+
|
|
51
|
+
let innerRowA = tid / BK;
|
|
52
|
+
let innerColA = tid % BK;
|
|
53
|
+
let innerRowB = tid / BN;
|
|
54
|
+
let innerColB = tid % BN;
|
|
55
|
+
|
|
56
|
+
var threadResults: array<f32, TM * TN>;
|
|
57
|
+
for (var i = 0u; i < TM * TN; i++) {
|
|
58
|
+
threadResults[i] = 0.0;
|
|
59
|
+
}
|
|
60
|
+
var regM: array<f32, TM>;
|
|
61
|
+
var regN: array<f32, TN>;
|
|
62
|
+
|
|
63
|
+
let numTiles = (params.k + BK - 1u) / BK;
|
|
64
|
+
for (var t = 0u; t < numTiles; t++) {
|
|
65
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
66
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
67
|
+
let gColA = t * BK + innerColA;
|
|
68
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
69
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
70
|
+
}
|
|
71
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
72
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
73
|
+
let gColB = blockCol + innerColB;
|
|
74
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
75
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
76
|
+
}
|
|
77
|
+
|
|
78
|
+
workgroupBarrier();
|
|
79
|
+
|
|
80
|
+
for (var dotIdx = 0u; dotIdx < BK; dotIdx++) {
|
|
81
|
+
for (var i = 0u; i < TM; i++) {
|
|
82
|
+
regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
|
|
83
|
+
}
|
|
84
|
+
for (var i = 0u; i < TN; i++) {
|
|
85
|
+
regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
|
|
86
|
+
}
|
|
87
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
88
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
89
|
+
threadResults[resIdxM * TN + resIdxN] += regM[resIdxM] * regN[resIdxN];
|
|
90
|
+
}
|
|
91
|
+
}
|
|
92
|
+
}
|
|
93
|
+
|
|
94
|
+
workgroupBarrier();
|
|
95
|
+
}
|
|
96
|
+
|
|
97
|
+
for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
|
|
98
|
+
let row = blockRow + threadRow * TM + resIdxM;
|
|
99
|
+
if (row < params.m) {
|
|
100
|
+
for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
|
|
101
|
+
let col = blockCol + threadCol * TN + resIdxN;
|
|
102
|
+
let inTriangle = select(col >= row, col <= row, params.uplo == 0u);
|
|
103
|
+
if (col < params.n && inTriangle) {
|
|
104
|
+
let cIdx = row * params.ldc + col;
|
|
105
|
+
C[cIdx] = params.alpha * threadResults[resIdxM * TN + resIdxN] + params.beta * C[cIdx];
|
|
106
|
+
}
|
|
107
|
+
}
|
|
108
|
+
}
|
|
109
|
+
}
|
|
110
|
+
}
|
|
@@ -0,0 +1,31 @@
|
|
|
1
|
+
// symmetrize: Adense := full dense expansion of a symmetric matrix stored
|
|
2
|
+
// with only its `uplo` triangle meaningful (the other triangle is implied
|
|
3
|
+
// by symmetry: A[i,j] = A[j,i]). A plain element-wise pass, no tiling or
|
|
4
|
+
// shared memory needed — used to materialize a dense operand for routines
|
|
5
|
+
// that read a symmetric matrix as a normal dense gemm input (e.g. ssymm),
|
|
6
|
+
// rather than teaching the tiled gemm kernel itself to mirror-read.
|
|
7
|
+
|
|
8
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
9
|
+
@group(0) @binding(1) var<storage, read_write> Adense: array<f32>;
|
|
10
|
+
|
|
11
|
+
struct Params {
|
|
12
|
+
n: u32,
|
|
13
|
+
lda: u32,
|
|
14
|
+
ldd: u32, // leading dimension of Adense
|
|
15
|
+
uplo: u32, // 0 = lower (stored where col <= row), 1 = upper (col >= row)
|
|
16
|
+
}
|
|
17
|
+
|
|
18
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
19
|
+
|
|
20
|
+
@compute @workgroup_size(8, 8)
|
|
21
|
+
fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
22
|
+
let row = gid.y;
|
|
23
|
+
let col = gid.x;
|
|
24
|
+
if (row >= params.n || col >= params.n) {
|
|
25
|
+
return;
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
let isStored = select(col >= row, col <= row, params.uplo == 0u);
|
|
29
|
+
let srcIdx = select(col * params.lda + row, row * params.lda + col, isStored);
|
|
30
|
+
Adense[row * params.ldd + col] = A[srcIdx];
|
|
31
|
+
}
|
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
// triangularize: Adense := dense expansion of op(A) (A or A^T per `trans`),
|
|
2
|
+
// zero-filling the unstored triangle (exact for a matmul) so strmm can reuse
|
|
3
|
+
// sgemm's kernel unchanged. `diag=1` substitutes 1.0 on the diagonal.
|
|
4
|
+
|
|
5
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
6
|
+
@group(0) @binding(1) var<storage, read_write> Adense: array<f32>;
|
|
7
|
+
|
|
8
|
+
struct Params {
|
|
9
|
+
n: u32,
|
|
10
|
+
lda: u32,
|
|
11
|
+
ldd: u32, // leading dimension of Adense
|
|
12
|
+
uplo: u32, // 0 = lower, 1 = upper
|
|
13
|
+
trans: u32, // 0 = no-transpose (op(A) = A), 1 = transpose (op(A) = A^T)
|
|
14
|
+
diag: u32, // 0 = non-unit, 1 = unit
|
|
15
|
+
}
|
|
16
|
+
|
|
17
|
+
@group(0) @binding(2) var<uniform> params: Params;
|
|
18
|
+
|
|
19
|
+
@compute @workgroup_size(8, 8)
|
|
20
|
+
fn main(@builtin(global_invocation_id) gid: vec3u) {
|
|
21
|
+
let row = gid.y;
|
|
22
|
+
let col = gid.x;
|
|
23
|
+
if (row >= params.n || col >= params.n) {
|
|
24
|
+
return;
|
|
25
|
+
}
|
|
26
|
+
|
|
27
|
+
if (row == col) {
|
|
28
|
+
Adense[row * params.ldd + col] = select(A[row * params.lda + row], 1.0, params.diag == 1u);
|
|
29
|
+
return;
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
var isMeaningful: bool;
|
|
33
|
+
var srcRow: u32;
|
|
34
|
+
var srcCol: u32;
|
|
35
|
+
if (params.trans == 0u) {
|
|
36
|
+
isMeaningful = select(col >= row, col <= row, params.uplo == 0u);
|
|
37
|
+
srcRow = row; srcCol = col;
|
|
38
|
+
} else {
|
|
39
|
+
isMeaningful = select(col <= row, col >= row, params.uplo == 0u);
|
|
40
|
+
srcRow = col; srcCol = row;
|
|
41
|
+
}
|
|
42
|
+
|
|
43
|
+
Adense[row * params.ldd + col] = select(0.0, A[srcRow * params.lda + srcCol], isMeaningful);
|
|
44
|
+
}
|
package/src/snrm2/snrm2.d.mts
CHANGED
|
@@ -26,7 +26,7 @@ export declare function snrm2(
|
|
|
26
26
|
/**
|
|
27
27
|
* Computes the Euclidean norm of a vector: result = sqrt(sum(x[i] * x[i]))
|
|
28
28
|
*
|
|
29
|
-
* {@includeCode ../../examples/snrm2/
|
|
29
|
+
* {@includeCode ../../examples/snrm2/gpu.snrm2.js}
|
|
30
30
|
*
|
|
31
31
|
* @param device - GPUDevice from `init()`
|
|
32
32
|
* @param n - number of elements (must be a positive integer)
|
package/src/srot/srot.d.mts
CHANGED
|
@@ -37,7 +37,7 @@ export declare function srot(
|
|
|
37
37
|
* x = c*x + s*y
|
|
38
38
|
* y = -s*x + c*y
|
|
39
39
|
*
|
|
40
|
-
* {@includeCode ../../examples/srot/
|
|
40
|
+
* {@includeCode ../../examples/srot/gpu.srot.js}
|
|
41
41
|
*
|
|
42
42
|
* @param device - GPUDevice from `init()`
|
|
43
43
|
* @param n - number of elements (must be a positive integer)
|
package/src/srotm/srotm.d.mts
CHANGED
|
@@ -36,7 +36,7 @@ export declare function srotm(
|
|
|
36
36
|
* x = H[0][0]*x + H[0][1]*y
|
|
37
37
|
* y = H[1][0]*x + H[1][1]*y
|
|
38
38
|
*
|
|
39
|
-
* {@includeCode ../../examples/srotm/
|
|
39
|
+
* {@includeCode ../../examples/srotm/gpu.srotm.js}
|
|
40
40
|
*
|
|
41
41
|
* @param device - GPUDevice from `init()`
|
|
42
42
|
* @param n - number of elements (must be a positive integer)
|
package/src/sscal/sscal.d.mts
CHANGED
|
@@ -27,7 +27,7 @@ export declare function sscal(
|
|
|
27
27
|
/**
|
|
28
28
|
* Scales a single-precision vector by a constant: x = alpha * x
|
|
29
29
|
*
|
|
30
|
-
* {@includeCode ../../examples/sscal/
|
|
30
|
+
* {@includeCode ../../examples/sscal/gpu.sscal.js}
|
|
31
31
|
*
|
|
32
32
|
* @param device - GPUDevice from `init()`
|
|
33
33
|
* @param n - number of elements to scale (must be a positive integer)
|
package/src/sswap/sswap.d.mts
CHANGED
|
@@ -29,7 +29,7 @@ export declare function sswap(
|
|
|
29
29
|
/**
|
|
30
30
|
* Swaps the elements of two single-precision vectors: x <-> y
|
|
31
31
|
*
|
|
32
|
-
* {@includeCode ../../examples/sswap/
|
|
32
|
+
* {@includeCode ../../examples/sswap/gpu.sswap.js}
|
|
33
33
|
*
|
|
34
34
|
* @param device - GPUDevice from `init()`
|
|
35
35
|
* @param n - number of elements to swap (must be a positive integer)
|
|
@@ -0,0 +1,103 @@
|
|
|
1
|
+
import { GpuMatrix } from "../classes/GpuMatrix.mjs";
|
|
2
|
+
|
|
3
|
+
/**
|
|
4
|
+
* Performs the symmetric matrix-matrix operation
|
|
5
|
+
* C := alpha * A * B + beta * C (`side='left'`) or
|
|
6
|
+
* C := alpha * B * A + beta * C (`side='right'`) — `A` is symmetric, only
|
|
7
|
+
* its `uplo` triangle stored; `B` and `C` are general m×n matrices.
|
|
8
|
+
*
|
|
9
|
+
* - `side='left'`: `A` is m×m — `A` premultiplies `B`
|
|
10
|
+
* - `side='right'`: `A` is n×n — `A` postmultiplies `B`
|
|
11
|
+
*
|
|
12
|
+
* No dedicated fused kernel — a `symmetrize` pass materializes a dense
|
|
13
|
+
* copy of `A` (mirroring the unstored triangle), then a plain `sgemm`
|
|
14
|
+
* pass (`sgemm_small.wgsl`/`sgemm_large.wgsl`, unmodified) does the
|
|
15
|
+
* actual multiply, both on one command encoder.
|
|
16
|
+
*
|
|
17
|
+
* {@includeCode ../../examples/ssymm/ssymm.js}
|
|
18
|
+
*
|
|
19
|
+
* **Browser (standalone HTML):**
|
|
20
|
+
* {@includeCode ../../examples/ssymm/web/ssymm.html}
|
|
21
|
+
*
|
|
22
|
+
* @param device - GPUDevice from `init()`
|
|
23
|
+
* @param side - `'left'` for A*B, `'right'` for B*A
|
|
24
|
+
* @param uplo - `'lower'` if only `A`'s lower triangle is stored, `'upper'` for upper
|
|
25
|
+
* @param m - rows of B and C
|
|
26
|
+
* @param n - columns of B and C
|
|
27
|
+
* @param alpha - scalar multiplier for the matrix product
|
|
28
|
+
* @param A - Float32Array, symmetric, row-major or column-major (see `layout`)
|
|
29
|
+
* @param lda - leading dimension of A as stored
|
|
30
|
+
* @param B - Float32Array, row-major or column-major (see `layout`)
|
|
31
|
+
* @param ldb - leading dimension of B as stored
|
|
32
|
+
* @param beta - scalar multiplier for C
|
|
33
|
+
* @param C - Float32Array input/output matrix, row-major or column-major
|
|
34
|
+
* @param ldc - leading dimension of C as stored
|
|
35
|
+
* @param layout - storage layout shared by A/B/C when they're Float32Array
|
|
36
|
+
* (default: `'row-major'`) — column-major A keeps representing the same
|
|
37
|
+
* symmetric matrix but flips which physical triangle looks stored, so
|
|
38
|
+
* `uplo` is adjusted internally to compensate
|
|
39
|
+
* @returns updated C as a Float32Array
|
|
40
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/ssymm/ssymm.mjs#L23">Source code: ssymm.mjs (L23)</a>
|
|
41
|
+
* @category BLAS Level 3
|
|
42
|
+
*/
|
|
43
|
+
export declare function ssymm(
|
|
44
|
+
device: GPUDevice,
|
|
45
|
+
side: 'left' | 'right',
|
|
46
|
+
uplo: 'lower' | 'upper',
|
|
47
|
+
m: number,
|
|
48
|
+
n: number,
|
|
49
|
+
alpha: number,
|
|
50
|
+
A: Float32Array,
|
|
51
|
+
lda: number,
|
|
52
|
+
B: Float32Array,
|
|
53
|
+
ldb: number,
|
|
54
|
+
beta: number,
|
|
55
|
+
C: Float32Array,
|
|
56
|
+
ldc: number,
|
|
57
|
+
layout?: 'row-major' | 'column-major',
|
|
58
|
+
): Promise<{ C: Float32Array; gpuTimeMs?: number }>;
|
|
59
|
+
|
|
60
|
+
/**
|
|
61
|
+
* Performs the symmetric matrix-matrix operation
|
|
62
|
+
* C := alpha * A * B + beta * C (`side='left'`) or
|
|
63
|
+
* C := alpha * B * A + beta * C (`side='right'`)
|
|
64
|
+
*
|
|
65
|
+
* A, B, and C are all kept GPU-resident. Each matrix's own `layout` (set at
|
|
66
|
+
* `GpuMatrix.from` time) determines the operation — there is no separate
|
|
67
|
+
* `layout` argument here. A and B must be GpuMatrix whenever C is, and vice
|
|
68
|
+
* versa — mixing a GpuMatrix with a plain Float32Array is not supported.
|
|
69
|
+
*
|
|
70
|
+
* {@includeCode ../../examples/ssymm/gpu.ssymm.js}
|
|
71
|
+
*
|
|
72
|
+
* @param device - GPUDevice from `init()`
|
|
73
|
+
* @param side - `'left'` for A*B, `'right'` for B*A
|
|
74
|
+
* @param uplo - `'lower'` if only `A`'s lower triangle is stored, `'upper'` for upper
|
|
75
|
+
* @param m - rows of B and C
|
|
76
|
+
* @param n - columns of B and C
|
|
77
|
+
* @param alpha - scalar multiplier for the matrix product
|
|
78
|
+
* @param A - GpuMatrix, symmetric
|
|
79
|
+
* @param lda - leading dimension of A (must equal A.lda)
|
|
80
|
+
* @param B - GpuMatrix
|
|
81
|
+
* @param ldb - leading dimension of B (must equal B.lda)
|
|
82
|
+
* @param beta - scalar multiplier for C
|
|
83
|
+
* @param C - GpuMatrix (mutated in place)
|
|
84
|
+
* @param ldc - leading dimension of C (must equal C.lda)
|
|
85
|
+
* @returns no C — it stays GPU-resident; call `C.read()` yourself for a CPU readback (see the example)
|
|
86
|
+
* @see <a href="https://github.com/manit2004/wgblas/blob/main/src/ssymm/ssymm.mjs#L23">Source code: ssymm.mjs (L23)</a>
|
|
87
|
+
* @category BLAS Level 3
|
|
88
|
+
*/
|
|
89
|
+
export declare function ssymm(
|
|
90
|
+
device: GPUDevice,
|
|
91
|
+
side: 'left' | 'right',
|
|
92
|
+
uplo: 'lower' | 'upper',
|
|
93
|
+
m: number,
|
|
94
|
+
n: number,
|
|
95
|
+
alpha: number,
|
|
96
|
+
A: GpuMatrix,
|
|
97
|
+
lda: number,
|
|
98
|
+
B: GpuMatrix,
|
|
99
|
+
ldb: number,
|
|
100
|
+
beta: number,
|
|
101
|
+
C: GpuMatrix,
|
|
102
|
+
ldc: number,
|
|
103
|
+
): Promise<{ gpuTimeMs?: number }>;
|