wgblas 2.0.0 → 2.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.
Files changed (66) hide show
  1. package/README.md +18 -18
  2. package/dist/wgblas.browser.js +1273 -1239
  3. package/index.d.mts +38 -6
  4. package/package.json +2 -1
  5. package/src/classes/GpuMatrix.mjs +17 -10
  6. package/src/classes/GpuVector.mjs +28 -10
  7. package/src/dasum/dasum.d.mts +4 -4
  8. package/src/dasum/dasum.mjs +19 -17
  9. package/src/devdocs.mjs +13 -0
  10. package/src/idamax/idamax.d.mts +20 -2
  11. package/src/idamax/idamax.mjs +18 -16
  12. package/src/init.mjs +114 -56
  13. package/src/isamax/isamax.d.mts +20 -2
  14. package/src/isamax/isamax.mjs +16 -14
  15. package/src/random/random.d.mts +1 -0
  16. package/src/sasum/sasum.d.mts +2 -2
  17. package/src/sasum/sasum.mjs +15 -13
  18. package/src/saxpy/saxpy.d.mts +2 -2
  19. package/src/saxpy/saxpy.mjs +10 -8
  20. package/src/scopy/scopy.d.mts +2 -2
  21. package/src/scopy/scopy.mjs +10 -8
  22. package/src/sdot/sdot.d.mts +2 -2
  23. package/src/sdot/sdot.mjs +16 -14
  24. package/src/sgemm/sgemm.mjs +28 -15
  25. package/src/sgemmtr/sgemmtr.mjs +16 -15
  26. package/src/sgemv/sgemv.mjs +38 -26
  27. package/src/sger/sger.mjs +10 -8
  28. package/src/shaders/index.mjs +164 -14
  29. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  30. package/src/shaders/sgemm_large.wgsl +107 -18
  31. package/src/shaders/sgemm_small.wgsl +115 -15
  32. package/src/shaders/sgemmtr_large.wgsl +4 -1
  33. package/src/shaders/sgemmtr_small.wgsl +4 -1
  34. package/src/shaders/sgemv_n.wgsl +3 -1
  35. package/src/shaders/sgemv_t.wgsl +3 -1
  36. package/src/shaders/snrm2.wgsl +72 -23
  37. package/src/shaders/ssymv.wgsl +3 -1
  38. package/src/snrm2/snrm2.d.mts +2 -2
  39. package/src/snrm2/snrm2.mjs +33 -21
  40. package/src/srot/srot.d.mts +2 -4
  41. package/src/srot/srot.mjs +11 -9
  42. package/src/srotm/srotm.d.mts +2 -4
  43. package/src/srotm/srotm.mjs +19 -10
  44. package/src/sscal/sscal.d.mts +3 -3
  45. package/src/sscal/sscal.mjs +12 -10
  46. package/src/sswap/sswap.d.mts +2 -2
  47. package/src/sswap/sswap.mjs +11 -9
  48. package/src/ssymm/ssymm.mjs +31 -22
  49. package/src/ssymv/ssymv.mjs +10 -8
  50. package/src/ssyr/ssyr.mjs +9 -7
  51. package/src/ssyr2/ssyr2.mjs +10 -8
  52. package/src/ssyr2k/ssyr2k.mjs +18 -17
  53. package/src/ssyrk/ssyrk.mjs +18 -17
  54. package/src/strmm/strmm.mjs +47 -32
  55. package/src/strmv/strmv.mjs +10 -8
  56. package/src/strsm/strsm.mjs +54 -36
  57. package/src/strsv/strsv.mjs +16 -12
  58. package/src/util/benchmark.mjs +4 -6
  59. package/src/util/bindgroup.mjs +1 -3
  60. package/src/util/buffer.mjs +113 -19
  61. package/src/util/compute.mjs +6 -9
  62. package/src/util/constants.mjs +57 -0
  63. package/src/util/device.mjs +34 -0
  64. package/src/util/pipeline.mjs +5 -6
  65. package/src/util/workgroup.mjs +55 -7
  66. package/src/shaders/browser-shaders.mjs +0 -81
@@ -0,0 +1,65 @@
1
+ // scaledSum reduction: collapses 2*WGS (scale, ssq) partials from
2
+ // snrm2.wgsl into the final norm — sqrt(scale² · ssq) == scale · sqrt(ssq).
3
+ // Mirrors reduction/sum.wgsl's shape exactly, merging via ssqMerge (see
4
+ // snrm2.wgsl for the derivation) instead of plain `+`, and taking the final
5
+ // sqrt here rather than on the CPU — unlike sasum/sdot's plain sum, "sum of
6
+ // squares" isn't a meaningful standalone value to hand back, only
7
+ // scale·sqrt(ssq) is.
8
+ // dispatch: 1 workgroup of WGS threads.
9
+ // partialsScale/partialsSsq must have exactly 2*WGS entries each.
10
+
11
+ @group(0) @binding(0) var<storage, read> partialsScale: array<f32>;
12
+ @group(0) @binding(1) var<storage, read> partialsSsq: array<f32>;
13
+ @group(0) @binding(2) var<storage, read_write> result: array<f32>;
14
+
15
+ const WGS: u32 = 64;
16
+
17
+ // True sum-of-squares represented so far == scale² · ssq — see snrm2.wgsl.
18
+ struct ScaleSsq {
19
+ scale: f32,
20
+ ssq: f32,
21
+ }
22
+
23
+ // Associative merge of two independent (scale, ssq) partials.
24
+ fn ssqMerge(a: ScaleSsq, b: ScaleSsq) -> ScaleSsq {
25
+ if (a.scale == 0.0 && b.scale == 0.0) { return ScaleSsq(0.0, 1.0); }
26
+ if (a.scale >= b.scale) {
27
+ let r = b.scale / a.scale;
28
+ return ScaleSsq(a.scale, a.ssq + b.ssq * r * r);
29
+ }
30
+ let r = a.scale / b.scale;
31
+ return ScaleSsq(b.scale, b.ssq + a.ssq * r * r);
32
+ }
33
+
34
+ var<workgroup> tileScale: array<f32, 64>;
35
+ var<workgroup> tileSsq: array<f32, 64>;
36
+
37
+ @compute @workgroup_size(64)
38
+ fn reduce_scaled(
39
+ @builtin(local_invocation_id) lid: vec3u,
40
+ ) {
41
+ let i = lid.x;
42
+ let merged0 = ssqMerge(
43
+ ScaleSsq(partialsScale[i], partialsSsq[i]),
44
+ ScaleSsq(partialsScale[i + WGS], partialsSsq[i + WGS]),
45
+ );
46
+ tileScale[i] = merged0.scale;
47
+ tileSsq[i] = merged0.ssq;
48
+ workgroupBarrier();
49
+
50
+ for (var s = WGS / 2u; s > 0u; s >>= 1u) {
51
+ if (i < s) {
52
+ let merged = ssqMerge(
53
+ ScaleSsq(tileScale[i], tileSsq[i]),
54
+ ScaleSsq(tileScale[i + s], tileSsq[i + s]),
55
+ );
56
+ tileScale[i] = merged.scale;
57
+ tileSsq[i] = merged.ssq;
58
+ }
59
+ workgroupBarrier();
60
+ }
61
+
62
+ if (i == 0u) {
63
+ result[0] = tileScale[0] * sqrt(tileSsq[0]);
64
+ }
65
+ }
@@ -6,9 +6,13 @@
6
6
  // single-tier baseline at n=512, +84% at n=1024. But BM=64 loses to BM=32
7
7
  // below a 6x6=36 workgroup grid (not enough workgroups to fill the GPU at
8
8
  // that tile size), hence the two-tier split rather than one global config.
9
- // Neither vectorized loads (kernel 6) nor warp-tiling (kernel 10) beat this
10
- // at the sizes tried, including warp-tiled variants in the same sweep at
11
- // BM=64/128.
9
+ //
10
+ // A and B are bound twice — scalar array<f32> and array<vec4<f32>> views of
11
+ // the same GPUBuffer (see vec4ViewBinding) — so each tile load can issue
12
+ // 16-byte vector reads along op(A)/op(B)'s contiguous dimension when the
13
+ // stride allows it (stride % 4 == 0 keeps every row base 16-byte aligned).
14
+ // Transposed or odd-stride operands take the scalar path; both paths
15
+ // zero-fill out-of-bounds components identically.
12
16
 
13
17
  const BM: u32 = 64u;
14
18
  const BN: u32 = 64u;
@@ -21,9 +25,11 @@ const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 128
21
25
  const STRIDE_A: u32 = NUM_THREADS / BK; // rows of As covered per load-loop step
22
26
  const STRIDE_B: u32 = NUM_THREADS / BN; // rows of Bs covered per load-loop step
23
27
 
24
- @group(0) @binding(0) var<storage, read> A: array<f32>;
25
- @group(0) @binding(1) var<storage, read> B: array<f32>;
26
- @group(0) @binding(2) var<storage, read_write> C: array<f32>;
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>;
27
33
 
28
34
  struct Params {
29
35
  m: u32,
@@ -36,9 +42,11 @@ struct Params {
36
42
  ldc: u32,
37
43
  transA: u32, // 0 = no-transpose, 1 = transpose
38
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
39
47
  }
40
48
 
41
- @group(0) @binding(3) var<uniform> params: Params;
49
+ @group(0) @binding(5) var<uniform> params: Params;
42
50
 
43
51
  var<workgroup> As: array<f32, BM * BK>;
44
52
  var<workgroup> Bs: array<f32, BK * BN>;
@@ -70,17 +78,95 @@ fn main(
70
78
 
71
79
  let numTiles = (params.k + BK - 1u) / BK;
72
80
  for (var t = 0u; t < numTiles; t++) {
73
- for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
74
- let gRowA = blockRow + innerRowA + loadOffset;
75
- let gColA = t * BK + innerColA;
76
- let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
77
- As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
81
+ // ── Load the BM×BK A tile into As (vectorized along op(A)'s fast dim
82
+ // when lda allows; every branch here is dispatch-uniform) ──
83
+ if (params.useVecA == 1u && params.transA == 0u) {
84
+ // No-transpose: columns contiguous. Each thread loads one vec4 of 4
85
+ // columns; 64 rows × 2 column-lanes = NUM_THREADS exactly, single pass.
86
+ let r4 = tid / (BK / 4u);
87
+ let c4 = tid % (BK / 4u);
88
+ let gRow = blockRow + r4;
89
+ let gCol = t * BK + c4 * 4u;
90
+ var v = A4[(gRow * params.lda + gCol) / 4u];
91
+ let rowOK = gRow < params.m;
92
+ v.x = select(0.0, v.x, rowOK && gCol < params.k);
93
+ v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.k);
94
+ v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.k);
95
+ v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.k);
96
+ As[r4 * BK + c4 * 4u] = v.x;
97
+ As[r4 * BK + c4 * 4u + 1u] = v.y;
98
+ As[r4 * BK + c4 * 4u + 2u] = v.z;
99
+ As[r4 * BK + c4 * 4u + 3u] = v.w;
100
+ } else if (params.useVecA == 1u && params.transA != 0u) {
101
+ // Transpose: rows contiguous within a column. Each thread loads one
102
+ // vec4 of 4 rows; 16 row-lanes × 8 columns = NUM_THREADS, single pass.
103
+ let r4 = tid % (BM / 4u);
104
+ let c = tid / (BM / 4u);
105
+ let gRow = blockRow + r4 * 4u;
106
+ let gCol = t * BK + c;
107
+ var v = A4[(gCol * params.lda + gRow) / 4u];
108
+ let colOK = gCol < params.k;
109
+ v.x = select(0.0, v.x, colOK && gRow < params.m);
110
+ v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.m);
111
+ v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.m);
112
+ v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.m);
113
+ As[(r4 * 4u) * BK + c] = v.x;
114
+ As[(r4 * 4u + 1u) * BK + c] = v.y;
115
+ As[(r4 * 4u + 2u) * BK + c] = v.z;
116
+ As[(r4 * 4u + 3u) * BK + c] = v.w;
117
+ } else {
118
+ // Scalar fallback: odd stride or unhandled orientation.
119
+ for (var loadOffset = 0u; loadOffset < BM; loadOffset += STRIDE_A) {
120
+ let gRowA = blockRow + innerRowA + loadOffset;
121
+ let gColA = t * BK + innerColA;
122
+ let aIdx = select(gRowA * params.lda + gColA, gColA * params.lda + gRowA, params.transA != 0u);
123
+ As[(innerRowA + loadOffset) * BK + innerColA] = select(0.0, A[aIdx], gRowA < params.m && gColA < params.k);
124
+ }
78
125
  }
79
- for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
80
- let gRowB = t * BK + innerRowB + loadOffset;
81
- let gColB = blockCol + innerColB;
82
- let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
83
- Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
126
+
127
+ // ── Load the BK×BN B tile into Bs ──
128
+ if (params.useVecB == 1u && params.transB == 0u) {
129
+ // No-transpose: columns contiguous. 8 rows × 16 column-lanes cover the
130
+ // tile in one pass (BK = NUM_THREADS / (BN/4)).
131
+ let r = tid / (BN / 4u);
132
+ let c4 = tid % (BN / 4u);
133
+ let gRow = t * BK + r;
134
+ let gCol = blockCol + c4 * 4u;
135
+ var v = B4[(gRow * params.ldb + gCol) / 4u];
136
+ let rowOK = gRow < params.k;
137
+ v.x = select(0.0, v.x, rowOK && gCol < params.n);
138
+ v.y = select(0.0, v.y, rowOK && (gCol + 1u) < params.n);
139
+ v.z = select(0.0, v.z, rowOK && (gCol + 2u) < params.n);
140
+ v.w = select(0.0, v.w, rowOK && (gCol + 3u) < params.n);
141
+ Bs[r * BN + c4 * 4u] = v.x;
142
+ Bs[r * BN + c4 * 4u + 1u] = v.y;
143
+ Bs[r * BN + c4 * 4u + 2u] = v.z;
144
+ Bs[r * BN + c4 * 4u + 3u] = v.w;
145
+ } else if (params.useVecB == 1u && params.transB != 0u) {
146
+ // Transpose: rows contiguous within a column. 2 row-lanes × 64 columns
147
+ // cover the tile in one pass (BN = NUM_THREADS / (BK/4)).
148
+ let r4 = tid % (BK / 4u);
149
+ let c = tid / (BK / 4u);
150
+ let gRow = t * BK + r4 * 4u;
151
+ let gCol = blockCol + c;
152
+ var v = B4[(gCol * params.ldb + gRow) / 4u];
153
+ let colOK = gCol < params.n;
154
+ v.x = select(0.0, v.x, colOK && gRow < params.k);
155
+ v.y = select(0.0, v.y, colOK && (gRow + 1u) < params.k);
156
+ v.z = select(0.0, v.z, colOK && (gRow + 2u) < params.k);
157
+ v.w = select(0.0, v.w, colOK && (gRow + 3u) < params.k);
158
+ Bs[(r4 * 4u) * BN + c] = v.x;
159
+ Bs[(r4 * 4u + 1u) * BN + c] = v.y;
160
+ Bs[(r4 * 4u + 2u) * BN + c] = v.z;
161
+ Bs[(r4 * 4u + 3u) * BN + c] = v.w;
162
+ } else {
163
+ // Scalar fallback.
164
+ for (var loadOffset = 0u; loadOffset < BK; loadOffset += STRIDE_B) {
165
+ let gRowB = t * BK + innerRowB + loadOffset;
166
+ let gColB = blockCol + innerColB;
167
+ let bIdx = select(gRowB * params.ldb + gColB, gColB * params.ldb + gRowB, params.transB != 0u);
168
+ Bs[(innerRowB + loadOffset) * BN + innerColB] = select(0.0, B[bIdx], gRowB < params.k && gColB < params.n);
169
+ }
84
170
  }
85
171
 
86
172
  workgroupBarrier();
@@ -109,7 +195,10 @@ fn main(
109
195
  let col = blockCol + threadCol * TN + resIdxN;
110
196
  if (col < params.n) {
111
197
  let cIdx = row * params.ldc + col;
112
- C[cIdx] = params.alpha * threadResults[resIdxM * TN + resIdxN] + params.beta * C[cIdx];
198
+ // BLAS beta==0 semantics: C is written, not accumulated — must not
199
+ // read C (stale NaN/Inf bits would survive 0 * C as NaN).
200
+ let acc = params.alpha * threadResults[resIdxM * TN + resIdxN];
201
+ C[cIdx] = select(acc, acc + params.beta * C[cIdx], params.beta != 0.0);
113
202
  }
114
203
  }
115
204
  }
@@ -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: 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>;
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(3) var<uniform> params: Params;
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
- 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);
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
- 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);
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
- C[cIdx] = params.alpha * threadResults[resIdxM * TN + resIdxN] + params.beta * C[cIdx];
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
- C[cIdx] = params.alpha * threadResults[resIdxM * TN + resIdxN] + params.beta * C[cIdx];
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
- C[cIdx] = params.alpha * threadResults[resIdxM * TN + resIdxN] + params.beta * C[cIdx];
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
  }
@@ -67,7 +67,9 @@ fn main(
67
67
 
68
68
  if lid.x == 0u {
69
69
  let yi = row * params.incy;
70
- y[yi] = params.alpha * scratch[0] + params.beta * y[yi];
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();
@@ -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
- y[yi] = params.alpha * (acc0 + acc1 + acc2 + acc3) + params.beta * y[yi];
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
  }
@@ -1,9 +1,22 @@
1
- // snrm2: result = sqrt(sum(x[i] * x[i]))
2
- // pass 1 dispatches exactly 2 * WGS workgroups; pass 2 uses reduction/sqsum.wgsl.
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: array<f32>;
5
- @group(0) @binding(1) var<storage, read_write> partials: array<f32>;
6
- @group(0) @binding(2) var<uniform> params: Params;
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
- var<workgroup> tile: array<f32, 64>;
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: f32 = 0.0;
25
- var acc1: f32 = 0.0;
26
- var acc2: f32 = 0.0;
27
- var acc3: f32 = 0.0;
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
- let v0 = x[ id * params.x_inc];
34
- let v1 = x[(id + stride) * params.x_inc];
35
- let v2 = x[(id + 2u * stride) * params.x_inc];
36
- let v3 = x[(id + 3u * stride) * params.x_inc];
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
- let v = x[id * params.x_inc];
44
- acc0 += v * v;
81
+ acc0 = ssqAccum(acc0, abs(x[id * params.x_inc]));
45
82
  }
46
83
 
47
- tile[lid.x] = acc0 + acc1 + acc2 + acc3;
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) { tile[lid.x] += tile[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) { partials[wgid.x] = tile[0]; }
101
+ if (lid.x == 0u) {
102
+ partialsScale[wgid.x] = tileScale[0];
103
+ partialsSsq[wgid.x] = tileSsq[0];
104
+ }
56
105
  }
@@ -63,7 +63,9 @@ fn main(
63
63
  }
64
64
 
65
65
  if lid.x == 0u {
66
- y[i * params.incy] = params.alpha * scratch[0] + params.beta * y[i * params.incy];
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
  }
@@ -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(sum(x[i] * x[i]))
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(sum(x[i] * x[i]))
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
  *