{ "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 } } ] } ] }