Download build/webgpu/manifest.json from webgpu-kernels/com.microsoft.MatMulNBits: direct link, hf CLI and curl.
- Browser
- Download file 54.8 kB
-
https://huggingface.co/kernels/webgpu-kernels/com.microsoft.MatMulNBits/resolve/v1/build/webgpu/manifest.json
- Command line
-
hf download hf://webgpu-kernels/com.microsoft.MatMulNBits@v1/build/webgpu/manifest.json
-
curl -L -o manifest.json https://huggingface.co/kernels/webgpu-kernels/com.microsoft.MatMulNBits/resolve/v1/build/webgpu/manifest.json
54.8 kB
| { | |
| "domain": "com.microsoft", | |
| "name": "MatMulNBits", | |
| "sinceVersion": 1, | |
| "inputs": { | |
| "aT": { "onnx": "A", "dtype": "T1" }, | |
| "bT": { "onnx": "B", "dtype": "uint8", "rank": 3, "layout": "packed" }, | |
| "scalesT": { "onnx": "scales", "dtype": "T1", "rank": 2 }, | |
| "zeroPointsT": { "onnx": "zero_points", "dtype": "T3", "rank": 2, "optional": true }, | |
| "biasT": { "onnx": "bias", "dtype": "T1", "rank": 1, "optional": true } | |
| }, | |
| "outputs": { "yT": { "onnx": "Y", "dtype": "T1", "rank": "ranks.aT", "shape": "shapes.aT[:-1] + [attrs.N]" } }, | |
| "attributes": { "accuracy_level": { "default": 0 }, "bits": { "default": 4 }, "K": {}, "N": {}, "block_size": {} }, | |
| "attributeConstraints": { | |
| "K": { "required": true }, | |
| "N": { "required": true }, | |
| "accuracy_level": { "values": [0, 1, 2, 3, 4] }, | |
| "bits": { "values": [2, 4, 8] }, | |
| "block_size": { "required": true } | |
| }, | |
| "typeConstraints": { "T1": ["float32", "float16"], "T3": ["float32", "float16"] }, | |
| "tunables": { | |
| "REGISTER_TILE_TAILK_MIN_N": { "default": 256 }, | |
| "WORKGROUP_SIZE": { "default": 64 }, | |
| "GEMV_N_COLS": { "default": 4 }, | |
| "REGISTER_TILE_MIN_N": { "default": 1024 }, | |
| "REGISTER_TILE_TALL_MIN_M": { "default": 512 }, | |
| "REGISTER_TILE_TALL_MIN_N": { "default": 256 }, | |
| "REGISTER_TILE_MIN_WORKGROUPS": { "default": 64 }, | |
| "REGISTER_TILE_LARGE_M": { "default": 64 }, | |
| "REGISTER_TILE_BK32_MIN_M": { "default": 128 }, | |
| "REGISTER_TILE_MIN_M": { "default": 48 }, | |
| "REGISTER_TILE_TAILK_MIN_M": { "default": 16 }, | |
| "REGISTER_TILE_SPLITK_TARGET_WORKGROUPS": { "default": 512 }, | |
| "REGISTER_TILE_SPLITK_MAX_TILES": { "default": 256 } | |
| }, | |
| "derive": { | |
| "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)", | |
| "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32", | |
| "subgroupsWave32": "device.features.has(\"subgroups\") and wave32Adapter", | |
| "narrowSubgroupRange": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize < device.adapterInfo.subgroupMaxSize and device.adapterInfo.subgroupMaxSize <= 16", | |
| "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", | |
| "pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter", | |
| "wave32Effective": "wave32Adapter or pinSubgroupSize32", | |
| "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", | |
| "packedFeature": "device.wgslLanguageFeatures.has(\"packed_4x8_integer_dot_product\")", | |
| "kBlocksExpected": "ceilDiv(attrs.K, attrs.block_size)", | |
| "blobSizeExpected": "ceilDiv(attrs.block_size * attrs.bits, 8)", | |
| "aRows": "numel(shapes.aT) / max(1, attrs.K)", | |
| "aRankOk": "(ranks.aT == 2 or ranks.aT == 3) and ranks.yT == ranks.aT and dim(shapes.aT, ranks.aT - 1) == attrs.K and dim(shapes.yT, ranks.yT - 1) == attrs.N and dim(shapes.yT, 0) == dim(shapes.aT, 0) and (ranks.aT == 2 or dim(shapes.yT, 1) == dim(shapes.aT, 1))", | |
| "gemvDispatchN": "ceilDiv(attrs.N, tunables.GEMV_N_COLS)", | |
| "dispatchN32": "ceilDiv(attrs.N, 32)", | |
| "dispatchN64": "ceilDiv(attrs.N, 64)", | |
| "dispatchM32": "ceilDiv(aRows, 32)", | |
| "dispatchM64": "ceilDiv(aRows, 64)", | |
| "sgmatTileRows": "64 if aRows >= 64 else 32", | |
| "sgmatWorkgroupSize": "256 if aRows >= 64 else 128", | |
| "sgmatRowSubtiles": "4 if aRows >= 64 else 2", | |
| "sgmatBLoadsPerRow": "sgmatRowSubtiles", | |
| "sgmatBLoadWidth": "8 if aRows >= 64 else 16", | |
| "sgmatNumSubgroups": "8 if aRows >= 64 else 4", | |
| "sgmatWorkgroupStorageBytes": "24576 if aRows >= 64 else 16384", | |
| "sgmatDispatchM": "dispatchM64 if aRows >= 64 else dispatchM32", | |
| "tiledRegBK": "32 if aRows >= tunables.REGISTER_TILE_BK32_MIN_M else 16", | |
| "aFloatOk": "(tensorDtypes.aT == \"float32\" or tensorDtypes.aT == \"float16\") and f16Ok(tensorDtypes.aT)", | |
| "portableWorkgroupSize": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)", | |
| "bitsSupported": "attrs.bits == 2 or attrs.bits == 4 or attrs.bits == 8", | |
| "blockSizeSupported": "attrs.block_size >= 16 and attrs.block_size == pow2ceil(attrs.block_size)", | |
| "commonShapeValid": "aRankOk and ranks.bT == 3 and ranks.scalesT == 2 and aFloatOk and blockSizeSupported and tensorDtypes.bT == \"uint8\" and tensorDtypes.scalesT == tensorDtypes.aT and tensorDtypes.yT == tensorDtypes.aT and attrs.K > 0 and attrs.N > 0 and dim(shapes.bT, 0) == attrs.N and dim(shapes.bT, 1) == kBlocksExpected and dim(shapes.bT, 2) == blobSizeExpected and dim(shapes.scalesT, 0) == attrs.N and dim(shapes.scalesT, 1) == dim(shapes.bT, 1)", | |
| "gemvShapeValid": "commonShapeValid and aRows == 1", | |
| "zeroPointsValid": "present.zeroPointsT and ranks.zeroPointsT == 2 and tensorDtypes.zeroPointsT == tensorDtypes.aT and dim(shapes.zeroPointsT, 0) == attrs.N and dim(shapes.zeroPointsT, 1) == dim(shapes.bT, 1)", | |
| "biasValid": "present.biasT and ranks.biasT == 1 and tensorDtypes.biasT == tensorDtypes.aT and dim(shapes.biasT, 0) == attrs.N", | |
| "defaultEpilogue": "not present.zeroPointsT and not present.biasT", | |
| "zeroBiasEpilogue": "zeroPointsValid and biasValid", | |
| "zeroOnlyEpilogue": "zeroPointsValid and not present.biasT", | |
| "biasOnlyEpilogue": "not present.zeroPointsT and biasValid", | |
| "portableWorkgroupFits": "portableWorkgroupSize > 0 and portableWorkgroupSize * 64 <= device.limits.maxComputeWorkgroupStorageSize", | |
| "tiledWorkgroupSide": 16, | |
| "dp4aQuantizeWorkgroupSize": 64, | |
| "tiledWorkgroupFits": "tiledWorkgroupSide <= device.limits.maxComputeWorkgroupSizeX and tiledWorkgroupSide <= device.limits.maxComputeWorkgroupSizeY and tiledWorkgroupSide * tiledWorkgroupSide <= device.limits.maxComputeInvocationsPerWorkgroup and 4096 <= device.limits.maxComputeWorkgroupStorageSize", | |
| "tiledRegWorkgroupFits": "tiledWorkgroupFits and 16384 <= device.limits.maxComputeWorkgroupStorageSize", | |
| "mediumTiledRegWorkgroupFits": "tiledWorkgroupFits and 6144 <= device.limits.maxComputeWorkgroupStorageSize", | |
| "sgmatWorkgroupFits": "sgmatWorkgroupSize <= deviceWorkgroupCap and sgmatWorkgroupStorageBytes <= device.limits.maxComputeWorkgroupStorageSize", | |
| "registerTileShape": "(aRows >= tunables.REGISTER_TILE_MIN_M or (attrs.K % attrs.block_size != 0 and aRows >= tunables.REGISTER_TILE_TAILK_MIN_M)) and ((attrs.N >= tunables.REGISTER_TILE_MIN_N or (aRows >= tunables.REGISTER_TILE_TALL_MIN_M and attrs.N >= tunables.REGISTER_TILE_TALL_MIN_N) or (attrs.K % attrs.block_size != 0 and attrs.N >= tunables.REGISTER_TILE_TAILK_MIN_N)) and (attrs.K % attrs.block_size != 0 or aRows >= tunables.REGISTER_TILE_LARGE_M or dispatchM64 * dispatchN64 >= tunables.REGISTER_TILE_MIN_WORKGROUPS))", | |
| "portableTile4Preferred": "registerTileShape and attrs.K % attrs.block_size == 0 and (not device.features.has(\"subgroups\") or narrowSubgroupRange) and has(device.adapterInfo, \"subgroupMinSize\") and device.adapterInfo.subgroupMinSize * 2 < portableWorkgroupSize", | |
| "portableMediumRegisterPreferred": "portableTile4Preferred and aRows >= 128 and attrs.K >= 128", | |
| "registerTilePreferred": "registerTileShape and not portableTile4Preferred", | |
| "mediumRegisterEligible": "registerTilePreferred or portableMediumRegisterPreferred", | |
| "largeTiledRegEligible": "registerTilePreferred and tiledRegWorkgroupFits and dispatchM64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "mediumTiledRegEligible": "mediumRegisterEligible and mediumTiledRegWorkgroupFits and dispatchM32 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "tiledRegVariantEligible": "largeTiledRegEligible or mediumTiledRegEligible", | |
| "tiledRegSelectedBK": "tiledRegBK if largeTiledRegEligible else 16", | |
| "tiledRegSelectedThreadRows": "4 if largeTiledRegEligible else 2", | |
| "tiledRegSelectedTileRows": "tiledRegSelectedThreadRows * tiledWorkgroupSide", | |
| "tiledRegSelectedDispatchM": "dispatchM64 if largeTiledRegEligible else dispatchM32", | |
| "blobWords": "blobSizeExpected / 4", | |
| "codesPerWord": "32 / attrs.bits", | |
| "gemvVecWords": "4 if blobSizeExpected % 16 == 0 else 1", | |
| "gemvCodesPerVec": "gemvVecWords * codesPerWord", | |
| "gemvVecPerBlock": "blobSizeExpected / (4 * gemvVecWords)", | |
| "gemvVecPerCol": "kBlocksExpected * gemvVecPerBlock", | |
| "gemvActVec4": "attrs.K % attrs.block_size == 0", | |
| "smallMWordsPerCol": "kBlocksExpected * blobWords", | |
| "smallMKLanes": "min(portableWorkgroupSize, pow2ceil(smallMWordsPerCol + 1) / 2)", | |
| "smallMColGroups": "portableWorkgroupSize / smallMKLanes", | |
| "smallMDispatchN": "ceilDiv(attrs.N, 4 * smallMColGroups)", | |
| "tiledRegVec4ThreadRows": "8 if aRows >= 256 else 4", | |
| "tiledRegVec4TileRows": "tiledRegVec4ThreadRows * tiledWorkgroupSide", | |
| "tiledRegVec4DispatchM": "ceilDiv(aRows, tiledRegVec4TileRows)", | |
| "tiledRegVec4Eligible": "largeTiledRegEligible and attrs.K % 4 == 0 and attrs.block_size % 32 == 0 and tiledRegVec4DispatchM <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "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", | |
| "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", | |
| "tiledRegSplitTiles": "tiledRegVec4DispatchM * dispatchN64", | |
| "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", | |
| "tiledRegSplitWant": "ceilDiv(tunables.REGISTER_TILE_SPLITK_TARGET_WORKGROUPS, max(1, tiledRegSplitTiles))", | |
| "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))", | |
| "tiledRegSplitTilesPerSplit": "ceilDiv(ceilDiv(attrs.K, 32), tiledRegSplitK)", | |
| "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)", | |
| "hasZero": "present.zeroPointsT", | |
| "hasBias": "present.biasT", | |
| "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", | |
| "bScalar": "\"u32\"", | |
| "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", | |
| "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", | |
| "usesF16": "tensorDtypes.aT == \"float16\"", | |
| "M": "aRows", | |
| "K": "attrs.K", | |
| "N": "attrs.N", | |
| "kBlocks": "dim(shapes.bT, 1)", | |
| "blockSize": "attrs.block_size", | |
| "blobSize": "dim(shapes.bT, 2)", | |
| "bits": "attrs.bits", | |
| "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)", | |
| "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", | |
| "workgroupSize": "portableWorkgroupSize", | |
| "B_LEN": "attrs.N * kBlocksExpected * blobSizeExpected / 4", | |
| "SCALES_LEN": "attrs.N * kBlocksExpected", | |
| "BIAS_LEN": "attrs.N" | |
| }, | |
| "bindings": { | |
| "a_a_t": { "arg": "aT", "name": "a", "elementType": "$aElement" }, | |
| "b_b_t": { "arg": "bT", "name": "b", "elementType": "$bElement", "length": "$B_VEC_LEN" }, | |
| "scales_main": { "arg": "scalesT", "name": "scales", "elementType": "$scaleScalar", "length": "$SCALES_LEN" }, | |
| "y_y_t": { "arg": "yT", "name": "y", "elementType": "$outputScalar" }, | |
| "params": { | |
| "struct": [ | |
| { "name": "K", "type": "u32", "value": "attrs.K" }, | |
| { "name": "N", "type": "u32", "value": "attrs.N" }, | |
| { "name": "kBlocks", "type": "u32", "value": "dim(shapes.bT, 1)" } | |
| ] | |
| }, | |
| "zero_points": { "arg": "zeroPointsT", "elementType": "$aScalar", "length": "$SCALES_LEN" }, | |
| "bias": { "arg": "biasT", "elementType": "$aScalar", "length": "$BIAS_LEN" }, | |
| "a_main": { "arg": "aT", "name": "a", "elementType": "$aScalar" }, | |
| "b_main": { "arg": "bT", "name": "b", "elementType": "$bScalar", "length": "$B_LEN" }, | |
| "a_partial": { "arg": "aT", "name": "a", "elementType": "$aVec4Element" }, | |
| "y_f32": { "scratch": "partials", "name": "y", "elementType": "f32" }, | |
| "partials": { "buffer": "read-only-storage", "elementType": "f32" }, | |
| "params_cols": { "name": "params", "struct": [{ "name": "cols", "type": "u32", "value": "aRows * attrs.N" }] }, | |
| "params_main": { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "rows", "type": "u32", "value": "aRows" }, | |
| { "name": "K", "type": "u32", "value": "attrs.K" }, | |
| { "name": "N", "type": "u32", "value": "attrs.N" }, | |
| { "name": "blockSize", "type": "u32", "value": "attrs.block_size" }, | |
| { "name": "kBlocks", "type": "u32", "value": "dim(shapes.bT, 1)" }, | |
| { "name": "blobSize", "type": "u32", "value": "dim(shapes.bT, 2)" } | |
| ] | |
| } | |
| }, | |
| "variants": [ | |
| { | |
| "id": "q4_dp4a_prefill", | |
| "priority": 19, | |
| "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"], | |
| "demoteWhen": ["has(device.adapterInfo, \"architecture\") and device.adapterInfo.architecture == \"maxwell\"", "device.features.has(\"chromium-experimental-subgroup-matrix\")", "device.adapterInfo.vendor == \"apple\""], | |
| "derive": { "vec4Count": "aRows * attrs.K / 4", "blockCount": "aRows * attrs.K / 128" }, | |
| "intermediates": [ | |
| { "id": "aQuant", "dtype": "uint32", "shape": "[aRows * attrs.K / 4]" }, | |
| { "id": "aScales", "dtype": "float32", "shape": "[aRows * attrs.K / 128]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "quantize", | |
| "name": "MatMulNBits.Dp4aQuantizeA", | |
| "shader": "matmul-nbits-dp4a-quantize.wgsl.jinja", | |
| "bindings": [ | |
| { "arg": "aT", "name": "a", "elementType": "f32" }, | |
| { "scratch": "aQuant", "name": "a_quant", "elementType": "u32" }, | |
| { "scratch": "aScales", "name": "a_scales", "elementType": "f32" } | |
| ], | |
| "dispatch": { | |
| "x": "min(ceilDiv((aRows * attrs.K / 4), (dp4aQuantizeWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((aRows * attrs.K / 4), (dp4aQuantizeWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "main", | |
| "name": "MatMulNBits.Dp4aPrefill", | |
| "shader": "matmul-nbits-q4-dp4a-prefill.wgsl.jinja", | |
| "bindings": [ | |
| { "scratch": "aQuant", "name": "a_quant", "buffer": "read-only-storage", "elementType": "u32" }, | |
| { "scratch": "aScales", "name": "a_scales", "buffer": "read-only-storage", "elementType": "f32" }, | |
| { "arg": "bT", "name": "b", "length": "$B_LEN" }, | |
| { "arg": "scalesT", "name": "scales", "elementType": "f32", "length": "$SCALES_LEN" }, | |
| { "arg": "yT", "name": "y", "elementType": "f32" } | |
| ], | |
| "dispatch": { "x": "ceilDiv(attrs.N, 64)", "y": "ceilDiv(aRows, 64)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "gemv_default_zero", | |
| "priority": 20, | |
| "when": ["gemvShapeValid", "defaultEpilogue", "bitsSupported", "portableWorkgroupFits"], | |
| "derive": { | |
| "gemvNCols": "tunables.GEMV_N_COLS", | |
| "useSubgroups": "device.features.has(\"subgroups\")", | |
| "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")", | |
| "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"", | |
| "B_VEC_LEN": "attrs.N * gemvVecPerCol", | |
| "vecWords": "gemvVecWords", | |
| "codesPerVec": "gemvCodesPerVec", | |
| "codesPerVec4": "gemvCodesPerVec / 4", | |
| "vecPerBlock": "gemvVecPerBlock", | |
| "vecPerCol": "gemvVecPerCol", | |
| "actVec4": "gemvActVec4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-gemv-q4.wgsl.jinja", | |
| "bindings": ["a_a_t", "b_b_t", "scales_main", "y_y_t", "params"], | |
| "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_sgmat_default_zero", | |
| "priority": 18, | |
| "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "attrs.block_size % 16 == 0", "attrs.K % 32 == 0", "aRows >= 32", "wave32Effective", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "sgmatDispatchM <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "portableWorkgroupFits", "sgmatWorkgroupFits"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "limits": { "maxComputeWorkgroupStorageSize": 16384 }, | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "tileRows": "sgmatTileRows", | |
| "workgroupSize": "sgmatWorkgroupSize", | |
| "rowSubtiles": "sgmatRowSubtiles", | |
| "bLoadsPerRow": "sgmatBLoadsPerRow", | |
| "bLoadWidth": "sgmatBLoadWidth", | |
| "numSubgroups": "sgmatNumSubgroups" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-sgmat.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_vec4_splitk_default_zero", | |
| "priority": 17, | |
| "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"], | |
| "demoteWhen": ["variableSubgroup16To32 and attrs.bits >= 4 and tiledRegSplitK >= 4 and aRows >= 2 * tiledRegVec4TileRows"], | |
| "derive": { | |
| "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"", | |
| "bk": 32, | |
| "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", | |
| "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows", | |
| "threadCols": "2 if q2WideMicroPreferred else 4", | |
| "tileRows": "threadRows * tileWorkgroupY", | |
| "tileCols": "threadCols * tileWorkgroupX", | |
| "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin", | |
| "alignedBlockLoads": true, | |
| "aVec4Loads": true, | |
| "splitK": "tiledRegSplitK", | |
| "tilesPerSplit": "tiledRegSplitTilesPerSplit", | |
| "split": "tiledRegSplitK", | |
| "workgroupSize": 256, | |
| "biasCols": "attrs.N" | |
| }, | |
| "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[tiledRegSplitK * aRows * attrs.N]" }], | |
| "passes": [ | |
| { | |
| "id": "partial", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_partial", "b_main", "scales_main", "y_f32"], | |
| "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)", "z": "tiledRegSplitK" } | |
| }, | |
| { | |
| "id": "combine", | |
| "shader": "reduce-axis0-splitk-combine.wgsl.jinja", | |
| "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "addBias": "hasBias" }, | |
| "bindings": ["partials", "y_y_t", "params_cols"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_vec4_default_zero", | |
| "priority": 16, | |
| "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"], | |
| "derive": { | |
| "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"", | |
| "bk": "16 if q2NarrowBkPreferred else 32", | |
| "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", | |
| "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows", | |
| "threadCols": "2 if q2WideMicroPreferred else 4", | |
| "tileRows": "threadRows * tileWorkgroupY", | |
| "tileCols": "threadCols * tileWorkgroupX", | |
| "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin", | |
| "alignedBlockLoads": true, | |
| "aVec4Loads": true | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_partial", "b_main", "scales_main", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_default_zero", | |
| "priority": 15, | |
| "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"], | |
| "derive": { | |
| "bk": "tiledRegSelectedBK", | |
| "tileRows": "tiledRegSelectedTileRows", | |
| "tileCols": 64, | |
| "threadRows": "tiledRegSelectedThreadRows", | |
| "threadCols": 4, | |
| "alignedBlockLoads": "aRows >= 128 and attrs.block_size % tiledRegSelectedBK == 0" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_default_zero", | |
| "priority": 14, | |
| "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"], | |
| "derive": {}, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "y_y_t"], | |
| "dispatch": { "x": "dispatchN32", "y": "dispatchM32" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tile4x4_default_zero", | |
| "priority": 13, | |
| "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"], | |
| "derive": { | |
| "wordsPerCol": "smallMWordsPerCol", | |
| "wordsPerBlock": "blobWords", | |
| "kLanes": "smallMKLanes", | |
| "colGroups": "smallMColGroups" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "y_y_t", "params_main"], | |
| "dispatch": { | |
| "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)", | |
| "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "z": "ceilDiv(smallMDispatchN, DISPATCH_FOLD_WIDTH)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "default_zero", | |
| "priority": 0, | |
| "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "portableWorkgroupFits"], | |
| "derive": {}, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "y_y_t", "params_main"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "gemv_zero_bias", | |
| "priority": 20, | |
| "when": ["gemvShapeValid", "zeroBiasEpilogue", "bitsSupported", "portableWorkgroupFits"], | |
| "derive": { | |
| "gemvNCols": "tunables.GEMV_N_COLS", | |
| "useSubgroups": "device.features.has(\"subgroups\")", | |
| "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")", | |
| "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"", | |
| "B_VEC_LEN": "attrs.N * gemvVecPerCol", | |
| "vecWords": "gemvVecWords", | |
| "codesPerVec": "gemvCodesPerVec", | |
| "codesPerVec4": "gemvCodesPerVec / 4", | |
| "vecPerBlock": "gemvVecPerBlock", | |
| "vecPerCol": "gemvVecPerCol", | |
| "actVec4": "gemvActVec4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-gemv-q4.wgsl.jinja", | |
| "bindings": ["a_a_t", "b_b_t", "scales_main", "zero_points", "bias", "y_y_t", "params"], | |
| "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_sgmat_zero_bias", | |
| "priority": 18, | |
| "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "attrs.block_size % 16 == 0", "attrs.K % 32 == 0", "aRows >= 32", "wave32Effective", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "sgmatDispatchM <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "portableWorkgroupFits", "sgmatWorkgroupFits"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "limits": { "maxComputeWorkgroupStorageSize": 16384 }, | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "tileRows": "sgmatTileRows", | |
| "workgroupSize": "sgmatWorkgroupSize", | |
| "rowSubtiles": "sgmatRowSubtiles", | |
| "bLoadsPerRow": "sgmatBLoadsPerRow", | |
| "bLoadWidth": "sgmatBLoadWidth", | |
| "numSubgroups": "sgmatNumSubgroups" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-sgmat.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "zero_points", "bias", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_vec4_splitk_zero_bias", | |
| "priority": 17, | |
| "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"], | |
| "demoteWhen": ["variableSubgroup16To32 and attrs.bits >= 4 and tiledRegSplitK >= 4 and aRows >= 2 * tiledRegVec4TileRows"], | |
| "derive": { | |
| "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"", | |
| "bk": 32, | |
| "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", | |
| "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows", | |
| "threadCols": "2 if q2WideMicroPreferred else 4", | |
| "tileRows": "threadRows * tileWorkgroupY", | |
| "tileCols": "threadCols * tileWorkgroupX", | |
| "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin", | |
| "alignedBlockLoads": true, | |
| "aVec4Loads": true, | |
| "splitK": "tiledRegSplitK", | |
| "tilesPerSplit": "tiledRegSplitTilesPerSplit", | |
| "split": "tiledRegSplitK", | |
| "workgroupSize": 256, | |
| "biasCols": "attrs.N" | |
| }, | |
| "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[tiledRegSplitK * aRows * attrs.N]" }], | |
| "passes": [ | |
| { | |
| "id": "partial", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_partial", "b_main", "scales_main", "zero_points", "y_f32"], | |
| "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)", "z": "tiledRegSplitK" } | |
| }, | |
| { | |
| "id": "combine", | |
| "shader": "reduce-axis0-splitk-combine.wgsl.jinja", | |
| "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "addBias": "hasBias" }, | |
| "bindings": ["partials", "bias", "y_y_t", "params_cols"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_vec4_zero_bias", | |
| "priority": 16, | |
| "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"], | |
| "derive": { | |
| "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"", | |
| "bk": "16 if q2NarrowBkPreferred else 32", | |
| "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", | |
| "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows", | |
| "threadCols": "2 if q2WideMicroPreferred else 4", | |
| "tileRows": "threadRows * tileWorkgroupY", | |
| "tileCols": "threadCols * tileWorkgroupX", | |
| "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin", | |
| "alignedBlockLoads": true, | |
| "aVec4Loads": true | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_partial", "b_main", "scales_main", "zero_points", "bias", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_zero_bias", | |
| "priority": 15, | |
| "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"], | |
| "derive": { | |
| "bk": "tiledRegSelectedBK", | |
| "tileRows": "tiledRegSelectedTileRows", | |
| "tileCols": 64, | |
| "threadRows": "tiledRegSelectedThreadRows", | |
| "threadCols": 4, | |
| "alignedBlockLoads": "aRows >= 128 and attrs.block_size % tiledRegSelectedBK == 0" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "zero_points", "bias", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_zero_bias", | |
| "priority": 14, | |
| "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"], | |
| "derive": {}, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "zero_points", "bias", "y_y_t"], | |
| "dispatch": { "x": "dispatchN32", "y": "dispatchM32" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tile4x4_zero_bias", | |
| "priority": 13, | |
| "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"], | |
| "derive": { | |
| "wordsPerCol": "smallMWordsPerCol", | |
| "wordsPerBlock": "blobWords", | |
| "kLanes": "smallMKLanes", | |
| "colGroups": "smallMColGroups" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "zero_points", "bias", "y_y_t", "params_main"], | |
| "dispatch": { | |
| "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)", | |
| "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "z": "ceilDiv(smallMDispatchN, DISPATCH_FOLD_WIDTH)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "zero_bias", | |
| "priority": 0, | |
| "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "portableWorkgroupFits"], | |
| "derive": {}, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "zero_points", "bias", "y_y_t", "params_main"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "gemv_zero_only", | |
| "priority": 20, | |
| "when": ["gemvShapeValid", "zeroOnlyEpilogue", "bitsSupported", "portableWorkgroupFits"], | |
| "derive": { | |
| "gemvNCols": "tunables.GEMV_N_COLS", | |
| "useSubgroups": "device.features.has(\"subgroups\")", | |
| "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")", | |
| "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"", | |
| "B_VEC_LEN": "attrs.N * gemvVecPerCol", | |
| "vecWords": "gemvVecWords", | |
| "codesPerVec": "gemvCodesPerVec", | |
| "codesPerVec4": "gemvCodesPerVec / 4", | |
| "vecPerBlock": "gemvVecPerBlock", | |
| "vecPerCol": "gemvVecPerCol", | |
| "actVec4": "gemvActVec4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-gemv-q4.wgsl.jinja", | |
| "bindings": ["a_a_t", "b_b_t", "scales_main", "zero_points", "y_y_t", "params"], | |
| "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_sgmat_zero_only", | |
| "priority": 18, | |
| "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "attrs.block_size % 16 == 0", "attrs.K % 32 == 0", "aRows >= 32", "wave32Effective", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "sgmatDispatchM <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "portableWorkgroupFits", "sgmatWorkgroupFits"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "limits": { "maxComputeWorkgroupStorageSize": 16384 }, | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "tileRows": "sgmatTileRows", | |
| "workgroupSize": "sgmatWorkgroupSize", | |
| "rowSubtiles": "sgmatRowSubtiles", | |
| "bLoadsPerRow": "sgmatBLoadsPerRow", | |
| "bLoadWidth": "sgmatBLoadWidth", | |
| "numSubgroups": "sgmatNumSubgroups" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-sgmat.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "zero_points", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_vec4_splitk_zero_only", | |
| "priority": 17, | |
| "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"], | |
| "demoteWhen": ["variableSubgroup16To32 and attrs.bits >= 4 and tiledRegSplitK >= 4 and aRows >= 2 * tiledRegVec4TileRows"], | |
| "derive": { | |
| "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"", | |
| "bk": 32, | |
| "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", | |
| "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows", | |
| "threadCols": "2 if q2WideMicroPreferred else 4", | |
| "tileRows": "threadRows * tileWorkgroupY", | |
| "tileCols": "threadCols * tileWorkgroupX", | |
| "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin", | |
| "alignedBlockLoads": true, | |
| "aVec4Loads": true, | |
| "splitK": "tiledRegSplitK", | |
| "tilesPerSplit": "tiledRegSplitTilesPerSplit", | |
| "split": "tiledRegSplitK", | |
| "workgroupSize": 256, | |
| "biasCols": "attrs.N" | |
| }, | |
| "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[tiledRegSplitK * aRows * attrs.N]" }], | |
| "passes": [ | |
| { | |
| "id": "partial", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_partial", "b_main", "scales_main", "zero_points", "y_f32"], | |
| "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)", "z": "tiledRegSplitK" } | |
| }, | |
| { | |
| "id": "combine", | |
| "shader": "reduce-axis0-splitk-combine.wgsl.jinja", | |
| "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "addBias": "hasBias" }, | |
| "bindings": ["partials", "y_y_t", "params_cols"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_vec4_zero_only", | |
| "priority": 16, | |
| "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"], | |
| "derive": { | |
| "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"", | |
| "bk": "16 if q2NarrowBkPreferred else 32", | |
| "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", | |
| "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows", | |
| "threadCols": "2 if q2WideMicroPreferred else 4", | |
| "tileRows": "threadRows * tileWorkgroupY", | |
| "tileCols": "threadCols * tileWorkgroupX", | |
| "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin", | |
| "alignedBlockLoads": true, | |
| "aVec4Loads": true | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_partial", "b_main", "scales_main", "zero_points", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_zero_only", | |
| "priority": 15, | |
| "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"], | |
| "derive": { | |
| "bk": "tiledRegSelectedBK", | |
| "tileRows": "tiledRegSelectedTileRows", | |
| "tileCols": 64, | |
| "threadRows": "tiledRegSelectedThreadRows", | |
| "threadCols": 4, | |
| "alignedBlockLoads": "aRows >= 128 and attrs.block_size % tiledRegSelectedBK == 0" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "zero_points", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_zero_only", | |
| "priority": 14, | |
| "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"], | |
| "derive": {}, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "zero_points", "y_y_t"], | |
| "dispatch": { "x": "dispatchN32", "y": "dispatchM32" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tile4x4_zero_only", | |
| "priority": 13, | |
| "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"], | |
| "derive": { | |
| "wordsPerCol": "smallMWordsPerCol", | |
| "wordsPerBlock": "blobWords", | |
| "kLanes": "smallMKLanes", | |
| "colGroups": "smallMColGroups" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "zero_points", "y_y_t", "params_main"], | |
| "dispatch": { | |
| "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)", | |
| "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "z": "ceilDiv(smallMDispatchN, DISPATCH_FOLD_WIDTH)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "zero_only", | |
| "priority": 0, | |
| "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "portableWorkgroupFits"], | |
| "derive": {}, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "zero_points", "y_y_t", "params_main"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "gemv_bias_only", | |
| "priority": 20, | |
| "when": ["gemvShapeValid", "biasOnlyEpilogue", "bitsSupported", "portableWorkgroupFits"], | |
| "derive": { | |
| "gemvNCols": "tunables.GEMV_N_COLS", | |
| "useSubgroups": "device.features.has(\"subgroups\")", | |
| "aElement": "(\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")", | |
| "bElement": "\"vec4<u32>\" if gemvVecWords == 4 else \"u32\"", | |
| "B_VEC_LEN": "attrs.N * gemvVecPerCol", | |
| "vecWords": "gemvVecWords", | |
| "codesPerVec": "gemvCodesPerVec", | |
| "codesPerVec4": "gemvCodesPerVec / 4", | |
| "vecPerBlock": "gemvVecPerBlock", | |
| "vecPerCol": "gemvVecPerCol", | |
| "actVec4": "gemvActVec4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-gemv-q4.wgsl.jinja", | |
| "bindings": ["a_a_t", "b_b_t", "scales_main", "bias", "y_y_t", "params"], | |
| "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_sgmat_bias_only", | |
| "priority": 18, | |
| "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "attrs.block_size % 16 == 0", "attrs.K % 32 == 0", "aRows >= 32", "wave32Effective", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "sgmatDispatchM <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "portableWorkgroupFits", "sgmatWorkgroupFits"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "limits": { "maxComputeWorkgroupStorageSize": 16384 }, | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "tileRows": "sgmatTileRows", | |
| "workgroupSize": "sgmatWorkgroupSize", | |
| "rowSubtiles": "sgmatRowSubtiles", | |
| "bLoadsPerRow": "sgmatBLoadsPerRow", | |
| "bLoadWidth": "sgmatBLoadWidth", | |
| "numSubgroups": "sgmatNumSubgroups" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-sgmat.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "bias", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "sgmatDispatchM" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_vec4_splitk_bias_only", | |
| "priority": 17, | |
| "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitEligible"], | |
| "demoteWhen": ["variableSubgroup16To32 and attrs.bits >= 4 and tiledRegSplitK >= 4 and aRows >= 2 * tiledRegVec4TileRows"], | |
| "derive": { | |
| "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"", | |
| "bk": 32, | |
| "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", | |
| "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows", | |
| "threadCols": "2 if q2WideMicroPreferred else 4", | |
| "tileRows": "threadRows * tileWorkgroupY", | |
| "tileCols": "threadCols * tileWorkgroupX", | |
| "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin", | |
| "alignedBlockLoads": true, | |
| "aVec4Loads": true, | |
| "splitK": "tiledRegSplitK", | |
| "tilesPerSplit": "tiledRegSplitTilesPerSplit", | |
| "split": "tiledRegSplitK", | |
| "workgroupSize": 256, | |
| "biasCols": "attrs.N" | |
| }, | |
| "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[tiledRegSplitK * aRows * attrs.N]" }], | |
| "passes": [ | |
| { | |
| "id": "partial", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_partial", "b_main", "scales_main", "y_f32"], | |
| "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)", "z": "tiledRegSplitK" } | |
| }, | |
| { | |
| "id": "combine", | |
| "shader": "reduce-axis0-splitk-combine.wgsl.jinja", | |
| "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "addBias": "hasBias" }, | |
| "bindings": ["partials", "bias", "y_y_t", "params_cols"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((aRows * attrs.N), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_vec4_bias_only", | |
| "priority": 16, | |
| "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"], | |
| "derive": { | |
| "aVec4Element": "\"vec4<f16>\" if tensorDtypes.aT == \"float16\" else \"vec4<f32>\"", | |
| "bk": "16 if q2NarrowBkPreferred else 32", | |
| "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", | |
| "tileWorkgroupX": "32 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "tileWorkgroupY": "8 if q2WideMicroPreferred else tiledWorkgroupSide", | |
| "threadRows": "8 if q2WideMicroPreferred else tiledRegVec4ThreadRows", | |
| "threadCols": "2 if q2WideMicroPreferred else 4", | |
| "tileRows": "threadRows * tileWorkgroupY", | |
| "tileCols": "threadCols * tileWorkgroupX", | |
| "tileSubgroupPin": "32 if q2WideMicroPreferred else tiledRegSubgroupPin", | |
| "alignedBlockLoads": true, | |
| "aVec4Loads": true | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_partial", "b_main", "scales_main", "bias", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "ceilDiv(aRows, tileRows)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_reg_bias_only", | |
| "priority": 15, | |
| "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"], | |
| "derive": { | |
| "bk": "tiledRegSelectedBK", | |
| "tileRows": "tiledRegSelectedTileRows", | |
| "tileCols": 64, | |
| "threadRows": "tiledRegSelectedThreadRows", | |
| "threadCols": 4, | |
| "alignedBlockLoads": "aRows >= 128 and attrs.block_size % tiledRegSelectedBK == 0" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "bias", "y_y_t"], | |
| "dispatch": { "x": "dispatchN64", "y": "tiledRegSelectedDispatchM" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tiled_bias_only", | |
| "priority": 14, | |
| "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"], | |
| "derive": {}, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "bias", "y_y_t"], | |
| "dispatch": { "x": "dispatchN32", "y": "dispatchM32" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "prefill_tile4x4_bias_only", | |
| "priority": 13, | |
| "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "attrs.K % attrs.block_size == 0", "aRows >= 2", "portableWorkgroupFits"], | |
| "derive": { | |
| "wordsPerCol": "smallMWordsPerCol", | |
| "wordsPerBlock": "blobWords", | |
| "kLanes": "smallMKLanes", | |
| "colGroups": "smallMColGroups" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "bias", "y_y_t", "params_main"], | |
| "dispatch": { | |
| "x": "min(smallMDispatchN, DISPATCH_FOLD_WIDTH)", | |
| "y": "min(ceilDiv(aRows, 4), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "z": "ceilDiv(smallMDispatchN, DISPATCH_FOLD_WIDTH)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "bias_only", | |
| "priority": 0, | |
| "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "portableWorkgroupFits"], | |
| "derive": {}, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-nbits.wgsl.jinja", | |
| "bindings": ["a_main", "b_main", "scales_main", "bias", "y_y_t", "params_main"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| } | |
| ] | |
| } | |