File size: 7,212 Bytes
fb9dcfe 746b1e6 fb9dcfe 746b1e6 fb9dcfe 746b1e6 bb8ee13 746b1e6 bb8ee13 746b1e6 bb8ee13 d8aeb25 bb8ee13 d8aeb25 bb8ee13 d8aeb25 bb8ee13 746b1e6 bb8ee13 746b1e6 d8aeb25 746b1e6 bb8ee13 746b1e6 8386dfb 746b1e6 bb8ee13 d8aeb25 bb8ee13 746b1e6 bb8ee13 746b1e6 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 | ---
library_name: kernels
license: apache-2.0
tags:
- kernel
- webgpu
- wgsl
---
# com.microsoft.FusedMatMul
`com.microsoft` · ONNX Runtime contrib operator · contrib since_version 1
## Description
Matrix product of two N-dimensional tensors `A` and `B`, following NumPy-style matrix-multiplication broadcasting. Supports optional transposition of either operand's last two dimensions, optional batch-dimension transposition, and a scalar `alpha` multiplier. Float32 and float16 are supported; double and bfloat16 are not.
See the [ONNX Runtime `FusedMatMul` contrib-operator spec](https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#com.microsoft.FusedMatMul) for the reference semantics.
## Inputs
| Name | Logical dtype | Rank | Shape | Description | Presence |
| --- | --- | --- | --- | --- | --- |
| `A` | `T` | — | — | N-dimensional matrix A. | required |
| `B` | `T` | — | — | N-dimensional matrix B. | required |
## Outputs
| Name | Logical dtype | Rank | Shape | Description | Presence |
| --- | --- | --- | --- | --- | --- |
| `Y` | `T` | derived | derived | Matrix-multiplication result whose shape follows NumPy-style rules after applying the requested batch and matrix transpositions. | required |
## Attributes
Default values (overridable per request):
| Attribute | Default | Description |
| --- | --- | --- |
| `alpha` | `1` | Scalar multiplier applied to the product of the input tensors. |
| `transA` | `0` | When non-zero, transposes `A` on its last two dimensions before multiplication. |
| `transB` | `0` | When non-zero, transposes `B` on its last two dimensions before multiplication. |
| `transBatchA` | `0` | When non-zero, transposes `A` on its first dimension and batch dimensions (dim-1 to dim-rank-2) before multiplication. |
| `transBatchB` | `0` | When non-zero, transposes `B` on its first dimension and batch dimensions (dim-1 to dim-rank-2) before multiplication. |
## Type constraints
| Variable | Allowed dtypes |
| --- | --- |
| `T` | `float32`, `float16` |
## Implementation variants
One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.
- `broadcast_transb_tiled_reg` — Register-blocked broadcast product with physically transposed B. Reuses the shared batch-addressing tile, keeps f32 accumulation, and preserves scalar K order for f16. Low tile count and excessive padding demote this otherwise correct path.
- `broadcast_transb_subgroup_matrix_f16` — Broadcast transposed-B product using a supported 8x8x8 subgroup-matrix configuration with f32 accumulation. Logical shapes and physical B strides share the existing matrix engine; insufficient output tiles or excessive padding retain the generic tile.
- `broadcast_transb_subgroup_matrix_f32` — Broadcast transposed-B product using a supported 8x8x8 subgroup-matrix configuration with f32 accumulation. Logical shapes and physical B strides share the existing matrix engine; insufficient output tiles or excessive padding retain the generic tile.
- `subgroup_matrix_transbatch_b_f16` — Aligned rank-3 transBatchB product with sufficient reduction depth and output tiles to amortize subgroup-matrix staging.
- `subgroup_matrix_transbatch_b_f32` — Aligned rank-3 transBatchB product with sufficient reduction depth and output tiles to amortize subgroup-matrix staging.
- `m1_gemv_vec4` — Vector-by-matrix specialization for a single output row: each workgroup owns 32 consecutive vec4 column groups and partitions the reduction across the workgroup's second dimension. The accumulator stays float32 for both tensor types.
- `rank2_band_vec4_splitk` — Splits the vec4 band's K axis across up to sixteen workgroups. Each range writes an f32 partial band with alpha applied, and a combine pass sums the partials.
- `rank2_band_vec4` — Few-row band for a rank-2 product without transposes. Each lane owns one vec4 column group and one accumulator per row, so every B word feeds all 2 to 16 rows; alpha is applied at the store.
- `rank2_band_vec4_f32_preferred` — Few-row band for a rank-2 product without transposes. Each lane owns one vec4 column group and one accumulator per row, so every B word feeds all 2 to 16 rows; alpha is applied at the store.
- `subgroup_matrix_splitk` — Partitions the K reduction of small-M rank-2 products across subgroup-matrix workgroups, then combines float32 partials that already include alpha.
- `subgroup_matrix` — Subgroup-matrix `Y = alpha * op(A) @ op(B)` over dense batches with float32 accumulation. An output width that is not a multiple of the 64-wide column tile switches the trailing tile to guarded addressing: its B columns clamp to N - 1 and its stores drop every column at or past N. Yields the shape when the padded column ratio exceeds its tunable ceiling.
- `plain_rank2_tiled_reg` — Register-blocked rank-2 `Y = alpha * A @ B` specialization for non-transposed inputs on tiers without subgroup-matrix support.
- `transbatch_b_tiled_reg` — Register-blocked logical rank3 product with an interleaved physical B batch axis.
## Device requirements
Some implementation variants require `subgroup-matrix` and `subgroups`. These are route-specific capabilities, not package-wide requirements; availability also depends on the request shape and dtype.
## Files
- [`metadata.json`](build/webgpu/metadata.json) — kernel metadata (id, digests, per-variant templates, provenance)
- [`manifest.json`](build/webgpu/manifest.json) — the op contract (source of truth)
- [`test.json`](build/webgpu/test.json) — correctness cases
- [`bench.json`](build/webgpu/bench.json) — benchmark cases
- [`fused-matmul-subgroup-matrix.wgsl.jinja`](build/webgpu/fused-matmul-subgroup-matrix.wgsl.jinja)
- [`matmul-band-vec4.wgsl.jinja`](build/webgpu/matmul-band-vec4.wgsl.jinja)
- [`matmul-subgroup-matrix-ext.wgsl.jinja`](build/webgpu/matmul-subgroup-matrix-ext.wgsl.jinja)
- [`matmul-tiled-general-reg.wgsl.jinja`](build/webgpu/matmul-tiled-general-reg.wgsl.jinja)
- [`matmul-tiled-general.wgsl.jinja`](build/webgpu/matmul-tiled-general.wgsl.jinja)
- [`matmul-vector-matrix-vec4.wgsl.jinja`](build/webgpu/matmul-vector-matrix-vec4.wgsl.jinja)
- [`reduce-axis0-splitk-combine.wgsl.jinja`](build/webgpu/reduce-axis0-splitk-combine.wgsl.jinja)
## Use with `@huggingface/kernels`
```sh
npm install --save-exact @huggingface/kernels@0.0.1-preview.3
```
Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.
The `version: 1` option selects the published kernel contract; it is independent of any operator opset, contrib `since_version`, or model version.
It follows the `v1` branch as fixes land. To pin exact artifact bytes, pass a 40-character commit `revision` instead of `version`.
Replace each `*Data` placeholder with a typed array containing the corresponding input data.
```js
import { getKernel } from "@huggingface/kernels";
const kernel = await getKernel("webgpu-kernels/com.microsoft.FusedMatMul", { version: 1 });
const { Y } = await kernel({ A: { data: AData, shape: [3] }, B: { data: BData, shape: [3] } });
```
|