sync 6fdf6301e2bb
Browse files- README.md +10 -10
- build/webgpu/bench.json +66 -0
- build/webgpu/manifest.json +192 -620
- build/webgpu/matmul-nbits-dp4a-quantize.wgsl.jinja +1 -1
- build/webgpu/matmul-nbits-gemv-q4.wgsl.jinja +6 -10
- build/webgpu/matmul-nbits-q4-dp4a-prefill.wgsl.jinja +3 -4
- build/webgpu/matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja +45 -9
- build/webgpu/matmul-nbits-q4-prefill-tiled.wgsl.jinja +5 -6
- build/webgpu/matmul-nbits-q4-sgmat.wgsl.jinja +2 -2
- build/webgpu/matmul-nbits.wgsl.jinja +8 -7
- build/webgpu/metadata.json +14 -14
- build/webgpu/reduce-axis0-splitk-combine.wgsl.jinja +4 -25
- build/webgpu/test.json +185 -15
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` —
|
| 59 |
-
- `prefill_tiled_reg_vec4_default_zero` —
|
| 60 |
-
- `prefill_tiled_reg_vec4_splitk_zero_bias` —
|
| 61 |
-
- `prefill_tiled_reg_vec4_zero_bias` —
|
| 62 |
-
- `prefill_tiled_reg_vec4_splitk_zero_only` —
|
| 63 |
-
- `prefill_tiled_reg_vec4_zero_only` —
|
| 64 |
-
- `prefill_tiled_reg_vec4_splitk_bias_only` —
|
| 65 |
-
- `prefill_tiled_reg_vec4_bias_only` —
|
| 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
|
| 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.
|
| 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 |
-
"
|
|
|
|
|
|
|
| 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 |
-
"
|
| 118 |
-
"
|
| 119 |
-
|
| 120 |
-
|
| 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 |
-
|
| 143 |
-
|
| 144 |
-
|
| 145 |
-
|
| 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 |
-
"
|
|
|
|
| 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)", "
|
| 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), (
|
| 203 |
-
"y": "ceilDiv(ceilDiv((aRows * attrs.K / 4), (
|
| 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": ["
|
| 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": ["
|
| 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 |
-
"
|
| 323 |
-
"
|
| 324 |
-
"
|
| 325 |
-
"
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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": ["
|
| 340 |
-
"dispatch": { "x": "dispatchN64", "y": "
|
| 341 |
},
|
| 342 |
{
|
| 343 |
"id": "combine",
|
| 344 |
"shader": "reduce-axis0-splitk-combine.wgsl.jinja",
|
| 345 |
-
"derive": {
|
| 346 |
-
|
| 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), (
|
| 354 |
-
"y": "ceilDiv(ceilDiv((aRows * attrs.N), (
|
| 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 |
-
"
|
| 370 |
-
"
|
| 371 |
-
"
|
| 372 |
-
"
|
| 373 |
-
"
|
| 374 |
-
"
|
| 375 |
-
"
|
| 376 |
-
"
|
| 377 |
-
"
|
| 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": ["
|
| 394 |
-
"dispatch": { "x": "dispatchN64", "y": "
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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 |
-
"
|
| 628 |
-
"
|
| 629 |
-
"
|
| 630 |
-
"
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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": ["
|
| 645 |
-
"dispatch": { "x": "dispatchN64", "y": "
|
| 646 |
},
|
| 647 |
{
|
| 648 |
"id": "combine",
|
| 649 |
"shader": "reduce-axis0-splitk-combine.wgsl.jinja",
|
| 650 |
-
"derive": {
|
| 651 |
-
|
| 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), (
|
| 659 |
-
"y": "ceilDiv(ceilDiv((aRows * attrs.N), (
|
| 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 |
-
"
|
| 675 |
-
"
|
| 676 |
-
"
|
| 677 |
-
"
|
| 678 |
-
"
|
| 679 |
-
"
|
| 680 |
-
"
|
| 681 |
-
"
|
| 682 |
-
"
|
| 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": ["
|
| 699 |
-
"dispatch": { "x": "dispatchN64", "y": "
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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 |
-
"
|
| 933 |
-
"
|
| 934 |
-
"
|
| 935 |
-
"
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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": ["
|
| 950 |
-
"dispatch": { "x": "dispatchN64", "y": "
|
| 951 |
},
|
| 952 |
{
|
| 953 |
"id": "combine",
|
| 954 |
"shader": "reduce-axis0-splitk-combine.wgsl.jinja",
|
| 955 |
-
"derive": {
|
| 956 |
-
|
| 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), (
|
| 964 |
-
"y": "ceilDiv(ceilDiv((aRows * attrs.N), (
|
| 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 |
-
"
|
| 980 |
-
"
|
| 981 |
-
"
|
| 982 |
-
"
|
| 983 |
-
"
|
| 984 |
-
"
|
| 985 |
-
"
|
| 986 |
-
"
|
| 987 |
-
"
|
| 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": ["
|
| 1004 |
-
"dispatch": { "x": "dispatchN64", "y": "
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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 |
-
"
|
| 1238 |
-
"
|
| 1239 |
-
"
|
| 1240 |
-
"
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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": ["
|
| 1255 |
-
"dispatch": { "x": "dispatchN64", "y": "
|
| 1256 |
},
|
| 1257 |
{
|
| 1258 |
"id": "combine",
|
| 1259 |
"shader": "reduce-axis0-splitk-combine.wgsl.jinja",
|
| 1260 |
-
"derive": {
|
| 1261 |
-
|
| 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), (
|
| 1269 |
-
"y": "ceilDiv(ceilDiv((aRows * attrs.N), (
|
| 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 |
-
"
|
| 1285 |
-
"
|
| 1286 |
-
"
|
| 1287 |
-
"
|
| 1288 |
-
"
|
| 1289 |
-
"
|
| 1290 |
-
"
|
| 1291 |
-
"
|
| 1292 |
-
"
|
| 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": ["
|
| 1309 |
-
"dispatch": { "x": "dispatchN64", "y": "
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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": ["
|
| 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 =
|
| 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 |
-
{
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
| 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 *
|
| 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 *
|
| 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.
|
| 10 |
-
//
|
|
|
|
| 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 |
-
|
| 28 |
-
|
|
|
|
|
|
|
| 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 |
-
{
|
| 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(
|
| 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 |
-
{
|
| 47 |
-
{{- matmul_nbits_packed_code(bits=bits, kBlocks="KBLOCKS", blobSize="BLOB_SIZE") }}
|
| 48 |
|
| 49 |
-
@compute @workgroup_size(
|
| 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 *
|
| 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 *
|
| 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 =
|
| 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 |
-
{
|
| 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 |
-
|
| 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": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"backend": { "type": "webgpu" },
|
| 7 |
"digest": {
|
| 8 |
"algorithm": "sha256",
|
| 9 |
"files": {
|
| 10 |
-
"bench.json": "
|
| 11 |
-
"manifest.json": "
|
| 12 |
-
"matmul-nbits-dp4a-quantize.wgsl.jinja": "
|
| 13 |
-
"matmul-nbits-gemv-q4.wgsl.jinja": "
|
| 14 |
-
"matmul-nbits-q4-dp4a-prefill.wgsl.jinja": "
|
| 15 |
"matmul-nbits-q4-prefill-tile4x4.wgsl.jinja": "8DVy3szxcxwVItIlEXoioQi2YxDmmd5TCFW4BCYPiCQ=",
|
| 16 |
-
"matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja": "
|
| 17 |
-
"matmul-nbits-q4-prefill-tiled.wgsl.jinja": "
|
| 18 |
-
"matmul-nbits-q4-sgmat.wgsl.jinja": "
|
| 19 |
-
"matmul-nbits.wgsl.jinja": "
|
| 20 |
-
"reduce-axis0-splitk-combine.wgsl.jinja": "
|
| 21 |
-
"test.json": "
|
| 22 |
}
|
| 23 |
},
|
| 24 |
-
"provenance": { "kernel": { "sha": "
|
| 25 |
"webgpu": {
|
| 26 |
-
"manifestSpec": "2.
|
| 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 |
-
{%
|
| 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 |
-
|
|
|
|
|
|
|
| 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": "
|
| 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": "
|
| 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
|
| 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": "
|
| 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": "
|
| 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": "
|
| 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
|
| 1223 |
},
|
| 1224 |
"inputs": {
|
| 1225 |
"aT": {
|
|
@@ -1254,7 +1253,7 @@
|
|
| 1254 |
{
|
| 1255 |
"name": "q4_prefill_tiled_reg_tailk_m32_k33_n4096",
|
| 1256 |
"provenance": {
|
| 1257 |
-
"notes": "
|
| 1258 |
},
|
| 1259 |
"inputs": {
|
| 1260 |
"aT": {
|
|
@@ -1367,7 +1366,7 @@
|
|
| 1367 |
{
|
| 1368 |
"name": "q8_no_zero_prefill_odd_n_fallback",
|
| 1369 |
"provenance": {
|
| 1370 |
-
"notes": "
|
| 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
|
| 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
|
| 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 |
}
|