Xenova's picture
Xenova HF Staff
sync 6fdf6301e2bb
902be69 verified
Raw History Blame
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
}
}
]
}
]
}