wgblas 2.0.0 → 2.2.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 +20 -18
- package/dist/wgblas.browser.js +2172 -1174
- package/index.d.mts +49 -44
- package/index.mjs +11 -0
- package/package.json +133 -63
- package/src/classes/Complex32.d.mts +43 -0
- package/src/classes/Complex32.mjs +82 -0
- package/src/classes/Complex64.d.mts +44 -0
- package/src/classes/Complex64.mjs +76 -0
- package/src/classes/GpuMatrix.d.mts +41 -41
- package/src/classes/GpuMatrix.mjs +126 -17
- package/src/classes/GpuVector.d.mts +36 -40
- package/src/classes/GpuVector.mjs +66 -11
- package/src/cscal/cscal.d.mts +47 -0
- package/src/cscal/cscal.mjs +98 -0
- package/src/dasum/dasum.d.mts +4 -4
- package/src/dasum/dasum.mjs +38 -20
- package/src/daxpy/daxpy.d.mts +56 -0
- package/src/daxpy/daxpy.mjs +150 -0
- package/src/dcopy/dcopy.d.mts +52 -0
- package/src/dcopy/dcopy.mjs +140 -0
- package/src/ddot/ddot.d.mts +62 -0
- package/src/ddot/ddot.mjs +184 -0
- package/src/devdocs.mjs +13 -0
- package/src/dnrm2/dnrm2.d.mts +50 -0
- package/src/dnrm2/dnrm2.mjs +189 -0
- package/src/drot/drot.d.mts +67 -0
- package/src/drot/drot.mjs +170 -0
- package/src/drotm/drotm.d.mts +67 -0
- package/src/drotm/drotm.mjs +171 -0
- package/src/dscal/dscal.d.mts +52 -0
- package/src/dscal/dscal.mjs +119 -0
- package/src/dswap/dswap.d.mts +57 -0
- package/src/dswap/dswap.mjs +155 -0
- package/src/idamax/idamax.d.mts +20 -2
- package/src/idamax/idamax.mjs +56 -24
- package/src/init.mjs +117 -56
- package/src/isamax/isamax.d.mts +20 -2
- package/src/isamax/isamax.mjs +21 -16
- package/src/random/random.d.mts +37 -39
- package/src/random/random.mjs +39 -7
- package/src/sasum/sasum.d.mts +2 -2
- package/src/sasum/sasum.mjs +20 -16
- package/src/saxpy/saxpy.d.mts +2 -2
- package/src/saxpy/saxpy.mjs +14 -11
- package/src/scopy/scopy.d.mts +2 -2
- package/src/scopy/scopy.mjs +13 -9
- package/src/sdot/sdot.d.mts +2 -2
- package/src/sdot/sdot.mjs +21 -17
- package/src/sgemm/sgemm.d.mts +2 -2
- package/src/sgemm/sgemm.mjs +109 -40
- package/src/sgemmtr/sgemmtr.d.mts +3 -2
- package/src/sgemmtr/sgemmtr.mjs +98 -40
- package/src/sgemv/sgemv.d.mts +2 -2
- package/src/sgemv/sgemv.mjs +69 -41
- package/src/sger/sger.d.mts +2 -2
- package/src/sger/sger.mjs +43 -19
- package/src/shaders/__test_pipeline_a.wgsl +3 -0
- package/src/shaders/__test_pipeline_b.wgsl +2 -0
- package/src/shaders/cscal.wgsl +33 -0
- package/src/shaders/daxpy.wgsl +66 -0
- package/src/shaders/dcopy.wgsl +34 -0
- package/src/shaders/ddot.wgsl +106 -0
- package/src/shaders/dnrm2.wgsl +167 -0
- package/src/shaders/drot.wgsl +81 -0
- package/src/shaders/drotm.wgsl +99 -0
- package/src/shaders/dscal.wgsl +60 -0
- package/src/shaders/dswap.wgsl +38 -0
- package/src/shaders/f64/utils/add.wgsl +6 -0
- package/src/shaders/f64/utils/divide.wgsl +45 -0
- package/src/shaders/f64/utils/multiply.wgsl +19 -10
- package/src/shaders/f64/utils/sqrt.wgsl +45 -0
- package/src/shaders/index.mjs +233 -14
- package/src/shaders/reduction/scaledSum.wgsl +65 -0
- package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
- package/src/shaders/sgemm_large.wgsl +107 -18
- package/src/shaders/sgemm_small.wgsl +115 -15
- package/src/shaders/sgemmtr_large.wgsl +4 -1
- package/src/shaders/sgemmtr_small.wgsl +4 -1
- package/src/shaders/sgemv_n.wgsl +3 -1
- package/src/shaders/sgemv_t.wgsl +3 -1
- package/src/shaders/snrm2.wgsl +72 -23
- package/src/shaders/ssymv.wgsl +3 -1
- package/src/snrm2/snrm2.d.mts +2 -2
- package/src/snrm2/snrm2.mjs +41 -23
- package/src/srot/srot.d.mts +2 -4
- package/src/srot/srot.mjs +16 -11
- package/src/srotm/srotm.d.mts +2 -4
- package/src/srotm/srotm.mjs +17 -11
- package/src/sscal/sscal.d.mts +3 -3
- package/src/sscal/sscal.mjs +14 -12
- package/src/sswap/sswap.d.mts +2 -2
- package/src/sswap/sswap.mjs +18 -10
- package/src/ssymm/ssymm.d.mts +5 -4
- package/src/ssymm/ssymm.mjs +150 -54
- package/src/ssymv/ssymv.d.mts +2 -2
- package/src/ssymv/ssymv.mjs +47 -26
- package/src/ssyr/ssyr.d.mts +2 -2
- package/src/ssyr/ssyr.mjs +38 -17
- package/src/ssyr2/ssyr2.d.mts +2 -2
- package/src/ssyr2/ssyr2.mjs +48 -21
- package/src/ssyr2k/ssyr2k.d.mts +3 -2
- package/src/ssyr2k/ssyr2k.mjs +140 -62
- package/src/ssyrk/ssyrk.d.mts +3 -2
- package/src/ssyrk/ssyrk.mjs +91 -39
- package/src/strmm/strmm.d.mts +5 -4
- package/src/strmm/strmm.mjs +174 -60
- package/src/strmv/strmv.d.mts +2 -2
- package/src/strmv/strmv.mjs +42 -20
- package/src/strsm/strsm.d.mts +6 -4
- package/src/strsm/strsm.mjs +438 -174
- package/src/strsv/strsv.d.mts +5 -3
- package/src/strsv/strsv.mjs +89 -34
- package/src/util/benchmark.mjs +9 -9
- package/src/util/bindgroup.mjs +1 -3
- package/src/util/buffer.mjs +139 -24
- package/src/util/complex.mjs +87 -0
- package/src/util/compute.mjs +19 -16
- package/src/util/constants.mjs +57 -0
- package/src/util/device.mjs +49 -0
- package/src/util/pipeline.mjs +44 -10
- package/src/util/workgroup.mjs +72 -7
- package/src/shaders/browser-shaders.mjs +0 -81
- package/src/shaders/f64add.wgsl +0 -281
|
@@ -5,6 +5,13 @@
|
|
|
5
5
|
// workgroups to fill the GPU. Same structure as sgemm_large.wgsl (2D
|
|
6
6
|
// register-blocked, shared-memory-tiled), just smaller.
|
|
7
7
|
//
|
|
8
|
+
// A and B are bound twice — scalar array<f32> and array<vec4<f32>> views of
|
|
9
|
+
// the same GPUBuffer (see vec4ViewBinding) — so each tile load can issue
|
|
10
|
+
// 16-byte vector reads along op(A)/op(B)'s contiguous dimension when the
|
|
11
|
+
// stride allows it (stride % 4 == 0 keeps every row base 16-byte aligned).
|
|
12
|
+
// NUM_THREADS (256) exceeds some small-tile load shapes, so the vectorized
|
|
13
|
+
// paths whose lane count doesn't tile exactly guard their As/Bs stores.
|
|
14
|
+
//
|
|
8
15
|
// col mapped to gid.x for coalesced B/C access (row-major: col contiguous).
|
|
9
16
|
|
|
10
17
|
const BM: u32 = 32u;
|
|
@@ -18,9 +25,11 @@ const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 256
|
|
|
18
25
|
const STRIDE_A: u32 = NUM_THREADS / BK;
|
|
19
26
|
const STRIDE_B: u32 = NUM_THREADS / BN;
|
|
20
27
|
|
|
21
|
-
@group(0) @binding(0) var<storage, read> A:
|
|
22
|
-
@group(0) @binding(1) var<storage, read>
|
|
23
|
-
@group(0) @binding(2) var<storage,
|
|
28
|
+
@group(0) @binding(0) var<storage, read> A: array<f32>;
|
|
29
|
+
@group(0) @binding(1) var<storage, read> A4: array<vec4<f32>>;
|
|
30
|
+
@group(0) @binding(2) var<storage, read> B: array<f32>;
|
|
31
|
+
@group(0) @binding(3) var<storage, read> B4: array<vec4<f32>>;
|
|
32
|
+
@group(0) @binding(4) var<storage, read_write> C: array<f32>;
|
|
24
33
|
|
|
25
34
|
struct Params {
|
|
26
35
|
m: u32,
|
|
@@ -33,9 +42,11 @@ struct Params {
|
|
|
33
42
|
ldc: u32,
|
|
34
43
|
transA: u32, // 0 = no-transpose, 1 = transpose
|
|
35
44
|
transB: u32,
|
|
45
|
+
useVecA: u32, // 1 = A's vec4 view covers every in-bounds element (see vec4Usable)
|
|
46
|
+
useVecB: u32, // 1 = B's vec4 view covers every in-bounds element
|
|
36
47
|
}
|
|
37
48
|
|
|
38
|
-
@group(0) @binding(
|
|
49
|
+
@group(0) @binding(5) var<uniform> params: Params;
|
|
39
50
|
|
|
40
51
|
var<workgroup> As: array<f32, BM * BK>;
|
|
41
52
|
var<workgroup> Bs: array<f32, BK * BN>;
|
|
@@ -65,17 +76,103 @@ fn main(
|
|
|
65
76
|
|
|
66
77
|
let numTiles = (params.k + BK - 1u) / BK;
|
|
67
78
|
for (var t = 0u; t < numTiles; t++) {
|
|
68
|
-
|
|
69
|
-
|
|
70
|
-
|
|
71
|
-
|
|
72
|
-
|
|
79
|
+
// ── Load the BM×BK A tile into As (vectorized along op(A)'s fast dim
|
|
80
|
+
// when lda allows; every branch here is dispatch-uniform) ──
|
|
81
|
+
if (params.useVecA == 1u && params.transA == 0u) {
|
|
82
|
+
// No-transpose: columns contiguous, one vec4 per thread. NUM_THREADS
|
|
83
|
+
// spans BM/4×(BK/4) several times over — guard the store.
|
|
84
|
+
let r4 = tid / (BK / 4u);
|
|
85
|
+
let c4 = tid % (BK / 4u);
|
|
86
|
+
if (r4 < BM) {
|
|
87
|
+
let gRow = blockRow + r4;
|
|
88
|
+
let gCol = t * BK + c4 * 4u;
|
|
89
|
+
var v = A4[(gRow * params.lda + gCol) / 4u];
|
|
90
|
+
let rowOK = gRow < params.m;
|
|
91
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.k);
|
|
92
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.k);
|
|
93
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.k);
|
|
94
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.k);
|
|
95
|
+
As[r4 * BK + c4 * 4u] = v.x;
|
|
96
|
+
As[r4 * BK + c4 * 4u + 1u] = v.y;
|
|
97
|
+
As[r4 * BK + c4 * 4u + 2u] = v.z;
|
|
98
|
+
As[r4 * BK + c4 * 4u + 3u] = v.w;
|
|
99
|
+
}
|
|
100
|
+
} else if (params.useVecA == 1u && params.transA != 0u) {
|
|
101
|
+
// Transpose: rows contiguous within a column. NUM_THREADS over-spans
|
|
102
|
+
// the BK-column tile — guard the store.
|
|
103
|
+
let r4 = tid % (BM / 4u);
|
|
104
|
+
let c = tid / (BM / 4u);
|
|
105
|
+
if (c < BK) {
|
|
106
|
+
let gRow = blockRow + r4 * 4u;
|
|
107
|
+
let gCol = t * BK + c;
|
|
108
|
+
var v = A4[(gCol * params.lda + gRow) / 4u];
|
|
109
|
+
let colOK = gCol < params.k;
|
|
110
|
+
v.x = select(0.0, v.x, colOK && gRow < params.m);
|
|
111
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.m);
|
|
112
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.m);
|
|
113
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.m);
|
|
114
|
+
As[(r4 * 4u) * BK + c] = v.x;
|
|
115
|
+
As[(r4 * 4u + 1u) * BK + c] = v.y;
|
|
116
|
+
As[(r4 * 4u + 2u) * BK + c] = v.z;
|
|
117
|
+
As[(r4 * 4u + 3u) * BK + c] = v.w;
|
|
118
|
+
}
|
|
119
|
+
} else {
|
|
120
|
+
// Scalar fallback: odd stride or unhandled orientation.
|
|
121
|
+
for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
|
|
122
|
+
let gRowA = blockRow + innerRowA + loadOffset;
|
|
123
|
+
let gColA = t * BK + innerColA;
|
|
124
|
+
let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
|
|
125
|
+
As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
|
|
126
|
+
}
|
|
73
127
|
}
|
|
74
|
-
|
|
75
|
-
|
|
76
|
-
|
|
77
|
-
|
|
78
|
-
|
|
128
|
+
|
|
129
|
+
// ── Load the BK×BN B tile into Bs ──
|
|
130
|
+
if (params.useVecB == 1u && params.transB == 0u) {
|
|
131
|
+
// No-transpose: columns contiguous, one vec4 per thread. NUM_THREADS
|
|
132
|
+
// over-spans the BK-row tile — guard the store.
|
|
133
|
+
let r = tid / (BN / 4u);
|
|
134
|
+
let c4 = tid % (BN / 4u);
|
|
135
|
+
if (r < BK) {
|
|
136
|
+
let gRow = t * BK + r;
|
|
137
|
+
let gCol = blockCol + c4 * 4u;
|
|
138
|
+
var v = B4[(gRow * params.ldb + gCol) / 4u];
|
|
139
|
+
let rowOK = gRow < params.k;
|
|
140
|
+
v.x = select(0.0, v.x, rowOK && gCol < params.n);
|
|
141
|
+
v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.n);
|
|
142
|
+
v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.n);
|
|
143
|
+
v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.n);
|
|
144
|
+
Bs[r * BN + c4 * 4u] = v.x;
|
|
145
|
+
Bs[r * BN + c4 * 4u + 1u] = v.y;
|
|
146
|
+
Bs[r * BN + c4 * 4u + 2u] = v.z;
|
|
147
|
+
Bs[r * BN + c4 * 4u + 3u] = v.w;
|
|
148
|
+
}
|
|
149
|
+
} else if (params.useVecB == 1u && params.transB != 0u) {
|
|
150
|
+
// Transpose: rows contiguous within a column, one vec4 per thread —
|
|
151
|
+
// NUM_THREADS over-spans the 32-column tile, so guard the store.
|
|
152
|
+
let r4 = tid % (BK / 4u);
|
|
153
|
+
let c = tid / (BK / 4u);
|
|
154
|
+
if (c < BN) {
|
|
155
|
+
let gRow = t * BK + r4 * 4u;
|
|
156
|
+
let gCol = blockCol + c;
|
|
157
|
+
var v = B4[(gCol * params.ldb + gRow) / 4u];
|
|
158
|
+
let colOK = gCol < params.n;
|
|
159
|
+
v.x = select(0.0, v.x, colOK && gRow < params.k);
|
|
160
|
+
v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.k);
|
|
161
|
+
v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.k);
|
|
162
|
+
v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.k);
|
|
163
|
+
Bs[(r4 * 4u) * BN + c] = v.x;
|
|
164
|
+
Bs[(r4 * 4u + 1u) * BN + c] = v.y;
|
|
165
|
+
Bs[(r4 * 4u + 2u) * BN + c] = v.z;
|
|
166
|
+
Bs[(r4 * 4u + 3u) * BN + c] = v.w;
|
|
167
|
+
}
|
|
168
|
+
} else {
|
|
169
|
+
// Scalar fallback.
|
|
170
|
+
for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
|
|
171
|
+
let gRowB = t * BK + innerRowB + loadOffset;
|
|
172
|
+
let gColB = blockCol + innerColB;
|
|
173
|
+
let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
|
|
174
|
+
Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
|
|
175
|
+
}
|
|
79
176
|
}
|
|
80
177
|
|
|
81
178
|
workgroupBarrier();
|
|
@@ -104,7 +201,10 @@ fn main(
|
|
|
104
201
|
let col = blockCol + threadCol * TN + resIdxN;
|
|
105
202
|
if (col < params.n) {
|
|
106
203
|
let cIdx = row * params.ldc + col;
|
|
107
|
-
|
|
204
|
+
// BLAS beta==0 semantics: C is written, not accumulated — must not
|
|
205
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
206
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
207
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
108
208
|
}
|
|
109
209
|
}
|
|
110
210
|
}
|
|
@@ -109,7 +109,10 @@ fn main(
|
|
|
109
109
|
let inTriangle = select(col >= row, col <= row, params.uplo == 0u);
|
|
110
110
|
if (col < params.n && inTriangle) {
|
|
111
111
|
let cIdx = row * params.ldc + col;
|
|
112
|
-
|
|
112
|
+
// BLAS beta==0 semantics: C is written, not accumulated — must not
|
|
113
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
114
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
115
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
113
116
|
}
|
|
114
117
|
}
|
|
115
118
|
}
|
|
@@ -102,7 +102,10 @@ fn main(
|
|
|
102
102
|
let inTriangle = select(col >= row, col <= row, params.uplo == 0u);
|
|
103
103
|
if (col < params.n && inTriangle) {
|
|
104
104
|
let cIdx = row * params.ldc + col;
|
|
105
|
-
|
|
105
|
+
// BLAS beta==0 semantics: C is written, not accumulated — must not
|
|
106
|
+
// read C (stale NaN/Inf bits would survive 0 * C as NaN).
|
|
107
|
+
let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
|
|
108
|
+
C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
|
|
106
109
|
}
|
|
107
110
|
}
|
|
108
111
|
}
|
package/src/shaders/sgemv_n.wgsl
CHANGED
|
@@ -67,7 +67,9 @@ fn main(
|
|
|
67
67
|
|
|
68
68
|
if lid.x == 0u {
|
|
69
69
|
let yi = row * params.incy;
|
|
70
|
-
y
|
|
70
|
+
// BLAS beta==0 semantics: y is written, not accumulated — must not read y.
|
|
71
|
+
let acc = params.alpha * scratch[0];
|
|
72
|
+
y[yi] = select(acc, acc + params.beta * y[yi], params.beta != 0.0);
|
|
71
73
|
}
|
|
72
74
|
// All 64 threads must agree before the next row reuses scratch[].
|
|
73
75
|
workgroupBarrier();
|
package/src/shaders/sgemv_t.wgsl
CHANGED
|
@@ -60,6 +60,8 @@ fn main(
|
|
|
60
60
|
acc0 += A[k * params.lda + col] * x[k * params.incx];
|
|
61
61
|
}
|
|
62
62
|
let yi = col * params.incy;
|
|
63
|
-
|
|
63
|
+
// BLAS beta==0 semantics: y is written, not accumulated — must not read y.
|
|
64
|
+
let acc = params.alpha * (acc0 + acc1 + acc2 + acc3);
|
|
65
|
+
y[yi] = select(acc, acc + params.beta * y[yi], params.beta != 0.0);
|
|
64
66
|
}
|
|
65
67
|
}
|
package/src/shaders/snrm2.wgsl
CHANGED
|
@@ -1,9 +1,22 @@
|
|
|
1
|
-
// snrm2: result = sqrt(sum(x[i] * x[i]))
|
|
2
|
-
//
|
|
1
|
+
// snrm2: result = sqrt(sum(x[i] * x[i])), computed via scaled accumulation
|
|
2
|
+
// (Blue's algorithm / reference BLAS's SLASSQ) rather than naive squaring —
|
|
3
|
+
// naive `sum += x_i * x_i` overflows to inf for |x_i| ≳ 1.8e19 (f32's
|
|
4
|
+
// squaring range is only sqrt(f32_max)) and loses precision on tiny
|
|
5
|
+
// magnitudes squaring into the denormal range. Running state is (scale,
|
|
6
|
+
// ssq) with true-sum-of-squares == scale² · ssq: scale tracks the largest
|
|
7
|
+
// |x_i| seen so far, and every other contribution is expressed *relative
|
|
8
|
+
// to* scale (never squared in absolute terms), so ssq stays near 1
|
|
9
|
+
// regardless of x's magnitude range. Merging two independent partials
|
|
10
|
+
// (ssqMerge) is associative, so this composes with the same 4-way-ILP +
|
|
11
|
+
// tree-reduction shape every other Level 1 reduction here uses — see
|
|
12
|
+
// reduction/scaledSum.wgsl for the pass-2 counterpart, which finishes with
|
|
13
|
+
// scale·sqrt(ssq).
|
|
14
|
+
// pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/scaledSum.wgsl.
|
|
3
15
|
|
|
4
|
-
@group(0) @binding(0) var<storage, read> x:
|
|
5
|
-
@group(0) @binding(1) var<storage, read_write>
|
|
6
|
-
@group(0) @binding(2) var<
|
|
16
|
+
@group(0) @binding(0) var<storage, read> x: array<f32>;
|
|
17
|
+
@group(0) @binding(1) var<storage, read_write> partialsScale: array<f32>;
|
|
18
|
+
@group(0) @binding(2) var<storage, read_write> partialsSsq: array<f32>;
|
|
19
|
+
@group(0) @binding(3) var<uniform> params: Params;
|
|
7
20
|
|
|
8
21
|
struct Params {
|
|
9
22
|
n: u32,
|
|
@@ -12,7 +25,36 @@ struct Params {
|
|
|
12
25
|
|
|
13
26
|
const WGS: u32 = 64;
|
|
14
27
|
|
|
15
|
-
|
|
28
|
+
struct ScaleSsq {
|
|
29
|
+
scale: f32,
|
|
30
|
+
ssq: f32,
|
|
31
|
+
}
|
|
32
|
+
|
|
33
|
+
// Folds one more |value| into a running (scale, ssq) pair.
|
|
34
|
+
fn ssqAccum(acc: ScaleSsq, absxi: f32) -> ScaleSsq {
|
|
35
|
+
if (absxi == 0.0) { return acc; }
|
|
36
|
+
if (absxi > acc.scale) {
|
|
37
|
+
let r = acc.scale / absxi; // 0/absxi == 0 on the first nonzero value — safe
|
|
38
|
+
return ScaleSsq(absxi, 1.0 + acc.ssq * r * r);
|
|
39
|
+
}
|
|
40
|
+
let r = absxi / acc.scale; // reached only once acc.scale > 0 (absxi <= acc.scale and absxi > 0)
|
|
41
|
+
return ScaleSsq(acc.scale, acc.ssq + r * r);
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
// Associative merge of two independent (scale, ssq) partials — lets this
|
|
45
|
+
// compose with a tree reduction exactly like a plain sum would.
|
|
46
|
+
fn ssqMerge(a: ScaleSsq, b: ScaleSsq) -> ScaleSsq {
|
|
47
|
+
if (a.scale == 0.0 && b.scale == 0.0) { return ScaleSsq(0.0, 1.0); }
|
|
48
|
+
if (a.scale >= b.scale) {
|
|
49
|
+
let r = b.scale / a.scale;
|
|
50
|
+
return ScaleSsq(a.scale, a.ssq + b.ssq * r * r);
|
|
51
|
+
}
|
|
52
|
+
let r = a.scale / b.scale;
|
|
53
|
+
return ScaleSsq(b.scale, b.ssq + a.ssq * r * r);
|
|
54
|
+
}
|
|
55
|
+
|
|
56
|
+
var<workgroup> tileScale: array<f32, 64>;
|
|
57
|
+
var<workgroup> tileSsq: array<f32, 64>;
|
|
16
58
|
|
|
17
59
|
@compute @workgroup_size(64)
|
|
18
60
|
fn main(
|
|
@@ -21,36 +63,43 @@ fn main(
|
|
|
21
63
|
@builtin(workgroup_id) wgid: vec3u,
|
|
22
64
|
@builtin(num_workgroups) num_wg: vec3u,
|
|
23
65
|
) {
|
|
24
|
-
var acc0
|
|
25
|
-
var acc1
|
|
26
|
-
var acc2
|
|
27
|
-
var acc3
|
|
66
|
+
var acc0 = ScaleSsq(0.0, 1.0);
|
|
67
|
+
var acc1 = ScaleSsq(0.0, 1.0);
|
|
68
|
+
var acc2 = ScaleSsq(0.0, 1.0);
|
|
69
|
+
var acc3 = ScaleSsq(0.0, 1.0);
|
|
28
70
|
|
|
29
71
|
let stride = num_wg.x * WGS;
|
|
30
72
|
let n4_floor = (params.n / (4u * stride)) * (4u * stride);
|
|
31
73
|
|
|
32
74
|
for (var id = gid.x; id < n4_floor; id += 4u * stride) {
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
acc0 += v0 * v0;
|
|
38
|
-
acc1 += v1 * v1;
|
|
39
|
-
acc2 += v2 * v2;
|
|
40
|
-
acc3 += v3 * v3;
|
|
75
|
+
acc0 = ssqAccum(acc0, abs(x[ id * params.x_inc]));
|
|
76
|
+
acc1 = ssqAccum(acc1, abs(x[(id + stride) * params.x_inc]));
|
|
77
|
+
acc2 = ssqAccum(acc2, abs(x[(id + 2u * stride) * params.x_inc]));
|
|
78
|
+
acc3 = ssqAccum(acc3, abs(x[(id + 3u * stride) * params.x_inc]));
|
|
41
79
|
}
|
|
42
80
|
for (var id = n4_floor + gid.x; id < params.n; id += stride) {
|
|
43
|
-
|
|
44
|
-
acc0 += v * v;
|
|
81
|
+
acc0 = ssqAccum(acc0, abs(x[id * params.x_inc]));
|
|
45
82
|
}
|
|
46
83
|
|
|
47
|
-
|
|
84
|
+
let combined = ssqMerge(ssqMerge(acc0, acc1), ssqMerge(acc2, acc3));
|
|
85
|
+
tileScale[lid.x] = combined.scale;
|
|
86
|
+
tileSsq[lid.x] = combined.ssq;
|
|
48
87
|
workgroupBarrier();
|
|
49
88
|
|
|
50
89
|
for (var s = WGS / 2u; s > 0u; s >>= 1u) {
|
|
51
|
-
if (lid.x < s) {
|
|
90
|
+
if (lid.x < s) {
|
|
91
|
+
let merged = ssqMerge(
|
|
92
|
+
ScaleSsq(tileScale[lid.x], tileSsq[lid.x]),
|
|
93
|
+
ScaleSsq(tileScale[lid.x + s], tileSsq[lid.x + s]),
|
|
94
|
+
);
|
|
95
|
+
tileScale[lid.x] = merged.scale;
|
|
96
|
+
tileSsq[lid.x] = merged.ssq;
|
|
97
|
+
}
|
|
52
98
|
workgroupBarrier();
|
|
53
99
|
}
|
|
54
100
|
|
|
55
|
-
if (lid.x == 0u) {
|
|
101
|
+
if (lid.x == 0u) {
|
|
102
|
+
partialsScale[wgid.x] = tileScale[0];
|
|
103
|
+
partialsSsq[wgid.x] = tileSsq[0];
|
|
104
|
+
}
|
|
56
105
|
}
|
package/src/shaders/ssymv.wgsl
CHANGED
|
@@ -63,7 +63,9 @@ fn main(
|
|
|
63
63
|
}
|
|
64
64
|
|
|
65
65
|
if lid.x == 0u {
|
|
66
|
-
|
|
66
|
+
// BLAS beta==0 semantics: y is written, not accumulated — must not read y.
|
|
67
|
+
let acc = params.alpha * scratch[0];
|
|
68
|
+
y[i * params.incy] = select(acc, acc + params.beta * y[i * params.incy], params.beta != 0.0);
|
|
67
69
|
}
|
|
68
70
|
}
|
|
69
71
|
}
|
package/src/snrm2/snrm2.d.mts
CHANGED
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
|
-
* Computes the Euclidean norm of a vector: result = sqrt
|
|
4
|
+
* Computes the Euclidean norm of a vector: $$\text{result} = \sqrt{\sum_{i} x_i^2}$$
|
|
5
5
|
*
|
|
6
6
|
* {@includeCode ../../examples/snrm2/snrm2.js}
|
|
7
7
|
*
|
|
@@ -24,7 +24,7 @@ export declare function snrm2(
|
|
|
24
24
|
): Promise<{ nrm2: number } | { nrm2: number; gpuTimeMs: number }>;
|
|
25
25
|
|
|
26
26
|
/**
|
|
27
|
-
* Computes the Euclidean norm of a vector: result = sqrt
|
|
27
|
+
* Computes the Euclidean norm of a vector: $$\text{result} = \sqrt{\sum_{i} x_i^2}$$
|
|
28
28
|
*
|
|
29
29
|
* {@includeCode ../../examples/snrm2/gpu.snrm2.js}
|
|
30
30
|
*
|
package/src/snrm2/snrm2.mjs
CHANGED
|
@@ -12,14 +12,14 @@ import { extractTimestamp } from "../util/benchmark.mjs";
|
|
|
12
12
|
import { extractResult } from "../util/result.mjs";
|
|
13
13
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
14
14
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
15
|
-
|
|
16
|
-
|
|
15
|
+
import { WGS } from "../util/constants.mjs";
|
|
16
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
17
17
|
|
|
18
18
|
export async function snrm2(device, n, x, incx) {
|
|
19
19
|
const xIsGpu = x instanceof GpuVector;
|
|
20
20
|
|
|
21
|
-
|
|
22
|
-
|
|
21
|
+
requireGpuDevice(device);
|
|
22
|
+
requireSameDevice(device, "snrm2", { x });
|
|
23
23
|
if (!Number.isInteger(n) || !Number.isInteger(incx))
|
|
24
24
|
throw new Error("n and incx must be integers.");
|
|
25
25
|
if (incx <= 0) throw new Error("incx must be positive.");
|
|
@@ -32,19 +32,31 @@ export async function snrm2(device, n, x, incx) {
|
|
|
32
32
|
);
|
|
33
33
|
|
|
34
34
|
const pipelineMain = await getPipeline(device, "snrm2");
|
|
35
|
-
const pipelineReduce = await getPipeline(device, "reduction/
|
|
35
|
+
const pipelineReduce = await getPipeline(device, "reduction/scaledSum");
|
|
36
36
|
|
|
37
37
|
let xBuffer = null;
|
|
38
|
-
let
|
|
38
|
+
let partialsScaleBuffer = null;
|
|
39
|
+
let partialsSsqBuffer = null;
|
|
39
40
|
let resultBuffer = null;
|
|
40
41
|
let paramsBuffer = null;
|
|
41
42
|
let readBuffer = null;
|
|
42
43
|
|
|
43
44
|
try {
|
|
44
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "snrm2-x", false);
|
|
45
|
-
|
|
46
|
-
|
|
45
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "snrm2-x", false);
|
|
46
|
+
// 2*WGS partial (scale, ssq) pairs — see snrm2.wgsl for what they represent.
|
|
47
|
+
partialsScaleBuffer = createStorageBuffer(
|
|
48
|
+
device,
|
|
49
|
+
2 * WGS * 4,
|
|
50
|
+
"snrm2-partials-scale",
|
|
51
|
+
);
|
|
52
|
+
partialsSsqBuffer = createStorageBuffer(
|
|
53
|
+
device,
|
|
54
|
+
2 * WGS * 4,
|
|
55
|
+
"snrm2-partials-ssq",
|
|
56
|
+
);
|
|
57
|
+
resultBuffer = createResultBuffer(device, 4, "snrm2-result"); // final f32 scalar
|
|
47
58
|
paramsBuffer = createParamsBuffer(
|
|
59
|
+
device,
|
|
48
60
|
[
|
|
49
61
|
{ value: n, type: "u32" },
|
|
50
62
|
{ value: incx, type: "u32" },
|
|
@@ -52,53 +64,59 @@ export async function snrm2(device, n, x, incx) {
|
|
|
52
64
|
"snrm2-params",
|
|
53
65
|
);
|
|
54
66
|
|
|
55
|
-
const bgMain = createBindGroup(pipelineMain.getBindGroupLayout(0), [
|
|
67
|
+
const bgMain = createBindGroup(device, pipelineMain.getBindGroupLayout(0), [
|
|
56
68
|
xBuffer,
|
|
57
|
-
|
|
69
|
+
partialsScaleBuffer,
|
|
70
|
+
partialsSsqBuffer,
|
|
58
71
|
paramsBuffer,
|
|
59
72
|
]);
|
|
60
73
|
const { commandEncoder: enc1, ts: ts1 } = runComputePass(
|
|
74
|
+
device,
|
|
61
75
|
pipelineMain,
|
|
62
76
|
bgMain,
|
|
63
77
|
2 * WGS,
|
|
64
78
|
); // dispatch 2*WGS workgroups
|
|
65
79
|
|
|
66
|
-
submit(enc1);
|
|
80
|
+
submit(device, enc1);
|
|
67
81
|
|
|
68
|
-
const bgReduce = createBindGroup(
|
|
69
|
-
|
|
70
|
-
|
|
71
|
-
|
|
82
|
+
const bgReduce = createBindGroup(
|
|
83
|
+
device,
|
|
84
|
+
pipelineReduce.getBindGroupLayout(0),
|
|
85
|
+
[partialsScaleBuffer, partialsSsqBuffer, resultBuffer],
|
|
86
|
+
);
|
|
72
87
|
const { commandEncoder: enc2, ts: ts2 } = runComputePass(
|
|
88
|
+
device,
|
|
73
89
|
pipelineReduce,
|
|
74
90
|
bgReduce,
|
|
75
91
|
1,
|
|
76
92
|
); // reduce partials to single result
|
|
77
|
-
readBuffer = stageReadback(enc2, resultBuffer);
|
|
93
|
+
readBuffer = stageReadback(device, enc2, resultBuffer);
|
|
78
94
|
|
|
79
|
-
submit(enc2);
|
|
95
|
+
submit(device, enc2);
|
|
80
96
|
|
|
81
97
|
const resultPromise = extractResult(readBuffer, Float32Array);
|
|
82
98
|
readBuffer = null; // ownership transferred — extractResult's own finally destroys it
|
|
83
99
|
|
|
84
|
-
const [gpuTime1, gpuTime2,
|
|
100
|
+
const [gpuTime1, gpuTime2, resultArr] = await Promise.all([
|
|
85
101
|
extractTimestamp(ts1),
|
|
86
102
|
extractTimestamp(ts2),
|
|
87
103
|
resultPromise,
|
|
88
104
|
]);
|
|
89
105
|
|
|
90
|
-
//
|
|
91
|
-
|
|
106
|
+
// reduction/scaledSum.wgsl already computes scale·sqrt(ssq) on the GPU —
|
|
107
|
+
// unlike the old naive-sum version, there's no separate sqrt step here.
|
|
108
|
+
const nrm2 = resultArr[0];
|
|
92
109
|
|
|
93
110
|
if (gpuTime1 !== undefined && gpuTime2 !== undefined)
|
|
94
111
|
return { nrm2, gpuTimeMs: gpuTime1 + gpuTime2 };
|
|
95
112
|
return { nrm2 };
|
|
96
113
|
} finally {
|
|
97
114
|
if (!xIsGpu && xBuffer) destroyBuffers(xBuffer);
|
|
98
|
-
if (
|
|
115
|
+
if (partialsScaleBuffer) destroyBuffers(partialsScaleBuffer);
|
|
116
|
+
if (partialsSsqBuffer) destroyBuffers(partialsSsqBuffer);
|
|
99
117
|
if (resultBuffer) destroyBuffers(resultBuffer);
|
|
100
118
|
if (paramsBuffer) destroyBuffers(paramsBuffer);
|
|
101
|
-
// Only reached if submit(enc2) threw before ownership was transferred above.
|
|
119
|
+
// Only reached if submit(device, enc2) threw before ownership was transferred above.
|
|
102
120
|
if (readBuffer) destroyBuffers(readBuffer);
|
|
103
121
|
}
|
|
104
122
|
}
|
package/src/srot/srot.d.mts
CHANGED
|
@@ -2,8 +2,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
|
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
4
|
* Applies a Givens plane rotation to vectors x and y:
|
|
5
|
-
*
|
|
6
|
-
* y = -s*x + c*y
|
|
5
|
+
* $$\begin{aligned} x &\leftarrow cx + sy \\\\ y &\leftarrow -sx + cy \end{aligned}$$
|
|
7
6
|
*
|
|
8
7
|
* {@includeCode ../../examples/srot/srot.js}
|
|
9
8
|
*
|
|
@@ -34,8 +33,7 @@ export declare function srot(
|
|
|
34
33
|
|
|
35
34
|
/**
|
|
36
35
|
* Applies a Givens plane rotation to vectors x and y:
|
|
37
|
-
*
|
|
38
|
-
* y = -s*x + c*y
|
|
36
|
+
* $$\begin{aligned} x &\leftarrow cx + sy \\\\ y &\leftarrow -sx + cy \end{aligned}$$
|
|
39
37
|
*
|
|
40
38
|
* {@includeCode ../../examples/srot/gpu.srot.js}
|
|
41
39
|
*
|
package/src/srot/srot.mjs
CHANGED
|
@@ -11,13 +11,14 @@ import { extractResult } from "../util/result.mjs";
|
|
|
11
11
|
import { getPipeline } from "../util/pipeline.mjs";
|
|
12
12
|
import { calcWorkgroups } from "../util/workgroup.mjs";
|
|
13
13
|
import { GpuVector } from "../classes/GpuVector.mjs";
|
|
14
|
+
import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
|
|
14
15
|
|
|
15
16
|
export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
16
17
|
const xIsGpu = x instanceof GpuVector;
|
|
17
18
|
const yIsGpu = y instanceof GpuVector;
|
|
18
19
|
|
|
19
|
-
|
|
20
|
-
|
|
20
|
+
requireGpuDevice(device);
|
|
21
|
+
requireSameDevice(device, "srot", { x, y });
|
|
21
22
|
if (
|
|
22
23
|
!Number.isInteger(n) ||
|
|
23
24
|
!Number.isInteger(incx) ||
|
|
@@ -26,7 +27,8 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
26
27
|
throw new Error("n, incx, and incy must be integers.");
|
|
27
28
|
if (typeof c !== "number") throw new Error("c must be a number.");
|
|
28
29
|
if (typeof s !== "number") throw new Error("s must be a number.");
|
|
29
|
-
if (Number.isNaN(c) || Number.isNaN(s))
|
|
30
|
+
if (Number.isNaN(c) || Number.isNaN(s))
|
|
31
|
+
throw new Error("c and s must not be NaN.");
|
|
30
32
|
if (!Number.isFinite(c)) throw new Error("c must be finite.");
|
|
31
33
|
if (!Number.isFinite(s)) throw new Error("s must be finite.");
|
|
32
34
|
if (incx <= 0 || incy <= 0)
|
|
@@ -58,9 +60,10 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
58
60
|
let readY = null;
|
|
59
61
|
|
|
60
62
|
try {
|
|
61
|
-
xBuffer = xIsGpu ? x._buf : uploadBuffer(x, "srot-x", true);
|
|
62
|
-
yBuffer = yIsGpu ? y._buf : uploadBuffer(y, "srot-y", true);
|
|
63
|
+
xBuffer = xIsGpu ? x._buf : uploadBuffer(device, x, "srot-x", true);
|
|
64
|
+
yBuffer = yIsGpu ? y._buf : uploadBuffer(device, y, "srot-y", true);
|
|
63
65
|
paramsBuffer = createParamsBuffer(
|
|
66
|
+
device,
|
|
64
67
|
[
|
|
65
68
|
{ value: n, type: "u32" },
|
|
66
69
|
{ value: c, type: "f32" },
|
|
@@ -71,23 +74,25 @@ export async function srot(device, n, x, incx, y, incy, c, s) {
|
|
|
71
74
|
"srot-params",
|
|
72
75
|
);
|
|
73
76
|
|
|
74
|
-
const bindGroup = createBindGroup(pipeline.getBindGroupLayout(0), [
|
|
77
|
+
const bindGroup = createBindGroup(device, pipeline.getBindGroupLayout(0), [
|
|
75
78
|
xBuffer,
|
|
76
79
|
yBuffer,
|
|
77
80
|
paramsBuffer,
|
|
78
81
|
]);
|
|
79
82
|
const { commandEncoder, ts } = runComputePass(
|
|
83
|
+
device,
|
|
80
84
|
pipeline,
|
|
81
85
|
bindGroup,
|
|
82
|
-
calcWorkgroups(n),
|
|
86
|
+
calcWorkgroups(device, n),
|
|
83
87
|
);
|
|
84
|
-
readX = xIsGpu ? null : stageReadback(commandEncoder, xBuffer);
|
|
85
|
-
readY = yIsGpu ? null : stageReadback(commandEncoder, yBuffer);
|
|
86
|
-
submit(commandEncoder);
|
|
88
|
+
readX = xIsGpu ? null : stageReadback(device, commandEncoder, xBuffer);
|
|
89
|
+
readY = yIsGpu ? null : stageReadback(device, commandEncoder, yBuffer);
|
|
90
|
+
submit(device, commandEncoder);
|
|
87
91
|
|
|
88
92
|
const gpuTimeMs = await extractTimestamp(ts);
|
|
89
93
|
|
|
90
|
-
if (xIsGpu
|
|
94
|
+
if (xIsGpu) {
|
|
95
|
+
// xIsGpu === yIsGpu, enforced above
|
|
91
96
|
if (gpuTimeMs !== undefined) return { gpuTimeMs };
|
|
92
97
|
return {};
|
|
93
98
|
}
|
package/src/srotm/srotm.d.mts
CHANGED
|
@@ -2,8 +2,7 @@ import { GpuVector } from "../classes/GpuVector.mjs";
|
|
|
2
2
|
|
|
3
3
|
/**
|
|
4
4
|
* Applies a modified Givens plane rotation H to vectors x and y:
|
|
5
|
-
*
|
|
6
|
-
* y = H[1][0]*x + H[1][1]*y
|
|
5
|
+
* $$\begin{pmatrix} x \\\\ y \end{pmatrix} \leftarrow \begin{pmatrix} h_{11} & h_{12} \\\\ h_{21} & h_{22} \end{pmatrix} \begin{pmatrix} x \\\\ y \end{pmatrix}$$
|
|
7
6
|
*
|
|
8
7
|
* {@includeCode ../../examples/srotm/srotm.js}
|
|
9
8
|
*
|
|
@@ -33,8 +32,7 @@ export declare function srotm(
|
|
|
33
32
|
|
|
34
33
|
/**
|
|
35
34
|
* Applies a modified Givens plane rotation H to vectors x and y:
|
|
36
|
-
*
|
|
37
|
-
* y = H[1][0]*x + H[1][1]*y
|
|
35
|
+
* $$\begin{pmatrix} x \\\\ y \end{pmatrix} \leftarrow \begin{pmatrix} h_{11} & h_{12} \\\\ h_{21} & h_{22} \end{pmatrix} \begin{pmatrix} x \\\\ y \end{pmatrix}$$
|
|
38
36
|
*
|
|
39
37
|
* {@includeCode ../../examples/srotm/gpu.srotm.js}
|
|
40
38
|
*
|