wgblas 1.2.1 → 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 (96) hide show
  1. package/LICENSE +1 -1
  2. package/README.md +54 -72
  3. package/dist/wgblas.browser.js +1637 -858
  4. package/index.d.mts +51 -6
  5. package/index.mjs +8 -0
  6. package/package.json +56 -2
  7. package/src/classes/GpuMatrix.mjs +17 -10
  8. package/src/classes/GpuVector.mjs +28 -10
  9. package/src/dasum/dasum.d.mts +6 -6
  10. package/src/dasum/dasum.mjs +25 -21
  11. package/src/devdocs.mjs +13 -0
  12. package/src/idamax/idamax.d.mts +69 -0
  13. package/src/idamax/idamax.mjs +130 -0
  14. package/src/init.mjs +115 -49
  15. package/src/isamax/isamax.d.mts +21 -3
  16. package/src/isamax/isamax.mjs +16 -14
  17. package/src/random/random.d.mts +1 -0
  18. package/src/sasum/sasum.d.mts +3 -3
  19. package/src/sasum/sasum.mjs +15 -13
  20. package/src/saxpy/saxpy.d.mts +3 -3
  21. package/src/saxpy/saxpy.mjs +10 -8
  22. package/src/scopy/scopy.d.mts +3 -3
  23. package/src/scopy/scopy.mjs +10 -8
  24. package/src/sdot/sdot.d.mts +3 -3
  25. package/src/sdot/sdot.mjs +16 -14
  26. package/src/sgemm/sgemm.d.mts +102 -0
  27. package/src/sgemm/sgemm.mjs +208 -0
  28. package/src/sgemmtr/sgemmtr.d.mts +104 -0
  29. package/src/sgemmtr/sgemmtr.mjs +204 -0
  30. package/src/sgemv/sgemv.d.mts +1 -38
  31. package/src/sgemv/sgemv.mjs +42 -26
  32. package/src/sger/sger.d.mts +1 -34
  33. package/src/sger/sger.mjs +12 -8
  34. package/src/shaders/block_transfer.wgsl +42 -0
  35. package/src/shaders/dasum.wgsl +3 -2
  36. package/src/shaders/f64/dekker.wgsl +4 -85
  37. package/src/shaders/f64/utils/abs.wgsl +10 -0
  38. package/src/shaders/f64/utils/add.wgsl +77 -0
  39. package/src/shaders/f64/utils/equal.wgsl +7 -0
  40. package/src/shaders/f64/utils/greater.wgsl +12 -0
  41. package/src/shaders/f64/utils/multiply.wgsl +81 -0
  42. package/src/shaders/idamax.wgsl +96 -0
  43. package/src/shaders/index.mjs +164 -14
  44. package/src/shaders/reduction/argmaxF64.wgsl +50 -0
  45. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  46. package/src/shaders/reduction/sumF64.wgsl +2 -2
  47. package/src/shaders/sgemm_large.wgsl +206 -0
  48. package/src/shaders/sgemm_small.wgsl +212 -0
  49. package/src/shaders/sgemmtr_large.wgsl +120 -0
  50. package/src/shaders/sgemmtr_small.wgsl +113 -0
  51. package/src/shaders/sgemv_n.wgsl +3 -1
  52. package/src/shaders/sgemv_t.wgsl +3 -1
  53. package/src/shaders/snrm2.wgsl +72 -23
  54. package/src/shaders/ssymv.wgsl +3 -1
  55. package/src/shaders/symmetrize.wgsl +31 -0
  56. package/src/shaders/triangularize.wgsl +44 -0
  57. package/src/snrm2/snrm2.d.mts +3 -3
  58. package/src/snrm2/snrm2.mjs +33 -21
  59. package/src/srot/srot.d.mts +3 -5
  60. package/src/srot/srot.mjs +11 -9
  61. package/src/srotm/srotm.d.mts +3 -5
  62. package/src/srotm/srotm.mjs +19 -10
  63. package/src/sscal/sscal.d.mts +4 -4
  64. package/src/sscal/sscal.mjs +12 -10
  65. package/src/sswap/sswap.d.mts +3 -3
  66. package/src/sswap/sswap.mjs +11 -9
  67. package/src/ssymm/ssymm.d.mts +103 -0
  68. package/src/ssymm/ssymm.mjs +218 -0
  69. package/src/ssymv/ssymv.d.mts +1 -36
  70. package/src/ssymv/ssymv.mjs +12 -8
  71. package/src/ssyr/ssyr.d.mts +1 -30
  72. package/src/ssyr/ssyr.mjs +11 -7
  73. package/src/ssyr2/ssyr2.d.mts +1 -34
  74. package/src/ssyr2/ssyr2.mjs +12 -8
  75. package/src/ssyr2k/ssyr2k.d.mts +100 -0
  76. package/src/ssyr2k/ssyr2k.mjs +202 -0
  77. package/src/ssyrk/ssyrk.d.mts +90 -0
  78. package/src/ssyrk/ssyrk.mjs +177 -0
  79. package/src/strmm/strmm.d.mts +100 -0
  80. package/src/strmm/strmm.mjs +226 -0
  81. package/src/strmv/strmv.d.mts +1 -36
  82. package/src/strmv/strmv.mjs +12 -8
  83. package/src/strsm/strsm.d.mts +99 -0
  84. package/src/strsm/strsm.mjs +360 -0
  85. package/src/strsv/strsv.d.mts +1 -32
  86. package/src/strsv/strsv.mjs +18 -12
  87. package/src/util/benchmark.mjs +4 -6
  88. package/src/util/bindgroup.mjs +1 -3
  89. package/src/util/buffer.mjs +116 -20
  90. package/src/util/compute.mjs +12 -12
  91. package/src/util/constants.mjs +57 -0
  92. package/src/util/device.mjs +34 -0
  93. package/src/util/f64.mjs +3 -3
  94. package/src/util/pipeline.mjs +5 -6
  95. package/src/util/workgroup.mjs +55 -7
  96. package/src/shaders/browser-shaders.mjs +0 -55
@@ -0,0 +1,206 @@
1
+ // sgemm_large: C = alpha * op(A) * op(B) + beta * C — large-tile half of
2
+ // the two-tier autotuned dispatch (see sgemm.mjs and sgemm_small.wgsl).
3
+ // BM=BN=64, BK=8, TM=8, TN=4 (128 threads/workgroup) — the kernel 9
4
+ // autotuning winner (temp/autotune_sweep.mjs, temp/gen_sweep_kernel.mjs,
5
+ // swept BM/BN/BK/TM/TN and warp-tiled variants), +69% over the old BM=32
6
+ // single-tier baseline at n=512, +84% at n=1024. But BM=64 loses to BM=32
7
+ // below a 6x6=36 workgroup grid (not enough workgroups to fill the GPU at
8
+ // that tile size), hence the two-tier split rather than one global config.
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.
16
+
17
+ const BM: u32 = 64u;
18
+ const BN: u32 = 64u;
19
+ const BK: u32 = 8u;
20
+ const TM: u32 = 8u;
21
+ const TN: u32 = 4u;
22
+ const THREADS_X: u32 = BN / TN;
23
+ const THREADS_Y: u32 = BM / TM;
24
+ const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 128
25
+ const STRIDE_A: u32 = NUM_THREADS / BK; // rows of As covered per load-loop step
26
+ const STRIDE_B: u32 = NUM_THREADS / BN; // rows of Bs covered per load-loop step
27
+
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>;
33
+
34
+ struct Params {
35
+ m: u32,
36
+ n: u32,
37
+ k: u32,
38
+ alpha: f32,
39
+ beta: f32,
40
+ lda: u32,
41
+ ldb: u32,
42
+ ldc: u32,
43
+ transA: u32, // 0 = no-transpose, 1 = transpose
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
47
+ }
48
+
49
+ @group(0) @binding(5) var<uniform> params: Params;
50
+
51
+ var<workgroup> As: array<f32, BM * BK>;
52
+ var<workgroup> Bs: array<f32, BK * BN>;
53
+
54
+ @compute @workgroup_size(THREADS_X, THREADS_Y)
55
+ fn main(
56
+ @builtin(workgroup_id) wid: vec3u,
57
+ @builtin(local_invocation_id) lid: vec3u,
58
+ @builtin(local_invocation_index) tid: u32,
59
+ ) {
60
+ let blockRow = wid.y * BM;
61
+ let blockCol = wid.x * BN;
62
+ let threadCol = lid.x;
63
+ let threadRow = lid.y;
64
+
65
+ // Load indices, independent of the compute thread shape — a loop since
66
+ // NUM_THREADS doesn't match the tile size 1:1 at this config.
67
+ let innerRowA = tid / BK;
68
+ let innerColA = tid % BK;
69
+ let innerRowB = tid / BN;
70
+ let innerColB = tid % BN;
71
+
72
+ var threadResults: array<f32, TM * TN>;
73
+ for (var i = 0u; i < TM * TN; i++) {
74
+ threadResults[i] = 0.0;
75
+ }
76
+ var regM: array<f32, TM>;
77
+ var regN: array<f32, TN>;
78
+
79
+ let numTiles = (params.k + BK - 1u) / BK;
80
+ for (var t = 0u; t < numTiles; t++) {
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
+ }
125
+ }
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
+ }
170
+ }
171
+
172
+ workgroupBarrier();
173
+
174
+ for (var dotIdx = 0u; dotIdx < BK; dotIdx++) {
175
+ for (var i = 0u; i < TM; i++) {
176
+ regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
177
+ }
178
+ for (var i = 0u; i < TN; i++) {
179
+ regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
180
+ }
181
+ for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
182
+ for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
183
+ threadResults[resIdxM * TN + resIdxN] += regM[resIdxM] * regN[resIdxN];
184
+ }
185
+ }
186
+ }
187
+
188
+ workgroupBarrier();
189
+ }
190
+
191
+ for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
192
+ let row = blockRow + threadRow * TM + resIdxM;
193
+ if (row < params.m) {
194
+ for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
195
+ let col = blockCol + threadCol * TN + resIdxN;
196
+ if (col < params.n) {
197
+ let cIdx = row * params.ldc + col;
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);
202
+ }
203
+ }
204
+ }
205
+ }
206
+ }
@@ -0,0 +1,212 @@
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
+ // 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
+ //
15
+ // col mapped to gid.x for coalesced B/C access (row-major: col contiguous).
16
+
17
+ const BM: u32 = 32u;
18
+ const BN: u32 = 32u;
19
+ const BK: u32 = 8u;
20
+ const TM: u32 = 2u;
21
+ const TN: u32 = 2u;
22
+ const THREADS_X: u32 = BN / TN;
23
+ const THREADS_Y: u32 = BM / TM;
24
+ const NUM_THREADS: u32 = THREADS_X * THREADS_Y; // 256
25
+ const STRIDE_A: u32 = NUM_THREADS / BK;
26
+ const STRIDE_B: u32 = NUM_THREADS / BN;
27
+
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>;
33
+
34
+ struct Params {
35
+ m: u32,
36
+ n: u32,
37
+ k: u32,
38
+ alpha: f32,
39
+ beta: f32,
40
+ lda: u32,
41
+ ldb: u32,
42
+ ldc: u32,
43
+ transA: u32, // 0 = no-transpose, 1 = transpose
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
47
+ }
48
+
49
+ @group(0) @binding(5) var<uniform> params: Params;
50
+
51
+ var<workgroup> As: array<f32, BM * BK>;
52
+ var<workgroup> Bs: array<f32, BK * BN>;
53
+
54
+ @compute @workgroup_size(THREADS_X, THREADS_Y)
55
+ fn main(
56
+ @builtin(workgroup_id) wid: vec3u,
57
+ @builtin(local_invocation_id) lid: vec3u,
58
+ @builtin(local_invocation_index) tid: u32,
59
+ ) {
60
+ let blockRow = wid.y * BM;
61
+ let blockCol = wid.x * BN;
62
+ let threadCol = lid.x;
63
+ let threadRow = lid.y;
64
+
65
+ let innerRowA = tid / BK;
66
+ let innerColA = tid % BK;
67
+ let innerRowB = tid / BN;
68
+ let innerColB = tid % BN;
69
+
70
+ var threadResults: array<f32, TM * TN>;
71
+ for (var i = 0u; i < TM * TN; i++) {
72
+ threadResults[i] = 0.0;
73
+ }
74
+ var regM: array<f32, TM>;
75
+ var regN: array<f32, TN>;
76
+
77
+ let numTiles = (params.k + BK - 1u) / BK;
78
+ for (var t = 0u; t < numTiles; t++) {
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
+ }
127
+ }
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
+ }
176
+ }
177
+
178
+ workgroupBarrier();
179
+
180
+ for (var dotIdx = 0u; dotIdx < BK; dotIdx++) {
181
+ for (var i = 0u; i < TM; i++) {
182
+ regM[i] = As[(threadRow * TM + i) * BK + dotIdx];
183
+ }
184
+ for (var i = 0u; i < TN; i++) {
185
+ regN[i] = Bs[dotIdx * BN + threadCol * TN + i];
186
+ }
187
+ for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
188
+ for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
189
+ threadResults[resIdxM * TN + resIdxN] += regM[resIdxM] * regN[resIdxN];
190
+ }
191
+ }
192
+ }
193
+
194
+ workgroupBarrier();
195
+ }
196
+
197
+ for (var resIdxM = 0u; resIdxM < TM; resIdxM++) {
198
+ let row = blockRow + threadRow * TM + resIdxM;
199
+ if (row < params.m) {
200
+ for (var resIdxN = 0u; resIdxN < TN; resIdxN++) {
201
+ let col = blockCol + threadCol * TN + resIdxN;
202
+ if (col < params.n) {
203
+ let cIdx = row * params.ldc + col;
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);
208
+ }
209
+ }
210
+ }
211
+ }
212
+ }
@@ -0,0 +1,120 @@
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
+ // 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);
116
+ }
117
+ }
118
+ }
119
+ }
120
+ }
@@ -0,0 +1,113 @@
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
+ // 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);
109
+ }
110
+ }
111
+ }
112
+ }
113
+ }
@@ -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
  }