Xenova HF Staff commited on
Commit
abb89f2
·
verified ·
1 Parent(s): 929af3e

sync 6fdf6301e2bb

Browse files
README.md CHANGED
@@ -55,14 +55,14 @@ Attributes and default values (overridable per request):
55
 
56
  One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
57
 
58
- - `prefill_tiled_reg_vec4_splitk_default_zero` — Cuts the K reduction of the four-wide register-blocked prefill tile into power-of-two slices across dispatch.z, retaining at least 512 K values per slice, then sums the f32 partials and adds bias in a second pass.
59
- - `prefill_tiled_reg_vec4_default_zero` — The register-blocked prefill tile with four-wide activation loads and a 128-row tile above 256 rows: the activation slice of every K tile is staged with one vector load per lane instead of four scalar loads, and the taller tile halves the dequantization work per multiply-add.
60
- - `prefill_tiled_reg_vec4_splitk_zero_bias` — Cuts the K reduction of the four-wide register-blocked prefill tile into power-of-two slices across dispatch.z, retaining at least 512 K values per slice, then sums the f32 partials and adds bias in a second pass.
61
- - `prefill_tiled_reg_vec4_zero_bias` — The register-blocked prefill tile with four-wide activation loads and a 128-row tile above 256 rows: the activation slice of every K tile is staged with one vector load per lane instead of four scalar loads, and the taller tile halves the dequantization work per multiply-add.
62
- - `prefill_tiled_reg_vec4_splitk_zero_only` — Cuts the K reduction of the four-wide register-blocked prefill tile into power-of-two slices across dispatch.z, retaining at least 512 K values per slice, then sums the f32 partials and adds bias in a second pass.
63
- - `prefill_tiled_reg_vec4_zero_only` — The register-blocked prefill tile with four-wide activation loads and a 128-row tile above 256 rows: the activation slice of every K tile is staged with one vector load per lane instead of four scalar loads, and the taller tile halves the dequantization work per multiply-add.
64
- - `prefill_tiled_reg_vec4_splitk_bias_only` — Cuts the K reduction of the four-wide register-blocked prefill tile into power-of-two slices across dispatch.z, retaining at least 512 K values per slice, then sums the f32 partials and adds bias in a second pass.
65
- - `prefill_tiled_reg_vec4_bias_only` — The register-blocked prefill tile with four-wide activation loads and a 128-row tile above 256 rows: the activation slice of every K tile is staged with one vector load per lane instead of four scalar loads, and the taller tile halves the dequantization work per multiply-add.
66
 
67
  ## Device requirements
68
 
@@ -73,7 +73,7 @@ Some implementation variants require `subgroup-matrix` and `subgroups`. These ar
73
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
74
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
75
  - [`test.json`](build/webgpu/test.json) — correctness cases
76
- - [`bench.json`](build/webgpu/bench.json) — benchmark + tuning cases
77
  - [`matmul-nbits-dp4a-quantize.wgsl.jinja`](build/webgpu/matmul-nbits-dp4a-quantize.wgsl.jinja)
78
  - [`matmul-nbits-gemv-q4.wgsl.jinja`](build/webgpu/matmul-nbits-gemv-q4.wgsl.jinja)
79
  - [`matmul-nbits-q4-dp4a-prefill.wgsl.jinja`](build/webgpu/matmul-nbits-q4-dp4a-prefill.wgsl.jinja)
@@ -87,7 +87,7 @@ Some implementation variants require `subgroup-matrix` and `subgroups`. These ar
87
  ## Use with `@huggingface/kernels`
88
 
89
  ```sh
90
- npm install --save-exact @huggingface/kernels@0.0.1-preview.2
91
  ```
92
 
93
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
 
55
 
56
  One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
57
 
58
+ - `prefill_tiled_reg_vec4_splitk_default_zero` — Splits the vec4 register-tiled K reduction across dispatch.z, then combines f32 partials and bias. Large 2-bit f16 tiles pin the narrowest subgroup width of a variable range under subgroup-size control; fixed 32-lane paths reuse each packed word, scale, and zero point across four vectors.
59
+ - `prefill_tiled_reg_vec4_default_zero` — Register-blocked prefill with vec4 activation loads. Taller inputs use 128-row tiles. Large 2-bit f16 grids pin the narrowest width of a variable subgroup range under size control, with BK=16 on 16–32 lanes; other tiers use BK=32. Deep 2-bit reductions on fixed 32-lane subgroups unpack four vectors per packed word and reuse its scale and zero point.
60
+ - `prefill_tiled_reg_vec4_splitk_zero_bias` — Splits the vec4 register-tiled K reduction across dispatch.z, then combines f32 partials and bias. Large 2-bit f16 tiles pin the narrowest subgroup width of a variable range under subgroup-size control; fixed 32-lane paths reuse each packed word, scale, and zero point across four vectors.
61
+ - `prefill_tiled_reg_vec4_zero_bias` — Register-blocked prefill with vec4 activation loads. Taller inputs use 128-row tiles. Large 2-bit f16 grids pin the narrowest width of a variable subgroup range under size control, with BK=16 on 16–32 lanes; other tiers use BK=32. Deep 2-bit reductions on fixed 32-lane subgroups unpack four vectors per packed word and reuse its scale and zero point.
62
+ - `prefill_tiled_reg_vec4_splitk_zero_only` — Splits the vec4 register-tiled K reduction across dispatch.z, then combines f32 partials and bias. Large 2-bit f16 tiles pin the narrowest subgroup width of a variable range under subgroup-size control; fixed 32-lane paths reuse each packed word, scale, and zero point across four vectors.
63
+ - `prefill_tiled_reg_vec4_zero_only` — Register-blocked prefill with vec4 activation loads. Taller inputs use 128-row tiles. Large 2-bit f16 grids pin the narrowest width of a variable subgroup range under size control, with BK=16 on 16–32 lanes; other tiers use BK=32. Deep 2-bit reductions on fixed 32-lane subgroups unpack four vectors per packed word and reuse its scale and zero point.
64
+ - `prefill_tiled_reg_vec4_splitk_bias_only` — Splits the vec4 register-tiled K reduction across dispatch.z, then combines f32 partials and bias. Large 2-bit f16 tiles pin the narrowest subgroup width of a variable range under subgroup-size control; fixed 32-lane paths reuse each packed word, scale, and zero point across four vectors.
65
+ - `prefill_tiled_reg_vec4_bias_only` — Register-blocked prefill with vec4 activation loads. Taller inputs use 128-row tiles. Large 2-bit f16 grids pin the narrowest width of a variable subgroup range under size control, with BK=16 on 16–32 lanes; other tiers use BK=32. Deep 2-bit reductions on fixed 32-lane subgroups unpack four vectors per packed word and reuse its scale and zero point.
66
 
67
  ## Device requirements
68
 
 
73
  - [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
74
  - [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
75
  - [`test.json`](build/webgpu/test.json) — correctness cases
76
+ - [`bench.json`](build/webgpu/bench.json) — benchmark cases
77
  - [`matmul-nbits-dp4a-quantize.wgsl.jinja`](build/webgpu/matmul-nbits-dp4a-quantize.wgsl.jinja)
78
  - [`matmul-nbits-gemv-q4.wgsl.jinja`](build/webgpu/matmul-nbits-gemv-q4.wgsl.jinja)
79
  - [`matmul-nbits-q4-dp4a-prefill.wgsl.jinja`](build/webgpu/matmul-nbits-q4-dp4a-prefill.wgsl.jinja)
 
87
  ## Use with `@huggingface/kernels`
88
 
89
  ```sh
90
+ npm install --save-exact @huggingface/kernels@0.0.1-preview.3
91
  ```
92
 
93
  Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
build/webgpu/bench.json CHANGED
@@ -859,6 +859,72 @@
859
  "outputs": { "yT": { "shape": [48, 4096], "dtype": "float32" } },
860
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * args.M * args.K * args.N" }] },
861
  "attrs": { "K": 4096, "N": 4096, "bits": 4, "block_size": 32 }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
862
  }
863
  ]
864
  }
 
859
  "outputs": { "yT": { "shape": [48, 4096], "dtype": "float32" } },
860
  "bench": { "metrics": [{ "type": "gflops", "value": "2 * args.M * args.K * args.N" }] },
861
  "attrs": { "K": 4096, "N": 4096, "bits": 4, "block_size": 32 }
862
+ },
863
+ {
864
+ "name": "bonsai2-ffn-gate-prefill-m128-q2g128-k5120-n17408-zero1-f16",
865
+ "preset": "model",
866
+ "vars": { "M": 128, "K": 5120, "N": 17408, "bits": 2, "blockSize": 128 },
867
+ "inputs": {
868
+ "aT": { "shape": [128, 5120], "dtype": "float16", "dist": "normal", "seed": 721, "scale": 0.2 },
869
+ "bT": { "shape": [17408, 40, 32], "dtype": "uint8", "dist": "linearMod", "mod": 256, "step": 7 },
870
+ "scalesT": {
871
+ "shape": [17408, 40],
872
+ "dtype": "float16",
873
+ "dist": "uniform",
874
+ "seed": 722,
875
+ "offset": 0.04,
876
+ "scale": 0.01,
877
+ "signed": false
878
+ },
879
+ "zeroPointsT": { "shape": [17408, 40], "dtype": "float16", "dist": "constant", "value": 1 }
880
+ },
881
+ "outputs": { "yT": { "shape": [128, 17408], "dtype": "float16" } },
882
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * args.M * args.K * args.N" }] },
883
+ "attrs": { "K": 5120, "N": 17408, "bits": 2, "block_size": 128 }
884
+ },
885
+ {
886
+ "name": "bonsai2-ffn-down-prefill-m128-q2g128-k17408-n5120-zero1-f16",
887
+ "preset": "model",
888
+ "vars": { "M": 128, "K": 17408, "N": 5120, "bits": 2, "blockSize": 128 },
889
+ "inputs": {
890
+ "aT": { "shape": [128, 17408], "dtype": "float16", "dist": "normal", "seed": 721, "scale": 0.2 },
891
+ "bT": { "shape": [5120, 136, 32], "dtype": "uint8", "dist": "linearMod", "mod": 256, "step": 7 },
892
+ "scalesT": {
893
+ "shape": [5120, 136],
894
+ "dtype": "float16",
895
+ "dist": "uniform",
896
+ "seed": 722,
897
+ "offset": 0.04,
898
+ "scale": 0.01,
899
+ "signed": false
900
+ },
901
+ "zeroPointsT": { "shape": [5120, 136], "dtype": "float16", "dist": "constant", "value": 1 }
902
+ },
903
+ "outputs": { "yT": { "shape": [128, 5120], "dtype": "float16" } },
904
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * args.M * args.K * args.N" }] },
905
+ "attrs": { "K": 17408, "N": 5120, "bits": 2, "block_size": 128 }
906
+ },
907
+ {
908
+ "name": "bonsai2-attn-qkv-prefill-m128-q2g128-k5120-n10240-zero1-f16",
909
+ "preset": "model",
910
+ "vars": { "M": 128, "K": 5120, "N": 10240, "bits": 2, "blockSize": 128 },
911
+ "inputs": {
912
+ "aT": { "shape": [128, 5120], "dtype": "float16", "dist": "normal", "seed": 721, "scale": 0.2 },
913
+ "bT": { "shape": [10240, 40, 32], "dtype": "uint8", "dist": "linearMod", "mod": 256, "step": 7 },
914
+ "scalesT": {
915
+ "shape": [10240, 40],
916
+ "dtype": "float16",
917
+ "dist": "uniform",
918
+ "seed": 722,
919
+ "offset": 0.04,
920
+ "scale": 0.01,
921
+ "signed": false
922
+ },
923
+ "zeroPointsT": { "shape": [10240, 40], "dtype": "float16", "dist": "constant", "value": 1 }
924
+ },
925
+ "outputs": { "yT": { "shape": [128, 10240], "dtype": "float16" } },
926
+ "bench": { "metrics": [{ "type": "gflops", "value": "2 * args.M * args.K * args.N" }] },
927
+ "attrs": { "K": 5120, "N": 10240, "bits": 2, "block_size": 128 }
928
  }
929
  ]
930
  }
build/webgpu/manifest.json CHANGED
@@ -37,10 +37,12 @@
37
  "derive": {
38
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
39
  "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32",
 
40
  "narrowSubgroupRange": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize < device.adapterInfo.subgroupMaxSize and device.adapterInfo.subgroupMaxSize <= 16",
41
  "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",
42
  "pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter",
43
  "wave32Effective": "wave32Adapter or pinSubgroupSize32",
 
44
  "packedFeature": "device.wgslLanguageFeatures.has(\"packed_4x8_integer_dot_product\")",
45
  "kBlocksExpected": "ceilDiv(attrs.K, attrs.block_size)",
46
  "blobSizeExpected": "ceilDiv(attrs.block_size * attrs.bits, 8)",
@@ -73,7 +75,9 @@
73
  "zeroOnlyEpilogue": "zeroPointsValid and not present.biasT",
74
  "biasOnlyEpilogue": "not present.zeroPointsT and biasValid",
75
  "portableWorkgroupFits": "portableWorkgroupSize > 0 and portableWorkgroupSize * 64 <= device.limits.maxComputeWorkgroupStorageSize",
76
- "tiledWorkgroupFits": "16 <= device.limits.maxComputeWorkgroupSizeX and 16 <= device.limits.maxComputeWorkgroupSizeY and 256 <= device.limits.maxComputeInvocationsPerWorkgroup and 4096 <= device.limits.maxComputeWorkgroupStorageSize",
 
 
77
  "tiledRegWorkgroupFits": "tiledWorkgroupFits and 16384 <= device.limits.maxComputeWorkgroupStorageSize",
78
  "mediumTiledRegWorkgroupFits": "tiledWorkgroupFits and 6144 <= device.limits.maxComputeWorkgroupStorageSize",
79
  "sgmatWorkgroupFits": "sgmatWorkgroupSize <= deviceWorkgroupCap and sgmatWorkgroupStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
@@ -86,8 +90,8 @@
86
  "mediumTiledRegEligible": "mediumRegisterEligible and mediumTiledRegWorkgroupFits and dispatchM32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
87
  "tiledRegVariantEligible": "largeTiledRegEligible or mediumTiledRegEligible",
88
  "tiledRegSelectedBK": "tiledRegBK if largeTiledRegEligible else 16",
89
- "tiledRegSelectedTileRows": "64 if largeTiledRegEligible else 32",
90
  "tiledRegSelectedThreadRows": "4 if largeTiledRegEligible else 2",
 
91
  "tiledRegSelectedDispatchM": "dispatchM64 if largeTiledRegEligible else dispatchM32",
92
  "blobWords": "blobSizeExpected / 4",
93
  "codesPerWord": "32 / attrs.bits",
@@ -100,64 +104,61 @@
100
  "smallMKLanes": "min(portableWorkgroupSize, pow2ceil(smallMWordsPerCol + 1) / 2)",
101
  "smallMColGroups": "portableWorkgroupSize / smallMKLanes",
102
  "smallMDispatchN": "ceilDiv(attrs.N, 4 * smallMColGroups)",
103
- "tiledRegVec4TileRows": "128 if aRows >= 256 else 64",
104
  "tiledRegVec4ThreadRows": "8 if aRows >= 256 else 4",
 
105
  "tiledRegVec4DispatchM": "ceilDiv(aRows, tiledRegVec4TileRows)",
106
  "tiledRegVec4Eligible": "largeTiledRegEligible and attrs.K % 4 == 0 and attrs.block_size % 32 == 0 and tiledRegVec4DispatchM <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
 
 
107
  "tiledRegSplitTiles": "tiledRegVec4DispatchM * dispatchN64",
 
108
  "tiledRegSplitWant": "ceilDiv(tunables.REGISTER_TILE_SPLITK_TARGET_WORKGROUPS, max(1, tiledRegSplitTiles))",
109
  "tiledRegSplitK": "8 if (tiledRegSplitWant >= 8 and attrs.K >= 4096) else (4 if (tiledRegSplitWant >= 4 and attrs.K >= 2048) else (2 if (tiledRegSplitWant >= 2 and attrs.K >= 1024) else 1))",
110
  "tiledRegSplitTilesPerSplit": "ceilDiv(ceilDiv(attrs.K, 32), tiledRegSplitK)",
111
- "tiledRegSplitEligible": "tiledRegVec4Eligible and tiledRegSplitK >= 2 and tiledRegSplitTiles <= tunables.REGISTER_TILE_SPLITK_MAX_TILES and tiledRegSplitK * aRows * attrs.N * 4 <= device.limits.maxStorageBufferBindingSize and tiledRegSplitK <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
112
  "B_LEN": "attrs.N * kBlocksExpected * blobSizeExpected / 4",
113
  "SCALES_LEN": "attrs.N * kBlocksExpected",
114
  "BIAS_LEN": "attrs.N"
115
  },
116
  "bindings": {
117
- "a_2": { "arg": "aT", "name": "a", "buffer": "read-only-storage", "elementType": "$aElement" },
118
- "b_2": {
119
- "arg": "bT",
120
- "name": "b",
121
- "buffer": "read-only-storage",
122
- "elementType": "$bElement",
123
- "length": "$B_VEC_LEN"
124
- },
125
- "scales_2": {
126
- "arg": "scalesT",
127
- "name": "scales",
128
- "buffer": "read-only-storage",
129
- "elementType": "$scaleScalar",
130
- "length": "$SCALES_LEN"
131
- },
132
- "y_2": { "arg": "yT", "name": "y", "buffer": "storage", "elementType": "$outputScalar" },
133
  "params": {
134
- "buffer": "uniform",
135
  "struct": [
136
  { "name": "K", "type": "u32", "value": "attrs.K" },
137
  { "name": "N", "type": "u32", "value": "attrs.N" },
138
  { "name": "kBlocks", "type": "u32", "value": "dim(shapes.bT, 1)" }
139
  ]
140
  },
141
- "zero_points": {
142
- "arg": "zeroPointsT",
143
- "buffer": "read-only-storage",
144
- "elementType": "$aScalar",
145
- "length": "$SCALES_LEN"
146
- },
147
- "bias": { "arg": "biasT", "buffer": "read-only-storage", "elementType": "$aScalar", "length": "$BIAS_LEN" },
148
- "a_3": { "arg": "aT", "name": "a", "buffer": "read-only-storage", "elementType": "$aScalar" },
149
- "b_3": { "arg": "bT", "name": "b", "buffer": "read-only-storage", "elementType": "$bScalar", "length": "$B_LEN" },
150
- "a_4": { "arg": "aT", "name": "a", "buffer": "read-only-storage", "elementType": "$aVec4Element" },
151
- "y_3": { "scratch": "partials", "name": "y", "buffer": "storage", "elementType": "f32" },
152
  "partials": { "buffer": "read-only-storage", "elementType": "f32" },
153
- "params_2": {
 
154
  "name": "params",
155
- "buffer": "uniform",
156
- "struct": [{ "name": "cols", "type": "u32", "value": "aRows * attrs.N" }]
157
- },
158
- "params_3": {
159
- "name": "params",
160
- "buffer": "uniform",
161
  "struct": [
162
  { "name": "rows", "type": "u32", "value": "aRows" },
163
  { "name": "K", "type": "u32", "value": "attrs.K" },
@@ -172,18 +173,9 @@
172
  {
173
  "id": "q4_dp4a_prefill",
174
  "priority": 19,
175
- "when": ["packedFeature", "commonShapeValid", "defaultEpilogue", "attrs.bits == 4", "attrs.accuracy_level == 4", "tensorDtypes.aT == \"float32\"", "attrs.block_size % 32 == 0", "attrs.K % 128 == 0", "attrs.N % 16 == 0", "aRows >= 32", "ceilDiv(attrs.N, 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(aRows, 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "16 <= device.limits.maxComputeWorkgroupSizeX", "16 <= device.limits.maxComputeWorkgroupSizeY", "64 <= device.limits.maxComputeWorkgroupSizeX", "256 <= device.limits.maxComputeInvocationsPerWorkgroup", "4608 <= device.limits.maxComputeWorkgroupStorageSize"],
176
  "demoteWhen": ["has(device.adapterInfo, \"architecture\") and device.adapterInfo.architecture == \"maxwell\"", "device.features.has(\"chromium-experimental-subgroup-matrix\")", "device.adapterInfo.vendor == \"apple\""],
177
- "derive": {
178
- "M": "aRows",
179
- "K": "attrs.K",
180
- "N": "attrs.N",
181
- "kBlocks": "dim(shapes.bT, 1)",
182
- "blockSize": "attrs.block_size",
183
- "blobSize": "dim(shapes.bT, 2)",
184
- "vec4Count": "aRows * attrs.K / 4",
185
- "blockCount": "aRows * attrs.K / 128"
186
- },
187
  "intermediates": [
188
  { "id": "aQuant", "dtype": "uint32", "shape": "[aRows * attrs.K / 4]" },
189
  { "id": "aScales", "dtype": "float32", "shape": "[aRows * attrs.K / 128]" }
@@ -199,8 +191,8 @@
199
  { "scratch": "aScales", "name": "a_scales", "elementType": "f32" }
200
  ],
201
  "dispatch": {
202
- "x": "min(ceilDiv((aRows * attrs.K / 4), (64)), 65535)",
203
- "y": "ceilDiv(ceilDiv((aRows * attrs.K / 4), (64)), 65535)",
204
  "z": 1
205
  }
206
  },
@@ -226,33 +218,22 @@
226
  "derive": {
227
  "gemvNCols": "tunables.GEMV_N_COLS",
228
  "useSubgroups": "device.features.has(\"subgroups\")",
229
- "hasZero": false,
230
- "hasBias": false,
231
- "workgroupSize": "portableWorkgroupSize",
232
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
233
  "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")",
234
  "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"",
235
  "B_VEC_LEN": "attrs.N * gemvVecPerCol",
236
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
237
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
238
- "bits": "attrs.bits",
239
- "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
240
  "vecWords": "gemvVecWords",
241
  "codesPerVec": "gemvCodesPerVec",
242
  "codesPerVec4": "gemvCodesPerVec / 4",
243
  "vecPerBlock": "gemvVecPerBlock",
244
  "vecPerCol": "gemvVecPerCol",
245
- "actVec4": "gemvActVec4",
246
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
247
- "usesF16": "tensorDtypes.aT == \"float16\""
248
  },
249
  "passes": [
250
  {
251
  "id": "main",
252
  "shader": "matmul-nbits-gemv-q4.wgsl.jinja",
253
- "bindings": ["a_2", "b_2", "scales_2", "y_2", "params"],
254
- "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 },
255
- "subgroupCollectivesWidth": "portable"
256
  }
257
  ]
258
  },
@@ -266,21 +247,6 @@
266
  "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
267
  },
268
  "derive": {
269
- "hasZero": false,
270
- "hasBias": false,
271
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
272
- "bScalar": "\"u32\"",
273
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
274
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
275
- "M": "aRows",
276
- "K": "attrs.K",
277
- "N": "attrs.N",
278
- "kBlocks": "dim(shapes.bT, 1)",
279
- "blockSize": "attrs.block_size",
280
- "blobSize": "dim(shapes.bT, 2)",
281
- "bits": "attrs.bits",
282
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
283
- "usesF16": "tensorDtypes.aT == \"float16\"",
284
  "tileRows": "sgmatTileRows",
285
  "workgroupSize": "sgmatWorkgroupSize",
286
  "rowSubtiles": "sgmatRowSubtiles",
@@ -292,7 +258,7 @@
292
  {
293
  "id": "main",
294
  "shader": "matmul-nbits-q4-sgmat.wgsl.jinja",
295
- "bindings": ["a_3", "b_3", "scales_2", "y_2"],
296
  "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" }
297
  }
298
  ]
@@ -301,28 +267,18 @@
301
  "id": "prefill_tiled_reg_vec4_splitk_default_zero",
302
  "priority": 17,
303
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"],
 
304
  "derive": {
305
- "hasZero": false,
306
- "hasBias": false,
307
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
308
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
309
- "bScalar": "\"u32\"",
310
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
311
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
312
- "M": "aRows",
313
- "K": "attrs.K",
314
- "N": "attrs.N",
315
- "kBlocks": "dim(shapes.bT, 1)",
316
- "blockSize": "attrs.block_size",
317
- "blobSize": "dim(shapes.bT, 2)",
318
- "bits": "attrs.bits",
319
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
320
- "usesF16": "tensorDtypes.aT == \"float16\"",
321
  "bk": 32,
322
- "tileRows": "tiledRegVec4TileRows",
323
- "tileCols": 64,
324
- "threadRows": "tiledRegVec4ThreadRows",
325
- "threadCols": 4,
 
 
 
 
326
  "alignedBlockLoads": true,
327
  "aVec4Loads": true,
328
  "splitK": "tiledRegSplitK",
@@ -336,22 +292,17 @@
336
  {
337
  "id": "partial",
338
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
339
- "bindings": ["a_4", "b_3", "scales_2", "y_3"],
340
- "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM", "z": "tiledRegSplitK" }
341
  },
342
  {
343
  "id": "combine",
344
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
345
- "derive": {
346
- "op": "\"sum\"",
347
- "outputF16": "tensorDtypes.aT == \"float16\"",
348
- "intMode": false,
349
- "addBias": false
350
- },
351
- "bindings": ["partials", "y_2", "params_2"],
352
  "dispatch": {
353
- "x": "min(ceilDiv((aRows * attrs.N), (256)), 65535)",
354
- "y": "ceilDiv(ceilDiv((aRows * attrs.N), (256)), 65535)",
355
  "z": 1
356
  }
357
  }
@@ -362,27 +313,16 @@
362
  "priority": 16,
363
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"],
364
  "derive": {
365
- "hasZero": false,
366
- "hasBias": false,
367
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
368
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
369
- "bScalar": "\"u32\"",
370
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
371
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
372
- "M": "aRows",
373
- "K": "attrs.K",
374
- "N": "attrs.N",
375
- "kBlocks": "dim(shapes.bT, 1)",
376
- "blockSize": "attrs.block_size",
377
- "blobSize": "dim(shapes.bT, 2)",
378
- "bits": "attrs.bits",
379
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
380
- "usesF16": "tensorDtypes.aT == \"float16\"",
381
- "bk": 32,
382
- "tileRows": "tiledRegVec4TileRows",
383
- "tileCols": 64,
384
- "threadRows": "tiledRegVec4ThreadRows",
385
- "threadCols": 4,
386
  "alignedBlockLoads": true,
387
  "aVec4Loads": true
388
  },
@@ -390,8 +330,8 @@
390
  {
391
  "id": "main",
392
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
393
- "bindings": ["a_4", "b_3", "scales_2", "y_2"],
394
- "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM" }
395
  }
396
  ]
397
  },
@@ -400,21 +340,6 @@
400
  "priority": 15,
401
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"],
402
  "derive": {
403
- "hasZero": false,
404
- "hasBias": false,
405
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
406
- "bScalar": "\"u32\"",
407
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
408
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
409
- "M": "aRows",
410
- "K": "attrs.K",
411
- "N": "attrs.N",
412
- "kBlocks": "dim(shapes.bT, 1)",
413
- "blockSize": "attrs.block_size",
414
- "blobSize": "dim(shapes.bT, 2)",
415
- "bits": "attrs.bits",
416
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
417
- "usesF16": "tensorDtypes.aT == \"float16\"",
418
  "bk": "tiledRegSelectedBK",
419
  "tileRows": "tiledRegSelectedTileRows",
420
  "tileCols": 64,
@@ -426,7 +351,7 @@
426
  {
427
  "id": "main",
428
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
429
- "bindings": ["a_3", "b_3", "scales_2", "y_2"],
430
  "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" }
431
  }
432
  ]
@@ -435,28 +360,12 @@
435
  "id": "prefill_tiled_default_zero",
436
  "priority": 14,
437
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 64", "dispatchN32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dispatchM32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledWorkgroupFits"],
438
- "derive": {
439
- "hasZero": false,
440
- "hasBias": false,
441
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
442
- "bScalar": "\"u32\"",
443
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
444
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
445
- "M": "aRows",
446
- "K": "attrs.K",
447
- "N": "attrs.N",
448
- "kBlocks": "dim(shapes.bT, 1)",
449
- "blockSize": "attrs.block_size",
450
- "blobSize": "dim(shapes.bT, 2)",
451
- "bits": "attrs.bits",
452
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
453
- "usesF16": "tensorDtypes.aT == \"float16\""
454
- },
455
  "passes": [
456
  {
457
  "id": "main",
458
  "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja",
459
- "bindings": ["a_3", "b_3", "scales_2", "y_2"],
460
  "dispatch": { "x": "dispatchN32", "y": "dispatchM32" }
461
  }
462
  ]
@@ -466,27 +375,16 @@
466
  "priority": 13,
467
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"],
468
  "derive": {
469
- "hasZero": false,
470
- "hasBias": false,
471
- "workgroupSize": "portableWorkgroupSize",
472
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
473
- "bScalar": "\"u32\"",
474
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
475
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
476
- "bits": "attrs.bits",
477
- "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
478
  "wordsPerCol": "smallMWordsPerCol",
479
  "wordsPerBlock": "blobWords",
480
  "kLanes": "smallMKLanes",
481
- "colGroups": "smallMColGroups",
482
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
483
- "usesF16": "tensorDtypes.aT == \"float16\""
484
  },
485
  "passes": [
486
  {
487
  "id": "main",
488
  "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja",
489
- "bindings": ["a_3", "b_3", "scales_2", "y_2", "params_3"],
490
  "dispatch": {
491
  "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)",
492
  "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
@@ -499,23 +397,12 @@
499
  "id": "default_zero",
500
  "priority": 0,
501
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "portableWorkgroupFits"],
502
- "derive": {
503
- "hasZero": false,
504
- "hasBias": false,
505
- "workgroupSize": "portableWorkgroupSize",
506
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
507
- "bScalar": "\"u32\"",
508
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
509
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
510
- "bits": "attrs.bits",
511
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
512
- "usesF16": "tensorDtypes.aT == \"float16\""
513
- },
514
  "passes": [
515
  {
516
  "id": "main",
517
  "shader": "matmul-nbits.wgsl.jinja",
518
- "bindings": ["a_3", "b_3", "scales_2", "y_2", "params_3"],
519
  "dispatch": {
520
  "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
521
  "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
@@ -531,33 +418,22 @@
531
  "derive": {
532
  "gemvNCols": "tunables.GEMV_N_COLS",
533
  "useSubgroups": "device.features.has(\"subgroups\")",
534
- "hasZero": true,
535
- "hasBias": true,
536
- "workgroupSize": "portableWorkgroupSize",
537
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
538
  "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")",
539
  "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"",
540
  "B_VEC_LEN": "attrs.N * gemvVecPerCol",
541
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
542
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
543
- "bits": "attrs.bits",
544
- "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
545
  "vecWords": "gemvVecWords",
546
  "codesPerVec": "gemvCodesPerVec",
547
  "codesPerVec4": "gemvCodesPerVec / 4",
548
  "vecPerBlock": "gemvVecPerBlock",
549
  "vecPerCol": "gemvVecPerCol",
550
- "actVec4": "gemvActVec4",
551
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
552
- "usesF16": "tensorDtypes.aT == \"float16\""
553
  },
554
  "passes": [
555
  {
556
  "id": "main",
557
  "shader": "matmul-nbits-gemv-q4.wgsl.jinja",
558
- "bindings": ["a_2", "b_2", "scales_2", "zero_points", "bias", "y_2", "params"],
559
- "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 },
560
- "subgroupCollectivesWidth": "portable"
561
  }
562
  ]
563
  },
@@ -571,21 +447,6 @@
571
  "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
572
  },
573
  "derive": {
574
- "hasZero": true,
575
- "hasBias": true,
576
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
577
- "bScalar": "\"u32\"",
578
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
579
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
580
- "M": "aRows",
581
- "K": "attrs.K",
582
- "N": "attrs.N",
583
- "kBlocks": "dim(shapes.bT, 1)",
584
- "blockSize": "attrs.block_size",
585
- "blobSize": "dim(shapes.bT, 2)",
586
- "bits": "attrs.bits",
587
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
588
- "usesF16": "tensorDtypes.aT == \"float16\"",
589
  "tileRows": "sgmatTileRows",
590
  "workgroupSize": "sgmatWorkgroupSize",
591
  "rowSubtiles": "sgmatRowSubtiles",
@@ -597,7 +458,7 @@
597
  {
598
  "id": "main",
599
  "shader": "matmul-nbits-q4-sgmat.wgsl.jinja",
600
- "bindings": ["a_3", "b_3", "scales_2", "zero_points", "bias", "y_2"],
601
  "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" }
602
  }
603
  ]
@@ -606,28 +467,18 @@
606
  "id": "prefill_tiled_reg_vec4_splitk_zero_bias",
607
  "priority": 17,
608
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"],
 
609
  "derive": {
610
- "hasZero": true,
611
- "hasBias": true,
612
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
613
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
614
- "bScalar": "\"u32\"",
615
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
616
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
617
- "M": "aRows",
618
- "K": "attrs.K",
619
- "N": "attrs.N",
620
- "kBlocks": "dim(shapes.bT, 1)",
621
- "blockSize": "attrs.block_size",
622
- "blobSize": "dim(shapes.bT, 2)",
623
- "bits": "attrs.bits",
624
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
625
- "usesF16": "tensorDtypes.aT == \"float16\"",
626
  "bk": 32,
627
- "tileRows": "tiledRegVec4TileRows",
628
- "tileCols": 64,
629
- "threadRows": "tiledRegVec4ThreadRows",
630
- "threadCols": 4,
 
 
 
 
631
  "alignedBlockLoads": true,
632
  "aVec4Loads": true,
633
  "splitK": "tiledRegSplitK",
@@ -641,22 +492,17 @@
641
  {
642
  "id": "partial",
643
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
644
- "bindings": ["a_4", "b_3", "scales_2", "zero_points", "y_3"],
645
- "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM", "z": "tiledRegSplitK" }
646
  },
647
  {
648
  "id": "combine",
649
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
650
- "derive": {
651
- "op": "\"sum\"",
652
- "outputF16": "tensorDtypes.aT == \"float16\"",
653
- "intMode": false,
654
- "addBias": true
655
- },
656
- "bindings": ["partials", "bias", "y_2", "params_2"],
657
  "dispatch": {
658
- "x": "min(ceilDiv((aRows * attrs.N), (256)), 65535)",
659
- "y": "ceilDiv(ceilDiv((aRows * attrs.N), (256)), 65535)",
660
  "z": 1
661
  }
662
  }
@@ -667,27 +513,16 @@
667
  "priority": 16,
668
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"],
669
  "derive": {
670
- "hasZero": true,
671
- "hasBias": true,
672
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
673
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
674
- "bScalar": "\"u32\"",
675
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
676
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
677
- "M": "aRows",
678
- "K": "attrs.K",
679
- "N": "attrs.N",
680
- "kBlocks": "dim(shapes.bT, 1)",
681
- "blockSize": "attrs.block_size",
682
- "blobSize": "dim(shapes.bT, 2)",
683
- "bits": "attrs.bits",
684
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
685
- "usesF16": "tensorDtypes.aT == \"float16\"",
686
- "bk": 32,
687
- "tileRows": "tiledRegVec4TileRows",
688
- "tileCols": 64,
689
- "threadRows": "tiledRegVec4ThreadRows",
690
- "threadCols": 4,
691
  "alignedBlockLoads": true,
692
  "aVec4Loads": true
693
  },
@@ -695,8 +530,8 @@
695
  {
696
  "id": "main",
697
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
698
- "bindings": ["a_4", "b_3", "scales_2", "zero_points", "bias", "y_2"],
699
- "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM" }
700
  }
701
  ]
702
  },
@@ -705,21 +540,6 @@
705
  "priority": 15,
706
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"],
707
  "derive": {
708
- "hasZero": true,
709
- "hasBias": true,
710
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
711
- "bScalar": "\"u32\"",
712
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
713
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
714
- "M": "aRows",
715
- "K": "attrs.K",
716
- "N": "attrs.N",
717
- "kBlocks": "dim(shapes.bT, 1)",
718
- "blockSize": "attrs.block_size",
719
- "blobSize": "dim(shapes.bT, 2)",
720
- "bits": "attrs.bits",
721
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
722
- "usesF16": "tensorDtypes.aT == \"float16\"",
723
  "bk": "tiledRegSelectedBK",
724
  "tileRows": "tiledRegSelectedTileRows",
725
  "tileCols": 64,
@@ -731,7 +551,7 @@
731
  {
732
  "id": "main",
733
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
734
- "bindings": ["a_3", "b_3", "scales_2", "zero_points", "bias", "y_2"],
735
  "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" }
736
  }
737
  ]
@@ -740,28 +560,12 @@
740
  "id": "prefill_tiled_zero_bias",
741
  "priority": 14,
742
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 64", "dispatchN32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dispatchM32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledWorkgroupFits"],
743
- "derive": {
744
- "hasZero": true,
745
- "hasBias": true,
746
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
747
- "bScalar": "\"u32\"",
748
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
749
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
750
- "M": "aRows",
751
- "K": "attrs.K",
752
- "N": "attrs.N",
753
- "kBlocks": "dim(shapes.bT, 1)",
754
- "blockSize": "attrs.block_size",
755
- "blobSize": "dim(shapes.bT, 2)",
756
- "bits": "attrs.bits",
757
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
758
- "usesF16": "tensorDtypes.aT == \"float16\""
759
- },
760
  "passes": [
761
  {
762
  "id": "main",
763
  "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja",
764
- "bindings": ["a_3", "b_3", "scales_2", "zero_points", "bias", "y_2"],
765
  "dispatch": { "x": "dispatchN32", "y": "dispatchM32" }
766
  }
767
  ]
@@ -771,27 +575,16 @@
771
  "priority": 13,
772
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"],
773
  "derive": {
774
- "hasZero": true,
775
- "hasBias": true,
776
- "workgroupSize": "portableWorkgroupSize",
777
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
778
- "bScalar": "\"u32\"",
779
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
780
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
781
- "bits": "attrs.bits",
782
- "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
783
  "wordsPerCol": "smallMWordsPerCol",
784
  "wordsPerBlock": "blobWords",
785
  "kLanes": "smallMKLanes",
786
- "colGroups": "smallMColGroups",
787
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
788
- "usesF16": "tensorDtypes.aT == \"float16\""
789
  },
790
  "passes": [
791
  {
792
  "id": "main",
793
  "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja",
794
- "bindings": ["a_3", "b_3", "scales_2", "zero_points", "bias", "y_2", "params_3"],
795
  "dispatch": {
796
  "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)",
797
  "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
@@ -804,23 +597,12 @@
804
  "id": "zero_bias",
805
  "priority": 0,
806
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "portableWorkgroupFits"],
807
- "derive": {
808
- "hasZero": true,
809
- "hasBias": true,
810
- "workgroupSize": "portableWorkgroupSize",
811
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
812
- "bScalar": "\"u32\"",
813
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
814
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
815
- "bits": "attrs.bits",
816
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
817
- "usesF16": "tensorDtypes.aT == \"float16\""
818
- },
819
  "passes": [
820
  {
821
  "id": "main",
822
  "shader": "matmul-nbits.wgsl.jinja",
823
- "bindings": ["a_3", "b_3", "scales_2", "zero_points", "bias", "y_2", "params_3"],
824
  "dispatch": {
825
  "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
826
  "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
@@ -836,33 +618,22 @@
836
  "derive": {
837
  "gemvNCols": "tunables.GEMV_N_COLS",
838
  "useSubgroups": "device.features.has(\"subgroups\")",
839
- "hasZero": true,
840
- "hasBias": false,
841
- "workgroupSize": "portableWorkgroupSize",
842
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
843
  "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")",
844
  "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"",
845
  "B_VEC_LEN": "attrs.N * gemvVecPerCol",
846
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
847
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
848
- "bits": "attrs.bits",
849
- "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
850
  "vecWords": "gemvVecWords",
851
  "codesPerVec": "gemvCodesPerVec",
852
  "codesPerVec4": "gemvCodesPerVec / 4",
853
  "vecPerBlock": "gemvVecPerBlock",
854
  "vecPerCol": "gemvVecPerCol",
855
- "actVec4": "gemvActVec4",
856
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
857
- "usesF16": "tensorDtypes.aT == \"float16\""
858
  },
859
  "passes": [
860
  {
861
  "id": "main",
862
  "shader": "matmul-nbits-gemv-q4.wgsl.jinja",
863
- "bindings": ["a_2", "b_2", "scales_2", "zero_points", "y_2", "params"],
864
- "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 },
865
- "subgroupCollectivesWidth": "portable"
866
  }
867
  ]
868
  },
@@ -876,21 +647,6 @@
876
  "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
877
  },
878
  "derive": {
879
- "hasZero": true,
880
- "hasBias": false,
881
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
882
- "bScalar": "\"u32\"",
883
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
884
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
885
- "M": "aRows",
886
- "K": "attrs.K",
887
- "N": "attrs.N",
888
- "kBlocks": "dim(shapes.bT, 1)",
889
- "blockSize": "attrs.block_size",
890
- "blobSize": "dim(shapes.bT, 2)",
891
- "bits": "attrs.bits",
892
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
893
- "usesF16": "tensorDtypes.aT == \"float16\"",
894
  "tileRows": "sgmatTileRows",
895
  "workgroupSize": "sgmatWorkgroupSize",
896
  "rowSubtiles": "sgmatRowSubtiles",
@@ -902,7 +658,7 @@
902
  {
903
  "id": "main",
904
  "shader": "matmul-nbits-q4-sgmat.wgsl.jinja",
905
- "bindings": ["a_3", "b_3", "scales_2", "zero_points", "y_2"],
906
  "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" }
907
  }
908
  ]
@@ -911,28 +667,18 @@
911
  "id": "prefill_tiled_reg_vec4_splitk_zero_only",
912
  "priority": 17,
913
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"],
 
914
  "derive": {
915
- "hasZero": true,
916
- "hasBias": false,
917
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
918
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
919
- "bScalar": "\"u32\"",
920
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
921
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
922
- "M": "aRows",
923
- "K": "attrs.K",
924
- "N": "attrs.N",
925
- "kBlocks": "dim(shapes.bT, 1)",
926
- "blockSize": "attrs.block_size",
927
- "blobSize": "dim(shapes.bT, 2)",
928
- "bits": "attrs.bits",
929
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
930
- "usesF16": "tensorDtypes.aT == \"float16\"",
931
  "bk": 32,
932
- "tileRows": "tiledRegVec4TileRows",
933
- "tileCols": 64,
934
- "threadRows": "tiledRegVec4ThreadRows",
935
- "threadCols": 4,
 
 
 
 
936
  "alignedBlockLoads": true,
937
  "aVec4Loads": true,
938
  "splitK": "tiledRegSplitK",
@@ -946,22 +692,17 @@
946
  {
947
  "id": "partial",
948
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
949
- "bindings": ["a_4", "b_3", "scales_2", "zero_points", "y_3"],
950
- "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM", "z": "tiledRegSplitK" }
951
  },
952
  {
953
  "id": "combine",
954
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
955
- "derive": {
956
- "op": "\"sum\"",
957
- "outputF16": "tensorDtypes.aT == \"float16\"",
958
- "intMode": false,
959
- "addBias": false
960
- },
961
- "bindings": ["partials", "y_2", "params_2"],
962
  "dispatch": {
963
- "x": "min(ceilDiv((aRows * attrs.N), (256)), 65535)",
964
- "y": "ceilDiv(ceilDiv((aRows * attrs.N), (256)), 65535)",
965
  "z": 1
966
  }
967
  }
@@ -972,27 +713,16 @@
972
  "priority": 16,
973
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"],
974
  "derive": {
975
- "hasZero": true,
976
- "hasBias": false,
977
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
978
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
979
- "bScalar": "\"u32\"",
980
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
981
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
982
- "M": "aRows",
983
- "K": "attrs.K",
984
- "N": "attrs.N",
985
- "kBlocks": "dim(shapes.bT, 1)",
986
- "blockSize": "attrs.block_size",
987
- "blobSize": "dim(shapes.bT, 2)",
988
- "bits": "attrs.bits",
989
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
990
- "usesF16": "tensorDtypes.aT == \"float16\"",
991
- "bk": 32,
992
- "tileRows": "tiledRegVec4TileRows",
993
- "tileCols": 64,
994
- "threadRows": "tiledRegVec4ThreadRows",
995
- "threadCols": 4,
996
  "alignedBlockLoads": true,
997
  "aVec4Loads": true
998
  },
@@ -1000,8 +730,8 @@
1000
  {
1001
  "id": "main",
1002
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
1003
- "bindings": ["a_4", "b_3", "scales_2", "zero_points", "y_2"],
1004
- "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM" }
1005
  }
1006
  ]
1007
  },
@@ -1010,21 +740,6 @@
1010
  "priority": 15,
1011
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"],
1012
  "derive": {
1013
- "hasZero": true,
1014
- "hasBias": false,
1015
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1016
- "bScalar": "\"u32\"",
1017
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1018
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1019
- "M": "aRows",
1020
- "K": "attrs.K",
1021
- "N": "attrs.N",
1022
- "kBlocks": "dim(shapes.bT, 1)",
1023
- "blockSize": "attrs.block_size",
1024
- "blobSize": "dim(shapes.bT, 2)",
1025
- "bits": "attrs.bits",
1026
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1027
- "usesF16": "tensorDtypes.aT == \"float16\"",
1028
  "bk": "tiledRegSelectedBK",
1029
  "tileRows": "tiledRegSelectedTileRows",
1030
  "tileCols": 64,
@@ -1036,7 +751,7 @@
1036
  {
1037
  "id": "main",
1038
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
1039
- "bindings": ["a_3", "b_3", "scales_2", "zero_points", "y_2"],
1040
  "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" }
1041
  }
1042
  ]
@@ -1045,28 +760,12 @@
1045
  "id": "prefill_tiled_zero_only",
1046
  "priority": 14,
1047
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 64", "dispatchN32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dispatchM32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledWorkgroupFits"],
1048
- "derive": {
1049
- "hasZero": true,
1050
- "hasBias": false,
1051
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1052
- "bScalar": "\"u32\"",
1053
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1054
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1055
- "M": "aRows",
1056
- "K": "attrs.K",
1057
- "N": "attrs.N",
1058
- "kBlocks": "dim(shapes.bT, 1)",
1059
- "blockSize": "attrs.block_size",
1060
- "blobSize": "dim(shapes.bT, 2)",
1061
- "bits": "attrs.bits",
1062
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1063
- "usesF16": "tensorDtypes.aT == \"float16\""
1064
- },
1065
  "passes": [
1066
  {
1067
  "id": "main",
1068
  "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja",
1069
- "bindings": ["a_3", "b_3", "scales_2", "zero_points", "y_2"],
1070
  "dispatch": { "x": "dispatchN32", "y": "dispatchM32" }
1071
  }
1072
  ]
@@ -1076,27 +775,16 @@
1076
  "priority": 13,
1077
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"],
1078
  "derive": {
1079
- "hasZero": true,
1080
- "hasBias": false,
1081
- "workgroupSize": "portableWorkgroupSize",
1082
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1083
- "bScalar": "\"u32\"",
1084
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1085
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1086
- "bits": "attrs.bits",
1087
- "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
1088
  "wordsPerCol": "smallMWordsPerCol",
1089
  "wordsPerBlock": "blobWords",
1090
  "kLanes": "smallMKLanes",
1091
- "colGroups": "smallMColGroups",
1092
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1093
- "usesF16": "tensorDtypes.aT == \"float16\""
1094
  },
1095
  "passes": [
1096
  {
1097
  "id": "main",
1098
  "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja",
1099
- "bindings": ["a_3", "b_3", "scales_2", "zero_points", "y_2", "params_3"],
1100
  "dispatch": {
1101
  "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)",
1102
  "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
@@ -1109,23 +797,12 @@
1109
  "id": "zero_only",
1110
  "priority": 0,
1111
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "portableWorkgroupFits"],
1112
- "derive": {
1113
- "hasZero": true,
1114
- "hasBias": false,
1115
- "workgroupSize": "portableWorkgroupSize",
1116
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1117
- "bScalar": "\"u32\"",
1118
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1119
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1120
- "bits": "attrs.bits",
1121
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1122
- "usesF16": "tensorDtypes.aT == \"float16\""
1123
- },
1124
  "passes": [
1125
  {
1126
  "id": "main",
1127
  "shader": "matmul-nbits.wgsl.jinja",
1128
- "bindings": ["a_3", "b_3", "scales_2", "zero_points", "y_2", "params_3"],
1129
  "dispatch": {
1130
  "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
1131
  "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
@@ -1141,33 +818,22 @@
1141
  "derive": {
1142
  "gemvNCols": "tunables.GEMV_N_COLS",
1143
  "useSubgroups": "device.features.has(\"subgroups\")",
1144
- "hasZero": false,
1145
- "hasBias": true,
1146
- "workgroupSize": "portableWorkgroupSize",
1147
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1148
  "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")",
1149
  "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"",
1150
  "B_VEC_LEN": "attrs.N * gemvVecPerCol",
1151
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1152
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1153
- "bits": "attrs.bits",
1154
- "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
1155
  "vecWords": "gemvVecWords",
1156
  "codesPerVec": "gemvCodesPerVec",
1157
  "codesPerVec4": "gemvCodesPerVec / 4",
1158
  "vecPerBlock": "gemvVecPerBlock",
1159
  "vecPerCol": "gemvVecPerCol",
1160
- "actVec4": "gemvActVec4",
1161
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1162
- "usesF16": "tensorDtypes.aT == \"float16\""
1163
  },
1164
  "passes": [
1165
  {
1166
  "id": "main",
1167
  "shader": "matmul-nbits-gemv-q4.wgsl.jinja",
1168
- "bindings": ["a_2", "b_2", "scales_2", "bias", "y_2", "params"],
1169
- "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 },
1170
- "subgroupCollectivesWidth": "portable"
1171
  }
1172
  ]
1173
  },
@@ -1181,21 +847,6 @@
1181
  "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
1182
  },
1183
  "derive": {
1184
- "hasZero": false,
1185
- "hasBias": true,
1186
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1187
- "bScalar": "\"u32\"",
1188
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1189
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1190
- "M": "aRows",
1191
- "K": "attrs.K",
1192
- "N": "attrs.N",
1193
- "kBlocks": "dim(shapes.bT, 1)",
1194
- "blockSize": "attrs.block_size",
1195
- "blobSize": "dim(shapes.bT, 2)",
1196
- "bits": "attrs.bits",
1197
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1198
- "usesF16": "tensorDtypes.aT == \"float16\"",
1199
  "tileRows": "sgmatTileRows",
1200
  "workgroupSize": "sgmatWorkgroupSize",
1201
  "rowSubtiles": "sgmatRowSubtiles",
@@ -1207,7 +858,7 @@
1207
  {
1208
  "id": "main",
1209
  "shader": "matmul-nbits-q4-sgmat.wgsl.jinja",
1210
- "bindings": ["a_3", "b_3", "scales_2", "bias", "y_2"],
1211
  "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" }
1212
  }
1213
  ]
@@ -1216,28 +867,18 @@
1216
  "id": "prefill_tiled_reg_vec4_splitk_bias_only",
1217
  "priority": 17,
1218
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"],
 
1219
  "derive": {
1220
- "hasZero": false,
1221
- "hasBias": true,
1222
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1223
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
1224
- "bScalar": "\"u32\"",
1225
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1226
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1227
- "M": "aRows",
1228
- "K": "attrs.K",
1229
- "N": "attrs.N",
1230
- "kBlocks": "dim(shapes.bT, 1)",
1231
- "blockSize": "attrs.block_size",
1232
- "blobSize": "dim(shapes.bT, 2)",
1233
- "bits": "attrs.bits",
1234
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1235
- "usesF16": "tensorDtypes.aT == \"float16\"",
1236
  "bk": 32,
1237
- "tileRows": "tiledRegVec4TileRows",
1238
- "tileCols": 64,
1239
- "threadRows": "tiledRegVec4ThreadRows",
1240
- "threadCols": 4,
 
 
 
 
1241
  "alignedBlockLoads": true,
1242
  "aVec4Loads": true,
1243
  "splitK": "tiledRegSplitK",
@@ -1251,22 +892,17 @@
1251
  {
1252
  "id": "partial",
1253
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
1254
- "bindings": ["a_4", "b_3", "scales_2", "y_3"],
1255
- "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM", "z": "tiledRegSplitK" }
1256
  },
1257
  {
1258
  "id": "combine",
1259
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
1260
- "derive": {
1261
- "op": "\"sum\"",
1262
- "outputF16": "tensorDtypes.aT == \"float16\"",
1263
- "intMode": false,
1264
- "addBias": true
1265
- },
1266
- "bindings": ["partials", "bias", "y_2", "params_2"],
1267
  "dispatch": {
1268
- "x": "min(ceilDiv((aRows * attrs.N), (256)), 65535)",
1269
- "y": "ceilDiv(ceilDiv((aRows * attrs.N), (256)), 65535)",
1270
  "z": 1
1271
  }
1272
  }
@@ -1277,27 +913,16 @@
1277
  "priority": 16,
1278
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"],
1279
  "derive": {
1280
- "hasZero": false,
1281
- "hasBias": true,
1282
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1283
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
1284
- "bScalar": "\"u32\"",
1285
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1286
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1287
- "M": "aRows",
1288
- "K": "attrs.K",
1289
- "N": "attrs.N",
1290
- "kBlocks": "dim(shapes.bT, 1)",
1291
- "blockSize": "attrs.block_size",
1292
- "blobSize": "dim(shapes.bT, 2)",
1293
- "bits": "attrs.bits",
1294
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1295
- "usesF16": "tensorDtypes.aT == \"float16\"",
1296
- "bk": 32,
1297
- "tileRows": "tiledRegVec4TileRows",
1298
- "tileCols": 64,
1299
- "threadRows": "tiledRegVec4ThreadRows",
1300
- "threadCols": 4,
1301
  "alignedBlockLoads": true,
1302
  "aVec4Loads": true
1303
  },
@@ -1305,8 +930,8 @@
1305
  {
1306
  "id": "main",
1307
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
1308
- "bindings": ["a_4", "b_3", "scales_2", "bias", "y_2"],
1309
- "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM" }
1310
  }
1311
  ]
1312
  },
@@ -1315,21 +940,6 @@
1315
  "priority": 15,
1316
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"],
1317
  "derive": {
1318
- "hasZero": false,
1319
- "hasBias": true,
1320
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1321
- "bScalar": "\"u32\"",
1322
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1323
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1324
- "M": "aRows",
1325
- "K": "attrs.K",
1326
- "N": "attrs.N",
1327
- "kBlocks": "dim(shapes.bT, 1)",
1328
- "blockSize": "attrs.block_size",
1329
- "blobSize": "dim(shapes.bT, 2)",
1330
- "bits": "attrs.bits",
1331
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1332
- "usesF16": "tensorDtypes.aT == \"float16\"",
1333
  "bk": "tiledRegSelectedBK",
1334
  "tileRows": "tiledRegSelectedTileRows",
1335
  "tileCols": 64,
@@ -1341,7 +951,7 @@
1341
  {
1342
  "id": "main",
1343
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
1344
- "bindings": ["a_3", "b_3", "scales_2", "bias", "y_2"],
1345
  "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" }
1346
  }
1347
  ]
@@ -1350,28 +960,12 @@
1350
  "id": "prefill_tiled_bias_only",
1351
  "priority": 14,
1352
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 64", "dispatchN32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dispatchM32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledWorkgroupFits"],
1353
- "derive": {
1354
- "hasZero": false,
1355
- "hasBias": true,
1356
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1357
- "bScalar": "\"u32\"",
1358
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1359
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1360
- "M": "aRows",
1361
- "K": "attrs.K",
1362
- "N": "attrs.N",
1363
- "kBlocks": "dim(shapes.bT, 1)",
1364
- "blockSize": "attrs.block_size",
1365
- "blobSize": "dim(shapes.bT, 2)",
1366
- "bits": "attrs.bits",
1367
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1368
- "usesF16": "tensorDtypes.aT == \"float16\""
1369
- },
1370
  "passes": [
1371
  {
1372
  "id": "main",
1373
  "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja",
1374
- "bindings": ["a_3", "b_3", "scales_2", "bias", "y_2"],
1375
  "dispatch": { "x": "dispatchN32", "y": "dispatchM32" }
1376
  }
1377
  ]
@@ -1381,27 +975,16 @@
1381
  "priority": 13,
1382
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"],
1383
  "derive": {
1384
- "hasZero": false,
1385
- "hasBias": true,
1386
- "workgroupSize": "portableWorkgroupSize",
1387
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1388
- "bScalar": "\"u32\"",
1389
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1390
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1391
- "bits": "attrs.bits",
1392
- "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
1393
  "wordsPerCol": "smallMWordsPerCol",
1394
  "wordsPerBlock": "blobWords",
1395
  "kLanes": "smallMKLanes",
1396
- "colGroups": "smallMColGroups",
1397
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1398
- "usesF16": "tensorDtypes.aT == \"float16\""
1399
  },
1400
  "passes": [
1401
  {
1402
  "id": "main",
1403
  "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja",
1404
- "bindings": ["a_3", "b_3", "scales_2", "bias", "y_2", "params_3"],
1405
  "dispatch": {
1406
  "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)",
1407
  "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
@@ -1414,23 +997,12 @@
1414
  "id": "bias_only",
1415
  "priority": 0,
1416
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "portableWorkgroupFits"],
1417
- "derive": {
1418
- "hasZero": false,
1419
- "hasBias": true,
1420
- "workgroupSize": "portableWorkgroupSize",
1421
- "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1422
- "bScalar": "\"u32\"",
1423
- "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1424
- "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
1425
- "bits": "attrs.bits",
1426
- "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
1427
- "usesF16": "tensorDtypes.aT == \"float16\""
1428
- },
1429
  "passes": [
1430
  {
1431
  "id": "main",
1432
  "shader": "matmul-nbits.wgsl.jinja",
1433
- "bindings": ["a_3", "b_3", "scales_2", "bias", "y_2", "params_3"],
1434
  "dispatch": {
1435
  "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
1436
  "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
 
37
  "derive": {
38
  "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
39
  "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32",
40
+ "subgroupsWave32": "device.features.has(\"subgroups\") and wave32Adapter",
41
  "narrowSubgroupRange": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize < device.adapterInfo.subgroupMaxSize and device.adapterInfo.subgroupMaxSize <= 16",
42
  "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",
43
  "pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter",
44
  "wave32Effective": "wave32Adapter or pinSubgroupSize32",
45
+ "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",
46
  "packedFeature": "device.wgslLanguageFeatures.has(\"packed_4x8_integer_dot_product\")",
47
  "kBlocksExpected": "ceilDiv(attrs.K, attrs.block_size)",
48
  "blobSizeExpected": "ceilDiv(attrs.block_size * attrs.bits, 8)",
 
75
  "zeroOnlyEpilogue": "zeroPointsValid and not present.biasT",
76
  "biasOnlyEpilogue": "not present.zeroPointsT and biasValid",
77
  "portableWorkgroupFits": "portableWorkgroupSize > 0 and portableWorkgroupSize * 64 <= device.limits.maxComputeWorkgroupStorageSize",
78
+ "tiledWorkgroupSide": 16,
79
+ "dp4aQuantizeWorkgroupSize": 64,
80
+ "tiledWorkgroupFits": "tiledWorkgroupSide <= device.limits.maxComputeWorkgroupSizeX and tiledWorkgroupSide <= device.limits.maxComputeWorkgroupSizeY and tiledWorkgroupSide * tiledWorkgroupSide <= device.limits.maxComputeInvocationsPerWorkgroup and 4096 <= device.limits.maxComputeWorkgroupStorageSize",
81
  "tiledRegWorkgroupFits": "tiledWorkgroupFits and 16384 <= device.limits.maxComputeWorkgroupStorageSize",
82
  "mediumTiledRegWorkgroupFits": "tiledWorkgroupFits and 6144 <= device.limits.maxComputeWorkgroupStorageSize",
83
  "sgmatWorkgroupFits": "sgmatWorkgroupSize <= deviceWorkgroupCap and sgmatWorkgroupStorageBytes <= device.limits.maxComputeWorkgroupStorageSize",
 
90
  "mediumTiledRegEligible": "mediumRegisterEligible and mediumTiledRegWorkgroupFits and dispatchM32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
91
  "tiledRegVariantEligible": "largeTiledRegEligible or mediumTiledRegEligible",
92
  "tiledRegSelectedBK": "tiledRegBK if largeTiledRegEligible else 16",
 
93
  "tiledRegSelectedThreadRows": "4 if largeTiledRegEligible else 2",
94
+ "tiledRegSelectedTileRows": "tiledRegSelectedThreadRows * tiledWorkgroupSide",
95
  "tiledRegSelectedDispatchM": "dispatchM64 if largeTiledRegEligible else dispatchM32",
96
  "blobWords": "blobSizeExpected / 4",
97
  "codesPerWord": "32 / attrs.bits",
 
104
  "smallMKLanes": "min(portableWorkgroupSize, pow2ceil(smallMWordsPerCol + 1) / 2)",
105
  "smallMColGroups": "portableWorkgroupSize / smallMKLanes",
106
  "smallMDispatchN": "ceilDiv(attrs.N, 4 * smallMColGroups)",
 
107
  "tiledRegVec4ThreadRows": "8 if aRows >= 256 else 4",
108
+ "tiledRegVec4TileRows": "tiledRegVec4ThreadRows * tiledWorkgroupSide",
109
  "tiledRegVec4DispatchM": "ceilDiv(aRows, tiledRegVec4TileRows)",
110
  "tiledRegVec4Eligible": "largeTiledRegEligible and attrs.K % 4 == 0 and attrs.block_size % 32 == 0 and tiledRegVec4DispatchM <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
111
+ "tiledRegSubgroupPin": "device.adapterInfo.subgroupMinSize if (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 < device.adapterInfo.subgroupMaxSize and device.adapterInfo.subgroupMinSize <= tiledWorkgroupSide and attrs.bits == 2 and tensorDtypes.aT == \"float16\" and tiledRegVec4Eligible and aRows >= 2 * tiledRegVec4TileRows) else 0",
112
+ "q2WideMicroPreferred": "variableSubgroup16To32 and device.features.has(\"subgroup-size-control\") and attrs.bits == 2 and tensorDtypes.aT == \"float16\" and tiledRegVec4ThreadRows == 4 and aRows >= 2 * tiledRegVec4TileRows and device.limits.maxComputeWorkgroupSizeX >= 32 and device.limits.maxComputeWorkgroupSizeY >= 8 and device.limits.maxComputeInvocationsPerWorkgroup >= 256",
113
  "tiledRegSplitTiles": "tiledRegVec4DispatchM * dispatchN64",
114
+ "q2NarrowBkPreferred": "variableSubgroup16To32 and attrs.bits == 2 and tensorDtypes.aT == \"float16\" and attrs.block_size >= 128 and attrs.block_size % 128 == 0 and attrs.K % attrs.block_size == 0 and tiledRegVec4TileRows == 64 and aRows >= 2 * tiledRegVec4TileRows and tiledRegSplitTiles >= tunables.REGISTER_TILE_SPLITK_MAX_TILES / 2",
115
  "tiledRegSplitWant": "ceilDiv(tunables.REGISTER_TILE_SPLITK_TARGET_WORKGROUPS, max(1, tiledRegSplitTiles))",
116
  "tiledRegSplitK": "8 if (tiledRegSplitWant >= 8 and attrs.K >= 4096) else (4 if (tiledRegSplitWant >= 4 and attrs.K >= 2048) else (2 if (tiledRegSplitWant >= 2 and attrs.K >= 1024) else 1))",
117
  "tiledRegSplitTilesPerSplit": "ceilDiv(ceilDiv(attrs.K, 32), tiledRegSplitK)",
118
+ "tiledRegSplitEligible": "tiledRegVec4Eligible and tiledRegSplitK >= 2 and tiledRegSplitTiles <= tunables.REGISTER_TILE_SPLITK_MAX_TILES and not (q2NarrowBkPreferred and attrs.block_size >= 256) and tiledRegSplitK * aRows * attrs.N * 4 <= device.limits.maxStorageBufferBindingSize and tiledRegSplitK <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
119
+ "hasZero": "present.zeroPointsT",
120
+ "hasBias": "present.biasT",
121
+ "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
122
+ "bScalar": "\"u32\"",
123
+ "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
124
+ "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"",
125
+ "usesF16": "tensorDtypes.aT == \"float16\"",
126
+ "M": "aRows",
127
+ "K": "attrs.K",
128
+ "N": "attrs.N",
129
+ "kBlocks": "dim(shapes.bT, 1)",
130
+ "blockSize": "attrs.block_size",
131
+ "blobSize": "dim(shapes.bT, 2)",
132
+ "bits": "attrs.bits",
133
+ "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)",
134
+ "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")",
135
+ "workgroupSize": "portableWorkgroupSize",
136
  "B_LEN": "attrs.N * kBlocksExpected * blobSizeExpected / 4",
137
  "SCALES_LEN": "attrs.N * kBlocksExpected",
138
  "BIAS_LEN": "attrs.N"
139
  },
140
  "bindings": {
141
+ "a_a_t": { "arg": "aT", "name": "a", "elementType": "$aElement" },
142
+ "b_b_t": { "arg": "bT", "name": "b", "elementType": "$bElement", "length": "$B_VEC_LEN" },
143
+ "scales_main": { "arg": "scalesT", "name": "scales", "elementType": "$scaleScalar", "length": "$SCALES_LEN" },
144
+ "y_y_t": { "arg": "yT", "name": "y", "elementType": "$outputScalar" },
 
 
 
 
 
 
 
 
 
 
 
 
145
  "params": {
 
146
  "struct": [
147
  { "name": "K", "type": "u32", "value": "attrs.K" },
148
  { "name": "N", "type": "u32", "value": "attrs.N" },
149
  { "name": "kBlocks", "type": "u32", "value": "dim(shapes.bT, 1)" }
150
  ]
151
  },
152
+ "zero_points": { "arg": "zeroPointsT", "elementType": "$aScalar", "length": "$SCALES_LEN" },
153
+ "bias": { "arg": "biasT", "elementType": "$aScalar", "length": "$BIAS_LEN" },
154
+ "a_main": { "arg": "aT", "name": "a", "elementType": "$aScalar" },
155
+ "b_main": { "arg": "bT", "name": "b", "elementType": "$bScalar", "length": "$B_LEN" },
156
+ "a_partial": { "arg": "aT", "name": "a", "elementType": "$aVec4Element" },
157
+ "y_f32": { "scratch": "partials", "name": "y", "elementType": "f32" },
 
 
 
 
 
158
  "partials": { "buffer": "read-only-storage", "elementType": "f32" },
159
+ "params_cols": { "name": "params", "struct": [{ "name": "cols", "type": "u32", "value": "aRows * attrs.N" }] },
160
+ "params_main": {
161
  "name": "params",
 
 
 
 
 
 
162
  "struct": [
163
  { "name": "rows", "type": "u32", "value": "aRows" },
164
  { "name": "K", "type": "u32", "value": "attrs.K" },
 
173
  {
174
  "id": "q4_dp4a_prefill",
175
  "priority": 19,
176
+ "when": ["packedFeature", "commonShapeValid", "defaultEpilogue", "attrs.bits == 4", "attrs.accuracy_level == 4", "tensorDtypes.aT == \"float32\"", "attrs.block_size % 32 == 0", "attrs.K % 128 == 0", "attrs.N % 16 == 0", "aRows >= 32", "ceilDiv(attrs.N, 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(aRows, 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledWorkgroupFits", "dp4aQuantizeWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX", "4608 <= device.limits.maxComputeWorkgroupStorageSize"],
177
  "demoteWhen": ["has(device.adapterInfo, \"architecture\") and device.adapterInfo.architecture == \"maxwell\"", "device.features.has(\"chromium-experimental-subgroup-matrix\")", "device.adapterInfo.vendor == \"apple\""],
178
+ "derive": { "vec4Count": "aRows * attrs.K / 4", "blockCount": "aRows * attrs.K / 128" },
 
 
 
 
 
 
 
 
 
179
  "intermediates": [
180
  { "id": "aQuant", "dtype": "uint32", "shape": "[aRows * attrs.K / 4]" },
181
  { "id": "aScales", "dtype": "float32", "shape": "[aRows * attrs.K / 128]" }
 
191
  { "scratch": "aScales", "name": "a_scales", "elementType": "f32" }
192
  ],
193
  "dispatch": {
194
+ "x": "min(ceilDiv((aRows * attrs.K / 4), (dp4aQuantizeWorkgroupSize)), 65535)",
195
+ "y": "ceilDiv(ceilDiv((aRows * attrs.K / 4), (dp4aQuantizeWorkgroupSize)), 65535)",
196
  "z": 1
197
  }
198
  },
 
218
  "derive": {
219
  "gemvNCols": "tunables.GEMV_N_COLS",
220
  "useSubgroups": "device.features.has(\"subgroups\")",
 
 
 
 
221
  "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")",
222
  "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"",
223
  "B_VEC_LEN": "attrs.N * gemvVecPerCol",
 
 
 
 
224
  "vecWords": "gemvVecWords",
225
  "codesPerVec": "gemvCodesPerVec",
226
  "codesPerVec4": "gemvCodesPerVec / 4",
227
  "vecPerBlock": "gemvVecPerBlock",
228
  "vecPerCol": "gemvVecPerCol",
229
+ "actVec4": "gemvActVec4"
 
 
230
  },
231
  "passes": [
232
  {
233
  "id": "main",
234
  "shader": "matmul-nbits-gemv-q4.wgsl.jinja",
235
+ "bindings": ["a_a_t", "b_b_t", "scales_main", "y_y_t", "params"],
236
+ "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 }
 
237
  }
238
  ]
239
  },
 
247
  "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
248
  },
249
  "derive": {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
250
  "tileRows": "sgmatTileRows",
251
  "workgroupSize": "sgmatWorkgroupSize",
252
  "rowSubtiles": "sgmatRowSubtiles",
 
258
  {
259
  "id": "main",
260
  "shader": "matmul-nbits-q4-sgmat.wgsl.jinja",
261
+ "bindings": ["a_main", "b_main", "scales_main", "y_y_t"],
262
  "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" }
263
  }
264
  ]
 
267
  "id": "prefill_tiled_reg_vec4_splitk_default_zero",
268
  "priority": 17,
269
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"],
270
+ "demoteWhen": ["variableSubgroup16To32 and attrs.bits >= 4 and tiledRegSplitK >= 4 and aRows >= 2 * tiledRegVec4TileRows"],
271
  "derive": {
 
 
 
272
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
 
 
 
 
 
 
 
 
 
 
 
 
273
  "bk": 32,
274
+ "q2ChunkLoads": "attrs.bits == 2 and attrs.K % 16 == 0 and attrs.block_size % 16 == 0 and subgroupsWave32 and kBlocksExpected >= device.adapterInfo.subgroupMinSize and dispatchN64 >= device.adapterInfo.subgroupMinSize and aRows >= 2 * tiledRegVec4TileRows",
275
+ "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide",
276
+ "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide",
277
+ "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows",
278
+ "threadCols": "2 if q2WideMicroPreferred else 4",
279
+ "tileRows": "threadRows * tileWorkgroupY",
280
+ "tileCols": "threadCols * tileWorkgroupX",
281
+ "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin",
282
  "alignedBlockLoads": true,
283
  "aVec4Loads": true,
284
  "splitK": "tiledRegSplitK",
 
292
  {
293
  "id": "partial",
294
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
295
+ "bindings": ["a_partial", "b_main", "scales_main", "y_f32"],
296
+ "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)", "z": "tiledRegSplitK" }
297
  },
298
  {
299
  "id": "combine",
300
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
301
+ "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "addBias": "hasBias" },
302
+ "bindings": ["partials", "y_y_t", "params_cols"],
 
 
 
 
 
303
  "dispatch": {
304
+ "x": "min(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)",
305
+ "y": "ceilDiv(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)",
306
  "z": 1
307
  }
308
  }
 
313
  "priority": 16,
314
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"],
315
  "derive": {
 
 
 
316
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
317
+ "bk": "16 if q2NarrowBkPreferred else 32",
318
+ "q2ChunkLoads": "attrs.bits == 2 and attrs.K % 16 == 0 and attrs.block_size % 16 == 0 and subgroupsWave32 and kBlocksExpected >= device.adapterInfo.subgroupMinSize and dispatchN64 >= device.adapterInfo.subgroupMinSize and aRows >= 2 * tiledRegVec4TileRows",
319
+ "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide",
320
+ "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide",
321
+ "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows",
322
+ "threadCols": "2 if q2WideMicroPreferred else 4",
323
+ "tileRows": "threadRows * tileWorkgroupY",
324
+ "tileCols": "threadCols * tileWorkgroupX",
325
+ "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin",
 
 
 
 
 
 
 
 
326
  "alignedBlockLoads": true,
327
  "aVec4Loads": true
328
  },
 
330
  {
331
  "id": "main",
332
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
333
+ "bindings": ["a_partial", "b_main", "scales_main", "y_y_t"],
334
+ "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)" }
335
  }
336
  ]
337
  },
 
340
  "priority": 15,
341
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"],
342
  "derive": {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
343
  "bk": "tiledRegSelectedBK",
344
  "tileRows": "tiledRegSelectedTileRows",
345
  "tileCols": 64,
 
351
  {
352
  "id": "main",
353
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
354
+ "bindings": ["a_main", "b_main", "scales_main", "y_y_t"],
355
  "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" }
356
  }
357
  ]
 
360
  "id": "prefill_tiled_default_zero",
361
  "priority": 14,
362
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 64", "dispatchN32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dispatchM32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledWorkgroupFits"],
363
+ "derive": {},
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
364
  "passes": [
365
  {
366
  "id": "main",
367
  "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja",
368
+ "bindings": ["a_main", "b_main", "scales_main", "y_y_t"],
369
  "dispatch": { "x": "dispatchN32", "y": "dispatchM32" }
370
  }
371
  ]
 
375
  "priority": 13,
376
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"],
377
  "derive": {
 
 
 
 
 
 
 
 
 
378
  "wordsPerCol": "smallMWordsPerCol",
379
  "wordsPerBlock": "blobWords",
380
  "kLanes": "smallMKLanes",
381
+ "colGroups": "smallMColGroups"
 
 
382
  },
383
  "passes": [
384
  {
385
  "id": "main",
386
  "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja",
387
+ "bindings": ["a_main", "b_main", "scales_main", "y_y_t", "params_main"],
388
  "dispatch": {
389
  "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)",
390
  "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
 
397
  "id": "default_zero",
398
  "priority": 0,
399
  "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "portableWorkgroupFits"],
400
+ "derive": {},
 
 
 
 
 
 
 
 
 
 
 
401
  "passes": [
402
  {
403
  "id": "main",
404
  "shader": "matmul-nbits.wgsl.jinja",
405
+ "bindings": ["a_main", "b_main", "scales_main", "y_y_t", "params_main"],
406
  "dispatch": {
407
  "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
408
  "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
 
418
  "derive": {
419
  "gemvNCols": "tunables.GEMV_N_COLS",
420
  "useSubgroups": "device.features.has(\"subgroups\")",
 
 
 
 
421
  "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")",
422
  "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"",
423
  "B_VEC_LEN": "attrs.N * gemvVecPerCol",
 
 
 
 
424
  "vecWords": "gemvVecWords",
425
  "codesPerVec": "gemvCodesPerVec",
426
  "codesPerVec4": "gemvCodesPerVec / 4",
427
  "vecPerBlock": "gemvVecPerBlock",
428
  "vecPerCol": "gemvVecPerCol",
429
+ "actVec4": "gemvActVec4"
 
 
430
  },
431
  "passes": [
432
  {
433
  "id": "main",
434
  "shader": "matmul-nbits-gemv-q4.wgsl.jinja",
435
+ "bindings": ["a_a_t", "b_b_t", "scales_main", "zero_points", "bias", "y_y_t", "params"],
436
+ "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 }
 
437
  }
438
  ]
439
  },
 
447
  "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
448
  },
449
  "derive": {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
450
  "tileRows": "sgmatTileRows",
451
  "workgroupSize": "sgmatWorkgroupSize",
452
  "rowSubtiles": "sgmatRowSubtiles",
 
458
  {
459
  "id": "main",
460
  "shader": "matmul-nbits-q4-sgmat.wgsl.jinja",
461
+ "bindings": ["a_main", "b_main", "scales_main", "zero_points", "bias", "y_y_t"],
462
  "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" }
463
  }
464
  ]
 
467
  "id": "prefill_tiled_reg_vec4_splitk_zero_bias",
468
  "priority": 17,
469
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"],
470
+ "demoteWhen": ["variableSubgroup16To32 and attrs.bits >= 4 and tiledRegSplitK >= 4 and aRows >= 2 * tiledRegVec4TileRows"],
471
  "derive": {
 
 
 
472
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
 
 
 
 
 
 
 
 
 
 
 
 
473
  "bk": 32,
474
+ "q2ChunkLoads": "attrs.bits == 2 and attrs.K % 16 == 0 and attrs.block_size % 16 == 0 and subgroupsWave32 and kBlocksExpected >= device.adapterInfo.subgroupMinSize and dispatchN64 >= device.adapterInfo.subgroupMinSize and aRows >= 2 * tiledRegVec4TileRows",
475
+ "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide",
476
+ "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide",
477
+ "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows",
478
+ "threadCols": "2 if q2WideMicroPreferred else 4",
479
+ "tileRows": "threadRows * tileWorkgroupY",
480
+ "tileCols": "threadCols * tileWorkgroupX",
481
+ "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin",
482
  "alignedBlockLoads": true,
483
  "aVec4Loads": true,
484
  "splitK": "tiledRegSplitK",
 
492
  {
493
  "id": "partial",
494
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
495
+ "bindings": ["a_partial", "b_main", "scales_main", "zero_points", "y_f32"],
496
+ "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)", "z": "tiledRegSplitK" }
497
  },
498
  {
499
  "id": "combine",
500
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
501
+ "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "addBias": "hasBias" },
502
+ "bindings": ["partials", "bias", "y_y_t", "params_cols"],
 
 
 
 
 
503
  "dispatch": {
504
+ "x": "min(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)",
505
+ "y": "ceilDiv(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)",
506
  "z": 1
507
  }
508
  }
 
513
  "priority": 16,
514
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"],
515
  "derive": {
 
 
 
516
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
517
+ "bk": "16 if q2NarrowBkPreferred else 32",
518
+ "q2ChunkLoads": "attrs.bits == 2 and attrs.K % 16 == 0 and attrs.block_size % 16 == 0 and subgroupsWave32 and kBlocksExpected >= device.adapterInfo.subgroupMinSize and dispatchN64 >= device.adapterInfo.subgroupMinSize and aRows >= 2 * tiledRegVec4TileRows",
519
+ "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide",
520
+ "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide",
521
+ "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows",
522
+ "threadCols": "2 if q2WideMicroPreferred else 4",
523
+ "tileRows": "threadRows * tileWorkgroupY",
524
+ "tileCols": "threadCols * tileWorkgroupX",
525
+ "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin",
 
 
 
 
 
 
 
 
526
  "alignedBlockLoads": true,
527
  "aVec4Loads": true
528
  },
 
530
  {
531
  "id": "main",
532
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
533
+ "bindings": ["a_partial", "b_main", "scales_main", "zero_points", "bias", "y_y_t"],
534
+ "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)" }
535
  }
536
  ]
537
  },
 
540
  "priority": 15,
541
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"],
542
  "derive": {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
543
  "bk": "tiledRegSelectedBK",
544
  "tileRows": "tiledRegSelectedTileRows",
545
  "tileCols": 64,
 
551
  {
552
  "id": "main",
553
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
554
+ "bindings": ["a_main", "b_main", "scales_main", "zero_points", "bias", "y_y_t"],
555
  "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" }
556
  }
557
  ]
 
560
  "id": "prefill_tiled_zero_bias",
561
  "priority": 14,
562
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 64", "dispatchN32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dispatchM32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledWorkgroupFits"],
563
+ "derive": {},
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
564
  "passes": [
565
  {
566
  "id": "main",
567
  "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja",
568
+ "bindings": ["a_main", "b_main", "scales_main", "zero_points", "bias", "y_y_t"],
569
  "dispatch": { "x": "dispatchN32", "y": "dispatchM32" }
570
  }
571
  ]
 
575
  "priority": 13,
576
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"],
577
  "derive": {
 
 
 
 
 
 
 
 
 
578
  "wordsPerCol": "smallMWordsPerCol",
579
  "wordsPerBlock": "blobWords",
580
  "kLanes": "smallMKLanes",
581
+ "colGroups": "smallMColGroups"
 
 
582
  },
583
  "passes": [
584
  {
585
  "id": "main",
586
  "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja",
587
+ "bindings": ["a_main", "b_main", "scales_main", "zero_points", "bias", "y_y_t", "params_main"],
588
  "dispatch": {
589
  "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)",
590
  "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
 
597
  "id": "zero_bias",
598
  "priority": 0,
599
  "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "portableWorkgroupFits"],
600
+ "derive": {},
 
 
 
 
 
 
 
 
 
 
 
601
  "passes": [
602
  {
603
  "id": "main",
604
  "shader": "matmul-nbits.wgsl.jinja",
605
+ "bindings": ["a_main", "b_main", "scales_main", "zero_points", "bias", "y_y_t", "params_main"],
606
  "dispatch": {
607
  "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
608
  "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
 
618
  "derive": {
619
  "gemvNCols": "tunables.GEMV_N_COLS",
620
  "useSubgroups": "device.features.has(\"subgroups\")",
 
 
 
 
621
  "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")",
622
  "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"",
623
  "B_VEC_LEN": "attrs.N * gemvVecPerCol",
 
 
 
 
624
  "vecWords": "gemvVecWords",
625
  "codesPerVec": "gemvCodesPerVec",
626
  "codesPerVec4": "gemvCodesPerVec / 4",
627
  "vecPerBlock": "gemvVecPerBlock",
628
  "vecPerCol": "gemvVecPerCol",
629
+ "actVec4": "gemvActVec4"
 
 
630
  },
631
  "passes": [
632
  {
633
  "id": "main",
634
  "shader": "matmul-nbits-gemv-q4.wgsl.jinja",
635
+ "bindings": ["a_a_t", "b_b_t", "scales_main", "zero_points", "y_y_t", "params"],
636
+ "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 }
 
637
  }
638
  ]
639
  },
 
647
  "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
648
  },
649
  "derive": {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
650
  "tileRows": "sgmatTileRows",
651
  "workgroupSize": "sgmatWorkgroupSize",
652
  "rowSubtiles": "sgmatRowSubtiles",
 
658
  {
659
  "id": "main",
660
  "shader": "matmul-nbits-q4-sgmat.wgsl.jinja",
661
+ "bindings": ["a_main", "b_main", "scales_main", "zero_points", "y_y_t"],
662
  "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" }
663
  }
664
  ]
 
667
  "id": "prefill_tiled_reg_vec4_splitk_zero_only",
668
  "priority": 17,
669
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"],
670
+ "demoteWhen": ["variableSubgroup16To32 and attrs.bits >= 4 and tiledRegSplitK >= 4 and aRows >= 2 * tiledRegVec4TileRows"],
671
  "derive": {
 
 
 
672
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
 
 
 
 
 
 
 
 
 
 
 
 
673
  "bk": 32,
674
+ "q2ChunkLoads": "attrs.bits == 2 and attrs.K % 16 == 0 and attrs.block_size % 16 == 0 and subgroupsWave32 and kBlocksExpected >= device.adapterInfo.subgroupMinSize and dispatchN64 >= device.adapterInfo.subgroupMinSize and aRows >= 2 * tiledRegVec4TileRows",
675
+ "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide",
676
+ "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide",
677
+ "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows",
678
+ "threadCols": "2 if q2WideMicroPreferred else 4",
679
+ "tileRows": "threadRows * tileWorkgroupY",
680
+ "tileCols": "threadCols * tileWorkgroupX",
681
+ "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin",
682
  "alignedBlockLoads": true,
683
  "aVec4Loads": true,
684
  "splitK": "tiledRegSplitK",
 
692
  {
693
  "id": "partial",
694
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
695
+ "bindings": ["a_partial", "b_main", "scales_main", "zero_points", "y_f32"],
696
+ "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)", "z": "tiledRegSplitK" }
697
  },
698
  {
699
  "id": "combine",
700
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
701
+ "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "addBias": "hasBias" },
702
+ "bindings": ["partials", "y_y_t", "params_cols"],
 
 
 
 
 
703
  "dispatch": {
704
+ "x": "min(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)",
705
+ "y": "ceilDiv(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)",
706
  "z": 1
707
  }
708
  }
 
713
  "priority": 16,
714
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"],
715
  "derive": {
 
 
 
716
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
717
+ "bk": "16 if q2NarrowBkPreferred else 32",
718
+ "q2ChunkLoads": "attrs.bits == 2 and attrs.K % 16 == 0 and attrs.block_size % 16 == 0 and subgroupsWave32 and kBlocksExpected >= device.adapterInfo.subgroupMinSize and dispatchN64 >= device.adapterInfo.subgroupMinSize and aRows >= 2 * tiledRegVec4TileRows",
719
+ "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide",
720
+ "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide",
721
+ "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows",
722
+ "threadCols": "2 if q2WideMicroPreferred else 4",
723
+ "tileRows": "threadRows * tileWorkgroupY",
724
+ "tileCols": "threadCols * tileWorkgroupX",
725
+ "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin",
 
 
 
 
 
 
 
 
726
  "alignedBlockLoads": true,
727
  "aVec4Loads": true
728
  },
 
730
  {
731
  "id": "main",
732
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
733
+ "bindings": ["a_partial", "b_main", "scales_main", "zero_points", "y_y_t"],
734
+ "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)" }
735
  }
736
  ]
737
  },
 
740
  "priority": 15,
741
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"],
742
  "derive": {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
743
  "bk": "tiledRegSelectedBK",
744
  "tileRows": "tiledRegSelectedTileRows",
745
  "tileCols": 64,
 
751
  {
752
  "id": "main",
753
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
754
+ "bindings": ["a_main", "b_main", "scales_main", "zero_points", "y_y_t"],
755
  "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" }
756
  }
757
  ]
 
760
  "id": "prefill_tiled_zero_only",
761
  "priority": 14,
762
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 64", "dispatchN32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dispatchM32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledWorkgroupFits"],
763
+ "derive": {},
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
764
  "passes": [
765
  {
766
  "id": "main",
767
  "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja",
768
+ "bindings": ["a_main", "b_main", "scales_main", "zero_points", "y_y_t"],
769
  "dispatch": { "x": "dispatchN32", "y": "dispatchM32" }
770
  }
771
  ]
 
775
  "priority": 13,
776
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"],
777
  "derive": {
 
 
 
 
 
 
 
 
 
778
  "wordsPerCol": "smallMWordsPerCol",
779
  "wordsPerBlock": "blobWords",
780
  "kLanes": "smallMKLanes",
781
+ "colGroups": "smallMColGroups"
 
 
782
  },
783
  "passes": [
784
  {
785
  "id": "main",
786
  "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja",
787
+ "bindings": ["a_main", "b_main", "scales_main", "zero_points", "y_y_t", "params_main"],
788
  "dispatch": {
789
  "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)",
790
  "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
 
797
  "id": "zero_only",
798
  "priority": 0,
799
  "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "portableWorkgroupFits"],
800
+ "derive": {},
 
 
 
 
 
 
 
 
 
 
 
801
  "passes": [
802
  {
803
  "id": "main",
804
  "shader": "matmul-nbits.wgsl.jinja",
805
+ "bindings": ["a_main", "b_main", "scales_main", "zero_points", "y_y_t", "params_main"],
806
  "dispatch": {
807
  "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
808
  "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
 
818
  "derive": {
819
  "gemvNCols": "tunables.GEMV_N_COLS",
820
  "useSubgroups": "device.features.has(\"subgroups\")",
 
 
 
 
821
  "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")",
822
  "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"",
823
  "B_VEC_LEN": "attrs.N * gemvVecPerCol",
 
 
 
 
824
  "vecWords": "gemvVecWords",
825
  "codesPerVec": "gemvCodesPerVec",
826
  "codesPerVec4": "gemvCodesPerVec / 4",
827
  "vecPerBlock": "gemvVecPerBlock",
828
  "vecPerCol": "gemvVecPerCol",
829
+ "actVec4": "gemvActVec4"
 
 
830
  },
831
  "passes": [
832
  {
833
  "id": "main",
834
  "shader": "matmul-nbits-gemv-q4.wgsl.jinja",
835
+ "bindings": ["a_a_t", "b_b_t", "scales_main", "bias", "y_y_t", "params"],
836
+ "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 }
 
837
  }
838
  ]
839
  },
 
847
  "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }]
848
  },
849
  "derive": {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
850
  "tileRows": "sgmatTileRows",
851
  "workgroupSize": "sgmatWorkgroupSize",
852
  "rowSubtiles": "sgmatRowSubtiles",
 
858
  {
859
  "id": "main",
860
  "shader": "matmul-nbits-q4-sgmat.wgsl.jinja",
861
+ "bindings": ["a_main", "b_main", "scales_main", "bias", "y_y_t"],
862
  "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" }
863
  }
864
  ]
 
867
  "id": "prefill_tiled_reg_vec4_splitk_bias_only",
868
  "priority": 17,
869
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"],
870
+ "demoteWhen": ["variableSubgroup16To32 and attrs.bits >= 4 and tiledRegSplitK >= 4 and aRows >= 2 * tiledRegVec4TileRows"],
871
  "derive": {
 
 
 
872
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
 
 
 
 
 
 
 
 
 
 
 
 
873
  "bk": 32,
874
+ "q2ChunkLoads": "attrs.bits == 2 and attrs.K % 16 == 0 and attrs.block_size % 16 == 0 and subgroupsWave32 and kBlocksExpected >= device.adapterInfo.subgroupMinSize and dispatchN64 >= device.adapterInfo.subgroupMinSize and aRows >= 2 * tiledRegVec4TileRows",
875
+ "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide",
876
+ "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide",
877
+ "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows",
878
+ "threadCols": "2 if q2WideMicroPreferred else 4",
879
+ "tileRows": "threadRows * tileWorkgroupY",
880
+ "tileCols": "threadCols * tileWorkgroupX",
881
+ "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin",
882
  "alignedBlockLoads": true,
883
  "aVec4Loads": true,
884
  "splitK": "tiledRegSplitK",
 
892
  {
893
  "id": "partial",
894
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
895
+ "bindings": ["a_partial", "b_main", "scales_main", "y_f32"],
896
+ "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)", "z": "tiledRegSplitK" }
897
  },
898
  {
899
  "id": "combine",
900
  "shader": "reduce-axis0-splitk-combine.wgsl.jinja",
901
+ "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "addBias": "hasBias" },
902
+ "bindings": ["partials", "bias", "y_y_t", "params_cols"],
 
 
 
 
 
903
  "dispatch": {
904
+ "x": "min(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)",
905
+ "y": "ceilDiv(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)",
906
  "z": 1
907
  }
908
  }
 
913
  "priority": 16,
914
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"],
915
  "derive": {
 
 
 
916
  "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"",
917
+ "bk": "16 if q2NarrowBkPreferred else 32",
918
+ "q2ChunkLoads": "attrs.bits == 2 and attrs.K % 16 == 0 and attrs.block_size % 16 == 0 and subgroupsWave32 and kBlocksExpected >= device.adapterInfo.subgroupMinSize and dispatchN64 >= device.adapterInfo.subgroupMinSize and aRows >= 2 * tiledRegVec4TileRows",
919
+ "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide",
920
+ "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide",
921
+ "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows",
922
+ "threadCols": "2 if q2WideMicroPreferred else 4",
923
+ "tileRows": "threadRows * tileWorkgroupY",
924
+ "tileCols": "threadCols * tileWorkgroupX",
925
+ "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin",
 
 
 
 
 
 
 
 
926
  "alignedBlockLoads": true,
927
  "aVec4Loads": true
928
  },
 
930
  {
931
  "id": "main",
932
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
933
+ "bindings": ["a_partial", "b_main", "scales_main", "bias", "y_y_t"],
934
+ "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)" }
935
  }
936
  ]
937
  },
 
940
  "priority": 15,
941
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"],
942
  "derive": {
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
943
  "bk": "tiledRegSelectedBK",
944
  "tileRows": "tiledRegSelectedTileRows",
945
  "tileCols": 64,
 
951
  {
952
  "id": "main",
953
  "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja",
954
+ "bindings": ["a_main", "b_main", "scales_main", "bias", "y_y_t"],
955
  "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" }
956
  }
957
  ]
 
960
  "id": "prefill_tiled_bias_only",
961
  "priority": 14,
962
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 64", "dispatchN32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "dispatchM32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledWorkgroupFits"],
963
+ "derive": {},
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
964
  "passes": [
965
  {
966
  "id": "main",
967
  "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja",
968
+ "bindings": ["a_main", "b_main", "scales_main", "bias", "y_y_t"],
969
  "dispatch": { "x": "dispatchN32", "y": "dispatchM32" }
970
  }
971
  ]
 
975
  "priority": 13,
976
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"],
977
  "derive": {
 
 
 
 
 
 
 
 
 
978
  "wordsPerCol": "smallMWordsPerCol",
979
  "wordsPerBlock": "blobWords",
980
  "kLanes": "smallMKLanes",
981
+ "colGroups": "smallMColGroups"
 
 
982
  },
983
  "passes": [
984
  {
985
  "id": "main",
986
  "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja",
987
+ "bindings": ["a_main", "b_main", "scales_main", "bias", "y_y_t", "params_main"],
988
  "dispatch": {
989
  "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)",
990
  "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
 
997
  "id": "bias_only",
998
  "priority": 0,
999
  "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "portableWorkgroupFits"],
1000
+ "derive": {},
 
 
 
 
 
 
 
 
 
 
 
1001
  "passes": [
1002
  {
1003
  "id": "main",
1004
  "shader": "matmul-nbits.wgsl.jinja",
1005
+ "bindings": ["a_main", "b_main", "scales_main", "bias", "y_y_t", "params_main"],
1006
  "dispatch": {
1007
  "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
1008
  "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)",
build/webgpu/matmul-nbits-dp4a-quantize.wgsl.jinja CHANGED
@@ -6,7 +6,7 @@
6
  // multiple of 128, so blocks never straddle rows and the flat layout is exact.
7
  const VEC4_COUNT: u32 = {{ vec4Count }}u;
8
  const BLOCK_COUNT: u32 = {{ blockCount }}u;
9
- const WG: u32 = 64u;
10
 
11
  var<workgroup> maxAbs: array<f32, WG>;
12
 
 
6
  // multiple of 128, so blocks never straddle rows and the flat layout is exact.
7
  const VEC4_COUNT: u32 = {{ vec4Count }}u;
8
  const BLOCK_COUNT: u32 = {{ blockCount }}u;
9
+ const WG: u32 = {{ dp4aQuantizeWorkgroupSize }}u;
10
 
11
  var<workgroup> maxAbs: array<f32, WG>;
12
 
build/webgpu/matmul-nbits-gemv-q4.wgsl.jinja CHANGED
@@ -46,15 +46,6 @@ fn zero_point({% if hasZero %}n: u32, block: u32{% endif %}) -> f32 {
46
  {% endif %}
47
  }
48
 
49
- {% macro vec_dot(words) %}
50
- {% for w in range(vecWords) %}
51
- {% set word = (words ~ "." ~ comps[w]) if vecWords == 4 else words %}
52
- {% for h in range(codesPerWord) %}
53
- dot = dot + a{{ w * codesPerWord + h }} * f32(({{ word }} >> {{ h * bits }}u) & {{ codeMask }}u);
54
- {% endfor %}
55
- {% endfor %}
56
- {%- endmacro %}
57
-
58
  @compute @workgroup_size({{ workgroupSize }}, 1, 1)
59
  fn main(
60
  @builtin(workgroup_id) wid: vec3<u32>,
@@ -108,7 +99,12 @@ fn main(
108
  let scale = f32(scales[n * params.kBlocks + block]);
109
  let zero = zero_point({% if hasZero %}n, block{% endif %});
110
  var dot = 0.0;
111
- {{ vec_dot("words") }}
 
 
 
 
 
112
  acc{{ sfx }}.{{ comps[c] }} = acc{{ sfx }}.{{ comps[c] }} + (dot - zero * asum) * scale;
113
  }
114
  {% endfor %}
 
46
  {% endif %}
47
  }
48
 
 
 
 
 
 
 
 
 
 
49
  @compute @workgroup_size({{ workgroupSize }}, 1, 1)
50
  fn main(
51
  @builtin(workgroup_id) wid: vec3<u32>,
 
99
  let scale = f32(scales[n * params.kBlocks + block]);
100
  let zero = zero_point({% if hasZero %}n, block{% endif %});
101
  var dot = 0.0;
102
+ {% for w in range(vecWords) %}
103
+ {% set word = ("words." ~ comps[w]) if vecWords == 4 else "words" %}
104
+ {% for h in range(codesPerWord) %}
105
+ dot = dot + a{{ w * codesPerWord + h }} * f32(({{ word }} >> {{ h * bits }}u) & {{ codeMask }}u);
106
+ {% endfor %}
107
+ {% endfor %}
108
  acc{{ sfx }}.{{ comps[c] }} = acc{{ sfx }}.{{ comps[c] }} + (dot - zero * asum) * scale;
109
  }
110
  {% endfor %}
build/webgpu/matmul-nbits-q4-dp4a-prefill.wgsl.jinja CHANGED
@@ -5,7 +5,6 @@ fn dot4_packed(a_word: u32, b_word: u32) -> i32 {
5
  return dot4I8Packed(a_word, b_word);
6
  }
7
 
8
-
9
  // com.microsoft.MatMulNBits q4 prefill with int8-quantized activations
10
  // (accuracy_level 4). A arrives pre-quantized as packed int8 words with one
11
  // scale per 128-element block; each weight nibble is rebiased by -8 and packed
@@ -27,14 +26,14 @@ var<workgroup> tB: array<array<u32, 8u>, 64u>;
27
  var<workgroup> tAscale: array<f32, 64u>;
28
  var<workgroup> tBscale: array<f32, 64u>;
29
 
30
- @compute @workgroup_size(16, 16, 1)
31
  fn main(
32
  @builtin(workgroup_id) wg: vec3<u32>,
33
  @builtin(local_invocation_id) lid: vec3<u32>
34
  ) {
35
  let mBase = wg.y * 64u;
36
  let nBase = wg.x * 64u;
37
- let li = lid.y * 16u + lid.x;
38
 
39
  var acc: array<f32, 16u>;
40
  for (var t = 0u; t < 16u; t = t + 1u) { acc[t] = 0.0; }
@@ -46,7 +45,7 @@ fn main(
46
  // Stage 64 rows x 8 packed A words and 64 cols x 8 packed B words; each of
47
  // the 256 threads loads two of each.
48
  for (var e = 0u; e < 2u; e = e + 1u) {
49
- let idx = li + e * 256u;
50
  let r = idx / 8u;
51
  let w = idx % 8u;
52
  let am = mBase + r;
 
5
  return dot4I8Packed(a_word, b_word);
6
  }
7
 
 
8
  // com.microsoft.MatMulNBits q4 prefill with int8-quantized activations
9
  // (accuracy_level 4). A arrives pre-quantized as packed int8 words with one
10
  // scale per 128-element block; each weight nibble is rebiased by -8 and packed
 
26
  var<workgroup> tAscale: array<f32, 64u>;
27
  var<workgroup> tBscale: array<f32, 64u>;
28
 
29
+ @compute @workgroup_size({{ tiledWorkgroupSide }}, {{ tiledWorkgroupSide }}, 1)
30
  fn main(
31
  @builtin(workgroup_id) wg: vec3<u32>,
32
  @builtin(local_invocation_id) lid: vec3<u32>
33
  ) {
34
  let mBase = wg.y * 64u;
35
  let nBase = wg.x * 64u;
36
+ let li = lid.y * {{ tiledWorkgroupSide }}u + lid.x;
37
 
38
  var acc: array<f32, 16u>;
39
  for (var t = 0u; t < 16u; t = t + 1u) { acc[t] = 0.0; }
 
45
  // Stage 64 rows x 8 packed A words and 64 cols x 8 packed B words; each of
46
  // the 256 threads loads two of each.
47
  for (var e = 0u; e < 2u; e = e + 1u) {
48
+ let idx = li + e * {{ tiledWorkgroupSide * tiledWorkgroupSide }}u;
49
  let r = idx / 8u;
50
  let w = idx % 8u;
51
  let am = mBase + r;
build/webgpu/matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja CHANGED
@@ -1,3 +1,7 @@
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  // com.microsoft.MatMulNBits q4/q8 prefill, no-subgroup-matrix tier — register-blocked.
@@ -6,8 +10,9 @@
6
  // blocking. Both tiles are indexed by their own output axis and group four K
7
  // values per vector word, so the micro-tile accumulates through dot() and the
8
  // column-owning loader writes whole words instead of a BN-strided column.
9
- // The full geometry computes a 4x4 micro-tile over a 64x64 output tile. The
10
- // portable geometry computes 2x4 over 32x64 to bound the per-lane accumulator
 
11
  // set. K_TILE specializes the K tile. For standard 32/64-element quant blocks,
12
  // one lane owns one output column and the full BK slice: scale and zero are
13
  // loaded once, and each stored byte is read once for the K-adjacent codes it
@@ -24,8 +29,10 @@ const BM: u32 = {{ tileRows }}u;
24
  const BN: u32 = {{ tileCols }}u;
25
  const TM: u32 = {{ threadRows }}u;
26
  const TN: u32 = {{ threadCols }}u;
27
- const WG_X: u32 = 16u;
28
- const WG_Y: u32 = 16u;
 
 
29
  const WG_THREADS: u32 = WG_X * WG_Y;
30
  {% set splitKValue = splitK if splitK is defined else 1 %}
31
  {% set tilesPerSplitValue = tilesPerSplit if tilesPerSplit is defined else 0 %}
@@ -61,14 +68,12 @@ fn {{ fn }}(n: u32, block: u32, offset: u32) -> u32 {
61
  let byte_index = (n * {{ kBlocks }} + block) * {{ blobSize }} + offset;
62
  return ({{ buffer }}[byte_index >> 2u] >> ((byte_index & 3u) * 8u)) & 255u;
63
  {% endif %}
64
- }
65
- {%- endmacro %}
66
- {{- matmul_nbits_packed_code(bits=bits, kBlocks="KBLOCKS", blobSize="BLOB_SIZE") }}
67
-
68
  {% endif %}
69
  {% macro zero_of(blockExpr) %}{% if hasZero %}f32(zero_points[bn * KBLOCKS + {{ blockExpr }}]){% else %}{{ defaultZero }}{% endif %}{% endmacro %}
70
 
71
- @compute @workgroup_size(16, 16, 1)
72
  fn main(
73
  @builtin(workgroup_id) wg: vec3<u32>,
74
  @builtin(local_invocation_id) lid: vec3<u32>
@@ -115,6 +120,36 @@ fn main(
115
  tileA[ar][ac4] = aWord;
116
  }
117
  {% if alignedBlockLoads %}
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
118
  // Every lane materializes one vector word of one output column. BLOCK_SIZE is
119
  // a multiple of BK, so the whole tile slice of a column shares one scale and
120
  // zero point. The blob is packed four bytes per u32 word and the four
@@ -162,6 +197,7 @@ fn main(
162
  }
163
  tileB[bc][kv] = word;
164
  }
 
165
  {% else %}
166
  // Quant blocks that do not contain a whole BK tile use the element-wise loader.
167
  for (var idx: u32 = li; idx < BN * K_VECS; idx = idx + WG_THREADS) {
 
1
+ {% set subgroupPin = tileSubgroupPin if tileSubgroupPin is defined else 0 %}
2
+ {% if subgroupPin %}
3
+ enable subgroup_size_control;
4
+ {% endif %}
5
  {{ env.wgsl.resourceDeclarations }}
6
 
7
  // com.microsoft.MatMulNBits q4/q8 prefill, no-subgroup-matrix tier — register-blocked.
 
10
  // blocking. Both tiles are indexed by their own output axis and group four K
11
  // values per vector word, so the micro-tile accumulates through dot() and the
12
  // column-owning loader writes whole words instead of a BN-strided column.
13
+ // The full geometry computes a 4x4 micro-tile over a 64x64 output tile. A
14
+ // device-selected 32x8 workgroup can use an 8x2 micro-tile on that same output
15
+ // tile. The portable geometry computes 2x4 over 32x64 to bound the accumulator
16
  // set. K_TILE specializes the K tile. For standard 32/64-element quant blocks,
17
  // one lane owns one output column and the full BK slice: scale and zero are
18
  // loaded once, and each stored byte is read once for the K-adjacent codes it
 
29
  const BN: u32 = {{ tileCols }}u;
30
  const TM: u32 = {{ threadRows }}u;
31
  const TN: u32 = {{ threadCols }}u;
32
+ {% set wgX = tileWorkgroupX if tileWorkgroupX is defined else tiledWorkgroupSide %}
33
+ {% set wgY = tileWorkgroupY if tileWorkgroupY is defined else tiledWorkgroupSide %}
34
+ const WG_X: u32 = {{ wgX }}u;
35
+ const WG_Y: u32 = {{ wgY }}u;
36
  const WG_THREADS: u32 = WG_X * WG_Y;
37
  {% set splitKValue = splitK if splitK is defined else 1 %}
38
  {% set tilesPerSplitValue = tilesPerSplit if tilesPerSplit is defined else 0 %}
 
68
  let byte_index = (n * {{ kBlocks }} + block) * {{ blobSize }} + offset;
69
  return ({{ buffer }}[byte_index >> 2u] >> ((byte_index & 3u) * 8u)) & 255u;
70
  {% endif %}
71
+ }{% endmacro %}
72
+ {{ matmul_nbits_packed_code(bits=bits, kBlocks="KBLOCKS", blobSize="BLOB_SIZE") }}
 
 
73
  {% endif %}
74
  {% macro zero_of(blockExpr) %}{% if hasZero %}f32(zero_points[bn * KBLOCKS + {{ blockExpr }}]){% else %}{{ defaultZero }}{% endif %}{% endmacro %}
75
 
76
+ @compute @workgroup_size({{ wgX }}, {{ wgY }}, 1){{ (" @subgroup_size(" ~ subgroupPin ~ ")") if subgroupPin else "" }}
77
  fn main(
78
  @builtin(workgroup_id) wg: vec3<u32>,
79
  @builtin(local_invocation_id) lid: vec3<u32>
 
120
  tileA[ar][ac4] = aWord;
121
  }
122
  {% if alignedBlockLoads %}
123
+ {% if q2ChunkLoads is defined and q2ChunkLoads and bits == 2 and bk % 16 == 0 %}
124
+ // A 32-bit word carries sixteen adjacent 2-bit codes. Decode it once per
125
+ // column and quant chunk, reusing the scale and zero for four staged vec4s.
126
+ for (var idx: u32 = li; idx < BN * (BK / 16u); idx = idx + WG_THREADS) {
127
+ let chunksPerCol = BK / 16u;
128
+ let bc = idx / chunksPerCol;
129
+ let chunk = idx % chunksPerCol;
130
+ let bn = nBase + bc;
131
+ let baseK = kBase + chunk * 16u;
132
+ var packed: u32 = 0u;
133
+ var scale: f32 = 0.0;
134
+ var zero: f32 = 0.0;
135
+ if (bn < N && baseK < K) {
136
+ let block = baseK / BLOCK_SIZE;
137
+ let offset = baseK % BLOCK_SIZE;
138
+ let byteIndex = (bn * KBLOCKS + block) * BLOB_SIZE + (offset >> 2u);
139
+ packed = b[byteIndex >> 2u];
140
+ scale = f32(scales[bn * KBLOCKS + block]);
141
+ zero = {{ zero_of("block") }};
142
+ }
143
+ {% for vec in range(4) %}
144
+ let word{{ vec }} = vec4<f32>(
145
+ {% for component in range(4) %}
146
+ (f32((packed >> {{ 2 * (vec * 4 + component) }}u) & 3u) - zero) * scale{% if not loop.last %},{% endif %}
147
+ {% endfor %}
148
+ );
149
+ tileB[bc][chunk * 4u + {{ vec }}u] = word{{ vec }};
150
+ {% endfor %}
151
+ }
152
+ {% else %}
153
  // Every lane materializes one vector word of one output column. BLOCK_SIZE is
154
  // a multiple of BK, so the whole tile slice of a column shares one scale and
155
  // zero point. The blob is packed four bytes per u32 word and the four
 
197
  }
198
  tileB[bc][kv] = word;
199
  }
200
+ {% endif %}
201
  {% else %}
202
  // Quant blocks that do not contain a whole BK tile use the element-wise loader.
203
  for (var idx: u32 = li; idx < BN * K_VECS; idx = idx + WG_THREADS) {
build/webgpu/matmul-nbits-q4-prefill-tiled.wgsl.jinja CHANGED
@@ -42,18 +42,17 @@ fn {{ fn }}(n: u32, block: u32, offset: u32) -> u32 {
42
  let byte_index = (n * {{ kBlocks }} + block) * {{ blobSize }} + offset;
43
  return ({{ buffer }}[byte_index >> 2u] >> ((byte_index & 3u) * 8u)) & 255u;
44
  {% endif %}
45
- }
46
- {%- endmacro %}
47
- {{- matmul_nbits_packed_code(bits=bits, kBlocks="KBLOCKS", blobSize="BLOB_SIZE") }}
48
 
49
- @compute @workgroup_size(16, 16, 1)
50
  fn main(
51
  @builtin(workgroup_id) wg: vec3<u32>,
52
  @builtin(local_invocation_id) lid: vec3<u32>
53
  ) {
54
  let mBase = wg.y * BM;
55
  let nBase = wg.x * BN;
56
- let li = lid.y * 16u + lid.x;
57
 
58
  var acc00: f32 = 0.0;
59
  var acc01: f32 = 0.0;
@@ -65,7 +64,7 @@ fn main(
65
  // Cooperative load: 32x16 A tile + 16x32 B tile, 256 threads x 2 each. The B
66
  // tile is dequantized from the packed q blob during the load.
67
  for (var e: u32 = 0u; e < 2u; e = e + 1u) {
68
- let idx = li + e * 256u;
69
  let ar = idx / BK;
70
  let ac = idx % BK;
71
  let am = mBase + ar;
 
42
  let byte_index = (n * {{ kBlocks }} + block) * {{ blobSize }} + offset;
43
  return ({{ buffer }}[byte_index >> 2u] >> ((byte_index & 3u) * 8u)) & 255u;
44
  {% endif %}
45
+ }{% endmacro %}
46
+ {{ matmul_nbits_packed_code(bits=bits, kBlocks="KBLOCKS", blobSize="BLOB_SIZE") }}
 
47
 
48
+ @compute @workgroup_size({{ tiledWorkgroupSide }}, {{ tiledWorkgroupSide }}, 1)
49
  fn main(
50
  @builtin(workgroup_id) wg: vec3<u32>,
51
  @builtin(local_invocation_id) lid: vec3<u32>
52
  ) {
53
  let mBase = wg.y * BM;
54
  let nBase = wg.x * BN;
55
+ let li = lid.y * {{ tiledWorkgroupSide }}u + lid.x;
56
 
57
  var acc00: f32 = 0.0;
58
  var acc01: f32 = 0.0;
 
64
  // Cooperative load: 32x16 A tile + 16x32 B tile, 256 threads x 2 each. The B
65
  // tile is dequantized from the packed q blob during the load.
66
  for (var e: u32 = 0u; e < 2u; e = e + 1u) {
67
+ let idx = li + e * {{ tiledWorkgroupSide * tiledWorkgroupSide }}u;
68
  let ar = idx / BK;
69
  let ac = idx % BK;
70
  let am = mBase + ar;
build/webgpu/matmul-nbits-q4-sgmat.wgsl.jinja CHANGED
@@ -13,7 +13,6 @@ enable subgroup_size_control;
13
  enable chromium_experimental_subgroup_matrix;
14
  diagnostic(off, chromium.subgroup_matrix_uniformity);
15
 
16
-
17
  {{ env.wgsl.resourceDeclarations }}
18
 
19
  const M: u32 = {{ M }}u;
@@ -150,7 +149,7 @@ fn main(
150
  workgroupBarrier();
151
 
152
  for (var step = 0u; step < TILE_K; step = step + 8u) {
153
- {% set operandScalar = "f32" %}{% set directInputs = directMatrixInputs is defined and directMatrixInputs %}
154
  let matrix_a_offset = subtile_idy * SUB_ROWS * TILE_K + step;
155
  {% for r in range(2) %}
156
  var matA{{ r }}: subgroup_matrix_left<{{ operandScalar }}, 8, 8> = subgroupMatrixLoad<subgroup_matrix_left<{{ operandScalar }}, 8, 8>, row_major>(&{{ "w" if directInputs else "tile_A" }}, matrix_a_offset{% if r > 0 %} + 8u * {{ "K" if directInputs else "TILE_K" }}{% endif %}, {{ "K" if directInputs else "TILE_K" }});
@@ -169,6 +168,7 @@ fn main(
169
  matC11 = subgroupMatrixMultiplyAccumulate(matA1, matB1, matC11);
170
  matC12 = subgroupMatrixMultiplyAccumulate(matA1, matB2, matC12);
171
  matC13 = subgroupMatrixMultiplyAccumulate(matA1, matB3, matC13);
 
172
  }
173
  workgroupBarrier();
174
  }
 
13
  enable chromium_experimental_subgroup_matrix;
14
  diagnostic(off, chromium.subgroup_matrix_uniformity);
15
 
 
16
  {{ env.wgsl.resourceDeclarations }}
17
 
18
  const M: u32 = {{ M }}u;
 
149
  workgroupBarrier();
150
 
151
  for (var step = 0u; step < TILE_K; step = step + 8u) {
152
+ {% set operandScalar = "f32" %}{% set directInputs = false %}
153
  let matrix_a_offset = subtile_idy * SUB_ROWS * TILE_K + step;
154
  {% for r in range(2) %}
155
  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" }});
 
168
  matC11 = subgroupMatrixMultiplyAccumulate(matA1, matB1, matC11);
169
  matC12 = subgroupMatrixMultiplyAccumulate(matA1, matB2, matC12);
170
  matC13 = subgroupMatrixMultiplyAccumulate(matA1, matB3, matC13);
171
+
172
  }
173
  workgroupBarrier();
174
  }
build/webgpu/matmul-nbits.wgsl.jinja CHANGED
@@ -1,3 +1,8 @@
 
 
 
 
 
1
  {{ env.wgsl.resourceDeclarations }}
2
 
3
  const WG: u32 = {{ workgroupSize }}u;
@@ -20,16 +25,12 @@ fn {{ fn }}(n: u32, block: u32, offset: u32) -> u32 {
20
  let byte_index = (n * {{ kBlocks }} + block) * {{ blobSize }} + offset;
21
  return ({{ buffer }}[byte_index >> 2u] >> ((byte_index & 3u) * 8u)) & 255u;
22
  {% endif %}
23
- }
24
- {%- endmacro %}
25
- {{- matmul_nbits_packed_code(bits=bits) }}
26
 
27
  @compute @workgroup_size(WG, 1, 1)
28
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
29
- // 2D-folded flat output-element index: gid.y carries the high bits past the
30
- // per-axis dispatch fold width. With no fold this
31
- // reduces to gid.x; the index >= total guard drops the tail.
32
- let index = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG;
33
  let total = params.rows * params.N;
34
 
35
  if (index >= total) {
 
1
+ {% macro flat_index_2d(workgroupSize, name="i", bound="params.count", guardInline=false) %}
2
+ {% set wgTerm = workgroupSize ~ "u" if workgroupSize is number else workgroupSize %}
3
+ // 2D-folded flat index: gid.y carries the high bits past the dispatch's
4
+ // per-axis workgroup fold width.
5
+ let {{ name }} = gid.x + gid.y * {{ DISPATCH_FOLD_WIDTH }}u * {{ wgTerm }};{% endmacro %}
6
  {{ env.wgsl.resourceDeclarations }}
7
 
8
  const WG: u32 = {{ workgroupSize }}u;
 
25
  let byte_index = (n * {{ kBlocks }} + block) * {{ blobSize }} + offset;
26
  return ({{ buffer }}[byte_index >> 2u] >> ((byte_index & 3u) * 8u)) & 255u;
27
  {% endif %}
28
+ }{% endmacro %}
29
+ {{ matmul_nbits_packed_code(bits=bits) }}
 
30
 
31
  @compute @workgroup_size(WG, 1, 1)
32
  fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
33
+ {{ flat_index_2d("WG", "index", "") }}
 
 
 
34
  let total = params.rows * params.N;
35
 
36
  if (index >= total) {
build/webgpu/metadata.json CHANGED
@@ -1,29 +1,29 @@
1
  {
2
  "name": "com.microsoft.MatMulNBits",
3
- "id": "_com_microsoft_matmulnbits_webgpu_6f18c00",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
- "bench.json": "P8wNyT1cbaSLn1VS4hVRFQxLYwji4Z7gp69Mt2pQaFo=",
11
- "manifest.json": "QPPdzEGo1lJtp16tQaaqA9UHKPlvBY3fq66gG4XPvyQ=",
12
- "matmul-nbits-dp4a-quantize.wgsl.jinja": "0gqEvgBRzH2Demz/RvrFkUOV487GyCq7Ujd26GyD+xw=",
13
- "matmul-nbits-gemv-q4.wgsl.jinja": "ewlLPcW7t3UdnoLymZBoXnV+1oOeB9+YTyaOAnPhhfA=",
14
- "matmul-nbits-q4-dp4a-prefill.wgsl.jinja": "pd9dWWUdIgYFVCRayU5OlvjYIvyqPvkOGgwtEPNMHas=",
15
  "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja": "8DVy3szxcxwVItIlEXoioQi2YxDmmd5TCFW4BCYPiCQ=",
16
- "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja": "W6JEeRleux7Jm3wmaflEusYFNuCiFtXnBTEW5cfS4Fk=",
17
- "matmul-nbits-q4-prefill-tiled.wgsl.jinja": "eAEIbpkW0qziqyXhYNRiVbCuDRghjoz8EjnOrP8AgD8=",
18
- "matmul-nbits-q4-sgmat.wgsl.jinja": "SFEMothi+irkTIclMjeTF1sCmB61Ai/2hGgnoOyNMUM=",
19
- "matmul-nbits.wgsl.jinja": "UndxgqiV/O19Plpxda1d388lnQWyeCugbBvJGcYr9t8=",
20
- "reduce-axis0-splitk-combine.wgsl.jinja": "Yz1hjK55R/kndUrw3ugPqgPaOBwKmYupqZoVdJTHO5Q=",
21
- "test.json": "Hnb4z+ExdwYZb4EjSBR0gdGi5eY1kh7DaoGj3UTC5Ug="
22
  }
23
  },
24
- "provenance": { "kernel": { "sha": "91d990483a174128daf7673f3f37a7c890493ae1", "dirty": false } },
25
  "webgpu": {
26
- "manifestSpec": "2.0",
27
  "variants": {
28
  "q4_dp4a_prefill": ["matmul-nbits-dp4a-quantize.wgsl.jinja", "matmul-nbits-q4-dp4a-prefill.wgsl.jinja"],
29
  "gemv_default_zero": ["matmul-nbits-gemv-q4.wgsl.jinja"],
 
1
  {
2
  "name": "com.microsoft.MatMulNBits",
3
+ "id": "_com_microsoft_matmulnbits_webgpu_544d369",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "backend": { "type": "webgpu" },
7
  "digest": {
8
  "algorithm": "sha256",
9
  "files": {
10
+ "bench.json": "jsZTXUF4MjY52ewIINIKTdgXYZvvFfi+Ch6E3mCRIbo=",
11
+ "manifest.json": "C1+nZAAnjkvjvNZMOCbeF0zCAyPt3GIyqvbRUnz4KJo=",
12
+ "matmul-nbits-dp4a-quantize.wgsl.jinja": "7Ywrm3pRk4HfGeL5D0YcObG/kKLAgpr4dfu9vUJw8SM=",
13
+ "matmul-nbits-gemv-q4.wgsl.jinja": "0NkQ1G4eYRQGokjjMW7XKRvU76j8Dno9mrCBoXzeOc8=",
14
+ "matmul-nbits-q4-dp4a-prefill.wgsl.jinja": "w6qZuxXnJFteFjjDq+HonDn70RMHtPHreWxRtDiPIdU=",
15
  "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja": "8DVy3szxcxwVItIlEXoioQi2YxDmmd5TCFW4BCYPiCQ=",
16
+ "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja": "Z22EQeNnROI5Vpy6mNgPK733TkxWXZB9rlp7qKV6fps=",
17
+ "matmul-nbits-q4-prefill-tiled.wgsl.jinja": "RzrALqTo5135awvgG96cJ029Hl8gsQtwujplaGYzcUs=",
18
+ "matmul-nbits-q4-sgmat.wgsl.jinja": "tUpNHvX4uEflOZ1HuFy4dezyL83kWO/5fYlxsHe5yaA=",
19
+ "matmul-nbits.wgsl.jinja": "i+euJpHwxE67Ejb8Pi9XzHDrK3FS/d2nT1odJJxR9rk=",
20
+ "reduce-axis0-splitk-combine.wgsl.jinja": "6S5tsaqhzGAZ66UOQ8u9LfKIlYeGnWu6B722kLl/auc=",
21
+ "test.json": "xR0ewiJSnSNGdmdIovbtC36Xsm1yv4aC8i3l8CaXdaY="
22
  }
23
  },
24
+ "provenance": { "kernel": { "sha": "6fdf6301e2bbcc2f03bf1eaf493b7ad55ef33afc", "dirty": false } },
25
  "webgpu": {
26
+ "manifestSpec": "2.1",
27
  "variants": {
28
  "q4_dp4a_prefill": ["matmul-nbits-dp4a-quantize.wgsl.jinja", "matmul-nbits-q4-dp4a-prefill.wgsl.jinja"],
29
  "gemv_default_zero": ["matmul-nbits-gemv-q4.wgsl.jinja"],
build/webgpu/reduce-axis0-splitk-combine.wgsl.jinja CHANGED
@@ -2,31 +2,16 @@
2
  // folds the segment partials and applies the selected reduction's final step.
3
  // Segments are folded in ascending order for deterministic results. This order
4
  // differs from the single-pass reduction but remains within the f32 tolerance.
5
- {% set addBias = addBias is defined and addBias %}
6
- {% set biasCols = biasCols | default(0) %}
7
- {% set intMode = intMode is defined and intMode %}
8
  {% set yv = "f16(" if outputF16 else "" %}
9
  {% set vy = ")" if outputF16 else "" %}
10
- {% if outputF16 %}
11
- enable f16;
12
- {% endif %}
13
  {{ env.wgsl.resourceDeclarations }}
14
  {% macro wgsl_minmax_identity(name, op, scalar="f32") %}
15
  /* Exact {{ op }} reduction identity. WGSL rejects infinity during constant
16
  * evaluation, so an f32 identity is constructed from its IEEE-754 bits. */
17
- {% set logicalBoolIdentity = logicalBool is defined and logicalBool %}
18
  fn {{ name }}() -> {{ scalar }} {
19
- {% if scalar == "i32" %}
20
- return {{ "-2147483647i - 1i" if op == "max" else "2147483647i" }};
21
- {% elif scalar == "u32" %}
22
- return {{ "1u" if logicalBoolIdentity and op == "min" else "0u" if op == "max" else "4294967295u" }};
23
- {% else %}
24
  var bits = {{ "0xff800000u" if op == "max" else "0x7f800000u" }};
25
  return bitcast<f32>(bits);
26
- {% endif %}
27
- }
28
- {%- endmacro %}
29
-
30
 
31
  const WG: u32 = {{ workgroupSize }}u;
32
  const SPLIT: u32 = {{ split }}u;
@@ -48,7 +33,9 @@ fn is_nan_f32(value: f32) -> bool {
48
  @compute @workgroup_size(WG, 1, 1)
49
  fn main(@builtin(global_invocation_id) gid: vec3<u32>,
50
  @builtin(num_workgroups) nwg: vec3<u32>) {
51
- let stride = nwg.x * WG;
 
 
52
  let start = (gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG) + gid.x;
53
  for (var col = start; col < params.cols; col = col + stride) {
54
  {% if op == "logsumexp" %}
@@ -74,13 +61,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>,
74
  let finite_or_inf = select(global_max + log(sum), global_max, has_positive_inf);
75
  y[col] = {{ yv }}select(finite_or_inf, nan_value, has_nan){{ vy }};
76
  {% else %}
77
- {% if intMode %}
78
- {% if op == "prod" %}
79
- var total = 1i;
80
- {% else %}
81
- var total = 0i;
82
- {% endif %}
83
- {% else %}
84
  {% if op == "max" %}
85
  var total = reduction_identity();
86
  {% elif op == "min" %}
@@ -89,7 +69,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>,
89
  var total = 1.0;
90
  {% else %}
91
  var total = 0.0;
92
- {% endif %}
93
  {% endif %}
94
  for (var seg = 0u; seg < SPLIT; seg = seg + 1u) {
95
  let p = partials[seg * params.cols + col];
 
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;
 
33
  @compute @workgroup_size(WG, 1, 1)
34
  fn main(@builtin(global_invocation_id) gid: vec3<u32>,
35
  @builtin(num_workgroups) nwg: vec3<u32>) {
36
+ // The start already folds gid.y in, so the stride must span every y row too;
37
+ // an x-only stride would send y = 0 lanes over columns the y >= 1 rows own.
38
+ let stride = nwg.x * nwg.y * WG;
39
  let start = (gid.y * {{ DISPATCH_FOLD_WIDTH }}u * WG) + gid.x;
40
  for (var col = start; col < params.cols; col = col + stride) {
41
  {% if op == "logsumexp" %}
 
61
  let finite_or_inf = select(global_max + log(sum), global_max, has_positive_inf);
62
  y[col] = {{ yv }}select(finite_or_inf, nan_value, has_nan){{ vy }};
63
  {% else %}
 
 
 
 
 
 
 
64
  {% if op == "max" %}
65
  var total = reduction_identity();
66
  {% elif op == "min" %}
 
69
  var total = 1.0;
70
  {% else %}
71
  var total = 0.0;
 
72
  {% endif %}
73
  for (var seg = 0u; seg < SPLIT; seg = seg + 1u) {
74
  let p = partials[seg * params.cols + col];
build/webgpu/test.json CHANGED
@@ -4,7 +4,8 @@
4
  "q4_weight_cycle_b_t": [16, 50, 84, 118, 152, 186, 220, 254, 33, 67, 101, 135, 169, 203, 237, 15],
5
  "quant_scale_cycle_t": [0.04, 0.055, 0.05, 0.065, 0.06, 0.075, 0.07, 0.085],
6
  "mixed_weight_cycle_b_t": [17, 200, 91, 45, 233, 128, 7, 176, 250, 33, 142, 99, 210, 64, 188, 121],
7
- "q8_zero_bias_gemv_m1_tail_n5_input_bT": [19, 56, 93, 130, 167, 204, 241, 22, 59, 96, 133, 170, 207, 244, 25, 62]
 
8
  },
9
  "cases": [
10
  {
@@ -42,7 +43,7 @@
42
  {
43
  "name": "q4_zero_bias_prefill_tile4x4_small_m8",
44
  "provenance": {
45
- "notes": "Small-M (M=8) q4 prefill with bias and zero points. M<64 excludes prefill_tiled_zero_bias, while the row-guarded prefill_tile4x4_zero_bias route admits M>=2 when N is divisible by 4. This pins the tile4x4 lower-bound contract and its bias/zero-point arithmetic."
46
  },
47
  "inputs": {
48
  "aT": {
@@ -406,7 +407,7 @@
406
  {
407
  "name": "q8_zero_bias_prefill_reg_vec4_splitk_m128_k1024_n1024",
408
  "provenance": {
409
- "notes": "With zero points and bias, the split-K four-wide route applies zero points in each partial pass and adds bias once in the combine."
410
  },
411
  "inputs": {
412
  "aT": {
@@ -440,9 +441,7 @@
440
  },
441
  {
442
  "name": "q8_zero_only_prefill_reg_vec4_splitk_m128_k1024_n1024",
443
- "provenance": {
444
- "notes": "With zero points and no bias, the split-K four-wide route applies zero points in each partial pass and only sums in the combine."
445
- },
446
  "inputs": {
447
  "aT": {
448
  "dtype": "float32",
@@ -932,7 +931,7 @@
932
  {
933
  "name": "q4_no_zero_prefill_tile4x4_partial_row_tile_m6",
934
  "provenance": {
935
- "notes": "M=6 is not a multiple of TILE_M=4, so the tile4x4 kernel's second row-tile (row_base=4) has valid rows 4,5 and guarded rows 6,7. Verifies the store_row partial-row-tile guard writes rows 4,5 correctly and does not corrupt/OOB rows 6,7. N=8 (%4==0), K=64, blockSize=32 routes to prefill_tile4x4_default_zero."
936
  },
937
  "inputs": {
938
  "aT": {
@@ -1033,7 +1032,7 @@
1033
  "name": "q4_gemv_default_zero_m1_n13_ncols8",
1034
  "tunables": { "GEMV_N_COLS": 8 },
1035
  "provenance": {
1036
- "notes": "GEMV_N_COLS=8 with N=13: two workgroups, first fully live, second with a partially live first group and one live column in the second."
1037
  },
1038
  "inputs": {
1039
  "aT": {
@@ -1059,7 +1058,7 @@
1059
  "name": "q4_gemv_default_zero_m1_tail_n7_ncols8",
1060
  "tunables": { "GEMV_N_COLS": 8 },
1061
  "provenance": {
1062
- "notes": "With `GEMV_N_COLS = 8` and N=7, one workgroup has active columns 4 through 6 and a fully guarded column 7 in its second group."
1063
  },
1064
  "inputs": {
1065
  "aT": {
@@ -1184,7 +1183,7 @@
1184
  "name": "q8_zero_bias_gemv_m1_tail_n5_ncols8",
1185
  "tunables": { "GEMV_N_COLS": 8 },
1186
  "provenance": {
1187
- "notes": "With `GEMV_N_COLS = 8` and N=5, q8 unpacking, explicit zero points, and bias run with one live column in the second group."
1188
  },
1189
  "inputs": {
1190
  "aT": {
@@ -1219,7 +1218,7 @@
1219
  {
1220
  "name": "q8_zero_bias_naive_fallback_tailK_m3_n6",
1221
  "provenance": {
1222
- "notes": "K=17 with block size 16 leaves a partial block; M=3 bypasses GEMV and N=6 remains below the tiled floors. The scalar fallback handles q8 unpacking, per-block zero points, bias, and the K tail."
1223
  },
1224
  "inputs": {
1225
  "aT": {
@@ -1254,7 +1253,7 @@
1254
  {
1255
  "name": "q4_prefill_tiled_reg_tailk_m32_k33_n4096",
1256
  "provenance": {
1257
- "notes": "Compact tail-block lock for the register-tiled prefill path used by the realistic K=2561 benchmark; the final 31 padded weights must not read past A."
1258
  },
1259
  "inputs": {
1260
  "aT": {
@@ -1367,7 +1366,7 @@
1367
  {
1368
  "name": "q8_no_zero_prefill_odd_n_fallback",
1369
  "provenance": {
1370
- "notes": "M>1 q8 prefill with N=17 and no zero_points/bias. Odd N excludes subgroup-matrix execution; the portable tile4x4 tail guards handle the final output column used by the odd-column benchmark guardrail."
1371
  },
1372
  "inputs": {
1373
  "aT": {
@@ -2077,7 +2076,7 @@
2077
  {
2078
  "name": "q8_zero_only_naive_fallback_tailk_m3_n6",
2079
  "provenance": {
2080
- "notes": "K=17 leaves a partial final block and M=3 bypasses GEMV. With explicit zero points but no bias, the aligned prefill paths are ineligible and the zero-only scalar fallback handles the tail block."
2081
  },
2082
  "inputs": {
2083
  "aT": {
@@ -2107,7 +2106,7 @@
2107
  {
2108
  "name": "q8_bias_only_naive_fallback_tailk_m3_n6",
2109
  "provenance": {
2110
- "notes": "K=17 leaves a partial final block and M=3 bypasses GEMV. With bias and the schema-default zero point, the aligned prefill paths are ineligible and the bias-only scalar fallback handles the tail block."
2111
  },
2112
  "inputs": {
2113
  "aT": {
@@ -2278,6 +2277,177 @@
2278
  },
2279
  "outputs": { "yT": { "dtype": "float32", "shape": [1, 5], "tolerance": 0.0001 } },
2280
  "attrs": { "K": 40, "N": 5, "bits": 4, "block_size": 16 }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2281
  }
2282
  ]
2283
  }
 
4
  "q4_weight_cycle_b_t": [16, 50, 84, 118, 152, 186, 220, 254, 33, 67, 101, 135, 169, 203, 237, 15],
5
  "quant_scale_cycle_t": [0.04, 0.055, 0.05, 0.065, 0.06, 0.075, 0.07, 0.085],
6
  "mixed_weight_cycle_b_t": [17, 200, 91, 45, 233, 128, 7, 176, 250, 33, 142, 99, 210, 64, 188, 121],
7
+ "q8_zero_bias_gemv_m1_tail_n5_input_bT": [19, 56, 93, 130, 167, 204, 241, 22, 59, 96, 133, 170, 207, 244, 25, 62],
8
+ "ort_f16_large_k_accumulator_cancellation_m1_k8192_n8_input_scalesT": [16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16, -16]
9
  },
10
  "cases": [
11
  {
 
43
  {
44
  "name": "q4_zero_bias_prefill_tile4x4_small_m8",
45
  "provenance": {
46
+ "notes": "M=8 rows of 4-bit quantized MatMulNBits with K=128, N=64, block_size=32, explicit zero points [6,7,8,9], and bias check correct dequantization and bias addition at a small row count."
47
  },
48
  "inputs": {
49
  "aT": {
 
407
  {
408
  "name": "q8_zero_bias_prefill_reg_vec4_splitk_m128_k1024_n1024",
409
  "provenance": {
410
+ "notes": "Zero points affect the complete reduction, and bias is added exactly once to the final output."
411
  },
412
  "inputs": {
413
  "aT": {
 
441
  },
442
  {
443
  "name": "q8_zero_only_prefill_reg_vec4_splitk_m128_k1024_n1024",
444
+ "provenance": { "notes": "Zero points affect the complete reduction without an output bias." },
 
 
445
  "inputs": {
446
  "aT": {
447
  "dtype": "float32",
 
931
  {
932
  "name": "q4_no_zero_prefill_tile4x4_partial_row_tile_m6",
933
  "provenance": {
934
+ "notes": "M=6 rows of 4-bit MatMulNBits (K=64, N=8, block_size=32, default zero point, no bias) check a row count that is not a multiple of four."
935
  },
936
  "inputs": {
937
  "aT": {
 
1032
  "name": "q4_gemv_default_zero_m1_n13_ncols8",
1033
  "tunables": { "GEMV_N_COLS": 8 },
1034
  "provenance": {
1035
+ "notes": "N=13 with M=1, K=32, block_size=32, 4-bit weights and the default zero point checks a column count that is not a multiple of eight."
1036
  },
1037
  "inputs": {
1038
  "aT": {
 
1058
  "name": "q4_gemv_default_zero_m1_tail_n7_ncols8",
1059
  "tunables": { "GEMV_N_COLS": 8 },
1060
  "provenance": {
1061
+ "notes": "N=7 with M=1, K=32, block_size=32, 4-bit weights and the default zero point checks that all seven output columns are dequantized and written correctly."
1062
  },
1063
  "inputs": {
1064
  "aT": {
 
1183
  "name": "q8_zero_bias_gemv_m1_tail_n5_ncols8",
1184
  "tunables": { "GEMV_N_COLS": 8 },
1185
  "provenance": {
1186
+ "notes": "N=5 with M=1, K=16, block_size=16, 8-bit weights, explicit per-block zero points and a bias checks that all five output columns are dequantized and written correctly."
1187
  },
1188
  "inputs": {
1189
  "aT": {
 
1218
  {
1219
  "name": "q8_zero_bias_naive_fallback_tailK_m3_n6",
1220
  "provenance": {
1221
+ "notes": "K=17 with block size 16 leaves a partial quantization block. M=3 and N=6 check q8 unpacking, per-block zero points, bias, and the K tail."
1222
  },
1223
  "inputs": {
1224
  "aT": {
 
1253
  {
1254
  "name": "q4_prefill_tiled_reg_tailk_m32_k33_n4096",
1255
  "provenance": {
1256
+ "notes": "A 33-element reduction leaves a partial final quantization block; padded weights must not contribute to the result or read past the activation input."
1257
  },
1258
  "inputs": {
1259
  "aT": {
 
1366
  {
1367
  "name": "q8_no_zero_prefill_odd_n_fallback",
1368
  "provenance": {
1369
+ "notes": "An eight-row q8 projection with N=17 and no zero points or bias checks the final odd output column."
1370
  },
1371
  "inputs": {
1372
  "aT": {
 
2076
  {
2077
  "name": "q8_zero_only_naive_fallback_tailk_m3_n6",
2078
  "provenance": {
2079
+ "notes": "K=17 with block_size=16 leaves one element in a partial final quantization block (M=3, N=6, 8-bit, explicit per-block zero points, no bias); checks that the partial block is weighted correctly without bias."
2080
  },
2081
  "inputs": {
2082
  "aT": {
 
2106
  {
2107
  "name": "q8_bias_only_naive_fallback_tailk_m3_n6",
2108
  "provenance": {
2109
+ "notes": "K=17 with block_size=16 leaves one element in a partial final quantization block (M=3, N=6, 8-bit, default zero point, explicit bias); checks that the partial block and bias combine correctly."
2110
  },
2111
  "inputs": {
2112
  "aT": {
 
2277
  },
2278
  "outputs": { "yT": { "dtype": "float32", "shape": [1, 5], "tolerance": 0.0001 } },
2279
  "attrs": { "K": 40, "N": 5, "bits": 4, "block_size": 16 }
2280
+ },
2281
+ {
2282
+ "name": "ort_f16_large_k_accumulator_cancellation_m1_k8192_n8",
2283
+ "provenance": {
2284
+ "notes": "Transcribed from ORT MatMulNBits.Float16_LargeK_AccumulatorOverflow (M=1 arm). A = 8 everywhere; the dequantized weight is +112 over the first half of K and -112 over the second, so the exact result is 0 while the running partial sum crosses the f16 ceiling (65504) at 114688 in any kernel that walks K in one accumulator. ORT puts the sign on B's codes (0xFF then 0x11); this package puts it on the block scale instead (+16 then -16) with every code 15, which yields a bit-identical dequantized weight matrix and lets B be one constant instead of a 4096-long cycle."
2285
+ },
2286
+ "attrs": { "K": 8192, "N": 8, "bits": 4, "block_size": 32 },
2287
+ "inputs": {
2288
+ "aT": { "dtype": "float16", "shape": [1, 8192], "data": { "kind": "constant", "value": 8.0 } },
2289
+ "bT": { "dtype": "uint8", "shape": [8, 256, 16], "data": { "kind": "constant", "value": 255 } },
2290
+ "scalesT": {
2291
+ "dtype": "float16",
2292
+ "shape": [8, 256],
2293
+ "data": {
2294
+ "kind": "cycle",
2295
+ "values": { "$ref": "#/fixtureArrays/ort_f16_large_k_accumulator_cancellation_m1_k8192_n8_input_scalesT" }
2296
+ }
2297
+ }
2298
+ },
2299
+ "outputs": { "yT": { "dtype": "float16", "shape": [1, 8], "tolerance": 0.05 } }
2300
+ },
2301
+ {
2302
+ "name": "ort_f16_large_k_accumulator_cancellation_m8_k8192_n8",
2303
+ "provenance": {
2304
+ "notes": "Transcribed from ORT MatMulNBits.Float16_LargeK_AccumulatorOverflow (M=8 arm). The same cancellation construction as the M=1 case checks float32 accumulation for eight output rows."
2305
+ },
2306
+ "attrs": { "K": 8192, "N": 8, "bits": 4, "block_size": 32 },
2307
+ "inputs": {
2308
+ "aT": { "dtype": "float16", "shape": [8, 8192], "data": { "kind": "constant", "value": 8.0 } },
2309
+ "bT": { "dtype": "uint8", "shape": [8, 256, 16], "data": { "kind": "constant", "value": 255 } },
2310
+ "scalesT": {
2311
+ "dtype": "float16",
2312
+ "shape": [8, 256],
2313
+ "data": {
2314
+ "kind": "cycle",
2315
+ "values": { "$ref": "#/fixtureArrays/ort_f16_large_k_accumulator_cancellation_m1_k8192_n8_input_scalesT" }
2316
+ }
2317
+ }
2318
+ },
2319
+ "outputs": { "yT": { "dtype": "float16", "shape": [8, 8], "tolerance": 0.05 } }
2320
+ },
2321
+ {
2322
+ "name": "f16_full_k_accumulator_cancellation_prefill_m64_k4096_n64",
2323
+ "provenance": {
2324
+ "notes": "Adapted from ORT MatMulNBits.Float16_LargeK_AccumulatorOverflow. With K=4,096, partial sums rise above the float16 finite ceiling (peak 1,835,008) before cancelling to an exact zero output."
2325
+ },
2326
+ "attrs": { "K": 4096, "N": 64, "bits": 4, "block_size": 32 },
2327
+ "inputs": {
2328
+ "aT": { "dtype": "float16", "shape": [64, 4096], "data": { "kind": "constant", "value": 8.0 } },
2329
+ "bT": { "dtype": "uint8", "shape": [64, 128, 16], "data": { "kind": "constant", "value": 255 } },
2330
+ "scalesT": {
2331
+ "dtype": "float16",
2332
+ "shape": [64, 128],
2333
+ "data": {
2334
+ "kind": "cycle",
2335
+ "values": [16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, 16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0, -16.0]
2336
+ }
2337
+ }
2338
+ },
2339
+ "outputs": { "yT": { "dtype": "float16", "shape": [64, 64], "tolerance": 0.05 } }
2340
+ },
2341
+ {
2342
+ "name": "ort_f16_decode_bias_m1_k1024_n128_b32",
2343
+ "provenance": {
2344
+ "notes": "Transcribed from ORT MatMulNBits.Float16_AccumulatorPrecisionOption_AllPaths, case {M=1, N=128, K=1024, block 32, accuracy_level 0, bias}. The provider option is not applicable here (accumulators are always f32); the shape is kept because it is the generic decode dispatch with a bias and there was no other float16 M=1 case. Weight/scale data reuse this file's existing fixture arrays."
2345
+ },
2346
+ "attrs": { "K": 1024, "N": 128, "bits": 4, "block_size": 32 },
2347
+ "inputs": {
2348
+ "aT": {
2349
+ "dtype": "float16",
2350
+ "shape": [1, 1024],
2351
+ "data": { "kind": "fillFloat32", "sinStep": 0.11, "cosStep": 0.23, "scale": 0.25 }
2352
+ },
2353
+ "bT": {
2354
+ "dtype": "uint8",
2355
+ "shape": [128, 32, 16],
2356
+ "data": { "kind": "cycle", "values": { "$ref": "#/fixtureArrays/q4_weight_cycle_b_t" } }
2357
+ },
2358
+ "scalesT": {
2359
+ "dtype": "float16",
2360
+ "shape": [128, 32],
2361
+ "data": { "kind": "cycle", "values": { "$ref": "#/fixtureArrays/quant_scale_cycle_t" } }
2362
+ },
2363
+ "biasT": {
2364
+ "dtype": "float16",
2365
+ "shape": [128],
2366
+ "data": { "kind": "fillFloat32", "sinStep": 0.07, "cosStep": 0.41, "scale": 0.125 }
2367
+ }
2368
+ },
2369
+ "outputs": { "yT": { "dtype": "float16", "shape": [1, 128], "tolerance": 0.01, "relTolerance": 0.005 } }
2370
+ },
2371
+ {
2372
+ "name": "ort_f16_prefill_bias_m8_k1024_n128_b32",
2373
+ "provenance": {
2374
+ "notes": "Transcribed from ORT MatMulNBits.Float16_AccumulatorPrecisionOption_AllPaths, case {M=8, N=128, K=1024, block 32, accuracy_level 0, bias} (ORT's wide-tile arm)."
2375
+ },
2376
+ "attrs": { "K": 1024, "N": 128, "bits": 4, "block_size": 32 },
2377
+ "inputs": {
2378
+ "aT": {
2379
+ "dtype": "float16",
2380
+ "shape": [8, 1024],
2381
+ "data": { "kind": "fillFloat32", "sinStep": 0.11, "cosStep": 0.23, "scale": 0.25 }
2382
+ },
2383
+ "bT": {
2384
+ "dtype": "uint8",
2385
+ "shape": [128, 32, 16],
2386
+ "data": { "kind": "cycle", "values": { "$ref": "#/fixtureArrays/q4_weight_cycle_b_t" } }
2387
+ },
2388
+ "scalesT": {
2389
+ "dtype": "float16",
2390
+ "shape": [128, 32],
2391
+ "data": { "kind": "cycle", "values": { "$ref": "#/fixtureArrays/quant_scale_cycle_t" } }
2392
+ },
2393
+ "biasT": {
2394
+ "dtype": "float16",
2395
+ "shape": [128],
2396
+ "data": { "kind": "fillFloat32", "sinStep": 0.07, "cosStep": 0.41, "scale": 0.125 }
2397
+ }
2398
+ },
2399
+ "outputs": { "yT": { "dtype": "float16", "shape": [8, 128], "tolerance": 0.01, "relTolerance": 0.005 } }
2400
+ },
2401
+ {
2402
+ "name": "q2_block128_f16_zero1_prefill_m128_k1024_n1024",
2403
+ "provenance": {
2404
+ "notes": "M=128, K=1024, N=1024 with 2-bit weights in 128-value quantization blocks, a uniform explicit zero point of 1, and float16 operands check dequantization at a full-size projection shape."
2405
+ },
2406
+ "attrs": { "K": 1024, "N": 1024, "bits": 2, "block_size": 128 },
2407
+ "inputs": {
2408
+ "aT": {
2409
+ "dtype": "float16",
2410
+ "shape": [128, 1024],
2411
+ "data": { "kind": "fillFloat32", "sinStep": 0.11, "cosStep": 0.23, "scale": 0.25 }
2412
+ },
2413
+ "bT": {
2414
+ "dtype": "uint8",
2415
+ "shape": [1024, 8, 32],
2416
+ "data": { "kind": "cycle", "values": { "$ref": "#/fixtureArrays/mixed_weight_cycle_b_t" } }
2417
+ },
2418
+ "scalesT": {
2419
+ "dtype": "float16",
2420
+ "shape": [1024, 8],
2421
+ "data": { "kind": "cycle", "values": { "$ref": "#/fixtureArrays/quant_scale_cycle_t" } }
2422
+ },
2423
+ "zeroPointsT": { "dtype": "float16", "shape": [1024, 8], "data": { "kind": "constant", "value": 1.0 } }
2424
+ },
2425
+ "outputs": { "yT": { "dtype": "float16", "shape": [128, 1024], "tolerance": 0.01, "relTolerance": 0.01 } }
2426
+ },
2427
+ {
2428
+ "name": "ort_f16_accuracy_level4_m8_k4096_n128_generic_route",
2429
+ "provenance": {
2430
+ "notes": "Transcribed from ORT MatMulNBits.Float16_AccumulatorPrecisionOption_AllPaths: M=8, N=128, K=4,096, block size 32, accuracy_level 4. The float16 inputs check float32 accumulation independently of provider-specific execution choices."
2431
+ },
2432
+ "attrs": { "K": 4096, "N": 128, "bits": 4, "block_size": 32, "accuracy_level": 4 },
2433
+ "inputs": {
2434
+ "aT": {
2435
+ "dtype": "float16",
2436
+ "shape": [8, 4096],
2437
+ "data": { "kind": "fillFloat32", "sinStep": 0.11, "cosStep": 0.23, "scale": 0.25 }
2438
+ },
2439
+ "bT": {
2440
+ "dtype": "uint8",
2441
+ "shape": [128, 128, 16],
2442
+ "data": { "kind": "cycle", "values": { "$ref": "#/fixtureArrays/q4_weight_cycle_b_t" } }
2443
+ },
2444
+ "scalesT": {
2445
+ "dtype": "float16",
2446
+ "shape": [128, 128],
2447
+ "data": { "kind": "cycle", "values": { "$ref": "#/fixtureArrays/quant_scale_cycle_t" } }
2448
+ }
2449
+ },
2450
+ "outputs": { "yT": { "dtype": "float16", "shape": [8, 128], "tolerance": 0.02, "relTolerance": 0.01 } }
2451
  }
2452
  ]
2453
  }