| { |
| "domain": "com.microsoft", |
| "name": "MatMulBnb4", |
| "sinceVersion": 1, |
| "description": "Computes `A @ dequant(B)^T` where `B` uses bitsandbytes 4-bit quantization: `quant_type = 0` selects FP4 and `quant_type = 1` selects NF4. Supports rank-2 float16/float32 `A`, `transB = 1`, and `training_mode = 0`; rank-1 and rank-3-or-higher `A`, bfloat16, `transB = 0`, and training are not implemented. `B` is the flattened `[N, K]` weight, two codes per byte with the even flat index in the high nibble. Each code indexes a fixed 16-entry codebook, and the value is `codebook[code] * absmax[flat_index / block_size]`.", |
| "inputs": [ |
| { "role": "A", "dtype": "T1", "rank": 2, "description": "Float input matrix of shape `(M, K)`, not quantized." }, |
| { |
| "role": "B", |
| "dtype": "T2", |
| "rank": 1, |
| "description": "The `[N, K]` weight, flattened and quantized to 4 bits, stored as `(N * K + 1) / 2` bytes; the ONNX type is uint8 (this WebGPU implementation reads one widened u32 per stored byte)." |
| }, |
| { |
| "role": "absmax", |
| "dtype": "T1", |
| "rank": 1, |
| "description": "Per-block absolute-maximum dequantization scales of shape `((N * K + block_size - 1) / block_size)`, same dtype as A." |
| } |
| ], |
| "outputs": [ |
| { |
| "role": "Y", |
| "dtype": "T1", |
| "rank": 2, |
| "shape": "[dim(shapes.A, 0), attrs.N]", |
| "description": "Result of `A` multiplied by the dequantized, transposed weight matrix, with shape `(M, N)` and the same dtype as `A`." |
| } |
| ], |
| "attributes": { "training_mode": 0, "transB": 1 }, |
| "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] } |
| }, |
| "attributeDescriptions": { |
| "K": "Input feature count (the shared dimension).", |
| "N": "Output feature count.", |
| "block_size": "Number of weights sharing one absmax scale; a power of two, at least 16.", |
| "quant_type": "Codebook selector: 0 = FP4, 1 = NF4.", |
| "training_mode": "Whether training outputs are requested. This inference-only implementation supports the standard default value 0.", |
| "transB": "Whether the quantized weight is stored transposed. This implementation supports the standard default value 1." |
| }, |
| "typeConstraints": { "T1": ["float32", "float16"], "T2": ["uint8"] }, |
| "args": { |
| "aT": { "kind": "tensor", "semantic": "A", "role": "input" }, |
| "bT": { "kind": "tensor", "semantic": "B", "role": "input" }, |
| "absmaxT": { "kind": "tensor", "semantic": "absmax", "role": "input" }, |
| "yT": { "kind": "tensor", "semantic": "Y", "role": "output" } |
| }, |
| "tunables": { |
| "WORKGROUP_SIZE": 64, |
| "TILE_MIN_M": 16, |
| "PORTABLE_TILE_K": 16, |
| "SGMAT_TILE_ROWS": 64, |
| "SGMAT_TALL_TILE_ROWS": 128, |
| "SGMAT_TALL_MIN_M": 128, |
| "SGMAT_TILE_COLS": 64, |
| "SGMAT_TILE_K": 32 |
| }, |
| "bindingSets": { |
| "main": [ |
| { |
| "name": "a", |
| "arg": "aT", |
| "semantic": "A", |
| "buffer": { "type": "read-only-storage" }, |
| "elementType": "$aScalar" |
| }, |
| { "name": "b", "arg": "bT", "semantic": "B", "buffer": { "type": "read-only-storage" }, "elementType": "u32" }, |
| { |
| "name": "absmax", |
| "arg": "absmaxT", |
| "semantic": "absmax", |
| "buffer": { "type": "read-only-storage" }, |
| "elementType": "$absmaxScalar", |
| "length": "$ABSMAX_LEN" |
| }, |
| { "name": "y", "arg": "yT", "semantic": "Y", "buffer": { "type": "storage" }, "elementType": "$aScalar" }, |
| { |
| "name": "params", |
| "semantic": "kernel.params", |
| "buffer": { "type": "uniform" }, |
| "struct": { |
| "name": "Params", |
| "fields": [ |
| { "name": "rows", "type": "u32", "value": "dim(shapes.A, 0)" }, |
| { "name": "K", "type": "u32", "value": "attrs.K" }, |
| { "name": "N", "type": "u32", "value": "attrs.N" }, |
| { "name": "blockSize", "type": "u32", "value": "attrs.block_size" } |
| ] |
| } |
| } |
| ], |
| "gemv": [ |
| { |
| "name": "a", |
| "arg": "aT", |
| "semantic": "A", |
| "buffer": { "type": "read-only-storage" }, |
| "elementType": "$aScalar" |
| }, |
| { "name": "b", "arg": "bT", "semantic": "B", "buffer": { "type": "read-only-storage" }, "elementType": "u32" }, |
| { |
| "name": "absmax", |
| "arg": "absmaxT", |
| "semantic": "absmax", |
| "buffer": { "type": "read-only-storage" }, |
| "elementType": "$absmaxScalar", |
| "length": "$ABSMAX_LEN" |
| }, |
| { "name": "y", "arg": "yT", "semantic": "Y", "buffer": { "type": "storage" }, "elementType": "$aScalar" }, |
| { |
| "name": "params", |
| "semantic": "kernel.params", |
| "buffer": { "type": "uniform" }, |
| "struct": { |
| "name": "Params", |
| "fields": [ |
| { "name": "K", "type": "u32", "value": "attrs.K" }, |
| { "name": "N", "type": "u32", "value": "attrs.N" }, |
| { "name": "blockSize", "type": "u32", "value": "attrs.block_size" } |
| ] |
| } |
| } |
| ], |
| "sgmat": [ |
| { |
| "name": "a", |
| "arg": "aT", |
| "semantic": "A", |
| "buffer": { "type": "read-only-storage" }, |
| "elementType": "$aScalar" |
| }, |
| { "name": "b", "arg": "bT", "semantic": "B", "buffer": { "type": "read-only-storage" }, "elementType": "u32" }, |
| { |
| "name": "absmax", |
| "arg": "absmaxT", |
| "semantic": "absmax", |
| "buffer": { "type": "read-only-storage" }, |
| "elementType": "$absmaxScalar", |
| "length": "$ABSMAX_LEN" |
| }, |
| { "name": "y", "arg": "yT", "semantic": "Y", "buffer": { "type": "storage" }, "elementType": "$aScalar" } |
| ], |
| "castA": [ |
| { |
| "name": "x", |
| "arg": "aT", |
| "semantic": "A", |
| "buffer": { "type": "read-only-storage" }, |
| "elementType": "$srcScalar" |
| }, |
| { "name": "y", "semantic": "aF32", "buffer": { "type": "storage" }, "elementType": "f32" }, |
| { |
| "name": "params", |
| "semantic": "kernel.params", |
| "buffer": { "type": "uniform" }, |
| "struct": { "name": "Params", "fields": [{ "name": "count", "type": "u32", "value": "numel(shapes.A)" }] } |
| } |
| ], |
| "sgmatWiden": [ |
| { "name": "a", "semantic": "aF32", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, |
| { "name": "b", "arg": "bT", "semantic": "B", "buffer": { "type": "read-only-storage" }, "elementType": "u32" }, |
| { |
| "name": "absmax", |
| "arg": "absmaxT", |
| "semantic": "absmax", |
| "buffer": { "type": "read-only-storage" }, |
| "elementType": "$absmaxScalar", |
| "length": "$ABSMAX_LEN" |
| }, |
| { "name": "y", "semantic": "yF32", "buffer": { "type": "storage" }, "elementType": "f32" } |
| ], |
| "castY": [ |
| { "name": "x", "semantic": "yF32", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, |
| { "name": "y", "arg": "yT", "semantic": "Y", "buffer": { "type": "storage" }, "elementType": "$srcScalar" }, |
| { |
| "name": "params", |
| "semantic": "kernel.params", |
| "buffer": { "type": "uniform" }, |
| "struct": { |
| "name": "Params", |
| "fields": [{ "name": "count", "type": "u32", "value": "dim(shapes.A, 0) * attrs.N" }] |
| } |
| } |
| ] |
| }, |
| "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", |
| "packedBytesExpected": "ceilDiv(attrs.N * attrs.K, 2)", |
| "absmaxCountExpected": "ceilDiv(attrs.N * attrs.K, attrs.block_size)", |
| "aFloatOk": "(tensorDtypes.A == \"float32\" or tensorDtypes.A == \"float16\") and f16Ok(tensorDtypes.A)", |
| "portableWorkgroupSize": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)", |
| "commonShapeValid": "ranks.A == 2 and ranks.B == 1 and ranks.absmax == 1 and ranks.Y == 2 and aFloatOk and tensorDtypes.B == \"uint8\" and tensorDtypes.absmax == tensorDtypes.A and tensorDtypes.Y == tensorDtypes.A and attrs.K > 0 and attrs.N > 0 and attrs.block_size >= 16 and pow2ceil(attrs.block_size) == attrs.block_size and dim(shapes.A, 1) == attrs.K and dim(shapes.B, 0) == packedBytesExpected and dim(shapes.absmax, 0) == absmaxCountExpected and dim(shapes.Y, 0) == dim(shapes.A, 0) and dim(shapes.Y, 1) == attrs.N", |
| "gemvShapeValid": "commonShapeValid and dim(shapes.A, 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.A, 0) >= tunables.TILE_MIN_M and portableTileKValid", |
| "tileWorkgroupStorageBytes": "8 * 64 * 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) <= device.limits.maxComputeWorkgroupsPerDimension and ceilDiv(dim(shapes.A, 0), 64) <= device.limits.maxComputeWorkgroupsPerDimension", |
| "sgmatMatrixSize": "8", |
| "sgmatTileRows": "tunables.SGMAT_TALL_TILE_ROWS if dim(shapes.A, 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": "device.adapterInfo.subgroupMinSize if has(device.adapterInfo, \"subgroupMinSize\") else 1", |
| "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.A, 0), sgmatTileRows)", |
| "sgmatWorkgroupFits": "sgmatWorkgroupSize <= deviceWorkgroupCap and sgmatWorkgroupStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and sgmatDispatchN <= device.limits.maxComputeWorkgroupsPerDimension and sgmatDispatchM <= device.limits.maxComputeWorkgroupsPerDimension", |
| "sgmatWidenBytes": "numel(shapes.A) * 4", |
| "sgmatWidenOutBytes": "dim(shapes.A, 0) * attrs.N * 4", |
| "sgmatWidenFits": "sgmatWidenBytes <= device.limits.maxStorageBufferBindingSize and sgmatWidenBytes <= device.limits.maxBufferSize and sgmatWidenOutBytes <= device.limits.maxStorageBufferBindingSize and sgmatWidenOutBytes <= device.limits.maxBufferSize" |
| }, |
| "constants": { "quantType": "attrs.quant_type", "ABSMAX_LEN": "ceilDiv(attrs.N * attrs.K, attrs.block_size)" }, |
| "variants": [ |
| { |
| "id": "gemv", |
| "priority": 20, |
| "when": ["gemvShapeValid", "portableWorkgroupFits"], |
| "constants": { |
| "workgroupSize": "portableWorkgroupSize", |
| "aScalar": "\"f16\" if tensorDtypes.A == \"float16\" else \"f32\"", |
| "usesF16": "tensorDtypes.A == \"float16\"", |
| "absmaxScalar": "\"f16\" if tensorDtypes.absmax == \"float16\" else \"f32\"" |
| }, |
| "passes": [ |
| { |
| "id": "main", |
| "shader": "matmul-bnb4-gemv.wgsl.jinja", |
| "bindings": "gemv", |
| "dispatch": { "workgroups": "attrs.N" } |
| } |
| ] |
| }, |
| { |
| "id": "sgmat", |
| "priority": 15, |
| "requires": { |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], |
| "limits": { "maxComputeWorkgroupStorageSize": 8192 }, |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] |
| }, |
| "when": ["commonShapeValid", "tensorDtypes.A == \"float32\"", "dim(shapes.A, 0) >= sgmatTileRows", "dim(shapes.A, 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"], |
| "constants": { |
| "K": "attrs.K", |
| "N": "attrs.N", |
| "blockSize": "attrs.block_size", |
| "tileRows": "sgmatTileRows", |
| "tileCols": "tunables.SGMAT_TILE_COLS", |
| "tileK": "tunables.SGMAT_TILE_K", |
| "subRows": "sgmatSubRows", |
| "subCols": "sgmatSubCols", |
| "matrixSize": "sgmatMatrixSize", |
| "rowMatrices": "sgmatSubRows / sgmatMatrixSize", |
| "colMatrices": "sgmatSubCols / sgmatMatrixSize", |
| "loadWidth": "sgmatLoadWidth", |
| "rowSubtiles": "sgmatRowSubtiles", |
| "workgroupSize": "sgmatWorkgroupSize", |
| "bLoadsPerRow": "sgmatBLoadsPerRow", |
| "aScalar": "\"f32\"", |
| "usesF16": false, |
| "absmaxScalar": "\"f16\" if tensorDtypes.absmax == \"float16\" else \"f32\"" |
| }, |
| "passes": [ |
| { |
| "id": "main", |
| "shader": "matmul-bnb4-sgmat.wgsl.jinja", |
| "bindings": "sgmat", |
| "dispatch": { "x": "sgmatDispatchN", "y": "sgmatDispatchM" } |
| } |
| ] |
| }, |
| { |
| "id": "sgmat_widened", |
| "priority": 16, |
| "requires": { |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], |
| "limits": { "maxComputeWorkgroupStorageSize": 8192 }, |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] |
| }, |
| "when": ["commonShapeValid", "tensorDtypes.A == \"float16\"", "dim(shapes.A, 0) >= sgmatTileRows", "dim(shapes.A, 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"], |
| "constants": { |
| "K": "attrs.K", |
| "N": "attrs.N", |
| "blockSize": "attrs.block_size", |
| "tileRows": "sgmatTileRows", |
| "tileCols": "tunables.SGMAT_TILE_COLS", |
| "tileK": "tunables.SGMAT_TILE_K", |
| "subRows": "sgmatSubRows", |
| "subCols": "sgmatSubCols", |
| "matrixSize": "sgmatMatrixSize", |
| "rowMatrices": "sgmatSubRows / sgmatMatrixSize", |
| "colMatrices": "sgmatSubCols / sgmatMatrixSize", |
| "loadWidth": "sgmatLoadWidth", |
| "rowSubtiles": "sgmatRowSubtiles", |
| "workgroupSize": "sgmatWorkgroupSize", |
| "bLoadsPerRow": "sgmatBLoadsPerRow", |
| "aScalar": "\"f32\"", |
| "usesF16": true, |
| "absmaxScalar": "\"f16\" if tensorDtypes.absmax == \"float16\" else \"f32\"", |
| "srcScalar": "\"f16\"", |
| "outScalar": "\"f32\"", |
| "wrapNarrowInt": false, |
| "wrapSigned": false |
| }, |
| "passes": [ |
| { |
| "id": "widen_a", |
| "name": "MatMulBnb4.WidenActivations", |
| "shader": "cast-scalar-x4.wgsl.jinja", |
| "bindings": "castA", |
| "dispatch": { "threads": "ceilDiv(numel(shapes.A), 4)", "workgroupSize": "tunables.WORKGROUP_SIZE" } |
| }, |
| { |
| "id": "main", |
| "name": "MatMulBnb4.SubgroupMatrixWidened", |
| "shader": "matmul-bnb4-sgmat.wgsl.jinja", |
| "bindings": "sgmatWiden", |
| "dispatch": { "x": "sgmatDispatchN", "y": "sgmatDispatchM" } |
| }, |
| { |
| "id": "narrow_y", |
| "name": "MatMulBnb4.NarrowOutput", |
| "shader": "cast-scalar-x4.wgsl.jinja", |
| "bindings": "castY", |
| "constants": { "outScalar": "\"f16\"" }, |
| "dispatch": { |
| "threads": "ceilDiv(dim(shapes.A, 0) * attrs.N, 4)", |
| "workgroupSize": "tunables.WORKGROUP_SIZE" |
| } |
| } |
| ], |
| "description": "Runs the subgroup-matrix tier for a float16 request by widening the activations to float32 either side of the multiply. The matrix units this operator uses accumulate in float32 from float32 operands; the device's float16 configuration returns a float16 result, so feeding them float16 directly would drop the accumulator's precision.", |
| "intermediates": [ |
| { "id": "aF32", "dtype": "float32", "shape": "[numel(shapes.A)]" }, |
| { "id": "yF32", "dtype": "float32", "shape": "[dim(shapes.A, 0) * attrs.N]" } |
| ] |
| }, |
| { |
| "id": "tiled", |
| "priority": 10, |
| "when": ["tileEligible", "tileWorkgroupFits"], |
| "constants": { |
| "tileK": "tunables.PORTABLE_TILE_K", |
| "aScalar": "\"f16\" if tensorDtypes.A == \"float16\" else \"f32\"", |
| "usesF16": "tensorDtypes.A == \"float16\"", |
| "absmaxScalar": "\"f16\" if tensorDtypes.absmax == \"float16\" else \"f32\"" |
| }, |
| "passes": [ |
| { |
| "id": "main", |
| "shader": "matmul-bnb4-tiled.wgsl.jinja", |
| "bindings": "main", |
| "dispatch": { "x": "ceilDiv(attrs.N, 64)", "y": "ceilDiv(dim(shapes.A, 0), 64)" } |
| } |
| ] |
| }, |
| { |
| "id": "scalar", |
| "priority": 0, |
| "when": ["commonShapeValid", "portableWorkgroupSize > 0"], |
| "constants": { |
| "workgroupSize": "portableWorkgroupSize", |
| "aScalar": "\"f16\" if tensorDtypes.A == \"float16\" else \"f32\"", |
| "usesF16": "tensorDtypes.A == \"float16\"", |
| "absmaxScalar": "\"f16\" if tensorDtypes.absmax == \"float16\" else \"f32\"" |
| }, |
| "passes": [ |
| { |
| "id": "main", |
| "shader": "matmul-bnb4.wgsl.jinja", |
| "bindings": "main", |
| "dispatch": { "threads": "numel(shapes.Y)", "workgroupSize": "constants.workgroupSize" } |
| } |
| ] |
| } |
| ] |
| } |
|
|