Xenova HF Staff commited on
Commit
5a12de3
·
verified ·
1 Parent(s): 7fed370

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -1,3 +1,98 @@
1
  ---
 
2
  license: apache-2.0
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ library_name: kernels
3
  license: apache-2.0
4
+ tags:
5
+ - kernel
6
+ - webgpu
7
+ - wgsl
8
  ---
9
+ # com.microsoft.TransposeMatMul
10
+
11
+ `com.microsoft` · ONNX Runtime contrib operator · contrib since_version 1
12
+
13
+ ## Description
14
+
15
+ Matrix product of two N-dimensional tensors `A` and `B`, following NumPy-style matrix-multiplication broadcasting, with optional transposition of either operand's last two dimensions and a scalar `alpha` multiplier. This is the strict subset of `FusedMatMul` that omits batch-dimension transposition, and it uses the same kernels. Float32 and float16 are supported; double and bfloat16 are not implemented.
16
+
17
+ See the [ONNX Runtime `TransposeMatMul` contrib-operator spec](https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#com.microsoft.TransposeMatMul) for the reference semantics.
18
+
19
+ ## Inputs
20
+
21
+ | Name | Logical dtype | Rank | Shape | Description | Presence |
22
+ | --- | --- | --- | --- | --- | --- |
23
+ | `A` | `T` | — | — | N-dimensional matrix A. | required |
24
+ | `B` | `T` | — | — | N-dimensional matrix B. | required |
25
+
26
+ ## Outputs
27
+
28
+ | Name | Logical dtype | Rank | Shape | Description | Presence |
29
+ | --- | --- | --- | --- | --- | --- |
30
+ | `Y` | `T` | derived | derived | Matrix-multiplication result whose shape follows NumPy-style rules after applying the requested batch and matrix transpositions. | required |
31
+
32
+ ## Attributes
33
+
34
+ Default values (overridable per request):
35
+
36
+ | Attribute | Default | Description |
37
+ | --- | --- | --- |
38
+ | `alpha` | `1` | Scalar multiplier applied to the product of the input tensors. |
39
+ | `transA` | `0` | When non-zero, transposes `A` on its last two dimensions before multiplication. |
40
+ | `transB` | `0` | When non-zero, transposes `B` on its last two dimensions before multiplication. |
41
+
42
+ ## Type constraints
43
+
44
+ | Variable | Allowed dtypes |
45
+ | --- | --- |
46
+ | `T` | `float32`, `float16` |
47
+
48
+ ## Implementation variants
49
+
50
+ One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
51
+
52
+ - `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.
53
+ - `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.
54
+ - `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.
55
+ - `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.
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
+
62
+ ## Device requirements
63
+
64
+ Some implementation variants require `subgroup-matrix` and `subgroups`. These are route-specific capabilities, not package-wide requirements; availability also depends on the request shape and dtype.
65
+
66
+ ## Files
67
+
68
+ - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
69
+ - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
70
+ - [`test.json`](build/webgpu/test.json) — correctness cases
71
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
72
+ - [`fused-matmul-subgroup-matrix.wgsl.jinja`](build/webgpu/fused-matmul-subgroup-matrix.wgsl.jinja)
73
+ - [`matmul-band-vec4.wgsl.jinja`](build/webgpu/matmul-band-vec4.wgsl.jinja)
74
+ - [`matmul-subgroup-matrix-ext.wgsl.jinja`](build/webgpu/matmul-subgroup-matrix-ext.wgsl.jinja)
75
+ - [`matmul-tiled-general-reg.wgsl.jinja`](build/webgpu/matmul-tiled-general-reg.wgsl.jinja)
76
+ - [`matmul-tiled-general.wgsl.jinja`](build/webgpu/matmul-tiled-general.wgsl.jinja)
77
+ - [`matmul-vector-matrix-vec4.wgsl.jinja`](build/webgpu/matmul-vector-matrix-vec4.wgsl.jinja)
78
+ - [`reduce-axis0-splitk-combine.wgsl.jinja`](build/webgpu/reduce-axis0-splitk-combine.wgsl.jinja)
79
+
80
+ ## Use with `@huggingface/kernels`
81
+
82
+ ```sh
83
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
84
+ ```
85
+
86
+ Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
87
+
88
+ The `version: 1` option selects the published kernel contract; it is independent of any operator opset, contrib `since_version`, or model version.
89
+ It follows the `v1` branch as fixes land. To pin exact artifact bytes, pass a 40-character commit `revision` instead of `version`.
90
+
91
+ Replace each `*Data` placeholder with a typed array containing the corresponding input data.
92
+
93
+ ```js
94
+ import { getKernel } from "@huggingface/kernels";
95
+
96
+ const kernel = await getKernel("webgpu-kernels/com.microsoft.TransposeMatMul", { version: 1 });
97
+ const { Y } = await kernel({ A: { data: AData, shape: [3] }, B: { data: BData, shape: [3] } });
98
+ ```
build/webgpu/bench.json ADDED
@@ -0,0 +1,449 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "cases": [
3
+ {
4
+ "name": "transposematmul-f32-plain-512x2048x512",
5
+ "preset": "smoke",
6
+ "attrs": { "alpha": 1 },
7
+ "inputs": {
8
+ "A": { "shape": [512, 2048], "dtype": "float32", "dist": "normal", "seed": 520, "scale": 0.1 },
9
+ "B": { "shape": [2048, 512], "dtype": "float32", "dist": "normal", "seed": 521, "scale": 0.1 }
10
+ },
11
+ "outputs": { "Y": { "shape": [512, 512], "dtype": "float32" } },
12
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * 512 * 512 * 2048" }] }
13
+ },
14
+ {
15
+ "name": "transposematmul-f16-plain-512x2048x512",
16
+ "preset": "smoke",
17
+ "attrs": { "alpha": 1 },
18
+ "inputs": {
19
+ "A": { "shape": [512, 2048], "dtype": "float16", "dist": "normal", "seed": 530, "scale": 0.1 },
20
+ "B": { "shape": [2048, 512], "dtype": "float16", "dist": "normal", "seed": 531, "scale": 0.1 }
21
+ },
22
+ "outputs": { "Y": { "shape": [512, 512], "dtype": "float16" } },
23
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * 512 * 512 * 2048" }] }
24
+ },
25
+ {
26
+ "name": "transposematmul-f32-unaligned-500x2000x500",
27
+ "preset": "smoke",
28
+ "attrs": { "alpha": 1 },
29
+ "inputs": {
30
+ "A": { "shape": [500, 2000], "dtype": "float32", "dist": "normal", "seed": 540, "scale": 0.1 },
31
+ "B": { "shape": [2000, 500], "dtype": "float32", "dist": "normal", "seed": 541, "scale": 0.1 }
32
+ },
33
+ "outputs": { "Y": { "shape": [500, 500], "dtype": "float32" } },
34
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * 500 * 500 * 2000" }] }
35
+ },
36
+ {
37
+ "name": "transposematmul-f16-decode-gemv-m1-1x2048x512",
38
+ "preset": "smoke",
39
+ "attrs": { "alpha": 1 },
40
+ "vars": { "M": 1, "K": 2048, "N": 512 },
41
+ "inputs": {
42
+ "A": { "shape": [1, 2048], "dtype": "float16", "dist": "normal", "seed": 606, "scale": 0.1 },
43
+ "B": { "shape": [2048, 512], "dtype": "float16", "dist": "normal", "seed": 607, "scale": 0.1 }
44
+ },
45
+ "outputs": { "Y": { "shape": [1, 512], "dtype": "float16" } },
46
+ "bench": { "primary": true, "metrics": [{ "type": "gflops", "value": "2 * 1 * 512 * 2048" }] }
47
+ },
48
+ {
49
+ "name": "transposematmul-f16-attn-scores-transb-8x512x64",
50
+ "preset": "smoke",
51
+ "attrs": { "alpha": 0.125, "transB": 1 },
52
+ "vars": { "M": 512, "K": 64, "N": 512 },
53
+ "inputs": {
54
+ "A": { "shape": [8, 512, 64], "dtype": "float16", "dist": "normal", "seed": 608, "scale": 0.2 },
55
+ "B": { "shape": [8, 512, 64], "dtype": "float16", "dist": "normal", "seed": 609, "scale": 0.2 }
56
+ },
57
+ "outputs": { "Y": { "shape": [8, 512, 512], "dtype": "float16" } },
58
+ "bench": { "primary": true, "metrics": [{ "type": "gflops", "value": "2 * 8 * 512 * 512 * 64" }] }
59
+ },
60
+ {
61
+ "name": "transposematmul-f16-broadcast-batch-rank4x3-2x8x512x64",
62
+ "preset": "model",
63
+ "attrs": { "alpha": 0.125, "transB": 1 },
64
+ "vars": { "M": 512, "K": 64, "N": 512 },
65
+ "inputs": {
66
+ "A": { "shape": [2, 8, 512, 64], "dtype": "float16", "dist": "normal", "seed": 610, "scale": 0.2 },
67
+ "B": { "shape": [8, 512, 64], "dtype": "float16", "dist": "normal", "seed": 611, "scale": 0.2 }
68
+ },
69
+ "outputs": { "Y": { "shape": [2, 8, 512, 512], "dtype": "float16" } },
70
+ "bench": { "primary": true, "metrics": [{ "type": "gflops", "value": "2 * 2 * 8 * 512 * 512 * 64" }] }
71
+ },
72
+ {
73
+ "name": "transposematmul-f16-rank4-by-rank2-shared-weight-b2h8-m512-k2048-n512",
74
+ "preset": "model",
75
+ "provenance": {
76
+ "notes": "Measures TransposeMatMul over a batched (2x8) projection sharing one rank-2 [2048,512] weight, with M=512, K=2048, N=512 (float16)."
77
+ },
78
+ "attrs": { "alpha": 1 },
79
+ "vars": { "dtype": "float16", "M": 512, "K": 2048, "N": 512 },
80
+ "inputs": {
81
+ "A": { "shape": [2, 8, 512, 2048], "dtype": "float16", "dist": "normal", "seed": 744, "scale": 0.05 },
82
+ "B": { "shape": [2048, 512], "dtype": "float16", "dist": "normal", "seed": 745, "scale": 0.05 }
83
+ },
84
+ "outputs": { "Y": { "shape": [2, 8, 512, 512], "dtype": "float16", "dist": "empty" } },
85
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * args.K" }] }
86
+ },
87
+ {
88
+ "name": "transposematmul-f16-band-m16-k2560-n4096",
89
+ "preset": "smoke",
90
+ "attrs": { "alpha": 0.5 },
91
+ "inputs": {
92
+ "A": { "shape": [16, 2560], "dtype": "float16", "dist": "normal", "seed": 752, "scale": 0.1 },
93
+ "B": { "shape": [2560, 4096], "dtype": "float16", "dist": "normal", "seed": 753, "scale": 0.1 }
94
+ },
95
+ "outputs": { "Y": { "shape": [16, 4096], "dtype": "float16", "dist": "empty" } },
96
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * 16 * 2560 * 4096" }] }
97
+ },
98
+ {
99
+ "name": "broadcast-transb-tails-float16",
100
+ "preset": "stress",
101
+ "attrs": { "transB": 1, "alpha": -0.5 },
102
+ "inputs": {
103
+ "A": { "dtype": "float16", "shape": [2, 1, 129, 65], "dist": "normal", "seed": 7302, "scale": 0.1 },
104
+ "B": { "dtype": "float16", "shape": [3, 129, 65], "dist": "normal", "seed": 7303, "scale": 0.1 }
105
+ },
106
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 129, 129] } },
107
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] },
108
+ "provenance": {
109
+ "notes": "Measures TransposeMatMul over a rank-4 by rank-3 broadcast (batch dims 2x1 against 3) with transposed B: M=129, K=65, N=129 (float16)."
110
+ }
111
+ },
112
+ {
113
+ "name": "broadcast-transb-rank3x2-float16",
114
+ "preset": "stress",
115
+ "attrs": { "transB": 1, "alpha": -0.5 },
116
+ "inputs": {
117
+ "A": { "dtype": "float16", "shape": [4, 256, 128], "dist": "normal", "seed": 7304, "scale": 0.1 },
118
+ "B": { "dtype": "float16", "shape": [512, 128], "dist": "normal", "seed": 7305, "scale": 0.1 }
119
+ },
120
+ "outputs": { "Y": { "dtype": "float16", "shape": [4, 256, 512] } },
121
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] },
122
+ "provenance": {
123
+ "notes": "Measures TransposeMatMul over a rank-3 by rank-2 broadcast with transposed B: M=256, K=128, N=512, batch=4 (float16)."
124
+ }
125
+ },
126
+ {
127
+ "name": "broadcast-transb-rank5x3-float16",
128
+ "preset": "stress",
129
+ "attrs": { "transB": 1, "alpha": -0.5 },
130
+ "inputs": {
131
+ "A": { "dtype": "float16", "shape": [2, 1, 2, 128, 32], "dist": "normal", "seed": 7306, "scale": 0.1 },
132
+ "B": { "dtype": "float16", "shape": [2, 128, 32], "dist": "normal", "seed": 7307, "scale": 0.1 }
133
+ },
134
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 1, 2, 128, 128] } },
135
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] },
136
+ "provenance": {
137
+ "notes": "Measures TransposeMatMul over a rank-5 by rank-3 broadcast with transposed B: M=128, K=32, N=128, batch dims 2x1x2 (float16)."
138
+ }
139
+ },
140
+ {
141
+ "name": "broadcast-transb-low_tiles-float16",
142
+ "preset": "stress",
143
+ "attrs": { "transB": 1, "alpha": -0.5 },
144
+ "inputs": {
145
+ "A": { "dtype": "float16", "shape": [1, 64, 32], "dist": "normal", "seed": 7308, "scale": 0.1 },
146
+ "B": { "dtype": "float16", "shape": [64, 32], "dist": "normal", "seed": 7309, "scale": 0.1 }
147
+ },
148
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 64, 64] } },
149
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] },
150
+ "provenance": {
151
+ "notes": "Measures TransposeMatMul over a rank-3 by rank-2 broadcast with transposed B: M=64, K=32, N=64, batch=1 (float16)."
152
+ }
153
+ },
154
+ {
155
+ "name": "broadcast-transb-large-float32",
156
+ "preset": "stress",
157
+ "attrs": { "transB": 1, "alpha": -0.5 },
158
+ "inputs": {
159
+ "A": { "dtype": "float32", "shape": [2, 8, 512, 64], "dist": "normal", "seed": 7300, "scale": 0.1 },
160
+ "B": { "dtype": "float32", "shape": [8, 512, 64], "dist": "normal", "seed": 7301, "scale": 0.1 }
161
+ },
162
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 8, 512, 512] } },
163
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] },
164
+ "provenance": {
165
+ "notes": "Measures TransposeMatMul over a rank-4 by rank-3 broadcast with transposed B: M=512, K=64, N=512, batch dims 2x8 (float32)."
166
+ }
167
+ },
168
+ {
169
+ "name": "broadcast-transb-tails-float32",
170
+ "preset": "stress",
171
+ "attrs": { "transB": 1, "alpha": -0.5 },
172
+ "inputs": {
173
+ "A": { "dtype": "float32", "shape": [2, 1, 129, 65], "dist": "normal", "seed": 7302, "scale": 0.1 },
174
+ "B": { "dtype": "float32", "shape": [3, 129, 65], "dist": "normal", "seed": 7303, "scale": 0.1 }
175
+ },
176
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 129, 129] } },
177
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] },
178
+ "provenance": {
179
+ "notes": "Measures TransposeMatMul over a rank-4 by rank-3 broadcast (batch dims 2x1 against 3) with transposed B: M=129, K=65, N=129 (float32)."
180
+ }
181
+ },
182
+ {
183
+ "name": "broadcast-transb-rank3x2-float32",
184
+ "preset": "stress",
185
+ "attrs": { "transB": 1, "alpha": -0.5 },
186
+ "inputs": {
187
+ "A": { "dtype": "float32", "shape": [4, 256, 128], "dist": "normal", "seed": 7304, "scale": 0.1 },
188
+ "B": { "dtype": "float32", "shape": [512, 128], "dist": "normal", "seed": 7305, "scale": 0.1 }
189
+ },
190
+ "outputs": { "Y": { "dtype": "float32", "shape": [4, 256, 512] } },
191
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] },
192
+ "provenance": {
193
+ "notes": "Measures TransposeMatMul over a rank-3 by rank-2 broadcast with transposed B: M=256, K=128, N=512, batch=4 (float32)."
194
+ }
195
+ },
196
+ {
197
+ "name": "broadcast-transb-rank5x3-float32",
198
+ "preset": "stress",
199
+ "attrs": { "transB": 1, "alpha": -0.5 },
200
+ "inputs": {
201
+ "A": { "dtype": "float32", "shape": [2, 1, 2, 128, 32], "dist": "normal", "seed": 7306, "scale": 0.1 },
202
+ "B": { "dtype": "float32", "shape": [2, 128, 32], "dist": "normal", "seed": 7307, "scale": 0.1 }
203
+ },
204
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 1, 2, 128, 128] } },
205
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] },
206
+ "provenance": {
207
+ "notes": "Measures TransposeMatMul over a rank-5 by rank-3 broadcast with transposed B: M=128, K=32, N=128, batch dims 2x1x2 (float32)."
208
+ }
209
+ },
210
+ {
211
+ "name": "broadcast-transb-low_tiles-float32",
212
+ "preset": "stress",
213
+ "attrs": { "transB": 1, "alpha": -0.5 },
214
+ "inputs": {
215
+ "A": { "dtype": "float32", "shape": [1, 64, 32], "dist": "normal", "seed": 7308, "scale": 0.1 },
216
+ "B": { "dtype": "float32", "shape": [64, 32], "dist": "normal", "seed": 7309, "scale": 0.1 }
217
+ },
218
+ "outputs": { "Y": { "dtype": "float32", "shape": [1, 64, 64] } },
219
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] },
220
+ "provenance": {
221
+ "notes": "Measures TransposeMatMul over a rank-3 by rank-2 broadcast with transposed B: M=64, K=32, N=64, batch=1 (float32)."
222
+ }
223
+ },
224
+ {
225
+ "name": "broadcast-grid-b3-m128-k64-n256-float16-a0.125",
226
+ "preset": "stress",
227
+ "attrs": { "transB": 1, "alpha": 0.125 },
228
+ "inputs": {
229
+ "A": { "dtype": "float16", "shape": [3, 128, 64], "dist": "normal", "seed": 7304, "scale": 0.1 },
230
+ "B": { "dtype": "float16", "shape": [256, 64], "dist": "normal", "seed": 7305, "scale": 0.1 }
231
+ },
232
+ "outputs": { "Y": { "dtype": "float16", "shape": [3, 128, 256] } },
233
+ "provenance": {
234
+ "notes": "Measures TransposeMatMul over a batch=3 grid with transposed B: M=128, K=64, N=256, alpha=0.125 (float16)."
235
+ },
236
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
237
+ },
238
+ {
239
+ "name": "broadcast-grid-b4-m128-k64-n256-float16-a-0.375",
240
+ "preset": "stress",
241
+ "attrs": { "transB": 1, "alpha": -0.375 },
242
+ "inputs": {
243
+ "A": { "dtype": "float16", "shape": [4, 128, 64], "dist": "normal", "seed": 7304, "scale": 0.1 },
244
+ "B": { "dtype": "float16", "shape": [256, 64], "dist": "normal", "seed": 7305, "scale": 0.1 }
245
+ },
246
+ "outputs": { "Y": { "dtype": "float16", "shape": [4, 128, 256] } },
247
+ "provenance": {
248
+ "notes": "Measures TransposeMatMul over a batch=4 grid with transposed B: M=128, K=64, N=256, alpha=-0.375 (float16)."
249
+ },
250
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
251
+ },
252
+ {
253
+ "name": "broadcast-grid-b8-m128-k64-n256-float16-a0.5",
254
+ "preset": "stress",
255
+ "attrs": { "transB": 1, "alpha": 0.5 },
256
+ "inputs": {
257
+ "A": { "dtype": "float16", "shape": [8, 128, 64], "dist": "normal", "seed": 7304, "scale": 0.1 },
258
+ "B": { "dtype": "float16", "shape": [256, 64], "dist": "normal", "seed": 7305, "scale": 0.1 }
259
+ },
260
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 128, 256] } },
261
+ "provenance": {
262
+ "notes": "Measures TransposeMatMul over a batch=8 grid with transposed B: M=128, K=64, N=256, alpha=0.5 (float16)."
263
+ },
264
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
265
+ },
266
+ {
267
+ "name": "broadcast-grid-b8-m129-k65-n257-float16-a0.125",
268
+ "preset": "stress",
269
+ "attrs": { "transB": 1, "alpha": 0.125 },
270
+ "inputs": {
271
+ "A": { "dtype": "float16", "shape": [8, 129, 65], "dist": "normal", "seed": 7304, "scale": 0.1 },
272
+ "B": { "dtype": "float16", "shape": [257, 65], "dist": "normal", "seed": 7305, "scale": 0.1 }
273
+ },
274
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 129, 257] } },
275
+ "provenance": {
276
+ "notes": "Measures TransposeMatMul over a batch=8 grid with transposed B: M=129, K=65, N=257, alpha=0.125 (float16)."
277
+ },
278
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
279
+ },
280
+ {
281
+ "name": "broadcast-grid-b8-m128-k96-n128-float16-a0",
282
+ "preset": "stress",
283
+ "attrs": { "transB": 1, "alpha": 0 },
284
+ "inputs": {
285
+ "A": { "dtype": "float16", "shape": [8, 128, 96], "dist": "normal", "seed": 7304, "scale": 0.1 },
286
+ "B": { "dtype": "float16", "shape": [128, 96], "dist": "normal", "seed": 7305, "scale": 0.1 }
287
+ },
288
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 128, 128] } },
289
+ "provenance": {
290
+ "notes": "Measures TransposeMatMul over a batch=8 grid with transposed B: M=128, K=96, N=128, alpha=0 (float16)."
291
+ },
292
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
293
+ },
294
+ {
295
+ "name": "broadcast-grid-b4-m256-k127-n512-float16-a-0.5",
296
+ "preset": "stress",
297
+ "attrs": { "transB": 1, "alpha": -0.5 },
298
+ "inputs": {
299
+ "A": { "dtype": "float16", "shape": [4, 256, 127], "dist": "normal", "seed": 7304, "scale": 0.1 },
300
+ "B": { "dtype": "float16", "shape": [512, 127], "dist": "normal", "seed": 7305, "scale": 0.1 }
301
+ },
302
+ "outputs": { "Y": { "dtype": "float16", "shape": [4, 256, 512] } },
303
+ "provenance": {
304
+ "notes": "Measures TransposeMatMul over a batch=4 grid with transposed B: M=256, K=127, N=512, alpha=-0.5 (float16)."
305
+ },
306
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
307
+ },
308
+ {
309
+ "name": "broadcast-grid-b3-m128-k64-n256-float32-a0.125",
310
+ "preset": "stress",
311
+ "attrs": { "transB": 1, "alpha": 0.125 },
312
+ "inputs": {
313
+ "A": { "dtype": "float32", "shape": [3, 128, 64], "dist": "normal", "seed": 7304, "scale": 0.1 },
314
+ "B": { "dtype": "float32", "shape": [256, 64], "dist": "normal", "seed": 7305, "scale": 0.1 }
315
+ },
316
+ "outputs": { "Y": { "dtype": "float32", "shape": [3, 128, 256] } },
317
+ "provenance": {
318
+ "notes": "Measures TransposeMatMul over a batch=3 grid with transposed B: M=128, K=64, N=256, alpha=0.125 (float32)."
319
+ },
320
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
321
+ },
322
+ {
323
+ "name": "broadcast-grid-b4-m128-k64-n256-float32-a-0.375",
324
+ "preset": "stress",
325
+ "attrs": { "transB": 1, "alpha": -0.375 },
326
+ "inputs": {
327
+ "A": { "dtype": "float32", "shape": [4, 128, 64], "dist": "normal", "seed": 7304, "scale": 0.1 },
328
+ "B": { "dtype": "float32", "shape": [256, 64], "dist": "normal", "seed": 7305, "scale": 0.1 }
329
+ },
330
+ "outputs": { "Y": { "dtype": "float32", "shape": [4, 128, 256] } },
331
+ "provenance": {
332
+ "notes": "Measures TransposeMatMul over a batch=4 grid with transposed B: M=128, K=64, N=256, alpha=-0.375 (float32)."
333
+ },
334
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
335
+ },
336
+ {
337
+ "name": "broadcast-grid-b8-m128-k64-n256-float32-a0.5",
338
+ "preset": "stress",
339
+ "attrs": { "transB": 1, "alpha": 0.5 },
340
+ "inputs": {
341
+ "A": { "dtype": "float32", "shape": [8, 128, 64], "dist": "normal", "seed": 7304, "scale": 0.1 },
342
+ "B": { "dtype": "float32", "shape": [256, 64], "dist": "normal", "seed": 7305, "scale": 0.1 }
343
+ },
344
+ "outputs": { "Y": { "dtype": "float32", "shape": [8, 128, 256] } },
345
+ "provenance": {
346
+ "notes": "Measures TransposeMatMul over a batch=8 grid with transposed B: M=128, K=64, N=256, alpha=0.5 (float32)."
347
+ },
348
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
349
+ },
350
+ {
351
+ "name": "broadcast-grid-b8-m129-k65-n257-float32-a0.125",
352
+ "preset": "stress",
353
+ "attrs": { "transB": 1, "alpha": 0.125 },
354
+ "inputs": {
355
+ "A": { "dtype": "float32", "shape": [8, 129, 65], "dist": "normal", "seed": 7304, "scale": 0.1 },
356
+ "B": { "dtype": "float32", "shape": [257, 65], "dist": "normal", "seed": 7305, "scale": 0.1 }
357
+ },
358
+ "outputs": { "Y": { "dtype": "float32", "shape": [8, 129, 257] } },
359
+ "provenance": {
360
+ "notes": "Measures TransposeMatMul over a batch=8 grid with transposed B: M=129, K=65, N=257, alpha=0.125 (float32)."
361
+ },
362
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
363
+ },
364
+ {
365
+ "name": "broadcast-grid-b8-m128-k96-n128-float32-a0",
366
+ "preset": "stress",
367
+ "attrs": { "transB": 1, "alpha": 0 },
368
+ "inputs": {
369
+ "A": { "dtype": "float32", "shape": [8, 128, 96], "dist": "normal", "seed": 7304, "scale": 0.1 },
370
+ "B": { "dtype": "float32", "shape": [128, 96], "dist": "normal", "seed": 7305, "scale": 0.1 }
371
+ },
372
+ "outputs": { "Y": { "dtype": "float32", "shape": [8, 128, 128] } },
373
+ "provenance": {
374
+ "notes": "Measures TransposeMatMul over a batch=8 grid with transposed B: M=128, K=96, N=128, alpha=0 (float32)."
375
+ },
376
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
377
+ },
378
+ {
379
+ "name": "broadcast-grid-b4-m256-k127-n512-float32-a-0.5",
380
+ "preset": "stress",
381
+ "attrs": { "transB": 1, "alpha": -0.5 },
382
+ "inputs": {
383
+ "A": { "dtype": "float32", "shape": [4, 256, 127], "dist": "normal", "seed": 7304, "scale": 0.1 },
384
+ "B": { "dtype": "float32", "shape": [512, 127], "dist": "normal", "seed": 7305, "scale": 0.1 }
385
+ },
386
+ "outputs": { "Y": { "dtype": "float32", "shape": [4, 256, 512] } },
387
+ "provenance": {
388
+ "notes": "Measures TransposeMatMul over a batch=4 grid with transposed B: M=256, K=127, N=512, alpha=-0.5 (float32)."
389
+ },
390
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
391
+ },
392
+ {
393
+ "name": "broadcast-selected-mn-tail-float16",
394
+ "preset": "stress",
395
+ "attrs": { "transB": 1, "alpha": 0.125 },
396
+ "inputs": {
397
+ "A": { "dtype": "float16", "shape": [8, 129, 64], "dist": "normal", "seed": 7304, "scale": 0.1 },
398
+ "B": { "dtype": "float16", "shape": [257, 64], "dist": "normal", "seed": 7305, "scale": 0.1 }
399
+ },
400
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 129, 257] } },
401
+ "provenance": {
402
+ "notes": "Measures TransposeMatMul over a batch=8 grid with transposed B: M=129, K=64, N=257, alpha=0.125 (float16)."
403
+ },
404
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
405
+ },
406
+ {
407
+ "name": "broadcast-selected-mn-tail-float32",
408
+ "preset": "stress",
409
+ "attrs": { "transB": 1, "alpha": 0.125 },
410
+ "inputs": {
411
+ "A": { "dtype": "float32", "shape": [8, 129, 64], "dist": "normal", "seed": 7304, "scale": 0.1 },
412
+ "B": { "dtype": "float32", "shape": [257, 64], "dist": "normal", "seed": 7305, "scale": 0.1 }
413
+ },
414
+ "outputs": { "Y": { "dtype": "float32", "shape": [8, 129, 257] } },
415
+ "provenance": {
416
+ "notes": "Measures TransposeMatMul over a batch=8 grid with transposed B: M=129, K=64, N=257, alpha=0.125 (float32)."
417
+ },
418
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
419
+ },
420
+ {
421
+ "name": "broadcast-selected-broadcast-tail-float16",
422
+ "preset": "stress",
423
+ "attrs": { "transB": 1, "alpha": 0.125 },
424
+ "inputs": {
425
+ "A": { "dtype": "float16", "shape": [4, 1, 129, 64], "dist": "normal", "seed": 7304, "scale": 0.1 },
426
+ "B": { "dtype": "float16", "shape": [3, 257, 64], "dist": "normal", "seed": 7305, "scale": 0.1 }
427
+ },
428
+ "outputs": { "Y": { "dtype": "float16", "shape": [4, 3, 129, 257] } },
429
+ "provenance": {
430
+ "notes": "Measures TransposeMatMul over a rank-4 by rank-3 broadcast (batch dims 4x1 against 3) with transposed B: M=129, K=64, N=257, alpha=0.125 (float16)."
431
+ },
432
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
433
+ },
434
+ {
435
+ "name": "broadcast-selected-broadcast-tail-float32",
436
+ "preset": "stress",
437
+ "attrs": { "transB": 1, "alpha": 0.125 },
438
+ "inputs": {
439
+ "A": { "dtype": "float32", "shape": [4, 1, 129, 64], "dist": "normal", "seed": 7304, "scale": 0.1 },
440
+ "B": { "dtype": "float32", "shape": [3, 257, 64], "dist": "normal", "seed": 7305, "scale": 0.1 }
441
+ },
442
+ "outputs": { "Y": { "dtype": "float32", "shape": [4, 3, 129, 257] } },
443
+ "provenance": {
444
+ "notes": "Measures TransposeMatMul over a rank-4 by rank-3 broadcast (batch dims 4x1 against 3) with transposed B: M=129, K=64, N=257, alpha=0.125 (float32)."
445
+ },
446
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * numel(shapes.Y) * dim(shapes.A, ranks.A - 1)" }] }
447
+ }
448
+ ]
449
+ }
build/webgpu/fused-matmul-subgroup-matrix.wgsl.jinja ADDED
@@ -0,0 +1,177 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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;
18
+ {% endif %}
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
+
27
+ const K: u32 = {{ K }}u;
28
+ const N: u32 = {{ N }}u;
29
+ {% if not transA %}
30
+ const A_M_STRIDE: u32 = K;
31
+ {% endif %}
32
+ const B_BATCH_STRIDE: u32 = K * N;
33
+ const ALPHA: {{ accScalar }} = {{ accScalar }}({{ alpha }});
34
+ const TILE_COLS: u32 = 64u;
35
+ const TILE_ROWS: u32 = 32u;
36
+ const TILE_K: u32 = 32u;
37
+ const SUB_COLS: u32 = 32u;
38
+ const SUB_ROWS: u32 = 16u;
39
+
40
+ var<workgroup> tile_A: array<{{ operandScalar }}, 32 * 32>;
41
+ var<workgroup> tile_B: array<{{ operandScalar }}, 64 * 32>;
42
+ var<workgroup> scratch: array<array<array<{{ accScalar }}, 64>, 4>, 4>;
43
+
44
+ fn loadSHMA(a_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
45
+ let a_global = tile_base + row;
46
+ let col = c_idx * 8u;
47
+ for (var col_offset = 0u; col_offset < 8u; col_offset = col_offset + 1u) {
48
+ let k = k_idx + col + col_offset;
49
+ if (a_global < params.M) {
50
+ {% if transA %}
51
+ // op(A) = A^T: A stored [.., K, M], so op(A)[a_global, k] = A[k, a_global].
52
+ tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(a[a_base + k * params.M + a_global]);
53
+ {% else %}
54
+ tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(a[a_base + a_global * A_M_STRIDE + k]);
55
+ {% endif %}
56
+ } else {
57
+ tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(0.0);
58
+ }
59
+ }
60
+ }
61
+
62
+ fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
63
+ let b_col = tile_base + row;
64
+ let col = c_idx * 16u;
65
+ for (var i = 0u; i < 16u; i = i + 1u) {
66
+ let k = k_idx + col + i;
67
+ {% if transB %}
68
+ // op(B) = B^T: B stored [.., N, K], so op(B)[k, b_col] = B[b_col, k].
69
+ tile_B[row * TILE_K + col + i] = {{ operandScalar }}(b[b_base + b_col * K + k]);
70
+ {% else %}
71
+ tile_B[row * TILE_K + col + i] = {{ operandScalar }}(b[b_base + k * N + b_col]);
72
+ {% endif %}
73
+ }
74
+ }
75
+
76
+ fn storeOutput(offset: u32, row: u32, col: u32, src_slot: u32, row_limit: i32) {
77
+ if (row_limit > 0 && row < u32(row_limit)) {
78
+ let col2 = col + 1u;
79
+ y[offset + row * N + col] = {{ outScalar }}(ALPHA * scratch[src_slot][0][row * 8u + col]);
80
+ y[offset + row * N + col + 8u] = {{ outScalar }}(ALPHA * scratch[src_slot][1][row * 8u + col]);
81
+ y[offset + row * N + col + 16u] = {{ outScalar }}(ALPHA * scratch[src_slot][2][row * 8u + col]);
82
+ y[offset + row * N + col + 24u] = {{ outScalar }}(ALPHA * scratch[src_slot][3][row * 8u + col]);
83
+
84
+ y[offset + row * N + col2] = {{ outScalar }}(ALPHA * scratch[src_slot][0][row * 8u + col2]);
85
+ y[offset + row * N + col2 + 8u] = {{ outScalar }}(ALPHA * scratch[src_slot][1][row * 8u + col2]);
86
+ y[offset + row * N + col2 + 16u] = {{ outScalar }}(ALPHA * scratch[src_slot][2][row * 8u + col2]);
87
+ y[offset + row * N + col2 + 24u] = {{ outScalar }}(ALPHA * scratch[src_slot][3][row * 8u + col2]);
88
+ }
89
+ }
90
+
91
+ @compute @workgroup_size(128, 1, 1){{ " @subgroup_size(32)" if pinSubgroupSize32 else "" }}
92
+ fn main(
93
+ @builtin(workgroup_id) workgroup_id: vec3<u32>,
94
+ @builtin(local_invocation_index) local_idx: u32,
95
+ @builtin(subgroup_invocation_id) sg_id: u32,
96
+ @builtin(subgroup_size) sg_size: u32
97
+ ) {
98
+ let M = params.M;
99
+ let A_BATCH_STRIDE = M * K;
100
+ let C_BATCH_STRIDE = M * N;
101
+ let batch = workgroup_id.z;
102
+ let a_base = batch * A_BATCH_STRIDE;
103
+ let b_base = batch * B_BATCH_STRIDE;
104
+ let c_base = batch * C_BATCH_STRIDE;
105
+ let a_global_base = workgroup_id.y * TILE_ROWS;
106
+ let b_global_base = workgroup_id.x * TILE_COLS;
107
+
108
+ let subtile_id = local_idx / sg_size;
109
+ let subtile_idx = subtile_id / 2u;
110
+ let subtile_idy = subtile_id % 2u;
111
+ let base_A = subtile_idy * SUB_ROWS;
112
+ let base_B = subtile_idx * SUB_COLS;
113
+
114
+ var matC00: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
115
+ var matC01: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
116
+ var matC02: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
117
+ var matC03: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
118
+ var matC10: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
119
+ var matC11: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
120
+ var matC12: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
121
+ var matC13: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
122
+
123
+ for (var kidx = 0u; kidx < K; kidx = kidx + TILE_K) {
124
+ loadSHMA(a_base, a_global_base, kidx, local_idx / 4u, local_idx % 4u);
125
+ loadSHMB(b_base, b_global_base, kidx, local_idx / 2u, local_idx % 2u);
126
+ workgroupBarrier();
127
+
128
+ for (var step = 0u; step < TILE_K; step = step + 8u) {
129
+ {% set directInputs = false %}
130
+ let matrix_a_offset = subtile_idy * SUB_ROWS * TILE_K + step;
131
+ {% for r in range(2) %}
132
+ 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" }});
133
+ {% endfor %}
134
+
135
+ let matrix_b_offset = subtile_idx * SUB_COLS * TILE_K + step;
136
+ {% for c in range(4) %}
137
+ 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" }});
138
+ {% endfor %}
139
+
140
+ matC00 = subgroupMatrixMultiplyAccumulate(matA0, matB0, matC00);
141
+ matC01 = subgroupMatrixMultiplyAccumulate(matA0, matB1, matC01);
142
+ matC02 = subgroupMatrixMultiplyAccumulate(matA0, matB2, matC02);
143
+ matC03 = subgroupMatrixMultiplyAccumulate(matA0, matB3, matC03);
144
+ matC10 = subgroupMatrixMultiplyAccumulate(matA1, matB0, matC10);
145
+ matC11 = subgroupMatrixMultiplyAccumulate(matA1, matB1, matC11);
146
+ matC12 = subgroupMatrixMultiplyAccumulate(matA1, matB2, matC12);
147
+ matC13 = subgroupMatrixMultiplyAccumulate(matA1, matB3, matC13);
148
+ }
149
+ workgroupBarrier();
150
+ }
151
+
152
+ // The four scratch banks are reused across the two row-groups, and each is written
153
+ // by a collective subgroupMatrixStore then read across lanes by storeOutput. Barriers
154
+ // give the reads visibility of the store and stop the second row-group's store from
155
+ // clobbering the first's still-in-flight readback when a partial final M-tile
156
+ // diverges storeOutput's guard. Without both barriers the last valid row can be corrupted.
157
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][0], 0u, matC00, 8u);
158
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][1], 0u, matC01, 8u);
159
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][2], 0u, matC02, 8u);
160
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][3], 0u, matC03, 8u);
161
+ workgroupBarrier();
162
+ let row = sg_id / 4u;
163
+ let col = (sg_id % 4u) * 2u;
164
+ var matrix_c_offset = c_base + (a_global_base + base_A) * N + b_global_base + base_B;
165
+ var row_limit = i32(M) - i32(a_global_base + base_A);
166
+ storeOutput(matrix_c_offset, row, col, subtile_id, row_limit);
167
+ workgroupBarrier();
168
+
169
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][0], 0u, matC10, 8u);
170
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][1], 0u, matC11, 8u);
171
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][2], 0u, matC12, 8u);
172
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][3], 0u, matC13, 8u);
173
+ workgroupBarrier();
174
+ matrix_c_offset = matrix_c_offset + 8u * N;
175
+ row_limit = i32(M) - i32(a_global_base + base_A + 8u);
176
+ storeOutput(matrix_c_offset, row, col, subtile_id, row_limit);
177
+ }
build/webgpu/manifest.json ADDED
@@ -0,0 +1,503 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "domain": "com.microsoft",
3
+ "name": "TransposeMatMul",
4
+ "sinceVersion": 1,
5
+ "inputs": { "A": { "dtype": "T" }, "B": { "dtype": "T" } },
6
+ "outputs": {
7
+ "Y": {
8
+ "dtype": "T",
9
+ "rank": "max(ranks.A, ranks.B) - (1 if ranks.A == 1 or ranks.B == 1 else 0)",
10
+ "shape": "matmulShape(logicalAShape, logicalBShape)"
11
+ }
12
+ },
13
+ "attributes": { "alpha": { "default": 1 }, "transA": { "default": 0 }, "transB": { "default": 0 } },
14
+ "typeConstraints": { "T": ["float32", "float16"] },
15
+ "tunables": {
16
+ "TILED_REG_MIN_WORKGROUPS": { "default": 64 },
17
+ "PLAIN_RANK2_REG_DEEP_K_TILES": { "default": 128 },
18
+ "GEMV_TARGET_BLOCKS": { "default": 512 },
19
+ "SUBGROUP_MATRIX_MIN_M": { "default": 2 },
20
+ "SUBGROUP_MATRIX_SPLITK_TARGET_WGS": { "default": 512 },
21
+ "SUBGROUP_MATRIX_SPLITK_MIN_K": { "default": 1024 },
22
+ "SUBGROUP_MATRIX_SPLITK_MAX_TILES": { "default": 128 },
23
+ "BAND_VEC4_MAX_ROWS": { "default": 16 },
24
+ "BAND_SPLIT_TARGET_WORKGROUPS": { "default": 256 },
25
+ "BAND_SPLIT_MAX_COLUMN_GROUPS": { "default": 24 },
26
+ "BAND_SPLIT_SLICES": { "default": 8 },
27
+ "BAND_PREFER_MAX_ROWS": { "default": 8 },
28
+ "BAND_PREFER_DEEP_K": { "default": 4096 },
29
+ "BROADCAST_TRANSB_MIN_WORKGROUPS": { "default": 64 },
30
+ "BROADCAST_TRANSB_MAX_PADDING_RATIO": { "default": 2 }
31
+ },
32
+ "derive": {
33
+ "batchMovedAShape": "shapes.A",
34
+ "batchMovedBShape": "shapes.B",
35
+ "logicalAShape": "moveAxis(batchMovedAShape, -1, -2) if attrs.transA != 0 and ranks.A > 1 else batchMovedAShape",
36
+ "logicalBShape": "moveAxis(batchMovedBShape, -1, -2) if attrs.transB != 0 and ranks.B > 1 else batchMovedBShape",
37
+ "gemvN": "dim(shapes.B, 1)",
38
+ "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32",
39
+ "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",
40
+ "pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter",
41
+ "wave32Effective": "wave32Adapter or pinSubgroupSize32",
42
+ "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",
43
+ "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
44
+ "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\"))))",
45
+ "broadcastTransbM": "dim(shapes.A, ranks.A - 2)",
46
+ "broadcastTransbN": "dim(shapes.B, ranks.B - 2)",
47
+ "broadcastTransbK": "dim(shapes.A, ranks.A - 1)",
48
+ "broadcastTransbBatches": "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 2))",
49
+ "gemvLanes": 32,
50
+ "vec4OutputTile": "4 * gemvLanes",
51
+ "gemvWorkgroups": "ceilDiv(gemvN, vec4OutputTile)",
52
+ "gemvSliceCap": "min(32, device.limits.maxComputeWorkgroupSizeY, floor(device.limits.maxComputeInvocationsPerWorkgroup / gemvLanes), floor(device.limits.maxComputeWorkgroupStorageSize / (16 * gemvLanes)))",
53
+ "gemvSlices": "max(1, min(gemvSliceCap, max(8, pow2ceil(ceilDiv(tunables.GEMV_TARGET_BLOCKS, gemvWorkgroups)))))",
54
+ "gemvResourcesFit": "gemvLanes <= device.limits.maxComputeWorkgroupSizeX and gemvSlices <= device.limits.maxComputeWorkgroupSizeY and gemvLanes * gemvSlices <= device.limits.maxComputeInvocationsPerWorkgroup and 16 * gemvLanes * gemvSlices <= device.limits.maxComputeWorkgroupStorageSize",
55
+ "registerTile": 64,
56
+ "generalTile": 32,
57
+ "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",
58
+ "plainRank2RegDeepPreferredTier": "rank2DeepPortableTier and dim(shapes.A, 1) >= tunables.PLAIN_RANK2_REG_DEEP_K_TILES * 16 and dim(shapes.A, 1) % 16 == 0",
59
+ "subgroupMatrixResourcesFit": "128 <= deviceWorkgroupCap and ((32 * 32 + 64 * 32) * dtypeBytes(dtypes.T) + 4 * 4 * 64 * 4) <= device.limits.maxComputeWorkgroupStorageSize",
60
+ "sgmatSplitKDepth": "dim(shapes.A, ranks.A - 1)",
61
+ "sgmatSplitK32Ok": "sgmatSplitKDepth % 1024 == 0",
62
+ "sgmatSplitK16Ok": "sgmatSplitKDepth % 512 == 0",
63
+ "sgmatSplitK8Ok": "sgmatSplitKDepth % 256 == 0",
64
+ "sgmatSplitK4Ok": "sgmatSplitKDepth % 128 == 0",
65
+ "sgmatSplitK2Ok": "sgmatSplitKDepth % 64 == 0",
66
+ "bandSplitWant": "pow2ceil(ceilDiv(tunables.BAND_SPLIT_TARGET_WORKGROUPS, max(1, gemvWorkgroups)))",
67
+ "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)))",
68
+ "scalar": "dtypes.T",
69
+ "alpha": "attrs.alpha",
70
+ "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
71
+ "fScalar": "\"f16\" if dtypes.T == \"f16\" else \"f32\"",
72
+ "aRank": "ranks.A",
73
+ "bRank": "ranks.B",
74
+ "fusedSgmatRank2Ok": "ranks.A == 2 and ranks.B == 2 and ranks.Y == 2 and attrs.transA == 0 and attrs.transB == 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)",
75
+ "sgmatOutTiles": "ceilDiv(dim(shapes.A, 0), 32) * ceilDiv(dim(shapes.B, 1), 64) if fusedSgmatRank2Ok else 1",
76
+ "sgmatSplitKWant": "ceilDiv(tunables.SUBGROUP_MATRIX_SPLITK_TARGET_WGS, sgmatOutTiles)",
77
+ "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))))",
78
+ "bandRank2Ok": "ranks.A == 2 and ranks.B == 2 and ranks.Y == 2 and attrs.transA == 0 and attrs.transB == 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)",
79
+ "broadcastTransbContract": "(f16Ok(dtypes.T)) and (attrs.transA == 0 and attrs.transB != 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))"
80
+ },
81
+ "bindings": {
82
+ "a": { "arg": "A", "elementType": "$scalar" },
83
+ "b": { "arg": "B", "elementType": "$vectorScalar" },
84
+ "partials": { "buffer": "read-only-storage", "elementType": "f32" },
85
+ "y": { "arg": "Y", "elementType": "$scalar" },
86
+ "params": { "struct": [{ "name": "cols", "type": "u32", "value": "numel(shapes.Y)" }] },
87
+ "b_scalar": { "arg": "B", "name": "b", "elementType": "$scalar" },
88
+ "params_rows": { "name": "params", "struct": [{ "name": "M", "type": "u32", "value": "rowCount" }] }
89
+ },
90
+ "variants": [
91
+ {
92
+ "id": "broadcast_transb_tiled_reg",
93
+ "priority": 6,
94
+ "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)"],
95
+ "derive": {
96
+ "bShape": "logicalBShape",
97
+ "bTransposed": true,
98
+ "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else broadcastTransbM",
99
+ "N": "broadcastTransbN",
100
+ "batchCount": "broadcastTransbBatches if ranks.B > 2 else 1",
101
+ "transBatchA": false,
102
+ "regSequentialK": "dtypes.T == \"f16\""
103
+ },
104
+ "passes": [
105
+ {
106
+ "id": "main",
107
+ "name": "TransposeMatMul.BroadcastTransBTiledReg",
108
+ "shader": "matmul-tiled-general-reg.wgsl.jinja",
109
+ "derive": {
110
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []",
111
+ "aRank": "ranks.A if ranks.B > 2 else 2",
112
+ "K": "dim(shapes.A, ranks.A - 1)"
113
+ },
114
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
115
+ "dispatch": { "x": "ceilDiv(N, registerTile)", "y": "ceilDiv(rowCount, registerTile)", "z": "batchCount" }
116
+ }
117
+ ],
118
+ "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"]
119
+ },
120
+ {
121
+ "id": "broadcast_transb_subgroup_matrix_f16",
122
+ "priority": 11,
123
+ "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)"],
124
+ "requires": {
125
+ "features": ["subgroups", "chromium-experimental-subgroup-matrix"],
126
+ "subgroupMatrixConfigs": [{ "componentType": "f16", "M": 8, "N": 8, "K": 8 }]
127
+ },
128
+ "derive": {
129
+ "bShape": "logicalBShape",
130
+ "bTransposed": true,
131
+ "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else broadcastTransbM",
132
+ "K": "broadcastTransbK",
133
+ "N": "broadcastTransbN",
134
+ "batchCount": "broadcastTransbBatches if ranks.B > 2 else 1",
135
+ "hasBias": false,
136
+ "fScalar": "dtypes.T",
137
+ "outScalar": "dtypes.T",
138
+ "generalAddressing": true,
139
+ "tailSafe": "dim(shapes.A, ranks.A - 1) % 32 != 0 or dim(shapes.B, ranks.B - 2) % 64 != 0",
140
+ "outputBuffer": "\"y\""
141
+ },
142
+ "passes": [
143
+ {
144
+ "id": "main",
145
+ "name": "TransposeMatMul.BroadcastTransBSubgroupMatrix",
146
+ "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
147
+ "derive": {
148
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []",
149
+ "aRank": "ranks.A if ranks.B > 2 else 2"
150
+ },
151
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
152
+ "dispatch": { "x": "ceilDiv(N,64)", "y": "ceilDiv(rowCount,32)", "z": "batchCount" }
153
+ }
154
+ ],
155
+ "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"]
156
+ },
157
+ {
158
+ "id": "broadcast_transb_subgroup_matrix_f32",
159
+ "priority": 11,
160
+ "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)"],
161
+ "requires": {
162
+ "features": ["subgroups", "chromium-experimental-subgroup-matrix"],
163
+ "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
164
+ },
165
+ "derive": {
166
+ "bShape": "logicalBShape",
167
+ "bTransposed": true,
168
+ "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else broadcastTransbM",
169
+ "K": "broadcastTransbK",
170
+ "N": "broadcastTransbN",
171
+ "batchCount": "broadcastTransbBatches if ranks.B > 2 else 1",
172
+ "hasBias": false,
173
+ "fScalar": "dtypes.T",
174
+ "outScalar": "dtypes.T",
175
+ "generalAddressing": true,
176
+ "tailSafe": "dim(shapes.A, ranks.A - 1) % 32 != 0 or dim(shapes.B, ranks.B - 2) % 64 != 0",
177
+ "outputBuffer": "\"y\""
178
+ },
179
+ "passes": [
180
+ {
181
+ "id": "main",
182
+ "name": "TransposeMatMul.BroadcastTransBSubgroupMatrix",
183
+ "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
184
+ "derive": {
185
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []",
186
+ "aRank": "ranks.A if ranks.B > 2 else 2"
187
+ },
188
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
189
+ "dispatch": { "x": "ceilDiv(N,64)", "y": "ceilDiv(rowCount,32)", "z": "batchCount" }
190
+ }
191
+ ],
192
+ "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"]
193
+ },
194
+ {
195
+ "id": "m1_gemv_vec4",
196
+ "priority": 30,
197
+ "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\")", "attrs.transA == 0", "attrs.transB == 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)"],
198
+ "derive": {
199
+ "unrollK2": "dtypes.T == \"f16\"",
200
+ "gemvScalar": "dtypes.T",
201
+ "gemvVector": "\"vec4<\" ~ dtypes.T ~ \">\"",
202
+ "alphaScale": "attrs.alpha"
203
+ },
204
+ "passes": [
205
+ {
206
+ "id": "main",
207
+ "name": "TransposeMatMul.M1GemvVec4",
208
+ "shader": "matmul-vector-matrix-vec4.wgsl.jinja",
209
+ "bindings": [
210
+ { "arg": "A", "name": "a", "elementType": "$gemvScalar" },
211
+ { "arg": "B", "name": "b", "elementType": "$gemvVector" },
212
+ { "arg": "Y", "name": "c", "elementType": "$gemvVector" },
213
+ {
214
+ "name": "params",
215
+ "struct": [
216
+ { "name": "K", "type": "u32", "value": "dim(shapes.A, 1)" },
217
+ { "name": "N4", "type": "u32", "value": "dim(shapes.B, 1) / 4" }
218
+ ]
219
+ }
220
+ ],
221
+ "dispatch": { "x": "gemvWorkgroups" }
222
+ }
223
+ ]
224
+ },
225
+ {
226
+ "id": "rank2_band_vec4_splitk",
227
+ "priority": 11,
228
+ "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"],
229
+ "derive": {
230
+ "batched": false,
231
+ "outputBuffer": "\"y\"",
232
+ "M": "dim(shapes.A, 0)",
233
+ "K": "dim(shapes.A, 1)",
234
+ "N": "dim(shapes.B, 1)",
235
+ "gemvSlices": "tunables.BAND_SPLIT_SLICES",
236
+ "kSplits": "bandSplitK",
237
+ "split": "bandSplitK",
238
+ "workgroupSize": 256
239
+ },
240
+ "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[bandSplitK * numel(shapes.Y)]" }],
241
+ "passes": [
242
+ {
243
+ "id": "partial",
244
+ "name": "TransposeMatMul.Rank2BandVec4SplitK",
245
+ "shader": "matmul-band-vec4.wgsl.jinja",
246
+ "bindings": ["a", "b", { "scratch": "partials", "name": "y", "elementType": "vec4<f32>" }],
247
+ "dispatch": { "x": "gemvWorkgroups", "y": "bandSplitK" }
248
+ },
249
+ {
250
+ "id": "combine",
251
+ "name": "TransposeMatMul.Rank2BandVec4SplitKCombine",
252
+ "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
253
+ "derive": { "op": "\"sum\"", "outputF16": "dtypes.T == \"f16\"" },
254
+ "bindings": ["partials", "y", "params"],
255
+ "dispatch": {
256
+ "x": "min(ceilDiv((numel(shapes.Y)), (256)), 65535)",
257
+ "y": "ceilDiv(ceilDiv((numel(shapes.Y)), (256)), 65535)",
258
+ "z": 1
259
+ }
260
+ }
261
+ ]
262
+ },
263
+ {
264
+ "id": "rank2_band_vec4",
265
+ "priority": 11,
266
+ "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)"],
267
+ "derive": {
268
+ "batched": false,
269
+ "outputBuffer": "\"y\"",
270
+ "M": "dim(shapes.A, 0)",
271
+ "K": "dim(shapes.A, 1)",
272
+ "N": "dim(shapes.B, 1)"
273
+ },
274
+ "passes": [
275
+ {
276
+ "id": "main",
277
+ "name": "TransposeMatMul.Rank2BandVec4",
278
+ "shader": "matmul-band-vec4.wgsl.jinja",
279
+ "bindings": ["a", "b", { "arg": "Y", "name": "y", "elementType": "$vectorScalar" }],
280
+ "dispatch": { "x": "gemvWorkgroups" }
281
+ }
282
+ ]
283
+ },
284
+ {
285
+ "id": "rank2_band_vec4_f32_preferred",
286
+ "priority": 13,
287
+ "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)"],
288
+ "demoteWhen": ["dtypes.T != \"f32\" or (dim(shapes.A, 0) > tunables.BAND_PREFER_MAX_ROWS and dim(shapes.A, 1) >= tunables.BAND_PREFER_DEEP_K)"],
289
+ "derive": {
290
+ "batched": false,
291
+ "outputBuffer": "\"y\"",
292
+ "M": "dim(shapes.A, 0)",
293
+ "K": "dim(shapes.A, 1)",
294
+ "N": "dim(shapes.B, 1)"
295
+ },
296
+ "passes": [
297
+ {
298
+ "id": "main",
299
+ "name": "TransposeMatMul.Rank2BandVec4",
300
+ "shader": "matmul-band-vec4.wgsl.jinja",
301
+ "bindings": ["a", "b", { "arg": "Y", "name": "y", "elementType": "$vectorScalar" }],
302
+ "dispatch": { "x": "gemvWorkgroups" }
303
+ }
304
+ ]
305
+ },
306
+ {
307
+ "id": "subgroup_matrix_splitk",
308
+ "priority": 12,
309
+ "when": ["(dtypes.T == \"f16\" or dtypes.T == \"f32\") and f16Ok(dtypes.T)", "fusedSgmatRank2Ok", "dim(shapes.A, 0) >= tunables.SUBGROUP_MATRIX_MIN_M", "dim(shapes.A, 1) >= tunables.SUBGROUP_MATRIX_SPLITK_MIN_K", "dim(shapes.B, 1) % 64 == 0", "sgmatSplitK >= 2", "sgmatOutTiles < tunables.SUBGROUP_MATRIX_SPLITK_MAX_TILES", "sgmatSplitK * numel(shapes.Y) * 4 <= device.limits.maxStorageBufferBindingSize", "sgmatSplitK <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.Y, 1), 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.Y, 0), 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "subgroupMatrixResourcesFit", "wave32Effective"],
310
+ "requires": {
311
+ "features": ["subgroups", "chromium-experimental-subgroup-matrix"],
312
+ "subgroupMatrixConfigs": [
313
+ { "componentType": "f16", "M": 8, "N": 8, "K": 8 },
314
+ { "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }
315
+ ]
316
+ },
317
+ "derive": {
318
+ "hasBias": false,
319
+ "generalAddressing": true,
320
+ "tailSafe": false,
321
+ "outputBuffer": "\"partials\"",
322
+ "outScalar": "\"f32\"",
323
+ "rowCount": "dim(shapes.A, 0)",
324
+ "K": "dim(shapes.A, 1)",
325
+ "N": "dim(shapes.B, 1)",
326
+ "batchCount": 1,
327
+ "splitK": "sgmatSplitK",
328
+ "kPerSplit": "dim(shapes.A, 1) / sgmatSplitK",
329
+ "split": "sgmatSplitK",
330
+ "workgroupSize": 256
331
+ },
332
+ "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[sgmatSplitK * numel(shapes.Y)]" }],
333
+ "passes": [
334
+ {
335
+ "id": "partial",
336
+ "name": "TransposeMatMul.SubgroupMatrixSplitK",
337
+ "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
338
+ "derive": { "bShape": ["dim(shapes.B, 0)", "dim(shapes.B, 1)"], "aRank": 2, "bRank": 2 },
339
+ "bindings": ["a", "b_scalar", { "name": "partials", "elementType": "f32" }, "params_rows"],
340
+ "dispatch": { "x": "ceilDiv(dim(shapes.Y, 1), 64)", "y": "ceilDiv(dim(shapes.Y, 0), 32)", "z": "sgmatSplitK" }
341
+ },
342
+ {
343
+ "id": "combine",
344
+ "name": "TransposeMatMul.SubgroupMatrixSplitKCombine",
345
+ "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
346
+ "derive": { "op": "\"sum\"", "outputF16": "dtypes.T == \"f16\"" },
347
+ "bindings": ["partials", "y", "params"],
348
+ "dispatch": {
349
+ "x": "min(ceilDiv((numel(shapes.Y)), (256)), 65535)",
350
+ "y": "ceilDiv(ceilDiv((numel(shapes.Y)), (256)), 65535)",
351
+ "z": 1
352
+ }
353
+ }
354
+ ]
355
+ },
356
+ {
357
+ "id": "subgroup_matrix_tail_broadcast",
358
+ "priority": 11,
359
+ "when": ["dtypes.T == \"f16\"", "f16Ok(dtypes.T)", "attrs.transA == 0", "attrs.transB == 0", "(((ranks.A == 2 and ranks.B == 2 and ranks.Y == 2) or (ranks.A == 4 and ranks.B == 3 and ranks.Y == 4 and dim(shapes.Y, 0) == dim(shapes.A, 0) and (dim(shapes.A, 1) == dim(shapes.B, 0) or dim(shapes.A, 1) == 1 or dim(shapes.B, 0) == 1) and dim(shapes.Y, 1) == max(dim(shapes.A, 1), dim(shapes.B, 0)))) or (ranks.A == 4 and ranks.B == 2 and ranks.Y == 4 and sameShape(prefix(shapes.Y, 2), prefix(shapes.A, 2))))", "dim(shapes.A, ranks.A - 1) == dim(shapes.B, ranks.B - 2)", "dim(shapes.A, ranks.A - 2) >= tunables.SUBGROUP_MATRIX_MIN_M", "dim(shapes.A, ranks.A - 1) >= 32", "dim(shapes.B, ranks.B - 1) >= 64", "dim(shapes.Y, ranks.Y - 2) == dim(shapes.A, ranks.A - 2)", "dim(shapes.Y, ranks.Y - 1) == dim(shapes.B, ranks.B - 1)", "ceil(dim(shapes.B, ranks.B - 1) / 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil(dim(shapes.A, ranks.A - 2) / 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 1)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "wave32Effective"],
360
+ "requires": {
361
+ "features": ["subgroups", "chromium-experimental-subgroup-matrix"],
362
+ "subgroupMatrixConfigs": [{ "componentType": "f16", "M": 8, "N": 8, "K": 8 }]
363
+ },
364
+ "derive": {
365
+ "hasBias": false,
366
+ "fScalar": "\"f16\"",
367
+ "outScalar": "\"f16\"",
368
+ "generalAddressing": true,
369
+ "tailSafe": "dim(shapes.A, ranks.A - 1) % 32 != 0 or dim(shapes.B, ranks.B - 1) % 64 != 0",
370
+ "outputBuffer": "\"y\"",
371
+ "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else dim(shapes.A, ranks.A - 2)",
372
+ "K": "dim(shapes.A, ranks.A - 1)",
373
+ "N": "dim(shapes.B, ranks.B - 1)",
374
+ "batchCount": "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 1)) if ranks.B > 2 else 1"
375
+ },
376
+ "passes": [
377
+ {
378
+ "id": "main",
379
+ "name": "TransposeMatMul.SubgroupMatrixTailBroadcast",
380
+ "shader": "matmul-subgroup-matrix-ext.wgsl.jinja",
381
+ "derive": {
382
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []",
383
+ "aRank": "ranks.A if ranks.B > 2 else 2",
384
+ "bShape": "shapes.B"
385
+ },
386
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
387
+ "dispatch": { "x": "ceil(N / 64)", "y": "ceil(rowCount / 32)", "z": "batchCount" }
388
+ }
389
+ ]
390
+ },
391
+ {
392
+ "id": "subgroup_matrix",
393
+ "priority": 10,
394
+ "when": ["f16Ok(dtypes.T)", "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, 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, 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, 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, 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, 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"],
395
+ "requires": {
396
+ "features": ["subgroups", "chromium-experimental-subgroup-matrix"],
397
+ "subgroupMatrixConfigs": [
398
+ { "componentType": "f16", "M": 8, "N": 8, "K": 8 },
399
+ { "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }
400
+ ]
401
+ },
402
+ "derive": {
403
+ "outScalar": "\"f16\" if dtypes.T == \"f16\" else \"f32\"",
404
+ "transA": "attrs.transA != 0",
405
+ "transB": "attrs.transB != 0",
406
+ "transBatchA": "false",
407
+ "rowCount": "(dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2))",
408
+ "K": "(dim(shapes.A, ranks.A - 2) if attrs.transA != 0 else dim(shapes.A, ranks.A - 1))",
409
+ "N": "(dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1))",
410
+ "batchCount": "numel(shapes.Y) / ((dim(shapes.A, ranks.A - 1) if attrs.transA != 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)))"
411
+ },
412
+ "passes": [
413
+ {
414
+ "id": "main",
415
+ "name": "TransposeMatMul.SubgroupMatrix",
416
+ "shader": "fused-matmul-subgroup-matrix.wgsl.jinja",
417
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
418
+ "dispatch": { "x": "ceil(N / 64)", "y": "ceil(rowCount / 32)", "z": "numel(shapes.Y) / (rowCount * N)" }
419
+ }
420
+ ]
421
+ },
422
+ {
423
+ "id": "broadcast_rank4_tiled_reg",
424
+ "priority": 6,
425
+ "when": ["f16Ok(dtypes.T)", "attrs.transA == 0", "attrs.transB == 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)"],
426
+ "passes": [
427
+ {
428
+ "id": "main",
429
+ "name": "TransposeMatMul.BroadcastRank4TiledReg",
430
+ "shader": "matmul-tiled-general-reg.wgsl.jinja",
431
+ "derive": {
432
+ "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []",
433
+ "aRank": "ranks.A if ranks.B > 2 else 2",
434
+ "K": "dim(shapes.A, ranks.A - 1)",
435
+ "bShape": "shapes.B",
436
+ "transBatchA": "false"
437
+ },
438
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
439
+ "dispatch": {
440
+ "x": "ceil(dim(shapes.B, ranks.B - 1) / registerTile)",
441
+ "y": "ceil(rowCount / registerTile)",
442
+ "z": "batchCount"
443
+ }
444
+ }
445
+ ],
446
+ "derive": {
447
+ "rowCount": "outer(shapes.A, 3) if ranks.B == 2 else dim(shapes.A, 2)",
448
+ "batchCount": "numel(shapes.Y) / (dim(shapes.A, 2) * dim(shapes.B, ranks.B - 1)) if ranks.B > 2 else 1"
449
+ }
450
+ },
451
+ {
452
+ "id": "plain_rank2_tiled_reg",
453
+ "priority": 4,
454
+ "demoteWhen": ["plainRank2RegDeepPreferredTier"],
455
+ "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)"],
456
+ "passes": [
457
+ {
458
+ "id": "main",
459
+ "name": "TransposeMatMul.PlainRank2TiledReg",
460
+ "shader": "matmul-tiled-general-reg.wgsl.jinja",
461
+ "derive": {
462
+ "rowCount": "dim(shapes.A, 0)",
463
+ "K": "dim(shapes.A, 1)",
464
+ "bShape": "shapes.B",
465
+ "transBatchA": "false"
466
+ },
467
+ "bindings": ["a", "b_scalar", "y", "params_rows"],
468
+ "dispatch": {
469
+ "x": "ceil(dim(shapes.B, 1) / registerTile)",
470
+ "y": "ceil(dim(shapes.A, 0) / registerTile)",
471
+ "z": 1
472
+ }
473
+ }
474
+ ]
475
+ },
476
+ {
477
+ "id": "tiled",
478
+ "priority": 0,
479
+ "when": ["ranks.A >= 1", "ranks.B >= 1", "f16Ok(dtypes.T)", "(dim(shapes.A, 0) if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA == 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, ranks.B - 2)))", "ceil((1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 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, 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, ranks.A - 2))) * (1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else dim(shapes.B, ranks.B - 2)))) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
480
+ "passes": [
481
+ {
482
+ "id": "main",
483
+ "name": "TransposeMatMul.Tiled",
484
+ "shader": "matmul-tiled-general.wgsl.jinja",
485
+ "derive": {
486
+ "aShape": "shapes.A",
487
+ "bShape": "shapes.B",
488
+ "transA": "attrs.transA != 0",
489
+ "transB": "attrs.transB != 0",
490
+ "transBatchA": "false",
491
+ "transBatchB": "false"
492
+ },
493
+ "bindings": ["a", "b_scalar", "y"],
494
+ "dispatch": {
495
+ "x": "ceil((1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else dim(shapes.B, ranks.B - 2))) / generalTile)",
496
+ "y": "ceil((1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2))) / generalTile)",
497
+ "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, ranks.A - 2))) * (1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else dim(shapes.B, ranks.B - 2))))"
498
+ }
499
+ }
500
+ ]
501
+ }
502
+ ]
503
+ }
build/webgpu/matmul-band-vec4.wgsl.jinja ADDED
@@ -0,0 +1,80 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // Band GEMM y[M, N] = a[M, K] @ B[K, N]. Each lane owns one vec4 column group
2
+ // and carries one accumulator per row, so a loaded B word feeds all M row
3
+ // accumulators and each A value is reused across four adjacent output columns.
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 %}
15
+ {% set alphaValue = alpha if alpha is defined else 1 %}
16
+ {% if kSplitsValue > 1 %}
17
+ const K_SPLITS: u32 = {{ kSplitsValue }}u;
18
+ const K_PER_SPLIT: u32 = (K + K_SPLITS - 1u) / K_SPLITS;
19
+ {% endif %}
20
+
21
+ // The rows drain through one 32 x SLICES array in turn, so the workgroup
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>
29
+ ) {
30
+ let lane = lid.x;
31
+ let slice = lid.y;
32
+ let cg = workgroup_id.x * LANES + lane;
33
+ {% if kSplitsValue > 1 %}
34
+ let a_base = 0u;
35
+ let b_base = 0u;
36
+ let y_base = workgroup_id.y * ({{ M }}u * N4);
37
+ let k_begin = workgroup_id.y * K_PER_SPLIT;
38
+ let k_end = min(K, k_begin + K_PER_SPLIT);
39
+ {% else %}
40
+ let a_base = 0u;
41
+ let b_base = 0u;
42
+ let y_base = 0u;
43
+ {% endif %}
44
+ {% for r in range(M) %}
45
+ var acc{{ r }} = vec4<f32>(0.0);
46
+ {% endfor %}
47
+ if (cg < N4) {
48
+ {% if kSplitsValue > 1 %}
49
+ for (var k = k_begin + slice; k < k_end; k = k + SLICES) {
50
+ {% else %}
51
+ for (var k = slice; k < K; k = k + SLICES) {
52
+ {% endif %}
53
+ let bv = vec4<f32>(b[b_base + k * N4 + cg]);
54
+ {% for r in range(M) %}
55
+ acc{{ r }} = acc{{ r }} + f32(a[a_base + {{ r }}u * K + k]) * bv;
56
+ {% endfor %}
57
+ }
58
+ }
59
+ {% for r in range(M) %}
60
+ partials[slice * LANES + lane] = acc{{ r }};
61
+ workgroupBarrier();
62
+ if (slice == 0u && cg < N4) {
63
+ var total = partials[lane];
64
+ for (var s = 1u; s < SLICES; s = s + 1u) {
65
+ total = total + partials[s * LANES + lane];
66
+ }
67
+ {% if alphaValue != 1 %}
68
+ total = total * {{ alphaValue }};
69
+ {% endif %}
70
+ {% if kSplitsValue > 1 %}
71
+ {{ outputBuffer }}[y_base + {{ r }}u * N4 + cg] = total;
72
+ {% else %}
73
+ {{ outputBuffer }}[y_base + {{ r }}u * N4 + cg] = vec4<{{ T }}>(total);
74
+ {% endif %}
75
+ }
76
+ {% if not loop.last %}
77
+ workgroupBarrier();
78
+ {% endif %}
79
+ {% endfor %}
80
+ }
build/webgpu/matmul-subgroup-matrix-ext.wgsl.jinja ADDED
@@ -0,0 +1,355 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // Subgroup-matrix matmul over row-major, batch-outermost operands, with alpha,
2
+ // dense/broadcast batching and guarded K/N tails under `generalAddressing`, and
3
+ // an optional fused bias on the direct dense path that omits it.
4
+ enable subgroups;
5
+ {% if pinSubgroupSize32 %}
6
+ enable subgroup_size_control;
7
+ {% endif %}
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
+
34
+ {% if STATIC_M %}
35
+ const M: u32 = {{ M }}u;
36
+ {% endif %}
37
+ const K: u32 = {{ K }}u;
38
+ const N: u32 = {{ N }}u;
39
+ const BATCH_COUNT: u32 = {{ batchCount if batchCount is defined else 1 }}u;
40
+ {% if SPLIT_K > 1 %}
41
+ const SPLIT_K: u32 = {{ SPLIT_K }}u;
42
+ const K_PER_SPLIT: u32 = {{ kPerSplit }}u;
43
+ {% endif %}
44
+ const A_M_STRIDE: u32 = {{ aMStride }}u;
45
+ const A_K_STRIDE: u32 = {{ aKStride }}u;
46
+ const B_K_STRIDE: u32 = {{ bKStride }}u;
47
+ const B_N_STRIDE: u32 = {{ bNStride }}u;
48
+ {% if TAIL %}const K_FULL: u32 = (K / 32u) * 32u;
49
+ {% endif %}
50
+ const ALPHA: f32 = f32({{ alpha }});
51
+ {% if STATIC_M %}
52
+ const C_BATCH_STRIDE: u32 = M * N;
53
+ {% endif %}
54
+ const TILE_COLS: u32 = 64u;
55
+ const TILE_ROWS: u32 = 32u;
56
+ const TILE_K: u32 = 32u;
57
+ const SUB_COLS: u32 = 32u;
58
+ const SUB_ROWS: u32 = 16u;
59
+
60
+ var<workgroup> tile_A: array<{{ operandScalar }}, 32 * 32>;
61
+ var<workgroup> tile_B: array<{{ operandScalar }}, 64 * 32>;
62
+ var<workgroup> scratch: array<array<array<{{ accScalar }}, 64>, 4>, 4>;
63
+
64
+ fn loadSHMA(a_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
65
+ let a_global = tile_base + row;
66
+ {% if not STATIC_M %}
67
+ let row_in = a_global < params.M;
68
+ {% endif %}
69
+ let col = c_idx * 8u;
70
+ for (var col_offset = 0u; col_offset < 8u; col_offset = col_offset + 1u) {
71
+ let k = k_idx + col + col_offset;
72
+ if ({{ ROW_TEST }}) {
73
+ {% if operandScalar == "f16" %}
74
+ tile_A[row * TILE_K + col + col_offset] = f16(a[a_base + a_global * A_M_STRIDE + k * A_K_STRIDE]);
75
+ {% else %}
76
+ tile_A[row * TILE_K + col + col_offset] = f32(a[a_base + a_global * A_M_STRIDE + k * A_K_STRIDE]);
77
+ {% endif %}
78
+ } else {
79
+ {% if operandScalar == "f16" %}
80
+ tile_A[row * TILE_K + col + col_offset] = 0.0h;
81
+ {% else %}
82
+ tile_A[row * TILE_K + col + col_offset] = 0.0;
83
+ {% endif %}
84
+ }
85
+ }
86
+ }
87
+ {% if GENERAL and TAIL %}
88
+
89
+ fn loadSHMAKTail(a_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
90
+ let a_global = tile_base + row;
91
+ let row_in = a_global < {{ ROWS }};
92
+ let col = c_idx * 8u;
93
+ for (var col_offset = 0u; col_offset < 8u; col_offset = col_offset + 1u) {
94
+ let k = k_idx + col + col_offset;
95
+ if (row_in && k < K) {
96
+ tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(a[a_base + a_global * A_M_STRIDE + k * A_K_STRIDE]);
97
+ } else {
98
+ tile_A[row * TILE_K + col + col_offset] = {{ operandScalar }}(0);
99
+ }
100
+ }
101
+ }
102
+
103
+ {% endif %}
104
+ fn loadSHMB(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
105
+ let b_col = tile_base + row;
106
+ let col = c_idx * 16u;
107
+ for (var i = 0u; i < 16u; i = i + 1u) {
108
+ let k = k_idx + col + i;
109
+ {% if TAIL %}
110
+ let b_safe = min(b_col, N - 1u);
111
+ tile_B[row * TILE_K + col + i] = {{ operandScalar }}(b[b_base + k * B_K_STRIDE + b_safe * B_N_STRIDE]);
112
+ {% else %}
113
+ {% if operandScalar == "f16" %}
114
+ tile_B[row * TILE_K + col + i] = f16(b[b_base + k * B_K_STRIDE + b_col * B_N_STRIDE]);
115
+ {% else %}
116
+ tile_B[row * TILE_K + col + i] = f32(b[b_base + k * B_K_STRIDE + b_col * B_N_STRIDE]);
117
+ {% endif %}
118
+ {% endif %}
119
+ }
120
+ }
121
+ {% if GENERAL and TAIL %}
122
+
123
+ fn loadSHMBKTail(b_base: u32, tile_base: u32, k_idx: u32, row: u32, c_idx: u32) {
124
+ let b_col = min(tile_base + row, N - 1u);
125
+ let col = c_idx * 16u;
126
+ for (var i = 0u; i < 16u; i = i + 1u) {
127
+ let k = k_idx + col + i;
128
+ if (k < K) {
129
+ tile_B[row * TILE_K + col + i] = {{ operandScalar }}(b[b_base + k * B_K_STRIDE + b_col * B_N_STRIDE]);
130
+ } else {
131
+ tile_B[row * TILE_K + col + i] = {{ operandScalar }}(0);
132
+ }
133
+ }
134
+ }
135
+
136
+ {% endif %}
137
+ {% set needsColBase = hasBias or (GENERAL and TAIL) %}
138
+ fn storeOutput(offset: u32{% if needsColBase %}, col_base: u32{% endif %}, row: u32, col: u32, src_slot: u32, row_limit: i32) {
139
+ if (row_limit > 0 && row < u32(row_limit)) {
140
+ let col2 = col + 1u;
141
+ {% for block in range(4) %}
142
+ {% if TAIL %}
143
+ if (col_base + col + {{ block * 8 }}u < N) {
144
+ {% endif %}
145
+ {{ OUT }}[offset + row * N + col + {{ block * 8 }}u] = {{ OUT_SCALAR }}(
146
+ ALPHA * scratch[src_slot][{{ block }}][row * 8u + col]
147
+ );
148
+ {% if TAIL %}
149
+ }
150
+ if (col_base + col2 + {{ block * 8 }}u < N) {
151
+ {% endif %}
152
+ {{ OUT }}[offset + row * N + col2 + {{ block * 8 }}u] = {{ OUT_SCALAR }}(
153
+ ALPHA * scratch[src_slot][{{ block }}][row * 8u + col2]
154
+ );
155
+ {% if TAIL %}
156
+ }
157
+ {% endif %}
158
+ {% endfor %}
159
+ }
160
+ }
161
+
162
+ @compute @workgroup_size(128, 1, 1){{ " @subgroup_size(32)" if pinSubgroupSize32 else "" }}
163
+ fn main(
164
+ @builtin(workgroup_id) workgroup_id: vec3<u32>,
165
+ @builtin(num_workgroups) num_wg: vec3<u32>,
166
+ @builtin(local_invocation_index) local_idx: u32,
167
+ @builtin(subgroup_invocation_id) sg_id: u32,
168
+ @builtin(subgroup_size) sg_size: u32
169
+ ) {
170
+ {% if not STATIC_M %}
171
+ // The row count arrives per call; the strides it scales are formed here.
172
+ let M = params.M;
173
+ let C_BATCH_STRIDE = M * N;
174
+ {% endif %}
175
+ let b_global_base = workgroup_id.x * TILE_COLS;
176
+
177
+ let subtile_id = local_idx / sg_size;
178
+ let subtile_idx = subtile_id / 2u;
179
+ let subtile_idy = subtile_id % 2u;
180
+ let base_A = subtile_idy * SUB_ROWS;
181
+ let base_B = subtile_idx * SUB_COLS;
182
+
183
+ // Grid-stride over both the M-tile (y) and batch (z) axes so the dispatch stays
184
+ // <= maxComputeWorkgroupsPerDimension per dimension even when ceil(M/TILE_ROWS) or BATCH_COUNT exceed the
185
+ // limit. workgroup_size.z = 1 so num_wg.z is the batch dispatch stride, and
186
+ // num_wg.y * TILE_ROWS is the row-tile dispatch stride. The loop bounds (M and
187
+ // BATCH_COUNT are compile-time / uniform; num_wg and workgroup_id are uniform)
188
+ // are workgroup-uniform, so the trailing workgroupBarrier()s and the subgroup
189
+ // matrix operations stay reconverged. When neither axis is clamped,
190
+ // num_wg.y * TILE_ROWS > M and num_wg.z > BATCH_COUNT, so each loop executes
191
+ // exactly once at workgroup_id.y/workgroup_id.z.
192
+ let row_tile_stride = num_wg.y * TILE_ROWS;
193
+ for (var a_global_base = workgroup_id.y * TILE_ROWS; a_global_base < M; a_global_base += row_tile_stride) {
194
+ {% if SPLIT_K > 1 %}
195
+ // Split-K maps z to (batch, K segment), increasing the independent workgroup
196
+ // count for narrow matrices. Every segment owns a disjoint contiguous K range.
197
+ for (var batch_split = workgroup_id.z; batch_split < BATCH_COUNT * SPLIT_K; batch_split += num_wg.z) {
198
+ let batch = batch_split / SPLIT_K;
199
+ let split_id = batch_split - batch * SPLIT_K;
200
+ {% else %}
201
+ // workgroup_size.z = 1, so num_wg.z is the dispatch stride over the batch axis.
202
+ for (var batch = workgroup_id.z; batch < BATCH_COUNT; batch += num_wg.z) {
203
+ {% endif %}
204
+ {% set hasBatchCoord = namespace(value=false) %}
205
+ {% for i in range(batchRank) %}
206
+ {% set axis = batchRank - 1 - i %}
207
+ {% set aAxis = axis - (batchRank - aBatchLen) %}
208
+ {% set bAxis = axis - (batchRank - bBatchLen) %}
209
+ {% set aDim = aDims[aAxis] %}
210
+ {% set bDim = bShape[bAxis] if bAxis >= 0 else 1 %}
211
+ {% if aDim > 1 or bDim > 1 %}{% set hasBatchCoord.value = true %}{% endif %}
212
+ {% endfor %}
213
+ // Right-aligned broadcast offsets, decomposed from the flattened output batch.
214
+ {% if hasBatchCoord.value %}
215
+ var zTmp = batch;
216
+ {% endif %}
217
+ var a_base: u32 = 0u;
218
+ var b_base: u32 = 0u;
219
+ {% for i in range(batchRank) %}
220
+ {% set axis = batchRank - 1 - i %}
221
+ {% set aAxis = axis - (batchRank - aBatchLen) %}
222
+ {% set bAxis = axis - (batchRank - bBatchLen) %}
223
+ {% set aDim = aDims[aAxis] %}
224
+ {% set bDim = bShape[bAxis] if bAxis >= 0 else 1 %}
225
+ {% set outDim = aDim if aDim >= bDim else bDim %}
226
+ {% set aStride = namespace(v=1) %}
227
+ {% 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 %}
228
+ {% set bStride = namespace(v=1) %}
229
+ {% 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 %}
230
+ {% if outDim > 1 %}
231
+ let c{{ axis }} = zTmp % {{ outDim }}u;
232
+ zTmp = zTmp / {{ outDim }}u;
233
+ {% if aStride.v != 0 %} a_base = a_base + c{{ axis }} * {% if aStride.v != 1 %}{{ aStride.v }}u * {% endif %}M * K;
234
+ {% endif %}
235
+ {% if bStride.v != 0 %} b_base = b_base + c{{ axis }} * {{ bStride.v }}u;
236
+ {% endif %}
237
+ {% endif %}
238
+ {% endfor %}
239
+ {% if SPLIT_K > 1 %}
240
+ let c_base = (batch * SPLIT_K + split_id) * C_BATCH_STRIDE;
241
+ {% else %}
242
+ let c_base = batch * C_BATCH_STRIDE;
243
+ {% endif %}
244
+
245
+ var matC00: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
246
+ var matC01: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
247
+ var matC02: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
248
+ var matC03: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
249
+ var matC10: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
250
+ var matC11: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
251
+ var matC12: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
252
+ var matC13: subgroup_matrix_result<{{ accScalar }}, 8, 8>;
253
+
254
+ {% if SPLIT_K > 1 %}
255
+ let k_begin = split_id * K_PER_SPLIT;
256
+ let k_end = min(k_begin + K_PER_SPLIT, K);
257
+ for (var kidx = k_begin; kidx < k_end; kidx = kidx + TILE_K) {
258
+ {% else %}
259
+ for (var kidx = 0u; kidx < {% if GENERAL and TAIL %}K_FULL{% else %}K{% endif %}; kidx = kidx + TILE_K) {
260
+ {% endif %}
261
+ loadSHMA(a_base, a_global_base, kidx, local_idx / 4u, local_idx % 4u);
262
+ loadSHMB(b_base, b_global_base, kidx, local_idx / 2u, local_idx % 2u);
263
+ workgroupBarrier();
264
+
265
+ for (var step = 0u; step < TILE_K; step = step + 8u) {
266
+ {% set directInputs = false %}
267
+ let matrix_a_offset = subtile_idy * SUB_ROWS * TILE_K + step;
268
+ {% for r in range(2) %}
269
+ 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" }});
270
+ {% endfor %}
271
+
272
+ let matrix_b_offset = subtile_idx * SUB_COLS * TILE_K + step;
273
+ {% for c in range(4) %}
274
+ 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" }});
275
+ {% endfor %}
276
+
277
+ matC00 = subgroupMatrixMultiplyAccumulate(matA0, matB0, matC00);
278
+ matC01 = subgroupMatrixMultiplyAccumulate(matA0, matB1, matC01);
279
+ matC02 = subgroupMatrixMultiplyAccumulate(matA0, matB2, matC02);
280
+ matC03 = subgroupMatrixMultiplyAccumulate(matA0, matB3, matC03);
281
+ matC10 = subgroupMatrixMultiplyAccumulate(matA1, matB0, matC10);
282
+ matC11 = subgroupMatrixMultiplyAccumulate(matA1, matB1, matC11);
283
+ matC12 = subgroupMatrixMultiplyAccumulate(matA1, matB2, matC12);
284
+ matC13 = subgroupMatrixMultiplyAccumulate(matA1, matB3, matC13);
285
+ }
286
+ workgroupBarrier();
287
+ }
288
+ {% if GENERAL and TAIL %}
289
+ if (K_FULL < K) {
290
+ loadSHMAKTail(a_base, a_global_base, K_FULL, local_idx / 4u, local_idx % 4u);
291
+ loadSHMBKTail(b_base, b_global_base, K_FULL, local_idx / 2u, local_idx % 2u);
292
+ workgroupBarrier();
293
+
294
+ for (var step = 0u; step < TILE_K; step = step + 8u) {
295
+ {% set directInputs = false %}
296
+ let matrix_a_offset = subtile_idy * SUB_ROWS * TILE_K + step;
297
+ {% for r in range(2) %}
298
+ 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" }});
299
+ {% endfor %}
300
+
301
+ let matrix_b_offset = subtile_idx * SUB_COLS * TILE_K + step;
302
+ {% for c in range(4) %}
303
+ 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" }});
304
+ {% endfor %}
305
+
306
+ matC00 = subgroupMatrixMultiplyAccumulate(matA0, matB0, matC00);
307
+ matC01 = subgroupMatrixMultiplyAccumulate(matA0, matB1, matC01);
308
+ matC02 = subgroupMatrixMultiplyAccumulate(matA0, matB2, matC02);
309
+ matC03 = subgroupMatrixMultiplyAccumulate(matA0, matB3, matC03);
310
+ matC10 = subgroupMatrixMultiplyAccumulate(matA1, matB0, matC10);
311
+ matC11 = subgroupMatrixMultiplyAccumulate(matA1, matB1, matC11);
312
+ matC12 = subgroupMatrixMultiplyAccumulate(matA1, matB2, matC12);
313
+ matC13 = subgroupMatrixMultiplyAccumulate(matA1, matB3, matC13);
314
+ }
315
+ workgroupBarrier();
316
+ }
317
+
318
+ {% endif %}
319
+ // The four scratch banks are reused across the two row-groups, and each is written
320
+ // by a collective subgroupMatrixStore then read CROSS-LANE by storeOutput. Barriers
321
+ // give the reads visibility of the store AND stop the second row-group's store from
322
+ // clobbering the first's still-in-flight readback when a partial final M-tile
323
+ // diverges storeOutput's guard. Without both barriers the last valid row can be corrupted.
324
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][0], 0u, matC00, 8u);
325
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][1], 0u, matC01, 8u);
326
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][2], 0u, matC02, 8u);
327
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][3], 0u, matC03, 8u);
328
+ workgroupBarrier();
329
+ let row = sg_id / 4u;
330
+ let col = (sg_id % 4u) * 2u;
331
+ let col_base = b_global_base + base_B;
332
+ var matrix_c_offset = c_base + (a_global_base + base_A) * N + col_base;
333
+ var row_limit = i32(M) - i32(a_global_base + base_A);
334
+ storeOutput(matrix_c_offset{% if needsColBase %}, col_base{% endif %}, row, col, subtile_id, row_limit);
335
+ workgroupBarrier();
336
+
337
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][0], 0u, matC10, 8u);
338
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][1], 0u, matC11, 8u);
339
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][2], 0u, matC12, 8u);
340
+ subgroupMatrixStore<row_major>(&scratch[subtile_id][3], 0u, matC13, 8u);
341
+ workgroupBarrier();
342
+ matrix_c_offset = matrix_c_offset + 8u * N;
343
+ row_limit = i32(M) - i32(a_global_base + base_A + 8u);
344
+ storeOutput(matrix_c_offset{% if needsColBase %}, col_base{% endif %}, row, col, subtile_id, row_limit);
345
+
346
+ // Re-stage workgroup tiles/scratch before the next batch iteration reuses them.
347
+ workgroupBarrier();
348
+ }
349
+ // Re-stage workgroup tiles/scratch before the next M-tile iteration reuses them.
350
+ // The loop bound is workgroup-uniform (M is a compile-time const or a uniform
351
+ // read, and num_wg.y and workgroup_id.y are uniform), so every invocation
352
+ // reaches this barrier together.
353
+ workgroupBarrier();
354
+ }
355
+ }
build/webgpu/matmul-tiled-general-reg.wgsl.jinja ADDED
@@ -0,0 +1,210 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {{ env.wgsl.resourceDeclarations }}
2
+
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
+ {% set M = aShape[aR-2] %}{% set K = aShape[aR-1] %}{% endif %}
26
+ {% set N = bShape[bR-1] %}
27
+ {% set aMStride = K %}{% set aKStride = 1 %}{% set bKStride = 1 if bTransposed is defined and bTransposed else bShape[bR-1] %}
28
+ {% set bNStride = bShape[bR-2] if bTransposed is defined and bTransposed else 1 %}{% set is_int = (scalar == "i32" or scalar == "u32") %}
29
+ {% if is_int %}
30
+ // Integer operands accumulate in their integer type, avoiding f32 rounding of
31
+ // values outside the exact 24-bit significand range.
32
+ {% endif %}
33
+ {% set accT = scalar if is_int else "f32" %}
34
+ {% set outScalar = outScalar if outScalar is defined else scalar %}
35
+ {% set splitK = splitK if splitK is defined else 1 %}
36
+ {% if scalar == "f16" %}
37
+ // f16 operands stay packed in workgroup memory and widen on shared load;
38
+ // accumulation remains f32.
39
+ {% endif %}
40
+ {% set tileT = scalar if scalar == "f16" else accT %}
41
+ {% set kTile = 16 %}
42
+ {% if STATIC_M %}
43
+ const M: u32 = {{ M }}u;
44
+ {% endif %}
45
+ const K: u32 = {{ K }}u;
46
+ const N: u32 = {{ N }}u;
47
+ const A_M_STRIDE: u32 = {{ aMStride }}u;
48
+ const A_K_STRIDE: u32 = {{ aKStride }}u;
49
+ const B_K_STRIDE: u32 = {{ bKStride }}u;
50
+ const B_N_STRIDE: u32 = {{ bNStride }}u;
51
+ {% if is_int %}const ALPHA: {{ accT }} = {{ accT }}(1);{% else %}const ALPHA: f32 = f32({{ alpha }});{% endif %}
52
+ // A 4x4 micro-tile over a 64x64 output tile reuses each staged operand across
53
+ // four accumulators. It increases arithmetic work per load without the large
54
+ // per-thread accumulator footprint of an 8x8 micro-tile.
55
+ {% set microTile = 4 %}
56
+ {% set lanes = 16 %}
57
+ const BK: u32 = {{ kTile }}u;
58
+ const BM: u32 = {{ registerTile }}u;
59
+ const BN: u32 = {{ registerTile }}u;
60
+ const TM: u32 = {{ microTile }}u; // per-thread micro-tile rows
61
+ const TN: u32 = {{ microTile }}u; // per-thread micro-tile cols
62
+ {% if splitK > 1 %}
63
+ const SPLIT_K: u32 = {{ splitK }}u;
64
+ const K_PER_SPLIT: u32 = {{ kPerSplit }}u;
65
+
66
+ {% endif %}
67
+ const K_VECS: u32 = BK / 4u;
68
+ var<workgroup> tileA: array<array<vec4<{{ tileT }}>, K_VECS>, BM>; // A[m][k/4]
69
+ var<workgroup> tileB: array<array<vec4<{{ tileT }}>, K_VECS>, BN>; // B[n][k/4]
70
+ {% set hasBatchCoord = namespace(value=false) %}
71
+ {% for i in range(batchRank) %}
72
+ {% set axis = batchRank - 1 - i %}
73
+ {% set aAxis = axis - (batchRank - aBatchLen) %}
74
+ {% set bAxis = axis - (batchRank - bBatchLen) %}
75
+ {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
76
+ {% set bStored = bAxis %}
77
+ {% set aDim = aDims[aStored] %}
78
+ {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
79
+ {% if aDim > 1 or bDim > 1 %}{% set hasBatchCoord.value = true %}{% endif %}
80
+ {% endfor %}
81
+
82
+ @compute @workgroup_size({{ lanes }}, {{ lanes }}, 1)
83
+ fn main(
84
+ @builtin(workgroup_id) wg: vec3<u32>,
85
+ @builtin(local_invocation_id) lid: vec3<u32>
86
+ ) {
87
+ {% if not STATIC_M %}
88
+ let M = params.M;
89
+ {% endif %}
90
+ let mBase = wg.y * BM;
91
+ let nBase = wg.x * BN;
92
+ let li = lid.y * {{ lanes }}u + lid.x;
93
+
94
+ let zOut = wg.z;
95
+ {% if splitK > 1 %}
96
+ let splitId = wg.z % SPLIT_K;
97
+ {% else %}
98
+ {% if hasBatchCoord.value %}
99
+ var zTmp = wg.z;
100
+ {% endif %}
101
+ {% endif %}
102
+ var aBatchOff: u32 = 0u;
103
+ var bBatchOff: u32 = 0u;
104
+ {% for i in range(batchRank) %}
105
+ {% set axis = batchRank - 1 - i %}
106
+ {% set aAxis = axis - (batchRank - aBatchLen) %}
107
+ {% set bAxis = axis - (batchRank - bBatchLen) %}
108
+ {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
109
+ {% set bStored = bAxis %}
110
+ {% set aDim = aDims[aStored] %}
111
+ {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
112
+ {% set outDim = aDim if aDim >= bDim else bDim %}
113
+ {% set aStride = namespace(v=1) %}
114
+ {% 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 %}
115
+ {% set bStride = namespace(v=1) %}
116
+ {% 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 outDim > 1 %}
117
+ let c{{ axis }} = zTmp % {{ outDim }}u;
118
+ zTmp = zTmp / {{ outDim }}u;
119
+ {% 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 %};
120
+ {% endif %}
121
+ {% if bStride.v != 0 %} bBatchOff = bBatchOff + c{{ axis }} * {{ bStride.v }}u;
122
+ {% endif %}
123
+ {% endif %}
124
+ {% endfor %}
125
+
126
+ var acc: array<{{ accT }}, TM * TN>; // [ti*TN + tj] for the TMxTN micro-tile
127
+ for (var i: u32 = 0u; i < TM * TN; i = i + 1u) { acc[i] = {{ accT }}(0); }
128
+
129
+ {% if splitK > 1 %}
130
+ let kStart = splitId * K_PER_SPLIT;
131
+ let numTiles = K_PER_SPLIT / BK;
132
+ {% else %}
133
+ let numTiles = (K + BK - 1u) / BK;
134
+ {% endif %}
135
+ for (var kt: u32 = 0u; kt < numTiles; kt = kt + 1u) {
136
+ {% if splitK > 1 %}
137
+ let kBase = kStart + kt * BK;
138
+ {% else %}
139
+ let kBase = kt * BK;
140
+ {% endif %}
141
+ // Cooperative load: one vector word per lane per pass. A's lanes walk K, which
142
+ // it stores contiguously; B's walk N, which it stores contiguously.
143
+ for (var idx: u32 = li; idx < BM * K_VECS; idx = idx + {{ lanes * lanes }}u) {
144
+ let ar = idx / K_VECS;
145
+ let ac4 = idx % K_VECS;
146
+ let am = mBase + ar;
147
+ let ak = kBase + ac4 * 4u;
148
+ var aWord = vec4<{{ tileT }}>({{ tileT }}(0));
149
+ if (am < M) {
150
+ let aRowOff = aBatchOff + am * A_M_STRIDE;
151
+ {% for component in range(4) %}
152
+ if (ak + {{ component }}u < K) { aWord[{{ component }}u] = {{ tileT }}(a[aRowOff + (ak + {{ component }}u) * A_K_STRIDE]); }
153
+ {% endfor %}
154
+ }
155
+ tileA[ar][ac4] = aWord;
156
+ }
157
+ for (var idx: u32 = li; idx < BN * K_VECS; idx = idx + {{ lanes * lanes }}u) {
158
+ let bc = idx % BN;
159
+ let br4 = idx / BN;
160
+ let bn = nBase + bc;
161
+ let bk = kBase + br4 * 4u;
162
+ var bWord = vec4<{{ tileT }}>({{ tileT }}(0));
163
+ if (bn < N) {
164
+ let bColOff = bBatchOff + bn * B_N_STRIDE;
165
+ {% for component in range(4) %}
166
+ if (bk + {{ component }}u < K) { bWord[{{ component }}u] = {{ tileT }}(b[bColOff + (bk + {{ component }}u) * B_K_STRIDE]); }
167
+ {% endfor %}
168
+ }
169
+ tileB[bc][br4] = bWord;
170
+ }
171
+ workgroupBarrier();
172
+ {% set regT = accT %}{% filter indent(4, true) %}
173
+ {% set regAccumulator = "acc" %}
174
+ let aRow = lid.y * TM;
175
+ let bCol = lid.x * TN;
176
+ for (var kv: u32 = 0u; kv < BK / 4u; kv = kv + 1u) {
177
+ var av: array<vec4<{{ regT }}>, TM>;
178
+ var bv: array<vec4<{{ regT }}>, TN>;
179
+ for (var i: u32 = 0u; i < TM; i = i + 1u) { av[i] = vec4<{{ regT }}>(tileA[aRow + i][kv]); }
180
+ for (var j: u32 = 0u; j < TN; j = j + 1u) { bv[j] = vec4<{{ regT }}>(tileB[bCol + j][kv]); }
181
+ for (var i: u32 = 0u; i < TM; i = i + 1u) {
182
+ for (var j: u32 = 0u; j < TN; j = j + 1u) {
183
+ {% if regSequentialK is defined and regSequentialK %}
184
+ {% for component in range(4) %}
185
+ {{ regAccumulator }}[i * TN + j] = {{ regAccumulator }}[i * TN + j] + av[i][{{ component }}] * bv[j][{{ component }}];
186
+ {% endfor %}
187
+ {% else %}
188
+ {{ regAccumulator }}[i * TN + j] = {{ regAccumulator }}[i * TN + j] + dot(av[i], bv[j]);
189
+ {% endif %}
190
+ }
191
+ }
192
+ }
193
+ {% endfilter %}
194
+ workgroupBarrier();
195
+ }
196
+
197
+ let rowBase = zOut * M * N;
198
+ let m0 = mBase + lid.y * TM;
199
+ let n0 = nBase + lid.x * TN;
200
+ for (var ti: u32 = 0u; ti < TM; ti = ti + 1u) {
201
+ let m = m0 + ti;
202
+ if (m >= M) { continue; }
203
+ for (var tj: u32 = 0u; tj < TN; tj = tj + 1u) {
204
+ let n = n0 + tj;
205
+ if (n < N) {
206
+ y[rowBase + m * N + n] = {{ outScalar }}(ALPHA * acc[ti * TN + tj]);
207
+ }
208
+ }
209
+ }
210
+ }
build/webgpu/matmul-tiled-general.wgsl.jinja ADDED
@@ -0,0 +1,168 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {{ env.wgsl.resourceDeclarations }}
2
+
3
+ // Shared tiled matrix multiplication: Y = alpha * op(A) @ op(B), where op
4
+ // transposes the last two axes when requested. This bounds-checked kernel
5
+ // handles any M/K/N (no alignment requirement), all four transpose combinations,
6
+ // right-aligned batch broadcasting, 1-D operand promotion (M==1 / N==1), and empty
7
+ // K/M/N. transA/transB only change which stored stride the logical (m,k)/(k,n)
8
+ // walk, using compiled stride constants here; the batch broadcast strides are
9
+ // compiled the same way (a 0 literal means "broadcast / absent on that operand").
10
+ {% set aR = aRank %}
11
+ {% set bR = bRank %}
12
+ {% set aVec = (aR == 1) %}
13
+ {% set bVec = (bR == 1) %}
14
+ {% set aBatchLen = (aR - 2) if aR >= 2 else 0 %}
15
+ {% set bBatchLen = (bR - 2) if bR >= 2 else 0 %}
16
+ {% set batchRank = aBatchLen if aBatchLen >= bBatchLen else bBatchLen %}
17
+ /* transBatch maps stored [d0, d1, ..., dR-2, dR-1] to logical
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 transA %}{% set M = aShape[aR-1] %}{% set K = aShape[aR-2] %}
28
+ {% else %}{% set M = aShape[aR-2] %}{% set K = aShape[aR-1] %}{% endif %}
29
+ {% if bVec %}{% set N = 1 %}
30
+ {% elif transB %}{% set N = bShape[bR-2] %}
31
+ {% else %}{% set N = bShape[bR-1] %}{% endif %}
32
+ {% if aVec %}{% set aMStride = 0 %}{% set aKStride = 1 %}
33
+ {% elif transA %}{% set aMStride = 1 %}{% set aKStride = M %}
34
+ {% else %}{% set aMStride = K %}{% set aKStride = 1 %}{% endif %}
35
+ {% if bVec %}{% set bKStride = 1 %}{% set bNStride = 0 %}
36
+ {% elif transB %}{% set bKStride = 1 %}{% set bNStride = bShape[bR-1] %}
37
+ {% else %}{% set bKStride = bShape[bR-1] %}{% set bNStride = 1 %}{% endif %}
38
+
39
+ {% set is_int = (scalar == "i32" or scalar == "u32") %}
40
+ {% set accT = scalar if is_int else "f32" %}
41
+ {% set tileT = scalar if scalar == "f16" else accT %}
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;
46
+ const A_K_STRIDE: u32 = {{ aKStride }}u;
47
+ const B_K_STRIDE: u32 = {{ bKStride }}u;
48
+ const B_N_STRIDE: u32 = {{ bNStride }}u;
49
+ {% if is_int %}/* Integer matrix multiplication accumulates in the integer type. Widening
50
+ * through f32 would round values above 2^24. Integer MatMul has alpha = 1. */
51
+ const ALPHA: {{ accT }} = {{ accT }}(1);{% else %}const ALPHA: f32 = f32({{ alpha }});{% endif %}
52
+ // 2x2 register-blocked tile: 16x16 threads each compute a 2x2 micro-tile, for a
53
+ // 32x32 output tile per workgroup with K stepped in BK=16 chunks. Each loaded
54
+ // shared-mem element feeds 2 FMAs, favoring register reuse in the inner loop.
55
+ {% set lanes = 16 %}
56
+ const BK: u32 = {{ lanes }}u;
57
+ const BM: u32 = {{ generalTile }}u;
58
+ const BN: u32 = {{ generalTile }}u;
59
+
60
+ var<workgroup> tileA: array<array<{{ tileT }}, {{ lanes }}>, {{ generalTile }}>;
61
+ var<workgroup> tileB: array<array<{{ tileT }}, {{ generalTile }}>, {{ lanes }}>;
62
+ {% set hasBatchCoord = namespace(value=false) %}
63
+ {% for i in range(batchRank) %}
64
+ {% set axis = batchRank - 1 - i %}
65
+ {% set aAxis = axis - (batchRank - aBatchLen) %}
66
+ {% set bAxis = axis - (batchRank - bBatchLen) %}
67
+ {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
68
+ {% set bStored = (bAxis + 1) if (transBatchB and bAxis >= 0) else bAxis %}
69
+ {% set aDim = aDims[aStored] if aStored >= 0 else 1 %}
70
+ {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
71
+ {% if aDim > 1 or bDim > 1 %}{% set hasBatchCoord.value = true %}{% endif %}
72
+ {% endfor %}
73
+
74
+ @compute @workgroup_size({{ lanes }}, {{ lanes }}, 1)
75
+ fn main(
76
+ @builtin(workgroup_id) wg: vec3<u32>,
77
+ @builtin(local_invocation_id) lid: vec3<u32>
78
+ ) {
79
+ let mBase = wg.y * BM;
80
+ let nBase = wg.x * BN;
81
+ let li = lid.y * {{ lanes }}u + lid.x;
82
+
83
+ // Per-batch base offsets into A and B using right-aligned broadcast strides.
84
+ // Decompose the flat output-batch index from the innermost axis outward.
85
+ let zOut = wg.z;
86
+ {% if hasBatchCoord.value %}
87
+ var zTmp = wg.z;
88
+ {% endif %}
89
+ var aBatchOff: u32 = 0u;
90
+ var bBatchOff: u32 = 0u;
91
+ {% for i in range(batchRank) %}
92
+ {% set axis = batchRank - 1 - i %}
93
+ {% set aAxis = axis - (batchRank - aBatchLen) %}
94
+ {% set bAxis = axis - (batchRank - bBatchLen) %}
95
+ {% set aStored = (aAxis + 1) if (transBatchA and aAxis >= 0) else aAxis %}
96
+ {% set bStored = (bAxis + 1) if (transBatchB and bAxis >= 0) else bAxis %}
97
+ {% set aDim = aDims[aStored] if aStored >= 0 else 1 %}
98
+ {% set bDim = bShape[bStored] if bStored >= 0 else 1 %}
99
+ {% set outDim = aDim if aDim >= bDim else bDim %}
100
+ {% set aStride = namespace(v=1) %}
101
+ {% 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 %}
102
+ {% set bStride = namespace(v=1) %}
103
+ {% 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 %}
104
+ {% if outDim > 1 %}
105
+ {% if aStride.v != 0 or bStride.v != 0 %}
106
+ let c{{ axis }} = zTmp % {{ outDim }}u;
107
+ {% endif %}
108
+ zTmp = zTmp / {{ outDim }}u;
109
+ {% if aStride.v != 0 %} aBatchOff = aBatchOff + c{{ axis }} * {{ aStride.v }}u;
110
+ {% endif %}
111
+ {% if bStride.v != 0 %} bBatchOff = bBatchOff + c{{ axis }} * {{ bStride.v }}u;
112
+ {% endif %}
113
+ {% endif %}
114
+ {% endfor %}
115
+
116
+ var acc00: {{ accT }} = {{ accT }}(0);
117
+ var acc01: {{ accT }} = {{ accT }}(0);
118
+ var acc10: {{ accT }} = {{ accT }}(0);
119
+ var acc11: {{ accT }} = {{ accT }}(0);
120
+ let numTiles = (K + BK - 1u) / BK;
121
+ for (var kt: u32 = 0u; kt < numTiles; kt = kt + 1u) {
122
+ let kBase = kt * BK;
123
+ // Cooperative load: 32x16 A tile + 16x32 B tile, 256 threads x 2 each.
124
+ for (var e: u32 = 0u; e < 2u; e = e + 1u) {
125
+ let idx = li + e * {{ lanes * lanes }}u;
126
+ let ar = idx / BK;
127
+ let ac = idx % BK;
128
+ let am = mBase + ar;
129
+ let ak = kBase + ac;
130
+ if (am < M && ak < K) {
131
+ tileA[ar][ac] = {{ tileT }}(a[aBatchOff + am * A_M_STRIDE + ak * A_K_STRIDE]);
132
+ } else {
133
+ tileA[ar][ac] = {{ tileT }}(0);
134
+ }
135
+ let br = idx / BN;
136
+ let bc = idx % BN;
137
+ let bk = kBase + br;
138
+ let bn = nBase + bc;
139
+ if (bk < K && bn < N) {
140
+ tileB[br][bc] = {{ tileT }}(b[bBatchOff + bk * B_K_STRIDE + bn * B_N_STRIDE]);
141
+ } else {
142
+ tileB[br][bc] = {{ tileT }}(0);
143
+ }
144
+ }
145
+ workgroupBarrier();
146
+ for (var kk: u32 = 0u; kk < BK; kk = kk + 1u) {
147
+ let a0 = {{ accT }}(tileA[lid.y * 2u][kk]);
148
+ let a1 = {{ accT }}(tileA[lid.y * 2u + 1u][kk]);
149
+ let b0 = {{ accT }}(tileB[kk][lid.x * 2u]);
150
+ let b1 = {{ accT }}(tileB[kk][lid.x * 2u + 1u]);
151
+ acc00 = acc00 + a0 * b0;
152
+ acc01 = acc01 + a0 * b1;
153
+ acc10 = acc10 + a1 * b0;
154
+ acc11 = acc11 + a1 * b1;
155
+ }
156
+ workgroupBarrier();
157
+ }
158
+
159
+ let m0 = mBase + lid.y * 2u;
160
+ let m1 = m0 + 1u;
161
+ let n0 = nBase + lid.x * 2u;
162
+ let n1 = n0 + 1u;
163
+ let rowBase = zOut * M * N;
164
+ if (m0 < M && n0 < N) { y[rowBase + m0 * N + n0] = {{ scalar }}(ALPHA * acc00); }
165
+ if (m0 < M && n1 < N) { y[rowBase + m0 * N + n1] = {{ scalar }}(ALPHA * acc01); }
166
+ if (m1 < M && n0 < N) { y[rowBase + m1 * N + n0] = {{ scalar }}(ALPHA * acc10); }
167
+ if (m1 < M && n1 < N) { y[rowBase + m1 * N + n1] = {{ scalar }}(ALPHA * acc11); }
168
+ }
build/webgpu/matmul-vector-matrix-vec4.wgsl.jinja ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // GEMV specialization for y[N] = a[K] @ B[K, N]. Each workgroup owns 32
2
+ // consecutive vec4 column groups, or 128 output columns.
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>
17
+ ) {
18
+ let lane = lid.x;
19
+ let slice = lid.y;
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();
43
+ if (slice == 0u && cg < params.N4) {
44
+ var total = partials[lane];
45
+ for (var s = 1u; s < SLICES; s = s + 1u) {
46
+ total = total + partials[s * LANES + lane];
47
+ }
48
+ {% if alphaScale is defined and alphaScale != 1 %}
49
+ // This specialization bakes the output multiplier into the shader and
50
+ // applies it after combining the K slices.
51
+ total = total * f32({{ alphaScale }});
52
+ {% endif %}
53
+ {{ OUT }}[cg] = vec4<{{ T }}>(total);
54
+ }
55
+ }
build/webgpu/metadata.json ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "name": "com.microsoft.TransposeMatMul",
3
+ "id": "_com_microsoft_transposematmul_webgpu_45fbc93",
4
+ "version": 1,
5
+ "license": "Apache-2.0",
6
+ "backend": { "type": "webgpu" },
7
+ "digest": {
8
+ "algorithm": "sha256",
9
+ "files": {
10
+ "bench.json": "XEFUj4H5WRbzIOZky4TK16K+wPpc4vYBtmqD+f6fhvA=",
11
+ "fused-matmul-subgroup-matrix.wgsl.jinja": "fqGrZsIrSx70rZzmpWys4oTG+JSOnAbqtBdQ3xwQpSw=",
12
+ "manifest.json": "NeYkA+AcZ/38wnrBvqcXxIiJM296w8JwhwTrkX+l40w=",
13
+ "matmul-band-vec4.wgsl.jinja": "LKAs6A++OJF0ZEEM4JITZmKXkr/wo9gd5qrR3zDZy1g=",
14
+ "matmul-subgroup-matrix-ext.wgsl.jinja": "DJxg5GNgH4xumnum/RzuWA5x+JNzB0Tx7U5RMsnlm6o=",
15
+ "matmul-tiled-general-reg.wgsl.jinja": "X8kCdoMgwvyMiCnLl/UDpfpdVcYFqTLlj6+o+R3pJlQ=",
16
+ "matmul-tiled-general.wgsl.jinja": "9uRxra70oG70IJCiKqA0DG76W0zS46RxLdT6fndR48w=",
17
+ "matmul-vector-matrix-vec4.wgsl.jinja": "Qiv2AO8MWj1BAoPLMCH98oXBasrQYFWFh/IUhdhZp5I=",
18
+ "reduce-axis0-splitk-combine.wgsl.jinja": "Zf7tz8nrapi2KMZIbv8cT2f4pliJoJHj4iwS4hc/6ms=",
19
+ "test.json": "h+wSc5e5ergXuijy9nMDICwCmTvIi54cexuSorODaWM="
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
+ "m1_gemv_vec4": ["matmul-vector-matrix-vec4.wgsl.jinja"],
30
+ "rank2_band_vec4_splitk": ["matmul-band-vec4.wgsl.jinja", "reduce-axis0-splitk-combine.wgsl.jinja"],
31
+ "rank2_band_vec4": ["matmul-band-vec4.wgsl.jinja"],
32
+ "rank2_band_vec4_f32_preferred": ["matmul-band-vec4.wgsl.jinja"],
33
+ "subgroup_matrix_splitk": ["matmul-subgroup-matrix-ext.wgsl.jinja", "reduce-axis0-splitk-combine.wgsl.jinja"],
34
+ "subgroup_matrix_tail_broadcast": ["matmul-subgroup-matrix-ext.wgsl.jinja"],
35
+ "subgroup_matrix": ["fused-matmul-subgroup-matrix.wgsl.jinja"],
36
+ "broadcast_rank4_tiled_reg": ["matmul-tiled-general-reg.wgsl.jinja"],
37
+ "plain_rank2_tiled_reg": ["matmul-tiled-general-reg.wgsl.jinja"],
38
+ "tiled": ["matmul-tiled-general.wgsl.jinja"]
39
+ }
40
+ }
41
+ }
build/webgpu/reduce-axis0-splitk-combine.wgsl.jinja ADDED
@@ -0,0 +1,95 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // Pass 2 of the split-K column-wise reduction. One invocation per output column
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;
21
+
22
+ fn is_nan_f32(value: f32) -> bool {
23
+ let bits = bitcast<u32>(value);
24
+ return (bits & 0x7f800000u) == 0x7f800000u && (bits & 0x007fffffu) != 0u;
25
+ }
26
+ {% elif op == "max" or op == "min" %}
27
+ {{ wgsl_minmax_identity("reduction_identity", op) }}
28
+ {% endif %}
29
+
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" %}
39
+ // Merge SPLIT (segMax, segSumExp) pairs stably; carry NaN / +Inf markers.
40
+ var nan_value = 0.0;
41
+ var has_nan = false;
42
+ var global_max = F32_MIN;
43
+ for (var seg = 0u; seg < SPLIT; seg = seg + 1u) {
44
+ let nv = partials[(2u * SPLIT + seg) * params.cols + col];
45
+ if (nv != 0.0 || is_nan_f32(nv)) {
46
+ has_nan = true;
47
+ nan_value = nv;
48
+ }
49
+ global_max = max(global_max, partials[seg * params.cols + col]);
50
+ }
51
+ var sum = 0.0;
52
+ for (var seg = 0u; seg < SPLIT; seg = seg + 1u) {
53
+ let seg_max = partials[seg * params.cols + col];
54
+ let seg_sum = partials[(SPLIT + seg) * params.cols + col];
55
+ sum = sum + seg_sum * exp(seg_max - global_max);
56
+ }
57
+ let has_positive_inf = global_max > F32_MAX;
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" %}
64
+ var total = reduction_identity();
65
+ {% elif op == "prod" %}
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];
72
+ {% if op == "max" or op == "min" %}
73
+ total = {{ op }}(total, p);
74
+ {% elif op == "prod" %}
75
+ total = total * p;
76
+ {% else %}
77
+ total = total + p;
78
+ {% endif %}
79
+ }
80
+ {% if op == "l2" %}
81
+ y[col] = {{ yv }}sqrt(total){{ vy }};
82
+ {% elif op == "logsum" %}
83
+ y[col] = {{ yv }}log(total){{ vy }};
84
+ {% elif op == "mean" %}
85
+ y[col] = {{ yv }}total / f32(params.rows){{ vy }};
86
+ {% else %}
87
+ {% if outputF16 %}
88
+ y[col] = f16(total);
89
+ {% else %}
90
+ y[col] = total;
91
+ {% endif %}
92
+ {% endif %}
93
+ {% endif %}
94
+ }
95
+ }
build/webgpu/test.json ADDED
@@ -0,0 +1,1973 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "fixtureArrays": {
3
+ "ort_float32_broadcast_rank3_by_rank4_output_Y": [1, 3, 5, 33, 43, 53, 5, 23, 41, 85, 111, 137, 9, 43, 77, 137, 179, 221],
4
+ "ort_float32_rank3_by_rank2_output_Y": [20, 23, 26, 29, 56, 68, 80, 92, 92, 113, 134, 155, 128, 158, 188, 218],
5
+ "ort_float32_batched_rank4_input_A": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15]
6
+ },
7
+ "cases": [
8
+ {
9
+ "name": "ort_float32_broadcast_rank4_by_rank3",
10
+ "provenance": {
11
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
12
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
13
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
14
+ },
15
+ "inputs": {
16
+ "A": {
17
+ "dtype": "float32",
18
+ "shape": [3, 1, 1, 2],
19
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0] }
20
+ },
21
+ "B": {
22
+ "dtype": "float32",
23
+ "shape": [2, 2, 2],
24
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0] }
25
+ }
26
+ },
27
+ "outputs": {
28
+ "Y": {
29
+ "dtype": "float32",
30
+ "shape": [3, 2, 1, 2],
31
+ "tolerance": 0.000001,
32
+ "data": { "kind": "values", "values": [2.0, 3.0, 6.0, 7.0, 6.0, 11.0, 26.0, 31.0, 10.0, 19.0, 46.0, 55.0] }
33
+ }
34
+ }
35
+ },
36
+ {
37
+ "name": "ort_float32_broadcast_rank3_by_rank4",
38
+ "provenance": {
39
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
40
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
41
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
42
+ },
43
+ "inputs": {
44
+ "A": {
45
+ "dtype": "float32",
46
+ "shape": [2, 3, 2],
47
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0] }
48
+ },
49
+ "B": {
50
+ "dtype": "float32",
51
+ "shape": [3, 2, 2, 1],
52
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0] }
53
+ }
54
+ },
55
+ "outputs": {
56
+ "Y": {
57
+ "dtype": "float32",
58
+ "shape": [3, 2, 3, 1],
59
+ "tolerance": 0.000001,
60
+ "data": {
61
+ "kind": "values",
62
+ "values": { "$ref": "#/fixtureArrays/ort_float32_broadcast_rank3_by_rank4_output_Y" }
63
+ }
64
+ }
65
+ }
66
+ },
67
+ {
68
+ "name": "ort_float32_left_1d_batched_rhs",
69
+ "provenance": {
70
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
71
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
72
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
73
+ },
74
+ "inputs": {
75
+ "A": { "dtype": "float32", "shape": [2], "data": { "kind": "values", "values": [0.0, 1.0] } },
76
+ "B": {
77
+ "dtype": "float32",
78
+ "shape": [3, 2, 1],
79
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0] }
80
+ }
81
+ },
82
+ "outputs": {
83
+ "Y": {
84
+ "dtype": "float32",
85
+ "shape": [3, 1],
86
+ "tolerance": 0.000001,
87
+ "data": { "kind": "values", "values": [1.0, 3.0, 5.0] }
88
+ }
89
+ }
90
+ },
91
+ {
92
+ "name": "ort_float32_right_1d_batched_lhs",
93
+ "provenance": {
94
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
95
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
96
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
97
+ },
98
+ "inputs": {
99
+ "A": {
100
+ "dtype": "float32",
101
+ "shape": [3, 1, 2],
102
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0] }
103
+ },
104
+ "B": { "dtype": "float32", "shape": [2], "data": { "kind": "values", "values": [0.0, 1.0] } }
105
+ },
106
+ "outputs": {
107
+ "Y": {
108
+ "dtype": "float32",
109
+ "shape": [3, 1],
110
+ "tolerance": 0.000001,
111
+ "data": { "kind": "values", "values": [1.0, 3.0, 5.0] }
112
+ }
113
+ }
114
+ },
115
+ {
116
+ "name": "ort_float32_plain_2d",
117
+ "provenance": {
118
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
119
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
120
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
121
+ },
122
+ "inputs": {
123
+ "A": {
124
+ "dtype": "float32",
125
+ "shape": [3, 4],
126
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0] }
127
+ },
128
+ "B": {
129
+ "dtype": "float32",
130
+ "shape": [4, 3],
131
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0] }
132
+ }
133
+ },
134
+ "outputs": {
135
+ "Y": {
136
+ "dtype": "float32",
137
+ "shape": [3, 3],
138
+ "tolerance": 0.000001,
139
+ "data": { "kind": "values", "values": [42.0, 48.0, 54.0, 114.0, 136.0, 158.0, 186.0, 224.0, 262.0] }
140
+ }
141
+ }
142
+ },
143
+ {
144
+ "name": "ort_float32_rank3_by_rank2",
145
+ "provenance": {
146
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
147
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
148
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
149
+ },
150
+ "inputs": {
151
+ "A": {
152
+ "dtype": "float32",
153
+ "shape": [2, 2, 3],
154
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0] }
155
+ },
156
+ "B": {
157
+ "dtype": "float32",
158
+ "shape": [3, 4],
159
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0] }
160
+ }
161
+ },
162
+ "outputs": {
163
+ "Y": {
164
+ "dtype": "float32",
165
+ "shape": [2, 2, 4],
166
+ "tolerance": 0.000001,
167
+ "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_float32_rank3_by_rank2_output_Y" } }
168
+ }
169
+ }
170
+ },
171
+ {
172
+ "name": "ort_float32_rank3_by_broadcast_rank3",
173
+ "provenance": {
174
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
175
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
176
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
177
+ },
178
+ "inputs": {
179
+ "A": {
180
+ "dtype": "float32",
181
+ "shape": [2, 2, 3],
182
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0] }
183
+ },
184
+ "B": {
185
+ "dtype": "float32",
186
+ "shape": [1, 3, 4],
187
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0] }
188
+ }
189
+ },
190
+ "outputs": {
191
+ "Y": {
192
+ "dtype": "float32",
193
+ "shape": [2, 2, 4],
194
+ "tolerance": 0.000001,
195
+ "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_float32_rank3_by_rank2_output_Y" } }
196
+ }
197
+ }
198
+ },
199
+ {
200
+ "name": "ort_float32_singleton_rank3_by_rank3",
201
+ "provenance": {
202
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
203
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
204
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
205
+ },
206
+ "inputs": {
207
+ "A": {
208
+ "dtype": "float32",
209
+ "shape": [1, 2, 3],
210
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0] }
211
+ },
212
+ "B": {
213
+ "dtype": "float32",
214
+ "shape": [1, 3, 4],
215
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0] }
216
+ }
217
+ },
218
+ "outputs": {
219
+ "Y": {
220
+ "dtype": "float32",
221
+ "shape": [1, 2, 4],
222
+ "tolerance": 0.000001,
223
+ "data": { "kind": "values", "values": [20.0, 23.0, 26.0, 29.0, 56.0, 68.0, 80.0, 92.0] }
224
+ }
225
+ }
226
+ },
227
+ {
228
+ "name": "ort_float32_batched_rank4",
229
+ "provenance": {
230
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
231
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
232
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
233
+ },
234
+ "inputs": {
235
+ "A": {
236
+ "dtype": "float32",
237
+ "shape": [2, 2, 2, 2],
238
+ "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_float32_batched_rank4_input_A" } }
239
+ },
240
+ "B": {
241
+ "dtype": "float32",
242
+ "shape": [2, 2, 2, 2],
243
+ "data": { "kind": "values", "values": { "$ref": "#/fixtureArrays/ort_float32_batched_rank4_input_A" } }
244
+ }
245
+ },
246
+ "outputs": {
247
+ "Y": {
248
+ "dtype": "float32",
249
+ "shape": [2, 2, 2, 2],
250
+ "tolerance": 0.000001,
251
+ "data": {
252
+ "kind": "values",
253
+ "values": [2.0, 3.0, 6.0, 11.0, 46.0, 55.0, 66.0, 79.0, 154.0, 171.0, 190.0, 211.0, 326.0, 351.0, 378.0, 407.0]
254
+ }
255
+ }
256
+ }
257
+ },
258
+ {
259
+ "name": "ort_float32_broadcast_rank4_by_rank4",
260
+ "provenance": {
261
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
262
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
263
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
264
+ },
265
+ "inputs": {
266
+ "A": {
267
+ "dtype": "float32",
268
+ "shape": [1, 2, 3, 2],
269
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0] }
270
+ },
271
+ "B": {
272
+ "dtype": "float32",
273
+ "shape": [3, 2, 2, 1],
274
+ "data": { "kind": "values", "values": [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0] }
275
+ }
276
+ },
277
+ "outputs": {
278
+ "Y": {
279
+ "dtype": "float32",
280
+ "shape": [3, 2, 3, 1],
281
+ "tolerance": 0.000001,
282
+ "data": {
283
+ "kind": "values",
284
+ "values": { "$ref": "#/fixtureArrays/ort_float32_broadcast_rank3_by_rank4_output_Y" }
285
+ }
286
+ }
287
+ }
288
+ },
289
+ {
290
+ "name": "ort_float32_vector_dot_scalar_output",
291
+ "provenance": {
292
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
293
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
294
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
295
+ },
296
+ "inputs": {
297
+ "A": { "dtype": "float32", "shape": [3], "data": { "kind": "values", "values": [0.0, 1.0, 2.0] } },
298
+ "B": { "dtype": "float32", "shape": [3], "data": { "kind": "values", "values": [0.0, 1.0, 2.0] } }
299
+ },
300
+ "outputs": {
301
+ "Y": { "dtype": "float32", "shape": [], "tolerance": 0.000001, "data": { "kind": "values", "values": [5.0] } }
302
+ }
303
+ },
304
+ {
305
+ "name": "ort_float32_alpha_zero_outputs_zero",
306
+ "provenance": {
307
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
308
+ "test": "FusedMatMulOpTest.DoubleTypeAlphaZero",
309
+ "notes": "Diverges from the upstream test's inputs (inputs.A values [1.0, 2.0, 3.0, 4.0] -> constant 2.0; inputs.B values [5.0, 6.0, 7.0, 8.0] -> constant 3.0); the expected output is recomputed by the CPU reference for the new inputs. A zero alpha scales the whole product away, so no operand value can reach the result and both operands are uniform fills."
310
+ },
311
+ "attrs": { "alpha": 0 },
312
+ "inputs": {
313
+ "A": { "dtype": "float32", "shape": [2, 2], "data": { "kind": "constant", "value": 2.0 } },
314
+ "B": { "dtype": "float32", "shape": [2, 2], "data": { "kind": "constant", "value": 3.0 } }
315
+ },
316
+ "outputs": {
317
+ "Y": {
318
+ "dtype": "float32",
319
+ "shape": [2, 2],
320
+ "tolerance": 0.000001,
321
+ "data": { "kind": "values", "values": [0.0, 0.0, 0.0, 0.0] }
322
+ }
323
+ }
324
+ },
325
+ {
326
+ "name": "ort_float32_empty_k_dimension_outputs_zero",
327
+ "provenance": {
328
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
329
+ "test": "FusedMatMulOpTest.DoubleTypeEmptyKDim",
330
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
331
+ },
332
+ "inputs": {
333
+ "A": { "dtype": "float32", "shape": [2, 0], "data": { "kind": "values", "values": [] } },
334
+ "B": { "dtype": "float32", "shape": [0, 3], "data": { "kind": "values", "values": [] } }
335
+ },
336
+ "outputs": {
337
+ "Y": {
338
+ "dtype": "float32",
339
+ "shape": [2, 3],
340
+ "tolerance": 0.000001,
341
+ "data": { "kind": "values", "values": [0.0, 0.0, 0.0, 0.0, 0.0, 0.0] }
342
+ }
343
+ }
344
+ },
345
+ {
346
+ "name": "ort_float32_transpose_a_scaled",
347
+ "provenance": {
348
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
349
+ "test": "FusedMatMulOpTest.DoubleTypeScale",
350
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
351
+ },
352
+ "attrs": { "alpha": 0.5, "transA": 1 },
353
+ "inputs": {
354
+ "A": {
355
+ "dtype": "float32",
356
+ "shape": [2, 3],
357
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
358
+ },
359
+ "B": {
360
+ "dtype": "float32",
361
+ "shape": [2, 3],
362
+ "data": { "kind": "values", "values": [7.0, 8.0, 9.0, 10.0, 11.0, 12.0] }
363
+ }
364
+ },
365
+ "outputs": {
366
+ "Y": {
367
+ "dtype": "float32",
368
+ "shape": [3, 3],
369
+ "tolerance": 0.000001,
370
+ "data": { "kind": "values", "values": [23.5, 26.0, 28.5, 32.0, 35.5, 39.0, 40.5, 45.0, 49.5] }
371
+ }
372
+ }
373
+ },
374
+ {
375
+ "name": "ort_float32_transpose_b",
376
+ "provenance": {
377
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
378
+ "test": "FusedMatMulOpTest.FloatTypeTransposeB",
379
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
380
+ },
381
+ "attrs": { "transB": 1 },
382
+ "inputs": {
383
+ "A": {
384
+ "dtype": "float32",
385
+ "shape": [2, 3],
386
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
387
+ },
388
+ "B": {
389
+ "dtype": "float32",
390
+ "shape": [4, 3],
391
+ "data": { "kind": "values", "values": [7.0, 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, 17.0, 18.0] }
392
+ }
393
+ },
394
+ "outputs": {
395
+ "Y": {
396
+ "dtype": "float32",
397
+ "shape": [2, 4],
398
+ "tolerance": 0.000001,
399
+ "data": { "kind": "values", "values": [50.0, 68.0, 86.0, 104.0, 122.0, 167.0, 212.0, 257.0] }
400
+ }
401
+ }
402
+ },
403
+ {
404
+ "name": "ort_float32_transpose_ab_scaled",
405
+ "provenance": {
406
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
407
+ "test": "FusedMatMulOpTest.FloatTypeScale",
408
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
409
+ },
410
+ "attrs": { "alpha": 4, "transA": 1, "transB": 1 },
411
+ "inputs": {
412
+ "A": {
413
+ "dtype": "float32",
414
+ "shape": [3, 2],
415
+ "data": { "kind": "values", "values": [1.0, 4.0, 2.0, 5.0, 3.0, 6.0] }
416
+ },
417
+ "B": {
418
+ "dtype": "float32",
419
+ "shape": [4, 3],
420
+ "data": { "kind": "values", "values": [7.0, 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, 17.0, 18.0] }
421
+ }
422
+ },
423
+ "outputs": {
424
+ "Y": {
425
+ "dtype": "float32",
426
+ "shape": [2, 4],
427
+ "tolerance": 0.000001,
428
+ "data": { "kind": "values", "values": [200.0, 272.0, 344.0, 416.0, 488.0, 668.0, 848.0, 1028.0] }
429
+ }
430
+ }
431
+ },
432
+ {
433
+ "name": "ort_float32_scaled_no_transpose",
434
+ "provenance": {
435
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
436
+ "test": "FusedMatMulOpTest.FloatTypeScale",
437
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
438
+ },
439
+ "attrs": { "alpha": 0.5 },
440
+ "inputs": {
441
+ "A": {
442
+ "dtype": "float32",
443
+ "shape": [2, 3],
444
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0] }
445
+ },
446
+ "B": {
447
+ "dtype": "float32",
448
+ "shape": [3, 2],
449
+ "data": { "kind": "values", "values": [7.0, 8.0, 9.0, 10.0, 11.0, 12.0] }
450
+ }
451
+ },
452
+ "outputs": {
453
+ "Y": {
454
+ "dtype": "float32",
455
+ "shape": [2, 2],
456
+ "tolerance": 0.000001,
457
+ "data": { "kind": "values", "values": [29.0, 32.0, 69.5, 77.0] }
458
+ }
459
+ }
460
+ },
461
+ {
462
+ "name": "ort_float32_empty_input_m_zero",
463
+ "provenance": {
464
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
465
+ "test": "FusedMatMulOpTest.DoubleTypeEmptyInput",
466
+ "notes": "Derived from the com.microsoft.FusedMatMul corpus. Upstream registers one kernel for both operators and TransposeMatMul is the strict subset that omits transBatchA/transBatchB, so an expectation taken with those flags at zero describes this operator exactly."
467
+ },
468
+ "inputs": {
469
+ "A": { "dtype": "float32", "shape": [0, 3], "data": { "kind": "values", "values": [] } },
470
+ "B": {
471
+ "dtype": "float32",
472
+ "shape": [3, 4],
473
+ "data": { "kind": "values", "values": [1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0] }
474
+ }
475
+ },
476
+ "outputs": {
477
+ "Y": { "dtype": "float32", "shape": [0, 4], "tolerance": 0.000001, "data": { "kind": "values", "values": [] } }
478
+ }
479
+ },
480
+ {
481
+ "name": "aligned_plain_64x32x64",
482
+ "inputs": {
483
+ "A": {
484
+ "dtype": "float32",
485
+ "shape": [64, 32],
486
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2 }
487
+ },
488
+ "B": {
489
+ "dtype": "float32",
490
+ "shape": [32, 64],
491
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2 }
492
+ }
493
+ },
494
+ "outputs": { "Y": { "dtype": "float32", "shape": [64, 64], "tolerance": 0.0001 } }
495
+ },
496
+ {
497
+ "name": "register_blocked_plain_512x64x512_alpha_scaled",
498
+ "inputs": {
499
+ "A": {
500
+ "dtype": "float32",
501
+ "shape": [512, 64],
502
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2 }
503
+ },
504
+ "B": {
505
+ "dtype": "float32",
506
+ "shape": [64, 512],
507
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2 }
508
+ }
509
+ },
510
+ "outputs": { "Y": { "dtype": "float32", "shape": [512, 512], "tolerance": 0.0001 } },
511
+ "attrs": { "alpha": 0.5 },
512
+ "provenance": {
513
+ "notes": "Rank-2 M=N=512 and K=64 produce 64 aligned 64x64 workgroup tiles, exercising register-blocked vec4 staging and 4x4 per-thread accumulation. alpha=0.5 verifies scaling in the output epilogue."
514
+ }
515
+ },
516
+ {
517
+ "name": "aligned_transB_alpha_64x32",
518
+ "attrs": { "transB": 1, "alpha": 0.5 },
519
+ "inputs": {
520
+ "A": {
521
+ "dtype": "float32",
522
+ "shape": [64, 32],
523
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.023, "scale": 0.2 }
524
+ },
525
+ "B": {
526
+ "dtype": "float32",
527
+ "shape": [64, 32],
528
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.009, "scale": 0.2 }
529
+ }
530
+ },
531
+ "outputs": { "Y": { "dtype": "float32", "shape": [64, 64], "tolerance": 0.0001 } }
532
+ },
533
+ {
534
+ "name": "aligned_batched_plain_2x64x32x64",
535
+ "inputs": {
536
+ "A": {
537
+ "dtype": "float32",
538
+ "shape": [2, 64, 32],
539
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2 }
540
+ },
541
+ "B": {
542
+ "dtype": "float32",
543
+ "shape": [2, 32, 64],
544
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2 }
545
+ }
546
+ },
547
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 64, 64], "tolerance": 0.0001 } }
548
+ },
549
+ {
550
+ "name": "aligned_batched_transB_alpha_2x64x32",
551
+ "attrs": { "transB": 1, "alpha": 0.25 },
552
+ "inputs": {
553
+ "A": {
554
+ "dtype": "float32",
555
+ "shape": [2, 64, 32],
556
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.023, "scale": 0.2 }
557
+ },
558
+ "B": {
559
+ "dtype": "float32",
560
+ "shape": [2, 64, 32],
561
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.009, "scale": 0.2 }
562
+ }
563
+ },
564
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 64, 64], "tolerance": 0.0001 } }
565
+ },
566
+ {
567
+ "name": "aligned_mtail_50x32x128",
568
+ "inputs": {
569
+ "A": {
570
+ "dtype": "float32",
571
+ "shape": [50, 32],
572
+ "data": { "kind": "fillFloat32", "sinStep": 0.015, "cosStep": 0.021, "scale": 0.2 }
573
+ },
574
+ "B": {
575
+ "dtype": "float32",
576
+ "shape": [32, 128],
577
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.011, "scale": 0.2 }
578
+ }
579
+ },
580
+ "outputs": { "Y": { "dtype": "float32", "shape": [50, 128], "tolerance": 0.0001 } }
581
+ },
582
+ {
583
+ "name": "aligned_transA_64x32",
584
+ "attrs": { "transA": 1 },
585
+ "inputs": {
586
+ "A": {
587
+ "dtype": "float32",
588
+ "shape": [32, 64],
589
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2 }
590
+ },
591
+ "B": {
592
+ "dtype": "float32",
593
+ "shape": [32, 64],
594
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2 }
595
+ }
596
+ },
597
+ "outputs": { "Y": { "dtype": "float32", "shape": [64, 64], "tolerance": 0.0001 } }
598
+ },
599
+ {
600
+ "name": "aligned_transA_transB_alpha_64x32",
601
+ "attrs": { "transA": 1, "transB": 1, "alpha": 0.5 },
602
+ "inputs": {
603
+ "A": {
604
+ "dtype": "float32",
605
+ "shape": [32, 64],
606
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.023, "scale": 0.2 }
607
+ },
608
+ "B": {
609
+ "dtype": "float32",
610
+ "shape": [64, 32],
611
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.009, "scale": 0.2 }
612
+ }
613
+ },
614
+ "outputs": { "Y": { "dtype": "float32", "shape": [64, 64], "tolerance": 0.0001 } }
615
+ },
616
+ {
617
+ "name": "f32_subgroup_matrix_subnormal_dot_products_gpu_gap",
618
+ "skipGpu": {
619
+ "category": "permanent",
620
+ "reason": "Portable WGSL floating-point semantics do not guarantee preservation of the subnormal values required by this fixture. Backend evidence: WebGPU/Metal flushes subnormals to zero in f32; the ~3e-39 subnormal dot products collapse to zero (subgroup-matrix path)."
621
+ },
622
+ "provenance": {
623
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
624
+ "test": "FusedMatMulOpTest.FloatTypeNoTranspose",
625
+ "notes": "M=32, K=32, N=64 with finite subnormal dot products that must not flush to zero."
626
+ },
627
+ "inputs": {
628
+ "A": { "dtype": "float32", "shape": [32, 32], "data": { "kind": "constant", "value": 1e-20 } },
629
+ "B": { "dtype": "float32", "shape": [32, 64], "data": { "kind": "constant", "value": 1e-20 } }
630
+ },
631
+ "outputs": { "Y": { "dtype": "float32", "shape": [32, 64], "tolerance": 1e-43 } }
632
+ },
633
+ {
634
+ "name": "f32_subgroup_matrix_scaled_subnormal_dot_products_gpu_gap",
635
+ "skipGpu": {
636
+ "category": "permanent",
637
+ "reason": "Portable WGSL floating-point semantics do not guarantee preservation of the subnormal values required by this fixture. Backend evidence: WebGPU/Metal flushes subnormals to zero in f32; the ~3e-39 subnormal dot products collapse to zero before alpha scaling (subgroup-matrix path)."
638
+ },
639
+ "provenance": {
640
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
641
+ "test": "FusedMatMulOpTest.FloatTypeScale",
642
+ "notes": "Alpha scaling is applied after accumulation, so finite subnormal products remain valid nonzero outputs."
643
+ },
644
+ "attrs": { "alpha": 0.5 },
645
+ "inputs": {
646
+ "A": { "dtype": "float32", "shape": [32, 32], "data": { "kind": "constant", "value": 1e-20 } },
647
+ "B": { "dtype": "float32", "shape": [32, 64], "data": { "kind": "constant", "value": 1e-20 } }
648
+ },
649
+ "outputs": { "Y": { "dtype": "float32", "shape": [32, 64], "tolerance": 1e-43 } }
650
+ },
651
+ {
652
+ "name": "f32_subgroup_matrix_transB_subnormal_dot_products_gpu_gap",
653
+ "skipGpu": {
654
+ "category": "permanent",
655
+ "reason": "Portable WGSL floating-point semantics do not guarantee preservation of the subnormal values required by this fixture. Backend evidence: WebGPU/Metal flushes subnormals to zero in f32; the ~3e-39 subnormal dot products collapse to zero (subgroup-matrix transB path)."
656
+ },
657
+ "provenance": {
658
+ "source": "onnxruntime/test/contrib_ops/fused_matmul_op_test.cc",
659
+ "test": "FusedMatMulOpTest.FloatTypeTransposeB",
660
+ "notes": "The transposed-B subgroup-matrix path has the same finite subnormal accumulation requirement."
661
+ },
662
+ "attrs": { "transB": 1 },
663
+ "inputs": {
664
+ "A": { "dtype": "float32", "shape": [32, 32], "data": { "kind": "constant", "value": 1e-20 } },
665
+ "B": { "dtype": "float32", "shape": [64, 32], "data": { "kind": "constant", "value": 1e-20 } }
666
+ },
667
+ "outputs": { "Y": { "dtype": "float32", "shape": [32, 64], "tolerance": 1e-43 } }
668
+ },
669
+ {
670
+ "name": "aligned_f16_plain_64x32x64",
671
+ "inputs": {
672
+ "A": {
673
+ "dtype": "float16",
674
+ "shape": [64, 32],
675
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2 }
676
+ },
677
+ "B": {
678
+ "dtype": "float16",
679
+ "shape": [32, 64],
680
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2 }
681
+ }
682
+ },
683
+ "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.002 } }
684
+ },
685
+ {
686
+ "name": "aligned_f16_transB_alpha_64x32",
687
+ "attrs": { "transB": 1, "alpha": 0.5 },
688
+ "inputs": {
689
+ "A": {
690
+ "dtype": "float16",
691
+ "shape": [64, 32],
692
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.023, "scale": 0.2 }
693
+ },
694
+ "B": {
695
+ "dtype": "float16",
696
+ "shape": [64, 32],
697
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.009, "scale": 0.2 }
698
+ }
699
+ },
700
+ "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.02 } }
701
+ },
702
+ {
703
+ "name": "f16_unaligned_3x5x7",
704
+ "inputs": {
705
+ "A": {
706
+ "dtype": "float16",
707
+ "shape": [3, 5],
708
+ "data": { "kind": "fillFloat32", "sinStep": 0.015, "cosStep": 0.021, "scale": 0.2 }
709
+ },
710
+ "B": {
711
+ "dtype": "float16",
712
+ "shape": [5, 7],
713
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.011, "scale": 0.2 }
714
+ }
715
+ },
716
+ "outputs": { "Y": { "dtype": "float16", "shape": [3, 7], "tolerance": 0.003 } }
717
+ },
718
+ {
719
+ "name": "aligned_f16_transA_64x32",
720
+ "attrs": { "transA": 1 },
721
+ "inputs": {
722
+ "A": {
723
+ "dtype": "float16",
724
+ "shape": [32, 64],
725
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2 }
726
+ },
727
+ "B": {
728
+ "dtype": "float16",
729
+ "shape": [32, 64],
730
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2 }
731
+ }
732
+ },
733
+ "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.001 } }
734
+ },
735
+ {
736
+ "name": "aligned_f16_batched_plain_2x64x32x64",
737
+ "inputs": {
738
+ "A": {
739
+ "dtype": "float16",
740
+ "shape": [2, 64, 32],
741
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2 }
742
+ },
743
+ "B": {
744
+ "dtype": "float16",
745
+ "shape": [2, 32, 64],
746
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2 }
747
+ }
748
+ },
749
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 64, 64], "tolerance": 0.002 } }
750
+ },
751
+ {
752
+ "name": "f16_rank3_by_broadcast_rank3",
753
+ "inputs": {
754
+ "A": {
755
+ "dtype": "float16",
756
+ "shape": [2, 2, 3],
757
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2 }
758
+ },
759
+ "B": {
760
+ "dtype": "float16",
761
+ "shape": [1, 3, 4],
762
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2 }
763
+ }
764
+ },
765
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 2, 4], "tolerance": 0.002 } }
766
+ },
767
+ {
768
+ "name": "aligned_f16_transA_transB_alpha_64x32",
769
+ "attrs": { "transA": 1, "transB": 1, "alpha": 0.5 },
770
+ "inputs": {
771
+ "A": {
772
+ "dtype": "float16",
773
+ "shape": [32, 64],
774
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.023, "scale": 0.2 }
775
+ },
776
+ "B": {
777
+ "dtype": "float16",
778
+ "shape": [64, 32],
779
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.009, "scale": 0.2 }
780
+ }
781
+ },
782
+ "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.002 } }
783
+ },
784
+ {
785
+ "name": "subgroup_matrix_m_tail_57_partial_block_f16",
786
+ "inputs": {
787
+ "A": {
788
+ "dtype": "float16",
789
+ "shape": [57, 32],
790
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.023, "scale": 0.2 }
791
+ },
792
+ "B": {
793
+ "dtype": "float16",
794
+ "shape": [32, 64],
795
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.011, "scale": 0.2 }
796
+ }
797
+ },
798
+ "outputs": { "Y": { "dtype": "float16", "shape": [57, 64], "tolerance": 0.001 } }
799
+ },
800
+ {
801
+ "name": "subgroup_matrix_m_tail_33_alpha_scaled_f32",
802
+ "attrs": { "alpha": 0.5 },
803
+ "inputs": {
804
+ "A": {
805
+ "dtype": "float32",
806
+ "shape": [33, 32],
807
+ "data": { "kind": "fillFloat32", "sinStep": 0.015, "cosStep": 0.021, "scale": 0.2 }
808
+ },
809
+ "B": {
810
+ "dtype": "float32",
811
+ "shape": [32, 64],
812
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.019, "scale": 0.2 }
813
+ }
814
+ },
815
+ "outputs": { "Y": { "dtype": "float32", "shape": [33, 64], "tolerance": 0.0002 } }
816
+ },
817
+ {
818
+ "name": "empty_n_dimension_zero_width_output",
819
+ "inputs": {
820
+ "A": {
821
+ "dtype": "float32",
822
+ "shape": [3, 4],
823
+ "data": { "kind": "values", "values": [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0] }
824
+ },
825
+ "B": { "dtype": "float32", "shape": [4, 0], "data": { "kind": "values", "values": [] } }
826
+ },
827
+ "outputs": {
828
+ "Y": { "dtype": "float32", "shape": [3, 0], "tolerance": 0.000001, "data": { "kind": "values", "values": [] } }
829
+ }
830
+ },
831
+ {
832
+ "name": "transA_transB_subgroup_matrix_m_tail_50_f16",
833
+ "attrs": { "transA": 1, "transB": 1, "alpha": 1 },
834
+ "inputs": {
835
+ "A": {
836
+ "dtype": "float16",
837
+ "shape": [32, 50],
838
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.023, "scale": 0.2 }
839
+ },
840
+ "B": {
841
+ "dtype": "float16",
842
+ "shape": [64, 32],
843
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.009, "scale": 0.2 }
844
+ }
845
+ },
846
+ "outputs": { "Y": { "dtype": "float16", "shape": [50, 64], "tolerance": 0.002 } }
847
+ },
848
+ {
849
+ "name": "f32_decode_gemv_m1_k65_n68_vec4_compact",
850
+ "provenance": {
851
+ "notes": "A compact M=1 float32 matrix product with odd K and N=68 checks the reduction and final output-column tail."
852
+ },
853
+ "attrs": { "alpha": 1 },
854
+ "inputs": {
855
+ "A": {
856
+ "dtype": "float32",
857
+ "shape": [1, 65],
858
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.2 }
859
+ },
860
+ "B": {
861
+ "dtype": "float32",
862
+ "shape": [65, 68],
863
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.2 }
864
+ }
865
+ },
866
+ "outputs": { "Y": { "dtype": "float32", "shape": [1, 68], "tolerance": 0.0002 } }
867
+ },
868
+ {
869
+ "name": "f16_decode_gemv_m1_k65_n68_vec4_compact",
870
+ "provenance": {
871
+ "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."
872
+ },
873
+ "attrs": { "alpha": 1 },
874
+ "inputs": {
875
+ "A": {
876
+ "dtype": "float16",
877
+ "shape": [1, 65],
878
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.2 }
879
+ },
880
+ "B": {
881
+ "dtype": "float16",
882
+ "shape": [65, 68],
883
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.2 }
884
+ }
885
+ },
886
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 68], "tolerance": 0.004, "relTolerance": 0.004 } }
887
+ },
888
+ {
889
+ "name": "f16_decode_gemv_m1_k64_n128_alpha_half",
890
+ "provenance": {
891
+ "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."
892
+ },
893
+ "attrs": { "alpha": 0.5 },
894
+ "inputs": {
895
+ "A": {
896
+ "dtype": "float16",
897
+ "shape": [1, 64],
898
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.2 }
899
+ },
900
+ "B": {
901
+ "dtype": "float16",
902
+ "shape": [64, 128],
903
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.2 }
904
+ }
905
+ },
906
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 128], "tolerance": 0.0002, "relTolerance": 0.004 } }
907
+ },
908
+ {
909
+ "name": "f32_rank4_by_rank2_shared_weight_compact",
910
+ "provenance": {
911
+ "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."
912
+ },
913
+ "attrs": { "alpha": 0.5 },
914
+ "inputs": {
915
+ "A": {
916
+ "dtype": "float32",
917
+ "shape": [2, 3, 5, 7],
918
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.2 }
919
+ },
920
+ "B": {
921
+ "dtype": "float32",
922
+ "shape": [7, 9],
923
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.2 }
924
+ }
925
+ },
926
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 5, 9], "tolerance": 0.0002 } }
927
+ },
928
+ {
929
+ "name": "subgroup_matrix_kn_tail_f16_compact",
930
+ "provenance": {
931
+ "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."
932
+ },
933
+ "attrs": { "alpha": 0.5 },
934
+ "inputs": {
935
+ "A": {
936
+ "dtype": "float16",
937
+ "shape": [33, 34],
938
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.1 }
939
+ },
940
+ "B": {
941
+ "dtype": "float16",
942
+ "shape": [34, 66],
943
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.1 }
944
+ }
945
+ },
946
+ "outputs": { "Y": { "dtype": "float16", "shape": [33, 66], "tolerance": 0.0004 } }
947
+ },
948
+ {
949
+ "name": "subgroup_matrix_broadcast_rank4x3_f16_compact",
950
+ "provenance": {
951
+ "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."
952
+ },
953
+ "attrs": { "alpha": 1 },
954
+ "inputs": {
955
+ "A": {
956
+ "dtype": "float16",
957
+ "shape": [1, 2, 33, 32],
958
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.1 }
959
+ },
960
+ "B": {
961
+ "dtype": "float16",
962
+ "shape": [2, 32, 64],
963
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.1 }
964
+ }
965
+ },
966
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 2, 33, 64], "tolerance": 0.0005 } }
967
+ },
968
+ {
969
+ "name": "broadcast_rank4_tiled_reg_f16_compact",
970
+ "provenance": {
971
+ "notes": "Compact rank-4 by rank-3 float16 broadcast with odd M/K/N checks output and reduction tails."
972
+ },
973
+ "attrs": { "alpha": 0.5 },
974
+ "inputs": {
975
+ "A": {
976
+ "dtype": "float16",
977
+ "shape": [1, 2, 65, 33],
978
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.1 }
979
+ },
980
+ "B": {
981
+ "dtype": "float16",
982
+ "shape": [2, 33, 67],
983
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.1 }
984
+ }
985
+ },
986
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 2, 65, 67], "tolerance": 0.0004 } }
987
+ },
988
+ {
989
+ "name": "broadcast_rank4_tiled_reg_shared_f32_compact",
990
+ "provenance": {
991
+ "notes": "Float32 shared rank-2 weight counterpart for the register-blocked rank-4 path. Odd M/K/N cover all output and reduction tails."
992
+ },
993
+ "attrs": { "alpha": 0.5 },
994
+ "inputs": {
995
+ "A": {
996
+ "dtype": "float32",
997
+ "shape": [1, 2, 65, 33],
998
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.1 }
999
+ },
1000
+ "B": {
1001
+ "dtype": "float32",
1002
+ "shape": [33, 67],
1003
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.1 }
1004
+ }
1005
+ },
1006
+ "outputs": { "Y": { "dtype": "float32", "shape": [1, 2, 65, 67], "tolerance": 0.0003 } }
1007
+ },
1008
+ {
1009
+ "name": "rank5_three_batch_dims",
1010
+ "attrs": { "alpha": 1, "transA": 0, "transB": 0 },
1011
+ "inputs": {
1012
+ "A": {
1013
+ "dtype": "float32",
1014
+ "shape": [2, 1, 2, 2, 3],
1015
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 0.5 }
1016
+ },
1017
+ "B": {
1018
+ "dtype": "float32",
1019
+ "shape": [1, 3, 1, 3, 4],
1020
+ "data": { "kind": "fillFloat32", "sinStep": 0.17, "cosStep": 0.31, "scale": 0.25 }
1021
+ }
1022
+ },
1023
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 2, 2, 4], "tolerance": 0.000001 } }
1024
+ },
1025
+ {
1026
+ "name": "subgroup_matrix_kn_tail_f16_offset_alpha_scale",
1027
+ "provenance": {
1028
+ "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."
1029
+ },
1030
+ "attrs": { "alpha": 0.5 },
1031
+ "inputs": {
1032
+ "A": {
1033
+ "dtype": "float16",
1034
+ "shape": [33, 34],
1035
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.1, "offset": 0.5 }
1036
+ },
1037
+ "B": {
1038
+ "dtype": "float16",
1039
+ "shape": [34, 66],
1040
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.1, "offset": 0.4 }
1041
+ }
1042
+ },
1043
+ "outputs": { "Y": { "dtype": "float16", "shape": [33, 66], "tolerance": 0.03, "relTolerance": 0.01 } }
1044
+ },
1045
+ {
1046
+ "name": "subgroup_matrix_broadcast_rank4x3_f16_offset_scale",
1047
+ "provenance": {
1048
+ "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."
1049
+ },
1050
+ "attrs": { "alpha": 1 },
1051
+ "inputs": {
1052
+ "A": {
1053
+ "dtype": "float16",
1054
+ "shape": [1, 2, 33, 32],
1055
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.1, "offset": 0.5 }
1056
+ },
1057
+ "B": {
1058
+ "dtype": "float16",
1059
+ "shape": [2, 32, 64],
1060
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.1, "offset": 0.4 }
1061
+ }
1062
+ },
1063
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 2, 33, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1064
+ },
1065
+ {
1066
+ "name": "subgroup_matrix_a_batch_broadcast_rank4x3_f16",
1067
+ "provenance": {
1068
+ "notes": "A has batch extent 1 while B has extent 2, so one A slice feeds both output batches. Offset operands keep the expected magnitude nonzero, exposing a swapped or nonzero A batch stride."
1069
+ },
1070
+ "attrs": { "alpha": 1 },
1071
+ "inputs": {
1072
+ "A": {
1073
+ "dtype": "float16",
1074
+ "shape": [1, 1, 33, 32],
1075
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.019, "scale": 0.1, "offset": 0.5 }
1076
+ },
1077
+ "B": {
1078
+ "dtype": "float16",
1079
+ "shape": [2, 32, 64],
1080
+ "data": { "kind": "fillFloat32", "sinStep": 0.007, "cosStep": 0.031, "scale": 0.1, "offset": 0.4 }
1081
+ }
1082
+ },
1083
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 2, 33, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1084
+ },
1085
+ {
1086
+ "name": "broadcast_rank4_tiled_reg_f16_offset_alpha_scale",
1087
+ "provenance": {
1088
+ "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."
1089
+ },
1090
+ "attrs": { "alpha": 0.5 },
1091
+ "inputs": {
1092
+ "A": {
1093
+ "dtype": "float16",
1094
+ "shape": [1, 2, 65, 33],
1095
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.1, "offset": 0.5 }
1096
+ },
1097
+ "B": {
1098
+ "dtype": "float16",
1099
+ "shape": [2, 33, 67],
1100
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.1, "offset": 0.4 }
1101
+ }
1102
+ },
1103
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 2, 65, 67], "tolerance": 0.03, "relTolerance": 0.01 } }
1104
+ },
1105
+ {
1106
+ "name": "aligned_f16_transA_transB_alpha_offset_scale",
1107
+ "provenance": {
1108
+ "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."
1109
+ },
1110
+ "attrs": { "transA": 1, "transB": 1, "alpha": 0.5 },
1111
+ "inputs": {
1112
+ "A": {
1113
+ "dtype": "float16",
1114
+ "shape": [32, 64],
1115
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.023, "scale": 0.2, "offset": 0.5 }
1116
+ },
1117
+ "B": {
1118
+ "dtype": "float16",
1119
+ "shape": [64, 32],
1120
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.009, "scale": 0.2, "offset": 0.4 }
1121
+ }
1122
+ },
1123
+ "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1124
+ },
1125
+ {
1126
+ "name": "f16_unaligned_3x5x7_offset_scale",
1127
+ "provenance": {
1128
+ "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."
1129
+ },
1130
+ "inputs": {
1131
+ "A": {
1132
+ "dtype": "float16",
1133
+ "shape": [3, 5],
1134
+ "data": { "kind": "fillFloat32", "sinStep": 0.015, "cosStep": 0.021, "scale": 0.2, "offset": 0.6 }
1135
+ },
1136
+ "B": {
1137
+ "dtype": "float16",
1138
+ "shape": [5, 7],
1139
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.011, "scale": 0.2, "offset": 0.5 }
1140
+ }
1141
+ },
1142
+ "outputs": { "Y": { "dtype": "float16", "shape": [3, 7], "tolerance": 0.01, "relTolerance": 0.005 } }
1143
+ },
1144
+ {
1145
+ "name": "aligned_f16_plain_64x32x64_offset_scale",
1146
+ "provenance": {
1147
+ "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."
1148
+ },
1149
+ "inputs": {
1150
+ "A": {
1151
+ "dtype": "float16",
1152
+ "shape": [64, 32],
1153
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2, "offset": 0.5 }
1154
+ },
1155
+ "B": {
1156
+ "dtype": "float16",
1157
+ "shape": [32, 64],
1158
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2, "offset": 0.4 }
1159
+ }
1160
+ },
1161
+ "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1162
+ },
1163
+ {
1164
+ "name": "aligned_f16_batched_plain_2x64x32x64_offset_scale",
1165
+ "provenance": {
1166
+ "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."
1167
+ },
1168
+ "inputs": {
1169
+ "A": {
1170
+ "dtype": "float16",
1171
+ "shape": [2, 64, 32],
1172
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2, "offset": 0.5 }
1173
+ },
1174
+ "B": {
1175
+ "dtype": "float16",
1176
+ "shape": [2, 32, 64],
1177
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2, "offset": 0.4 }
1178
+ }
1179
+ },
1180
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1181
+ },
1182
+ {
1183
+ "name": "subgroup_matrix_m_tail_57_partial_block_f16_offset_scale",
1184
+ "provenance": {
1185
+ "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."
1186
+ },
1187
+ "inputs": {
1188
+ "A": {
1189
+ "dtype": "float16",
1190
+ "shape": [57, 32],
1191
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.023, "scale": 0.2, "offset": 0.5 }
1192
+ },
1193
+ "B": {
1194
+ "dtype": "float16",
1195
+ "shape": [32, 64],
1196
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.011, "scale": 0.2, "offset": 0.4 }
1197
+ }
1198
+ },
1199
+ "outputs": { "Y": { "dtype": "float16", "shape": [57, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1200
+ },
1201
+ {
1202
+ "name": "aligned_f16_transA_64x32_offset_scale",
1203
+ "provenance": {
1204
+ "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."
1205
+ },
1206
+ "attrs": { "transA": 1 },
1207
+ "inputs": {
1208
+ "A": {
1209
+ "dtype": "float16",
1210
+ "shape": [32, 64],
1211
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2, "offset": 0.5 }
1212
+ },
1213
+ "B": {
1214
+ "dtype": "float16",
1215
+ "shape": [32, 64],
1216
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2, "offset": 0.4 }
1217
+ }
1218
+ },
1219
+ "outputs": { "Y": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1220
+ },
1221
+ {
1222
+ "name": "transA_transB_subgroup_matrix_m_tail_50_f16_offset_scale",
1223
+ "provenance": {
1224
+ "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."
1225
+ },
1226
+ "attrs": { "transA": 1, "transB": 1, "alpha": 1 },
1227
+ "inputs": {
1228
+ "A": {
1229
+ "dtype": "float16",
1230
+ "shape": [32, 50],
1231
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.023, "scale": 0.2, "offset": 0.5 }
1232
+ },
1233
+ "B": {
1234
+ "dtype": "float16",
1235
+ "shape": [64, 32],
1236
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.009, "scale": 0.2, "offset": 0.4 }
1237
+ }
1238
+ },
1239
+ "outputs": { "Y": { "dtype": "float16", "shape": [50, 64], "tolerance": 0.03, "relTolerance": 0.01 } }
1240
+ },
1241
+ {
1242
+ "name": "f16_rank3_by_broadcast_rank3_offset_scale",
1243
+ "provenance": {
1244
+ "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."
1245
+ },
1246
+ "inputs": {
1247
+ "A": {
1248
+ "dtype": "float16",
1249
+ "shape": [2, 2, 3],
1250
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.027, "scale": 0.2, "offset": 1.0 }
1251
+ },
1252
+ "B": {
1253
+ "dtype": "float16",
1254
+ "shape": [1, 3, 4],
1255
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.007, "scale": 0.2, "offset": 0.8 }
1256
+ }
1257
+ },
1258
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 2, 4], "tolerance": 0.01, "relTolerance": 0.005 } }
1259
+ },
1260
+ {
1261
+ "name": "rank6_four_batch_dims",
1262
+ "attrs": { "alpha": 1, "transA": 0, "transB": 0 },
1263
+ "inputs": {
1264
+ "A": {
1265
+ "dtype": "float32",
1266
+ "shape": [2, 2, 1, 2, 2, 3],
1267
+ "data": { "kind": "fillFloat32", "sinStep": 0.13, "cosStep": 0.29, "scale": 0.5 }
1268
+ },
1269
+ "B": {
1270
+ "dtype": "float32",
1271
+ "shape": [1, 1, 3, 1, 3, 4],
1272
+ "data": { "kind": "fillFloat32", "sinStep": 0.17, "cosStep": 0.31, "scale": 0.25 }
1273
+ }
1274
+ },
1275
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 2, 3, 2, 2, 4], "tolerance": 0.000001 } }
1276
+ },
1277
+ {
1278
+ "name": "subgroup_matrix_band_m8_f16",
1279
+ "attrs": { "alpha": 2 },
1280
+ "inputs": {
1281
+ "A": {
1282
+ "dtype": "float16",
1283
+ "shape": [8, 64],
1284
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021 }
1285
+ },
1286
+ "B": {
1287
+ "dtype": "float16",
1288
+ "shape": [64, 64],
1289
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.029 }
1290
+ }
1291
+ },
1292
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 64], "tolerance": 0.005 } }
1293
+ },
1294
+ {
1295
+ "name": "subgroup_matrix_splitk_m_tail_f16",
1296
+ "attrs": { "alpha": 2 },
1297
+ "inputs": {
1298
+ "A": {
1299
+ "dtype": "float16",
1300
+ "shape": [10, 2048],
1301
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021 }
1302
+ },
1303
+ "B": {
1304
+ "dtype": "float16",
1305
+ "shape": [2048, 256],
1306
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.031 }
1307
+ }
1308
+ },
1309
+ "outputs": { "Y": { "dtype": "float16", "shape": [10, 256], "tolerance": 0.005 } }
1310
+ },
1311
+ {
1312
+ "name": "subgroup_matrix_splitk_alpha_scaled_f32",
1313
+ "attrs": { "alpha": 1.5 },
1314
+ "inputs": {
1315
+ "A": {
1316
+ "dtype": "float32",
1317
+ "shape": [16, 1024],
1318
+ "data": { "kind": "fillFloat32", "sinStep": 0.03, "cosStep": 0.07, "scale": 0.2 }
1319
+ },
1320
+ "B": {
1321
+ "dtype": "float32",
1322
+ "shape": [1024, 128],
1323
+ "data": { "kind": "fillFloat32", "sinStep": 0.041, "cosStep": 0.089, "scale": 0.2 }
1324
+ }
1325
+ },
1326
+ "outputs": { "Y": { "dtype": "float32", "shape": [16, 128], "tolerance": 0.0002 } }
1327
+ },
1328
+ {
1329
+ "name": "band_vec4_alpha_scaled_m8_k256_n512",
1330
+ "provenance": {
1331
+ "notes": "Eight rows, 256 reduction elements, and 512 output columns with alpha=0.5 verify that the non-unit scale factor is applied correctly to the matrix product."
1332
+ },
1333
+ "attrs": { "alpha": 0.5 },
1334
+ "inputs": {
1335
+ "A": {
1336
+ "dtype": "float32",
1337
+ "shape": [8, 256],
1338
+ "data": { "kind": "fillFloat32", "sinStep": 0.07, "cosStep": 0.13, "scale": 0.2 }
1339
+ },
1340
+ "B": {
1341
+ "dtype": "float32",
1342
+ "shape": [256, 512],
1343
+ "data": { "kind": "fillFloat32", "sinStep": 0.05, "cosStep": 0.19, "scale": 0.2 }
1344
+ }
1345
+ },
1346
+ "outputs": { "Y": { "dtype": "float32", "shape": [8, 512], "tolerance": 0.0001 } }
1347
+ },
1348
+ {
1349
+ "name": "subgroup_matrix_batched_transB_small_m_f16",
1350
+ "attrs": { "transB": 1, "alpha": 0.25 },
1351
+ "inputs": {
1352
+ "A": {
1353
+ "dtype": "float16",
1354
+ "shape": [2, 4, 32],
1355
+ "data": { "kind": "fillFloat32", "sinStep": 0.011, "cosStep": 0.023, "scale": 0.2 }
1356
+ },
1357
+ "B": {
1358
+ "dtype": "float16",
1359
+ "shape": [2, 64, 32],
1360
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.009, "scale": 0.2 }
1361
+ }
1362
+ },
1363
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 4, 64], "tolerance": 0.01 } }
1364
+ },
1365
+ {
1366
+ "name": "subgroup_matrix_transA_small_m_f16",
1367
+ "attrs": { "transA": 1, "alpha": 3 },
1368
+ "inputs": {
1369
+ "A": {
1370
+ "dtype": "float16",
1371
+ "shape": [64, 8],
1372
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021 }
1373
+ },
1374
+ "B": {
1375
+ "dtype": "float16",
1376
+ "shape": [64, 64],
1377
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.029 }
1378
+ }
1379
+ },
1380
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 64], "tolerance": 0.005 } }
1381
+ },
1382
+ {
1383
+ "name": "f32_decode_gemv_m1_k65_n68_alpha_half",
1384
+ "provenance": {
1385
+ "notes": "A compact M=1 GEMV with alpha=0.5 exercises a non-unit multiplier folded into the store as a baked constant. The expected output verifies that the multiplier is applied exactly once."
1386
+ },
1387
+ "attrs": { "alpha": 0.5 },
1388
+ "inputs": {
1389
+ "A": {
1390
+ "dtype": "float32",
1391
+ "shape": [1, 65],
1392
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021, "scale": 0.2 }
1393
+ },
1394
+ "B": {
1395
+ "dtype": "float32",
1396
+ "shape": [65, 68],
1397
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029, "scale": 0.2 }
1398
+ }
1399
+ },
1400
+ "outputs": { "Y": { "dtype": "float32", "shape": [1, 68], "tolerance": 0.0002 } }
1401
+ },
1402
+ {
1403
+ "name": "f16_rank4_by_rank2_shared_weight_m33_k34_n66",
1404
+ "provenance": {
1405
+ "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."
1406
+ },
1407
+ "attrs": { "alpha": 0.5 },
1408
+ "inputs": {
1409
+ "A": {
1410
+ "dtype": "float16",
1411
+ "shape": [2, 3, 33, 34],
1412
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021 }
1413
+ },
1414
+ "B": {
1415
+ "dtype": "float16",
1416
+ "shape": [34, 66],
1417
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029 }
1418
+ }
1419
+ },
1420
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 33, 66], "tolerance": 0.001 } }
1421
+ },
1422
+ {
1423
+ "name": "f16_rank4_by_rank2_shared_weight_m8_k32_n64",
1424
+ "provenance": {
1425
+ "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."
1426
+ },
1427
+ "attrs": { "alpha": 0.5 },
1428
+ "inputs": {
1429
+ "A": {
1430
+ "dtype": "float16",
1431
+ "shape": [2, 3, 8, 32],
1432
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021 }
1433
+ },
1434
+ "B": {
1435
+ "dtype": "float16",
1436
+ "shape": [32, 64],
1437
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029 }
1438
+ }
1439
+ },
1440
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 8, 64], "tolerance": 0.001 } }
1441
+ },
1442
+ {
1443
+ "name": "f16_rank4_by_rank2_shared_weight_m1_k32_n64",
1444
+ "provenance": {
1445
+ "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."
1446
+ },
1447
+ "attrs": { "alpha": 0.5 },
1448
+ "inputs": {
1449
+ "A": {
1450
+ "dtype": "float16",
1451
+ "shape": [2, 3, 1, 32],
1452
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.021 }
1453
+ },
1454
+ "B": {
1455
+ "dtype": "float16",
1456
+ "shape": [32, 64],
1457
+ "data": { "kind": "fillFloat32", "sinStep": 0.009, "cosStep": 0.029 }
1458
+ }
1459
+ },
1460
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 1, 64], "tolerance": 0.001 } }
1461
+ },
1462
+ {
1463
+ "name": "f32_band_preferred_m4_k2048_n4096",
1464
+ "attrs": { "alpha": 0.5 },
1465
+ "inputs": {
1466
+ "A": {
1467
+ "dtype": "float32",
1468
+ "shape": [4, 2048],
1469
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.1, "offset": 0.02 }
1470
+ },
1471
+ "B": {
1472
+ "dtype": "float32",
1473
+ "shape": [2048, 4096],
1474
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.041, "scale": 0.1, "offset": 0.03 }
1475
+ }
1476
+ },
1477
+ "outputs": { "Y": { "dtype": "float32", "shape": [4, 4096], "tolerance": 0.0001, "relTolerance": 0.00001 } },
1478
+ "provenance": {
1479
+ "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."
1480
+ }
1481
+ },
1482
+ {
1483
+ "name": "f32_band_preferred_m16_k2560_n4096",
1484
+ "attrs": { "alpha": 0.5 },
1485
+ "inputs": {
1486
+ "A": {
1487
+ "dtype": "float32",
1488
+ "shape": [16, 2560],
1489
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.1, "offset": 0.02 }
1490
+ },
1491
+ "B": {
1492
+ "dtype": "float32",
1493
+ "shape": [2560, 4096],
1494
+ "data": { "kind": "fillFloat32", "sinStep": 0.017, "cosStep": 0.041, "scale": 0.1, "offset": 0.03 }
1495
+ }
1496
+ },
1497
+ "outputs": { "Y": { "dtype": "float32", "shape": [16, 4096], "tolerance": 0.0001, "relTolerance": 0.00001 } },
1498
+ "provenance": {
1499
+ "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."
1500
+ }
1501
+ },
1502
+ {
1503
+ "name": "broadcast-transb-tails-float16",
1504
+ "attrs": { "transB": 1, "alpha": -0.5 },
1505
+ "inputs": {
1506
+ "A": {
1507
+ "dtype": "float16",
1508
+ "shape": [2, 1, 129, 65],
1509
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1510
+ },
1511
+ "B": {
1512
+ "dtype": "float16",
1513
+ "shape": [3, 129, 65],
1514
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1515
+ }
1516
+ },
1517
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 3, 129, 129], "tolerance": 0.001, "relTolerance": 0 } },
1518
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1519
+ },
1520
+ {
1521
+ "name": "broadcast-transb-rank3x2-float16",
1522
+ "attrs": { "transB": 1, "alpha": -0.5 },
1523
+ "inputs": {
1524
+ "A": {
1525
+ "dtype": "float16",
1526
+ "shape": [4, 256, 128],
1527
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1528
+ },
1529
+ "B": {
1530
+ "dtype": "float16",
1531
+ "shape": [512, 128],
1532
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1533
+ }
1534
+ },
1535
+ "outputs": { "Y": { "dtype": "float16", "shape": [4, 256, 512], "tolerance": 0.001, "relTolerance": 0 } },
1536
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1537
+ },
1538
+ {
1539
+ "name": "broadcast-transb-rank5x3-float16",
1540
+ "attrs": { "transB": 1, "alpha": -0.5 },
1541
+ "inputs": {
1542
+ "A": {
1543
+ "dtype": "float16",
1544
+ "shape": [2, 1, 2, 128, 32],
1545
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1546
+ },
1547
+ "B": {
1548
+ "dtype": "float16",
1549
+ "shape": [2, 128, 32],
1550
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1551
+ }
1552
+ },
1553
+ "outputs": { "Y": { "dtype": "float16", "shape": [2, 1, 2, 128, 128], "tolerance": 0.001, "relTolerance": 0 } },
1554
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1555
+ },
1556
+ {
1557
+ "name": "broadcast-transb-low_tiles-float16",
1558
+ "attrs": { "transB": 1, "alpha": -0.5 },
1559
+ "inputs": {
1560
+ "A": {
1561
+ "dtype": "float16",
1562
+ "shape": [1, 64, 32],
1563
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1564
+ },
1565
+ "B": {
1566
+ "dtype": "float16",
1567
+ "shape": [64, 32],
1568
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1569
+ }
1570
+ },
1571
+ "outputs": { "Y": { "dtype": "float16", "shape": [1, 64, 64], "tolerance": 0.001, "relTolerance": 0 } },
1572
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1573
+ },
1574
+ {
1575
+ "name": "broadcast-transb-large-float32",
1576
+ "attrs": { "transB": 1, "alpha": -0.5 },
1577
+ "inputs": {
1578
+ "A": {
1579
+ "dtype": "float32",
1580
+ "shape": [2, 8, 512, 64],
1581
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1582
+ },
1583
+ "B": {
1584
+ "dtype": "float32",
1585
+ "shape": [8, 512, 64],
1586
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1587
+ }
1588
+ },
1589
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 8, 512, 512], "tolerance": 0.00001, "relTolerance": 0 } },
1590
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1591
+ },
1592
+ {
1593
+ "name": "broadcast-transb-tails-float32",
1594
+ "attrs": { "transB": 1, "alpha": -0.5 },
1595
+ "inputs": {
1596
+ "A": {
1597
+ "dtype": "float32",
1598
+ "shape": [2, 1, 129, 65],
1599
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1600
+ },
1601
+ "B": {
1602
+ "dtype": "float32",
1603
+ "shape": [3, 129, 65],
1604
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1605
+ }
1606
+ },
1607
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 3, 129, 129], "tolerance": 0.00001, "relTolerance": 0 } },
1608
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1609
+ },
1610
+ {
1611
+ "name": "broadcast-transb-rank3x2-float32",
1612
+ "attrs": { "transB": 1, "alpha": -0.5 },
1613
+ "inputs": {
1614
+ "A": {
1615
+ "dtype": "float32",
1616
+ "shape": [4, 256, 128],
1617
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1618
+ },
1619
+ "B": {
1620
+ "dtype": "float32",
1621
+ "shape": [512, 128],
1622
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1623
+ }
1624
+ },
1625
+ "outputs": { "Y": { "dtype": "float32", "shape": [4, 256, 512], "tolerance": 0.00001, "relTolerance": 0 } },
1626
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1627
+ },
1628
+ {
1629
+ "name": "broadcast-transb-rank5x3-float32",
1630
+ "attrs": { "transB": 1, "alpha": -0.5 },
1631
+ "inputs": {
1632
+ "A": {
1633
+ "dtype": "float32",
1634
+ "shape": [2, 1, 2, 128, 32],
1635
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1636
+ },
1637
+ "B": {
1638
+ "dtype": "float32",
1639
+ "shape": [2, 128, 32],
1640
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1641
+ }
1642
+ },
1643
+ "outputs": { "Y": { "dtype": "float32", "shape": [2, 1, 2, 128, 128], "tolerance": 0.00001, "relTolerance": 0 } },
1644
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1645
+ },
1646
+ {
1647
+ "name": "broadcast-transb-low_tiles-float32",
1648
+ "attrs": { "transB": 1, "alpha": -0.5 },
1649
+ "inputs": {
1650
+ "A": {
1651
+ "dtype": "float32",
1652
+ "shape": [1, 64, 32],
1653
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1654
+ },
1655
+ "B": {
1656
+ "dtype": "float32",
1657
+ "shape": [64, 32],
1658
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1659
+ }
1660
+ },
1661
+ "outputs": { "Y": { "dtype": "float32", "shape": [1, 64, 64], "tolerance": 0.00001, "relTolerance": 0 } },
1662
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1663
+ },
1664
+ {
1665
+ "name": "broadcast-grid-b3-m128-k64-n256-float16-a0.125",
1666
+ "attrs": { "transB": 1, "alpha": 0.125 },
1667
+ "inputs": {
1668
+ "A": {
1669
+ "dtype": "float16",
1670
+ "shape": [3, 128, 64],
1671
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1672
+ },
1673
+ "B": {
1674
+ "dtype": "float16",
1675
+ "shape": [256, 64],
1676
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1677
+ }
1678
+ },
1679
+ "outputs": { "Y": { "dtype": "float16", "shape": [3, 128, 256], "tolerance": 0.001, "relTolerance": 0 } },
1680
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1681
+ },
1682
+ {
1683
+ "name": "broadcast-grid-b4-m128-k64-n256-float16-a-0.375",
1684
+ "attrs": { "transB": 1, "alpha": -0.375 },
1685
+ "inputs": {
1686
+ "A": {
1687
+ "dtype": "float16",
1688
+ "shape": [4, 128, 64],
1689
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1690
+ },
1691
+ "B": {
1692
+ "dtype": "float16",
1693
+ "shape": [256, 64],
1694
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1695
+ }
1696
+ },
1697
+ "outputs": { "Y": { "dtype": "float16", "shape": [4, 128, 256], "tolerance": 0.001, "relTolerance": 0 } },
1698
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1699
+ },
1700
+ {
1701
+ "name": "broadcast-grid-b8-m128-k64-n256-float16-a0.5",
1702
+ "attrs": { "transB": 1, "alpha": 0.5 },
1703
+ "inputs": {
1704
+ "A": {
1705
+ "dtype": "float16",
1706
+ "shape": [8, 128, 64],
1707
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1708
+ },
1709
+ "B": {
1710
+ "dtype": "float16",
1711
+ "shape": [256, 64],
1712
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1713
+ }
1714
+ },
1715
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 128, 256], "tolerance": 0.001, "relTolerance": 0 } },
1716
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1717
+ },
1718
+ {
1719
+ "name": "broadcast-grid-b8-m129-k65-n257-float16-a0.125",
1720
+ "attrs": { "transB": 1, "alpha": 0.125 },
1721
+ "inputs": {
1722
+ "A": {
1723
+ "dtype": "float16",
1724
+ "shape": [8, 129, 65],
1725
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1726
+ },
1727
+ "B": {
1728
+ "dtype": "float16",
1729
+ "shape": [257, 65],
1730
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1731
+ }
1732
+ },
1733
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 129, 257], "tolerance": 0.001, "relTolerance": 0 } },
1734
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1735
+ },
1736
+ {
1737
+ "name": "broadcast-grid-b8-m128-k96-n128-float16-a0",
1738
+ "attrs": { "transB": 1, "alpha": 0 },
1739
+ "inputs": {
1740
+ "A": {
1741
+ "dtype": "float16",
1742
+ "shape": [8, 128, 96],
1743
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1744
+ },
1745
+ "B": {
1746
+ "dtype": "float16",
1747
+ "shape": [128, 96],
1748
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1749
+ }
1750
+ },
1751
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 128, 128], "tolerance": 0.001, "relTolerance": 0 } },
1752
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1753
+ },
1754
+ {
1755
+ "name": "broadcast-grid-b4-m256-k127-n512-float16-a-0.5",
1756
+ "attrs": { "transB": 1, "alpha": -0.5 },
1757
+ "inputs": {
1758
+ "A": {
1759
+ "dtype": "float16",
1760
+ "shape": [4, 256, 127],
1761
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1762
+ },
1763
+ "B": {
1764
+ "dtype": "float16",
1765
+ "shape": [512, 127],
1766
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1767
+ }
1768
+ },
1769
+ "outputs": { "Y": { "dtype": "float16", "shape": [4, 256, 512], "tolerance": 0.001, "relTolerance": 0 } },
1770
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1771
+ },
1772
+ {
1773
+ "name": "broadcast-grid-b3-m128-k64-n256-float32-a0.125",
1774
+ "attrs": { "transB": 1, "alpha": 0.125 },
1775
+ "inputs": {
1776
+ "A": {
1777
+ "dtype": "float32",
1778
+ "shape": [3, 128, 64],
1779
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1780
+ },
1781
+ "B": {
1782
+ "dtype": "float32",
1783
+ "shape": [256, 64],
1784
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1785
+ }
1786
+ },
1787
+ "outputs": { "Y": { "dtype": "float32", "shape": [3, 128, 256], "tolerance": 0.00001, "relTolerance": 0 } },
1788
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1789
+ },
1790
+ {
1791
+ "name": "broadcast-grid-b4-m128-k64-n256-float32-a-0.375",
1792
+ "attrs": { "transB": 1, "alpha": -0.375 },
1793
+ "inputs": {
1794
+ "A": {
1795
+ "dtype": "float32",
1796
+ "shape": [4, 128, 64],
1797
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1798
+ },
1799
+ "B": {
1800
+ "dtype": "float32",
1801
+ "shape": [256, 64],
1802
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1803
+ }
1804
+ },
1805
+ "outputs": { "Y": { "dtype": "float32", "shape": [4, 128, 256], "tolerance": 0.00001, "relTolerance": 0 } },
1806
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1807
+ },
1808
+ {
1809
+ "name": "broadcast-grid-b8-m128-k64-n256-float32-a0.5",
1810
+ "attrs": { "transB": 1, "alpha": 0.5 },
1811
+ "inputs": {
1812
+ "A": {
1813
+ "dtype": "float32",
1814
+ "shape": [8, 128, 64],
1815
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1816
+ },
1817
+ "B": {
1818
+ "dtype": "float32",
1819
+ "shape": [256, 64],
1820
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1821
+ }
1822
+ },
1823
+ "outputs": { "Y": { "dtype": "float32", "shape": [8, 128, 256], "tolerance": 0.00001, "relTolerance": 0 } },
1824
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1825
+ },
1826
+ {
1827
+ "name": "broadcast-grid-b8-m129-k65-n257-float32-a0.125",
1828
+ "attrs": { "transB": 1, "alpha": 0.125 },
1829
+ "inputs": {
1830
+ "A": {
1831
+ "dtype": "float32",
1832
+ "shape": [8, 129, 65],
1833
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1834
+ },
1835
+ "B": {
1836
+ "dtype": "float32",
1837
+ "shape": [257, 65],
1838
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1839
+ }
1840
+ },
1841
+ "outputs": { "Y": { "dtype": "float32", "shape": [8, 129, 257], "tolerance": 0.00001, "relTolerance": 0 } },
1842
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1843
+ },
1844
+ {
1845
+ "name": "broadcast-grid-b8-m128-k96-n128-float32-a0",
1846
+ "attrs": { "transB": 1, "alpha": 0 },
1847
+ "inputs": {
1848
+ "A": {
1849
+ "dtype": "float32",
1850
+ "shape": [8, 128, 96],
1851
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1852
+ },
1853
+ "B": {
1854
+ "dtype": "float32",
1855
+ "shape": [128, 96],
1856
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1857
+ }
1858
+ },
1859
+ "outputs": { "Y": { "dtype": "float32", "shape": [8, 128, 128], "tolerance": 0.00001, "relTolerance": 0 } },
1860
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1861
+ },
1862
+ {
1863
+ "name": "broadcast-grid-b4-m256-k127-n512-float32-a-0.5",
1864
+ "attrs": { "transB": 1, "alpha": -0.5 },
1865
+ "inputs": {
1866
+ "A": {
1867
+ "dtype": "float32",
1868
+ "shape": [4, 256, 127],
1869
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1870
+ },
1871
+ "B": {
1872
+ "dtype": "float32",
1873
+ "shape": [512, 127],
1874
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1875
+ }
1876
+ },
1877
+ "outputs": { "Y": { "dtype": "float32", "shape": [4, 256, 512], "tolerance": 0.00001, "relTolerance": 0 } },
1878
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1879
+ },
1880
+ {
1881
+ "name": "broadcast-selected-mn-tail-float16",
1882
+ "attrs": { "transB": 1, "alpha": 0.125 },
1883
+ "inputs": {
1884
+ "A": {
1885
+ "dtype": "float16",
1886
+ "shape": [8, 129, 64],
1887
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1888
+ },
1889
+ "B": {
1890
+ "dtype": "float16",
1891
+ "shape": [257, 64],
1892
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1893
+ }
1894
+ },
1895
+ "outputs": { "Y": { "dtype": "float16", "shape": [8, 129, 257], "tolerance": 0.001, "relTolerance": 0 } },
1896
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1897
+ },
1898
+ {
1899
+ "name": "broadcast-selected-mn-tail-float32",
1900
+ "attrs": { "transB": 1, "alpha": 0.125 },
1901
+ "inputs": {
1902
+ "A": {
1903
+ "dtype": "float32",
1904
+ "shape": [8, 129, 64],
1905
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1906
+ },
1907
+ "B": {
1908
+ "dtype": "float32",
1909
+ "shape": [257, 64],
1910
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1911
+ }
1912
+ },
1913
+ "outputs": { "Y": { "dtype": "float32", "shape": [8, 129, 257], "tolerance": 0.00001, "relTolerance": 0 } },
1914
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1915
+ },
1916
+ {
1917
+ "name": "broadcast-selected-broadcast-tail-float16",
1918
+ "attrs": { "transB": 1, "alpha": 0.125 },
1919
+ "inputs": {
1920
+ "A": {
1921
+ "dtype": "float16",
1922
+ "shape": [4, 1, 129, 64],
1923
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1924
+ },
1925
+ "B": {
1926
+ "dtype": "float16",
1927
+ "shape": [3, 257, 64],
1928
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1929
+ }
1930
+ },
1931
+ "outputs": { "Y": { "dtype": "float16", "shape": [4, 3, 129, 257], "tolerance": 0.001, "relTolerance": 0 } },
1932
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1933
+ },
1934
+ {
1935
+ "name": "broadcast-selected-broadcast-tail-float32",
1936
+ "attrs": { "transB": 1, "alpha": 0.125 },
1937
+ "inputs": {
1938
+ "A": {
1939
+ "dtype": "float32",
1940
+ "shape": [4, 1, 129, 64],
1941
+ "data": { "kind": "fillFloat32", "sinStep": 0.013, "cosStep": 0.031, "scale": 0.075, "offset": 0.02 }
1942
+ },
1943
+ "B": {
1944
+ "dtype": "float32",
1945
+ "shape": [3, 257, 64],
1946
+ "data": { "kind": "fillFloat32", "sinStep": 0.019, "cosStep": 0.017, "scale": 0.075, "offset": -0.02 }
1947
+ }
1948
+ },
1949
+ "outputs": { "Y": { "dtype": "float32", "shape": [4, 3, 129, 257], "tolerance": 0.00001, "relTolerance": 0 } },
1950
+ "provenance": { "notes": "Broadcast transposed-B geometry, alpha, padding and output-grid boundary." }
1951
+ },
1952
+ {
1953
+ "name": "plain_rank2_tiled_reg_row_tail_m100_k64_n2048",
1954
+ "provenance": {
1955
+ "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."
1956
+ },
1957
+ "attrs": { "alpha": 0.5 },
1958
+ "inputs": {
1959
+ "A": {
1960
+ "dtype": "float32",
1961
+ "shape": [100, 64],
1962
+ "data": { "kind": "fillFloat32", "sinStep": 0.037, "cosStep": 0.061, "scale": 0.2 }
1963
+ },
1964
+ "B": {
1965
+ "dtype": "float32",
1966
+ "shape": [64, 2048],
1967
+ "data": { "kind": "fillFloat32", "sinStep": 0.043, "cosStep": 0.079, "scale": 0.2 }
1968
+ }
1969
+ },
1970
+ "outputs": { "Y": { "dtype": "float32", "shape": [100, 2048], "tolerance": 0.0001 } }
1971
+ }
1972
+ ]
1973
+ }