Download build/webgpu/manifest.json from webgpu-kernels/com.microsoft.MatMulBnb4: direct link, hf CLI and curl.
- Browser
- Download file 22.8 kB
-
https://huggingface.co/kernels/webgpu-kernels/com.microsoft.MatMulBnb4/resolve/v1/build/webgpu/manifest.json
- Command line
-
hf download hf://webgpu-kernels/com.microsoft.MatMulBnb4@v1/build/webgpu/manifest.json
-
curl -L -o manifest.json https://huggingface.co/kernels/webgpu-kernels/com.microsoft.MatMulBnb4/resolve/v1/build/webgpu/manifest.json
22.8 kB
| { | |
| "domain": "com.microsoft", | |
| "name": "MatMulBnb4", | |
| "sinceVersion": 1, | |
| "inputs": { | |
| "aT": { "onnx": "A", "dtype": "T1", "rank": 2 }, | |
| "bT": { "onnx": "B", "dtype": "T2", "rank": 1 }, | |
| "absmaxT": { "onnx": "absmax", "dtype": "T1", "rank": 1 } | |
| }, | |
| "outputs": { "yT": { "onnx": "Y", "dtype": "T1", "rank": 2, "shape": "[dim(shapes.aT, 0), attrs.N]" } }, | |
| "attributes": { | |
| "training_mode": { "default": 0 }, | |
| "transB": { "default": 1 }, | |
| "K": {}, | |
| "N": {}, | |
| "block_size": {}, | |
| "quant_type": {} | |
| }, | |
| "attributeConstraints": { | |
| "K": { "required": true }, | |
| "N": { "required": true }, | |
| "block_size": { "required": true }, | |
| "quant_type": { "required": true, "values": [0, 1] }, | |
| "training_mode": { "values": [0] }, | |
| "transB": { "values": [1] } | |
| }, | |
| "typeConstraints": { "T1": ["float32", "float16"], "T2": ["uint8"] }, | |
| "tunables": { | |
| "WORKGROUP_SIZE": { "default": 64 }, | |
| "TILE_MIN_M": { "default": 4 }, | |
| "PORTABLE_TILE_K": { "default": 16 }, | |
| "SGMAT_TILE_ROWS": { "default": 64 }, | |
| "SGMAT_TALL_TILE_ROWS": { "default": 128 }, | |
| "SGMAT_TALL_MIN_M": { "default": 128 }, | |
| "SGMAT_TILE_COLS": { "default": 64 }, | |
| "SGMAT_TILE_K": { "default": 32 } | |
| }, | |
| "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", | |
| "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", | |
| "portableTileRows": "8 * min(8, max(1, ceilDiv(dim(shapes.aT, 0), 8)))", | |
| "packedBytesExpected": "ceilDiv(attrs.N * attrs.K, 2)", | |
| "absmaxCountExpected": "ceilDiv(attrs.N * attrs.K, attrs.block_size)", | |
| "aFloatOk": "(tensorDtypes.aT == \"float32\" or tensorDtypes.aT == \"float16\") and f16Ok(tensorDtypes.aT)", | |
| "portableWorkgroupSize": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)", | |
| "commonShapeValid": "ranks.aT == 2 and ranks.bT == 1 and ranks.absmaxT == 1 and ranks.yT == 2 and aFloatOk and tensorDtypes.bT == \"uint8\" and tensorDtypes.absmaxT == tensorDtypes.aT and tensorDtypes.yT == tensorDtypes.aT and attrs.K > 0 and attrs.N > 0 and attrs.block_size >= 16 and pow2ceil(attrs.block_size) == attrs.block_size and dim(shapes.aT, 1) == attrs.K and dim(shapes.bT, 0) == packedBytesExpected and dim(shapes.absmaxT, 0) == absmaxCountExpected and dim(shapes.yT, 0) == dim(shapes.aT, 0) and dim(shapes.yT, 1) == attrs.N", | |
| "gemvShapeValid": "commonShapeValid and dim(shapes.aT, 0) == 1", | |
| "portableWorkgroupFits": "portableWorkgroupSize > 0 and portableWorkgroupSize * 16 <= device.limits.maxComputeWorkgroupStorageSize", | |
| "portableTileKValid": "tunables.PORTABLE_TILE_K >= 8 and tunables.PORTABLE_TILE_K % 8 == 0", | |
| "tileEligible": "commonShapeValid and dim(shapes.aT, 0) >= tunables.TILE_MIN_M and portableTileKValid", | |
| "tileWorkgroupStorageBytes": "(portableTileRows * dtypeBytes(tensorDtypes.aT) + 64 * 4) * tunables.PORTABLE_TILE_K", | |
| "tileWorkgroupFits": "16 <= device.limits.maxComputeWorkgroupSizeX and 8 <= device.limits.maxComputeWorkgroupSizeY and 128 <= device.limits.maxComputeInvocationsPerWorkgroup and tileWorkgroupStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and ceilDiv(attrs.N, 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.aT, 0), portableTileRows) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "sgmatMatrixSize": "8", | |
| "sgmatTileRows": "tunables.SGMAT_TALL_TILE_ROWS if dim(shapes.aT, 0) >= tunables.SGMAT_TALL_MIN_M else tunables.SGMAT_TILE_ROWS", | |
| "sgmatRowSubtiles": "4", | |
| "sgmatSubRows": "sgmatTileRows / sgmatRowSubtiles", | |
| "sgmatSubCols": "4 * sgmatMatrixSize", | |
| "sgmatLoadWidth": "sgmatMatrixSize", | |
| "sgmatColSubtiles": "tunables.SGMAT_TILE_COLS / sgmatSubCols", | |
| "sgmatNumSubgroups": "sgmatRowSubtiles * sgmatColSubtiles", | |
| "sgmatSubgroupSize": "32", | |
| "sgmatWorkgroupSize": "sgmatNumSubgroups * sgmatSubgroupSize", | |
| "sgmatBLoadsPerRow": "tunables.SGMAT_TILE_K / sgmatLoadWidth", | |
| "sgmatWorkgroupStorageBytes": "4 * tunables.SGMAT_TILE_COLS * tunables.SGMAT_TILE_K", | |
| "sgmatDispatchN": "ceilDiv(attrs.N, tunables.SGMAT_TILE_COLS)", | |
| "sgmatDispatchM": "ceilDiv(dim(shapes.aT, 0), sgmatTileRows)", | |
| "sgmatWorkgroupFits": "sgmatWorkgroupSize <= deviceWorkgroupCap and sgmatWorkgroupStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and sgmatDispatchN <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and sgmatDispatchM <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "sgmatWidenBytes": "numel(shapes.aT) * 4", | |
| "sgmatWidenOutBytes": "dim(shapes.aT, 0) * attrs.N * 4", | |
| "sgmatWidenFits": "sgmatWidenBytes <= device.limits.maxStorageBufferBindingSize and sgmatWidenBytes <= device.limits.maxBufferSize and sgmatWidenOutBytes <= device.limits.maxStorageBufferBindingSize and sgmatWidenOutBytes <= device.limits.maxBufferSize", | |
| "quantType": "attrs.quant_type", | |
| "K": "attrs.K", | |
| "N": "attrs.N", | |
| "blockSize": "attrs.block_size", | |
| "tileCols": "tunables.SGMAT_TILE_COLS", | |
| "tileK": "tunables.SGMAT_TILE_K", | |
| "matrixSize": "sgmatMatrixSize", | |
| "subCols": "sgmatSubCols", | |
| "colMatrices": "sgmatSubCols / sgmatMatrixSize", | |
| "loadWidth": "sgmatLoadWidth", | |
| "rowSubtiles": "sgmatRowSubtiles", | |
| "workgroupSize": "sgmatWorkgroupSize", | |
| "bLoadsPerRow": "sgmatBLoadsPerRow", | |
| "aScalar": "\"f16\" if tensorDtypes.aT == \"float16\" else \"f32\"", | |
| "usesF16": "tensorDtypes.aT == \"float16\"", | |
| "absmaxScalar": "\"f16\" if tensorDtypes.absmaxT == \"float16\" else \"f32\"", | |
| "ABSMAX_LEN": "ceilDiv(attrs.N * attrs.K, attrs.block_size)", | |
| "stagedGeometryValid": "tunables.SGMAT_TILE_COLS >= sgmatSubCols and tunables.SGMAT_TILE_COLS % sgmatSubCols == 0 and tunables.SGMAT_TILE_K >= sgmatMatrixSize and tunables.SGMAT_TILE_K % sgmatMatrixSize == 0", | |
| "stagedRowSubtiles": "max(1, min(sgmatRowSubtiles, pow2ceil(ceilDiv(dim(shapes.aT, 0), sgmatMatrixSize)), floor(deviceWorkgroupCap / (32 * max(1, sgmatColSubtiles))), floor(tunables.SGMAT_TILE_K / sgmatMatrixSize), floor((device.limits.maxComputeWorkgroupStorageSize - sgmatWorkgroupStorageBytes) / (4 * sgmatMatrixSize * max(1, tunables.SGMAT_TILE_K)))))", | |
| "stagedTileRows": "stagedRowSubtiles * sgmatMatrixSize", | |
| "stagedLoadWorkgroupTarget": "max(1, min(deviceWorkgroupCap, tunables.SGMAT_TILE_COLS * tunables.SGMAT_TILE_K / sgmatLoadWidth))", | |
| "stagedColumnMatrices": "min(4, pow2ceil(max(1, ceilDiv(stagedRowSubtiles * tunables.SGMAT_TILE_COLS * 32, stagedLoadWorkgroupTarget * sgmatMatrixSize))))", | |
| "stagedSubCols": "stagedColumnMatrices * sgmatMatrixSize", | |
| "stagedWorkgroupSize": "stagedRowSubtiles * (tunables.SGMAT_TILE_COLS / stagedSubCols) * 32", | |
| "stagedStorageBytes": "4 * (stagedTileRows + tunables.SGMAT_TILE_COLS) * tunables.SGMAT_TILE_K", | |
| "stagedResourcesFit": "stagedWorkgroupSize <= deviceWorkgroupCap and stagedStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and sgmatDispatchN <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.aT, 0), stagedTileRows) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "sgmatGeometryValid": "tunables.SGMAT_TILE_COLS >= sgmatSubCols and tunables.SGMAT_TILE_COLS % sgmatSubCols == 0 and tunables.SGMAT_TILE_K >= sgmatMatrixSize and tunables.SGMAT_TILE_K % sgmatMatrixSize == 0 and sgmatTileRows >= sgmatRowSubtiles * sgmatMatrixSize and sgmatTileRows % (sgmatRowSubtiles * sgmatMatrixSize) == 0", | |
| "prefixTileRows": "tunables.SGMAT_TILE_ROWS", | |
| "sgmatFullRows": "floor(dim(shapes.aT, 0) / max(1, prefixTileRows)) * prefixTileRows", | |
| "sgmatTailRows": "dim(shapes.aT, 0) - sgmatFullRows", | |
| "tailGeometryValid": "tunables.SGMAT_TILE_COLS >= sgmatSubCols and tunables.SGMAT_TILE_COLS % sgmatSubCols == 0 and tunables.SGMAT_TILE_K >= sgmatMatrixSize and tunables.SGMAT_TILE_K % sgmatMatrixSize == 0", | |
| "tailRowSubtiles": "max(1, min(sgmatRowSubtiles, pow2ceil(ceilDiv(sgmatTailRows, sgmatMatrixSize)), floor(deviceWorkgroupCap / (32 * max(1, sgmatColSubtiles))), floor(tunables.SGMAT_TILE_K / sgmatMatrixSize), floor((device.limits.maxComputeWorkgroupStorageSize - sgmatWorkgroupStorageBytes) / (4 * sgmatMatrixSize * max(1, tunables.SGMAT_TILE_K)))))", | |
| "prefixGeometryValid": "tunables.SGMAT_TILE_COLS >= sgmatSubCols and tunables.SGMAT_TILE_COLS % sgmatSubCols == 0 and tunables.SGMAT_TILE_K >= sgmatMatrixSize and tunables.SGMAT_TILE_K % sgmatMatrixSize == 0 and prefixTileRows >= sgmatRowSubtiles * sgmatMatrixSize and prefixTileRows % (sgmatRowSubtiles * sgmatMatrixSize) == 0", | |
| "prefixResourcesFit": "sgmatWorkgroupFits and ceilDiv(sgmatFullRows, max(1, prefixTileRows)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "hybridTailRowSubtiles": "pow2ceil(tailRowSubtiles) / (2 if pow2ceil(tailRowSubtiles) > tailRowSubtiles else 1)", | |
| "hybridTailTileRows": "hybridTailRowSubtiles * sgmatMatrixSize", | |
| "hybridTailSubCols": "sgmatSubCols * hybridTailRowSubtiles / sgmatRowSubtiles", | |
| "hybridStorageBytes": "4 * (tunables.SGMAT_TILE_COLS + hybridTailTileRows) * tunables.SGMAT_TILE_K", | |
| "hybridDispatchRows": "sgmatFullRows / max(1, prefixTileRows) + ceilDiv(sgmatTailRows, hybridTailTileRows)", | |
| "hybridResourcesFit": "prefixResourcesFit and hybridStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and hybridDispatchRows <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)" | |
| }, | |
| "bindings": { | |
| "a": { "arg": "aT", "elementType": "$aScalar" }, | |
| "b": { "arg": "bT", "elementType": "u32" }, | |
| "absmax": { "arg": "absmaxT", "elementType": "$absmaxScalar", "length": "$ABSMAX_LEN" }, | |
| "y": { "arg": "yT", "elementType": "$aScalar" }, | |
| "params": { | |
| "struct": [ | |
| { "name": "rows", "type": "u32", "value": "dim(shapes.aT, 0)" }, | |
| { "name": "K", "type": "u32", "value": "attrs.K" }, | |
| { "name": "N", "type": "u32", "value": "attrs.N" }, | |
| { "name": "blockSize", "type": "u32", "value": "attrs.block_size" } | |
| ] | |
| } | |
| }, | |
| "variants": [ | |
| { | |
| "id": "sgmat_hybrid_rows", | |
| "priority": 21, | |
| "when": ["prefixGeometryValid", "commonShapeValid", "tensorDtypes.aT == \"float32\"", "attrs.N % tunables.SGMAT_TILE_COLS == 0", "attrs.K % tunables.SGMAT_TILE_K == 0", "attrs.K % attrs.block_size == 0", "attrs.block_size % sgmatLoadWidth == 0", "wave32Effective", "prefixResourcesFit", "sgmatFullRows > 0", "sgmatTailRows > 0", "tailGeometryValid", "hybridResourcesFit"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "limits": { "maxComputeWorkgroupStorageSize": 8192 }, | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "tileRows": "prefixTileRows", | |
| "subRows": "prefixTileRows / sgmatRowSubtiles", | |
| "rowMatrices": "prefixTileRows / (sgmatRowSubtiles * sgmatMatrixSize)", | |
| "hybridRows": true, | |
| "rowCount": "dim(shapes.aT, 0)", | |
| "fullDispatchRows": "sgmatFullRows / prefixTileRows", | |
| "rowOffset": "sgmatFullRows" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-bnb4-sgmat.wgsl.jinja", | |
| "bindings": ["a", "b", "absmax", "y"], | |
| "dispatch": { "x": "sgmatDispatchN", "y": "hybridDispatchRows" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "sgmat_hybrid_rows_widened", | |
| "priority": 21, | |
| "when": ["prefixGeometryValid", "commonShapeValid", "tensorDtypes.aT == \"float16\"", "attrs.N % tunables.SGMAT_TILE_COLS == 0", "attrs.K % tunables.SGMAT_TILE_K == 0", "attrs.K % attrs.block_size == 0", "attrs.block_size % sgmatLoadWidth == 0", "wave32Effective", "prefixResourcesFit", "sgmatWidenFits", "sgmatFullRows > 0", "sgmatTailRows > 0", "tailGeometryValid", "hybridResourcesFit"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "limits": { "maxComputeWorkgroupStorageSize": 8192 }, | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "tileRows": "prefixTileRows", | |
| "subRows": "prefixTileRows / sgmatRowSubtiles", | |
| "rowMatrices": "prefixTileRows / (sgmatRowSubtiles * sgmatMatrixSize)", | |
| "aScalar": "\"f32\"", | |
| "srcScalar": "\"f16\"", | |
| "outScalar": "\"f32\"", | |
| "hybridRows": true, | |
| "rowCount": "dim(shapes.aT, 0)", | |
| "fullDispatchRows": "sgmatFullRows / prefixTileRows", | |
| "rowOffset": "sgmatFullRows" | |
| }, | |
| "intermediates": [ | |
| { "id": "aF32", "dtype": "float32", "shape": "[numel(shapes.aT)]" }, | |
| { "id": "yF32", "dtype": "float32", "shape": "[numel(shapes.yT)]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "widen_a", | |
| "name": "MatMulBnb4.WidenActivations", | |
| "shader": "cast-scalar-x4.wgsl.jinja", | |
| "bindings": [ | |
| { "arg": "aT", "name": "x", "elementType": "$srcScalar" }, | |
| { "scratch": "aF32", "name": "y", "elementType": "f32" }, | |
| { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.aT)" }] } | |
| ], | |
| "dispatch": { | |
| "x": "min(ceilDiv((ceilDiv(numel(shapes.aT), 4)), (tunables.WORKGROUP_SIZE)), 65535)", | |
| "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.aT), 4)), (tunables.WORKGROUP_SIZE)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "main", | |
| "name": "MatMulBnb4.SubgroupMatrixWidened", | |
| "shader": "matmul-bnb4-sgmat.wgsl.jinja", | |
| "bindings": [ | |
| { "scratch": "aF32", "name": "a", "buffer": "read-only-storage", "elementType": "f32" }, | |
| "b", | |
| "absmax", | |
| { "scratch": "yF32", "name": "y", "elementType": "f32" } | |
| ], | |
| "dispatch": { "x": "sgmatDispatchN", "y": "hybridDispatchRows" } | |
| }, | |
| { | |
| "id": "narrow_y", | |
| "name": "MatMulBnb4.NarrowOutput", | |
| "shader": "cast-scalar-x4.wgsl.jinja", | |
| "derive": { "outScalar": "\"f16\"" }, | |
| "bindings": [ | |
| { "scratch": "yF32", "name": "x", "buffer": "read-only-storage", "elementType": "f32" }, | |
| { "arg": "yT", "name": "y", "elementType": "$srcScalar" }, | |
| { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.yT)" }] } | |
| ], | |
| "dispatch": { | |
| "x": "min(ceilDiv((ceilDiv(numel(shapes.yT), 4)), (tunables.WORKGROUP_SIZE)), 65535)", | |
| "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.yT), 4)), (tunables.WORKGROUP_SIZE)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "sgmat_staged_rows", | |
| "priority": 19, | |
| "when": ["commonShapeValid", "dim(shapes.aT, 0) > 1", "dim(shapes.aT, 0) < sgmatTileRows or dim(shapes.aT, 0) % sgmatTileRows != 0", "stagedGeometryValid", "attrs.N % tunables.SGMAT_TILE_COLS == 0", "attrs.K % tunables.SGMAT_TILE_K == 0", "attrs.K % attrs.block_size == 0", "attrs.block_size % sgmatLoadWidth == 0", "wave32Effective", "stagedResourcesFit"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "tileRows": "stagedTileRows", | |
| "subRows": "sgmatMatrixSize", | |
| "subCols": "stagedSubCols", | |
| "rowMatrices": 1, | |
| "colMatrices": "stagedColumnMatrices", | |
| "rowSubtiles": "stagedRowSubtiles", | |
| "workgroupSize": "stagedWorkgroupSize", | |
| "stageActivations": true, | |
| "rowCount": "dim(shapes.aT, 0)" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-bnb4-sgmat.wgsl.jinja", | |
| "bindings": ["a", "b", "absmax", "y"], | |
| "dispatch": { "x": "sgmatDispatchN", "y": "ceilDiv(dim(shapes.aT, 0), stagedTileRows)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "gemv", | |
| "priority": 20, | |
| "when": ["gemvShapeValid", "portableWorkgroupFits"], | |
| "derive": { "workgroupSize": "portableWorkgroupSize" }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-bnb4-gemv.wgsl.jinja", | |
| "bindings": [ | |
| "a", | |
| "b", | |
| "absmax", | |
| "y", | |
| { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "K", "type": "u32", "value": "attrs.K" }, | |
| { "name": "N", "type": "u32", "value": "attrs.N" }, | |
| { "name": "blockSize", "type": "u32", "value": "attrs.block_size" } | |
| ] | |
| } | |
| ], | |
| "dispatch": { "x": "min(attrs.N, 65535)", "y": "ceilDiv(attrs.N, 65535)", "z": 1 } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "sgmat", | |
| "priority": 15, | |
| "when": ["sgmatGeometryValid", "commonShapeValid", "tensorDtypes.aT == \"float32\"", "dim(shapes.aT, 0) >= sgmatTileRows", "dim(shapes.aT, 0) % sgmatTileRows == 0", "attrs.N % tunables.SGMAT_TILE_COLS == 0", "attrs.K % tunables.SGMAT_TILE_K == 0", "attrs.K % attrs.block_size == 0", "attrs.block_size % sgmatLoadWidth == 0", "wave32Effective", "sgmatWorkgroupFits"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "limits": { "maxComputeWorkgroupStorageSize": 8192 }, | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "tileRows": "sgmatTileRows", | |
| "subRows": "sgmatSubRows", | |
| "rowMatrices": "sgmatSubRows / sgmatMatrixSize" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-bnb4-sgmat.wgsl.jinja", | |
| "bindings": ["a", "b", "absmax", "y"], | |
| "dispatch": { "x": "sgmatDispatchN", "y": "sgmatDispatchM" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "sgmat_widened", | |
| "priority": 16, | |
| "when": ["sgmatGeometryValid", "commonShapeValid", "tensorDtypes.aT == \"float16\"", "dim(shapes.aT, 0) >= sgmatTileRows", "dim(shapes.aT, 0) % sgmatTileRows == 0", "attrs.N % tunables.SGMAT_TILE_COLS == 0", "attrs.K % tunables.SGMAT_TILE_K == 0", "attrs.K % attrs.block_size == 0", "attrs.block_size % sgmatLoadWidth == 0", "wave32Effective", "sgmatWorkgroupFits", "sgmatWidenFits"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "limits": { "maxComputeWorkgroupStorageSize": 8192 }, | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "tileRows": "sgmatTileRows", | |
| "subRows": "sgmatSubRows", | |
| "rowMatrices": "sgmatSubRows / sgmatMatrixSize", | |
| "aScalar": "\"f32\"", | |
| "srcScalar": "\"f16\"", | |
| "outScalar": "\"f32\"" | |
| }, | |
| "intermediates": [ | |
| { "id": "aF32", "dtype": "float32", "shape": "[numel(shapes.aT)]" }, | |
| { "id": "yF32", "dtype": "float32", "shape": "[dim(shapes.aT, 0) * attrs.N]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "widen_a", | |
| "name": "MatMulBnb4.WidenActivations", | |
| "shader": "cast-scalar-x4.wgsl.jinja", | |
| "bindings": [ | |
| { "arg": "aT", "name": "x", "elementType": "$srcScalar" }, | |
| { "scratch": "aF32", "name": "y", "elementType": "f32" }, | |
| { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.aT)" }] } | |
| ], | |
| "dispatch": { | |
| "x": "min(ceilDiv((ceilDiv(numel(shapes.aT), 4)), (tunables.WORKGROUP_SIZE)), 65535)", | |
| "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.aT), 4)), (tunables.WORKGROUP_SIZE)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "main", | |
| "name": "MatMulBnb4.SubgroupMatrixWidened", | |
| "shader": "matmul-bnb4-sgmat.wgsl.jinja", | |
| "bindings": [ | |
| { "scratch": "aF32", "name": "a", "buffer": "read-only-storage", "elementType": "f32" }, | |
| "b", | |
| "absmax", | |
| { "scratch": "yF32", "name": "y", "elementType": "f32" } | |
| ], | |
| "dispatch": { "x": "sgmatDispatchN", "y": "sgmatDispatchM" } | |
| }, | |
| { | |
| "id": "narrow_y", | |
| "name": "MatMulBnb4.NarrowOutput", | |
| "shader": "cast-scalar-x4.wgsl.jinja", | |
| "derive": { "outScalar": "\"f16\"" }, | |
| "bindings": [ | |
| { "scratch": "yF32", "name": "x", "buffer": "read-only-storage", "elementType": "f32" }, | |
| { "arg": "yT", "name": "y", "elementType": "$srcScalar" }, | |
| { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "dim(shapes.aT, 0) * attrs.N" }] } | |
| ], | |
| "dispatch": { | |
| "x": "min(ceilDiv((ceilDiv(dim(shapes.aT, 0) * attrs.N, 4)), (tunables.WORKGROUP_SIZE)), 65535)", | |
| "y": "ceilDiv(ceilDiv((ceilDiv(dim(shapes.aT, 0) * attrs.N, 4)), (tunables.WORKGROUP_SIZE)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "tiled", | |
| "priority": 10, | |
| "when": ["tileEligible", "tileWorkgroupFits"], | |
| "derive": { "tileK": "tunables.PORTABLE_TILE_K", "tileRows": "portableTileRows" }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-bnb4-tiled.wgsl.jinja", | |
| "bindings": ["a", "b", "absmax", "y", "params"], | |
| "dispatch": { "x": "ceilDiv(attrs.N, 64)", "y": "ceilDiv(dim(shapes.aT, 0), portableTileRows)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "scalar", | |
| "priority": 0, | |
| "when": ["commonShapeValid", "portableWorkgroupSize > 0"], | |
| "derive": { "workgroupSize": "portableWorkgroupSize" }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "shader": "matmul-bnb4.wgsl.jinja", | |
| "bindings": ["a", "b", "absmax", "y", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numel(shapes.yT)), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| } | |
| ] | |
| } | |