Xenova HF Staff commited on
Commit
d8aeb25
·
verified ·
1 Parent(s): bb8ee13

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -51,12 +51,17 @@ Default values (overridable per request):
51
 
52
  One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
53
 
 
 
 
54
  - `subgroup_matrix_transbatch_b_f16` — Aligned rank-3 transBatchB product with sufficient reduction depth and output tiles to amortize subgroup-matrix staging.
55
  - `subgroup_matrix_transbatch_b_f32` — Aligned rank-3 transBatchB product with sufficient reduction depth and output tiles to amortize subgroup-matrix staging.
 
56
  - `rank2_band_vec4_splitk` — Splits the vec4 band's K axis across up to sixteen workgroups. Each range writes an f32 partial band with alpha applied, and a combine pass sums the partials.
57
  - `rank2_band_vec4` — Few-row band for a rank-2 product without transposes. Each lane owns one vec4 column group and one accumulator per row, so every B word feeds all 2 to 16 rows; alpha is applied at the store.
58
  - `rank2_band_vec4_f32_preferred` — Few-row band for a rank-2 product without transposes. Each lane owns one vec4 column group and one accumulator per row, so every B word feeds all 2 to 16 rows; alpha is applied at the store.
59
  - `subgroup_matrix_splitk` — Partitions the K reduction of small-M rank-2 products across subgroup-matrix workgroups, then combines float32 partials that already include alpha.
 
60
  - `plain_rank2_tiled_reg` — Register-blocked rank-2 `Y = alpha * A @ B` specialization for non-transposed inputs on tiers without subgroup-matrix support.
61
  - `transbatch_b_tiled_reg` — Register-blocked logical rank3 product with an interleaved physical B batch axis.
62
 
@@ -69,7 +74,7 @@ Some implementation variants require `subgroup-matrix` and `subgroups`. These ar
69
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
70
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
71
  - [`test.json`](build/webgpu/test.json) — correctness cases
72
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
73
  - [`fused-matmul-subgroup-matrix.wgsl.jinja`](build/webgpu/fused-matmul-subgroup-matrix.wgsl.jinja)
74
  - [`matmul-band-vec4.wgsl.jinja`](build/webgpu/matmul-band-vec4.wgsl.jinja)
75
  - [`matmul-subgroup-matrix-ext.wgsl.jinja`](build/webgpu/matmul-subgroup-matrix-ext.wgsl.jinja)
@@ -81,7 +86,7 @@ Some implementation variants require `subgroup-matrix` and `subgroups`. These ar
81
  ## Use with `@huggingface/kernels`
82
 
83
  ```sh
84
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
85
  ```
86
 
87
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
51
 
52
  One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
53
 
54
+ - `broadcast_transb_tiled_reg` — Register-blocked broadcast product with physically transposed B. Reuses the shared batch-addressing tile, keeps f32 accumulation, and preserves scalar K order for f16. Low tile count and excessive padding demote this otherwise correct path.
55
+ - `broadcast_transb_subgroup_matrix_f16` — Broadcast transposed-B product using a supported 8x8x8 subgroup-matrix configuration with f32 accumulation. Logical shapes and physical B strides share the existing matrix engine; insufficient output tiles or excessive padding retain the generic tile.
56
+ - `broadcast_transb_subgroup_matrix_f32` — Broadcast transposed-B product using a supported 8x8x8 subgroup-matrix configuration with f32 accumulation. Logical shapes and physical B strides share the existing matrix engine; insufficient output tiles or excessive padding retain the generic tile.
57
  - `subgroup_matrix_transbatch_b_f16` — Aligned rank-3 transBatchB product with sufficient reduction depth and output tiles to amortize subgroup-matrix staging.
58
  - `subgroup_matrix_transbatch_b_f32` — Aligned rank-3 transBatchB product with sufficient reduction depth and output tiles to amortize subgroup-matrix staging.
59
+ - `m1_gemv_vec4` — Vector-by-matrix specialization for a single output row: each workgroup owns 32 consecutive vec4 column groups and partitions the reduction across the workgroup's second dimension. The accumulator stays float32 for both tensor types.
60
  - `rank2_band_vec4_splitk` — Splits the vec4 band's K axis across up to sixteen workgroups. Each range writes an f32 partial band with alpha applied, and a combine pass sums the partials.
61
  - `rank2_band_vec4` — Few-row band for a rank-2 product without transposes. Each lane owns one vec4 column group and one accumulator per row, so every B word feeds all 2 to 16 rows; alpha is applied at the store.
62
  - `rank2_band_vec4_f32_preferred` — Few-row band for a rank-2 product without transposes. Each lane owns one vec4 column group and one accumulator per row, so every B word feeds all 2 to 16 rows; alpha is applied at the store.
63
  - `subgroup_matrix_splitk` — Partitions the K reduction of small-M rank-2 products across subgroup-matrix workgroups, then combines float32 partials that already include alpha.
64
+ - `subgroup_matrix` — Subgroup-matrix `Y = alpha * op(A) @ op(B)` over dense batches with float32 accumulation. An output width that is not a multiple of the 64-wide column tile switches the trailing tile to guarded addressing: its B columns clamp to N - 1 and its stores drop every column at or past N. Yields the shape when the padded column ratio exceeds its tunable ceiling.
65
  - `plain_rank2_tiled_reg` — Register-blocked rank-2 `Y = alpha * A @ B` specialization for non-transposed inputs on tiers without subgroup-matrix support.
66
  - `transbatch_b_tiled_reg` — Register-blocked logical rank3 product with an interleaved physical B batch axis.
67
 
 
74
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
75
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
76
  - [`test.json`](build/webgpu/test.json) — correctness cases
77
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
78
  - [`fused-matmul-subgroup-matrix.wgsl.jinja`](build/webgpu/fused-matmul-subgroup-matrix.wgsl.jinja)
79
  - [`matmul-band-vec4.wgsl.jinja`](build/webgpu/matmul-band-vec4.wgsl.jinja)
80
  - [`matmul-subgroup-matrix-ext.wgsl.jinja`](build/webgpu/matmul-subgroup-matrix-ext.wgsl.jinja)
 
86
  ## Use with `@huggingface/kernels`
87
 
88
  ```sh
89
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
90
  ```
91
 
92
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/bench.json CHANGED
@@ -56,7 +56,62 @@
56
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * 384 * 400 * 1536" }] }
57
  },
58
  {
59
- "name": "fusedmatmul-f16-aligned-n512-512x2048x512-healthy",
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
60
  "preset": "smoke",
61
  "attrs": { "alpha": 1 },
62
  "vars": { "M": 512, "K": 2048, "N": 512 },
@@ -163,7 +218,7 @@
163
  "name": "fusedmatmul-f16-rank4-by-rank2-shared-weight-b2h8-m512-k2048-n512-pathology",
164
  "preset": "stress",
165
  "provenance": {
166
- "notes": "A production-scale batched projection shares one rank-2 weight across batches. This NumPy-style broadcast shape exercises the register-blocked portable route."
167
  },
168
  "attrs": { "alpha": 1 },
169
  "vars": { "dtype": "float16", "M": 512, "K": 2048, "N": 512 },
@@ -175,7 +230,7 @@
175
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * args.K" }] }
176
  },
177
  {
178
- "name": "fusedmatmul-f16-transbatch-a-forces-tiled-8x512x2048x512-stress",
179
  "preset": "stress",
180
  "attrs": { "alpha": 1, "transBatchA": 1 },
181
  "vars": { "batch": 8, "M": 512, "K": 2048, "N": 512 },
@@ -190,7 +245,7 @@
190
  "name": "fusedmatmul-f16-transbatch-b-8x512x2048x512-pathology",
191
  "preset": "stress",
192
  "provenance": {
193
- "notes": "Interleaved B batches in a production-scale projection; compares stride-aware subgroup matrices with the portable tiled path."
194
  },
195
  "attrs": { "alpha": 1, "transBatchB": 1 },
196
  "vars": { "batch": 8, "M": 512, "K": 2048, "N": 512 },
@@ -202,7 +257,7 @@
202
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * 8 * 512 * 512 * 2048" }] }
203
  },
204
  {
205
- "name": "fusedmatmul-f16-broadcast-batch-forces-tiled-1x8x512x2048x512-stress",
206
  "preset": "stress",
207
  "attrs": { "alpha": 1 },
208
  "inputs": {
@@ -305,7 +360,7 @@
305
  "name": "fusedmatmul-f32-transbatch-b-8x512x2048x512-pathology",
306
  "preset": "stress",
307
  "provenance": {
308
- "notes": "Interleaved B batches in a production-scale projection; compares stride-aware subgroup matrices with the portable tiled path."
309
  },
310
  "attrs": { "alpha": 1, "transBatchB": 1 },
311
  "vars": { "batch": 8, "M": 512, "K": 2048, "N": 512 },
@@ -319,9 +374,7 @@
319
  {
320
  "name": "fusedmatmul-float16-transbatch-b-k128-workgroups64",
321
  "preset": "stress",
322
- "provenance": {
323
- "notes": "Aligned interleaved B layout at the reduction-length and 64-matrix-workgroup selector floors."
324
- },
325
  "attrs": { "alpha": 0.5, "transBatchB": 1 },
326
  "vars": { "batch": 2, "M": 128, "K": 128, "N": 512 },
327
  "inputs": {
@@ -334,9 +387,7 @@
334
  {
335
  "name": "fusedmatmul-float32-transbatch-b-k128-workgroups64",
336
  "preset": "stress",
337
- "provenance": {
338
- "notes": "Aligned interleaved B layout at the reduction-length and 64-matrix-workgroup selector floors."
339
- },
340
  "attrs": { "alpha": 0.5, "transBatchB": 1 },
341
  "vars": { "batch": 2, "M": 128, "K": 128, "N": 512 },
342
  "inputs": {
@@ -363,7 +414,7 @@
363
  },
364
  "outputs": { "Y": { "dtype": "float16", "shape": [4, 128, 512] } },
365
  "provenance": {
366
- "notes": "Portable register tile transBatchB: physical B[K,batch,N], natural tile count reaches the shared occupancy gate; negative alpha and positive inputs avoid cancellation."
367
  },
368
  "preset": "smoke",
369
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, 2)" }] }
@@ -385,7 +436,7 @@
385
  },
386
  "outputs": { "Y": { "dtype": "float16", "shape": [3, 129, 513] } },
387
  "provenance": {
388
- "notes": "Portable register tile transBatchB: physical B[K,batch,N], natural tile count reaches the shared occupancy gate; negative alpha and positive inputs avoid cancellation."
389
  },
390
  "preset": "smoke",
391
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, 2)" }] }
@@ -407,7 +458,7 @@
407
  },
408
  "outputs": { "Y": { "dtype": "float32", "shape": [4, 128, 512] } },
409
  "provenance": {
410
- "notes": "Portable register tile transBatchB: physical B[K,batch,N], natural tile count reaches the shared occupancy gate; negative alpha and positive inputs avoid cancellation."
411
  },
412
  "preset": "smoke",
413
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, 2)" }] }
@@ -429,7 +480,7 @@
429
  },
430
  "outputs": { "Y": { "dtype": "float32", "shape": [3, 129, 513] } },
431
  "provenance": {
432
- "notes": "Portable register tile transBatchB: physical B[K,batch,N], natural tile count reaches the shared occupancy gate; negative alpha and positive inputs avoid cancellation."
433
  },
434
  "preset": "smoke",
435
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, 2)" }] }
@@ -499,6 +550,244 @@
499
  },
500
  "outputs": { "Y": { "dtype": "float32", "shape": [16, 4096], "dist": "empty" } },
501
  "bench": { "metrics": [{ "type": "gflops", "value": 1073741824 }] }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
502
  }
503
  ]
504
  }
 
56
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * 384 * 400 * 1536" }] }
57
  },
58
  {
59
+ "name": "fusedmatmul-f16-deep-skew-m128-k2048-n2048",
60
+ "preset": "model",
61
+ "attrs": { "alpha": 1 },
62
+ "inputs": {
63
+ "A": { "shape": [128, 2048], "dtype": "float16", "dist": "normal", "seed": 596, "scale": 0.1 },
64
+ "B": { "shape": [2048, 2048], "dtype": "float16", "dist": "normal", "seed": 597, "scale": 0.1 }
65
+ },
66
+ "outputs": { "Y": { "shape": [128, 2048], "dtype": "float16" } },
67
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * 128 * 2048 * 2048" }] }
68
+ },
69
+ {
70
+ "name": "fusedmatmul-f16-deep-skew-m1024-k2048-n256",
71
+ "preset": "model",
72
+ "attrs": { "alpha": 1 },
73
+ "inputs": {
74
+ "A": { "shape": [1024, 2048], "dtype": "float16", "dist": "normal", "seed": 598, "scale": 0.1 },
75
+ "B": { "shape": [2048, 256], "dtype": "float16", "dist": "normal", "seed": 599, "scale": 0.1 }
76
+ },
77
+ "outputs": { "Y": { "shape": [1024, 256], "dtype": "float16" } },
78
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * 1024 * 256 * 2048" }] }
79
+ },
80
+ {
81
+ "name": "fusedmatmul-f16-depth-k64-512x64x512",
82
+ "preset": "model",
83
+ "attrs": { "alpha": 1 },
84
+ "inputs": {
85
+ "A": { "shape": [512, 64], "dtype": "float16", "dist": "normal", "seed": 590, "scale": 0.1 },
86
+ "B": { "shape": [64, 512], "dtype": "float16", "dist": "normal", "seed": 591, "scale": 0.1 }
87
+ },
88
+ "outputs": { "Y": { "shape": [512, 512], "dtype": "float16" } },
89
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * 512 * 512 * 64" }] }
90
+ },
91
+ {
92
+ "name": "fusedmatmul-f16-depth-k512-512x512x512",
93
+ "preset": "model",
94
+ "attrs": { "alpha": 1 },
95
+ "inputs": {
96
+ "A": { "shape": [512, 512], "dtype": "float16", "dist": "normal", "seed": 592, "scale": 0.1 },
97
+ "B": { "shape": [512, 512], "dtype": "float16", "dist": "normal", "seed": 593, "scale": 0.1 }
98
+ },
99
+ "outputs": { "Y": { "shape": [512, 512], "dtype": "float16" } },
100
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * 512 * 512 * 512" }] }
101
+ },
102
+ {
103
+ "name": "fusedmatmul-f16-depth-k1024-512x1024x512",
104
+ "preset": "model",
105
+ "attrs": { "alpha": 1 },
106
+ "inputs": {
107
+ "A": { "shape": [512, 1024], "dtype": "float16", "dist": "normal", "seed": 594, "scale": 0.1 },
108
+ "B": { "shape": [1024, 512], "dtype": "float16", "dist": "normal", "seed": 595, "scale": 0.1 }
109
+ },
110
+ "outputs": { "Y": { "shape": [512, 512], "dtype": "float16" } },
111
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * 512 * 512 * 1024" }] }
112
+ },
113
+ {
114
+ "name": "fusedmatmul-f16-aligned-n512-512x2048x512-control",
115
  "preset": "smoke",
116
  "attrs": { "alpha": 1 },
117
  "vars": { "M": 512, "K": 2048, "N": 512 },
 
218
  "name": "fusedmatmul-f16-rank4-by-rank2-shared-weight-b2h8-m512-k2048-n512-pathology",
219
  "preset": "stress",
220
  "provenance": {
221
+ "notes": "Measures FusedMatMul over a batched (2x8) projection sharing one rank-2 [2048,512] weight, with M=512, K=2048, N=512 (float16)."
222
  },
223
  "attrs": { "alpha": 1 },
224
  "vars": { "dtype": "float16", "M": 512, "K": 2048, "N": 512 },
 
230
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * args.K" }] }
231
  },
232
  {
233
+ "name": "fusedmatmul-f16-transbatch-a-tiled-8x512x2048x512-stress",
234
  "preset": "stress",
235
  "attrs": { "alpha": 1, "transBatchA": 1 },
236
  "vars": { "batch": 8, "M": 512, "K": 2048, "N": 512 },
 
245
  "name": "fusedmatmul-f16-transbatch-b-8x512x2048x512-pathology",
246
  "preset": "stress",
247
  "provenance": {
248
+ "notes": "Measures FusedMatMul over a batch=8 projection with transposed-batch B (physical [K=2048, batch=8, N=512]): M=512, K=2048, N=512 (float16)."
249
  },
250
  "attrs": { "alpha": 1, "transBatchB": 1 },
251
  "vars": { "batch": 8, "M": 512, "K": 2048, "N": 512 },
 
257
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * 8 * 512 * 512 * 2048" }] }
258
  },
259
  {
260
+ "name": "fusedmatmul-f16-broadcast-batch-tiled-1x8x512x2048x512-stress",
261
  "preset": "stress",
262
  "attrs": { "alpha": 1 },
263
  "inputs": {
 
360
  "name": "fusedmatmul-f32-transbatch-b-8x512x2048x512-pathology",
361
  "preset": "stress",
362
  "provenance": {
363
+ "notes": "Measures FusedMatMul over a batch=8 projection with transposed-batch B (physical [K=2048, batch=8, N=512]): M=512, K=2048, N=512 (float32)."
364
  },
365
  "attrs": { "alpha": 1, "transBatchB": 1 },
366
  "vars": { "batch": 8, "M": 512, "K": 2048, "N": 512 },
 
374
  {
375
  "name": "fusedmatmul-float16-transbatch-b-k128-workgroups64",
376
  "preset": "stress",
377
+ "provenance": { "notes": "Aligned interleaved B layout with a long reduction and 64 output tile groups." },
 
 
378
  "attrs": { "alpha": 0.5, "transBatchB": 1 },
379
  "vars": { "batch": 2, "M": 128, "K": 128, "N": 512 },
380
  "inputs": {
 
387
  {
388
  "name": "fusedmatmul-float32-transbatch-b-k128-workgroups64",
389
  "preset": "stress",
390
+ "provenance": { "notes": "Aligned interleaved B layout with a long reduction and 64 output tile groups." },
 
 
391
  "attrs": { "alpha": 0.5, "transBatchB": 1 },
392
  "vars": { "batch": 2, "M": 128, "K": 128, "N": 512 },
393
  "inputs": {
 
414
  },
415
  "outputs": { "Y": { "dtype": "float16", "shape": [4, 128, 512] } },
416
  "provenance": {
417
+ "notes": "Measures FusedMatMul over a fully tile-aligned M=128, K=128, N=512 shape with transposed-batch B (float16, batch=4)."
418
  },
419
  "preset": "smoke",
420
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, 2)" }] }
 
436
  },
437
  "outputs": { "Y": { "dtype": "float16", "shape": [3, 129, 513] } },
438
  "provenance": {
439
+ "notes": "Measures FusedMatMul over a transposed-batch-B shape with partial tiles: M=129, K=131, N=513 (float16, batch=3)."
440
  },
441
  "preset": "smoke",
442
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, 2)" }] }
 
458
  },
459
  "outputs": { "Y": { "dtype": "float32", "shape": [4, 128, 512] } },
460
  "provenance": {
461
+ "notes": "Measures FusedMatMul over a fully tile-aligned M=128, K=128, N=512 shape with transposed-batch B (float32, batch=4)."
462
  },
463
  "preset": "smoke",
464
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, 2)" }] }
 
480
  },
481
  "outputs": { "Y": { "dtype": "float32", "shape": [3, 129, 513] } },
482
  "provenance": {
483
+ "notes": "Measures FusedMatMul over a transposed-batch-B shape with partial tiles: M=129, K=131, N=513 (float32, batch=3)."
484
  },
485
  "preset": "smoke",
486
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, 2)" }] }
 
550
  },
551
  "outputs": { "Y": { "dtype": "float32", "shape": [16, 4096], "dist": "empty" } },
552
  "bench": { "metrics": [{ "type": "gflops", "value": 1073741824 }] }
553
+ },
554
+ {
555
+ "name": "fusedmatmul-transb-tails-float16",
556
+ "preset": "stress",
557
+ "attrs": { "transB": 1, "alpha": -0.5 },
558
+ "inputs": {
559
+ "A": { "dtype": "float16", "shape": [2, 1, 129, 65], "dist": "normal", "seed": 7302, "scale": 0.1 },
560
+ "B": { "dtype": "float16", "shape": [3, 129, 65], "dist": "normal", "seed": 7303, "scale": 0.1 }
561
+ },
562
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 129, 129] } },
563
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
564
+ "provenance": {
565
+ "notes": "Measures FusedMatMul over a rank-4 by rank-3 broadcast (batch dims 2x1 against 3) with transposed B: M=129, K=65, N=129 (float16)."
566
+ }
567
+ },
568
+ {
569
+ "name": "fusedmatmul-transb-rank3x2-float16",
570
+ "preset": "stress",
571
+ "attrs": { "transB": 1, "alpha": -0.5 },
572
+ "inputs": {
573
+ "A": { "dtype": "float16", "shape": [4, 256, 128], "dist": "normal", "seed": 7304, "scale": 0.1 },
574
+ "B": { "dtype": "float16", "shape": [512, 128], "dist": "normal", "seed": 7305, "scale": 0.1 }
575
+ },
576
+ "outputs": { "Y": { "dtype": "float16", "shape": [4, 256, 512] } },
577
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
578
+ "provenance": {
579
+ "notes": "Measures FusedMatMul over a rank-3 by rank-2 broadcast with transposed B: M=256, K=128, N=512, batch=4 (float16)."
580
+ }
581
+ },
582
+ {
583
+ "name": "fusedmatmul-transb-rank5x3-float16",
584
+ "preset": "stress",
585
+ "attrs": { "transB": 1, "alpha": -0.5 },
586
+ "inputs": {
587
+ "A": { "dtype": "float16", "shape": [2, 1, 2, 128, 32], "dist": "normal", "seed": 7306, "scale": 0.1 },
588
+ "B": { "dtype": "float16", "shape": [2, 128, 32], "dist": "normal", "seed": 7307, "scale": 0.1 }
589
+ },
590
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 1, 2, 128, 128] } },
591
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
592
+ "provenance": {
593
+ "notes": "Measures FusedMatMul over a rank-5 by rank-3 broadcast with transposed B: M=128, K=32, N=128, batch dims 2x1x2 (float16)."
594
+ }
595
+ },
596
+ {
597
+ "name": "fusedmatmul-transb-low_tiles-float16",
598
+ "preset": "stress",
599
+ "attrs": { "transB": 1, "alpha": -0.5 },
600
+ "inputs": {
601
+ "A": { "dtype": "float16", "shape": [1, 64, 32], "dist": "normal", "seed": 7308, "scale": 0.1 },
602
+ "B": { "dtype": "float16", "shape": [64, 32], "dist": "normal", "seed": 7309, "scale": 0.1 }
603
+ },
604
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 64, 64] } },
605
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
606
+ "provenance": {
607
+ "notes": "Measures FusedMatMul over a rank-3 by rank-2 broadcast with transposed B: M=64, K=32, N=64, batch=1 (float16)."
608
+ }
609
+ },
610
+ {
611
+ "name": "fusedmatmul-transb-large-float32",
612
+ "preset": "stress",
613
+ "attrs": { "transB": 1, "alpha": -0.5 },
614
+ "inputs": {
615
+ "A": { "dtype": "float32", "shape": [2, 8, 512, 64], "dist": "normal", "seed": 7300, "scale": 0.1 },
616
+ "B": { "dtype": "float32", "shape": [8, 512, 64], "dist": "normal", "seed": 7301, "scale": 0.1 }
617
+ },
618
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 8, 512, 512] } },
619
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
620
+ "provenance": {
621
+ "notes": "Measures FusedMatMul over a rank-4 by rank-3 broadcast with transposed B: M=512, K=64, N=512, batch dims 2x8 (float32)."
622
+ }
623
+ },
624
+ {
625
+ "name": "fusedmatmul-transb-tails-float32",
626
+ "preset": "stress",
627
+ "attrs": { "transB": 1, "alpha": -0.5 },
628
+ "inputs": {
629
+ "A": { "dtype": "float32", "shape": [2, 1, 129, 65], "dist": "normal", "seed": 7302, "scale": 0.1 },
630
+ "B": { "dtype": "float32", "shape": [3, 129, 65], "dist": "normal", "seed": 7303, "scale": 0.1 }
631
+ },
632
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 129, 129] } },
633
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
634
+ "provenance": {
635
+ "notes": "Measures FusedMatMul over a rank-4 by rank-3 broadcast (batch dims 2x1 against 3) with transposed B: M=129, K=65, N=129 (float32)."
636
+ }
637
+ },
638
+ {
639
+ "name": "fusedmatmul-transb-rank3x2-float32",
640
+ "preset": "stress",
641
+ "attrs": { "transB": 1, "alpha": -0.5 },
642
+ "inputs": {
643
+ "A": { "dtype": "float32", "shape": [4, 256, 128], "dist": "normal", "seed": 7304, "scale": 0.1 },
644
+ "B": { "dtype": "float32", "shape": [512, 128], "dist": "normal", "seed": 7305, "scale": 0.1 }
645
+ },
646
+ "outputs": { "Y": { "dtype": "float32", "shape": [4, 256, 512] } },
647
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
648
+ "provenance": {
649
+ "notes": "Measures FusedMatMul over a rank-3 by rank-2 broadcast with transposed B: M=256, K=128, N=512, batch=4 (float32)."
650
+ }
651
+ },
652
+ {
653
+ "name": "fusedmatmul-transb-rank5x3-float32",
654
+ "preset": "stress",
655
+ "attrs": { "transB": 1, "alpha": -0.5 },
656
+ "inputs": {
657
+ "A": { "dtype": "float32", "shape": [2, 1, 2, 128, 32], "dist": "normal", "seed": 7306, "scale": 0.1 },
658
+ "B": { "dtype": "float32", "shape": [2, 128, 32], "dist": "normal", "seed": 7307, "scale": 0.1 }
659
+ },
660
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 1, 2, 128, 128] } },
661
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
662
+ "provenance": {
663
+ "notes": "Measures FusedMatMul over a rank-5 by rank-3 broadcast with transposed B: M=128, K=32, N=128, batch dims 2x1x2 (float32)."
664
+ }
665
+ },
666
+ {
667
+ "name": "fusedmatmul-transb-low_tiles-float32",
668
+ "preset": "stress",
669
+ "attrs": { "transB": 1, "alpha": -0.5 },
670
+ "inputs": {
671
+ "A": { "dtype": "float32", "shape": [1, 64, 32], "dist": "normal", "seed": 7308, "scale": 0.1 },
672
+ "B": { "dtype": "float32", "shape": [64, 32], "dist": "normal", "seed": 7309, "scale": 0.1 }
673
+ },
674
+ "outputs": { "Y": { "dtype": "float32", "shape": [1, 64, 64] } },
675
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
676
+ "provenance": {
677
+ "notes": "Measures FusedMatMul over a rank-3 by rank-2 broadcast with transposed B: M=64, K=32, N=64, batch=1 (float32)."
678
+ }
679
+ },
680
+ {
681
+ "name": "fusedmatmul-transb-matrix_tile_floor-float16",
682
+ "preset": "stress",
683
+ "attrs": { "transB": 1, "alpha": -0.5 },
684
+ "inputs": {
685
+ "A": { "dtype": "float16", "shape": [2, 256, 64], "dist": "normal", "seed": 7320, "scale": 0.1 },
686
+ "B": { "dtype": "float16", "shape": [256, 64], "dist": "normal", "seed": 7321, "scale": 0.1 }
687
+ },
688
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 256, 256] } },
689
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
690
+ "provenance": {
691
+ "notes": "Profitability-boundary probe for output tile count and padded arithmetic, with a shared transposed B matrix."
692
+ }
693
+ },
694
+ {
695
+ "name": "fusedmatmul-transb-register_tile_floor-float16",
696
+ "preset": "stress",
697
+ "attrs": { "transB": 1, "alpha": -0.5 },
698
+ "inputs": {
699
+ "A": { "dtype": "float16", "shape": [4, 256, 64], "dist": "normal", "seed": 7322, "scale": 0.1 },
700
+ "B": { "dtype": "float16", "shape": [256, 64], "dist": "normal", "seed": 7323, "scale": 0.1 }
701
+ },
702
+ "outputs": { "Y": { "dtype": "float16", "shape": [4, 256, 256] } },
703
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
704
+ "provenance": {
705
+ "notes": "Profitability-boundary probe for output tile count and padded arithmetic, with a shared transposed B matrix."
706
+ }
707
+ },
708
+ {
709
+ "name": "fusedmatmul-transb-padding_inside-float16",
710
+ "preset": "stress",
711
+ "attrs": { "transB": 1, "alpha": -0.5 },
712
+ "inputs": {
713
+ "A": { "dtype": "float16", "shape": [8, 256, 128], "dist": "normal", "seed": 7324, "scale": 0.1 },
714
+ "B": { "dtype": "float16", "shape": [129, 128], "dist": "normal", "seed": 7325, "scale": 0.1 }
715
+ },
716
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 256, 129] } },
717
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
718
+ "provenance": {
719
+ "notes": "Profitability-boundary probe for output tile count and padded arithmetic, with a shared transposed B matrix."
720
+ }
721
+ },
722
+ {
723
+ "name": "fusedmatmul-transb-padding_cross_engine-float16",
724
+ "preset": "stress",
725
+ "attrs": { "transB": 1, "alpha": -0.5 },
726
+ "inputs": {
727
+ "A": { "dtype": "float16", "shape": [8, 256, 65], "dist": "normal", "seed": 7326, "scale": 0.1 },
728
+ "B": { "dtype": "float16", "shape": [129, 65], "dist": "normal", "seed": 7327, "scale": 0.1 }
729
+ },
730
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 256, 129] } },
731
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
732
+ "provenance": {
733
+ "notes": "Profitability-boundary probe for output tile count and padded arithmetic, with a shared transposed B matrix."
734
+ }
735
+ },
736
+ {
737
+ "name": "fusedmatmul-transb-matrix_tile_floor-float32",
738
+ "preset": "stress",
739
+ "attrs": { "transB": 1, "alpha": -0.5 },
740
+ "inputs": {
741
+ "A": { "dtype": "float32", "shape": [2, 256, 64], "dist": "normal", "seed": 7320, "scale": 0.1 },
742
+ "B": { "dtype": "float32", "shape": [256, 64], "dist": "normal", "seed": 7321, "scale": 0.1 }
743
+ },
744
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 256, 256] } },
745
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
746
+ "provenance": {
747
+ "notes": "Profitability-boundary probe for output tile count and padded arithmetic, with a shared transposed B matrix."
748
+ }
749
+ },
750
+ {
751
+ "name": "fusedmatmul-transb-register_tile_floor-float32",
752
+ "preset": "stress",
753
+ "attrs": { "transB": 1, "alpha": -0.5 },
754
+ "inputs": {
755
+ "A": { "dtype": "float32", "shape": [4, 256, 64], "dist": "normal", "seed": 7322, "scale": 0.1 },
756
+ "B": { "dtype": "float32", "shape": [256, 64], "dist": "normal", "seed": 7323, "scale": 0.1 }
757
+ },
758
+ "outputs": { "Y": { "dtype": "float32", "shape": [4, 256, 256] } },
759
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
760
+ "provenance": {
761
+ "notes": "Profitability-boundary probe for output tile count and padded arithmetic, with a shared transposed B matrix."
762
+ }
763
+ },
764
+ {
765
+ "name": "fusedmatmul-transb-padding_inside-float32",
766
+ "preset": "stress",
767
+ "attrs": { "transB": 1, "alpha": -0.5 },
768
+ "inputs": {
769
+ "A": { "dtype": "float32", "shape": [8, 256, 128], "dist": "normal", "seed": 7324, "scale": 0.1 },
770
+ "B": { "dtype": "float32", "shape": [129, 128], "dist": "normal", "seed": 7325, "scale": 0.1 }
771
+ },
772
+ "outputs": { "Y": { "dtype": "float32", "shape": [8, 256, 129] } },
773
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
774
+ "provenance": {
775
+ "notes": "Profitability-boundary probe for output tile count and padded arithmetic, with a shared transposed B matrix."
776
+ }
777
+ },
778
+ {
779
+ "name": "fusedmatmul-transb-padding_cross_engine-float32",
780
+ "preset": "stress",
781
+ "attrs": { "transB": 1, "alpha": -0.5 },
782
+ "inputs": {
783
+ "A": { "dtype": "float32", "shape": [8, 256, 65], "dist": "normal", "seed": 7326, "scale": 0.1 },
784
+ "B": { "dtype": "float32", "shape": [129, 65], "dist": "normal", "seed": 7327, "scale": 0.1 }
785
+ },
786
+ "outputs": { "Y": { "dtype": "float32", "shape": [8, 256, 129] } },
787
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, -1)" }] },
788
+ "provenance": {
789
+ "notes": "Profitability-boundary probe for output tile count and padded arithmetic, with a shared transposed B matrix."
790
+ }
791
  }
792
  ]
793
  }
build/webgpu/fused-matmul-subgroup-matrix.wgsl.jinja CHANGED
@@ -1,9 +1,17 @@
1
  // com.microsoft.FusedMatMul subgroup-matrix specialization: Y = alpha * op(A) @ op(B).
2
  // transA and transB transpose the corresponding matrix operand on load.
3
  // Dense batches map through workgroup_id.z, and M-tail rows are guarded by row_limit.
4
- // Alignment gates keep K % 32 == 0 and N % 64 == 0 so subgroupMatrixLoad never
5
- // sees partial 8x8 tiles. The batch is required to match between A and B
6
- // (no broadcast) because a_base/b_base both index by the same workgroup_id.z.
 
 
 
 
 
 
 
 
7
  enable subgroups;
8
  {% if pinSubgroupSize32 %}
9
  enable subgroup_size_control;
@@ -11,27 +19,22 @@ enable subgroup_size_control;
11
  enable chromium_experimental_subgroup_matrix;
12
  diagnostic(off, chromium.subgroup_matrix_uniformity);
13
 
14
-
15
  {{ env.wgsl.resourceDeclarations }}
16
 
17
  {% set operandScalar = fScalar %}
18
  {% set accScalar = "f32" %}
 
19
 
20
- const M: u32 = {{ M }}u;
21
  const K: u32 = {{ K }}u;
22
  const N: u32 = {{ N }}u;
23
  {% if transBatchA %}
24
  const BATCH_COUNT: u32 = {{ batchCount }}u;
25
  const A_BATCH_STRIDE: u32 = K;
26
  const A_M_STRIDE: u32 = BATCH_COUNT * K;
27
- {% else %}
28
- const A_BATCH_STRIDE: u32 = M * K;
29
- {% if not transA %}
30
  const A_M_STRIDE: u32 = K;
31
  {% endif %}
32
- {% endif %}
33
  const B_BATCH_STRIDE: u32 = K * N;
34
- const C_BATCH_STRIDE: u32 = M * N;
35
  const ALPHA: {{ accScalar }} = {{ accScalar }}({{ alpha }});
36
  const TILE_COLS: u32 = 64u;
37
  const TILE_ROWS: u32 = 32u;
@@ -48,10 +51,10 @@ fn loadSHMA(a_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
48
  let col = c_idx * 8u;
49
  for (var col_offset = 0u; col_offset < 8u; col_offset = col_offset + 1u) {
50
  let k = k_idx + col + col_offset;
51
- if (a_global < M) {
52
  {% if transA %}
53
  // op(A) = A^T: A stored [.., K, M], so op(A)[a_global, k] = A[k, a_global].
54
- tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(a[a_base + k * M + a_global]);
55
  {% else %}
56
  tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(a[a_base + a_global * A_M_STRIDE + k]);
57
  {% endif %}
@@ -62,7 +65,13 @@ fn loadSHMA(a_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
62
  }
63
 
64
  fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
 
 
 
 
 
65
  let b_col = tile_base + row;
 
66
  let col = c_idx * 16u;
67
  for (var i = 0u; i < 16u; i = i + 1u) {
68
  let k = k_idx + col + i;
@@ -75,9 +84,19 @@ fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
75
  }
76
  }
77
 
78
- fn storeOutput(offset: u32, row: u32, col: u32, src_slot: u32, row_limit: i32) {
79
  if (row_limit > 0 && row < u32(row_limit)) {
80
  let col2 = col + 1u;
 
 
 
 
 
 
 
 
 
 
81
  y[offset + row * N + col] = {{ outScalar }}(ALPHA * scratch[src_slot][0][row * 8u + col]);
82
  y[offset + row * N + col + 8u] = {{ outScalar }}(ALPHA * scratch[src_slot][1][row * 8u + col]);
83
  y[offset + row * N + col + 16u] = {{ outScalar }}(ALPHA * scratch[src_slot][2][row * 8u + col]);
@@ -87,6 +106,7 @@ fn storeOutput(offset: u32, row: u32, col: u32, src_slot: u32, row_limit: i32) {
87
  y[offset + row * N + col2 + 8u] = {{ outScalar }}(ALPHA * scratch[src_slot][1][row * 8u + col2]);
88
  y[offset + row * N + col2 + 16u] = {{ outScalar }}(ALPHA * scratch[src_slot][2][row * 8u + col2]);
89
  y[offset + row * N + col2 + 24u] = {{ outScalar }}(ALPHA * scratch[src_slot][3][row * 8u + col2]);
 
90
  }
91
  }
92
 
@@ -97,6 +117,11 @@ fn main(
97
  @builtin(subgroup_invocation_id) sg_id: u32,
98
  @builtin(subgroup_size) sg_size: u32
99
  ) {
 
 
 
 
 
100
  let batch = workgroup_id.z;
101
  let a_base = batch * A_BATCH_STRIDE;
102
  let b_base = batch * B_BATCH_STRIDE;
@@ -125,7 +150,7 @@ fn main(
125
  workgroupBarrier();
126
 
127
  for (var step = 0u; step < TILE_K; step = step + 8u) {
128
- {% set directInputs = directMatrixInputs is defined and directMatrixInputs %}
129
  let matrix_a_offset = subtile_idy * SUB_ROWS * TILE_K + step;
130
  {% for r in range(2) %}
131
  var matA{{ r }}: subgroup_matrix_left<{{ operandScalar }}, 8, 8> = subgroupMatrixLoad<subgroup_matrix_left<{{ operandScalar }}, 8, 8>, row_major>(&{{ "w" if directInputs else "tile_A" }}, matrix_a_offset{% if r > 0 %} + 8u * {{ "K" if directInputs else "TILE_K" }}{% endif %}, {{ "K" if directInputs else "TILE_K" }});
@@ -160,9 +185,12 @@ fn main(
160
  workgroupBarrier();
161
  let row = sg_id / 4u;
162
  let col = (sg_id % 4u) * 2u;
 
 
 
163
  var matrix_c_offset = c_base + (a_global_base + base_A) * N + b_global_base + base_B;
164
  var row_limit = i32(M) - i32(a_global_base + base_A);
165
- storeOutput(matrix_c_offset, row, col, subtile_id, row_limit);
166
  workgroupBarrier();
167
 
168
  subgroupMatrixStore<row_major>(&scratch[subtile_id][0], 0u, matC10, 8u);
@@ -172,5 +200,5 @@ fn main(
172
  workgroupBarrier();
173
  matrix_c_offset = matrix_c_offset + 8u * N;
174
  row_limit = i32(M) - i32(a_global_base + base_A + 8u);
175
- storeOutput(matrix_c_offset, row, col, subtile_id, row_limit);
176
  }
 
1
  // com.microsoft.FusedMatMul subgroup-matrix specialization: Y = alpha * op(A) @ op(B).
2
  // transA and transB transpose the corresponding matrix operand on load.
3
  // Dense batches map through workgroup_id.z, and M-tail rows are guarded by row_limit.
4
+ // The row count M arrives per call in `params.M`, so one pipeline serves every M;
5
+ // K, N and the batch layout compile in.
6
+ // A K % 32 == 0 gate keeps the reduction loop whole; N is free of the 64-wide column
7
+ // tile because nTailSafe clamps the trailing tile's B columns to N - 1 and guards
8
+ // every store on col < N. Both operands stage through workgroup memory, so
9
+ // subgroupMatrixLoad only ever reads the full tile_A/tile_B arrays and never sees a
10
+ // partial 8x8 tile at any M or N. The clamp is not interchangeable with a zero fill:
11
+ // an out-of-bounds subgroupMatrixLoad resets to offset 0 and returns a different
12
+ // valid tile, and a duplicated real column keeps the discarded accumulators finite.
13
+ // The batch is required to match between A and B (no broadcast) because a_base and
14
+ // b_base both index by the same workgroup_id.z.
15
  enable subgroups;
16
  {% if pinSubgroupSize32 %}
17
  enable subgroup_size_control;
 
19
  enable chromium_experimental_subgroup_matrix;
20
  diagnostic(off, chromium.subgroup_matrix_uniformity);
21
 
 
22
  {{ env.wgsl.resourceDeclarations }}
23
 
24
  {% set operandScalar = fScalar %}
25
  {% set accScalar = "f32" %}
26
+ {% set N_TAIL = nTailSafe is defined and nTailSafe %}
27
 
 
28
  const K: u32 = {{ K }}u;
29
  const N: u32 = {{ N }}u;
30
  {% if transBatchA %}
31
  const BATCH_COUNT: u32 = {{ batchCount }}u;
32
  const A_BATCH_STRIDE: u32 = K;
33
  const A_M_STRIDE: u32 = BATCH_COUNT * K;
34
+ {% elif not transA %}
 
 
35
  const A_M_STRIDE: u32 = K;
36
  {% endif %}
 
37
  const B_BATCH_STRIDE: u32 = K * N;
 
38
  const ALPHA: {{ accScalar }} = {{ accScalar }}({{ alpha }});
39
  const TILE_COLS: u32 = 64u;
40
  const TILE_ROWS: u32 = 32u;
 
51
  let col = c_idx * 8u;
52
  for (var col_offset = 0u; col_offset < 8u; col_offset = col_offset + 1u) {
53
  let k = k_idx + col + col_offset;
54
+ if (a_global < params.M) {
55
  {% if transA %}
56
  // op(A) = A^T: A stored [.., K, M], so op(A)[a_global, k] = A[k, a_global].
57
+ tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(a[a_base + k * params.M + a_global]);
58
  {% else %}
59
  tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(a[a_base + a_global * A_M_STRIDE + k]);
60
  {% endif %}
 
65
  }
66
 
67
  fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
68
+ {% if N_TAIL %}
69
+ // The trailing column tile reads column N - 1 in place of every column past N.
70
+ // storeOutput discards those lanes; duplicating a real column keeps them finite.
71
+ let b_col = min(tile_base + row, N - 1u);
72
+ {% else %}
73
  let b_col = tile_base + row;
74
+ {% endif %}
75
  let col = c_idx * 16u;
76
  for (var i = 0u; i < 16u; i = i + 1u) {
77
  let k = k_idx + col + i;
 
84
  }
85
  }
86
 
87
+ fn storeOutput(offset: u32{% if N_TAIL %}, col_base: u32{% endif %}, row: u32, col: u32, src_slot: u32, row_limit: i32) {
88
  if (row_limit > 0 && row < u32(row_limit)) {
89
  let col2 = col + 1u;
90
+ {% if N_TAIL %}
91
+ {% for block in range(4) %}
92
+ if (col_base + col + {{ block * 8 }}u < N) {
93
+ y[offset + row * N + col + {{ block * 8 }}u] = {{ outScalar }}(ALPHA * scratch[src_slot][{{ block }}][row * 8u + col]);
94
+ }
95
+ if (col_base + col2 + {{ block * 8 }}u < N) {
96
+ y[offset + row * N + col2 + {{ block * 8 }}u] = {{ outScalar }}(ALPHA * scratch[src_slot][{{ block }}][row * 8u + col2]);
97
+ }
98
+ {% endfor %}
99
+ {% else %}
100
  y[offset + row * N + col] = {{ outScalar }}(ALPHA * scratch[src_slot][0][row * 8u + col]);
101
  y[offset + row * N + col + 8u] = {{ outScalar }}(ALPHA * scratch[src_slot][1][row * 8u + col]);
102
  y[offset + row * N + col + 16u] = {{ outScalar }}(ALPHA * scratch[src_slot][2][row * 8u + col]);
 
106
  y[offset + row * N + col2 + 8u] = {{ outScalar }}(ALPHA * scratch[src_slot][1][row * 8u + col2]);
107
  y[offset + row * N + col2 + 16u] = {{ outScalar }}(ALPHA * scratch[src_slot][2][row * 8u + col2]);
108
  y[offset + row * N + col2 + 24u] = {{ outScalar }}(ALPHA * scratch[src_slot][3][row * 8u + col2]);
109
+ {% endif %}
110
  }
111
  }
112
 
 
117
  @builtin(subgroup_invocation_id) sg_id: u32,
118
  @builtin(subgroup_size) sg_size: u32
119
  ) {
120
+ let M = params.M;
121
+ {% if not transBatchA %}
122
+ let A_BATCH_STRIDE = M * K;
123
+ {% endif %}
124
+ let C_BATCH_STRIDE = M * N;
125
  let batch = workgroup_id.z;
126
  let a_base = batch * A_BATCH_STRIDE;
127
  let b_base = batch * B_BATCH_STRIDE;
 
150
  workgroupBarrier();
151
 
152
  for (var step = 0u; step < TILE_K; step = step + 8u) {
153
+ {% set directInputs = false %}
154
  let matrix_a_offset = subtile_idy * SUB_ROWS * TILE_K + step;
155
  {% for r in range(2) %}
156
  var matA{{ r }}: subgroup_matrix_left<{{ operandScalar }}, 8, 8> = subgroupMatrixLoad<subgroup_matrix_left<{{ operandScalar }}, 8, 8>, row_major>(&{{ "w" if directInputs else "tile_A" }}, matrix_a_offset{% if r > 0 %} + 8u * {{ "K" if directInputs else "TILE_K" }}{% endif %}, {{ "K" if directInputs else "TILE_K" }});
 
185
  workgroupBarrier();
186
  let row = sg_id / 4u;
187
  let col = (sg_id % 4u) * 2u;
188
+ {% if N_TAIL %}
189
+ let col_base = b_global_base + base_B;
190
+ {% endif %}
191
  var matrix_c_offset = c_base + (a_global_base + base_A) * N + b_global_base + base_B;
192
  var row_limit = i32(M) - i32(a_global_base + base_A);
193
+ storeOutput(matrix_c_offset{% if N_TAIL %}, col_base{% endif %}, row, col, subtile_id, row_limit);
194
  workgroupBarrier();
195
 
196
  subgroupMatrixStore<row_major>(&scratch[subtile_id][0], 0u, matC10, 8u);
 
200
  workgroupBarrier();
201
  matrix_c_offset = matrix_c_offset + 8u * N;
202
  row_limit = i32(M) - i32(a_global_base + base_A + 8u);
203
+ storeOutput(matrix_c_offset{% if N_TAIL %}, col_base{% endif %}, row, col, subtile_id, row_limit);
204
  }
build/webgpu/manifest.json CHANGED
@@ -20,8 +20,10 @@
20
  "typeConstraints": { "T": ["float32", "float16"] },
21
  "tunables": {
22
  "TILED_REG_MIN_WORKGROUPS": { "default": 64 },
 
23
  "GEMV_TARGET_BLOCKS": { "default": 512 },
24
  "SUBGROUP_MATRIX_MIN_M": { "default": 2 },
 
25
  "SUBGROUP_MATRIX_SPLITK_TARGET_WGS": { "default": 512 },
26
  "SUBGROUP_MATRIX_SPLITK_MIN_K": { "default": 1024 },
27
  "SUBGROUP_MATRIX_SPLITK_MAX_TILES": { "default": 128 },
@@ -33,46 +35,176 @@
33
  "TRANSBATCH_B_SUBGROUP_MATRIX_MIN_WORKGROUPS": { "default": 64 },
34
  "TRANSBATCH_REG_MIN_K": { "default": 128 },
35
  "BAND_PREFER_MAX_ROWS": { "default": 8 },
36
- "BAND_PREFER_DEEP_K": { "default": 4096 }
 
 
37
  },
38
  "derive": {
39
- "gemvWorkgroups": "ceilDiv(dim(shapes.B, 1), 128)",
40
- "gemvSliceCap": "min(32, device.limits.maxComputeWorkgroupSizeY, floor(device.limits.maxComputeInvocationsPerWorkgroup / 32), floor(device.limits.maxComputeWorkgroupStorageSize / 512))",
41
- "gemvSlicesPlan": "max(1, min(gemvSliceCap, max(8, pow2ceil(ceilDiv(tunables.GEMV_TARGET_BLOCKS, gemvWorkgroups)))))",
42
  "batchMovedAShape": "moveAxis(shapes.A, 0, -2) if attrs.transBatchA != 0 else shapes.A",
43
  "batchMovedBShape": "moveAxis(shapes.B, 0, -2) if attrs.transBatchB != 0 else shapes.B",
44
  "logicalAShape": "moveAxis(batchMovedAShape, -1, -2) if attrs.transA != 0 and ranks.A > 1 else batchMovedAShape",
45
  "logicalBShape": "moveAxis(batchMovedBShape, -1, -2) if attrs.transB != 0 and ranks.B > 1 else batchMovedBShape",
46
  "transBatchContract": "(attrs.transBatchA == 0 and attrs.transBatchB == 0) or (ranks.A == ranks.B and ranks.A >= 3)",
 
47
  "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32",
48
  "canPinSubgroupSize32": "device.features.has(\"subgroups\") and device.features.has(\"subgroup-size-control\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize <= 32 and device.adapterInfo.subgroupMaxSize >= 32",
49
  "pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter",
50
  "wave32Effective": "wave32Adapter or pinSubgroupSize32",
 
51
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
52
- "fusedSgmatRank2Ok": "ranks.A == 2 and ranks.B == 2 and ranks.Y == 2 and attrs.transA == 0 and attrs.transB == 0 and attrs.transBatchA == 0 and attrs.transBatchB == 0 and dim(shapes.A, 1) == dim(shapes.B, 0) and dim(shapes.Y, 0) == dim(shapes.A, 0) and dim(shapes.Y, 1) == dim(shapes.B, 1)",
 
 
 
 
 
 
 
 
 
 
 
 
 
 
53
  "subgroupMatrixResourcesFit": "128 <= deviceWorkgroupCap and ((32 * 32 + 64 * 32) * dtypeBytes(dtypes.T) + 4 * 4 * 64 * 4) <= device.limits.maxComputeWorkgroupStorageSize",
54
  "sgmatSplitKDepth": "dim(shapes.A, ranks.A - 1)",
55
- "sgmatOutTiles": "ceilDiv(dim(shapes.A, 0), 32) * ceilDiv(dim(shapes.B, 1), 64) if fusedSgmatRank2Ok else 1",
56
- "sgmatSplitKWant": "ceilDiv(tunables.SUBGROUP_MATRIX_SPLITK_TARGET_WGS, sgmatOutTiles)",
57
  "sgmatSplitK32Ok": "sgmatSplitKDepth % 1024 == 0",
58
  "sgmatSplitK16Ok": "sgmatSplitKDepth % 512 == 0",
59
  "sgmatSplitK8Ok": "sgmatSplitKDepth % 256 == 0",
60
  "sgmatSplitK4Ok": "sgmatSplitKDepth % 128 == 0",
61
  "sgmatSplitK2Ok": "sgmatSplitKDepth % 64 == 0",
 
 
 
 
 
 
 
 
 
 
 
 
 
 
62
  "sgmatSplitK": "32 if (sgmatSplitKWant > 16 and sgmatSplitK32Ok) else (16 if (sgmatSplitKWant > 8 and sgmatSplitK16Ok) else (8 if (sgmatSplitKWant > 4 and sgmatSplitK8Ok) else (4 if (sgmatSplitKWant > 2 and sgmatSplitK4Ok) else (2 if sgmatSplitK2Ok else 1))))",
63
  "bandRank2Ok": "ranks.A == 2 and ranks.B == 2 and ranks.Y == 2 and attrs.transA == 0 and attrs.transB == 0 and attrs.transBatchA == 0 and attrs.transBatchB == 0 and dim(shapes.A, 1) == dim(shapes.B, 0) and dim(shapes.Y, 0) == dim(shapes.A, 0) and dim(shapes.Y, 1) == dim(shapes.B, 1)",
64
- "bandSplitWant": "pow2ceil(ceilDiv(tunables.BAND_SPLIT_TARGET_WORKGROUPS, max(1, gemvWorkgroups)))",
65
- "bandSplitK": "16 if (bandSplitWant >= 16 and dim(shapes.A, ranks.A - 1) >= 4096) else (8 if (bandSplitWant >= 8 and dim(shapes.A, ranks.A - 1) >= 2048) else (4 if (bandSplitWant >= 4 and dim(shapes.A, ranks.A - 1) >= 1024) else (2 if (bandSplitWant >= 2 and dim(shapes.A, ranks.A - 1) >= 512) else 1)))"
66
  },
67
  "bindings": {
68
- "a": { "arg": "A", "buffer": "read-only-storage", "elementType": "$scalar" },
69
- "b": { "arg": "B", "buffer": "read-only-storage", "elementType": "$vectorScalar" },
70
  "partials": { "buffer": "read-only-storage", "elementType": "f32" },
71
- "y": { "arg": "Y", "buffer": "storage", "elementType": "$scalar" },
72
- "params": { "buffer": "uniform", "struct": [{ "name": "cols", "type": "u32", "value": "numel(shapes.Y)" }] },
73
- "b_3": { "arg": "B", "name": "b", "buffer": "read-only-storage", "elementType": "$scalar" }
 
74
  },
75
  "variants": [
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
76
  {
77
  "id": "subgroup_matrix_transbatch_b_f16",
78
  "priority": 11,
@@ -83,14 +215,11 @@
83
  },
84
  "derive": {
85
  "hasBias": false,
86
- "usesF16": "dtypes.T == \"f16\"",
87
  "fScalar": "dtypes.T",
88
  "outScalar": "dtypes.T",
89
- "scalar": "dtypes.T",
90
  "generalAddressing": true,
91
  "outputBuffer": "\"y\"",
92
- "alpha": "attrs.alpha",
93
- "M": "dim(shapes.A, ranks.A - 2)",
94
  "K": "dim(shapes.A, ranks.A - 1)",
95
  "N": "dim(shapes.B, ranks.B - 1)",
96
  "batchCount": "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 1))"
@@ -101,14 +230,12 @@
101
  "name": "FusedMatMul.SubgroupMatrixTransBatchB",
102
  "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
103
  "derive": {
104
- "aShape": "shapes.A",
105
  "bShape": "logicalBShape",
106
- "aRank": "ranks.A",
107
- "bRank": "ranks.B",
108
  "bStorageStrides": "[dim(shapes.B, 2), dim(shapes.B, 1) * dim(shapes.B, 2), 1]"
109
  },
110
- "bindings": ["a", "b_3", "y"],
111
- "dispatch": { "x": "ceil(N / 64)", "y": "ceil(M / 32)", "z": "batchCount" }
112
  }
113
  ]
114
  },
@@ -122,14 +249,11 @@
122
  },
123
  "derive": {
124
  "hasBias": false,
125
- "usesF16": "dtypes.T == \"f16\"",
126
  "fScalar": "dtypes.T",
127
  "outScalar": "dtypes.T",
128
- "scalar": "dtypes.T",
129
  "generalAddressing": true,
130
  "outputBuffer": "\"y\"",
131
- "alpha": "attrs.alpha",
132
- "M": "dim(shapes.A, ranks.A - 2)",
133
  "K": "dim(shapes.A, ranks.A - 1)",
134
  "N": "dim(shapes.B, ranks.B - 1)",
135
  "batchCount": "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 1))"
@@ -140,31 +264,34 @@
140
  "name": "FusedMatMul.SubgroupMatrixTransBatchB",
141
  "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
142
  "derive": {
143
- "aShape": "shapes.A",
144
  "bShape": "logicalBShape",
145
- "aRank": "ranks.A",
146
- "bRank": "ranks.B",
147
  "bStorageStrides": "[dim(shapes.B, 2), dim(shapes.B, 1) * dim(shapes.B, 2), 1]"
148
  },
149
- "bindings": ["a", "b_3", "y"],
150
- "dispatch": { "x": "ceil(N / 64)", "y": "ceil(M / 32)", "z": "batchCount" }
151
  }
152
  ]
153
  },
154
  {
155
- "id": "f32_m1_gemv_vec4",
156
  "priority": 30,
157
- "when": ["dtypes.T == \"f32\"", "attrs.transA == 0", "attrs.transB == 0", "attrs.transBatchA == 0", "attrs.transBatchB == 0", "ranks.A == 2", "ranks.B == 2", "ranks.Y == 2", "dim(shapes.A, 0) == 1", "dim(shapes.Y, 0) == 1", "dim(shapes.A, 1) == dim(shapes.B, 0)", "dim(shapes.Y, 1) == dim(shapes.B, 1)", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "ceil(dim(shapes.B, 1) / 128) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
158
- "derive": { "usesF16": false, "gemvSlices": "gemvSlicesPlan", "alphaScale": "attrs.alpha" },
 
 
 
 
 
159
  "passes": [
160
  {
161
  "id": "main",
162
- "name": "FusedMatMul.F32M1GemvVec4",
163
  "shader": "matmul-vector-matrix-vec4.wgsl.jinja",
164
  "bindings": [
165
- { "arg": "A", "name": "a", "elementType": "f32" },
166
- { "arg": "B", "name": "b", "elementType": "vec4<f32>" },
167
- { "arg": "Y", "name": "c", "elementType": "vec4<f32>" },
168
  {
169
  "name": "params",
170
  "struct": [
@@ -173,21 +300,17 @@
173
  ]
174
  }
175
  ],
176
- "dispatch": { "x": "ceil(dim(shapes.B, 1) / 128)" }
177
  }
178
  ]
179
  },
180
  {
181
  "id": "rank2_band_vec4_splitk",
182
  "priority": 11,
183
- "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T)", "bandRank2Ok", "dim(shapes.A, 0) >= 2", "dim(shapes.A, 0) <= tunables.BAND_VEC4_MAX_ROWS", "dim(shapes.A, 1) > 0", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "gemvWorkgroups <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "32 <= device.limits.maxComputeWorkgroupSizeX", "gemvWorkgroups <= tunables.BAND_SPLIT_MAX_COLUMN_GROUPS", "bandSplitK >= 2", "bandSplitK * numel(shapes.Y) * 4 <= device.limits.maxStorageBufferBindingSize", "bandSplitK <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tunables.BAND_SPLIT_SLICES <= device.limits.maxComputeWorkgroupSizeY", "32 * tunables.BAND_SPLIT_SLICES <= device.limits.maxComputeInvocationsPerWorkgroup"],
184
  "derive": {
185
- "usesF16": "dtypes.T == \"f16\"",
186
- "scalar": "dtypes.T",
187
- "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
188
  "batched": false,
189
  "outputBuffer": "\"y\"",
190
- "alpha": "attrs.alpha",
191
  "M": "dim(shapes.A, 0)",
192
  "K": "dim(shapes.A, 1)",
193
  "N": "dim(shapes.B, 1)",
@@ -209,7 +332,7 @@
209
  "id": "combine",
210
  "name": "FusedMatMul.Rank2BandVec4SplitKCombine",
211
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
212
- "derive": { "op": "\"sum\"", "outputF16": "dtypes.T == \"f16\"", "intMode": false },
213
  "bindings": ["partials", "y", "params"],
214
  "dispatch": {
215
  "x": "min(ceilDiv((numel(shapes.Y)), (256)), 65535)",
@@ -222,18 +345,13 @@
222
  {
223
  "id": "rank2_band_vec4",
224
  "priority": 11,
225
- "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T)", "bandRank2Ok", "dim(shapes.A, 0) >= 2", "dim(shapes.A, 0) <= tunables.BAND_VEC4_MAX_ROWS", "dim(shapes.A, 1) > 0", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "gemvWorkgroups <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "32 <= device.limits.maxComputeWorkgroupSizeX", "gemvSlicesPlan <= device.limits.maxComputeWorkgroupSizeY", "32 * gemvSlicesPlan <= device.limits.maxComputeInvocationsPerWorkgroup", "32 * gemvSlicesPlan * 16 <= device.limits.maxComputeWorkgroupStorageSize", "not (gemvWorkgroups <= tunables.BAND_SPLIT_MAX_COLUMN_GROUPS and bandSplitK >= 2)"],
226
  "derive": {
227
- "usesF16": "dtypes.T == \"f16\"",
228
- "scalar": "dtypes.T",
229
- "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
230
  "batched": false,
231
  "outputBuffer": "\"y\"",
232
- "alpha": "attrs.alpha",
233
  "M": "dim(shapes.A, 0)",
234
  "K": "dim(shapes.A, 1)",
235
- "N": "dim(shapes.B, 1)",
236
- "gemvSlices": "gemvSlicesPlan"
237
  },
238
  "passes": [
239
  {
@@ -243,24 +361,19 @@
243
  "bindings": ["a", "b", { "arg": "Y", "name": "y", "elementType": "$vectorScalar" }],
244
  "dispatch": { "x": "gemvWorkgroups" }
245
  }
246
- ],
247
- "demoteWhen": ["false"]
248
  },
249
  {
250
  "id": "rank2_band_vec4_f32_preferred",
251
  "priority": 13,
252
- "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T)", "bandRank2Ok", "dim(shapes.A, 0) >= 2", "dim(shapes.A, 0) <= tunables.BAND_VEC4_MAX_ROWS", "dim(shapes.A, 1) > 0", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "gemvWorkgroups <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "32 <= device.limits.maxComputeWorkgroupSizeX", "gemvSlicesPlan <= device.limits.maxComputeWorkgroupSizeY", "32 * gemvSlicesPlan <= device.limits.maxComputeInvocationsPerWorkgroup", "32 * gemvSlicesPlan * 16 <= device.limits.maxComputeWorkgroupStorageSize", "not (gemvWorkgroups <= tunables.BAND_SPLIT_MAX_COLUMN_GROUPS and bandSplitK >= 2)"],
 
253
  "derive": {
254
- "usesF16": "dtypes.T == \"f16\"",
255
- "scalar": "dtypes.T",
256
- "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
257
  "batched": false,
258
  "outputBuffer": "\"y\"",
259
- "alpha": "attrs.alpha",
260
  "M": "dim(shapes.A, 0)",
261
  "K": "dim(shapes.A, 1)",
262
- "N": "dim(shapes.B, 1)",
263
- "gemvSlices": "gemvSlicesPlan"
264
  },
265
  "passes": [
266
  {
@@ -270,8 +383,7 @@
270
  "bindings": ["a", "b", { "arg": "Y", "name": "y", "elementType": "$vectorScalar" }],
271
  "dispatch": { "x": "gemvWorkgroups" }
272
  }
273
- ],
274
- "demoteWhen": ["dtypes.T != \"f32\" or (dim(shapes.A, 0) > tunables.BAND_PREFER_MAX_ROWS and dim(shapes.A, 1) >= tunables.BAND_PREFER_DEEP_K)"]
275
  },
276
  {
277
  "id": "subgroup_matrix_splitk",
@@ -285,16 +397,12 @@
285
  ]
286
  },
287
  "derive": {
288
- "usesF16": "dtypes.T == \"f16\"",
289
- "fScalar": "\"f16\" if dtypes.T == \"f16\" else \"f32\"",
290
- "scalar": "dtypes.T",
291
  "hasBias": false,
292
  "generalAddressing": true,
293
  "tailSafe": false,
294
  "outputBuffer": "\"partials\"",
295
  "outScalar": "\"f32\"",
296
- "alpha": "attrs.alpha",
297
- "M": "dim(shapes.A, 0)",
298
  "K": "dim(shapes.A, 1)",
299
  "N": "dim(shapes.B, 1)",
300
  "batchCount": 1,
@@ -309,20 +417,15 @@
309
  "id": "partial",
310
  "name": "FusedMatMul.SubgroupMatrixSplitK",
311
  "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
312
- "derive": {
313
- "aShape": ["dim(shapes.A, 0)", "dim(shapes.A, 1)"],
314
- "bShape": ["dim(shapes.B, 0)", "dim(shapes.B, 1)"],
315
- "aRank": 2,
316
- "bRank": 2
317
- },
318
- "bindings": ["a", "b_3", { "name": "partials", "buffer": "storage", "elementType": "f32" }],
319
  "dispatch": { "x": "ceilDiv(dim(shapes.Y, 1), 64)", "y": "ceilDiv(dim(shapes.Y, 0), 32)", "z": "sgmatSplitK" }
320
  },
321
  {
322
  "id": "combine",
323
  "name": "FusedMatMul.SubgroupMatrixSplitKCombine",
324
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
325
- "derive": { "op": "\"sum\"", "outputF16": "dtypes.T == \"f16\"", "intMode": false },
326
  "bindings": ["partials", "y", "params"],
327
  "dispatch": {
328
  "x": "min(ceilDiv((numel(shapes.Y)), (256)), 65535)",
@@ -342,34 +445,35 @@
342
  },
343
  "derive": {
344
  "hasBias": false,
345
- "usesF16": true,
346
  "fScalar": "\"f16\"",
347
  "outScalar": "\"f16\"",
348
- "scalar": "dtypes.T",
349
  "generalAddressing": true,
350
  "tailSafe": "dim(shapes.A, ranks.A - 1) % 32 != 0 or dim(shapes.B, ranks.B - 1) % 64 != 0",
351
  "outputBuffer": "\"y\"",
352
- "alpha": "attrs.alpha",
353
- "M": "dim(shapes.A, ranks.A - 2)",
354
  "K": "dim(shapes.A, ranks.A - 1)",
355
  "N": "dim(shapes.B, ranks.B - 1)",
356
- "batchCount": "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 1))"
357
  },
358
  "passes": [
359
  {
360
  "id": "main",
361
  "name": "FusedMatMul.SubgroupMatrixTailBroadcast",
362
  "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
363
- "derive": { "aShape": "shapes.A", "bShape": "shapes.B", "aRank": "ranks.A", "bRank": "ranks.B" },
364
- "bindings": ["a", "b_3", "y"],
365
- "dispatch": { "x": "ceil(N / 64)", "y": "ceil(M / 32)", "z": "batchCount" }
 
 
 
 
366
  }
367
  ]
368
  },
369
  {
370
  "id": "subgroup_matrix",
371
  "priority": 10,
372
- "when": ["f16Ok(dtypes.T)", "attrs.transBatchA == 0 or (attrs.transA == 0 and ranks.A == 3)", "attrs.transBatchB == 0", "ranks.A >= 2", "ranks.B == ranks.A", "ranks.Y == ranks.A", "(dim(shapes.A, ranks.A - 2) if attrs.transA != 0 else dim(shapes.A, ranks.A - 1)) == (dim(shapes.B, ranks.B - 1) if attrs.transB != 0 else dim(shapes.B, ranks.B - 2))", "(dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2))) >= tunables.SUBGROUP_MATRIX_MIN_M", "(dim(shapes.A, ranks.A - 2) if attrs.transA != 0 else dim(shapes.A, ranks.A - 1)) % 32 == 0", "(dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1)) % 64 == 0", "(ranks.A == 2 or (ranks.A == 3 and (dim(shapes.A, 1) if attrs.transBatchA != 0 else dim(shapes.A, 0)) == dim(shapes.B, 0) and dim(shapes.Y, 0) == dim(shapes.B, 0)) or (ranks.A == 4 and dim(shapes.A, 0) == dim(shapes.B, 0) and dim(shapes.A, 1) == dim(shapes.B, 1) and dim(shapes.Y, 0) == dim(shapes.A, 0) and dim(shapes.Y, 1) == dim(shapes.A, 1)))", "dim(shapes.Y, ranks.Y - 2) == (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))", "dim(shapes.Y, ranks.Y - 1) == (dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1))", "ceil((dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1)) / 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil((dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2))) / 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "numel(shapes.Y) / ((dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2))) * (dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1))) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "wave32Effective"],
373
  "requires": {
374
  "features": ["subgroups", "chromium-experimental-subgroup-matrix"],
375
  "subgroupMatrixConfigs": [
@@ -378,14 +482,12 @@
378
  ]
379
  },
380
  "derive": {
381
- "usesF16": "dtypes.T == \"f16\"",
382
- "fScalar": "\"f16\" if dtypes.T == \"f16\" else \"f32\"",
383
  "outScalar": "\"f16\" if dtypes.T == \"f16\" else \"f32\"",
384
- "scalar": "dtypes.T",
385
  "transA": "attrs.transA != 0",
386
  "transB": "attrs.transB != 0",
387
  "transBatchA": "attrs.transBatchA != 0",
388
- "M": "(dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))",
389
  "K": "(dim(shapes.A, ranks.A - 2) if attrs.transA != 0 else dim(shapes.A, ranks.A - 1))",
390
  "N": "(dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1))",
391
  "batchCount": "numel(shapes.Y) / ((dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2))) * (dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1)))"
@@ -395,85 +497,80 @@
395
  "id": "main",
396
  "name": "FusedMatMul.SubgroupMatrix",
397
  "shader": "fused-matmul-subgroup-matrix.wgsl.jinja",
398
- "derive": { "alpha": "attrs.alpha" },
399
- "bindings": ["a", "b_3", "y"],
400
- "dispatch": { "x": "ceil(N / 64)", "y": "ceil(M / 32)", "z": "numel(shapes.Y) / (M * N)" }
401
  }
402
- ]
 
403
  },
404
  {
405
  "id": "broadcast_rank4_tiled_reg",
406
  "priority": 6,
407
- "when": ["f16Ok(dtypes.T)", "attrs.transA == 0", "attrs.transB == 0", "attrs.transBatchA == 0", "attrs.transBatchB == 0", "ranks.A == 4", "(ranks.B == 2 or ranks.B == 3)", "ranks.Y == 4", "dim(shapes.Y, 0) == dim(shapes.A, 0)", "(ranks.B == 2 or dim(shapes.A, 1) == dim(shapes.B, 0) or dim(shapes.A, 1) == 1 or dim(shapes.B, 0) == 1)", "dim(shapes.Y, 1) == (dim(shapes.A, 1) if ranks.B == 2 else max(dim(shapes.A, 1), dim(shapes.B, 0)))", "dim(shapes.A, 3) == dim(shapes.B, ranks.B - 2)", "dim(shapes.Y, 2) == dim(shapes.A, 2)", "dim(shapes.Y, 3) == dim(shapes.B, ranks.B - 1)", "dim(shapes.A, 2) >= 64", "dim(shapes.A, 3) >= 32", "dim(shapes.B, ranks.B - 1) >= 64", "ceil(dim(shapes.B, ranks.B - 1) / 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil(dim(shapes.A, 2) / 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "numel(shapes.Y) / (dim(shapes.A, 2) * dim(shapes.B, ranks.B - 1)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
408
- "derive": { "scalar": "dtypes.T", "usesF16": "dtypes.T == \"f16\"" },
409
  "passes": [
410
  {
411
  "id": "main",
412
  "name": "FusedMatMul.BroadcastRank4TiledReg",
413
  "shader": "matmul-tiled-general-reg.wgsl.jinja",
414
  "derive": {
415
- "aShape": "shapes.A",
 
 
416
  "bShape": "shapes.B",
417
- "alpha": "attrs.alpha",
418
- "aRank": "ranks.A",
419
- "bRank": "ranks.B",
420
  "transBatchA": "false"
421
  },
422
- "bindings": ["a", "b_3", "y"],
423
  "dispatch": {
424
- "x": "ceil(dim(shapes.B, ranks.B - 1) / 64)",
425
- "y": "ceil(dim(shapes.A, 2) / 64)",
426
- "z": "numel(shapes.Y) / (dim(shapes.A, 2) * dim(shapes.B, ranks.B - 1))"
427
  }
428
  }
429
- ]
 
 
 
 
430
  },
431
  {
432
  "id": "plain_rank2_tiled_reg",
433
  "priority": 4,
434
- "when": ["f16Ok(dtypes.T)", "fusedSgmatRank2Ok", "dim(shapes.A, 0) >= 64", "dim(shapes.A, 1) >= 32", "dim(shapes.B, 1) >= 64", "ceil(dim(shapes.A, 0) / 64) * ceil(dim(shapes.B, 1) / 64) >= tunables.TILED_REG_MIN_WORKGROUPS", "ceil(dim(shapes.B, 1) / 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil(dim(shapes.A, 0) / 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
435
- "derive": { "scalar": "dtypes.T", "usesF16": "dtypes.T == \"f16\"" },
436
  "passes": [
437
  {
438
  "id": "main",
439
  "name": "FusedMatMul.PlainRank2TiledReg",
440
  "shader": "matmul-tiled-general-reg.wgsl.jinja",
441
  "derive": {
442
- "aShape": "shapes.A",
 
443
  "bShape": "shapes.B",
444
- "alpha": "attrs.alpha",
445
- "aRank": "ranks.A",
446
- "bRank": "ranks.B",
447
  "transBatchA": "false"
448
  },
449
- "bindings": ["a", "b_3", "y"],
450
- "dispatch": { "x": "ceil(dim(shapes.B, 1) / 64)", "y": "ceil(dim(shapes.A, 0) / 64)", "z": 1 }
 
 
 
 
451
  }
452
  ]
453
  },
454
  {
455
  "id": "transbatch_a_tiled_reg",
456
  "priority": 5,
457
- "when": ["f16Ok(dtypes.T)", "attrs.transBatchA != 0", "attrs.transBatchB == 0", "attrs.transA == 0", "attrs.transB == 0", "ranks.A == 3", "ranks.B == 3", "ranks.Y == 3", "dim(shapes.A, 1) == dim(shapes.B, 0)", "dim(shapes.Y, 0) == dim(shapes.B, 0)", "dim(shapes.A, 2) == dim(shapes.B, 1)", "dim(shapes.Y, 1) == dim(shapes.A, 0)", "dim(shapes.Y, 2) == dim(shapes.B, 2)", "dim(shapes.A, 0) >= 64", "dim(shapes.A, 2) >= 32", "dim(shapes.B, 2) >= 64", "ceil(dim(shapes.B, 2) / 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil(dim(shapes.A, 0) / 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dim(shapes.Y, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
458
- "derive": { "scalar": "dtypes.T", "usesF16": "dtypes.T == \"f16\"" },
459
  "passes": [
460
  {
461
  "id": "main",
462
  "name": "FusedMatMul.TransBatchATiledReg",
463
  "shader": "matmul-tiled-general-reg.wgsl.jinja",
464
- "derive": {
465
- "aShape": "shapes.A",
466
- "bShape": "shapes.B",
467
- "alpha": "attrs.alpha",
468
- "aRank": "ranks.A",
469
- "bRank": "ranks.B",
470
- "transBatchA": "true",
471
- "kTile": "4"
472
- },
473
- "bindings": ["a", "b_3", "y"],
474
  "dispatch": {
475
- "x": "ceil(dim(shapes.B, 2) / 64)",
476
- "y": "ceil(dim(shapes.A, 0) / 64)",
477
  "z": "dim(shapes.Y, 0)"
478
  }
479
  }
@@ -483,7 +580,6 @@
483
  "id": "tiled",
484
  "priority": 0,
485
  "when": ["ranks.A >= 1", "ranks.B >= 1", "f16Ok(dtypes.T)", "transBatchContract", "(dim(shapes.A, 0) if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA == 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))) == (dim(shapes.B, 0) if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB != 0 else (dim(shapes.B, 0) if attrs.transBatchB != 0 else dim(shapes.B, ranks.B - 2))))", "ceil((1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))) / 16) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil((1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else (dim(shapes.B, 0) if attrs.transBatchB != 0 else dim(shapes.B, ranks.B - 2)))) / 16) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "numel(shapes.Y) / max(1, (1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))) * (1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else (dim(shapes.B, 0) if attrs.transBatchB != 0 else dim(shapes.B, ranks.B - 2))))) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
486
- "derive": { "scalar": "dtypes.T", "usesF16": "dtypes.T == \"f16\"" },
487
  "passes": [
488
  {
489
  "id": "main",
@@ -494,16 +590,13 @@
494
  "bShape": "shapes.B",
495
  "transA": "attrs.transA != 0",
496
  "transB": "attrs.transB != 0",
497
- "alpha": "attrs.alpha",
498
- "aRank": "ranks.A",
499
- "bRank": "ranks.B",
500
  "transBatchA": "attrs.transBatchA != 0",
501
  "transBatchB": "attrs.transBatchB != 0"
502
  },
503
- "bindings": ["a", "b_3", "y"],
504
  "dispatch": {
505
- "x": "ceil((1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else (dim(shapes.B, 0) if attrs.transBatchB != 0 else dim(shapes.B, ranks.B - 2)))) / 32)",
506
- "y": "ceil((1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))) / 32)",
507
  "z": "numel(shapes.Y) / max(1, (1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))) * (1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else (dim(shapes.B, 0) if attrs.transBatchB != 0 else dim(shapes.B, ranks.B - 2)))))"
508
  }
509
  }
@@ -512,27 +605,25 @@
512
  {
513
  "id": "transbatch_b_tiled_reg",
514
  "priority": 5,
515
- "when": ["f16Ok(dtypes.T)", "ranks.A == 3 and ranks.B == 3 and ranks.Y == 3", "attrs.transA == 0 and attrs.transB == 0 and attrs.transBatchA == 0 and attrs.transBatchB != 0", "dim(shapes.A, 0) == dim(shapes.B, 1) and dim(shapes.Y, 0) == dim(shapes.A, 0)", "dim(shapes.A, 2) == dim(shapes.B, 0)", "dim(shapes.Y, 1) == dim(shapes.A, 1) and dim(shapes.Y, 2) == dim(shapes.B, 2)", "dim(shapes.A, 0) > 0", "dim(shapes.A, 1) >= 64", "dim(shapes.A, 2) >= tunables.TRANSBATCH_REG_MIN_K", "dim(shapes.B, 2) >= 64", "ceilDiv(dim(shapes.B, 2), 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.A, 1), 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dim(shapes.A, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dim(shapes.A, 0) * ceilDiv(dim(shapes.A, 1), 64) * ceilDiv(dim(shapes.B, 2), 64) >= tunables.TILED_REG_MIN_WORKGROUPS", "16 <= device.limits.maxComputeWorkgroupSizeX and 16 <= device.limits.maxComputeWorkgroupSizeY and 256 <= device.limits.maxComputeInvocationsPerWorkgroup", "(64 * 16 + 64 * 16) * dtypeBytes(dtypes.T) <= device.limits.maxComputeWorkgroupStorageSize"],
516
- "derive": { "scalar": "dtypes.T", "usesF16": "dtypes.T == \"f16\"" },
517
  "passes": [
518
  {
519
  "id": "main",
520
  "name": "FusedMatMul.TransBatchBTiledReg",
521
  "shader": "matmul-tiled-general-reg.wgsl.jinja",
522
  "derive": {
523
- "aShape": "shapes.A",
 
 
524
  "bShape": "logicalBShape",
525
- "alpha": "attrs.alpha",
526
- "aRank": "ranks.A",
527
- "bRank": "ranks.B",
528
  "transBatchA": false,
529
  "bStorageStrides": "[dim(shapes.B, 2), dim(shapes.B, 1) * dim(shapes.B, 2), 1]",
530
  "regSequentialK": "dtypes.T == \"f16\""
531
  },
532
- "bindings": ["a", "b_3", "y"],
533
  "dispatch": {
534
- "x": "ceilDiv(dim(shapes.B, 2), 64)",
535
- "y": "ceilDiv(dim(shapes.A, 1), 64)",
536
  "z": "dim(shapes.Y, 0)"
537
  }
538
  }
 
20
  "typeConstraints": { "T": ["float32", "float16"] },
21
  "tunables": {
22
  "TILED_REG_MIN_WORKGROUPS": { "default": 64 },
23
+ "PLAIN_RANK2_REG_DEEP_K_TILES": { "default": 128 },
24
  "GEMV_TARGET_BLOCKS": { "default": 512 },
25
  "SUBGROUP_MATRIX_MIN_M": { "default": 2 },
26
+ "SUBGROUP_MATRIX_MAX_N_PADDING_RATIO": { "default": 1 },
27
  "SUBGROUP_MATRIX_SPLITK_TARGET_WGS": { "default": 512 },
28
  "SUBGROUP_MATRIX_SPLITK_MIN_K": { "default": 1024 },
29
  "SUBGROUP_MATRIX_SPLITK_MAX_TILES": { "default": 128 },
 
35
  "TRANSBATCH_B_SUBGROUP_MATRIX_MIN_WORKGROUPS": { "default": 64 },
36
  "TRANSBATCH_REG_MIN_K": { "default": 128 },
37
  "BAND_PREFER_MAX_ROWS": { "default": 8 },
38
+ "BAND_PREFER_DEEP_K": { "default": 4096 },
39
+ "BROADCAST_TRANSB_MIN_WORKGROUPS": { "default": 64 },
40
+ "BROADCAST_TRANSB_MAX_PADDING_RATIO": { "default": 2 }
41
  },
42
  "derive": {
 
 
 
43
  "batchMovedAShape": "moveAxis(shapes.A, 0, -2) if attrs.transBatchA != 0 else shapes.A",
44
  "batchMovedBShape": "moveAxis(shapes.B, 0, -2) if attrs.transBatchB != 0 else shapes.B",
45
  "logicalAShape": "moveAxis(batchMovedAShape, -1, -2) if attrs.transA != 0 and ranks.A > 1 else batchMovedAShape",
46
  "logicalBShape": "moveAxis(batchMovedBShape, -1, -2) if attrs.transB != 0 and ranks.B > 1 else batchMovedBShape",
47
  "transBatchContract": "(attrs.transBatchA == 0 and attrs.transBatchB == 0) or (ranks.A == ranks.B and ranks.A >= 3)",
48
+ "gemvN": "dim(shapes.B, 1)",
49
  "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32",
50
  "canPinSubgroupSize32": "device.features.has(\"subgroups\") and device.features.has(\"subgroup-size-control\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize <= 32 and device.adapterInfo.subgroupMaxSize >= 32",
51
  "pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter",
52
  "wave32Effective": "wave32Adapter or pinSubgroupSize32",
53
+ "variableSubgroup16To32": "device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 16 and device.adapterInfo.subgroupMaxSize == 32",
54
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
55
+ "rank2DeepPortableTier": "variableSubgroup16To32 or (has(device.adapterInfo, \"architecture\") and (device.adapterInfo.architecture == \"pascal\" or (not device.features.has(\"subgroups\") and (device.adapterInfo.architecture == \"apple\" or device.adapterInfo.architecture == \"gen-9\"))))",
56
+ "broadcastTransbM": "dim(shapes.A, ranks.A - 2)",
57
+ "broadcastTransbN": "dim(shapes.B, ranks.B - 2)",
58
+ "broadcastTransbK": "dim(shapes.A, ranks.A - 1)",
59
+ "broadcastTransbBatches": "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 2))",
60
+ "gemvLanes": 32,
61
+ "vec4OutputTile": "4 * gemvLanes",
62
+ "gemvWorkgroups": "ceilDiv(gemvN, vec4OutputTile)",
63
+ "gemvSliceCap": "min(32, device.limits.maxComputeWorkgroupSizeY, floor(device.limits.maxComputeInvocationsPerWorkgroup / gemvLanes), floor(device.limits.maxComputeWorkgroupStorageSize / (16 * gemvLanes)))",
64
+ "gemvSlices": "max(1, min(gemvSliceCap, max(8, pow2ceil(ceilDiv(tunables.GEMV_TARGET_BLOCKS, gemvWorkgroups)))))",
65
+ "gemvResourcesFit": "gemvLanes <= device.limits.maxComputeWorkgroupSizeX and gemvSlices <= device.limits.maxComputeWorkgroupSizeY and gemvLanes * gemvSlices <= device.limits.maxComputeInvocationsPerWorkgroup and 16 * gemvLanes * gemvSlices <= device.limits.maxComputeWorkgroupStorageSize",
66
+ "registerTile": 64,
67
+ "generalTile": 32,
68
+ "tiledRegResourcesFit": "registerTile / 4 <= device.limits.maxComputeWorkgroupSizeX and registerTile / 4 <= device.limits.maxComputeWorkgroupSizeY and registerTile * registerTile / 16 <= device.limits.maxComputeInvocationsPerWorkgroup and 32 * registerTile * dtypeBytes(dtypes.T) <= device.limits.maxComputeWorkgroupStorageSize",
69
+ "plainRank2RegDeepPreferredTier": "rank2DeepPortableTier and dim(shapes.A, 1) >= tunables.PLAIN_RANK2_REG_DEEP_K_TILES * 16 and dim(shapes.A, 1) % 16 == 0",
70
  "subgroupMatrixResourcesFit": "128 <= deviceWorkgroupCap and ((32 * 32 + 64 * 32) * dtypeBytes(dtypes.T) + 4 * 4 * 64 * 4) <= device.limits.maxComputeWorkgroupStorageSize",
71
  "sgmatSplitKDepth": "dim(shapes.A, ranks.A - 1)",
 
 
72
  "sgmatSplitK32Ok": "sgmatSplitKDepth % 1024 == 0",
73
  "sgmatSplitK16Ok": "sgmatSplitKDepth % 512 == 0",
74
  "sgmatSplitK8Ok": "sgmatSplitKDepth % 256 == 0",
75
  "sgmatSplitK4Ok": "sgmatSplitKDepth % 128 == 0",
76
  "sgmatSplitK2Ok": "sgmatSplitKDepth % 64 == 0",
77
+ "bandSplitWant": "pow2ceil(ceilDiv(tunables.BAND_SPLIT_TARGET_WORKGROUPS, max(1, gemvWorkgroups)))",
78
+ "bandSplitK": "16 if (bandSplitWant >= 16 and dim(shapes.A, ranks.A - 1) >= 4096) else (8 if (bandSplitWant >= 8 and dim(shapes.A, ranks.A - 1) >= 2048) else (4 if (bandSplitWant >= 4 and dim(shapes.A, ranks.A - 1) >= 1024) else (2 if (bandSplitWant >= 2 and dim(shapes.A, ranks.A - 1) >= 512) else 1)))",
79
+ "scalar": "dtypes.T",
80
+ "alpha": "attrs.alpha",
81
+ "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
82
+ "fScalar": "\"f16\" if dtypes.T == \"f16\" else \"f32\"",
83
+ "aRank": "ranks.A",
84
+ "bRank": "ranks.B",
85
+ "fusedSgmatRank2Ok": "ranks.A == 2 and ranks.B == 2 and ranks.Y == 2 and attrs.transA == 0 and attrs.transB == 0 and attrs.transBatchA == 0 and attrs.transBatchB == 0 and dim(shapes.A, 1) == dim(shapes.B, 0) and dim(shapes.Y, 0) == dim(shapes.A, 0) and dim(shapes.Y, 1) == dim(shapes.B, 1)",
86
+ "fusedSgmatN": "dim(shapes.B, ranks.B - 2) if (attrs.transB != 0 and ranks.B >= 2) else dim(shapes.B, ranks.B - 1)",
87
+ "fusedSgmatNTail": "fusedSgmatN % 64 != 0",
88
+ "fusedSgmatNPaddingHigh": "ceilDiv(fusedSgmatN, 64) * 64 > tunables.SUBGROUP_MATRIX_MAX_N_PADDING_RATIO * fusedSgmatN",
89
+ "sgmatOutTiles": "ceilDiv(dim(shapes.A, 0), 32) * ceilDiv(dim(shapes.B, 1), 64) if fusedSgmatRank2Ok else 1",
90
+ "sgmatSplitKWant": "ceilDiv(tunables.SUBGROUP_MATRIX_SPLITK_TARGET_WGS, sgmatOutTiles)",
91
  "sgmatSplitK": "32 if (sgmatSplitKWant > 16 and sgmatSplitK32Ok) else (16 if (sgmatSplitKWant > 8 and sgmatSplitK16Ok) else (8 if (sgmatSplitKWant > 4 and sgmatSplitK8Ok) else (4 if (sgmatSplitKWant > 2 and sgmatSplitK4Ok) else (2 if sgmatSplitK2Ok else 1))))",
92
  "bandRank2Ok": "ranks.A == 2 and ranks.B == 2 and ranks.Y == 2 and attrs.transA == 0 and attrs.transB == 0 and attrs.transBatchA == 0 and attrs.transBatchB == 0 and dim(shapes.A, 1) == dim(shapes.B, 0) and dim(shapes.Y, 0) == dim(shapes.A, 0) and dim(shapes.Y, 1) == dim(shapes.B, 1)",
93
+ "broadcastTransbContract": "(f16Ok(dtypes.T)) and (attrs.transA == 0 and attrs.transB != 0 and attrs.transBatchA == 0 and attrs.transBatchB == 0) and (ranks.A > ranks.B and ranks.B >= 2) and (ranks.Y == ranks.A) and (sameShape(shapes.Y, matmulShape(logicalAShape, logicalBShape))) and (dim(shapes.A, ranks.A - 1) == dim(shapes.B, ranks.B - 1)) and (dim(shapes.A, ranks.A - 2) >= 64) and (dim(shapes.A, ranks.A - 1) >= 32) and (dim(shapes.B, ranks.B - 2) >= 64) and (numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 2)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535))"
 
94
  },
95
  "bindings": {
96
+ "a": { "arg": "A", "elementType": "$scalar" },
97
+ "b": { "arg": "B", "elementType": "$vectorScalar" },
98
  "partials": { "buffer": "read-only-storage", "elementType": "f32" },
99
+ "y": { "arg": "Y", "elementType": "$scalar" },
100
+ "params": { "struct": [{ "name": "cols", "type": "u32", "value": "numel(shapes.Y)" }] },
101
+ "b_scalar": { "arg": "B", "name": "b", "elementType": "$scalar" },
102
+ "params_rows": { "name": "params", "struct": [{ "name": "M", "type": "u32", "value": "rowCount" }] }
103
  },
104
  "variants": [
105
+ {
106
+ "id": "broadcast_transb_tiled_reg",
107
+ "priority": 6,
108
+ "when": ["broadcastTransbContract", "tiledRegResourcesFit", "ceilDiv(dim(shapes.A, ranks.A - 2), registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.B, ranks.B - 2), registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
109
+ "derive": {
110
+ "bShape": "logicalBShape",
111
+ "bTransposed": true,
112
+ "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else broadcastTransbM",
113
+ "N": "broadcastTransbN",
114
+ "batchCount": "broadcastTransbBatches if ranks.B > 2 else 1",
115
+ "transBatchA": false,
116
+ "regSequentialK": "dtypes.T == \"f16\""
117
+ },
118
+ "passes": [
119
+ {
120
+ "id": "main",
121
+ "name": "FusedMatMul.BroadcastTransBTiledReg",
122
+ "shader": "matmul-tiled-general-reg.wgsl.jinja",
123
+ "derive": {
124
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []",
125
+ "aRank": "ranks.A if ranks.B > 2 else 2",
126
+ "K": "dim(shapes.A, ranks.A - 1)"
127
+ },
128
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
129
+ "dispatch": { "x": "ceilDiv(N, registerTile)", "y": "ceilDiv(rowCount, registerTile)", "z": "batchCount" }
130
+ }
131
+ ],
132
+ "demoteWhen": ["broadcastTransbBatches * ceilDiv(broadcastTransbM, registerTile) * ceilDiv(broadcastTransbN, registerTile) < tunables.BROADCAST_TRANSB_MIN_WORKGROUPS", "ceilDiv(broadcastTransbM, registerTile) * registerTile * ceilDiv(broadcastTransbN, registerTile) * registerTile * ceilDiv(broadcastTransbK,16) * 16 > tunables.BROADCAST_TRANSB_MAX_PADDING_RATIO * broadcastTransbM * broadcastTransbN * broadcastTransbK"]
133
+ },
134
+ {
135
+ "id": "broadcast_transb_subgroup_matrix_f16",
136
+ "priority": 11,
137
+ "when": ["broadcastTransbContract", "dtypes.T == \"f16\"", "wave32Effective", "subgroupMatrixResourcesFit", "ceilDiv(dim(shapes.A, ranks.A - 2),32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.B, ranks.B - 2),64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
138
+ "requires": {
139
+ "features": ["subgroups", "chromium-experimental-subgroup-matrix"],
140
+ "subgroupMatrixConfigs": [{ "componentType": "f16", "M": 8, "N": 8, "K": 8 }]
141
+ },
142
+ "derive": {
143
+ "bShape": "logicalBShape",
144
+ "bTransposed": true,
145
+ "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else broadcastTransbM",
146
+ "K": "broadcastTransbK",
147
+ "N": "broadcastTransbN",
148
+ "batchCount": "broadcastTransbBatches if ranks.B > 2 else 1",
149
+ "hasBias": false,
150
+ "fScalar": "dtypes.T",
151
+ "outScalar": "dtypes.T",
152
+ "generalAddressing": true,
153
+ "tailSafe": "dim(shapes.A, ranks.A - 1) % 32 != 0 or dim(shapes.B, ranks.B - 2) % 64 != 0",
154
+ "outputBuffer": "\"y\""
155
+ },
156
+ "passes": [
157
+ {
158
+ "id": "main",
159
+ "name": "FusedMatMul.BroadcastTransBSubgroupMatrix",
160
+ "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
161
+ "derive": {
162
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []",
163
+ "aRank": "ranks.A if ranks.B > 2 else 2"
164
+ },
165
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
166
+ "dispatch": { "x": "ceilDiv(N,64)", "y": "ceilDiv(rowCount,32)", "z": "batchCount" }
167
+ }
168
+ ],
169
+ "demoteWhen": ["broadcastTransbBatches * ceilDiv(broadcastTransbM,32) * ceilDiv(broadcastTransbN,64) < tunables.BROADCAST_TRANSB_MIN_WORKGROUPS", "ceilDiv(broadcastTransbM,32) * 32 * ceilDiv(broadcastTransbN,64) * 64 * ceilDiv(broadcastTransbK,32) * 32 > tunables.BROADCAST_TRANSB_MAX_PADDING_RATIO * broadcastTransbM * broadcastTransbN * broadcastTransbK"]
170
+ },
171
+ {
172
+ "id": "broadcast_transb_subgroup_matrix_f32",
173
+ "priority": 11,
174
+ "when": ["broadcastTransbContract", "dtypes.T == \"f32\"", "wave32Effective", "subgroupMatrixResourcesFit", "ceilDiv(dim(shapes.A, ranks.A - 2),32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.B, ranks.B - 2),64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
175
+ "requires": {
176
+ "features": ["subgroups", "chromium-experimental-subgroup-matrix"],
177
+ "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
178
+ },
179
+ "derive": {
180
+ "bShape": "logicalBShape",
181
+ "bTransposed": true,
182
+ "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else broadcastTransbM",
183
+ "K": "broadcastTransbK",
184
+ "N": "broadcastTransbN",
185
+ "batchCount": "broadcastTransbBatches if ranks.B > 2 else 1",
186
+ "hasBias": false,
187
+ "fScalar": "dtypes.T",
188
+ "outScalar": "dtypes.T",
189
+ "generalAddressing": true,
190
+ "tailSafe": "dim(shapes.A, ranks.A - 1) % 32 != 0 or dim(shapes.B, ranks.B - 2) % 64 != 0",
191
+ "outputBuffer": "\"y\""
192
+ },
193
+ "passes": [
194
+ {
195
+ "id": "main",
196
+ "name": "FusedMatMul.BroadcastTransBSubgroupMatrix",
197
+ "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
198
+ "derive": {
199
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []",
200
+ "aRank": "ranks.A if ranks.B > 2 else 2"
201
+ },
202
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
203
+ "dispatch": { "x": "ceilDiv(N,64)", "y": "ceilDiv(rowCount,32)", "z": "batchCount" }
204
+ }
205
+ ],
206
+ "demoteWhen": ["broadcastTransbBatches * ceilDiv(broadcastTransbM,32) * ceilDiv(broadcastTransbN,64) < tunables.BROADCAST_TRANSB_MIN_WORKGROUPS", "ceilDiv(broadcastTransbM,32) * 32 * ceilDiv(broadcastTransbN,64) * 64 * ceilDiv(broadcastTransbK,32) * 32 > tunables.BROADCAST_TRANSB_MAX_PADDING_RATIO * broadcastTransbM * broadcastTransbN * broadcastTransbK"]
207
+ },
208
  {
209
  "id": "subgroup_matrix_transbatch_b_f16",
210
  "priority": 11,
 
215
  },
216
  "derive": {
217
  "hasBias": false,
 
218
  "fScalar": "dtypes.T",
219
  "outScalar": "dtypes.T",
 
220
  "generalAddressing": true,
221
  "outputBuffer": "\"y\"",
222
+ "rowCount": "dim(shapes.A, ranks.A - 2)",
 
223
  "K": "dim(shapes.A, ranks.A - 1)",
224
  "N": "dim(shapes.B, ranks.B - 1)",
225
  "batchCount": "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 1))"
 
230
  "name": "FusedMatMul.SubgroupMatrixTransBatchB",
231
  "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
232
  "derive": {
233
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2)",
234
  "bShape": "logicalBShape",
 
 
235
  "bStorageStrides": "[dim(shapes.B, 2), dim(shapes.B, 1) * dim(shapes.B, 2), 1]"
236
  },
237
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
238
+ "dispatch": { "x": "ceil(N / 64)", "y": "ceil(rowCount / 32)", "z": "batchCount" }
239
  }
240
  ]
241
  },
 
249
  },
250
  "derive": {
251
  "hasBias": false,
 
252
  "fScalar": "dtypes.T",
253
  "outScalar": "dtypes.T",
 
254
  "generalAddressing": true,
255
  "outputBuffer": "\"y\"",
256
+ "rowCount": "dim(shapes.A, ranks.A - 2)",
 
257
  "K": "dim(shapes.A, ranks.A - 1)",
258
  "N": "dim(shapes.B, ranks.B - 1)",
259
  "batchCount": "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 1))"
 
264
  "name": "FusedMatMul.SubgroupMatrixTransBatchB",
265
  "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
266
  "derive": {
267
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2)",
268
  "bShape": "logicalBShape",
 
 
269
  "bStorageStrides": "[dim(shapes.B, 2), dim(shapes.B, 1) * dim(shapes.B, 2), 1]"
270
  },
271
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
272
+ "dispatch": { "x": "ceil(N / 64)", "y": "ceil(rowCount / 32)", "z": "batchCount" }
273
  }
274
  ]
275
  },
276
  {
277
+ "id": "m1_gemv_vec4",
278
  "priority": 30,
279
+ "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\")", "attrs.transA == 0", "attrs.transB == 0", "attrs.transBatchA == 0", "attrs.transBatchB == 0", "ranks.A == 2", "ranks.B == 2", "ranks.Y == 2", "dim(shapes.A, 0) == 1", "dim(shapes.Y, 0) == 1", "dim(shapes.A, 1) == dim(shapes.B, 0)", "dim(shapes.Y, 1) == dim(shapes.B, 1)", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "gemvWorkgroups <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "f16Ok(dtypes.T)"],
280
+ "derive": {
281
+ "unrollK2": "dtypes.T == \"f16\"",
282
+ "gemvScalar": "dtypes.T",
283
+ "gemvVector": "\"vec4<\" ~ dtypes.T ~ \">\"",
284
+ "alphaScale": "attrs.alpha"
285
+ },
286
  "passes": [
287
  {
288
  "id": "main",
289
+ "name": "FusedMatMul.M1GemvVec4",
290
  "shader": "matmul-vector-matrix-vec4.wgsl.jinja",
291
  "bindings": [
292
+ { "arg": "A", "name": "a", "elementType": "$gemvScalar" },
293
+ { "arg": "B", "name": "b", "elementType": "$gemvVector" },
294
+ { "arg": "Y", "name": "c", "elementType": "$gemvVector" },
295
  {
296
  "name": "params",
297
  "struct": [
 
300
  ]
301
  }
302
  ],
303
+ "dispatch": { "x": "gemvWorkgroups" }
304
  }
305
  ]
306
  },
307
  {
308
  "id": "rank2_band_vec4_splitk",
309
  "priority": 11,
310
+ "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T)", "bandRank2Ok", "dim(shapes.A, 0) >= 2", "dim(shapes.A, 0) <= tunables.BAND_VEC4_MAX_ROWS", "dim(shapes.A, 1) > 0", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "gemvWorkgroups <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "gemvLanes <= device.limits.maxComputeWorkgroupSizeX", "gemvWorkgroups <= tunables.BAND_SPLIT_MAX_COLUMN_GROUPS", "bandSplitK >= 2", "bandSplitK * numel(shapes.Y) * 4 <= device.limits.maxStorageBufferBindingSize", "bandSplitK <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tunables.BAND_SPLIT_SLICES <= device.limits.maxComputeWorkgroupSizeY", "gemvLanes * tunables.BAND_SPLIT_SLICES <= device.limits.maxComputeInvocationsPerWorkgroup"],
311
  "derive": {
 
 
 
312
  "batched": false,
313
  "outputBuffer": "\"y\"",
 
314
  "M": "dim(shapes.A, 0)",
315
  "K": "dim(shapes.A, 1)",
316
  "N": "dim(shapes.B, 1)",
 
332
  "id": "combine",
333
  "name": "FusedMatMul.Rank2BandVec4SplitKCombine",
334
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
335
+ "derive": { "op": "\"sum\"", "outputF16": "dtypes.T == \"f16\"" },
336
  "bindings": ["partials", "y", "params"],
337
  "dispatch": {
338
  "x": "min(ceilDiv((numel(shapes.Y)), (256)), 65535)",
 
345
  {
346
  "id": "rank2_band_vec4",
347
  "priority": 11,
348
+ "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T)", "bandRank2Ok", "dim(shapes.A, 0) >= 2", "dim(shapes.A, 0) <= tunables.BAND_VEC4_MAX_ROWS", "dim(shapes.A, 1) > 0", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "gemvWorkgroups <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "gemvResourcesFit", "not (gemvWorkgroups <= tunables.BAND_SPLIT_MAX_COLUMN_GROUPS and bandSplitK >= 2)"],
349
  "derive": {
 
 
 
350
  "batched": false,
351
  "outputBuffer": "\"y\"",
 
352
  "M": "dim(shapes.A, 0)",
353
  "K": "dim(shapes.A, 1)",
354
+ "N": "dim(shapes.B, 1)"
 
355
  },
356
  "passes": [
357
  {
 
361
  "bindings": ["a", "b", { "arg": "Y", "name": "y", "elementType": "$vectorScalar" }],
362
  "dispatch": { "x": "gemvWorkgroups" }
363
  }
364
+ ]
 
365
  },
366
  {
367
  "id": "rank2_band_vec4_f32_preferred",
368
  "priority": 13,
369
+ "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T)", "bandRank2Ok", "dim(shapes.A, 0) >= 2", "dim(shapes.A, 0) <= tunables.BAND_VEC4_MAX_ROWS", "dim(shapes.A, 1) > 0", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "gemvWorkgroups <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "gemvResourcesFit", "not (gemvWorkgroups <= tunables.BAND_SPLIT_MAX_COLUMN_GROUPS and bandSplitK >= 2)"],
370
+ "demoteWhen": ["dtypes.T != \"f32\" or (dim(shapes.A, 0) > tunables.BAND_PREFER_MAX_ROWS and dim(shapes.A, 1) >= tunables.BAND_PREFER_DEEP_K)"],
371
  "derive": {
 
 
 
372
  "batched": false,
373
  "outputBuffer": "\"y\"",
 
374
  "M": "dim(shapes.A, 0)",
375
  "K": "dim(shapes.A, 1)",
376
+ "N": "dim(shapes.B, 1)"
 
377
  },
378
  "passes": [
379
  {
 
383
  "bindings": ["a", "b", { "arg": "Y", "name": "y", "elementType": "$vectorScalar" }],
384
  "dispatch": { "x": "gemvWorkgroups" }
385
  }
386
+ ]
 
387
  },
388
  {
389
  "id": "subgroup_matrix_splitk",
 
397
  ]
398
  },
399
  "derive": {
 
 
 
400
  "hasBias": false,
401
  "generalAddressing": true,
402
  "tailSafe": false,
403
  "outputBuffer": "\"partials\"",
404
  "outScalar": "\"f32\"",
405
+ "rowCount": "dim(shapes.A, 0)",
 
406
  "K": "dim(shapes.A, 1)",
407
  "N": "dim(shapes.B, 1)",
408
  "batchCount": 1,
 
417
  "id": "partial",
418
  "name": "FusedMatMul.SubgroupMatrixSplitK",
419
  "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
420
+ "derive": { "bShape": ["dim(shapes.B, 0)", "dim(shapes.B, 1)"], "aRank": 2, "bRank": 2 },
421
+ "bindings": ["a", "b_scalar", { "name": "partials", "elementType": "f32" }, "params_rows"],
 
 
 
 
 
422
  "dispatch": { "x": "ceilDiv(dim(shapes.Y, 1), 64)", "y": "ceilDiv(dim(shapes.Y, 0), 32)", "z": "sgmatSplitK" }
423
  },
424
  {
425
  "id": "combine",
426
  "name": "FusedMatMul.SubgroupMatrixSplitKCombine",
427
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
428
+ "derive": { "op": "\"sum\"", "outputF16": "dtypes.T == \"f16\"" },
429
  "bindings": ["partials", "y", "params"],
430
  "dispatch": {
431
  "x": "min(ceilDiv((numel(shapes.Y)), (256)), 65535)",
 
445
  },
446
  "derive": {
447
  "hasBias": false,
 
448
  "fScalar": "\"f16\"",
449
  "outScalar": "\"f16\"",
 
450
  "generalAddressing": true,
451
  "tailSafe": "dim(shapes.A, ranks.A - 1) % 32 != 0 or dim(shapes.B, ranks.B - 1) % 64 != 0",
452
  "outputBuffer": "\"y\"",
453
+ "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else dim(shapes.A, ranks.A - 2)",
 
454
  "K": "dim(shapes.A, ranks.A - 1)",
455
  "N": "dim(shapes.B, ranks.B - 1)",
456
+ "batchCount": "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 1)) if ranks.B > 2 else 1"
457
  },
458
  "passes": [
459
  {
460
  "id": "main",
461
  "name": "FusedMatMul.SubgroupMatrixTailBroadcast",
462
  "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
463
+ "derive": {
464
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []",
465
+ "aRank": "ranks.A if ranks.B > 2 else 2",
466
+ "bShape": "shapes.B"
467
+ },
468
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
469
+ "dispatch": { "x": "ceil(N / 64)", "y": "ceil(rowCount / 32)", "z": "batchCount" }
470
  }
471
  ]
472
  },
473
  {
474
  "id": "subgroup_matrix",
475
  "priority": 10,
476
+ "when": ["f16Ok(dtypes.T)", "attrs.transBatchA == 0 or (attrs.transA == 0 and ranks.A == 3)", "attrs.transBatchB == 0", "ranks.A >= 2", "ranks.B == ranks.A", "ranks.Y == ranks.A", "(dim(shapes.A, ranks.A - 2) if attrs.transA != 0 else dim(shapes.A, ranks.A - 1)) == (dim(shapes.B, ranks.B - 1) if attrs.transB != 0 else dim(shapes.B, ranks.B - 2))", "(dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2))) >= tunables.SUBGROUP_MATRIX_MIN_M", "(dim(shapes.A, ranks.A - 2) if attrs.transA != 0 else dim(shapes.A, ranks.A - 1)) % 32 == 0", "(ranks.A == 2 or (ranks.A == 3 and (dim(shapes.A, 1) if attrs.transBatchA != 0 else dim(shapes.A, 0)) == dim(shapes.B, 0) and dim(shapes.Y, 0) == dim(shapes.B, 0)) or (ranks.A == 4 and dim(shapes.A, 0) == dim(shapes.B, 0) and dim(shapes.A, 1) == dim(shapes.B, 1) and dim(shapes.Y, 0) == dim(shapes.A, 0) and dim(shapes.Y, 1) == dim(shapes.A, 1)))", "dim(shapes.Y, ranks.Y - 2) == (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))", "dim(shapes.Y, ranks.Y - 1) == (dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1))", "ceil((dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1)) / 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil((dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2))) / 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "numel(shapes.Y) / ((dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2))) * (dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1))) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "wave32Effective"],
477
  "requires": {
478
  "features": ["subgroups", "chromium-experimental-subgroup-matrix"],
479
  "subgroupMatrixConfigs": [
 
482
  ]
483
  },
484
  "derive": {
 
 
485
  "outScalar": "\"f16\" if dtypes.T == \"f16\" else \"f32\"",
486
+ "nTailSafe": "fusedSgmatNTail",
487
  "transA": "attrs.transA != 0",
488
  "transB": "attrs.transB != 0",
489
  "transBatchA": "attrs.transBatchA != 0",
490
+ "rowCount": "(dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))",
491
  "K": "(dim(shapes.A, ranks.A - 2) if attrs.transA != 0 else dim(shapes.A, ranks.A - 1))",
492
  "N": "(dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1))",
493
  "batchCount": "numel(shapes.Y) / ((dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2))) * (dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1)))"
 
497
  "id": "main",
498
  "name": "FusedMatMul.SubgroupMatrix",
499
  "shader": "fused-matmul-subgroup-matrix.wgsl.jinja",
500
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
501
+ "dispatch": { "x": "ceil(N / 64)", "y": "ceil(rowCount / 32)", "z": "numel(shapes.Y) / (rowCount * N)" }
 
502
  }
503
+ ],
504
+ "demoteWhen": ["fusedSgmatNPaddingHigh"]
505
  },
506
  {
507
  "id": "broadcast_rank4_tiled_reg",
508
  "priority": 6,
509
+ "when": ["f16Ok(dtypes.T)", "attrs.transA == 0", "attrs.transB == 0", "attrs.transBatchA == 0", "attrs.transBatchB == 0", "ranks.A == 4", "(ranks.B == 2 or ranks.B == 3)", "ranks.Y == 4", "dim(shapes.Y, 0) == dim(shapes.A, 0)", "(ranks.B == 2 or dim(shapes.A, 1) == dim(shapes.B, 0) or dim(shapes.A, 1) == 1 or dim(shapes.B, 0) == 1)", "dim(shapes.Y, 1) == (dim(shapes.A, 1) if ranks.B == 2 else max(dim(shapes.A, 1), dim(shapes.B, 0)))", "dim(shapes.A, 3) == dim(shapes.B, ranks.B - 2)", "dim(shapes.Y, 2) == dim(shapes.A, 2)", "dim(shapes.Y, 3) == dim(shapes.B, ranks.B - 1)", "dim(shapes.A, 2) >= 64", "dim(shapes.A, 3) >= 32", "dim(shapes.B, ranks.B - 1) >= 64", "ceil(dim(shapes.B, ranks.B - 1) / registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil(dim(shapes.A, 2) / registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "numel(shapes.Y) / (dim(shapes.A, 2) * dim(shapes.B, ranks.B - 1)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
 
510
  "passes": [
511
  {
512
  "id": "main",
513
  "name": "FusedMatMul.BroadcastRank4TiledReg",
514
  "shader": "matmul-tiled-general-reg.wgsl.jinja",
515
  "derive": {
516
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []",
517
+ "aRank": "ranks.A if ranks.B > 2 else 2",
518
+ "K": "dim(shapes.A, ranks.A - 1)",
519
  "bShape": "shapes.B",
 
 
 
520
  "transBatchA": "false"
521
  },
522
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
523
  "dispatch": {
524
+ "x": "ceil(dim(shapes.B, ranks.B - 1) / registerTile)",
525
+ "y": "ceil(rowCount / registerTile)",
526
+ "z": "batchCount"
527
  }
528
  }
529
+ ],
530
+ "derive": {
531
+ "rowCount": "outer(shapes.A, 3) if ranks.B == 2 else dim(shapes.A, 2)",
532
+ "batchCount": "numel(shapes.Y) / (dim(shapes.A, 2) * dim(shapes.B, ranks.B - 1)) if ranks.B > 2 else 1"
533
+ }
534
  },
535
  {
536
  "id": "plain_rank2_tiled_reg",
537
  "priority": 4,
538
+ "demoteWhen": ["plainRank2RegDeepPreferredTier"],
539
+ "when": ["f16Ok(dtypes.T)", "fusedSgmatRank2Ok", "dim(shapes.A, 0) >= 64", "dim(shapes.A, 1) >= 32", "dim(shapes.B, 1) >= 64", "ceil(dim(shapes.A, 0) / registerTile) * ceil(dim(shapes.B, 1) / registerTile) >= tunables.TILED_REG_MIN_WORKGROUPS", "ceil(dim(shapes.B, 1) / registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil(dim(shapes.A, 0) / registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
540
  "passes": [
541
  {
542
  "id": "main",
543
  "name": "FusedMatMul.PlainRank2TiledReg",
544
  "shader": "matmul-tiled-general-reg.wgsl.jinja",
545
  "derive": {
546
+ "rowCount": "dim(shapes.A, 0)",
547
+ "K": "dim(shapes.A, 1)",
548
  "bShape": "shapes.B",
 
 
 
549
  "transBatchA": "false"
550
  },
551
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
552
+ "dispatch": {
553
+ "x": "ceil(dim(shapes.B, 1) / registerTile)",
554
+ "y": "ceil(dim(shapes.A, 0) / registerTile)",
555
+ "z": 1
556
+ }
557
  }
558
  ]
559
  },
560
  {
561
  "id": "transbatch_a_tiled_reg",
562
  "priority": 5,
563
+ "when": ["f16Ok(dtypes.T)", "attrs.transBatchA != 0", "attrs.transBatchB == 0", "attrs.transA == 0", "attrs.transB == 0", "ranks.A == 3", "ranks.B == 3", "ranks.Y == 3", "dim(shapes.A, 1) == dim(shapes.B, 0)", "dim(shapes.Y, 0) == dim(shapes.B, 0)", "dim(shapes.A, 2) == dim(shapes.B, 1)", "dim(shapes.Y, 1) == dim(shapes.A, 0)", "dim(shapes.Y, 2) == dim(shapes.B, 2)", "dim(shapes.A, 0) >= 64", "dim(shapes.A, 2) >= 32", "dim(shapes.B, 2) >= 64", "ceil(dim(shapes.B, 2) / registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil(dim(shapes.A, 0) / registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dim(shapes.Y, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
 
564
  "passes": [
565
  {
566
  "id": "main",
567
  "name": "FusedMatMul.TransBatchATiledReg",
568
  "shader": "matmul-tiled-general-reg.wgsl.jinja",
569
+ "derive": { "aShape": "shapes.A", "bShape": "shapes.B", "transBatchA": "true", "kTile": "4" },
570
+ "bindings": ["a", "b_scalar", "y"],
 
 
 
 
 
 
 
 
571
  "dispatch": {
572
+ "x": "ceil(dim(shapes.B, 2) / registerTile)",
573
+ "y": "ceil(dim(shapes.A, 0) / registerTile)",
574
  "z": "dim(shapes.Y, 0)"
575
  }
576
  }
 
580
  "id": "tiled",
581
  "priority": 0,
582
  "when": ["ranks.A >= 1", "ranks.B >= 1", "f16Ok(dtypes.T)", "transBatchContract", "(dim(shapes.A, 0) if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA == 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))) == (dim(shapes.B, 0) if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB != 0 else (dim(shapes.B, 0) if attrs.transBatchB != 0 else dim(shapes.B, ranks.B - 2))))", "ceil((1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))) / 16) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil((1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else (dim(shapes.B, 0) if attrs.transBatchB != 0 else dim(shapes.B, ranks.B - 2)))) / 16) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "numel(shapes.Y) / max(1, (1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))) * (1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else (dim(shapes.B, 0) if attrs.transBatchB != 0 else dim(shapes.B, ranks.B - 2))))) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
 
583
  "passes": [
584
  {
585
  "id": "main",
 
590
  "bShape": "shapes.B",
591
  "transA": "attrs.transA != 0",
592
  "transB": "attrs.transB != 0",
 
 
 
593
  "transBatchA": "attrs.transBatchA != 0",
594
  "transBatchB": "attrs.transBatchB != 0"
595
  },
596
+ "bindings": ["a", "b_scalar", "y"],
597
  "dispatch": {
598
+ "x": "ceil((1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else (dim(shapes.B, 0) if attrs.transBatchB != 0 else dim(shapes.B, ranks.B - 2)))) / generalTile)",
599
+ "y": "ceil((1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))) / generalTile)",
600
  "z": "numel(shapes.Y) / max(1, (1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else (dim(shapes.A, 0) if attrs.transBatchA != 0 else dim(shapes.A, ranks.A - 2)))) * (1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else (dim(shapes.B, 0) if attrs.transBatchB != 0 else dim(shapes.B, ranks.B - 2)))))"
601
  }
602
  }
 
605
  {
606
  "id": "transbatch_b_tiled_reg",
607
  "priority": 5,
608
+ "when": ["f16Ok(dtypes.T)", "ranks.A == 3 and ranks.B == 3 and ranks.Y == 3", "attrs.transA == 0 and attrs.transB == 0 and attrs.transBatchA == 0 and attrs.transBatchB != 0", "dim(shapes.A, 0) == dim(shapes.B, 1) and dim(shapes.Y, 0) == dim(shapes.A, 0)", "dim(shapes.A, 2) == dim(shapes.B, 0)", "dim(shapes.Y, 1) == dim(shapes.A, 1) and dim(shapes.Y, 2) == dim(shapes.B, 2)", "dim(shapes.A, 0) > 0", "dim(shapes.A, 1) >= 64", "dim(shapes.A, 2) >= tunables.TRANSBATCH_REG_MIN_K", "dim(shapes.B, 2) >= 64", "ceilDiv(dim(shapes.B, 2), registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.A, 1), registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dim(shapes.A, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dim(shapes.A, 0) * ceilDiv(dim(shapes.A, 1), registerTile) * ceilDiv(dim(shapes.B, 2), registerTile) >= tunables.TILED_REG_MIN_WORKGROUPS", "tiledRegResourcesFit"],
 
609
  "passes": [
610
  {
611
  "id": "main",
612
  "name": "FusedMatMul.TransBatchBTiledReg",
613
  "shader": "matmul-tiled-general-reg.wgsl.jinja",
614
  "derive": {
615
+ "rowCount": "dim(shapes.A, 1)",
616
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2)",
617
+ "K": "dim(shapes.A, ranks.A - 1)",
618
  "bShape": "logicalBShape",
 
 
 
619
  "transBatchA": false,
620
  "bStorageStrides": "[dim(shapes.B, 2), dim(shapes.B, 1) * dim(shapes.B, 2), 1]",
621
  "regSequentialK": "dtypes.T == \"f16\""
622
  },
623
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
624
  "dispatch": {
625
+ "x": "ceilDiv(dim(shapes.B, 2), registerTile)",
626
+ "y": "ceilDiv(dim(shapes.A, 1), registerTile)",
627
  "z": "dim(shapes.Y, 0)"
628
  }
629
  }
build/webgpu/matmul-band-vec4.wgsl.jinja CHANGED
@@ -4,15 +4,11 @@
4
  //
5
  // A batched consumer runs one band per workgroup row: workgroup_id.y selects
6
  // the matrix, and every operand is offset by its per-matrix extent.
7
- {% if usesF16 %}
8
- enable f16;
9
-
10
- {% endif %}
11
  {{ env.wgsl.resourceDeclarations }}
12
 
13
  const K: u32 = {{ K }}u;
14
  const N4: u32 = {{ N }}u / 4u;
15
- const LANES: u32 = 32u;
16
  // SLICES partitions the K reduction across the workgroup's second dimension.
17
  const SLICES: u32 = {{ gemvSlices }}u;
18
  {% set kSplitsValue = kSplits if kSplits is defined else 1 %}
@@ -26,7 +22,7 @@ const K_PER_SPLIT: u32 = (K + K_SPLITS - 1u) / K_SPLITS;
26
  // footprint does not grow with the band.
27
  var<workgroup> partials: array<vec4<f32>, LANES * SLICES>;
28
 
29
- @compute @workgroup_size(32, {{ gemvSlices }}, 1)
30
  fn main(
31
  @builtin(workgroup_id) workgroup_id: vec3<u32>,
32
  @builtin(local_invocation_id) lid: vec3<u32>
 
4
  //
5
  // A batched consumer runs one band per workgroup row: workgroup_id.y selects
6
  // the matrix, and every operand is offset by its per-matrix extent.
 
 
 
 
7
  {{ env.wgsl.resourceDeclarations }}
8
 
9
  const K: u32 = {{ K }}u;
10
  const N4: u32 = {{ N }}u / 4u;
11
+ const LANES: u32 = {{ gemvLanes }}u;
12
  // SLICES partitions the K reduction across the workgroup's second dimension.
13
  const SLICES: u32 = {{ gemvSlices }}u;
14
  {% set kSplitsValue = kSplits if kSplits is defined else 1 %}
 
22
  // footprint does not grow with the band.
23
  var<workgroup> partials: array<vec4<f32>, LANES * SLICES>;
24
 
25
+ @compute @workgroup_size({{ gemvLanes }}, {{ gemvSlices }}, 1)
26
  fn main(
27
  @builtin(workgroup_id) workgroup_id: vec3<u32>,
28
  @builtin(local_invocation_id) lid: vec3<u32>
build/webgpu/matmul-subgroup-matrix-ext.wgsl.jinja CHANGED
@@ -8,28 +8,33 @@ enable subgroup_size_control;
8
  enable chromium_experimental_subgroup_matrix;
9
  diagnostic(off, chromium.subgroup_matrix_uniformity);
10
 
11
-
12
  {{ env.wgsl.resourceDeclarations }}
13
-
14
  {% set operandScalar = fScalar %}
15
  {% set accScalar = "f32" %}
16
- {% set GENERAL = generalAddressing is defined and generalAddressing %}
17
  {% set TAIL = tailSafe is defined and tailSafe %}
18
  {% set SPLIT_K = splitK if splitK is defined else 1 %}
19
  {% set OUT = outputBuffer if outputBuffer is defined else "c" %}
20
  {% set OUT_SCALAR = outScalar if outScalar is defined else T %}
 
 
 
 
 
21
  {% set aR = aRank %}
22
  {% set bR = bRank %}
23
  {% set aBatchLen = aR - 2 %}
24
  {% set bBatchLen = bR - 2 %}
25
  {% set batchRank = aBatchLen %}
26
- {% set aMStride = aShape[aR-1] %}
27
  {% set aKStride = 1 %}
28
- {% set bKStride = bShape[bR-1] %}
29
- {% set bNStride = 1 %}
30
  {% if bStorageStrides is defined %}{% set bKStride = bStorageStrides[bR-2] %}{% set bNStride = bStorageStrides[bR-1] %}{% endif %}
31
 
 
32
  const M: u32 = {{ M }}u;
 
33
  const K: u32 = {{ K }}u;
34
  const N: u32 = {{ N }}u;
35
  const BATCH_COUNT: u32 = {{ batchCount if batchCount is defined else 1 }}u;
@@ -44,7 +49,9 @@ const B_N_STRIDE: u32 = {{ bNStride }}u;
44
  {% if TAIL %}const K_FULL: u32 = (K / 32u) * 32u;
45
  {% endif %}
46
  const ALPHA: f32 = f32({{ alpha }});
 
47
  const C_BATCH_STRIDE: u32 = M * N;
 
48
  const TILE_COLS: u32 = 64u;
49
  const TILE_ROWS: u32 = 32u;
50
  const TILE_K: u32 = 32u;
@@ -57,10 +64,13 @@ var<workgroup> scratch: array<array<array<{{ accScalar }}, 64>, 4>, 4>;
57
 
58
  fn loadSHMA(a_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
59
  let a_global = tile_base + row;
 
 
 
60
  let col = c_idx * 8u;
61
  for (var col_offset = 0u; col_offset < 8u; col_offset = col_offset + 1u) {
62
  let k = k_idx + col + col_offset;
63
- if (a_global < M) {
64
  {% if operandScalar == "f16" %}
65
  tile_A[row * TILE_K + col + col_offset] = f16(a[a_base + a_global * A_M_STRIDE + k * A_K_STRIDE]);
66
  {% else %}
@@ -79,10 +89,11 @@ fn loadSHMA(a_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
79
 
80
  fn loadSHMAKTail(a_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
81
  let a_global = tile_base + row;
 
82
  let col = c_idx * 8u;
83
  for (var col_offset = 0u; col_offset < 8u; col_offset = col_offset + 1u) {
84
  let k = k_idx + col + col_offset;
85
- if (a_global < M && k < K) {
86
  tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(a[a_base + a_global * A_M_STRIDE + k * A_K_STRIDE]);
87
  } else {
88
  tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(0);
@@ -157,6 +168,11 @@ fn main(
157
  @builtin(subgroup_invocation_id) sg_id: u32,
158
  @builtin(subgroup_size) sg_size: u32
159
  ) {
 
 
 
 
 
160
  let b_global_base = workgroup_id.x * TILE_COLS;
161
 
162
  let subtile_id = local_idx / sg_size;
@@ -191,7 +207,7 @@ fn main(
191
  {% set axis = batchRank - 1 - i %}
192
  {% set aAxis = axis - (batchRank - aBatchLen) %}
193
  {% set bAxis = axis - (batchRank - bBatchLen) %}
194
- {% set aDim = aShape[aAxis] %}
195
  {% set bDim = bShape[bAxis] if bAxis >= 0 else 1 %}
196
  {% if aDim > 1 or bDim > 1 %}{% set hasBatchCoord.value = true %}{% endif %}
197
  {% endfor %}
@@ -205,18 +221,18 @@ fn main(
205
  {% set axis = batchRank - 1 - i %}
206
  {% set aAxis = axis - (batchRank - aBatchLen) %}
207
  {% set bAxis = axis - (batchRank - bBatchLen) %}
208
- {% set aDim = aShape[aAxis] %}
209
  {% set bDim = bShape[bAxis] if bAxis >= 0 else 1 %}
210
  {% set outDim = aDim if aDim >= bDim else bDim %}
211
  {% set aStride = namespace(v=1) %}
212
- {% if aDim != 1 %}{% for j in range(aAxis + 1, aR) %}{% set aStride.v = aStride.v * aShape[j] %}{% endfor %}{% else %}{% set aStride.v = 0 %}{% endif %}
213
  {% set bStride = namespace(v=1) %}
214
  {% if bAxis >= 0 and bDim != 1 %}{% for j in range(bAxis + 1, bR) %}{% set bStride.v = bStride.v * bShape[j] %}{% endfor %}{% else %}{% set bStride.v = 0 %}{% endif %}
215
  {% if bStorageStrides is defined and bAxis >= 0 and bDim != 1 %}{% set bStride.v = bStorageStrides[bAxis] %}{% endif %}
216
  {% if outDim > 1 %}
217
  let c{{ axis }} = zTmp % {{ outDim }}u;
218
  zTmp = zTmp / {{ outDim }}u;
219
- {% if aStride.v != 0 %} a_base = a_base + c{{ axis }} * {{ aStride.v }}u;
220
  {% endif %}
221
  {% if bStride.v != 0 %} b_base = b_base + c{{ axis }} * {{ bStride.v }}u;
222
  {% endif %}
@@ -249,7 +265,7 @@ fn main(
249
  workgroupBarrier();
250
 
251
  for (var step = 0u; step < TILE_K; step = step + 8u) {
252
- {% set directInputs = directMatrixInputs is defined and directMatrixInputs %}
253
  let matrix_a_offset = subtile_idy * SUB_ROWS * TILE_K + step;
254
  {% for r in range(2) %}
255
  var matA{{ r }}: subgroup_matrix_left<{{ operandScalar }}, 8, 8> = subgroupMatrixLoad<subgroup_matrix_left<{{ operandScalar }}, 8, 8>, row_major>(&{{ "w" if directInputs else "tile_A" }}, matrix_a_offset{% if r > 0 %} + 8u * {{ "K" if directInputs else "TILE_K" }}{% endif %}, {{ "K" if directInputs else "TILE_K" }});
@@ -278,15 +294,15 @@ fn main(
278
  workgroupBarrier();
279
 
280
  for (var step = 0u; step < TILE_K; step = step + 8u) {
281
- {% set directInputs = directMatrixInputs is defined and directMatrixInputs %}
282
- let matrix_a_offset = {% if directInputs %}(a_global_base + subtile_idy * SUB_ROWS) * K + kidx + step{% else %}subtile_idy * SUB_ROWS * TILE_K + step{% endif %};
283
  {% for r in range(2) %}
284
  var matA{{ r }}: subgroup_matrix_left<{{ operandScalar }}, 8, 8> = subgroupMatrixLoad<subgroup_matrix_left<{{ operandScalar }}, 8, 8>, row_major>(&{{ "w" if directInputs else "tile_A" }}, matrix_a_offset{% if r > 0 %} + 8u * {{ "K" if directInputs else "TILE_K" }}{% endif %}, {{ "K" if directInputs else "TILE_K" }});
285
  {% endfor %}
286
 
287
- let matrix_b_offset = {% if directInputs %}b_base + (kidx + step) * N + b_global_base + subtile_idx * SUB_COLS{% else %}subtile_idx * SUB_COLS * TILE_K + step{% endif %};
288
  {% for c in range(4) %}
289
- var matB{{ c }}: subgroup_matrix_right<{{ operandScalar }}, 8, 8> = subgroupMatrixLoad<subgroup_matrix_right<{{ operandScalar }}, 8, 8>, {% if directInputs %}row_major{% else %}col_major{% endif %}>(&{{ "xm" if directInputs else "tile_B" }}, matrix_b_offset{% if c > 0 %} + {{ c * 8 }}u{% if not directInputs %} * TILE_K{% endif %}{% endif %}, {{ "N" if directInputs else "TILE_K" }});
290
  {% endfor %}
291
 
292
  matC00 = subgroupMatrixMultiplyAccumulate(matA0, matB0, matC00);
@@ -333,8 +349,9 @@ fn main(
333
  workgroupBarrier();
334
  }
335
  // Re-stage workgroup tiles/scratch before the next M-tile iteration reuses them.
336
- // The loop bound is workgroup-uniform (M is a compile-time const, num_wg.y and
337
- // workgroup_id.y are uniform), so every invocation reaches this barrier together.
 
338
  workgroupBarrier();
339
  }
340
  }
 
8
  enable chromium_experimental_subgroup_matrix;
9
  diagnostic(off, chromium.subgroup_matrix_uniformity);
10
 
 
11
  {{ env.wgsl.resourceDeclarations }}
 
12
  {% set operandScalar = fScalar %}
13
  {% set accScalar = "f32" %}
14
+ {% set GENERAL = true %}
15
  {% set TAIL = tailSafe is defined and tailSafe %}
16
  {% set SPLIT_K = splitK if splitK is defined else 1 %}
17
  {% set OUT = outputBuffer if outputBuffer is defined else "c" %}
18
  {% set OUT_SCALAR = outScalar if outScalar is defined else T %}
19
+ {% set STATIC_M = M is defined %}
20
+ {% set ROWS = "M" if STATIC_M else "params.M" %}
21
+ {% set ROW_TEST = "a_global < M" if STATIC_M else "row_in" %}
22
+ {% set aDims = (aShape | default([])) if STATIC_M else (aBatchShape | default([])) %}
23
+ {% set kPerSplit = kPerSplit | default(0) %}
24
  {% set aR = aRank %}
25
  {% set bR = bRank %}
26
  {% set aBatchLen = aR - 2 %}
27
  {% set bBatchLen = bR - 2 %}
28
  {% set batchRank = aBatchLen %}
29
+ {% set aMStride = K %}
30
  {% set aKStride = 1 %}
31
+ {% set bKStride = 1 if bTransposed is defined and bTransposed else bShape[bR-1] %}
32
+ {% set bNStride = bShape[bR-2] if bTransposed is defined and bTransposed else 1 %}
33
  {% if bStorageStrides is defined %}{% set bKStride = bStorageStrides[bR-2] %}{% set bNStride = bStorageStrides[bR-1] %}{% endif %}
34
 
35
+ {% if STATIC_M %}
36
  const M: u32 = {{ M }}u;
37
+ {% endif %}
38
  const K: u32 = {{ K }}u;
39
  const N: u32 = {{ N }}u;
40
  const BATCH_COUNT: u32 = {{ batchCount if batchCount is defined else 1 }}u;
 
49
  {% if TAIL %}const K_FULL: u32 = (K / 32u) * 32u;
50
  {% endif %}
51
  const ALPHA: f32 = f32({{ alpha }});
52
+ {% if STATIC_M %}
53
  const C_BATCH_STRIDE: u32 = M * N;
54
+ {% endif %}
55
  const TILE_COLS: u32 = 64u;
56
  const TILE_ROWS: u32 = 32u;
57
  const TILE_K: u32 = 32u;
 
64
 
65
  fn loadSHMA(a_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
66
  let a_global = tile_base + row;
67
+ {% if not STATIC_M %}
68
+ let row_in = a_global < params.M;
69
+ {% endif %}
70
  let col = c_idx * 8u;
71
  for (var col_offset = 0u; col_offset < 8u; col_offset = col_offset + 1u) {
72
  let k = k_idx + col + col_offset;
73
+ if ({{ ROW_TEST }}) {
74
  {% if operandScalar == "f16" %}
75
  tile_A[row * TILE_K + col + col_offset] = f16(a[a_base + a_global * A_M_STRIDE + k * A_K_STRIDE]);
76
  {% else %}
 
89
 
90
  fn loadSHMAKTail(a_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
91
  let a_global = tile_base + row;
92
+ let row_in = a_global < {{ ROWS }};
93
  let col = c_idx * 8u;
94
  for (var col_offset = 0u; col_offset < 8u; col_offset = col_offset + 1u) {
95
  let k = k_idx + col + col_offset;
96
+ if (row_in && k < K) {
97
  tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(a[a_base + a_global * A_M_STRIDE + k * A_K_STRIDE]);
98
  } else {
99
  tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(0);
 
168
  @builtin(subgroup_invocation_id) sg_id: u32,
169
  @builtin(subgroup_size) sg_size: u32
170
  ) {
171
+ {% if not STATIC_M %}
172
+ // The row count arrives per call; the strides it scales are formed here.
173
+ let M = params.M;
174
+ let C_BATCH_STRIDE = M * N;
175
+ {% endif %}
176
  let b_global_base = workgroup_id.x * TILE_COLS;
177
 
178
  let subtile_id = local_idx / sg_size;
 
207
  {% set axis = batchRank - 1 - i %}
208
  {% set aAxis = axis - (batchRank - aBatchLen) %}
209
  {% set bAxis = axis - (batchRank - bBatchLen) %}
210
+ {% set aDim = aDims[aAxis] %}
211
  {% set bDim = bShape[bAxis] if bAxis >= 0 else 1 %}
212
  {% if aDim > 1 or bDim > 1 %}{% set hasBatchCoord.value = true %}{% endif %}
213
  {% endfor %}
 
221
  {% set axis = batchRank - 1 - i %}
222
  {% set aAxis = axis - (batchRank - aBatchLen) %}
223
  {% set bAxis = axis - (batchRank - bBatchLen) %}
224
+ {% set aDim = aDims[aAxis] %}
225
  {% set bDim = bShape[bAxis] if bAxis >= 0 else 1 %}
226
  {% set outDim = aDim if aDim >= bDim else bDim %}
227
  {% set aStride = namespace(v=1) %}
228
+ {% if aDim != 1 %}{% for j in range(aAxis + 1, aR - 2) %}{% set aStride.v = aStride.v * aDims[j] %}{% endfor %}{% else %}{% set aStride.v = 0 %}{% endif %}
229
  {% set bStride = namespace(v=1) %}
230
  {% if bAxis >= 0 and bDim != 1 %}{% for j in range(bAxis + 1, bR) %}{% set bStride.v = bStride.v * bShape[j] %}{% endfor %}{% else %}{% set bStride.v = 0 %}{% endif %}
231
  {% if bStorageStrides is defined and bAxis >= 0 and bDim != 1 %}{% set bStride.v = bStorageStrides[bAxis] %}{% endif %}
232
  {% if outDim > 1 %}
233
  let c{{ axis }} = zTmp % {{ outDim }}u;
234
  zTmp = zTmp / {{ outDim }}u;
235
+ {% if aStride.v != 0 %} a_base = a_base + c{{ axis }} * {% if aStride.v != 1 %}{{ aStride.v }}u * {% endif %}M * K;
236
  {% endif %}
237
  {% if bStride.v != 0 %} b_base = b_base + c{{ axis }} * {{ bStride.v }}u;
238
  {% endif %}
 
265
  workgroupBarrier();
266
 
267
  for (var step = 0u; step < TILE_K; step = step + 8u) {
268
+ {% set directInputs = false %}
269
  let matrix_a_offset = subtile_idy * SUB_ROWS * TILE_K + step;
270
  {% for r in range(2) %}
271
  var matA{{ r }}: subgroup_matrix_left<{{ operandScalar }}, 8, 8> = subgroupMatrixLoad<subgroup_matrix_left<{{ operandScalar }}, 8, 8>, row_major>(&{{ "w" if directInputs else "tile_A" }}, matrix_a_offset{% if r > 0 %} + 8u * {{ "K" if directInputs else "TILE_K" }}{% endif %}, {{ "K" if directInputs else "TILE_K" }});
 
294
  workgroupBarrier();
295
 
296
  for (var step = 0u; step < TILE_K; step = step + 8u) {
297
+ {% set directInputs = false %}
298
+ let matrix_a_offset = subtile_idy * SUB_ROWS * TILE_K + step;
299
  {% for r in range(2) %}
300
  var matA{{ r }}: subgroup_matrix_left<{{ operandScalar }}, 8, 8> = subgroupMatrixLoad<subgroup_matrix_left<{{ operandScalar }}, 8, 8>, row_major>(&{{ "w" if directInputs else "tile_A" }}, matrix_a_offset{% if r > 0 %} + 8u * {{ "K" if directInputs else "TILE_K" }}{% endif %}, {{ "K" if directInputs else "TILE_K" }});
301
  {% endfor %}
302
 
303
+ let matrix_b_offset = subtile_idx * SUB_COLS * TILE_K + step;
304
  {% for c in range(4) %}
305
+ var matB{{ c }}: subgroup_matrix_right<{{ operandScalar }}, 8, 8> = subgroupMatrixLoad<subgroup_matrix_right<{{ operandScalar }}, 8, 8>, col_major>(&{{ "xm" if directInputs else "tile_B" }}, matrix_b_offset{% if c > 0 %} + {{ c * 8 }}u * TILE_K{% endif %}, {{ "N" if directInputs else "TILE_K" }});
306
  {% endfor %}
307
 
308
  matC00 = subgroupMatrixMultiplyAccumulate(matA0, matB0, matC00);
 
349
  workgroupBarrier();
350
  }
351
  // Re-stage workgroup tiles/scratch before the next M-tile iteration reuses them.
352
+ // The loop bound is workgroup-uniform (M is a compile-time const or a uniform
353
+ // read, and num_wg.y and workgroup_id.y are uniform), so every invocation
354
+ // reaches this barrier together.
355
  workgroupBarrier();
356
  }
357
  }
build/webgpu/matmul-tiled-general-reg.wgsl.jinja CHANGED
@@ -3,28 +3,33 @@
3
  // Register-blocked MatMul for the no-subgroup-matrix
4
  // tier: Y = alpha * A @ B. It retains the bounds-checked addressing and
5
  // batch-broadcast of the general kernel, and its transposed-batch-A layout.
6
- // Transposed and 1-D operands select other variants and are not handled here.
 
7
  // Each thread computes a 4x4 micro-tile within a 64x64 workgroup tile, reusing
8
  // each staged operand across four accumulators. Both tiles are indexed by their
9
  // own output axis and group four K values per vector word, so the micro-tile
10
  // accumulates through dot() and one step reads TM + TN words rather than
11
- // 4 * (TM + TN) scalars. A stores K contiguously and B stores N contiguously,
12
- // so each staging lane walks the axis its operand already has.
13
  {% set aR = aRank %}
14
  {% set bR = bRank %}
15
  {% set aBatchLen = aR - 2 %}
16
  {% set bBatchLen = bR - 2 %}
17
  {% set batchRank = aBatchLen %}
 
 
18
  {% set aTailStride = namespace(v=1) %}
19
- {% for j in range(1, aR) %}{% set aTailStride.v = aTailStride.v * aShape[j] %}{% endfor %}
20
  {% set bTailStride = namespace(v=1) %}
21
  {% for j in range(1, bR) %}{% set bTailStride.v = bTailStride.v * bShape[j] %}{% endfor %}
 
22
  {% if transBatchA %}{% set M = aShape[0] %}{% set K = aShape[aR-1] %}
23
  {% else %}{% set M = aShape[aR-2] %}{% set K = aShape[aR-1] %}{% endif %}
 
24
  {% set N = bShape[bR-1] %}
25
  {% if transBatchA %}{% set aMStride = aTailStride.v %}{% set aKStride = 1 %}
26
- {% else %}{% set aMStride = aShape[aR-1] %}{% set aKStride = 1 %}{% endif %}
27
- {% set bKStride = bShape[bR-1] %}{% set bNStride = 1 %}{% if bStorageStrides is defined %}{% set bKStride = bStorageStrides[bR-2] %}{% set bNStride = bStorageStrides[bR-1] %}{% endif %}
 
28
  {% set is_int = (scalar == "i32" or scalar == "u32") %}
29
  {% if is_int %}
30
  // Integer operands accumulate in their integer type, avoiding f32 rounding of
@@ -39,7 +44,9 @@
39
  {% endif %}
40
  {% set tileT = scalar if scalar == "f16" else accT %}
41
  {% set kTile = kTile if kTile is defined else 16 %}
 
42
  const M: u32 = {{ M }}u;
 
43
  const K: u32 = {{ K }}u;
44
  const N: u32 = {{ N }}u;
45
  const A_M_STRIDE: u32 = {{ aMStride }}u;
@@ -50,11 +57,13 @@ const B_N_STRIDE: u32 = {{ bNStride }}u;
50
  // A 4x4 micro-tile over a 64x64 output tile reuses each staged operand across
51
  // four accumulators. It increases arithmetic work per load without the large
52
  // per-thread accumulator footprint of an 8x8 micro-tile.
 
 
53
  const BK: u32 = {{ kTile }}u;
54
- const BM: u32 = 64u;
55
- const BN: u32 = 64u;
56
- const TM: u32 = 4u; // per-thread micro-tile rows
57
- const TN: u32 = 4u; // per-thread micro-tile cols
58
  {% if splitK > 1 %}
59
  const SPLIT_K: u32 = {{ splitK }}u;
60
  const K_PER_SPLIT: u32 = {{ kPerSplit }}u;
@@ -70,19 +79,22 @@ var<workgroup> tileB: array<array<vec4<{{ tileT }}>, K_VECS>, BN>; // B[n][k/4
70
  {% set bAxis = axis - (batchRank - bBatchLen) %}
71
  {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
72
  {% set bStored = bAxis %}
73
- {% set aDim = aShape[aStored] %}
74
  {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
75
  {% if aDim > 1 or bDim > 1 %}{% set hasBatchCoord.value = true %}{% endif %}
76
  {% endfor %}
77
 
78
- @compute @workgroup_size(16, 16, 1)
79
  fn main(
80
  @builtin(workgroup_id) wg: vec3<u32>,
81
  @builtin(local_invocation_id) lid: vec3<u32>
82
  ) {
 
 
 
83
  let mBase = wg.y * BM;
84
  let nBase = wg.x * BN;
85
- let li = lid.y * 16u + lid.x;
86
 
87
  let zOut = wg.z;
88
  {% if splitK > 1 %}
@@ -100,17 +112,17 @@ fn main(
100
  {% set bAxis = axis - (batchRank - bBatchLen) %}
101
  {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
102
  {% set bStored = bAxis %}
103
- {% set aDim = aShape[aStored] %}
104
  {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
105
  {% set outDim = aDim if aDim >= bDim else bDim %}
106
  {% set aStride = namespace(v=1) %}
107
- {% if aStored >= 0 and aDim != 1 %}{% for j in range(aStored + 1, aR) %}{% set aStride.v = aStride.v * aShape[j] %}{% endfor %}{% else %}{% set aStride.v = 0 %}{% endif %}
108
  {% set bStride = namespace(v=1) %}
109
  {% if bStored >= 0 and bDim != 1 %}{% for j in range(bStored + 1, bR) %}{% set bStride.v = bStride.v * bShape[j] %}{% endfor %}{% else %}{% set bStride.v = 0 %}{% endif %}{% if bStorageStrides is defined and bStored >= 0 and bDim != 1 %}{% set bStride.v = bStorageStrides[bStored] %}{% endif %}
110
  {% if outDim > 1 %}
111
  let c{{ axis }} = zTmp % {{ outDim }}u;
112
  zTmp = zTmp / {{ outDim }}u;
113
- {% if aStride.v != 0 %} aBatchOff = aBatchOff + c{{ axis }} * {{ aStride.v }}u;
114
  {% endif %}
115
  {% if bStride.v != 0 %} bBatchOff = bBatchOff + c{{ axis }} * {{ bStride.v }}u;
116
  {% endif %}
@@ -134,7 +146,7 @@ fn main(
134
  {% endif %}
135
  // Cooperative load: one vector word per lane per pass. A's lanes walk K, which
136
  // it stores contiguously; B's walk N, which it stores contiguously.
137
- for (var idx: u32 = li; idx < BM * K_VECS; idx = idx + 256u) {
138
  let ar = idx / K_VECS;
139
  let ac4 = idx % K_VECS;
140
  let am = mBase + ar;
@@ -148,7 +160,7 @@ fn main(
148
  }
149
  tileA[ar][ac4] = aWord;
150
  }
151
- for (var idx: u32 = li; idx < BN * K_VECS; idx = idx + 256u) {
152
  let bc = idx % BN;
153
  let br4 = idx / BN;
154
  let bn = nBase + bc;
@@ -164,6 +176,7 @@ fn main(
164
  }
165
  workgroupBarrier();
166
  {% set regT = accT %}{% filter indent(4, true) %}
 
167
  let aRow = lid.y * TM;
168
  let bCol = lid.x * TN;
169
  for (var kv: u32 = 0u; kv < BK / 4u; kv = kv + 1u) {
@@ -173,9 +186,12 @@ for (var kv: u32 = 0u; kv < BK / 4u; kv = kv + 1u) {
173
  for (var j: u32 = 0u; j < TN; j = j + 1u) { bv[j] = vec4<{{ regT }}>(tileB[bCol + j][kv]); }
174
  for (var i: u32 = 0u; i < TM; i = i + 1u) {
175
  for (var j: u32 = 0u; j < TN; j = j + 1u) {
176
- {% if regSequentialK is defined and regSequentialK %}{% for component in range(4) %}
177
- acc[i * TN + j] = acc[i * TN + j] + av[i][{{ component }}] * bv[j][{{ component }}];
178
- {% endfor %}{% else %} acc[i * TN + j] = acc[i * TN + j] + dot(av[i], bv[j]);
 
 
 
179
  {% endif %}
180
  }
181
  }
 
3
  // Register-blocked MatMul for the no-subgroup-matrix
4
  // tier: Y = alpha * A @ B. It retains the bounds-checked addressing and
5
  // batch-broadcast of the general kernel, and its transposed-batch-A layout.
6
+ // A retains its matrix-axis order; B may transpose its matrix axes.
7
+ // 1-D operands use other variants.
8
  // Each thread computes a 4x4 micro-tile within a 64x64 workgroup tile, reusing
9
  // each staged operand across four accumulators. Both tiles are indexed by their
10
  // own output axis and group four K values per vector word, so the micro-tile
11
  // accumulates through dot() and one step reads TM + TN words rather than
12
+ // 4 * (TM + TN) scalars. Logical matrix axes specialize to physical operand strides.
 
13
  {% set aR = aRank %}
14
  {% set bR = bRank %}
15
  {% set aBatchLen = aR - 2 %}
16
  {% set bBatchLen = bR - 2 %}
17
  {% set batchRank = aBatchLen %}
18
+ {% set STATIC_M = aShape is defined %}
19
+ {% set aDims = aShape if aShape is defined else (aBatchShape | default([])) %}
20
  {% set aTailStride = namespace(v=1) %}
21
+ {% if aShape is defined %}{% for j in range(1, aR) %}{% set aTailStride.v = aTailStride.v * aShape[j] %}{% endfor %}{% endif %}
22
  {% set bTailStride = namespace(v=1) %}
23
  {% for j in range(1, bR) %}{% set bTailStride.v = bTailStride.v * bShape[j] %}{% endfor %}
24
+ {% if aShape is defined %}
25
  {% if transBatchA %}{% set M = aShape[0] %}{% set K = aShape[aR-1] %}
26
  {% else %}{% set M = aShape[aR-2] %}{% set K = aShape[aR-1] %}{% endif %}
27
+ {% endif %}
28
  {% set N = bShape[bR-1] %}
29
  {% if transBatchA %}{% set aMStride = aTailStride.v %}{% set aKStride = 1 %}
30
+ {% else %}{% set aMStride = K %}{% set aKStride = 1 %}{% endif %}
31
+ {% set bKStride = 1 if bTransposed is defined and bTransposed else bShape[bR-1] %}
32
+ {% set bNStride = bShape[bR-2] if bTransposed is defined and bTransposed else 1 %}{% if bStorageStrides is defined %}{% set bKStride = bStorageStrides[bR-2] %}{% set bNStride = bStorageStrides[bR-1] %}{% endif %}
33
  {% set is_int = (scalar == "i32" or scalar == "u32") %}
34
  {% if is_int %}
35
  // Integer operands accumulate in their integer type, avoiding f32 rounding of
 
44
  {% endif %}
45
  {% set tileT = scalar if scalar == "f16" else accT %}
46
  {% set kTile = kTile if kTile is defined else 16 %}
47
+ {% if STATIC_M %}
48
  const M: u32 = {{ M }}u;
49
+ {% endif %}
50
  const K: u32 = {{ K }}u;
51
  const N: u32 = {{ N }}u;
52
  const A_M_STRIDE: u32 = {{ aMStride }}u;
 
57
  // A 4x4 micro-tile over a 64x64 output tile reuses each staged operand across
58
  // four accumulators. It increases arithmetic work per load without the large
59
  // per-thread accumulator footprint of an 8x8 micro-tile.
60
+ {% set microTile = 4 %}
61
+ {% set lanes = 16 %}
62
  const BK: u32 = {{ kTile }}u;
63
+ const BM: u32 = {{ registerTile }}u;
64
+ const BN: u32 = {{ registerTile }}u;
65
+ const TM: u32 = {{ microTile }}u; // per-thread micro-tile rows
66
+ const TN: u32 = {{ microTile }}u; // per-thread micro-tile cols
67
  {% if splitK > 1 %}
68
  const SPLIT_K: u32 = {{ splitK }}u;
69
  const K_PER_SPLIT: u32 = {{ kPerSplit }}u;
 
79
  {% set bAxis = axis - (batchRank - bBatchLen) %}
80
  {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
81
  {% set bStored = bAxis %}
82
+ {% set aDim = aDims[aStored] %}
83
  {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
84
  {% if aDim > 1 or bDim > 1 %}{% set hasBatchCoord.value = true %}{% endif %}
85
  {% endfor %}
86
 
87
+ @compute @workgroup_size({{ lanes }}, {{ lanes }}, 1)
88
  fn main(
89
  @builtin(workgroup_id) wg: vec3<u32>,
90
  @builtin(local_invocation_id) lid: vec3<u32>
91
  ) {
92
+ {% if not STATIC_M %}
93
+ let M = params.M;
94
+ {% endif %}
95
  let mBase = wg.y * BM;
96
  let nBase = wg.x * BN;
97
+ let li = lid.y * {{ lanes }}u + lid.x;
98
 
99
  let zOut = wg.z;
100
  {% if splitK > 1 %}
 
112
  {% set bAxis = axis - (batchRank - bBatchLen) %}
113
  {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
114
  {% set bStored = bAxis %}
115
+ {% set aDim = aDims[aStored] %}
116
  {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
117
  {% set outDim = aDim if aDim >= bDim else bDim %}
118
  {% set aStride = namespace(v=1) %}
119
+ {% if aStored >= 0 and aDim != 1 %}{% for j in range(aStored + 1, aR if STATIC_M else aR - 2) %}{% set aStride.v = aStride.v * aDims[j] %}{% endfor %}{% else %}{% set aStride.v = 0 %}{% endif %}
120
  {% set bStride = namespace(v=1) %}
121
  {% if bStored >= 0 and bDim != 1 %}{% for j in range(bStored + 1, bR) %}{% set bStride.v = bStride.v * bShape[j] %}{% endfor %}{% else %}{% set bStride.v = 0 %}{% endif %}{% if bStorageStrides is defined and bStored >= 0 and bDim != 1 %}{% set bStride.v = bStorageStrides[bStored] %}{% endif %}
122
  {% if outDim > 1 %}
123
  let c{{ axis }} = zTmp % {{ outDim }}u;
124
  zTmp = zTmp / {{ outDim }}u;
125
+ {% if aStride.v != 0 %} aBatchOff = aBatchOff + c{{ axis }} * {% if STATIC_M %}{{ aStride.v }}u{% else %}{% if aStride.v != 1 %}{{ aStride.v }}u * {% endif %}M * K{% endif %};
126
  {% endif %}
127
  {% if bStride.v != 0 %} bBatchOff = bBatchOff + c{{ axis }} * {{ bStride.v }}u;
128
  {% endif %}
 
146
  {% endif %}
147
  // Cooperative load: one vector word per lane per pass. A's lanes walk K, which
148
  // it stores contiguously; B's walk N, which it stores contiguously.
149
+ for (var idx: u32 = li; idx < BM * K_VECS; idx = idx + {{ lanes * lanes }}u) {
150
  let ar = idx / K_VECS;
151
  let ac4 = idx % K_VECS;
152
  let am = mBase + ar;
 
160
  }
161
  tileA[ar][ac4] = aWord;
162
  }
163
+ for (var idx: u32 = li; idx < BN * K_VECS; idx = idx + {{ lanes * lanes }}u) {
164
  let bc = idx % BN;
165
  let br4 = idx / BN;
166
  let bn = nBase + bc;
 
176
  }
177
  workgroupBarrier();
178
  {% set regT = accT %}{% filter indent(4, true) %}
179
+ {% set regAccumulator = "acc" %}
180
  let aRow = lid.y * TM;
181
  let bCol = lid.x * TN;
182
  for (var kv: u32 = 0u; kv < BK / 4u; kv = kv + 1u) {
 
186
  for (var j: u32 = 0u; j < TN; j = j + 1u) { bv[j] = vec4<{{ regT }}>(tileB[bCol + j][kv]); }
187
  for (var i: u32 = 0u; i < TM; i = i + 1u) {
188
  for (var j: u32 = 0u; j < TN; j = j + 1u) {
189
+ {% if regSequentialK is defined and regSequentialK %}
190
+ {% for component in range(4) %}
191
+ {{ regAccumulator }}[i * TN + j] = {{ regAccumulator }}[i * TN + j] + av[i][{{ component }}] * bv[j][{{ component }}];
192
+ {% endfor %}
193
+ {% else %}
194
+ {{ regAccumulator }}[i * TN + j] = {{ regAccumulator }}[i * TN + j] + dot(av[i], bv[j]);
195
  {% endif %}
196
  }
197
  }
build/webgpu/matmul-tiled-general.wgsl.jinja CHANGED
@@ -18,9 +18,10 @@
18
  * [d1, ..., dR-2, d0, dR-1]. Stored axis zero becomes the M axis, the final K
19
  * axis is unchanged, and the remaining axes form the batch. This stride
20
  * permutation composes with the ordinary last-two-axis transpose. */
 
 
21
  {% set aTailStride = namespace(v=1) %}
22
- {% for j in range(1, aR) %}{% set aTailStride.v = aTailStride.v * aShape[j] %}{% endfor %}
23
- {% set bTailStride = namespace(v=1) %}
24
  {% for j in range(1, bR) %}{% set bTailStride.v = bTailStride.v * bShape[j] %}{% endfor %}
25
  {% if aVec %}{% set M = 1 %}{% set K = aShape[0] %}
26
  {% elif transBatchA and transA %}{% set M = aShape[aR-1] %}{% set K = aShape[0] %}
@@ -34,8 +35,8 @@
34
  {% if aVec %}{% set aMStride = 0 %}{% set aKStride = 1 %}
35
  {% elif transBatchA and transA %}{% set aMStride = 1 %}{% set aKStride = aTailStride.v %}
36
  {% elif transBatchA %}{% set aMStride = aTailStride.v %}{% set aKStride = 1 %}
37
- {% elif transA %}{% set aMStride = 1 %}{% set aKStride = aShape[aR-1] %}
38
- {% else %}{% set aMStride = aShape[aR-1] %}{% set aKStride = 1 %}{% endif %}
39
  {% if bVec %}{% set bKStride = 1 %}{% set bNStride = 0 %}
40
  {% elif transBatchB and transB %}{% set bKStride = 1 %}{% set bNStride = bTailStride.v %}
41
  {% elif transBatchB %}{% set bKStride = bTailStride.v %}{% set bNStride = 1 %}
@@ -58,12 +59,13 @@ const ALPHA: {{ accT }} = {{ accT }}(1);{% else %}const ALPHA: f32 = f32({{ alph
58
  // 2x2 register-blocked tile: 16x16 threads each compute a 2x2 micro-tile, for a
59
  // 32x32 output tile per workgroup with K stepped in BK=16 chunks. Each loaded
60
  // shared-mem element feeds 2 FMAs, favoring register reuse in the inner loop.
61
- const BK: u32 = 16u;
62
- const BM: u32 = 32u;
63
- const BN: u32 = 32u;
 
64
 
65
- var<workgroup> tileA: array<array<{{ tileT }}, 16>, 32>;
66
- var<workgroup> tileB: array<array<{{ tileT }}, 32>, 16>;
67
  {% set hasBatchCoord = namespace(value=false) %}
68
  {% for i in range(batchRank) %}
69
  {% set axis = batchRank - 1 - i %}
@@ -71,19 +73,19 @@ var<workgroup> tileB: array<array<{{ tileT }}, 32>, 16>;
71
  {% set bAxis = axis - (batchRank - bBatchLen) %}
72
  {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
73
  {% set bStored = (bAxis + 1) if (transBatchB and bAxis >= 0) else bAxis %}
74
- {% set aDim = aShape[aStored] if aStored >= 0 else 1 %}
75
  {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
76
  {% if aDim > 1 or bDim > 1 %}{% set hasBatchCoord.value = true %}{% endif %}
77
  {% endfor %}
78
 
79
- @compute @workgroup_size(16, 16, 1)
80
  fn main(
81
  @builtin(workgroup_id) wg: vec3<u32>,
82
  @builtin(local_invocation_id) lid: vec3<u32>
83
  ) {
84
  let mBase = wg.y * BM;
85
  let nBase = wg.x * BN;
86
- let li = lid.y * 16u + lid.x;
87
 
88
  // Per-batch base offsets into A and B using right-aligned broadcast strides.
89
  // Decompose the flat output-batch index from the innermost axis outward.
@@ -99,11 +101,11 @@ fn main(
99
  {% set bAxis = axis - (batchRank - bBatchLen) %}
100
  {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
101
  {% set bStored = (bAxis + 1) if (transBatchB and bAxis >= 0) else bAxis %}
102
- {% set aDim = aShape[aStored] if aStored >= 0 else 1 %}
103
  {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
104
  {% set outDim = aDim if aDim >= bDim else bDim %}
105
  {% set aStride = namespace(v=1) %}
106
- {% if aStored >= 0 and aDim != 1 %}{% for j in range(aStored + 1, aR) %}{% set aStride.v = aStride.v * aShape[j] %}{% endfor %}{% else %}{% set aStride.v = 0 %}{% endif %}
107
  {% set bStride = namespace(v=1) %}
108
  {% if bStored >= 0 and bDim != 1 %}{% for j in range(bStored + 1, bR) %}{% set bStride.v = bStride.v * bShape[j] %}{% endfor %}{% else %}{% set bStride.v = 0 %}{% endif %}
109
  {% if outDim > 1 %}
@@ -127,7 +129,7 @@ fn main(
127
  let kBase = kt * BK;
128
  // Cooperative load: 32x16 A tile + 16x32 B tile, 256 threads x 2 each.
129
  for (var e: u32 = 0u; e < 2u; e = e + 1u) {
130
- let idx = li + e * 256u;
131
  let ar = idx / BK;
132
  let ac = idx % BK;
133
  let am = mBase + ar;
 
18
  * [d1, ..., dR-2, d0, dR-1]. Stored axis zero becomes the M axis, the final K
19
  * axis is unchanged, and the remaining axes form the batch. This stride
20
  * permutation composes with the ordinary last-two-axis transpose. */
21
+ {% set STATIC_M = true %}
22
+ {% set aDims = aShape if aShape is defined else (aBatchShape | default([])) %}
23
  {% set aTailStride = namespace(v=1) %}
24
+ {% for j in range(1, aR) %}{% set aTailStride.v = aTailStride.v * aShape[j] %}{% endfor %}{% set bTailStride = namespace(v=1) %}
 
25
  {% for j in range(1, bR) %}{% set bTailStride.v = bTailStride.v * bShape[j] %}{% endfor %}
26
  {% if aVec %}{% set M = 1 %}{% set K = aShape[0] %}
27
  {% elif transBatchA and transA %}{% set M = aShape[aR-1] %}{% set K = aShape[0] %}
 
35
  {% if aVec %}{% set aMStride = 0 %}{% set aKStride = 1 %}
36
  {% elif transBatchA and transA %}{% set aMStride = 1 %}{% set aKStride = aTailStride.v %}
37
  {% elif transBatchA %}{% set aMStride = aTailStride.v %}{% set aKStride = 1 %}
38
+ {% elif transA %}{% set aMStride = 1 %}{% set aKStride = M %}
39
+ {% else %}{% set aMStride = K %}{% set aKStride = 1 %}{% endif %}
40
  {% if bVec %}{% set bKStride = 1 %}{% set bNStride = 0 %}
41
  {% elif transBatchB and transB %}{% set bKStride = 1 %}{% set bNStride = bTailStride.v %}
42
  {% elif transBatchB %}{% set bKStride = bTailStride.v %}{% set bNStride = 1 %}
 
59
  // 2x2 register-blocked tile: 16x16 threads each compute a 2x2 micro-tile, for a
60
  // 32x32 output tile per workgroup with K stepped in BK=16 chunks. Each loaded
61
  // shared-mem element feeds 2 FMAs, favoring register reuse in the inner loop.
62
+ {% set lanes = 16 %}
63
+ const BK: u32 = {{ lanes }}u;
64
+ const BM: u32 = {{ generalTile }}u;
65
+ const BN: u32 = {{ generalTile }}u;
66
 
67
+ var<workgroup> tileA: array<array<{{ tileT }}, {{ lanes }}>, {{ generalTile }}>;
68
+ var<workgroup> tileB: array<array<{{ tileT }}, {{ generalTile }}>, {{ lanes }}>;
69
  {% set hasBatchCoord = namespace(value=false) %}
70
  {% for i in range(batchRank) %}
71
  {% set axis = batchRank - 1 - i %}
 
73
  {% set bAxis = axis - (batchRank - bBatchLen) %}
74
  {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
75
  {% set bStored = (bAxis + 1) if (transBatchB and bAxis >= 0) else bAxis %}
76
+ {% set aDim = aDims[aStored] if aStored >= 0 else 1 %}
77
  {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
78
  {% if aDim > 1 or bDim > 1 %}{% set hasBatchCoord.value = true %}{% endif %}
79
  {% endfor %}
80
 
81
+ @compute @workgroup_size({{ lanes }}, {{ lanes }}, 1)
82
  fn main(
83
  @builtin(workgroup_id) wg: vec3<u32>,
84
  @builtin(local_invocation_id) lid: vec3<u32>
85
  ) {
86
  let mBase = wg.y * BM;
87
  let nBase = wg.x * BN;
88
+ let li = lid.y * {{ lanes }}u + lid.x;
89
 
90
  // Per-batch base offsets into A and B using right-aligned broadcast strides.
91
  // Decompose the flat output-batch index from the innermost axis outward.
 
101
  {% set bAxis = axis - (batchRank - bBatchLen) %}
102
  {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
103
  {% set bStored = (bAxis + 1) if (transBatchB and bAxis >= 0) else bAxis %}
104
+ {% set aDim = aDims[aStored] if aStored >= 0 else 1 %}
105
  {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
106
  {% set outDim = aDim if aDim >= bDim else bDim %}
107
  {% set aStride = namespace(v=1) %}
108
+ {% if aStored >= 0 and aDim != 1 %}{% for j in range(aStored + 1, aR if STATIC_M else aR - 2) %}{% set aStride.v = aStride.v * aDims[j] %}{% endfor %}{% else %}{% set aStride.v = 0 %}{% endif %}
109
  {% set bStride = namespace(v=1) %}
110
  {% if bStored >= 0 and bDim != 1 %}{% for j in range(bStored + 1, bR) %}{% set bStride.v = bStride.v * bShape[j] %}{% endfor %}{% else %}{% set bStride.v = 0 %}{% endif %}
111
  {% if outDim > 1 %}
 
129
  let kBase = kt * BK;
130
  // Cooperative load: 32x16 A tile + 16x32 B tile, 256 threads x 2 each.
131
  for (var e: u32 = 0u; e < 2u; e = e + 1u) {
132
+ let idx = li + e * {{ lanes * lanes }}u;
133
  let ar = idx / BK;
134
  let ac = idx % BK;
135
  let am = mBase + ar;
build/webgpu/matmul-vector-matrix-vec4.wgsl.jinja CHANGED
@@ -3,14 +3,14 @@
3
  {{ env.wgsl.resourceDeclarations }}
4
  {% set OUT = outputBuffer if outputBuffer is defined else "c" %}
5
 
6
- const LANES: u32 = 32u;
7
  // SLICES partitions the K reduction across the workgroup's second dimension.
8
  // Thread zero of each column group combines the slice partials in index order.
9
  const SLICES: u32 = {{ gemvSlices }}u;
10
 
11
  var<workgroup> partials: array<vec4<f32>, LANES * SLICES>;
12
 
13
- @compute @workgroup_size(32, {{ gemvSlices }}, 1)
14
  fn main(
15
  @builtin(workgroup_id) workgroup_id: vec3<u32>,
16
  @builtin(local_invocation_id) lid: vec3<u32>
@@ -20,9 +20,23 @@ fn main(
20
  let cg = workgroup_id.x * LANES + lane;
21
  var acc = vec4<f32>(0.0);
22
  if (cg < params.N4) {
 
 
 
 
 
 
 
 
 
 
 
 
 
23
  for (var k = slice; k < params.K; k = k + SLICES) {
24
  acc = acc + f32(a[k]) * vec4<f32>(b[k * params.N4 + cg]);
25
  }
 
26
  }
27
  partials[slice * LANES + lane] = acc;
28
  workgroupBarrier();
 
3
  {{ env.wgsl.resourceDeclarations }}
4
  {% set OUT = outputBuffer if outputBuffer is defined else "c" %}
5
 
6
+ const LANES: u32 = {{ gemvLanes }}u;
7
  // SLICES partitions the K reduction across the workgroup's second dimension.
8
  // Thread zero of each column group combines the slice partials in index order.
9
  const SLICES: u32 = {{ gemvSlices }}u;
10
 
11
  var<workgroup> partials: array<vec4<f32>, LANES * SLICES>;
12
 
13
+ @compute @workgroup_size({{ gemvLanes }}, {{ gemvSlices }}, 1)
14
  fn main(
15
  @builtin(workgroup_id) workgroup_id: vec3<u32>,
16
  @builtin(local_invocation_id) lid: vec3<u32>
 
20
  let cg = workgroup_id.x * LANES + lane;
21
  var acc = vec4<f32>(0.0);
22
  if (cg < params.N4) {
23
+ {% if unrollK2 is defined and unrollK2 %}
24
+ // Two K positions per iteration amortize loop/address arithmetic on long
25
+ // decode projections while preserving each slice's exact strided order.
26
+ var k = slice;
27
+ for (; k + SLICES < params.K; k = k + 2u * SLICES) {
28
+ acc = acc + f32(a[k]) * vec4<f32>(b[k * params.N4 + cg]);
29
+ let k1 = k + SLICES;
30
+ acc = acc + f32(a[k1]) * vec4<f32>(b[k1 * params.N4 + cg]);
31
+ }
32
+ if (k < params.K) {
33
+ acc = acc + f32(a[k]) * vec4<f32>(b[k * params.N4 + cg]);
34
+ }
35
+ {% else %}
36
  for (var k = slice; k < params.K; k = k + SLICES) {
37
  acc = acc + f32(a[k]) * vec4<f32>(b[k * params.N4 + cg]);
38
  }
39
+ {% endif %}
40
  }
41
  partials[slice * LANES + lane] = acc;
42
  workgroupBarrier();
build/webgpu/metadata.json CHANGED
@@ -1,31 +1,34 @@
1
  {
2
  "name": "com.microsoft.FusedMatMul",
3
- "id": "_com_microsoft_fusedmatmul_webgpu_8b2cb77",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "ATSeCpumMpLFypVxXuSfXbcCPQhVBBmfOof4PcQd89o=",
11
- "fused-matmul-subgroup-matrix.wgsl.jinja": "0G1fPWFTk4qK1SiuAcriccOujenbWQfIgE0F0iXPlFo=",
12
- "manifest.json": "pl4JzUzE2Z8hDvUoxFEkFv1qWvP/cMYQvOPN0s3asPs=",
13
- "matmul-band-vec4.wgsl.jinja": "CN34bTiT4RnjmEoH/zPKP3a0tRpfYqP4l/X/bKCvHUk=",
14
- "matmul-subgroup-matrix-ext.wgsl.jinja": "9hswKFk/g2HxvNxKwowsc7cJtEUGvrL53ISLi46Y3Co=",
15
- "matmul-tiled-general-reg.wgsl.jinja": "Td1g9ghExH+uVbEhwxTEo1jkBpmkgKjBlkWGMD6JZ5I=",
16
- "matmul-tiled-general.wgsl.jinja": "7vW3jelY9YQHhc92bE/yQZFQk4yeuSRby23ELji4rFs=",
17
- "matmul-vector-matrix-vec4.wgsl.jinja": "TxVXDiRp6BURoOhTbh2ILDAAL9D9cVDfUhk8e5ViSfU=",
18
- "reduce-axis0-splitk-combine.wgsl.jinja": "Yz1hjK55R/kndUrw3ugPqgPaOBwKmYupqZoVdJTHO5Q=",
19
- "test.json": "CuQ/9/rI0a9gPgB8JwczBXvdF7PIuQoXpiFFKVmkeC4="
20
  }
21
  },
22
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
23
  "webgpu": {
24
- "manifestSpec": "2.0",
25
  "variants": {
 
 
 
26
  "subgroup_matrix_transbatch_b_f16": ["matmul-subgroup-matrix-ext.wgsl.jinja"],
27
  "subgroup_matrix_transbatch_b_f32": ["matmul-subgroup-matrix-ext.wgsl.jinja"],
28
- "f32_m1_gemv_vec4": ["matmul-vector-matrix-vec4.wgsl.jinja"],
29
  "rank2_band_vec4_splitk": ["matmul-band-vec4.wgsl.jinja", "reduce-axis0-splitk-combine.wgsl.jinja"],
30
  "rank2_band_vec4": ["matmul-band-vec4.wgsl.jinja"],
31
  "rank2_band_vec4_f32_preferred": ["matmul-band-vec4.wgsl.jinja"],
 
1
  {
2
  "name": "com.microsoft.FusedMatMul",
3
+ "id": "_com_microsoft_fusedmatmul_webgpu_17d6f48",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "bench.json": "hPORMosjCjoEdszLuJCkanMmm8j/WxJFR8aS4bZRdWw=",
11
+ "fused-matmul-subgroup-matrix.wgsl.jinja": "idFWg6H0GgKnO1AG8WFLO2rXeN4cqOwW+V9dJii9Fn4=",
12
+ "manifest.json": "epnZxQxVo55octrjWqNWM/cs0DsR74Cz59XvUB+MyOI=",
13
+ "matmul-band-vec4.wgsl.jinja": "LKAs6A++OJF0ZEEM4JITZmKXkr/wo9gd5qrR3zDZy1g=",
14
+ "matmul-subgroup-matrix-ext.wgsl.jinja": "MGpbY/2ZIE3EEnG0SmKB74zCrPmVGbnNBSw9Ml81+7c=",
15
+ "matmul-tiled-general-reg.wgsl.jinja": "DGpeWkL4fuu22CI4+Yw3/QD2NON89wCggbotpFNWobE=",
16
+ "matmul-tiled-general.wgsl.jinja": "j+8KKAFZDVh9M3zNFxJCU4ihsiLgI0KxpE9CD8TqF6Q=",
17
+ "matmul-vector-matrix-vec4.wgsl.jinja": "Qiv2AO8MWj1BAoPLMCH98oXBasrQYFWFh/IUhdhZp5I=",
18
+ "reduce-axis0-splitk-combine.wgsl.jinja": "Zf7tz8nrapi2KMZIbv8cT2f4pliJoJHj4iwS4hc/6ms=",
19
+ "test.json": "LCHLBZTTSQFp/nqb55OJuX7oiL59ThEfHhocijVk5DE="
20
  }
21
  },
22
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
23
  "webgpu": {
24
+ "manifestSpec": "2.1",
25
  "variants": {
26
+ "broadcast_transb_tiled_reg": ["matmul-tiled-general-reg.wgsl.jinja"],
27
+ "broadcast_transb_subgroup_matrix_f16": ["matmul-subgroup-matrix-ext.wgsl.jinja"],
28
+ "broadcast_transb_subgroup_matrix_f32": ["matmul-subgroup-matrix-ext.wgsl.jinja"],
29
  "subgroup_matrix_transbatch_b_f16": ["matmul-subgroup-matrix-ext.wgsl.jinja"],
30
  "subgroup_matrix_transbatch_b_f32": ["matmul-subgroup-matrix-ext.wgsl.jinja"],
31
+ "m1_gemv_vec4": ["matmul-vector-matrix-vec4.wgsl.jinja"],
32
  "rank2_band_vec4_splitk": ["matmul-band-vec4.wgsl.jinja", "reduce-axis0-splitk-combine.wgsl.jinja"],
33
  "rank2_band_vec4": ["matmul-band-vec4.wgsl.jinja"],
34
  "rank2_band_vec4_f32_preferred": ["matmul-band-vec4.wgsl.jinja"],
build/webgpu/reduce-axis0-splitk-combine.wgsl.jinja CHANGED
@@ -2,37 +2,19 @@
2
  // folds the segment partials and applies the selected reduction's final step.
3
  // Segments are folded in ascending order for deterministic results. This order
4
  // differs from the single-pass reduction but remains within the f32 tolerance.
5
- {% set addBias = addBias is defined and addBias %}
6
- {% set biasCols = biasCols | default(0) %}
7
- {% set intMode = intMode is defined and intMode %}
8
  {% set yv = "f16(" if outputF16 else "" %}
9
  {% set vy = ")" if outputF16 else "" %}
10
- {% if outputF16 %}
11
- enable f16;
12
- {% endif %}
13
  {{ env.wgsl.resourceDeclarations }}
14
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
15
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
16
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
17
- {% set logicalBoolIdentity = logicalBool is defined and logicalBool %}
18
  fn {{ name }}() -> {{ scalar }} {
19
- {% if scalar == "i32" %}
20
- return {{ "-2147483647i - 1i" if op == "max" else "2147483647i" }};
21
- {% elif scalar == "u32" %}
22
- return {{ "1u" if logicalBoolIdentity and op == "min" else "0u" if op == "max" else "4294967295u" }};
23
- {% else %}
24
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
25
  return bitcast<f32>(bits);
26
- {% endif %}
27
- }
28
- {%- endmacro %}
29
-
30
 
31
  const WG: u32 = {{ workgroupSize }}u;
32
  const SPLIT: u32 = {{ split }}u;
33
- {% if addBias %}
34
- const BIAS_COLS: u32 = {{ biasCols }}u;
35
- {% endif %}
36
  {% if op == "logsumexp" %}
37
  const F32_MIN: f32 = -3.4028234663852886e38;
38
  const F32_MAX: f32 = 3.4028234663852886e38;
@@ -48,7 +30,9 @@ fn is_nan_f32(value: f32) -> bool {
48
  @compute @workgroup_size(WG, 1, 1)
49
  fn main(@builtin(global_invocation_id) gid: vec3<u32>,
50
  @builtin(num_workgroups) nwg: vec3<u32>) {
51
- let stride = nwg.x * WG;
 
 
52
  let start = (gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG) + gid.x;
53
  for (var col = start; col < params.cols; col = col + stride) {
54
  {% if op == "logsumexp" %}
@@ -74,13 +58,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>,
74
  let finite_or_inf = select(global_max + log(sum), global_max, has_positive_inf);
75
  y[col] = {{ yv }}select(finite_or_inf, nan_value, has_nan){{ vy }};
76
  {% else %}
77
- {% if intMode %}
78
- {% if op == "prod" %}
79
- var total = 1i;
80
- {% else %}
81
- var total = 0i;
82
- {% endif %}
83
- {% else %}
84
  {% if op == "max" %}
85
  var total = reduction_identity();
86
  {% elif op == "min" %}
@@ -89,7 +66,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>,
89
  var total = 1.0;
90
  {% else %}
91
  var total = 0.0;
92
- {% endif %}
93
  {% endif %}
94
  for (var seg = 0u; seg < SPLIT; seg = seg + 1u) {
95
  let p = partials[seg * params.cols + col];
@@ -101,9 +77,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>,
101
  total = total + p;
102
  {% endif %}
103
  }
104
- {% if addBias %}
105
- total = total + f32(bias[col % BIAS_COLS]);
106
- {% endif %}
107
  {% if op == "l2" %}
108
  y[col] = {{ yv }}sqrt(total){{ vy }};
109
  {% elif op == "logsum" %}
 
2
  // folds the segment partials and applies the selected reduction's final step.
3
  // Segments are folded in ascending order for deterministic results. This order
4
  // differs from the single-pass reduction but remains within the f32 tolerance.
 
 
 
5
  {% set yv = "f16(" if outputF16 else "" %}
6
  {% set vy = ")" if outputF16 else "" %}
 
 
 
7
  {{ env.wgsl.resourceDeclarations }}
8
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
9
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
10
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
 
11
  fn {{ name }}() -> {{ scalar }} {
 
 
 
 
 
12
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
13
  return bitcast<f32>(bits);
14
+ }{% endmacro %}
 
 
 
15
 
16
  const WG: u32 = {{ workgroupSize }}u;
17
  const SPLIT: u32 = {{ split }}u;
 
 
 
18
  {% if op == "logsumexp" %}
19
  const F32_MIN: f32 = -3.4028234663852886e38;
20
  const F32_MAX: f32 = 3.4028234663852886e38;
 
30
  @compute @workgroup_size(WG, 1, 1)
31
  fn main(@builtin(global_invocation_id) gid: vec3<u32>,
32
  @builtin(num_workgroups) nwg: vec3<u32>) {
33
+ // The start already folds gid.y in, so the stride must span every y row too;
34
+ // an x-only stride would send y = 0 lanes over columns the y >= 1 rows own.
35
+ let stride = nwg.x * nwg.y * WG;
36
  let start = (gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG) + gid.x;
37
  for (var col = start; col < params.cols; col = col + stride) {
38
  {% if op == "logsumexp" %}
 
58
  let finite_or_inf = select(global_max + log(sum), global_max, has_positive_inf);
59
  y[col] = {{ yv }}select(finite_or_inf, nan_value, has_nan){{ vy }};
60
  {% else %}
 
 
 
 
 
 
 
61
  {% if op == "max" %}
62
  var total = reduction_identity();
63
  {% elif op == "min" %}
 
66
  var total = 1.0;
67
  {% else %}
68
  var total = 0.0;
 
69
  {% endif %}
70
  for (var seg = 0u; seg < SPLIT; seg = seg + 1u) {
71
  let p = partials[seg * params.cols + col];
 
77
  total = total + p;
78
  {% endif %}
79
  }
 
 
 
80
  {% if op == "l2" %}
81
  y[col] = {{ yv }}sqrt(total){{ vy }};
82
  {% elif op == "logsum" %}
build/webgpu/test.json CHANGED
@@ -7,6 +7,25 @@
7
  "ort_float32_trans_batch_b_input_B": [1, 0, 1, 2, 0, 1, -1, 0, 1, 1, 0, 1, 2, -1, 1, 1]
8
  },
9
  "cases": [
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
10
  {
11
  "name": "ort_float32_broadcast_rank4_by_rank3",
12
  "provenance": {
@@ -693,7 +712,7 @@
693
  "provenance": {
694
  "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
695
  "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
696
- "notes": "M=32, K=32, N=64 selects the subgroup-matrix path; finite subnormal dot products must not flush to zero."
697
  },
698
  "inputs": {
699
  "A": { "dtype": "float32", "shape": [32, 32], "data": { "kind": "constant", "value": 1e-20 } },
@@ -936,7 +955,7 @@
936
  {
937
  "name": "f32_decode_gemv_m1_k65_n68_vec4_compact",
938
  "provenance": {
939
- "notes": "Compact M=1 float32 GEMV correctness lock for the model-shaped K=4096,N=4096 bandwidth-bound benchmark. Odd K preserves the sliced reduction while N=68 exercises the final partial 128-column workgroup."
940
  },
941
  "attrs": { "alpha": 1 },
942
  "inputs": {
@@ -953,10 +972,50 @@
953
  },
954
  "outputs": { "Y": { "dtype": "float32", "shape": [1, 68], "tolerance": 0.0002 } }
955
  },
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
956
  {
957
  "name": "f32_rank4_by_rank2_shared_weight_compact",
958
  "provenance": {
959
- "notes": "Compact rank-4 by rank-2 shared-weight broadcast lock for the attention-shaped benchmark. Odd M/K/N exercise batch offset and tile-tail handling."
960
  },
961
  "attrs": { "alpha": 0.5 },
962
  "inputs": {
@@ -975,7 +1034,9 @@
975
  },
976
  {
977
  "name": "subgroup_matrix_kn_tail_f16_compact",
978
- "provenance": { "notes": "Compact f16 subgroup-matrix lock with both a partial K=34 tile and N=66 output tail." },
 
 
979
  "attrs": { "alpha": 0.5 },
980
  "inputs": {
981
  "A": {
@@ -994,7 +1055,7 @@
994
  {
995
  "name": "subgroup_matrix_broadcast_rank4x3_f16_compact",
996
  "provenance": {
997
- "notes": "Compact rank-4 by rank-3 broadcast lock for the model-shaped [1,8,M,K] x [8,K,N] stress case."
998
  },
999
  "attrs": { "alpha": 1 },
1000
  "inputs": {
@@ -1014,7 +1075,7 @@
1014
  {
1015
  "name": "broadcast_rank4_tiled_reg_f16_compact",
1016
  "provenance": {
1017
- "notes": "Compact rank-4 by rank-3 broadcast lock for the register-blocked non-subgroup-matrix path. Odd M/K/N exercise every output and reduction tail."
1018
  },
1019
  "attrs": { "alpha": 0.5 },
1020
  "inputs": {
@@ -1054,7 +1115,7 @@
1054
  {
1055
  "name": "transbatch_a_dense_m_tail_f16_compact",
1056
  "provenance": {
1057
- "notes": "Compact lock for stored [M,batch,K] transBatchA addressing. M=65 exercises the subgroup-matrix row tail; all/no-mma/no-subgroups select the MMA/register-blocked portable paths used by the model-shaped stress case."
1058
  },
1059
  "attrs": { "alpha": 0.5, "transBatchA": 1 },
1060
  "inputs": {
@@ -1109,7 +1170,7 @@
1109
  "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 2, 2, 4], "tolerance": 0.000001 } }
1110
  },
1111
  {
1112
- "name": "subgroup_matrix_kn_tail_f16_offset_alpha_scale_lock",
1113
  "provenance": {
1114
  "notes": "Offset float16 operands keep each output near `alpha * K * aOffset * bOffset` (about 3.4), making the K=34 reduction tail, N=66 column tail, and `alpha = 0.5` epilogue observable on the subgroup-matrix route."
1115
  },
@@ -1129,7 +1190,7 @@
1129
  "outputs": { "Y": { "dtype": "float16", "shape": [33, 66], "tolerance": 0.03, "relTolerance": 0.01 } }
1130
  },
1131
  {
1132
- "name": "subgroup_matrix_broadcast_rank4x3_f16_offset_scale_lock",
1133
  "provenance": {
1134
  "notes": "Offset operands keep outputs near `K * aOffset * bOffset`, making the per-batch B slice and K=32 contraction observable in a rank-4 by rank-3 broadcast."
1135
  },
@@ -1169,7 +1230,7 @@
1169
  "outputs": { "Y": { "dtype": "float16", "shape": [1, 2, 33, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1170
  },
1171
  {
1172
- "name": "broadcast_rank4_tiled_reg_f16_offset_alpha_scale_lock",
1173
  "provenance": {
1174
  "notes": "Offset operands keep outputs near `alpha * K * aOffset * bOffset` with K=33, making the one-element reduction tail and `alpha` multiplier observable on the register-blocked rank-4 route."
1175
  },
@@ -1189,7 +1250,7 @@
1189
  "outputs": { "Y": { "dtype": "float16", "shape": [1, 2, 65, 67], "tolerance": 0.03, "relTolerance": 0.01 } }
1190
  },
1191
  {
1192
- "name": "transbatch_a_dense_m_tail_f16_offset_alpha_scale_lock",
1193
  "provenance": {
1194
  "notes": "Offset operands make each output proportional to `alpha * K`, exposing the [M, batch, K] transBatchA stride, K=32 contraction, and `alpha = 0.5` scale."
1195
  },
@@ -1209,7 +1270,7 @@
1209
  "outputs": { "Y": { "dtype": "float16", "shape": [2, 65, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1210
  },
1211
  {
1212
- "name": "aligned_f16_transA_transB_alpha_offset_scale_lock",
1213
  "provenance": {
1214
  "notes": "Offset operands keep the doubly transposed output proportional to `alpha * K`, making the K=32 contraction and `alpha = 0.5` epilogue observable on subgroup-matrix and portable tiled routes."
1215
  },
@@ -1229,7 +1290,7 @@
1229
  "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1230
  },
1231
  {
1232
- "name": "f16_unaligned_3x5x7_offset_scale_lock",
1233
  "provenance": {
1234
  "notes": "Offset operands in a 3x5 by 5x7 multiply keep outputs proportional to K, exposing dropped reduction elements or doubled tails on the unaligned scalar and tiled routes."
1235
  },
@@ -1248,7 +1309,7 @@
1248
  "outputs": { "Y": { "dtype": "float16", "shape": [3, 7], "tolerance": 0.01, "relTolerance": 0.005 } }
1249
  },
1250
  {
1251
- "name": "aligned_f16_plain_64x32x64_offset_scale_lock",
1252
  "provenance": {
1253
  "notes": "Offset operands keep each fully aligned M=64, K=32, N=64 output proportional to K, making the subgroup-matrix reduction count and scratch drain observable."
1254
  },
@@ -1267,7 +1328,7 @@
1267
  "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1268
  },
1269
  {
1270
- "name": "aligned_f16_batched_plain_2x64x32x64_offset_scale_lock",
1271
  "provenance": {
1272
  "notes": "Two batches carry distinct offset operands, keeping outputs proportional to K and making both the batch stride and aligned subgroup-matrix reduction count observable."
1273
  },
@@ -1286,7 +1347,7 @@
1286
  "outputs": { "Y": { "dtype": "float16", "shape": [2, 64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1287
  },
1288
  {
1289
- "name": "subgroup_matrix_m_tail_57_partial_block_f16_offset_scale_lock",
1290
  "provenance": {
1291
  "notes": "M=57 leaves 25 rows after one full 32-row tile. Offset operands require tail rows to match the full rows' expected magnitude, exposing a short reduction or stale scratch value."
1292
  },
@@ -1305,7 +1366,7 @@
1305
  "outputs": { "Y": { "dtype": "float16", "shape": [57, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1306
  },
1307
  {
1308
- "name": "aligned_f16_transA_64x32_offset_scale_lock",
1309
  "provenance": {
1310
  "notes": "With only A transposed, offset operands keep each output proportional to K and expose both a transposed-A stride error and an incorrect reduction count."
1311
  },
@@ -1325,7 +1386,7 @@
1325
  "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1326
  },
1327
  {
1328
- "name": "transA_transB_subgroup_matrix_m_tail_50_f16_offset_scale_lock",
1329
  "provenance": {
1330
  "notes": "Both operands are transposed and M=50 leaves an 18-row tail. Offset operands make the guarded tail rows' magnitude and placement independently observable."
1331
  },
@@ -1345,7 +1406,7 @@
1345
  "outputs": { "Y": { "dtype": "float16", "shape": [50, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1346
  },
1347
  {
1348
- "name": "f16_rank3_by_broadcast_rank3_offset_scale_lock",
1349
  "provenance": {
1350
  "notes": "f16_rank3_by_broadcast_rank3 cancels to 0.076 under a 0.02 absolute tolerance (26% blind). Offsetting both operands makes each element ~K * aOffset * bOffset over K=3, so the shared single-batch B - read by both output batches - is pinned for value as well as for broadcast addressing."
1351
  },
@@ -1470,7 +1531,7 @@
1470
  {
1471
  "name": "band_vec4_alpha_scaled_m8_k256_n512",
1472
  "provenance": {
1473
- "notes": "Route lock for the few-row band without a K split: eight rows over four 128-column groups, alpha folded into the vec4 store."
1474
  },
1475
  "attrs": { "alpha": 0.5 },
1476
  "inputs": {
@@ -1544,7 +1605,7 @@
1544
  {
1545
  "name": "f16_rank4_by_rank2_shared_weight_m33_k34_n66",
1546
  "provenance": {
1547
- "notes": "Shared rank-2 B across both rank-4 batch axes; covers subgroup-matrix admission, partial tiles and its small-M fallback without changing accumulation precision."
1548
  },
1549
  "attrs": { "alpha": 0.5 },
1550
  "inputs": {
@@ -1564,7 +1625,7 @@
1564
  {
1565
  "name": "f16_rank4_by_rank2_shared_weight_m8_k32_n64",
1566
  "provenance": {
1567
- "notes": "Shared rank-2 B across both rank-4 batch axes; covers subgroup-matrix admission, partial tiles and its small-M fallback without changing accumulation precision."
1568
  },
1569
  "attrs": { "alpha": 0.5 },
1570
  "inputs": {
@@ -1584,7 +1645,7 @@
1584
  {
1585
  "name": "f16_rank4_by_rank2_shared_weight_m1_k32_n64",
1586
  "provenance": {
1587
- "notes": "Shared rank-2 B across both rank-4 batch axes; covers subgroup-matrix admission, partial tiles and its small-M fallback without changing accumulation precision."
1588
  },
1589
  "attrs": { "alpha": 0.5 },
1590
  "inputs": {
@@ -1652,7 +1713,7 @@
1652
  },
1653
  "outputs": { "Y": { "dtype": "float16", "shape": [4, 128, 512], "tolerance": 0.001, "relTolerance": 0.001 } },
1654
  "provenance": {
1655
- "notes": "Portable register tile transBatchB: physical B[K,batch,N], natural tile count reaches the shared occupancy gate; negative alpha and positive inputs avoid cancellation."
1656
  }
1657
  },
1658
  {
@@ -1672,7 +1733,7 @@
1672
  },
1673
  "outputs": { "Y": { "dtype": "float16", "shape": [3, 129, 513], "tolerance": 0.001, "relTolerance": 0.001 } },
1674
  "provenance": {
1675
- "notes": "Portable register tile transBatchB: physical B[K,batch,N], natural tile count reaches the shared occupancy gate; negative alpha and positive inputs avoid cancellation."
1676
  }
1677
  },
1678
  {
@@ -1692,7 +1753,7 @@
1692
  },
1693
  "outputs": { "Y": { "dtype": "float32", "shape": [4, 128, 512], "tolerance": 0.0001, "relTolerance": 0.00002 } },
1694
  "provenance": {
1695
- "notes": "Portable register tile transBatchB: physical B[K,batch,N], natural tile count reaches the shared occupancy gate; negative alpha and positive inputs avoid cancellation."
1696
  }
1697
  },
1698
  {
@@ -1712,7 +1773,7 @@
1712
  },
1713
  "outputs": { "Y": { "dtype": "float32", "shape": [3, 129, 513], "tolerance": 0.0001, "relTolerance": 0.00002 } },
1714
  "provenance": {
1715
- "notes": "Portable register tile transBatchB: physical B[K,batch,N], natural tile count reaches the shared occupancy gate; negative alpha and positive inputs avoid cancellation."
1716
  }
1717
  },
1718
  {
@@ -1732,7 +1793,7 @@
1732
  },
1733
  "outputs": { "Y": { "dtype": "float32", "shape": [4, 4096], "tolerance": 0.0001, "relTolerance": 0.00001 } },
1734
  "provenance": {
1735
- "notes": "Natural f32 preferred single-band geometry with nonuniform operands and alpha0.5. N4096 keeps the existing single-band path eligible at deep K; N2048 would select the split-band family."
1736
  }
1737
  },
1738
  {
@@ -1752,8 +1813,612 @@
1752
  },
1753
  "outputs": { "Y": { "dtype": "float32", "shape": [16, 4096], "tolerance": 0.0001, "relTolerance": 0.00001 } },
1754
  "provenance": {
1755
- "notes": "Natural f32 preferred single-band geometry with nonuniform operands and alpha0.5. N4096 keeps the existing single-band path eligible at deep K; N2048 would select the split-band family."
1756
  }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1757
  }
1758
  ]
1759
  }
 
7
  "ort_float32_trans_batch_b_input_B": [1, 0, 1, 2, 0, 1, -1, 0, 1, 1, 0, 1, 2, -1, 1, 1]
8
  },
9
  "cases": [
10
+ {
11
+ "name": "f16_plain_register_tile_m512_k64_n512",
12
+ "provenance": {
13
+ "notes": "A compact rank-two float16 matrix product with offset operands exposes dropped reduction terms across output tiles."
14
+ },
15
+ "inputs": {
16
+ "A": {
17
+ "dtype": "float16",
18
+ "shape": [512, 64],
19
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.05, "offset": 0.2 }
20
+ },
21
+ "B": {
22
+ "dtype": "float16",
23
+ "shape": [64, 512],
24
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
25
+ }
26
+ },
27
+ "outputs": { "Y": { "dtype": "float16", "shape": [512, 512], "tolerance": 0.01, "relTolerance": 0.001 } }
28
+ },
29
  {
30
  "name": "ort_float32_broadcast_rank4_by_rank3",
31
  "provenance": {
 
712
  "provenance": {
713
  "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
714
  "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
715
+ "notes": "M=32, K=32, N=64 with finite subnormal dot products that must not flush to zero."
716
  },
717
  "inputs": {
718
  "A": { "dtype": "float32", "shape": [32, 32], "data": { "kind": "constant", "value": 1e-20 } },
 
955
  {
956
  "name": "f32_decode_gemv_m1_k65_n68_vec4_compact",
957
  "provenance": {
958
+ "notes": "A compact M=1 float32 matrix product with odd K and N=68 checks the reduction and final output-column tail."
959
  },
960
  "attrs": { "alpha": 1 },
961
  "inputs": {
 
972
  },
973
  "outputs": { "Y": { "dtype": "float32", "shape": [1, 68], "tolerance": 0.0002 } }
974
  },
975
+ {
976
+ "name": "f16_decode_gemv_m1_k65_n68_vec4_compact",
977
+ "provenance": {
978
+ "notes": "A float16 M=1 matrix product with odd K and N=68 checks both reduction and output-column tails. The dot product accumulates in float32 and narrows only at the store."
979
+ },
980
+ "attrs": { "alpha": 1 },
981
+ "inputs": {
982
+ "A": {
983
+ "dtype": "float16",
984
+ "shape": [1, 65],
985
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.2 }
986
+ },
987
+ "B": {
988
+ "dtype": "float16",
989
+ "shape": [65, 68],
990
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.2 }
991
+ }
992
+ },
993
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 68], "tolerance": 0.004, "relTolerance": 0.004 } }
994
+ },
995
+ {
996
+ "name": "f16_decode_gemv_m1_k64_n128_alpha_half",
997
+ "provenance": {
998
+ "notes": "A float16 M=1 matrix product with even K and non-unit alpha checks complete reduction. The dot product accumulates in float32 and narrows only at the store."
999
+ },
1000
+ "attrs": { "alpha": 0.5 },
1001
+ "inputs": {
1002
+ "A": {
1003
+ "dtype": "float16",
1004
+ "shape": [1, 64],
1005
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.2 }
1006
+ },
1007
+ "B": {
1008
+ "dtype": "float16",
1009
+ "shape": [64, 128],
1010
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.2 }
1011
+ }
1012
+ },
1013
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 128], "tolerance": 0.0002, "relTolerance": 0.004 } }
1014
+ },
1015
  {
1016
  "name": "f32_rank4_by_rank2_shared_weight_compact",
1017
  "provenance": {
1018
+ "notes": "A rank-2 weight shared across a 2x3 batch multiplies M=5, K=7, N=9 (alpha=0.5); the odd dimensions check batch-offset indexing and non-power-of-two tails."
1019
  },
1020
  "attrs": { "alpha": 0.5 },
1021
  "inputs": {
 
1034
  },
1035
  {
1036
  "name": "subgroup_matrix_kn_tail_f16_compact",
1037
+ "provenance": {
1038
+ "notes": "M=33 and K=34 are one and two past a 32-element boundary and N=66 is two past a 64-element boundary (float16, alpha=0.5), checking partial tiles in all three dimensions."
1039
+ },
1040
  "attrs": { "alpha": 0.5 },
1041
  "inputs": {
1042
  "A": {
 
1055
  {
1056
  "name": "subgroup_matrix_broadcast_rank4x3_f16_compact",
1057
  "provenance": {
1058
+ "notes": "A rank-4 [1,2,33,32] by rank-3 [2,32,64] float16 broadcast multiplies M=33, K=32, N=64 across a batch of 2, checking batch-dimension broadcasting at a compact scale."
1059
  },
1060
  "attrs": { "alpha": 1 },
1061
  "inputs": {
 
1075
  {
1076
  "name": "broadcast_rank4_tiled_reg_f16_compact",
1077
  "provenance": {
1078
+ "notes": "Compact rank-4 by rank-3 float16 broadcast with odd M/K/N checks output and reduction tails."
1079
  },
1080
  "attrs": { "alpha": 0.5 },
1081
  "inputs": {
 
1115
  {
1116
  "name": "transbatch_a_dense_m_tail_f16_compact",
1117
  "provenance": {
1118
+ "notes": "A stored [M=65, batch=2, K=32] with transBatchA (float16, alpha=0.5): M=65 is one row past a 64-element boundary (K=32,N=64), checking the partial final row across both batches under this transposed-batch addressing."
1119
  },
1120
  "attrs": { "alpha": 0.5, "transBatchA": 1 },
1121
  "inputs": {
 
1170
  "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 2, 2, 4], "tolerance": 0.000001 } }
1171
  },
1172
  {
1173
+ "name": "subgroup_matrix_kn_tail_f16_offset_alpha_scale",
1174
  "provenance": {
1175
  "notes": "Offset float16 operands keep each output near `alpha * K * aOffset * bOffset` (about 3.4), making the K=34 reduction tail, N=66 column tail, and `alpha = 0.5` epilogue observable on the subgroup-matrix route."
1176
  },
 
1190
  "outputs": { "Y": { "dtype": "float16", "shape": [33, 66], "tolerance": 0.03, "relTolerance": 0.01 } }
1191
  },
1192
  {
1193
+ "name": "subgroup_matrix_broadcast_rank4x3_f16_offset_scale",
1194
  "provenance": {
1195
  "notes": "Offset operands keep outputs near `K * aOffset * bOffset`, making the per-batch B slice and K=32 contraction observable in a rank-4 by rank-3 broadcast."
1196
  },
 
1230
  "outputs": { "Y": { "dtype": "float16", "shape": [1, 2, 33, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1231
  },
1232
  {
1233
+ "name": "broadcast_rank4_tiled_reg_f16_offset_alpha_scale",
1234
  "provenance": {
1235
  "notes": "Offset operands keep outputs near `alpha * K * aOffset * bOffset` with K=33, making the one-element reduction tail and `alpha` multiplier observable on the register-blocked rank-4 route."
1236
  },
 
1250
  "outputs": { "Y": { "dtype": "float16", "shape": [1, 2, 65, 67], "tolerance": 0.03, "relTolerance": 0.01 } }
1251
  },
1252
  {
1253
+ "name": "transbatch_a_dense_m_tail_f16_offset_alpha_scale",
1254
  "provenance": {
1255
  "notes": "Offset operands make each output proportional to `alpha * K`, exposing the [M, batch, K] transBatchA stride, K=32 contraction, and `alpha = 0.5` scale."
1256
  },
 
1270
  "outputs": { "Y": { "dtype": "float16", "shape": [2, 65, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1271
  },
1272
  {
1273
+ "name": "aligned_f16_transA_transB_alpha_offset_scale",
1274
  "provenance": {
1275
  "notes": "Offset operands keep the doubly transposed output proportional to `alpha * K`, making the K=32 contraction and `alpha = 0.5` epilogue observable on subgroup-matrix and portable tiled routes."
1276
  },
 
1290
  "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1291
  },
1292
  {
1293
+ "name": "f16_unaligned_3x5x7_offset_scale",
1294
  "provenance": {
1295
  "notes": "Offset operands in a 3x5 by 5x7 multiply keep outputs proportional to K, exposing dropped reduction elements or doubled tails on the unaligned scalar and tiled routes."
1296
  },
 
1309
  "outputs": { "Y": { "dtype": "float16", "shape": [3, 7], "tolerance": 0.01, "relTolerance": 0.005 } }
1310
  },
1311
  {
1312
+ "name": "aligned_f16_plain_64x32x64_offset_scale",
1313
  "provenance": {
1314
  "notes": "Offset operands keep each fully aligned M=64, K=32, N=64 output proportional to K, making the subgroup-matrix reduction count and scratch drain observable."
1315
  },
 
1328
  "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1329
  },
1330
  {
1331
+ "name": "aligned_f16_batched_plain_2x64x32x64_offset_scale",
1332
  "provenance": {
1333
  "notes": "Two batches carry distinct offset operands, keeping outputs proportional to K and making both the batch stride and aligned subgroup-matrix reduction count observable."
1334
  },
 
1347
  "outputs": { "Y": { "dtype": "float16", "shape": [2, 64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1348
  },
1349
  {
1350
+ "name": "subgroup_matrix_m_tail_57_partial_block_f16_offset_scale",
1351
  "provenance": {
1352
  "notes": "M=57 leaves 25 rows after one full 32-row tile. Offset operands require tail rows to match the full rows' expected magnitude, exposing a short reduction or stale scratch value."
1353
  },
 
1366
  "outputs": { "Y": { "dtype": "float16", "shape": [57, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1367
  },
1368
  {
1369
+ "name": "aligned_f16_transA_64x32_offset_scale",
1370
  "provenance": {
1371
  "notes": "With only A transposed, offset operands keep each output proportional to K and expose both a transposed-A stride error and an incorrect reduction count."
1372
  },
 
1386
  "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1387
  },
1388
  {
1389
+ "name": "transA_transB_subgroup_matrix_m_tail_50_f16_offset_scale",
1390
  "provenance": {
1391
  "notes": "Both operands are transposed and M=50 leaves an 18-row tail. Offset operands make the guarded tail rows' magnitude and placement independently observable."
1392
  },
 
1406
  "outputs": { "Y": { "dtype": "float16", "shape": [50, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1407
  },
1408
  {
1409
+ "name": "f16_rank3_by_broadcast_rank3_offset_scale",
1410
  "provenance": {
1411
  "notes": "f16_rank3_by_broadcast_rank3 cancels to 0.076 under a 0.02 absolute tolerance (26% blind). Offsetting both operands makes each element ~K * aOffset * bOffset over K=3, so the shared single-batch B - read by both output batches - is pinned for value as well as for broadcast addressing."
1412
  },
 
1531
  {
1532
  "name": "band_vec4_alpha_scaled_m8_k256_n512",
1533
  "provenance": {
1534
+ "notes": "Eight rows over 512 output columns with non-unit alpha check independent matrix-product results."
1535
  },
1536
  "attrs": { "alpha": 0.5 },
1537
  "inputs": {
 
1605
  {
1606
  "name": "f16_rank4_by_rank2_shared_weight_m33_k34_n66",
1607
  "provenance": {
1608
+ "notes": "A [2, 3, 33, 34] float16 A shares one [34, 66] B across both batch axes, with alpha 0.5; none of M=33, K=34 or N=66 is a multiple of 32."
1609
  },
1610
  "attrs": { "alpha": 0.5 },
1611
  "inputs": {
 
1625
  {
1626
  "name": "f16_rank4_by_rank2_shared_weight_m8_k32_n64",
1627
  "provenance": {
1628
+ "notes": "A [2, 3, 8, 32] float16 A shares one [32, 64] B across both batch axes, with alpha 0.5, at an 8-row M."
1629
  },
1630
  "attrs": { "alpha": 0.5 },
1631
  "inputs": {
 
1645
  {
1646
  "name": "f16_rank4_by_rank2_shared_weight_m1_k32_n64",
1647
  "provenance": {
1648
+ "notes": "A [2, 3, 1, 32] float16 A shares one [32, 64] B across both batch axes, with alpha 0.5, at a single-row M."
1649
  },
1650
  "attrs": { "alpha": 0.5 },
1651
  "inputs": {
 
1713
  },
1714
  "outputs": { "Y": { "dtype": "float16", "shape": [4, 128, 512], "tolerance": 0.001, "relTolerance": 0.001 } },
1715
  "provenance": {
1716
+ "notes": "Physical B stored [K=128, batch=4, N=512] with transBatchB (float16, alpha=-0.5, negative to avoid cancellation with positive inputs): M=128, K=128, N=512 are all tile-aligned, checking this transposed-batch addressing at a fully aligned shape."
1717
  }
1718
  },
1719
  {
 
1733
  },
1734
  "outputs": { "Y": { "dtype": "float16", "shape": [3, 129, 513], "tolerance": 0.001, "relTolerance": 0.001 } },
1735
  "provenance": {
1736
+ "notes": "Physical B stored [K=131, batch=3, N=513] with transBatchB (float16, alpha=-0.5): M=129, K=131, N=513 are each one to three past a 128/512 boundary, checking partial final tiles under transposed-batch addressing."
1737
  }
1738
  },
1739
  {
 
1753
  },
1754
  "outputs": { "Y": { "dtype": "float32", "shape": [4, 128, 512], "tolerance": 0.0001, "relTolerance": 0.00002 } },
1755
  "provenance": {
1756
+ "notes": "Physical B stored [K=128, batch=4, N=512] with transBatchB (float32, alpha=-0.5, negative to avoid cancellation with positive inputs): M=128, K=128, N=512 are all tile-aligned, checking this transposed-batch addressing at a fully aligned shape."
1757
  }
1758
  },
1759
  {
 
1773
  },
1774
  "outputs": { "Y": { "dtype": "float32", "shape": [3, 129, 513], "tolerance": 0.0001, "relTolerance": 0.00002 } },
1775
  "provenance": {
1776
+ "notes": "Physical B stored [K=131, batch=3, N=513] with transBatchB (float32, alpha=-0.5): M=129, K=131, N=513 are each one to three past a 128/512 boundary, checking partial final tiles under transposed-batch addressing."
1777
  }
1778
  },
1779
  {
 
1793
  },
1794
  "outputs": { "Y": { "dtype": "float32", "shape": [4, 4096], "tolerance": 0.0001, "relTolerance": 0.00001 } },
1795
  "provenance": {
1796
+ "notes": "M=4, K=2048, N=4096 float32 operands with alpha=0.5 and different per-operand offsets (0.02 vs 0.03) avoid cancellation in the matrix product."
1797
  }
1798
  },
1799
  {
 
1813
  },
1814
  "outputs": { "Y": { "dtype": "float32", "shape": [16, 4096], "tolerance": 0.0001, "relTolerance": 0.00001 } },
1815
  "provenance": {
1816
+ "notes": "M=16, K=2560, N=4096 float32 operands with alpha=0.5 and different per-operand offsets (0.02 vs 0.03) avoid cancellation in the matrix product."
1817
  }
1818
+ },
1819
+ {
1820
+ "name": "broadcast_transb_rank4x3_aligned_float16",
1821
+ "attrs": { "transB": 2, "alpha": -0.5 },
1822
+ "inputs": {
1823
+ "A": {
1824
+ "dtype": "float16",
1825
+ "shape": [2, 3, 64, 32],
1826
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
1827
+ },
1828
+ "B": {
1829
+ "dtype": "float16",
1830
+ "shape": [3, 64, 32],
1831
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
1832
+ }
1833
+ },
1834
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 64, 64], "tolerance": 0.0005, "relTolerance": 0.001 } },
1835
+ "provenance": {
1836
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
1837
+ }
1838
+ },
1839
+ {
1840
+ "name": "broadcast_transb_rank4x3_tails_float16",
1841
+ "attrs": { "transB": 2, "alpha": -0.5 },
1842
+ "inputs": {
1843
+ "A": {
1844
+ "dtype": "float16",
1845
+ "shape": [2, 3, 65, 33],
1846
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
1847
+ },
1848
+ "B": {
1849
+ "dtype": "float16",
1850
+ "shape": [3, 67, 33],
1851
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
1852
+ }
1853
+ },
1854
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 65, 67], "tolerance": 0.0005, "relTolerance": 0.001 } },
1855
+ "provenance": {
1856
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
1857
+ }
1858
+ },
1859
+ {
1860
+ "name": "broadcast_transb_broadcast_a_float16",
1861
+ "attrs": { "transB": 2, "alpha": -0.5 },
1862
+ "inputs": {
1863
+ "A": {
1864
+ "dtype": "float16",
1865
+ "shape": [2, 1, 65, 33],
1866
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
1867
+ },
1868
+ "B": {
1869
+ "dtype": "float16",
1870
+ "shape": [3, 67, 33],
1871
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
1872
+ }
1873
+ },
1874
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 65, 67], "tolerance": 0.0005, "relTolerance": 0.001 } },
1875
+ "provenance": {
1876
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
1877
+ }
1878
+ },
1879
+ {
1880
+ "name": "broadcast_transb_broadcast_b_float16",
1881
+ "attrs": { "transB": 2, "alpha": -0.5 },
1882
+ "inputs": {
1883
+ "A": {
1884
+ "dtype": "float16",
1885
+ "shape": [2, 3, 65, 33],
1886
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
1887
+ },
1888
+ "B": {
1889
+ "dtype": "float16",
1890
+ "shape": [1, 67, 33],
1891
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
1892
+ }
1893
+ },
1894
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 65, 67], "tolerance": 0.0005, "relTolerance": 0.001 } },
1895
+ "provenance": {
1896
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
1897
+ }
1898
+ },
1899
+ {
1900
+ "name": "broadcast_transb_rank3x2_float16",
1901
+ "attrs": { "transB": 2, "alpha": -0.5 },
1902
+ "inputs": {
1903
+ "A": {
1904
+ "dtype": "float16",
1905
+ "shape": [3, 64, 32],
1906
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
1907
+ },
1908
+ "B": {
1909
+ "dtype": "float16",
1910
+ "shape": [64, 32],
1911
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
1912
+ }
1913
+ },
1914
+ "outputs": { "Y": { "dtype": "float16", "shape": [3, 64, 64], "tolerance": 0.0005, "relTolerance": 0.001 } },
1915
+ "provenance": {
1916
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
1917
+ }
1918
+ },
1919
+ {
1920
+ "name": "broadcast_transb_rank5x3_float16",
1921
+ "attrs": { "transB": 2, "alpha": -0.5 },
1922
+ "inputs": {
1923
+ "A": {
1924
+ "dtype": "float16",
1925
+ "shape": [2, 1, 3, 65, 33],
1926
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
1927
+ },
1928
+ "B": {
1929
+ "dtype": "float16",
1930
+ "shape": [3, 67, 33],
1931
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
1932
+ }
1933
+ },
1934
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 1, 3, 65, 67], "tolerance": 0.0005, "relTolerance": 0.001 } },
1935
+ "provenance": {
1936
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
1937
+ }
1938
+ },
1939
+ {
1940
+ "name": "broadcast_transb_m_below_tile_float16",
1941
+ "attrs": { "transB": 2, "alpha": -0.5 },
1942
+ "inputs": {
1943
+ "A": {
1944
+ "dtype": "float16",
1945
+ "shape": [2, 3, 63, 32],
1946
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
1947
+ },
1948
+ "B": {
1949
+ "dtype": "float16",
1950
+ "shape": [3, 64, 32],
1951
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
1952
+ }
1953
+ },
1954
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 63, 64], "tolerance": 0.0005, "relTolerance": 0.001 } },
1955
+ "provenance": {
1956
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
1957
+ }
1958
+ },
1959
+ {
1960
+ "name": "broadcast_transb_k_below_tile_float16",
1961
+ "attrs": { "transB": 2, "alpha": -0.5 },
1962
+ "inputs": {
1963
+ "A": {
1964
+ "dtype": "float16",
1965
+ "shape": [2, 3, 64, 31],
1966
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
1967
+ },
1968
+ "B": {
1969
+ "dtype": "float16",
1970
+ "shape": [3, 64, 31],
1971
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
1972
+ }
1973
+ },
1974
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 64, 64], "tolerance": 0.0005, "relTolerance": 0.001 } },
1975
+ "provenance": {
1976
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
1977
+ }
1978
+ },
1979
+ {
1980
+ "name": "broadcast_transb_n_below_tile_float16",
1981
+ "attrs": { "transB": 2, "alpha": -0.5 },
1982
+ "inputs": {
1983
+ "A": {
1984
+ "dtype": "float16",
1985
+ "shape": [2, 3, 64, 32],
1986
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
1987
+ },
1988
+ "B": {
1989
+ "dtype": "float16",
1990
+ "shape": [3, 63, 32],
1991
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
1992
+ }
1993
+ },
1994
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 64, 63], "tolerance": 0.0005, "relTolerance": 0.001 } },
1995
+ "provenance": {
1996
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
1997
+ }
1998
+ },
1999
+ {
2000
+ "name": "broadcast_transb_last_A_NaN_float16",
2001
+ "attrs": { "transB": 2, "alpha": -0.5 },
2002
+ "inputs": {
2003
+ "A": {
2004
+ "dtype": "float16",
2005
+ "shape": [2, 3, 65, 33],
2006
+ "data": {
2007
+ "kind": "fillFloat32",
2008
+ "sinStep": 0.013,
2009
+ "cosStep": 0.029,
2010
+ "scale": 0.05,
2011
+ "offset": 0.15,
2012
+ "nanStart": 12869.0,
2013
+ "nanCount": 1.0
2014
+ }
2015
+ },
2016
+ "B": {
2017
+ "dtype": "float16",
2018
+ "shape": [3, 67, 33],
2019
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
2020
+ }
2021
+ },
2022
+ "outputs": {
2023
+ "Y": {
2024
+ "dtype": "float16",
2025
+ "shape": [2, 3, 65, 67],
2026
+ "tolerance": 0.0005,
2027
+ "relTolerance": 0.001,
2028
+ "allowNaN": true
2029
+ }
2030
+ },
2031
+ "provenance": {
2032
+ "notes": "NaN at the final physical element of a tailed transposed-B broadcast product; checks propagation stays confined to the rows/columns reached by that operand."
2033
+ }
2034
+ },
2035
+ {
2036
+ "name": "broadcast_transb_last_B_NaN_float16",
2037
+ "attrs": { "transB": 2, "alpha": -0.5 },
2038
+ "inputs": {
2039
+ "A": {
2040
+ "dtype": "float16",
2041
+ "shape": [2, 3, 65, 33],
2042
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
2043
+ },
2044
+ "B": {
2045
+ "dtype": "float16",
2046
+ "shape": [3, 67, 33],
2047
+ "data": {
2048
+ "kind": "fillFloat32",
2049
+ "sinStep": 0.019,
2050
+ "cosStep": 0.007,
2051
+ "scale": 0.07,
2052
+ "offset": -0.1,
2053
+ "nanStart": 6632.0,
2054
+ "nanCount": 1.0
2055
+ }
2056
+ }
2057
+ },
2058
+ "outputs": {
2059
+ "Y": {
2060
+ "dtype": "float16",
2061
+ "shape": [2, 3, 65, 67],
2062
+ "tolerance": 0.0005,
2063
+ "relTolerance": 0.001,
2064
+ "allowNaN": true
2065
+ }
2066
+ },
2067
+ "provenance": {
2068
+ "notes": "NaN at the final physical element of a tailed transposed-B broadcast product; checks propagation stays confined to the rows/columns reached by that operand."
2069
+ }
2070
+ },
2071
+ {
2072
+ "name": "broadcast_transb_rank4x3_aligned_float32",
2073
+ "attrs": { "transB": 2, "alpha": -0.5 },
2074
+ "inputs": {
2075
+ "A": {
2076
+ "dtype": "float32",
2077
+ "shape": [2, 3, 64, 32],
2078
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
2079
+ },
2080
+ "B": {
2081
+ "dtype": "float32",
2082
+ "shape": [3, 64, 32],
2083
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
2084
+ }
2085
+ },
2086
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 64, 64], "tolerance": 0.00001, "relTolerance": 0.0001 } },
2087
+ "provenance": {
2088
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
2089
+ }
2090
+ },
2091
+ {
2092
+ "name": "broadcast_transb_rank4x3_tails_float32",
2093
+ "attrs": { "transB": 2, "alpha": -0.5 },
2094
+ "inputs": {
2095
+ "A": {
2096
+ "dtype": "float32",
2097
+ "shape": [2, 3, 65, 33],
2098
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
2099
+ },
2100
+ "B": {
2101
+ "dtype": "float32",
2102
+ "shape": [3, 67, 33],
2103
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
2104
+ }
2105
+ },
2106
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 65, 67], "tolerance": 0.00001, "relTolerance": 0.0001 } },
2107
+ "provenance": {
2108
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
2109
+ }
2110
+ },
2111
+ {
2112
+ "name": "broadcast_transb_broadcast_a_float32",
2113
+ "attrs": { "transB": 2, "alpha": -0.5 },
2114
+ "inputs": {
2115
+ "A": {
2116
+ "dtype": "float32",
2117
+ "shape": [2, 1, 65, 33],
2118
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
2119
+ },
2120
+ "B": {
2121
+ "dtype": "float32",
2122
+ "shape": [3, 67, 33],
2123
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
2124
+ }
2125
+ },
2126
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 65, 67], "tolerance": 0.00001, "relTolerance": 0.0001 } },
2127
+ "provenance": {
2128
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
2129
+ }
2130
+ },
2131
+ {
2132
+ "name": "broadcast_transb_broadcast_b_float32",
2133
+ "attrs": { "transB": 2, "alpha": -0.5 },
2134
+ "inputs": {
2135
+ "A": {
2136
+ "dtype": "float32",
2137
+ "shape": [2, 3, 65, 33],
2138
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
2139
+ },
2140
+ "B": {
2141
+ "dtype": "float32",
2142
+ "shape": [1, 67, 33],
2143
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
2144
+ }
2145
+ },
2146
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 65, 67], "tolerance": 0.00001, "relTolerance": 0.0001 } },
2147
+ "provenance": {
2148
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
2149
+ }
2150
+ },
2151
+ {
2152
+ "name": "broadcast_transb_rank3x2_float32",
2153
+ "attrs": { "transB": 2, "alpha": -0.5 },
2154
+ "inputs": {
2155
+ "A": {
2156
+ "dtype": "float32",
2157
+ "shape": [3, 64, 32],
2158
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
2159
+ },
2160
+ "B": {
2161
+ "dtype": "float32",
2162
+ "shape": [64, 32],
2163
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
2164
+ }
2165
+ },
2166
+ "outputs": { "Y": { "dtype": "float32", "shape": [3, 64, 64], "tolerance": 0.00001, "relTolerance": 0.0001 } },
2167
+ "provenance": {
2168
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
2169
+ }
2170
+ },
2171
+ {
2172
+ "name": "broadcast_transb_rank5x3_float32",
2173
+ "attrs": { "transB": 2, "alpha": -0.5 },
2174
+ "inputs": {
2175
+ "A": {
2176
+ "dtype": "float32",
2177
+ "shape": [2, 1, 3, 65, 33],
2178
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
2179
+ },
2180
+ "B": {
2181
+ "dtype": "float32",
2182
+ "shape": [3, 67, 33],
2183
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
2184
+ }
2185
+ },
2186
+ "outputs": {
2187
+ "Y": { "dtype": "float32", "shape": [2, 1, 3, 65, 67], "tolerance": 0.00001, "relTolerance": 0.0001 }
2188
+ },
2189
+ "provenance": {
2190
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
2191
+ }
2192
+ },
2193
+ {
2194
+ "name": "broadcast_transb_m_below_tile_float32",
2195
+ "attrs": { "transB": 2, "alpha": -0.5 },
2196
+ "inputs": {
2197
+ "A": {
2198
+ "dtype": "float32",
2199
+ "shape": [2, 3, 63, 32],
2200
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
2201
+ },
2202
+ "B": {
2203
+ "dtype": "float32",
2204
+ "shape": [3, 64, 32],
2205
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
2206
+ }
2207
+ },
2208
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 63, 64], "tolerance": 0.00001, "relTolerance": 0.0001 } },
2209
+ "provenance": {
2210
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
2211
+ }
2212
+ },
2213
+ {
2214
+ "name": "broadcast_transb_k_below_tile_float32",
2215
+ "attrs": { "transB": 2, "alpha": -0.5 },
2216
+ "inputs": {
2217
+ "A": {
2218
+ "dtype": "float32",
2219
+ "shape": [2, 3, 64, 31],
2220
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
2221
+ },
2222
+ "B": {
2223
+ "dtype": "float32",
2224
+ "shape": [3, 64, 31],
2225
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
2226
+ }
2227
+ },
2228
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 64, 64], "tolerance": 0.00001, "relTolerance": 0.0001 } },
2229
+ "provenance": {
2230
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
2231
+ }
2232
+ },
2233
+ {
2234
+ "name": "broadcast_transb_n_below_tile_float32",
2235
+ "attrs": { "transB": 2, "alpha": -0.5 },
2236
+ "inputs": {
2237
+ "A": {
2238
+ "dtype": "float32",
2239
+ "shape": [2, 3, 64, 32],
2240
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
2241
+ },
2242
+ "B": {
2243
+ "dtype": "float32",
2244
+ "shape": [3, 63, 32],
2245
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
2246
+ }
2247
+ },
2248
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 64, 63], "tolerance": 0.00001, "relTolerance": 0.0001 } },
2249
+ "provenance": {
2250
+ "notes": "Transposed-B broadcast addressing: rank, singleton-axis, and tile-boundary coverage; negative alpha and non-unit true transB."
2251
+ }
2252
+ },
2253
+ {
2254
+ "name": "broadcast_transb_last_A_NaN_float32",
2255
+ "attrs": { "transB": 2, "alpha": -0.5 },
2256
+ "inputs": {
2257
+ "A": {
2258
+ "dtype": "float32",
2259
+ "shape": [2, 3, 65, 33],
2260
+ "data": {
2261
+ "kind": "fillFloat32",
2262
+ "sinStep": 0.013,
2263
+ "cosStep": 0.029,
2264
+ "scale": 0.05,
2265
+ "offset": 0.15,
2266
+ "nanStart": 12869.0,
2267
+ "nanCount": 1.0
2268
+ }
2269
+ },
2270
+ "B": {
2271
+ "dtype": "float32",
2272
+ "shape": [3, 67, 33],
2273
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.07, "offset": -0.1 }
2274
+ }
2275
+ },
2276
+ "outputs": {
2277
+ "Y": {
2278
+ "dtype": "float32",
2279
+ "shape": [2, 3, 65, 67],
2280
+ "tolerance": 0.00001,
2281
+ "relTolerance": 0.0001,
2282
+ "allowNaN": true
2283
+ }
2284
+ },
2285
+ "provenance": {
2286
+ "notes": "NaN at the final physical element of a tailed transposed-B broadcast product; checks propagation stays confined to the rows/columns reached by that operand."
2287
+ }
2288
+ },
2289
+ {
2290
+ "name": "broadcast_transb_last_B_NaN_float32",
2291
+ "attrs": { "transB": 2, "alpha": -0.5 },
2292
+ "inputs": {
2293
+ "A": {
2294
+ "dtype": "float32",
2295
+ "shape": [2, 3, 65, 33],
2296
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.05, "offset": 0.15 }
2297
+ },
2298
+ "B": {
2299
+ "dtype": "float32",
2300
+ "shape": [3, 67, 33],
2301
+ "data": {
2302
+ "kind": "fillFloat32",
2303
+ "sinStep": 0.019,
2304
+ "cosStep": 0.007,
2305
+ "scale": 0.07,
2306
+ "offset": -0.1,
2307
+ "nanStart": 6632.0,
2308
+ "nanCount": 1.0
2309
+ }
2310
+ }
2311
+ },
2312
+ "outputs": {
2313
+ "Y": {
2314
+ "dtype": "float32",
2315
+ "shape": [2, 3, 65, 67],
2316
+ "tolerance": 0.00001,
2317
+ "relTolerance": 0.0001,
2318
+ "allowNaN": true
2319
+ }
2320
+ },
2321
+ "provenance": {
2322
+ "notes": "NaN at the final physical element of a tailed transposed-B broadcast product; checks propagation stays confined to the rows/columns reached by that operand."
2323
+ }
2324
+ },
2325
+ {
2326
+ "name": "sgmat_n_tail_alpha_f32_64x64x65",
2327
+ "attrs": { "alpha": 0.5 },
2328
+ "provenance": {
2329
+ "notes": "The 64-wide subgroup-matrix column tile computes ceilDiv(N, 64) * 64 columns: its trailing tile reads column N - 1 in place of every column at or past N and drops those stores, so a wrong clamp or a missing store guard shows up as a wrong last column or a write past the output. alpha = 0.5 rides the same guarded store, so a tail that skipped the epilogue would show up as an unscaled trailing column."
2330
+ },
2331
+ "inputs": {
2332
+ "A": {
2333
+ "dtype": "float32",
2334
+ "shape": [64, 64],
2335
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2 }
2336
+ },
2337
+ "B": {
2338
+ "dtype": "float32",
2339
+ "shape": [64, 65],
2340
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2 }
2341
+ }
2342
+ },
2343
+ "outputs": { "Y": { "dtype": "float32", "shape": [64, 65], "tolerance": 0.0001 } }
2344
+ },
2345
+ {
2346
+ "name": "sgmat_n_tail_transB_f16_64x64x65",
2347
+ "attrs": { "transB": 1 },
2348
+ "provenance": {
2349
+ "notes": "A 65-column float16 transposed-B matrix product checks the final partial column tile, especially the last valid output and guarded stores."
2350
+ },
2351
+ "inputs": {
2352
+ "A": {
2353
+ "dtype": "float16",
2354
+ "shape": [64, 64],
2355
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.023, "scale": 0.2 }
2356
+ },
2357
+ "B": {
2358
+ "dtype": "float16",
2359
+ "shape": [65, 64],
2360
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.009, "scale": 0.2 }
2361
+ }
2362
+ },
2363
+ "outputs": { "Y": { "dtype": "float16", "shape": [64, 65], "tolerance": 0.005, "relTolerance": 0.02 } }
2364
+ },
2365
+ {
2366
+ "name": "sgmat_n_padding_ratio_admits_f32_64x64x32",
2367
+ "provenance": {
2368
+ "notes": "M=64, K=64, N=32 float32 operands check the matrix product at a power-of-two column count; its N=31 sibling is one column narrower."
2369
+ },
2370
+ "inputs": {
2371
+ "A": {
2372
+ "dtype": "float32",
2373
+ "shape": [64, 64],
2374
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.2 }
2375
+ },
2376
+ "B": {
2377
+ "dtype": "float32",
2378
+ "shape": [64, 32],
2379
+ "data": { "kind": "fillFloat32", "sinStep": 0.031, "cosStep": 0.017, "scale": 0.2 }
2380
+ }
2381
+ },
2382
+ "outputs": { "Y": { "dtype": "float32", "shape": [64, 32], "tolerance": 0.0001 } }
2383
+ },
2384
+ {
2385
+ "name": "sgmat_n_padding_ratio_rejects_f32_64x64x31",
2386
+ "provenance": {
2387
+ "notes": "M=64, K=64, N=31 float32 operands check the matrix product one column narrower than its N=32 sibling."
2388
+ },
2389
+ "inputs": {
2390
+ "A": {
2391
+ "dtype": "float32",
2392
+ "shape": [64, 64],
2393
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.029, "scale": 0.2 }
2394
+ },
2395
+ "B": {
2396
+ "dtype": "float32",
2397
+ "shape": [64, 31],
2398
+ "data": { "kind": "fillFloat32", "sinStep": 0.031, "cosStep": 0.017, "scale": 0.2 }
2399
+ }
2400
+ },
2401
+ "outputs": { "Y": { "dtype": "float32", "shape": [64, 31], "tolerance": 0.0001 } }
2402
+ },
2403
+ {
2404
+ "name": "plain_rank2_tiled_reg_row_tail_m100_k64_n2048",
2405
+ "provenance": {
2406
+ "notes": "M=100 leaves 36 rows past a 64-row boundary (K=64, N=2048), with alpha=0.5 scaling the output; the extra rows must be included correctly in the result."
2407
+ },
2408
+ "attrs": { "alpha": 0.5 },
2409
+ "inputs": {
2410
+ "A": {
2411
+ "dtype": "float32",
2412
+ "shape": [100, 64],
2413
+ "data": { "kind": "fillFloat32", "sinStep": 0.037, "cosStep": 0.061, "scale": 0.2 }
2414
+ },
2415
+ "B": {
2416
+ "dtype": "float32",
2417
+ "shape": [64, 2048],
2418
+ "data": { "kind": "fillFloat32", "sinStep": 0.043, "cosStep": 0.079, "scale": 0.2 }
2419
+ }
2420
+ },
2421
+ "outputs": { "Y": { "dtype": "float32", "shape": [100, 2048], "tolerance": 0.0001 } }
2422
  }
2423
  ]
2424
  }