Download build/webgpu/manifest.json from webgpu-kernels/com.microsoft.TransposeMatMul: direct link, hf CLI and curl.
- Browser
- Download file 33.1 kB
-
https://huggingface.co/kernels/webgpu-kernels/com.microsoft.TransposeMatMul/resolve/v1/build/webgpu/manifest.json
- Command line
-
hf download hf://webgpu-kernels/com.microsoft.TransposeMatMul@v1/build/webgpu/manifest.json
-
curl -L -o manifest.json https://huggingface.co/kernels/webgpu-kernels/com.microsoft.TransposeMatMul/resolve/v1/build/webgpu/manifest.json
33.1 kB
| { | |
| "domain": "com.microsoft", | |
| "name": "TransposeMatMul", | |
| "sinceVersion": 1, | |
| "inputs": { "A": { "dtype": "T" }, "B": { "dtype": "T" } }, | |
| "outputs": { | |
| "Y": { | |
| "dtype": "T", | |
| "rank": "max(ranks.A, ranks.B) - (1 if ranks.A == 1 or ranks.B == 1 else 0)", | |
| "shape": "matmulShape(logicalAShape, logicalBShape)" | |
| } | |
| }, | |
| "attributes": { "alpha": { "default": 1 }, "transA": { "default": 0 }, "transB": { "default": 0 } }, | |
| "typeConstraints": { "T": ["float32", "float16"] }, | |
| "tunables": { | |
| "TILED_REG_MIN_WORKGROUPS": { "default": 64 }, | |
| "PLAIN_RANK2_REG_DEEP_K_TILES": { "default": 128 }, | |
| "GEMV_TARGET_BLOCKS": { "default": 512 }, | |
| "SUBGROUP_MATRIX_MIN_M": { "default": 2 }, | |
| "SUBGROUP_MATRIX_SPLITK_TARGET_WGS": { "default": 512 }, | |
| "SUBGROUP_MATRIX_SPLITK_MIN_K": { "default": 1024 }, | |
| "SUBGROUP_MATRIX_SPLITK_MAX_TILES": { "default": 128 }, | |
| "BAND_VEC4_MAX_ROWS": { "default": 16 }, | |
| "BAND_SPLIT_TARGET_WORKGROUPS": { "default": 256 }, | |
| "BAND_SPLIT_MAX_COLUMN_GROUPS": { "default": 24 }, | |
| "BAND_SPLIT_SLICES": { "default": 8 }, | |
| "BAND_PREFER_MAX_ROWS": { "default": 8 }, | |
| "BAND_PREFER_DEEP_K": { "default": 4096 }, | |
| "BROADCAST_TRANSB_MIN_WORKGROUPS": { "default": 64 }, | |
| "BROADCAST_TRANSB_MAX_PADDING_RATIO": { "default": 2 } | |
| }, | |
| "derive": { | |
| "batchMovedAShape": "shapes.A", | |
| "batchMovedBShape": "shapes.B", | |
| "logicalAShape": "moveAxis(batchMovedAShape, -1, -2) if attrs.transA != 0 and ranks.A > 1 else batchMovedAShape", | |
| "logicalBShape": "moveAxis(batchMovedBShape, -1, -2) if attrs.transB != 0 and ranks.B > 1 else batchMovedBShape", | |
| "gemvN": "dim(shapes.B, 1)", | |
| "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", | |
| "variableSubgroup16To32": "device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 16 and device.adapterInfo.subgroupMaxSize == 32", | |
| "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)", | |
| "rank2DeepPortableTier": "variableSubgroup16To32 or (has(device.adapterInfo, \"architecture\") and (device.adapterInfo.architecture == \"pascal\" or (not device.features.has(\"subgroups\") and (device.adapterInfo.architecture == \"apple\" or device.adapterInfo.architecture == \"gen-9\"))))", | |
| "broadcastTransbM": "dim(shapes.A, ranks.A - 2)", | |
| "broadcastTransbN": "dim(shapes.B, ranks.B - 2)", | |
| "broadcastTransbK": "dim(shapes.A, ranks.A - 1)", | |
| "broadcastTransbBatches": "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 2))", | |
| "gemvLanes": 32, | |
| "vec4OutputTile": "4 * gemvLanes", | |
| "gemvWorkgroups": "ceilDiv(gemvN, vec4OutputTile)", | |
| "gemvSliceCap": "min(32, device.limits.maxComputeWorkgroupSizeY, floor(device.limits.maxComputeInvocationsPerWorkgroup / gemvLanes), floor(device.limits.maxComputeWorkgroupStorageSize / (16 * gemvLanes)))", | |
| "gemvSlices": "max(1, min(gemvSliceCap, max(8, pow2ceil(ceilDiv(tunables.GEMV_TARGET_BLOCKS, gemvWorkgroups)))))", | |
| "gemvResourcesFit": "gemvLanes <= device.limits.maxComputeWorkgroupSizeX and gemvSlices <= device.limits.maxComputeWorkgroupSizeY and gemvLanes * gemvSlices <= device.limits.maxComputeInvocationsPerWorkgroup and 16 * gemvLanes * gemvSlices <= device.limits.maxComputeWorkgroupStorageSize", | |
| "registerTile": 64, | |
| "generalTile": 32, | |
| "tiledRegResourcesFit": "registerTile / 4 <= device.limits.maxComputeWorkgroupSizeX and registerTile / 4 <= device.limits.maxComputeWorkgroupSizeY and registerTile * registerTile / 16 <= device.limits.maxComputeInvocationsPerWorkgroup and 32 * registerTile * dtypeBytes(dtypes.T) <= device.limits.maxComputeWorkgroupStorageSize", | |
| "plainRank2RegDeepPreferredTier": "rank2DeepPortableTier and dim(shapes.A, 1) >= tunables.PLAIN_RANK2_REG_DEEP_K_TILES * 16 and dim(shapes.A, 1) % 16 == 0", | |
| "subgroupMatrixResourcesFit": "128 <= deviceWorkgroupCap and ((32 * 32 + 64 * 32) * dtypeBytes(dtypes.T) + 4 * 4 * 64 * 4) <= device.limits.maxComputeWorkgroupStorageSize", | |
| "sgmatSplitKDepth": "dim(shapes.A, ranks.A - 1)", | |
| "sgmatSplitK32Ok": "sgmatSplitKDepth % 1024 == 0", | |
| "sgmatSplitK16Ok": "sgmatSplitKDepth % 512 == 0", | |
| "sgmatSplitK8Ok": "sgmatSplitKDepth % 256 == 0", | |
| "sgmatSplitK4Ok": "sgmatSplitKDepth % 128 == 0", | |
| "sgmatSplitK2Ok": "sgmatSplitKDepth % 64 == 0", | |
| "bandSplitWant": "pow2ceil(ceilDiv(tunables.BAND_SPLIT_TARGET_WORKGROUPS, max(1, gemvWorkgroups)))", | |
| "bandSplitK": "16 if (bandSplitWant >= 16 and dim(shapes.A, ranks.A - 1) >= 4096) else (8 if (bandSplitWant >= 8 and dim(shapes.A, ranks.A - 1) >= 2048) else (4 if (bandSplitWant >= 4 and dim(shapes.A, ranks.A - 1) >= 1024) else (2 if (bandSplitWant >= 2 and dim(shapes.A, ranks.A - 1) >= 512) else 1)))", | |
| "scalar": "dtypes.T", | |
| "alpha": "attrs.alpha", | |
| "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"", | |
| "fScalar": "\"f16\" if dtypes.T == \"f16\" else \"f32\"", | |
| "aRank": "ranks.A", | |
| "bRank": "ranks.B", | |
| "fusedSgmatRank2Ok": "ranks.A == 2 and ranks.B == 2 and ranks.Y == 2 and attrs.transA == 0 and attrs.transB == 0 and dim(shapes.A, 1) == dim(shapes.B, 0) and dim(shapes.Y, 0) == dim(shapes.A, 0) and dim(shapes.Y, 1) == dim(shapes.B, 1)", | |
| "sgmatOutTiles": "ceilDiv(dim(shapes.A, 0), 32) * ceilDiv(dim(shapes.B, 1), 64) if fusedSgmatRank2Ok else 1", | |
| "sgmatSplitKWant": "ceilDiv(tunables.SUBGROUP_MATRIX_SPLITK_TARGET_WGS, sgmatOutTiles)", | |
| "sgmatSplitK": "32 if (sgmatSplitKWant > 16 and sgmatSplitK32Ok) else (16 if (sgmatSplitKWant > 8 and sgmatSplitK16Ok) else (8 if (sgmatSplitKWant > 4 and sgmatSplitK8Ok) else (4 if (sgmatSplitKWant > 2 and sgmatSplitK4Ok) else (2 if sgmatSplitK2Ok else 1))))", | |
| "bandRank2Ok": "ranks.A == 2 and ranks.B == 2 and ranks.Y == 2 and attrs.transA == 0 and attrs.transB == 0 and dim(shapes.A, 1) == dim(shapes.B, 0) and dim(shapes.Y, 0) == dim(shapes.A, 0) and dim(shapes.Y, 1) == dim(shapes.B, 1)", | |
| "broadcastTransbContract": "(f16Ok(dtypes.T)) and (attrs.transA == 0 and attrs.transB != 0) and (ranks.A > ranks.B and ranks.B >= 2) and (ranks.Y == ranks.A) and (sameShape(shapes.Y, matmulShape(logicalAShape, logicalBShape))) and (dim(shapes.A, ranks.A - 1) == dim(shapes.B, ranks.B - 1)) and (dim(shapes.A, ranks.A - 2) >= 64) and (dim(shapes.A, ranks.A - 1) >= 32) and (dim(shapes.B, ranks.B - 2) >= 64) and (numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 2)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535))" | |
| }, | |
| "bindings": { | |
| "a": { "arg": "A", "elementType": "$scalar" }, | |
| "b": { "arg": "B", "elementType": "$vectorScalar" }, | |
| "partials": { "buffer": "read-only-storage", "elementType": "f32" }, | |
| "y": { "arg": "Y", "elementType": "$scalar" }, | |
| "params": { "struct": [{ "name": "cols", "type": "u32", "value": "numel(shapes.Y)" }] }, | |
| "b_scalar": { "arg": "B", "name": "b", "elementType": "$scalar" }, | |
| "params_rows": { "name": "params", "struct": [{ "name": "M", "type": "u32", "value": "rowCount" }] } | |
| }, | |
| "variants": [ | |
| { | |
| "id": "broadcast_transb_tiled_reg", | |
| "priority": 6, | |
| "when": ["broadcastTransbContract", "tiledRegResourcesFit", "ceilDiv(dim(shapes.A, ranks.A - 2), registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.B, ranks.B - 2), registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"], | |
| "derive": { | |
| "bShape": "logicalBShape", | |
| "bTransposed": true, | |
| "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else broadcastTransbM", | |
| "N": "broadcastTransbN", | |
| "batchCount": "broadcastTransbBatches if ranks.B > 2 else 1", | |
| "transBatchA": false, | |
| "regSequentialK": "dtypes.T == \"f16\"" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "TransposeMatMul.BroadcastTransBTiledReg", | |
| "shader": "matmul-tiled-general-reg.wgsl.jinja", | |
| "derive": { | |
| "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []", | |
| "aRank": "ranks.A if ranks.B > 2 else 2", | |
| "K": "dim(shapes.A, ranks.A - 1)" | |
| }, | |
| "bindings": ["a", "b_scalar", "y", "params_rows"], | |
| "dispatch": { "x": "ceilDiv(N, registerTile)", "y": "ceilDiv(rowCount, registerTile)", "z": "batchCount" } | |
| } | |
| ], | |
| "demoteWhen": ["broadcastTransbBatches * ceilDiv(broadcastTransbM, registerTile) * ceilDiv(broadcastTransbN, registerTile) < tunables.BROADCAST_TRANSB_MIN_WORKGROUPS", "ceilDiv(broadcastTransbM, registerTile) * registerTile * ceilDiv(broadcastTransbN, registerTile) * registerTile * ceilDiv(broadcastTransbK,16) * 16 > tunables.BROADCAST_TRANSB_MAX_PADDING_RATIO * broadcastTransbM * broadcastTransbN * broadcastTransbK"] | |
| }, | |
| { | |
| "id": "broadcast_transb_subgroup_matrix_f16", | |
| "priority": 11, | |
| "when": ["broadcastTransbContract", "dtypes.T == \"f16\"", "wave32Effective", "subgroupMatrixResourcesFit", "ceilDiv(dim(shapes.A, ranks.A - 2),32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.B, ranks.B - 2),64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f16", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "bShape": "logicalBShape", | |
| "bTransposed": true, | |
| "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else broadcastTransbM", | |
| "K": "broadcastTransbK", | |
| "N": "broadcastTransbN", | |
| "batchCount": "broadcastTransbBatches if ranks.B > 2 else 1", | |
| "hasBias": false, | |
| "fScalar": "dtypes.T", | |
| "outScalar": "dtypes.T", | |
| "generalAddressing": true, | |
| "tailSafe": "dim(shapes.A, ranks.A - 1) % 32 != 0 or dim(shapes.B, ranks.B - 2) % 64 != 0", | |
| "outputBuffer": "\"y\"" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "TransposeMatMul.BroadcastTransBSubgroupMatrix", | |
| "shader": "matmul-subgroup-matrix-ext.wgsl.jinja", | |
| "derive": { | |
| "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []", | |
| "aRank": "ranks.A if ranks.B > 2 else 2" | |
| }, | |
| "bindings": ["a", "b_scalar", "y", "params_rows"], | |
| "dispatch": { "x": "ceilDiv(N,64)", "y": "ceilDiv(rowCount,32)", "z": "batchCount" } | |
| } | |
| ], | |
| "demoteWhen": ["broadcastTransbBatches * ceilDiv(broadcastTransbM,32) * ceilDiv(broadcastTransbN,64) < tunables.BROADCAST_TRANSB_MIN_WORKGROUPS", "ceilDiv(broadcastTransbM,32) * 32 * ceilDiv(broadcastTransbN,64) * 64 * ceilDiv(broadcastTransbK,32) * 32 > tunables.BROADCAST_TRANSB_MAX_PADDING_RATIO * broadcastTransbM * broadcastTransbN * broadcastTransbK"] | |
| }, | |
| { | |
| "id": "broadcast_transb_subgroup_matrix_f32", | |
| "priority": 11, | |
| "when": ["broadcastTransbContract", "dtypes.T == \"f32\"", "wave32Effective", "subgroupMatrixResourcesFit", "ceilDiv(dim(shapes.A, ranks.A - 2),32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.B, ranks.B - 2),64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "bShape": "logicalBShape", | |
| "bTransposed": true, | |
| "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else broadcastTransbM", | |
| "K": "broadcastTransbK", | |
| "N": "broadcastTransbN", | |
| "batchCount": "broadcastTransbBatches if ranks.B > 2 else 1", | |
| "hasBias": false, | |
| "fScalar": "dtypes.T", | |
| "outScalar": "dtypes.T", | |
| "generalAddressing": true, | |
| "tailSafe": "dim(shapes.A, ranks.A - 1) % 32 != 0 or dim(shapes.B, ranks.B - 2) % 64 != 0", | |
| "outputBuffer": "\"y\"" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "TransposeMatMul.BroadcastTransBSubgroupMatrix", | |
| "shader": "matmul-subgroup-matrix-ext.wgsl.jinja", | |
| "derive": { | |
| "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []", | |
| "aRank": "ranks.A if ranks.B > 2 else 2" | |
| }, | |
| "bindings": ["a", "b_scalar", "y", "params_rows"], | |
| "dispatch": { "x": "ceilDiv(N,64)", "y": "ceilDiv(rowCount,32)", "z": "batchCount" } | |
| } | |
| ], | |
| "demoteWhen": ["broadcastTransbBatches * ceilDiv(broadcastTransbM,32) * ceilDiv(broadcastTransbN,64) < tunables.BROADCAST_TRANSB_MIN_WORKGROUPS", "ceilDiv(broadcastTransbM,32) * 32 * ceilDiv(broadcastTransbN,64) * 64 * ceilDiv(broadcastTransbK,32) * 32 > tunables.BROADCAST_TRANSB_MAX_PADDING_RATIO * broadcastTransbM * broadcastTransbN * broadcastTransbK"] | |
| }, | |
| { | |
| "id": "m1_gemv_vec4", | |
| "priority": 30, | |
| "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\")", "attrs.transA == 0", "attrs.transB == 0", "ranks.A == 2", "ranks.B == 2", "ranks.Y == 2", "dim(shapes.A, 0) == 1", "dim(shapes.Y, 0) == 1", "dim(shapes.A, 1) == dim(shapes.B, 0)", "dim(shapes.Y, 1) == dim(shapes.B, 1)", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "gemvWorkgroups <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "f16Ok(dtypes.T)"], | |
| "derive": { | |
| "unrollK2": "dtypes.T == \"f16\"", | |
| "gemvScalar": "dtypes.T", | |
| "gemvVector": "\"vec4<\" ~ dtypes.T ~ \">\"", | |
| "alphaScale": "attrs.alpha" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "TransposeMatMul.M1GemvVec4", | |
| "shader": "matmul-vector-matrix-vec4.wgsl.jinja", | |
| "bindings": [ | |
| { "arg": "A", "name": "a", "elementType": "$gemvScalar" }, | |
| { "arg": "B", "name": "b", "elementType": "$gemvVector" }, | |
| { "arg": "Y", "name": "c", "elementType": "$gemvVector" }, | |
| { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "K", "type": "u32", "value": "dim(shapes.A, 1)" }, | |
| { "name": "N4", "type": "u32", "value": "dim(shapes.B, 1) / 4" } | |
| ] | |
| } | |
| ], | |
| "dispatch": { "x": "gemvWorkgroups" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "rank2_band_vec4_splitk", | |
| "priority": 11, | |
| "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T)", "bandRank2Ok", "dim(shapes.A, 0) >= 2", "dim(shapes.A, 0) <= tunables.BAND_VEC4_MAX_ROWS", "dim(shapes.A, 1) > 0", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "gemvWorkgroups <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "gemvLanes <= device.limits.maxComputeWorkgroupSizeX", "gemvWorkgroups <= tunables.BAND_SPLIT_MAX_COLUMN_GROUPS", "bandSplitK >= 2", "bandSplitK * numel(shapes.Y) * 4 <= device.limits.maxStorageBufferBindingSize", "bandSplitK <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "tunables.BAND_SPLIT_SLICES <= device.limits.maxComputeWorkgroupSizeY", "gemvLanes * tunables.BAND_SPLIT_SLICES <= device.limits.maxComputeInvocationsPerWorkgroup"], | |
| "derive": { | |
| "batched": false, | |
| "outputBuffer": "\"y\"", | |
| "M": "dim(shapes.A, 0)", | |
| "K": "dim(shapes.A, 1)", | |
| "N": "dim(shapes.B, 1)", | |
| "gemvSlices": "tunables.BAND_SPLIT_SLICES", | |
| "kSplits": "bandSplitK", | |
| "split": "bandSplitK", | |
| "workgroupSize": 256 | |
| }, | |
| "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[bandSplitK * numel(shapes.Y)]" }], | |
| "passes": [ | |
| { | |
| "id": "partial", | |
| "name": "TransposeMatMul.Rank2BandVec4SplitK", | |
| "shader": "matmul-band-vec4.wgsl.jinja", | |
| "bindings": ["a", "b", { "scratch": "partials", "name": "y", "elementType": "vec4<f32>" }], | |
| "dispatch": { "x": "gemvWorkgroups", "y": "bandSplitK" } | |
| }, | |
| { | |
| "id": "combine", | |
| "name": "TransposeMatMul.Rank2BandVec4SplitKCombine", | |
| "shader": "reduce-axis0-splitk-combine.wgsl.jinja", | |
| "derive": { "op": "\"sum\"", "outputF16": "dtypes.T == \"f16\"" }, | |
| "bindings": ["partials", "y", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numel(shapes.Y)), (256)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numel(shapes.Y)), (256)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "rank2_band_vec4", | |
| "priority": 11, | |
| "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T)", "bandRank2Ok", "dim(shapes.A, 0) >= 2", "dim(shapes.A, 0) <= tunables.BAND_VEC4_MAX_ROWS", "dim(shapes.A, 1) > 0", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "gemvWorkgroups <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "gemvResourcesFit", "not (gemvWorkgroups <= tunables.BAND_SPLIT_MAX_COLUMN_GROUPS and bandSplitK >= 2)"], | |
| "derive": { | |
| "batched": false, | |
| "outputBuffer": "\"y\"", | |
| "M": "dim(shapes.A, 0)", | |
| "K": "dim(shapes.A, 1)", | |
| "N": "dim(shapes.B, 1)" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "TransposeMatMul.Rank2BandVec4", | |
| "shader": "matmul-band-vec4.wgsl.jinja", | |
| "bindings": ["a", "b", { "arg": "Y", "name": "y", "elementType": "$vectorScalar" }], | |
| "dispatch": { "x": "gemvWorkgroups" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "rank2_band_vec4_f32_preferred", | |
| "priority": 13, | |
| "when": ["(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T)", "bandRank2Ok", "dim(shapes.A, 0) >= 2", "dim(shapes.A, 0) <= tunables.BAND_VEC4_MAX_ROWS", "dim(shapes.A, 1) > 0", "dim(shapes.B, 1) > 0", "dim(shapes.B, 1) % 4 == 0", "gemvWorkgroups <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "gemvResourcesFit", "not (gemvWorkgroups <= tunables.BAND_SPLIT_MAX_COLUMN_GROUPS and bandSplitK >= 2)"], | |
| "demoteWhen": ["dtypes.T != \"f32\" or (dim(shapes.A, 0) > tunables.BAND_PREFER_MAX_ROWS and dim(shapes.A, 1) >= tunables.BAND_PREFER_DEEP_K)"], | |
| "derive": { | |
| "batched": false, | |
| "outputBuffer": "\"y\"", | |
| "M": "dim(shapes.A, 0)", | |
| "K": "dim(shapes.A, 1)", | |
| "N": "dim(shapes.B, 1)" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "TransposeMatMul.Rank2BandVec4", | |
| "shader": "matmul-band-vec4.wgsl.jinja", | |
| "bindings": ["a", "b", { "arg": "Y", "name": "y", "elementType": "$vectorScalar" }], | |
| "dispatch": { "x": "gemvWorkgroups" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "subgroup_matrix_splitk", | |
| "priority": 12, | |
| "when": ["(dtypes.T == \"f16\" or dtypes.T == \"f32\") and f16Ok(dtypes.T)", "fusedSgmatRank2Ok", "dim(shapes.A, 0) >= tunables.SUBGROUP_MATRIX_MIN_M", "dim(shapes.A, 1) >= tunables.SUBGROUP_MATRIX_SPLITK_MIN_K", "dim(shapes.B, 1) % 64 == 0", "sgmatSplitK >= 2", "sgmatOutTiles < tunables.SUBGROUP_MATRIX_SPLITK_MAX_TILES", "sgmatSplitK * numel(shapes.Y) * 4 <= device.limits.maxStorageBufferBindingSize", "sgmatSplitK <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.Y, 1), 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(dim(shapes.Y, 0), 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "subgroupMatrixResourcesFit", "wave32Effective"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [ | |
| { "componentType": "f16", "M": 8, "N": 8, "K": 8 }, | |
| { "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 } | |
| ] | |
| }, | |
| "derive": { | |
| "hasBias": false, | |
| "generalAddressing": true, | |
| "tailSafe": false, | |
| "outputBuffer": "\"partials\"", | |
| "outScalar": "\"f32\"", | |
| "rowCount": "dim(shapes.A, 0)", | |
| "K": "dim(shapes.A, 1)", | |
| "N": "dim(shapes.B, 1)", | |
| "batchCount": 1, | |
| "splitK": "sgmatSplitK", | |
| "kPerSplit": "dim(shapes.A, 1) / sgmatSplitK", | |
| "split": "sgmatSplitK", | |
| "workgroupSize": 256 | |
| }, | |
| "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[sgmatSplitK * numel(shapes.Y)]" }], | |
| "passes": [ | |
| { | |
| "id": "partial", | |
| "name": "TransposeMatMul.SubgroupMatrixSplitK", | |
| "shader": "matmul-subgroup-matrix-ext.wgsl.jinja", | |
| "derive": { "bShape": ["dim(shapes.B, 0)", "dim(shapes.B, 1)"], "aRank": 2, "bRank": 2 }, | |
| "bindings": ["a", "b_scalar", { "name": "partials", "elementType": "f32" }, "params_rows"], | |
| "dispatch": { "x": "ceilDiv(dim(shapes.Y, 1), 64)", "y": "ceilDiv(dim(shapes.Y, 0), 32)", "z": "sgmatSplitK" } | |
| }, | |
| { | |
| "id": "combine", | |
| "name": "TransposeMatMul.SubgroupMatrixSplitKCombine", | |
| "shader": "reduce-axis0-splitk-combine.wgsl.jinja", | |
| "derive": { "op": "\"sum\"", "outputF16": "dtypes.T == \"f16\"" }, | |
| "bindings": ["partials", "y", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numel(shapes.Y)), (256)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numel(shapes.Y)), (256)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "subgroup_matrix_tail_broadcast", | |
| "priority": 11, | |
| "when": ["dtypes.T == \"f16\"", "f16Ok(dtypes.T)", "attrs.transA == 0", "attrs.transB == 0", "(((ranks.A == 2 and ranks.B == 2 and ranks.Y == 2) or (ranks.A == 4 and ranks.B == 3 and ranks.Y == 4 and dim(shapes.Y, 0) == dim(shapes.A, 0) and (dim(shapes.A, 1) == dim(shapes.B, 0) or dim(shapes.A, 1) == 1 or dim(shapes.B, 0) == 1) and dim(shapes.Y, 1) == max(dim(shapes.A, 1), dim(shapes.B, 0)))) or (ranks.A == 4 and ranks.B == 2 and ranks.Y == 4 and sameShape(prefix(shapes.Y, 2), prefix(shapes.A, 2))))", "dim(shapes.A, ranks.A - 1) == dim(shapes.B, ranks.B - 2)", "dim(shapes.A, ranks.A - 2) >= tunables.SUBGROUP_MATRIX_MIN_M", "dim(shapes.A, ranks.A - 1) >= 32", "dim(shapes.B, ranks.B - 1) >= 64", "dim(shapes.Y, ranks.Y - 2) == dim(shapes.A, ranks.A - 2)", "dim(shapes.Y, ranks.Y - 1) == dim(shapes.B, ranks.B - 1)", "ceil(dim(shapes.B, ranks.B - 1) / 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil(dim(shapes.A, ranks.A - 2) / 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 1)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "wave32Effective"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f16", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "hasBias": false, | |
| "fScalar": "\"f16\"", | |
| "outScalar": "\"f16\"", | |
| "generalAddressing": true, | |
| "tailSafe": "dim(shapes.A, ranks.A - 1) % 32 != 0 or dim(shapes.B, ranks.B - 1) % 64 != 0", | |
| "outputBuffer": "\"y\"", | |
| "rowCount": "outer(shapes.A, ranks.A - 1) if ranks.B == 2 else dim(shapes.A, ranks.A - 2)", | |
| "K": "dim(shapes.A, ranks.A - 1)", | |
| "N": "dim(shapes.B, ranks.B - 1)", | |
| "batchCount": "numel(shapes.Y) / (dim(shapes.A, ranks.A - 2) * dim(shapes.B, ranks.B - 1)) if ranks.B > 2 else 1" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "TransposeMatMul.SubgroupMatrixTailBroadcast", | |
| "shader": "matmul-subgroup-matrix-ext.wgsl.jinja", | |
| "derive": { | |
| "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []", | |
| "aRank": "ranks.A if ranks.B > 2 else 2", | |
| "bShape": "shapes.B" | |
| }, | |
| "bindings": ["a", "b_scalar", "y", "params_rows"], | |
| "dispatch": { "x": "ceil(N / 64)", "y": "ceil(rowCount / 32)", "z": "batchCount" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "subgroup_matrix", | |
| "priority": 10, | |
| "when": ["f16Ok(dtypes.T)", "ranks.A >= 2", "ranks.B == ranks.A", "ranks.Y == ranks.A", "(dim(shapes.A, ranks.A - 2) if attrs.transA != 0 else dim(shapes.A, ranks.A - 1)) == (dim(shapes.B, ranks.B - 1) if attrs.transB != 0 else dim(shapes.B, ranks.B - 2))", "(dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2)) >= tunables.SUBGROUP_MATRIX_MIN_M", "(dim(shapes.A, ranks.A - 2) if attrs.transA != 0 else dim(shapes.A, ranks.A - 1)) % 32 == 0", "(dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1)) % 64 == 0", "(ranks.A == 2 or (ranks.A == 3 and dim(shapes.A, 0) == dim(shapes.B, 0) and dim(shapes.Y, 0) == dim(shapes.B, 0)) or (ranks.A == 4 and dim(shapes.A, 0) == dim(shapes.B, 0) and dim(shapes.A, 1) == dim(shapes.B, 1) and dim(shapes.Y, 0) == dim(shapes.A, 0) and dim(shapes.Y, 1) == dim(shapes.A, 1)))", "dim(shapes.Y, ranks.Y - 2) == (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2))", "dim(shapes.Y, ranks.Y - 1) == (dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1))", "ceil((dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1)) / 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil((dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2)) / 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "numel(shapes.Y) / ((dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2)) * (dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1))) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "wave32Effective"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [ | |
| { "componentType": "f16", "M": 8, "N": 8, "K": 8 }, | |
| { "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 } | |
| ] | |
| }, | |
| "derive": { | |
| "outScalar": "\"f16\" if dtypes.T == \"f16\" else \"f32\"", | |
| "transA": "attrs.transA != 0", | |
| "transB": "attrs.transB != 0", | |
| "transBatchA": "false", | |
| "rowCount": "(dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2))", | |
| "K": "(dim(shapes.A, ranks.A - 2) if attrs.transA != 0 else dim(shapes.A, ranks.A - 1))", | |
| "N": "(dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1))", | |
| "batchCount": "numel(shapes.Y) / ((dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2)) * (dim(shapes.B, ranks.B - 2) if attrs.transB != 0 else dim(shapes.B, ranks.B - 1)))" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "TransposeMatMul.SubgroupMatrix", | |
| "shader": "fused-matmul-subgroup-matrix.wgsl.jinja", | |
| "bindings": ["a", "b_scalar", "y", "params_rows"], | |
| "dispatch": { "x": "ceil(N / 64)", "y": "ceil(rowCount / 32)", "z": "numel(shapes.Y) / (rowCount * N)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "broadcast_rank4_tiled_reg", | |
| "priority": 6, | |
| "when": ["f16Ok(dtypes.T)", "attrs.transA == 0", "attrs.transB == 0", "ranks.A == 4", "(ranks.B == 2 or ranks.B == 3)", "ranks.Y == 4", "dim(shapes.Y, 0) == dim(shapes.A, 0)", "(ranks.B == 2 or dim(shapes.A, 1) == dim(shapes.B, 0) or dim(shapes.A, 1) == 1 or dim(shapes.B, 0) == 1)", "dim(shapes.Y, 1) == (dim(shapes.A, 1) if ranks.B == 2 else max(dim(shapes.A, 1), dim(shapes.B, 0)))", "dim(shapes.A, 3) == dim(shapes.B, ranks.B - 2)", "dim(shapes.Y, 2) == dim(shapes.A, 2)", "dim(shapes.Y, 3) == dim(shapes.B, ranks.B - 1)", "dim(shapes.A, 2) >= 64", "dim(shapes.A, 3) >= 32", "dim(shapes.B, ranks.B - 1) >= 64", "ceil(dim(shapes.B, ranks.B - 1) / registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil(dim(shapes.A, 2) / registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "numel(shapes.Y) / (dim(shapes.A, 2) * dim(shapes.B, ranks.B - 1)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"], | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "TransposeMatMul.BroadcastRank4TiledReg", | |
| "shader": "matmul-tiled-general-reg.wgsl.jinja", | |
| "derive": { | |
| "aBatchShape": "prefix(shapes.A, ranks.A - 2) if ranks.B > 2 else []", | |
| "aRank": "ranks.A if ranks.B > 2 else 2", | |
| "K": "dim(shapes.A, ranks.A - 1)", | |
| "bShape": "shapes.B", | |
| "transBatchA": "false" | |
| }, | |
| "bindings": ["a", "b_scalar", "y", "params_rows"], | |
| "dispatch": { | |
| "x": "ceil(dim(shapes.B, ranks.B - 1) / registerTile)", | |
| "y": "ceil(rowCount / registerTile)", | |
| "z": "batchCount" | |
| } | |
| } | |
| ], | |
| "derive": { | |
| "rowCount": "outer(shapes.A, 3) if ranks.B == 2 else dim(shapes.A, 2)", | |
| "batchCount": "numel(shapes.Y) / (dim(shapes.A, 2) * dim(shapes.B, ranks.B - 1)) if ranks.B > 2 else 1" | |
| } | |
| }, | |
| { | |
| "id": "plain_rank2_tiled_reg", | |
| "priority": 4, | |
| "demoteWhen": ["plainRank2RegDeepPreferredTier"], | |
| "when": ["f16Ok(dtypes.T)", "fusedSgmatRank2Ok", "dim(shapes.A, 0) >= 64", "dim(shapes.A, 1) >= 32", "dim(shapes.B, 1) >= 64", "ceil(dim(shapes.A, 0) / registerTile) * ceil(dim(shapes.B, 1) / registerTile) >= tunables.TILED_REG_MIN_WORKGROUPS", "ceil(dim(shapes.B, 1) / registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil(dim(shapes.A, 0) / registerTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"], | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "TransposeMatMul.PlainRank2TiledReg", | |
| "shader": "matmul-tiled-general-reg.wgsl.jinja", | |
| "derive": { | |
| "rowCount": "dim(shapes.A, 0)", | |
| "K": "dim(shapes.A, 1)", | |
| "bShape": "shapes.B", | |
| "transBatchA": "false" | |
| }, | |
| "bindings": ["a", "b_scalar", "y", "params_rows"], | |
| "dispatch": { | |
| "x": "ceil(dim(shapes.B, 1) / registerTile)", | |
| "y": "ceil(dim(shapes.A, 0) / registerTile)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "tiled", | |
| "priority": 0, | |
| "when": ["ranks.A >= 1", "ranks.B >= 1", "f16Ok(dtypes.T)", "(dim(shapes.A, 0) if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA == 0 else dim(shapes.A, ranks.A - 2))) == (dim(shapes.B, 0) if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB != 0 else dim(shapes.B, ranks.B - 2)))", "ceil((1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2))) / 16) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceil((1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else dim(shapes.B, ranks.B - 2))) / 16) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "numel(shapes.Y) / max(1, (1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2))) * (1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else dim(shapes.B, ranks.B - 2)))) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"], | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "TransposeMatMul.Tiled", | |
| "shader": "matmul-tiled-general.wgsl.jinja", | |
| "derive": { | |
| "aShape": "shapes.A", | |
| "bShape": "shapes.B", | |
| "transA": "attrs.transA != 0", | |
| "transB": "attrs.transB != 0", | |
| "transBatchA": "false", | |
| "transBatchB": "false" | |
| }, | |
| "bindings": ["a", "b_scalar", "y"], | |
| "dispatch": { | |
| "x": "ceil((1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else dim(shapes.B, ranks.B - 2))) / generalTile)", | |
| "y": "ceil((1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2))) / generalTile)", | |
| "z": "numel(shapes.Y) / max(1, (1 if ranks.A == 1 else (dim(shapes.A, ranks.A - 1) if attrs.transA != 0 else dim(shapes.A, ranks.A - 2))) * (1 if ranks.B == 1 else (dim(shapes.B, ranks.B - 1) if attrs.transB == 0 else dim(shapes.B, ranks.B - 2))))" | |
| } | |
| } | |
| ] | |
| } | |
| ] | |
| } | |