{ "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", "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", "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", "tiledWorkgroupFits": "16 <= device.limits.maxComputeWorkgroupSizeX and 16 <= device.limits.maxComputeWorkgroupSizeY and 256 <= 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", "tiledRegSelectedTileRows": "64 if largeTiledRegEligible else 32", "tiledRegSelectedThreadRows": "4 if largeTiledRegEligible else 2", "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)", "tiledRegVec4TileRows": "128 if aRows >= 256 else 64", "tiledRegVec4ThreadRows": "8 if aRows >= 256 else 4", "tiledRegVec4DispatchM": "ceilDiv(aRows, tiledRegVec4TileRows)", "tiledRegVec4Eligible": "largeTiledRegEligible and attrs.K % 4 == 0 and attrs.block_size % 32 == 0 and tiledRegVec4DispatchM <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegSplitTiles": "tiledRegVec4DispatchM * dispatchN64", "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 tiledRegSplitK * aRows * attrs.N * 4 <= device.limits.maxStorageBufferBindingSize and tiledRegSplitK <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "B_LEN": "attrs.N * kBlocksExpected * blobSizeExpected / 4", "SCALES_LEN": "attrs.N * kBlocksExpected", "BIAS_LEN": "attrs.N" }, "bindings": { "a_2": { "arg": "aT", "name": "a", "buffer": "read-only-storage", "elementType": "$aElement" }, "b_2": { "arg": "bT", "name": "b", "buffer": "read-only-storage", "elementType": "$bElement", "length": "$B_VEC_LEN" }, "scales_2": { "arg": "scalesT", "name": "scales", "buffer": "read-only-storage", "elementType": "$scaleScalar", "length": "$SCALES_LEN" }, "y_2": { "arg": "yT", "name": "y", "buffer": "storage", "elementType": "$outputScalar" }, "params": { "buffer": "uniform", "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", "buffer": "read-only-storage", "elementType": "$aScalar", "length": "$SCALES_LEN" }, "bias": { "arg": "biasT", "buffer": "read-only-storage", "elementType": "$aScalar", "length": "$BIAS_LEN" }, "a_3": { "arg": "aT", "name": "a", "buffer": "read-only-storage", "elementType": "$aScalar" }, "b_3": { "arg": "bT", "name": "b", "buffer": "read-only-storage", "elementType": "$bScalar", "length": "$B_LEN" }, "a_4": { "arg": "aT", "name": "a", "buffer": "read-only-storage", "elementType": "$aVec4Element" }, "y_3": { "scratch": "partials", "name": "y", "buffer": "storage", "elementType": "f32" }, "partials": { "buffer": "read-only-storage", "elementType": "f32" }, "params_2": { "name": "params", "buffer": "uniform", "struct": [{ "name": "cols", "type": "u32", "value": "aRows * attrs.N" }] }, "params_3": { "name": "params", "buffer": "uniform", "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)", "16 <= device.limits.maxComputeWorkgroupSizeX", "16 <= device.limits.maxComputeWorkgroupSizeY", "64 <= device.limits.maxComputeWorkgroupSizeX", "256 <= device.limits.maxComputeInvocationsPerWorkgroup", "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": { "M": "aRows", "K": "attrs.K", "N": "attrs.N", "kBlocks": "dim(shapes.bT, 1)", "blockSize": "attrs.block_size", "blobSize": "dim(shapes.bT, 2)", "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), (64)), 65535)", "y": "ceilDiv(ceilDiv((aRows * attrs.K / 4), (64)), 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\")", "hasZero": false, "hasBias": false, "workgroupSize": "portableWorkgroupSize", "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aElement": "(\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")", "bElement": "\"vec4\" if gemvVecWords == 4 else \"u32\"", "B_VEC_LEN": "attrs.N * gemvVecPerCol", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "bits": "attrs.bits", "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)", "vecWords": "gemvVecWords", "codesPerVec": "gemvCodesPerVec", "codesPerVec4": "gemvCodesPerVec / 4", "vecPerBlock": "gemvVecPerBlock", "vecPerCol": "gemvVecPerCol", "actVec4": "gemvActVec4", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-gemv-q4.wgsl.jinja", "bindings": ["a_2", "b_2", "scales_2", "y_2", "params"], "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 }, "subgroupCollectivesWidth": "portable" } ] }, { "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": { "hasZero": false, "hasBias": false, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "tileRows": "sgmatTileRows", "workgroupSize": "sgmatWorkgroupSize", "rowSubtiles": "sgmatRowSubtiles", "bLoadsPerRow": "sgmatBLoadsPerRow", "bLoadWidth": "sgmatBLoadWidth", "numSubgroups": "sgmatNumSubgroups" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-sgmat.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "y_2"], "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"], "derive": { "hasZero": false, "hasBias": false, "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aVec4Element": "\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\"", "bScalar": "\"u32\"", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "bk": 32, "tileRows": "tiledRegVec4TileRows", "tileCols": 64, "threadRows": "tiledRegVec4ThreadRows", "threadCols": 4, "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_4", "b_3", "scales_2", "y_3"], "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM", "z": "tiledRegSplitK" } }, { "id": "combine", "shader": "reduce-axis0-splitk-combine.wgsl.jinja", "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "intMode": false, "addBias": false }, "bindings": ["partials", "y_2", "params_2"], "dispatch": { "x": "min(ceilDiv((aRows * attrs.N), (256)), 65535)", "y": "ceilDiv(ceilDiv((aRows * attrs.N), (256)), 65535)", "z": 1 } } ] }, { "id": "prefill_tiled_reg_vec4_default_zero", "priority": 16, "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"], "derive": { "hasZero": false, "hasBias": false, "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aVec4Element": "\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\"", "bScalar": "\"u32\"", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "bk": 32, "tileRows": "tiledRegVec4TileRows", "tileCols": 64, "threadRows": "tiledRegVec4ThreadRows", "threadCols": 4, "alignedBlockLoads": true, "aVec4Loads": true }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", "bindings": ["a_4", "b_3", "scales_2", "y_2"], "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM" } } ] }, { "id": "prefill_tiled_reg_default_zero", "priority": 15, "when": ["commonShapeValid", "defaultEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"], "derive": { "hasZero": false, "hasBias": false, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "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_3", "b_3", "scales_2", "y_2"], "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": { "hasZero": false, "hasBias": false, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "y_2"], "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": { "hasZero": false, "hasBias": false, "workgroupSize": "portableWorkgroupSize", "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\"", "bits": "attrs.bits", "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)", "wordsPerCol": "smallMWordsPerCol", "wordsPerBlock": "blobWords", "kLanes": "smallMKLanes", "colGroups": "smallMColGroups", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "y_2", "params_3"], "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": { "hasZero": false, "hasBias": false, "workgroupSize": "portableWorkgroupSize", "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\"", "bits": "attrs.bits", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "y_2", "params_3"], "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\")", "hasZero": true, "hasBias": true, "workgroupSize": "portableWorkgroupSize", "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aElement": "(\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")", "bElement": "\"vec4\" if gemvVecWords == 4 else \"u32\"", "B_VEC_LEN": "attrs.N * gemvVecPerCol", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "bits": "attrs.bits", "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)", "vecWords": "gemvVecWords", "codesPerVec": "gemvCodesPerVec", "codesPerVec4": "gemvCodesPerVec / 4", "vecPerBlock": "gemvVecPerBlock", "vecPerCol": "gemvVecPerCol", "actVec4": "gemvActVec4", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-gemv-q4.wgsl.jinja", "bindings": ["a_2", "b_2", "scales_2", "zero_points", "bias", "y_2", "params"], "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 }, "subgroupCollectivesWidth": "portable" } ] }, { "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": { "hasZero": true, "hasBias": true, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "tileRows": "sgmatTileRows", "workgroupSize": "sgmatWorkgroupSize", "rowSubtiles": "sgmatRowSubtiles", "bLoadsPerRow": "sgmatBLoadsPerRow", "bLoadWidth": "sgmatBLoadWidth", "numSubgroups": "sgmatNumSubgroups" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-sgmat.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "zero_points", "bias", "y_2"], "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"], "derive": { "hasZero": true, "hasBias": true, "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aVec4Element": "\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\"", "bScalar": "\"u32\"", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "bk": 32, "tileRows": "tiledRegVec4TileRows", "tileCols": 64, "threadRows": "tiledRegVec4ThreadRows", "threadCols": 4, "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_4", "b_3", "scales_2", "zero_points", "y_3"], "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM", "z": "tiledRegSplitK" } }, { "id": "combine", "shader": "reduce-axis0-splitk-combine.wgsl.jinja", "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "intMode": false, "addBias": true }, "bindings": ["partials", "bias", "y_2", "params_2"], "dispatch": { "x": "min(ceilDiv((aRows * attrs.N), (256)), 65535)", "y": "ceilDiv(ceilDiv((aRows * attrs.N), (256)), 65535)", "z": 1 } } ] }, { "id": "prefill_tiled_reg_vec4_zero_bias", "priority": 16, "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"], "derive": { "hasZero": true, "hasBias": true, "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aVec4Element": "\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\"", "bScalar": "\"u32\"", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "bk": 32, "tileRows": "tiledRegVec4TileRows", "tileCols": 64, "threadRows": "tiledRegVec4ThreadRows", "threadCols": 4, "alignedBlockLoads": true, "aVec4Loads": true }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", "bindings": ["a_4", "b_3", "scales_2", "zero_points", "bias", "y_2"], "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM" } } ] }, { "id": "prefill_tiled_reg_zero_bias", "priority": 15, "when": ["commonShapeValid", "zeroBiasEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"], "derive": { "hasZero": true, "hasBias": true, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "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_3", "b_3", "scales_2", "zero_points", "bias", "y_2"], "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": { "hasZero": true, "hasBias": true, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "zero_points", "bias", "y_2"], "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": { "hasZero": true, "hasBias": true, "workgroupSize": "portableWorkgroupSize", "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\"", "bits": "attrs.bits", "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)", "wordsPerCol": "smallMWordsPerCol", "wordsPerBlock": "blobWords", "kLanes": "smallMKLanes", "colGroups": "smallMColGroups", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "zero_points", "bias", "y_2", "params_3"], "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": { "hasZero": true, "hasBias": true, "workgroupSize": "portableWorkgroupSize", "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\"", "bits": "attrs.bits", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "zero_points", "bias", "y_2", "params_3"], "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\")", "hasZero": true, "hasBias": false, "workgroupSize": "portableWorkgroupSize", "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aElement": "(\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")", "bElement": "\"vec4\" if gemvVecWords == 4 else \"u32\"", "B_VEC_LEN": "attrs.N * gemvVecPerCol", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "bits": "attrs.bits", "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)", "vecWords": "gemvVecWords", "codesPerVec": "gemvCodesPerVec", "codesPerVec4": "gemvCodesPerVec / 4", "vecPerBlock": "gemvVecPerBlock", "vecPerCol": "gemvVecPerCol", "actVec4": "gemvActVec4", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-gemv-q4.wgsl.jinja", "bindings": ["a_2", "b_2", "scales_2", "zero_points", "y_2", "params"], "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 }, "subgroupCollectivesWidth": "portable" } ] }, { "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": { "hasZero": true, "hasBias": false, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "tileRows": "sgmatTileRows", "workgroupSize": "sgmatWorkgroupSize", "rowSubtiles": "sgmatRowSubtiles", "bLoadsPerRow": "sgmatBLoadsPerRow", "bLoadWidth": "sgmatBLoadWidth", "numSubgroups": "sgmatNumSubgroups" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-sgmat.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "zero_points", "y_2"], "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"], "derive": { "hasZero": true, "hasBias": false, "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aVec4Element": "\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\"", "bScalar": "\"u32\"", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "bk": 32, "tileRows": "tiledRegVec4TileRows", "tileCols": 64, "threadRows": "tiledRegVec4ThreadRows", "threadCols": 4, "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_4", "b_3", "scales_2", "zero_points", "y_3"], "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM", "z": "tiledRegSplitK" } }, { "id": "combine", "shader": "reduce-axis0-splitk-combine.wgsl.jinja", "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "intMode": false, "addBias": false }, "bindings": ["partials", "y_2", "params_2"], "dispatch": { "x": "min(ceilDiv((aRows * attrs.N), (256)), 65535)", "y": "ceilDiv(ceilDiv((aRows * attrs.N), (256)), 65535)", "z": 1 } } ] }, { "id": "prefill_tiled_reg_vec4_zero_only", "priority": 16, "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"], "derive": { "hasZero": true, "hasBias": false, "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aVec4Element": "\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\"", "bScalar": "\"u32\"", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "bk": 32, "tileRows": "tiledRegVec4TileRows", "tileCols": 64, "threadRows": "tiledRegVec4ThreadRows", "threadCols": 4, "alignedBlockLoads": true, "aVec4Loads": true }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", "bindings": ["a_4", "b_3", "scales_2", "zero_points", "y_2"], "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM" } } ] }, { "id": "prefill_tiled_reg_zero_only", "priority": 15, "when": ["commonShapeValid", "zeroOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"], "derive": { "hasZero": true, "hasBias": false, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "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_3", "b_3", "scales_2", "zero_points", "y_2"], "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": { "hasZero": true, "hasBias": false, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "zero_points", "y_2"], "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": { "hasZero": true, "hasBias": false, "workgroupSize": "portableWorkgroupSize", "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\"", "bits": "attrs.bits", "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)", "wordsPerCol": "smallMWordsPerCol", "wordsPerBlock": "blobWords", "kLanes": "smallMKLanes", "colGroups": "smallMColGroups", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "zero_points", "y_2", "params_3"], "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": { "hasZero": true, "hasBias": false, "workgroupSize": "portableWorkgroupSize", "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\"", "bits": "attrs.bits", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "zero_points", "y_2", "params_3"], "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\")", "hasZero": false, "hasBias": true, "workgroupSize": "portableWorkgroupSize", "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aElement": "(\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\") if gemvActVec4 else (\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\")", "bElement": "\"vec4\" if gemvVecWords == 4 else \"u32\"", "B_VEC_LEN": "attrs.N * gemvVecPerCol", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "bits": "attrs.bits", "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)", "vecWords": "gemvVecWords", "codesPerVec": "gemvCodesPerVec", "codesPerVec4": "gemvCodesPerVec / 4", "vecPerBlock": "gemvVecPerBlock", "vecPerCol": "gemvVecPerCol", "actVec4": "gemvActVec4", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-gemv-q4.wgsl.jinja", "bindings": ["a_2", "b_2", "scales_2", "bias", "y_2", "params"], "dispatch": { "x": "min(gemvDispatchN, 65535)", "y": "ceilDiv(gemvDispatchN, 65535)", "z": 1 }, "subgroupCollectivesWidth": "portable" } ] }, { "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": { "hasZero": false, "hasBias": true, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "tileRows": "sgmatTileRows", "workgroupSize": "sgmatWorkgroupSize", "rowSubtiles": "sgmatRowSubtiles", "bLoadsPerRow": "sgmatBLoadsPerRow", "bLoadWidth": "sgmatBLoadWidth", "numSubgroups": "sgmatNumSubgroups" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-sgmat.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "bias", "y_2"], "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"], "derive": { "hasZero": false, "hasBias": true, "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aVec4Element": "\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\"", "bScalar": "\"u32\"", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "bk": 32, "tileRows": "tiledRegVec4TileRows", "tileCols": 64, "threadRows": "tiledRegVec4ThreadRows", "threadCols": 4, "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_4", "b_3", "scales_2", "y_3"], "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM", "z": "tiledRegSplitK" } }, { "id": "combine", "shader": "reduce-axis0-splitk-combine.wgsl.jinja", "derive": { "op": "\"sum\"", "outputF16": "tensorDtypes.aT == \"float16\"", "intMode": false, "addBias": true }, "bindings": ["partials", "bias", "y_2", "params_2"], "dispatch": { "x": "min(ceilDiv((aRows * attrs.N), (256)), 65535)", "y": "ceilDiv(ceilDiv((aRows * attrs.N), (256)), 65535)", "z": 1 } } ] }, { "id": "prefill_tiled_reg_vec4_bias_only", "priority": 16, "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVec4Eligible"], "derive": { "hasZero": false, "hasBias": true, "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "aVec4Element": "\"vec4\" if tensorDtypes.aT == \"float16\" else \"vec4\"", "bScalar": "\"u32\"", "scaleScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "outputScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "bk": 32, "tileRows": "tiledRegVec4TileRows", "tileCols": 64, "threadRows": "tiledRegVec4ThreadRows", "threadCols": 4, "alignedBlockLoads": true, "aVec4Loads": true }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tiled-reg.wgsl.jinja", "bindings": ["a_4", "b_3", "scales_2", "bias", "y_2"], "dispatch": { "x": "dispatchN64", "y": "tiledRegVec4DispatchM" } } ] }, { "id": "prefill_tiled_reg_bias_only", "priority": 15, "when": ["commonShapeValid", "biasOnlyEpilogue", "bitsSupported", "dispatchN64 <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tiledRegVariantEligible"], "derive": { "hasZero": false, "hasBias": true, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"", "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_3", "b_3", "scales_2", "bias", "y_2"], "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": { "hasZero": false, "hasBias": true, "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\"", "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", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tiled.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "bias", "y_2"], "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": { "hasZero": false, "hasBias": true, "workgroupSize": "portableWorkgroupSize", "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\"", "bits": "attrs.bits", "codeMask": "3 if attrs.bits == 2 else (15 if attrs.bits == 4 else 255)", "wordsPerCol": "smallMWordsPerCol", "wordsPerBlock": "blobWords", "kLanes": "smallMKLanes", "colGroups": "smallMColGroups", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits-q4-prefill-tile4x4.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "bias", "y_2", "params_3"], "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": { "hasZero": false, "hasBias": true, "workgroupSize": "portableWorkgroupSize", "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\"", "bits": "attrs.bits", "defaultZero": "\"2.0\" if attrs.bits == 2 else (\"8.0\" if attrs.bits == 4 else \"128.0\")", "usesF16": "tensorDtypes.aT == \"float16\"" }, "passes": [ { "id": "main", "shader": "matmul-nbits.wgsl.jinja", "bindings": ["a_3", "b_3", "scales_2", "bias", "y_2", "params_3"], "dispatch": { "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", "z": 1 } } ] } ] }