Download build/webgpu/manifest.json from webgpu-kernels/com.microsoft.MultiHeadAttention: direct link, hf CLI and curl.
- Browser
- Download file 113 kB
-
https://huggingface.co/kernels/webgpu-kernels/com.microsoft.MultiHeadAttention/resolve/v1/build/webgpu/manifest.json
- Command line
-
hf download hf://webgpu-kernels/com.microsoft.MultiHeadAttention@v1/build/webgpu/manifest.json
-
curl -L -o manifest.json https://huggingface.co/kernels/webgpu-kernels/com.microsoft.MultiHeadAttention/resolve/v1/build/webgpu/manifest.json
113 kB
| { | |
| "domain": "com.microsoft", | |
| "name": "MultiHeadAttention", | |
| "sinceVersion": 1, | |
| "inputs": { | |
| "queryT": { "onnx": "query", "dtype": "T", "rank": 3 }, | |
| "keyT": { "onnx": "key", "dtype": "T", "rank": 3 }, | |
| "valueT": { "onnx": "value", "dtype": "T", "rank": 3 }, | |
| "biasT": { "onnx": "bias", "dtype": "T", "rank": 1, "optional": true }, | |
| "attentionBiasT": { "onnx": "attention_bias", "dtype": "T", "rank": 4, "optional": true } | |
| }, | |
| "outputs": { | |
| "outputT": { | |
| "onnx": "output", | |
| "dtype": "T", | |
| "rank": 3, | |
| "shape": "[dim(shapes.queryT, 0), dim(shapes.queryT, 1), dim(shapes.valueT, 2)]" | |
| } | |
| }, | |
| "attributes": { "unidirectional": { "default": 0 }, "num_heads": {}, "scale": {} }, | |
| "attributeConstraints": { "num_heads": { "required": true }, "unidirectional": { "values": [0, 1] } }, | |
| "typeConstraints": { "T": ["float32", "float16"] }, | |
| "tunables": { | |
| "WORKGROUP_SIZE": { "default": 256 }, | |
| "SMALL_SEQ_MAX": { "default": 32 }, | |
| "SMALL_SEQ_BLOCKED_MAX_KV": { "default": 64 }, | |
| "SMALL_SEQ_BLOCKED_MAX_HEAD_DIM": { "default": 32 }, | |
| "SMALL_SEQ_BLOCKED_QUERY_BLOCK": { "default": 8 }, | |
| "SMALL_SEQ_MAX_PRIVATE_FLOATS": { "default": 96 }, | |
| "FLASH_MAX_TILE_K": { "default": 8 }, | |
| "FLASH_CLUSTER_WG_SMALL": { "default": 64 }, | |
| "FLASH_CLUSTER_WG_LARGE": { "default": 128 }, | |
| "FLASH_MIN_QUERY_HEADS": { "default": 248 }, | |
| "DECODE_MAX_SPLITS": { "default": 16 }, | |
| "DECODE_KEYS_PER_SPLIT": { "default": 128 }, | |
| "SPLITK_TARGET_WORKGROUPS": { "default": 128 }, | |
| "MATERIALIZED_INNER_TILE": { "default": 16 }, | |
| "MATERIALIZED_QUERY_TILE": { "default": 64 }, | |
| "MATERIALIZED_KEY_TILE": { "default": 64 }, | |
| "MATERIALIZED_VALUE_TILE": { "default": 64 }, | |
| "MATERIALIZED_VALUE_TILE_D128": { "default": 128 }, | |
| "MATERIALIZED_WORKGROUP_DIM": { "default": 16 }, | |
| "MATERIALIZED_SOFTMAX_WORKGROUP_SIZE": { "default": 256 }, | |
| "MATERIALIZED_SGMAT_QUERY_TILE": { "default": 64 }, | |
| "MATERIALIZED_SGMAT_KEY_TILE": { "default": 64 }, | |
| "MATERIALIZED_SGMAT_INNER_TILE": { "default": 32 }, | |
| "MATERIALIZED_CACHED_SOFTMAX_WORKGROUP_SIZE": { "default": 128 }, | |
| "MATERIALIZED_CACHED_SOFTMAX_MAX_VECS_PER_LANE": { "default": 4 }, | |
| "PREFILL_QUERY_TILE": { "default": 32 }, | |
| "MATERIALIZED_FUSED_SOFTMAX_MIN_SCORE_BYTES": { "default": 16777216 } | |
| }, | |
| "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", | |
| "subgroupsWave32": "device.features.has(\"subgroups\") and wave32Adapter", | |
| "narrowSubgroupRange": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize < device.adapterInfo.subgroupMaxSize and device.adapterInfo.subgroupMaxSize <= 16", | |
| "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", | |
| "wave32SubgroupsUsable": "subgroupsWave32 or pinSubgroupSize32", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads if (ranks.queryT == 3 and attrs.num_heads > 0) else 0", | |
| "qkvDtypesOk": "tensorDtypes.keyT == tensorDtypes.queryT and tensorDtypes.valueT == tensorDtypes.queryT and tensorDtypes.outputT == tensorDtypes.queryT", | |
| "floatDtypeOk": "(tensorDtypes.queryT == \"float32\" or tensorDtypes.queryT == \"float16\") and f16Ok(tensorDtypes.queryT)", | |
| "qkvShapeOk": "ranks.queryT == 3 and ranks.keyT == 3 and ranks.valueT == 3 and ranks.outputT == 3 and attrs.num_heads > 0 and dim(shapes.queryT, 2) % attrs.num_heads == 0 and dim(shapes.keyT, 2) == dim(shapes.queryT, 2) and dim(shapes.valueT, 2) == dim(shapes.queryT, 2) and dim(shapes.keyT, 1) == dim(shapes.valueT, 1) and dim(shapes.queryT, 0) == dim(shapes.keyT, 0) and dim(shapes.queryT, 0) == dim(shapes.valueT, 0) and dim(shapes.outputT, 0) == dim(shapes.queryT, 0) and dim(shapes.outputT, 1) == dim(shapes.queryT, 1) and dim(shapes.outputT, 2) == dim(shapes.valueT, 2)", | |
| "noAttnBias": "not present.attentionBiasT", | |
| "attnBiasOk": "present.attentionBiasT and ranks.attentionBiasT == 4 and tensorDtypes.attentionBiasT == tensorDtypes.queryT and (dim(shapes.attentionBiasT, 0) == dim(shapes.queryT, 0) or dim(shapes.attentionBiasT, 0) == 1) and (dim(shapes.attentionBiasT, 1) == attrs.num_heads or dim(shapes.attentionBiasT, 1) == 1) and dim(shapes.attentionBiasT, 2) == dim(shapes.queryT, 1) and dim(shapes.attentionBiasT, 3) == dim(shapes.keyT, 1)", | |
| "qkvContractOk": "qkvShapeOk and qkvDtypesOk and floatDtypeOk and noAttnBias", | |
| "qkvMaskContractOk": "qkvShapeOk and qkvDtypesOk and floatDtypeOk and attnBiasOk", | |
| "biasOk": "present.biasT and ranks.biasT == 1 and tensorDtypes.biasT == tensorDtypes.queryT and dim(shapes.biasT, 0) == 3 * dim(shapes.queryT, 2)", | |
| "hasBias": "present.biasT", | |
| "hasMask": "present.attentionBiasT", | |
| "q32BroadcastSubgroupLanes": "32 if wave32SubgroupsUsable else 0", | |
| "q32BroadcastF32HeadVectors": "headDim / 4 if headDim % 4 == 0 else 0", | |
| "q32BroadcastF32RegisterGeometry": "wave32SubgroupsUsable and q32BroadcastF32HeadVectors == q32BroadcastSubgroupLanes", | |
| "subgroupCluster4": "not narrowSubgroupRange and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize >= 4 and device.adapterInfo.subgroupMinSize % 4 == 0 and device.adapterInfo.subgroupMaxSize % 4 == 0", | |
| "subgroupCluster8": "not narrowSubgroupRange and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize >= 8 and device.adapterInfo.subgroupMinSize % 8 == 0 and device.adapterInfo.subgroupMaxSize % 8 == 0", | |
| "attentionDispatchFits": "dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 1) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "flashHeadOk": "headDim % 4 == 0 and headDim >= 32 and headDim <= 256", | |
| "flashSizeOk": "flashHeadOk and (dim(shapes.queryT, 1) * attrs.num_heads >= tunables.FLASH_MIN_QUERY_HEADS or (dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512) or (dim(shapes.queryT, 1) > 1 and dim(shapes.keyT, 1) >= 2048)) and attentionDispatchFits", | |
| "flashShapeOk": "qkvContractOk and flashSizeOk", | |
| "flashMaskShapeOk": "qkvMaskContractOk and flashSizeOk", | |
| "noBiasSplitKCount": "min(tunables.DECODE_MAX_SPLITS if dim(shapes.queryT, 1) == 1 else max(1, ceilDiv(tunables.SPLITK_TARGET_WORKGROUPS, dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads)), ceilDiv(dim(shapes.keyT, 1), tunables.DECODE_KEYS_PER_SPLIT))", | |
| "biasSplitKCount": "min(tunables.DECODE_MAX_SPLITS, ceilDiv(dim(shapes.keyT, 1), tunables.DECODE_KEYS_PER_SPLIT))", | |
| "noBiasPartialOutBytes": "dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount * headDim * 4", | |
| "noBiasStatsBytes": "2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount * 4", | |
| "noBiasSplitScratchFits": "noBiasPartialOutBytes <= device.limits.maxStorageBufferBindingSize and noBiasPartialOutBytes <= device.limits.maxBufferSize and noBiasStatsBytes <= device.limits.maxStorageBufferBindingSize and noBiasStatsBytes <= device.limits.maxBufferSize", | |
| "noBiasSplitDispatchFits": "dim(shapes.queryT, 1) * noBiasSplitKCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "biasPartialOutBytes": "dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount * headDim * 4", | |
| "biasStatsBytes": "2 * dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount * 4", | |
| "biasSplitScratchFits": "biasPartialOutBytes <= device.limits.maxStorageBufferBindingSize and biasPartialOutBytes <= device.limits.maxBufferSize and biasStatsBytes <= device.limits.maxStorageBufferBindingSize and biasStatsBytes <= device.limits.maxBufferSize", | |
| "biasSplitDispatchFits": "biasSplitKCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "decodeSplitKNoBiasOk": "qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512 and flashHeadOk and attentionDispatchFits and noBiasSplitDispatchFits and noBiasSplitScratchFits", | |
| "shortQuerySplitKNoBiasOk": "qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) >= 2 and dim(shapes.queryT, 1) <= 16 and dim(shapes.keyT, 1) >= 2048 and flashHeadOk and attentionDispatchFits and noBiasSplitDispatchFits and noBiasSplitScratchFits", | |
| "decodeSplitKBiasOk": "biasOk and qkvContractOk and attrs.unidirectional == 0 and dim(shapes.queryT, 1) == 1 and dim(shapes.keyT, 1) >= 512 and flashHeadOk and attentionDispatchFits and biasSplitDispatchFits and biasSplitScratchFits", | |
| "decodeSplitKPortablePreferred": "tensorDtypes.queryT == \"float32\" and device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and headDim / 4 < device.adapterInfo.subgroupMinSize", | |
| "materializedScoreBytes": "dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1) * 4", | |
| "materializedScoreFits": "materializedScoreBytes <= device.limits.maxStorageBufferBindingSize and materializedScoreBytes <= device.limits.maxBufferSize", | |
| "materializedWorkgroupSize": "tunables.MATERIALIZED_WORKGROUP_DIM * tunables.MATERIALIZED_WORKGROUP_DIM", | |
| "materializedScoreStorageBytes": "(tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE + tunables.MATERIALIZED_KEY_TILE * (tunables.MATERIALIZED_INNER_TILE + 4)) * 4", | |
| "materializedApplyTileN": "tunables.MATERIALIZED_VALUE_TILE_D128 if headDim == 128 else tunables.MATERIALIZED_VALUE_TILE", | |
| "materializedApplyStorageBytes": "(tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE + tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN) * 4", | |
| "materializedTileGeometryOk": "tunables.MATERIALIZED_INNER_TILE % 4 == 0 and (materializedApplyTileN / tunables.MATERIALIZED_WORKGROUP_DIM) % 4 == 0 and tunables.MATERIALIZED_INNER_TILE > 0 and tunables.MATERIALIZED_WORKGROUP_DIM > 0 and tunables.MATERIALIZED_QUERY_TILE % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and tunables.MATERIALIZED_KEY_TILE % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and materializedApplyTileN % tunables.MATERIALIZED_WORKGROUP_DIM == 0 and tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE >= materializedWorkgroupSize and tunables.MATERIALIZED_KEY_TILE * tunables.MATERIALIZED_INNER_TILE >= materializedWorkgroupSize and tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN >= materializedWorkgroupSize and (tunables.MATERIALIZED_QUERY_TILE * tunables.MATERIALIZED_INNER_TILE) % materializedWorkgroupSize == 0 and (tunables.MATERIALIZED_KEY_TILE * tunables.MATERIALIZED_INNER_TILE) % materializedWorkgroupSize == 0 and (tunables.MATERIALIZED_INNER_TILE * materializedApplyTileN) % materializedWorkgroupSize == 0", | |
| "materializedDeviceOk": "materializedTileGeometryOk and materializedWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.MATERIALIZED_WORKGROUP_DIM <= device.limits.maxComputeWorkgroupSizeX and tunables.MATERIALIZED_WORKGROUP_DIM <= device.limits.maxComputeWorkgroupSizeY and materializedScoreStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and materializedApplyStorageBytes <= device.limits.maxComputeWorkgroupStorageSize", | |
| "materializedWideSimdOk": "device.features.has(\"subgroups\") or (has(device.adapterInfo, \"subgroupMinSize\") and device.adapterInfo.subgroupMinSize >= 16)", | |
| "materializedF32CoreOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and attrs.unidirectional == 0 and headDim >= 64 and headDim <= 128 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and dim(shapes.queryT, 0) * attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and materializedScoreFits and materializedDeviceOk", | |
| "materializedSoftmaxStorageBytes": "tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE * 8 + 8", | |
| "materializedSoftmaxResourcesFit": "tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE > 0 and pow2ceil(tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE) == tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE and tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE <= deviceWorkgroupCap and materializedSoftmaxStorageBytes <= device.limits.maxComputeWorkgroupStorageSize", | |
| "materializedSgmatQueryTile": "tunables.MATERIALIZED_SGMAT_QUERY_TILE", | |
| "materializedSgmatKeyTile": "tunables.MATERIALIZED_SGMAT_KEY_TILE", | |
| "materializedSgmatInnerTile": "tunables.MATERIALIZED_SGMAT_INNER_TILE", | |
| "materializedSgmatSubgroupRows": "floor(materializedSgmatQueryTile / 16)", | |
| "materializedSgmatSubgroupCols": "floor(materializedSgmatKeyTile / 32)", | |
| "materializedSgmatStatSlots": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile) * materializedSgmatSubgroupCols", | |
| "materializedRowStatsWg": "min(tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE, deviceWorkgroupCap)", | |
| "materializedRowStatsElements": "dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * 2", | |
| "materializedScorePartialElements": "dim(shapes.queryT, 0) * attrs.num_heads * materializedSgmatStatSlots * dim(shapes.queryT, 1) * 2", | |
| "materializedGemmStatSlots": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)", | |
| "materializedGemmScorePartialElements": "dim(shapes.queryT, 0) * attrs.num_heads * materializedGemmStatSlots * dim(shapes.queryT, 1) * 2", | |
| "materializedGemmFusedSoftmaxWorthIt": "materializedScoreBytes >= tunables.MATERIALIZED_FUSED_SOFTMAX_MIN_SCORE_BYTES and materializedGemmScorePartialElements * 4 <= device.limits.maxStorageBufferBindingSize and materializedGemmScorePartialElements * 4 <= device.limits.maxBufferSize and materializedRowStatsElements * 4 <= device.limits.maxStorageBufferBindingSize and materializedRowStatsWg > 0 and tunables.MATERIALIZED_QUERY_TILE <= tunables.MATERIALIZED_WORKGROUP_DIM * tunables.MATERIALIZED_WORKGROUP_DIM", | |
| "materializedFusedSoftmaxTilesOk": "materializedSgmatQueryTile >= 64 and materializedSgmatKeyTile >= 64", | |
| "materializedFusedSoftmaxWorthIt": "materializedScoreBytes >= tunables.MATERIALIZED_FUSED_SOFTMAX_MIN_SCORE_BYTES and materializedFusedSoftmaxTilesOk", | |
| "materializedSgmatWorkgroupSize": "materializedSgmatSubgroupRows * materializedSgmatSubgroupCols * 32", | |
| "materializedSgmatCompactStorageBytes": "(materializedSgmatQueryTile + materializedSgmatKeyTile) * materializedSgmatInnerTile * 4", | |
| "materializedSgmatGeometryOk": "materializedSgmatQueryTile >= 16 and materializedSgmatQueryTile % 16 == 0 and materializedSgmatKeyTile >= 32 and materializedSgmatKeyTile <= 64 and materializedSgmatKeyTile % 32 == 0 and materializedSgmatInnerTile == 32 and materializedSgmatWorkgroupSize > 0", | |
| "materializedSgmatBuffersFit": "numel(shapes.queryT) * 4 <= device.limits.maxStorageBufferBindingSize and numel(shapes.queryT) * 4 <= device.limits.maxBufferSize and numel(shapes.keyT) * 4 <= device.limits.maxStorageBufferBindingSize and numel(shapes.keyT) * 4 <= device.limits.maxBufferSize and numel(shapes.valueT) * 4 <= device.limits.maxStorageBufferBindingSize and numel(shapes.valueT) * 4 <= device.limits.maxBufferSize and numel(shapes.outputT) * 4 <= device.limits.maxStorageBufferBindingSize and numel(shapes.outputT) * 4 <= device.limits.maxBufferSize and materializedScoreFits", | |
| "materializedSgmatResourcesFit": "materializedSgmatGeometryOk and materializedSgmatWorkgroupSize <= deviceWorkgroupCap and materializedSgmatCompactStorageBytes <= device.limits.maxComputeWorkgroupStorageSize", | |
| "materializedSgmatDispatchFits": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) * attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "materializedSgmatDirectScoreStore": "dim(shapes.queryT, 1) % materializedSgmatQueryTile == 0 and dim(shapes.keyT, 1) % materializedSgmatKeyTile == 0", | |
| "materializedSgmatDirectApplyStore": "dim(shapes.queryT, 1) % materializedSgmatQueryTile == 0 and headDim % materializedSgmatKeyTile == 0", | |
| "materializedSgmatRuntimeDirectStore": "dim(shapes.queryT, 1) >= 2 * materializedSgmatQueryTile and dim(shapes.keyT, 1) >= 2 * materializedSgmatKeyTile", | |
| "materializedSgmatCoreOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and attrs.unidirectional == 0 and headDim >= 32 and headDim <= 256 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and device.features.has(\"subgroups\") and wave32Effective and device.features.has(\"chromium-experimental-subgroup-matrix\") and materializedSgmatBuffersFit and materializedSgmatResourcesFit and materializedSgmatDispatchFits", | |
| "materializedCachedSoftmaxWg": "min(tunables.MATERIALIZED_CACHED_SOFTMAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)", | |
| "materializedCachedSoftmaxVecsPerLane": "ceilDiv(ceilDiv(dim(shapes.keyT, 1), 4), max(1, materializedCachedSoftmaxWg))", | |
| "materializedCachedSoftmaxStorageBytes": "materializedCachedSoftmaxWg * 8 + 8", | |
| "materializedCachedSoftmaxResourcesFit": "materializedCachedSoftmaxWg > 0 and pow2ceil(materializedCachedSoftmaxWg) == materializedCachedSoftmaxWg and materializedCachedSoftmaxStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and tunables.MATERIALIZED_CACHED_SOFTMAX_MAX_VECS_PER_LANE > 0 and materializedCachedSoftmaxVecsPerLane <= tunables.MATERIALIZED_CACHED_SOFTMAX_MAX_VECS_PER_LANE", | |
| "materializedCachedSoftmaxOk": "dim(shapes.keyT, 1) % 4 == 0 and materializedCachedSoftmaxResourcesFit", | |
| "materializedAdaptiveSoftmaxOk": "materializedCachedSoftmaxOk or materializedSoftmaxResourcesFit", | |
| "materializedSgmatOk": "materializedSgmatCoreOk and materializedAdaptiveSoftmaxOk", | |
| "materializedSgmatFusedOk": "materializedSgmatCoreOk", | |
| "materializedF32Ok": "materializedF32CoreOk and materializedAdaptiveSoftmaxOk", | |
| "clusterTileKWg64": "max(1, min(tunables.FLASH_MAX_TILE_K, floor(device.limits.maxComputeWorkgroupStorageSize / (headDim * (8 if tensorDtypes.queryT != \"float16\" else 4) + tunables.FLASH_CLUSTER_WG_SMALL * 4))))", | |
| "clusterTileKWg128": "max(1, min(tunables.FLASH_MAX_TILE_K, floor(device.limits.maxComputeWorkgroupStorageSize / (headDim * (8 if tensorDtypes.queryT != \"float16\" else 4) + tunables.FLASH_CLUSTER_WG_LARGE * 4))))", | |
| "smallHeadShapeOk": "qkvContractOk and headDim < 32 and dim(shapes.keyT, 1) >= 64 and attrs.unidirectional == 0 and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "smallHeadParallelOk": "smallHeadShapeOk and dim(shapes.keyT, 1) <= 2048", | |
| "prefillTiledStorageBytes": "headDim * tunables.PREFILL_QUERY_TILE * 4", | |
| "prefillTiledDeviceOk": "tunables.PREFILL_QUERY_TILE <= deviceWorkgroupCap and prefillTiledStorageBytes <= device.limits.maxComputeWorkgroupStorageSize", | |
| "portableWorkgroupSize": "min(tunables.WORKGROUP_SIZE, pow2ceil(max(1, headDim)))", | |
| "portableWorkgroupStorageBytes": "portableWorkgroupSize * 8 + max(1, headDim) * 4 + 16", | |
| "portableWorkgroupOk": "tunables.WORKGROUP_SIZE > 0 and pow2ceil(tunables.WORKGROUP_SIZE) == tunables.WORKGROUP_SIZE and portableWorkgroupSize <= deviceWorkgroupCap and portableWorkgroupStorageBytes <= device.limits.maxComputeWorkgroupStorageSize", | |
| "fallbackShapeOk": "qkvContractOk and portableWorkgroupOk", | |
| "fallbackMaskShapeOk": "qkvMaskContractOk and portableWorkgroupOk", | |
| "smallSeqShapeOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and dim(shapes.queryT, 1) >= 1 and dim(shapes.queryT, 1) <= tunables.SMALL_SEQ_MAX and dim(shapes.keyT, 1) >= 1 and dim(shapes.keyT, 1) <= tunables.SMALL_SEQ_MAX and headDim >= 1 and attrs.unidirectional == 0", | |
| "smallSeqPrivateFloats": "dim(shapes.keyT, 1) + headDim", | |
| "smallSeqWorkgroupSize": "max(32, pow2ceil(dim(shapes.queryT, 1)))", | |
| "smallSeqSharedBytes": "dim(shapes.keyT, 1) * headDim * 8", | |
| "smallSeqResourcesFit": "smallSeqPrivateFloats <= tunables.SMALL_SEQ_MAX_PRIVATE_FLOATS and smallSeqWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and smallSeqWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX and smallSeqSharedBytes <= device.limits.maxComputeWorkgroupStorageSize", | |
| "smallSeqBlockedKvBytes": "dim(shapes.keyT, 1) * headDim * 8", | |
| "smallSeqBlockedLaneBytes": "8 + headDim * 4", | |
| "smallSeqBlockedKeyLanes": "min(pow2ceil(dim(shapes.keyT, 1)), 16 if smallSeqBlockedKvBytes + tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * 16 * smallSeqBlockedLaneBytes <= device.limits.maxComputeWorkgroupStorageSize else (8 if smallSeqBlockedKvBytes + tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * 8 * smallSeqBlockedLaneBytes <= device.limits.maxComputeWorkgroupStorageSize else 4))", | |
| "smallSeqBlockedWorkgroupSize": "tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK * smallSeqBlockedKeyLanes", | |
| "smallSeqBlockedSharedBytes": "smallSeqBlockedKvBytes + smallSeqBlockedWorkgroupSize * smallSeqBlockedLaneBytes", | |
| "smallSeqBlockedShapeOk": "qkvContractOk and tensorDtypes.queryT == \"float32\" and headDim % 4 == 0 and headDim >= 4 and headDim <= tunables.SMALL_SEQ_BLOCKED_MAX_HEAD_DIM and dim(shapes.keyT, 1) >= 1 and dim(shapes.keyT, 1) <= tunables.SMALL_SEQ_BLOCKED_MAX_KV and dim(shapes.queryT, 1) >= 1", | |
| "smallSeqBlockedFits": "smallSeqBlockedWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup and smallSeqBlockedWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX and smallSeqBlockedSharedBytes <= device.limits.maxComputeWorkgroupStorageSize and ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "smallSeqDispatchFits": "attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.queryT, 0) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "materializedSgmatCoreF16Ok": "qkvContractOk and tensorDtypes.queryT == \"float16\" and tensorDtypes.keyT == \"float16\" and tensorDtypes.valueT == \"float16\" and device.features.has(\"shader-f16\") and attrs.unidirectional == 0 and headDim >= 32 and headDim <= 256 and dim(shapes.queryT, 1) >= 512 and dim(shapes.keyT, 1) >= 512 and device.features.has(\"subgroups\") and wave32Effective and device.features.has(\"chromium-experimental-subgroup-matrix\") and materializedSgmatBuffersFit and materializedSgmatResourcesFit and materializedSgmatDispatchFits", | |
| "materializedSgmatFusedF16Ok": "materializedSgmatCoreF16Ok", | |
| "attentionScaleExpression": "\"select(inverseSqrt(f32(HEAD_DIM)), params.scale, params.scale != 0.0)\"", | |
| "smallHeadValueWg": "pow(2, log2ceil(min(64, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX) + 1) - 1)", | |
| "smallHeadValueStorage": "max(dim(shapes.keyT, 1), smallHeadValueWg * headDim)", | |
| "smallHeadValueQueryBlock": "2 if dim(shapes.queryT, 1) >= 2 and (smallHeadValueStorage + smallHeadValueWg) * 8 <= device.limits.maxComputeWorkgroupStorageSize else 1", | |
| "smallHeadValueEstimatedSteps": "headDim * (ceilDiv(dim(shapes.keyT, 1), smallHeadValueWg) + 2 * log2ceil(smallHeadValueWg) * (2 if tensorDtypes.queryT == \"float16\" else 1) + 2)", | |
| "smallHeadValueSgReductionSteps": "ceilDiv(smallHeadValueWg, max(1, device.adapterInfo.subgroupMinSize)) if has(device.adapterInfo, \"subgroupMinSize\") else log2ceil(smallHeadValueWg)", | |
| "smallHeadValueSgEstimatedSteps": "headDim * (ceilDiv(dim(shapes.keyT, 1), smallHeadValueWg) + 2 * smallHeadValueSgReductionSteps * (2 if tensorDtypes.queryT == \"float16\" else 1) + 2)" | |
| }, | |
| "bindings": { | |
| "query": { "arg": "queryT", "elementType": "$inputElement" }, | |
| "key": { "arg": "keyT", "elementType": "$inputElement" }, | |
| "value": { "arg": "valueT", "elementType": "$inputElement" }, | |
| "bias": { "arg": "biasT", "elementType": "$inputScalar" }, | |
| "output": { "arg": "outputT", "elementType": "$outputElement" }, | |
| "params": { | |
| "struct": [ | |
| { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" }, | |
| { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" }, | |
| { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }, | |
| { "name": "isCausal", "type": "u32", "value": "attrs.unidirectional" } | |
| ] | |
| }, | |
| "params_main": { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" }, | |
| { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" } | |
| ] | |
| }, | |
| "q": { "arg": "queryT", "elementType": "$scalar" }, | |
| "k": { "arg": "keyT", "elementType": "$scalar" }, | |
| "v": { "arg": "valueT", "elementType": "$scalar" }, | |
| "y": { "arg": "outputT", "elementType": "$scalar" }, | |
| "params_scores": { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" }, | |
| { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" }, | |
| { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" } | |
| ] | |
| }, | |
| "query_f32": { "arg": "queryT", "name": "query", "elementType": "f32" }, | |
| "key_f32": { "arg": "keyT", "name": "key", "elementType": "f32" }, | |
| "scores": { "scratch": "materializedScores", "elementType": "f32" }, | |
| "scorePartials": { "scratch": "materializedScorePartials", "elementType": "f32" }, | |
| "bias_f32": { "arg": "biasT", "name": "bias", "elementType": "f32" }, | |
| "scorePartials_f32": { | |
| "scratch": "materializedScorePartials", | |
| "name": "scorePartials", | |
| "buffer": "read-only-storage", | |
| "elementType": "f32" | |
| }, | |
| "rowStats": { "scratch": "materializedRowStats", "elementType": "f32" }, | |
| "params_rows": { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "rows", "type": "u32", "value": "dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)" } | |
| ] | |
| }, | |
| "scores_f32": { | |
| "scratch": "materializedScores", | |
| "name": "scores", | |
| "buffer": "read-only-storage", | |
| "elementType": "f32" | |
| }, | |
| "value_f32": { "arg": "valueT", "name": "value", "elementType": "f32" }, | |
| "output_f32": { "arg": "outputT", "name": "output", "elementType": "f32" }, | |
| "rowStats_f32": { | |
| "scratch": "materializedRowStats", | |
| "name": "rowStats", | |
| "buffer": "read-only-storage", | |
| "elementType": "f32" | |
| }, | |
| "params_apply": { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" }, | |
| { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" } | |
| ] | |
| }, | |
| "attn_mask_main": { "arg": "attentionBiasT", "name": "attn_mask", "elementType": "$maskElement" }, | |
| "params__uniform": { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" }, | |
| { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" }, | |
| { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }, | |
| { "name": "isCausal", "type": "u32", "value": "attrs.unidirectional" }, | |
| { | |
| "name": "maskBatchStride", | |
| "type": "u32", | |
| "value": "0 if dim(shapes.attentionBiasT, 0) == 1 else dim(shapes.attentionBiasT, 1) * dim(shapes.attentionBiasT, 2) * dim(shapes.attentionBiasT, 3)" | |
| }, | |
| { | |
| "name": "maskHeadStride", | |
| "type": "u32", | |
| "value": "0 if dim(shapes.attentionBiasT, 1) == 1 else dim(shapes.attentionBiasT, 2) * dim(shapes.attentionBiasT, 3)" | |
| }, | |
| { "name": "maskSeqStride", "type": "u32", "value": "dim(shapes.attentionBiasT, 3)" } | |
| ] | |
| }, | |
| "query_query_t": { "arg": "queryT", "name": "query", "elementType": "$inputVec4" }, | |
| "key_key_t": { "arg": "keyT", "name": "key", "elementType": "$inputVec4" }, | |
| "value_value_t": { "arg": "valueT", "name": "value", "elementType": "$inputVec4" }, | |
| "partial_out": { "scratch": "partialOut", "elementType": "vec4<f32>" }, | |
| "partial_stats": { "scratch": "partialStats", "elementType": "vec2<f32>" }, | |
| "params_kv_seq_scale": { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" }, | |
| { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" } | |
| ] | |
| }, | |
| "partial_out_merge": { | |
| "scratch": "partialOut", | |
| "name": "partial_out", | |
| "buffer": "read-only-storage", | |
| "elementType": "vec4<f32>" | |
| }, | |
| "partial_stats_merge": { | |
| "scratch": "partialStats", | |
| "name": "partial_stats", | |
| "buffer": "read-only-storage", | |
| "elementType": "vec2<f32>" | |
| }, | |
| "output_merge": { "arg": "outputT", "name": "output", "elementType": "$inputVec4" }, | |
| "scores_softmax": { "scratch": "materializedScores", "name": "scores", "elementType": "$softmaxElementType" } | |
| }, | |
| "variants": [ | |
| { | |
| "id": "qkv_no_bias_small_head_value_subgroups", | |
| "priority": 12, | |
| "when": ["not present.biasT", "smallHeadShapeOk", "headDim > 0", "headDim <= smallHeadValueWg", "(smallHeadValueStorage + smallHeadValueWg) * 4 * smallHeadValueQueryBlock <= device.limits.maxComputeWorkgroupStorageSize", "device.wgslLanguageFeatures.has(\"subgroup_id\")"], | |
| "derive": { | |
| "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "kvSeq": "dim(shapes.keyT, 1)", | |
| "valueWorkgroupSize": "smallHeadValueWg", | |
| "scoreStorageElements": "smallHeadValueStorage", | |
| "queryBlock": "smallHeadValueQueryBlock", | |
| "queryTail": "dim(shapes.queryT, 1) % smallHeadValueQueryBlock != 0", | |
| "useValueSubgroups": true | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.SmallHeadValueSubgroups", | |
| "shader": "attn-small-head-value.wgsl.jinja", | |
| "bindings": ["query", "key", "value", "output", "params_main"], | |
| "dispatch": { | |
| "x": "min(ceilDiv(dim(shapes.queryT, 1), smallHeadValueQueryBlock), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ], | |
| "requires": { "features": ["subgroups"] }, | |
| "demoteWhen": ["dim(shapes.keyT, 1) <= smallHeadValueSgEstimatedSteps"] | |
| }, | |
| { | |
| "id": "qkv_no_bias_small_head_value_tree", | |
| "priority": 11, | |
| "when": ["not present.biasT", "smallHeadShapeOk", "headDim > 0", "headDim <= smallHeadValueWg", "(smallHeadValueStorage + smallHeadValueWg) * 4 * smallHeadValueQueryBlock <= device.limits.maxComputeWorkgroupStorageSize", "true"], | |
| "derive": { | |
| "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "kvSeq": "dim(shapes.keyT, 1)", | |
| "valueWorkgroupSize": "smallHeadValueWg", | |
| "scoreStorageElements": "smallHeadValueStorage", | |
| "queryBlock": "smallHeadValueQueryBlock", | |
| "queryTail": "dim(shapes.queryT, 1) % smallHeadValueQueryBlock != 0", | |
| "useValueSubgroups": false | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.SmallHeadValueTree", | |
| "shader": "attn-small-head-value.wgsl.jinja", | |
| "bindings": ["query", "key", "value", "output", "params_main"], | |
| "dispatch": { | |
| "x": "min(ceilDiv(dim(shapes.queryT, 1), smallHeadValueQueryBlock), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ], | |
| "requires": { "features": [] }, | |
| "demoteWhen": ["dim(shapes.keyT, 1) <= smallHeadValueEstimatedSteps"] | |
| }, | |
| { | |
| "id": "qkv_bias_small_seq_blocked", | |
| "priority": 40, | |
| "when": ["biasOk", "smallSeqBlockedShapeOk", "not flashShapeOk", "smallSeqBlockedFits"], | |
| "derive": { | |
| "inputElement": "\"vec4<f32>\"", | |
| "outputElement": "\"vec4<f32>\"", | |
| "inputScalar": "\"f32\"", | |
| "headDimV4": "headDim / 4", | |
| "hidden": "dim(shapes.queryT, 2)", | |
| "hiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvSeq": "dim(shapes.keyT, 1)", | |
| "queryBlock": "tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK", | |
| "keyLanes": "smallSeqBlockedKeyLanes" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.SmallSeqBlockedBias", | |
| "shader": "mha-small-seq-blocked.wgsl.jinja", | |
| "bindings": ["query", "key", "value", "bias", "output", "params"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_small_seq_blocked", | |
| "priority": 40, | |
| "when": ["not present.biasT", "smallSeqBlockedShapeOk", "not flashShapeOk", "smallSeqBlockedFits"], | |
| "derive": { | |
| "inputElement": "\"vec4<f32>\"", | |
| "outputElement": "\"vec4<f32>\"", | |
| "headDimV4": "headDim / 4", | |
| "hidden": "dim(shapes.queryT, 2)", | |
| "hiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvSeq": "dim(shapes.keyT, 1)", | |
| "queryBlock": "tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK", | |
| "keyLanes": "smallSeqBlockedKeyLanes" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.SmallSeqBlocked", | |
| "shader": "mha-small-seq-blocked.wgsl.jinja", | |
| "bindings": ["query", "key", "value", "output", "params"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), tunables.SMALL_SEQ_BLOCKED_QUERY_BLOCK)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_small_seq", | |
| "priority": 60, | |
| "when": ["not present.biasT", "smallSeqShapeOk", "not flashShapeOk", "smallSeqResourcesFit", "smallSeqDispatchFits"], | |
| "derive": { | |
| "inputElement": "\"f32\"", | |
| "outputElement": "\"f32\"", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "kvSeq": "dim(shapes.keyT, 1)", | |
| "hidden": "dim(shapes.queryT, 2)", | |
| "workgroupSize": "smallSeqWorkgroupSize" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention", | |
| "shader": "mha-small-seq.wgsl.jinja", | |
| "bindings": ["query", "key", "value", "output", "params_main"], | |
| "dispatch": { "x": "attrs.num_heads", "y": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_tiled_nosg", | |
| "priority": 19, | |
| "when": ["not present.biasT", "qkvContractOk", "not smallHeadParallelOk", "headDim % 4 == 0", "headDim <= 128", "dim(shapes.queryT, 1) >= 31", "prefillTiledDeviceOk"], | |
| "supersededBy": ["qkv_no_bias_flash_cluster_nosg", "qkv_no_bias_flash_cluster_lpq4_nosg"], | |
| "derive": { | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "blockM": "tunables.PREFILL_QUERY_TILE", | |
| "vHeadCap": "dim(shapes.queryT, 2) / attrs.num_heads" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.PrefillTiledNoSg", | |
| "shader": "attention-rank4-tiled.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": [ | |
| "q", | |
| "k", | |
| "v", | |
| "y", | |
| { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "count", "type": "u32", "value": "numel(shapes.outputT)" }, | |
| { "name": "qHeads", "type": "u32", "value": "attrs.num_heads" }, | |
| { "name": "kvHeads", "type": "u32", "value": "attrs.num_heads" }, | |
| { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" }, | |
| { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" }, | |
| { "name": "headSize", "type": "u32", "value": "dim(shapes.queryT, 2) / attrs.num_heads" }, | |
| { "name": "vHeadSize", "type": "u32", "value": "dim(shapes.valueT, 2) / attrs.num_heads" }, | |
| { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }, | |
| { "name": "softcap", "type": "f32", "value": "0" }, | |
| { "name": "isCausal", "type": "u32", "value": "attrs.unidirectional" }, | |
| { "name": "qHidden", "type": "u32", "value": "dim(shapes.queryT, 2)" }, | |
| { "name": "kvHidden", "type": "u32", "value": "dim(shapes.keyT, 2)" }, | |
| { "name": "vHidden", "type": "u32", "value": "dim(shapes.valueT, 2)" } | |
| ] | |
| } | |
| ], | |
| "dispatch": { | |
| "x": "min(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)", | |
| "y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_flash_q32_broadcast_f32_d128", | |
| "priority": 30, | |
| "when": ["tensorDtypes.queryT == \"float32\"", "biasOk", "attrs.unidirectional == 0", "flashShapeOk", "q32BroadcastF32RegisterGeometry", "dim(shapes.queryT, 1) >= 31", "ceilDiv(dim(shapes.queryT, 1), 32) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "wave32SubgroupsUsable"], | |
| "requires": { "features": ["subgroups"] }, | |
| "derive": { | |
| "hasCausal": false, | |
| "usesF16": false, | |
| "scalar": "\"f32\"", | |
| "inputElement": "\"vec4<f32>\"", | |
| "outputElement": "\"vec4<f32>\"", | |
| "inputScalar": "\"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "kStep": 64, | |
| "qkGroups": 16 | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.FlashQ32BroadcastF32Bias", | |
| "shader": "attn-flash-q32-broadcast.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "bias", "output", "params_scores"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), 32)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_small_head_parallel", | |
| "priority": 10, | |
| "when": ["not present.biasT", "smallHeadParallelOk"], | |
| "derive": { | |
| "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "kvSeq": "dim(shapes.keyT, 1)" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.SmallHeadParallel", | |
| "shader": "attn-small-head-parallel.wgsl.jinja", | |
| "bindings": ["query", "key", "value", "output", "params_main"], | |
| "dispatch": { | |
| "x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_tiled_attn_bias_nosg", | |
| "priority": 17, | |
| "when": ["not present.biasT", "qkvMaskContractOk", "not smallHeadParallelOk", "headDim % 4 == 0", "headDim <= 128", "dim(shapes.queryT, 1) >= 31", "prefillTiledDeviceOk"], | |
| "derive": { | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "blockM": "tunables.PREFILL_QUERY_TILE", | |
| "vHeadCap": "dim(shapes.queryT, 2) / attrs.num_heads" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.PrefillTiledAttnBiasNoSg", | |
| "shader": "attention-rank4-tiled.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": [ | |
| "q", | |
| "k", | |
| "v", | |
| { "arg": "attentionBiasT", "name": "attn_mask", "elementType": "$scalar" }, | |
| "y", | |
| { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "count", "type": "u32", "value": "numel(shapes.outputT)" }, | |
| { "name": "qHeads", "type": "u32", "value": "attrs.num_heads" }, | |
| { "name": "kvHeads", "type": "u32", "value": "attrs.num_heads" }, | |
| { "name": "qSeq", "type": "u32", "value": "dim(shapes.queryT, 1)" }, | |
| { "name": "kvSeq", "type": "u32", "value": "dim(shapes.keyT, 1)" }, | |
| { "name": "headSize", "type": "u32", "value": "dim(shapes.queryT, 2) / attrs.num_heads" }, | |
| { "name": "vHeadSize", "type": "u32", "value": "dim(shapes.valueT, 2) / attrs.num_heads" }, | |
| { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" }, | |
| { "name": "softcap", "type": "f32", "value": "0" }, | |
| { "name": "isCausal", "type": "u32", "value": "attrs.unidirectional" }, | |
| { "name": "qHidden", "type": "u32", "value": "dim(shapes.queryT, 2)" }, | |
| { "name": "kvHidden", "type": "u32", "value": "dim(shapes.keyT, 2)" }, | |
| { "name": "vHidden", "type": "u32", "value": "dim(shapes.valueT, 2)" }, | |
| { | |
| "name": "maskBatchStride", | |
| "type": "u32", | |
| "value": "0 if dim(shapes.attentionBiasT, 0) == 1 else dim(shapes.attentionBiasT, 1) * dim(shapes.attentionBiasT, 2) * dim(shapes.attentionBiasT, 3)" | |
| }, | |
| { | |
| "name": "maskHeadStride", | |
| "type": "u32", | |
| "value": "0 if dim(shapes.attentionBiasT, 1) == 1 else dim(shapes.attentionBiasT, 2) * dim(shapes.attentionBiasT, 3)" | |
| }, | |
| { "name": "maskSeqStride", "type": "u32", "value": "dim(shapes.attentionBiasT, 3)" } | |
| ] | |
| } | |
| ], | |
| "dispatch": { | |
| "x": "min(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)", | |
| "y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * attrs.num_heads * ceil(dim(shapes.outputT, 1) / blockM) * blockM), (blockM)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_materialized_sgmat_fused_f32", | |
| "priority": 52, | |
| "when": ["materializedSgmatFusedOk", "not present.biasT", "materializedFusedSoftmaxWorthIt"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "qNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "statSlots": "materializedSgmatStatSlots", | |
| "statQuerySeq": "dim(shapes.queryT, 1)" | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "materializedScores", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1)]" | |
| }, | |
| { "id": "materializedRowStats", "dtype": "float32", "shape": "[materializedRowStatsElements]" }, | |
| { "id": "materializedScorePartials", "dtype": "float32", "shape": "[materializedScorePartialElements]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "scores", | |
| "name": "MultiHeadAttention.MaterializedScoresSgmat", | |
| "shader": "attn-materialized-sgmat-f32.wgsl.jinja", | |
| "derive": { "phase": "\"score\"", "emitRowStats": true }, | |
| "bindings": ["query_f32", "key_f32", "scores", "scorePartials", "params_scores"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| }, | |
| { | |
| "id": "rowstats", | |
| "name": "MultiHeadAttention.MaterializedRowStatsCombine", | |
| "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja", | |
| "bindings": ["scorePartials_f32", "rowStats", "params_rows"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": 1, | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "apply", | |
| "name": "MultiHeadAttention.MaterializedApplySgmat", | |
| "shader": "attn-materialized-sgmat-f32.wgsl.jinja", | |
| "derive": { "phase": "\"apply\"", "fusedSoftmax": true }, | |
| "bindings": ["scores_f32", "value_f32", "output_f32", "rowStats_f32", "params_apply"], | |
| "dispatch": { | |
| "x": "ceilDiv(headDim, materializedSgmatKeyTile)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_materialized_sgmat_fused_f32", | |
| "priority": 52, | |
| "when": ["materializedSgmatFusedOk", "biasOk", "materializedFusedSoftmaxWorthIt"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "qNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "statSlots": "materializedSgmatStatSlots", | |
| "statQuerySeq": "dim(shapes.queryT, 1)" | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "materializedScores", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1)]" | |
| }, | |
| { "id": "materializedRowStats", "dtype": "float32", "shape": "[materializedRowStatsElements]" }, | |
| { "id": "materializedScorePartials", "dtype": "float32", "shape": "[materializedScorePartialElements]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "scores", | |
| "name": "MultiHeadAttention.MaterializedScoresSgmatBias", | |
| "shader": "attn-materialized-sgmat-f32.wgsl.jinja", | |
| "derive": { "phase": "\"score\"", "emitRowStats": true }, | |
| "bindings": ["query_f32", "key_f32", "bias_f32", "scores", "scorePartials", "params_scores"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| }, | |
| { | |
| "id": "rowstats", | |
| "name": "MultiHeadAttention.MaterializedRowStatsCombineBias", | |
| "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja", | |
| "bindings": ["scorePartials_f32", "rowStats", "params_rows"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": 1, | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "apply", | |
| "name": "MultiHeadAttention.MaterializedApplySgmatBias", | |
| "shader": "attn-materialized-sgmat-f32.wgsl.jinja", | |
| "derive": { "phase": "\"apply\"", "fusedSoftmax": true }, | |
| "bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "rowStats_f32", "params_apply"], | |
| "dispatch": { | |
| "x": "ceilDiv(headDim, materializedSgmatKeyTile)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_materialized_sgmat_fused_f16", | |
| "priority": 52, | |
| "when": ["materializedSgmatFusedF16Ok", "not present.biasT", "materializedFusedSoftmaxWorthIt"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f16", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "qNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "statSlots": "materializedSgmatStatSlots", | |
| "statQuerySeq": "dim(shapes.queryT, 1)", | |
| "operandF16": true, | |
| "inputElement": "\"f16\"", | |
| "outputElement": "\"f16\"" | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "materializedScores", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1)]" | |
| }, | |
| { "id": "materializedRowStats", "dtype": "float32", "shape": "[materializedRowStatsElements]" }, | |
| { "id": "materializedScorePartials", "dtype": "float32", "shape": "[materializedScorePartialElements]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "scores", | |
| "name": "MultiHeadAttention.MaterializedScoresSgmatF16", | |
| "shader": "attn-materialized-sgmat-f32.wgsl.jinja", | |
| "derive": { "phase": "\"score\"", "emitRowStats": true }, | |
| "bindings": ["query", "key", "scores", "scorePartials", "params_scores"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| }, | |
| { | |
| "id": "rowstats", | |
| "name": "MultiHeadAttention.MaterializedRowStatsCombine", | |
| "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja", | |
| "bindings": ["scorePartials_f32", "rowStats", "params_rows"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": 1, | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "apply", | |
| "name": "MultiHeadAttention.MaterializedApplySgmatF16", | |
| "shader": "attn-materialized-sgmat-f32.wgsl.jinja", | |
| "derive": { "phase": "\"apply\"", "fusedSoftmax": true }, | |
| "bindings": ["scores_f32", "value", "output", "rowStats_f32", "params_apply"], | |
| "dispatch": { | |
| "x": "ceilDiv(headDim, materializedSgmatKeyTile)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_cluster_lpq4_nosg", | |
| "priority": 20, | |
| "when": ["not present.biasT", "flashShapeOk", "headDim % 16 == 0", "headDim % 32 != 0", "dim(shapes.queryT, 1) >= 31"], | |
| "requires": {}, | |
| "derive": { | |
| "hasCausal": true, | |
| "usesF16": "tensorDtypes.queryT == \"float16\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "TILE_Q": 16, | |
| "TILE_K": "clusterTileKWg64", | |
| "batchNoSgReduction": "tensorDtypes.queryT == \"float16\"", | |
| "LPQ": 4, | |
| "useSubgroups": false | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-prefill-cluster.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "output", "params"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_cluster_nosg", | |
| "priority": 20, | |
| "when": ["not present.biasT", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31"], | |
| "requires": {}, | |
| "derive": { | |
| "hasCausal": true, | |
| "usesF16": "tensorDtypes.queryT == \"float16\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "TILE_Q": 16, | |
| "TILE_K": "clusterTileKWg128", | |
| "batchNoSgReduction": "tensorDtypes.queryT == \"float16\"", | |
| "LPQ": 8, | |
| "useSubgroups": false | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-prefill-cluster.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "output", "params"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_flash_cluster_nosg", | |
| "priority": 19, | |
| "when": ["biasOk", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31"], | |
| "requires": {}, | |
| "derive": { | |
| "hasCausal": true, | |
| "usesF16": "tensorDtypes.queryT == \"float16\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "TILE_Q": 16, | |
| "TILE_K": "clusterTileKWg128", | |
| "batchNoSgReduction": "tensorDtypes.queryT == \"float16\"", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "LPQ": 8, | |
| "useSubgroups": false | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-prefill-cluster.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "bias", "output", "params"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_cluster_lpq4", | |
| "priority": 22, | |
| "when": ["not present.biasT", "flashShapeOk", "headDim % 16 == 0", "headDim % 32 != 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster4"], | |
| "requires": { "features": ["subgroups"] }, | |
| "derive": { | |
| "hasCausal": true, | |
| "usesF16": "tensorDtypes.queryT == \"float16\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "TILE_Q": 16, | |
| "TILE_K": "32 if (tensorDtypes.queryT != \"float16\" and (dim(shapes.queryT, 2) / attrs.num_heads) <= 64) else 8", | |
| "LPQ": 4 | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-prefill-cluster.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "output", "params"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_cluster", | |
| "priority": 22, | |
| "when": ["not present.biasT", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"], | |
| "requires": { "features": ["subgroups"] }, | |
| "derive": { | |
| "hasCausal": true, | |
| "usesF16": "tensorDtypes.queryT == \"float16\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "TILE_Q": 16, | |
| "TILE_K": "32 if (tensorDtypes.queryT != \"float16\" and (dim(shapes.queryT, 2) / attrs.num_heads) <= 64) else 8", | |
| "LPQ": 8 | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-prefill-cluster.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "output", "params"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_flash_cluster", | |
| "priority": 21, | |
| "when": ["biasOk", "flashShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"], | |
| "requires": { "features": ["subgroups"] }, | |
| "derive": { | |
| "hasCausal": true, | |
| "usesF16": "tensorDtypes.queryT == \"float16\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "TILE_Q": 16, | |
| "TILE_K": "32 if (tensorDtypes.queryT != \"float16\" and (dim(shapes.queryT, 2) / attrs.num_heads) <= 64) else 8", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "LPQ": 8 | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-prefill-cluster.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "bias", "output", "params"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_cluster_attn_bias", | |
| "priority": 22, | |
| "when": ["not present.biasT", "flashMaskShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"], | |
| "requires": { "features": ["subgroups"] }, | |
| "derive": { | |
| "hasCausal": true, | |
| "usesF16": "tensorDtypes.queryT == \"float16\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "TILE_Q": 16, | |
| "TILE_K": "32 if (tensorDtypes.queryT != \"float16\" and (dim(shapes.queryT, 2) / attrs.num_heads) <= 64) else 8", | |
| "LPQ": 8 | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-prefill-cluster.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_flash_cluster_attn_bias", | |
| "priority": 21, | |
| "when": ["biasOk", "flashMaskShapeOk", "headDim % 32 == 0", "dim(shapes.queryT, 1) >= 31", "subgroupCluster8"], | |
| "requires": { "features": ["subgroups"] }, | |
| "derive": { | |
| "hasCausal": true, | |
| "usesF16": "tensorDtypes.queryT == \"float16\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "TILE_Q": 16, | |
| "TILE_K": "32 if (tensorDtypes.queryT != \"float16\" and (dim(shapes.queryT, 2) / attrs.num_heads) <= 64) else 8", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "LPQ": 8 | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-prefill-cluster.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), TILE_Q)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_splitk_nosg", | |
| "priority": 21, | |
| "when": ["not present.biasT", "decodeSplitKNoBiasOk or shortQuerySplitKNoBiasOk"], | |
| "requires": {}, | |
| "derive": { | |
| "useSubgroups": false, | |
| "splitQueries": true, | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "qSeq": "dim(shapes.queryT, 1)", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "numSplits": "noBiasSplitKCount", | |
| "usesF16": "tensorDtypes.queryT == \"float16\"" | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "partialOut", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount * (dim(shapes.queryT, 2) / attrs.num_heads)]" | |
| }, | |
| { | |
| "id": "partialStats", | |
| "dtype": "float32", | |
| "shape": "[2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount]" | |
| } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "split_attention", | |
| "name": "MultiHeadAttention.DecodeSplitKNoSg", | |
| "shader": "attn-flash-decode-splitk.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query_query_t", "key_key_t", "value_value_t", "partial_out", "partial_stats", "params_kv_seq_scale"], | |
| "dispatch": { | |
| "x": "dim(shapes.queryT, 1) * noBiasSplitKCount", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| }, | |
| { | |
| "id": "merge", | |
| "name": "MultiHeadAttention.DecodeSplitKMergeNoSg", | |
| "shader": "attn-flash-decode-splitk-merge.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["partial_out_merge", "partial_stats_merge", "output_merge"], | |
| "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_flash_splitk_nosg", | |
| "priority": 17, | |
| "when": ["biasOk", "decodeSplitKBiasOk"], | |
| "requires": {}, | |
| "derive": { | |
| "useSubgroups": false, | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "numSplits": "biasSplitKCount", | |
| "usesF16": "tensorDtypes.queryT == \"float16\"" | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "partialOut", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount * (dim(shapes.queryT, 2) / attrs.num_heads)]" | |
| }, | |
| { | |
| "id": "partialStats", | |
| "dtype": "float32", | |
| "shape": "[2 * dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount]" | |
| } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "split_attention", | |
| "name": "MultiHeadAttention.DecodeSplitKBiasNoSg", | |
| "shader": "attn-flash-decode-splitk.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query_query_t", "key_key_t", "value_value_t", "bias", "partial_out", "partial_stats", "params_kv_seq_scale"], | |
| "dispatch": { "x": "biasSplitKCount", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| }, | |
| { | |
| "id": "merge", | |
| "name": "MultiHeadAttention.DecodeSplitKMergeBiasNoSg", | |
| "shader": "attn-flash-decode-splitk-merge.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["partial_out_merge", "partial_stats_merge", "bias", "output_merge"], | |
| "dispatch": { "x": 1, "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_splitk", | |
| "priority": 25, | |
| "when": ["not present.biasT", "decodeSplitKNoBiasOk"], | |
| "demoteWhen": ["decodeSplitKPortablePreferred"], | |
| "requires": { "features": ["subgroups"] }, | |
| "derive": { | |
| "splitQueries": true, | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "qSeq": "dim(shapes.queryT, 1)", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "numSplits": "noBiasSplitKCount", | |
| "usesF16": "tensorDtypes.queryT == \"float16\"" | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "partialOut", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount * (dim(shapes.queryT, 2) / attrs.num_heads)]" | |
| }, | |
| { | |
| "id": "partialStats", | |
| "dtype": "float32", | |
| "shape": "[2 * dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * attrs.num_heads * noBiasSplitKCount]" | |
| } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "split_attention", | |
| "name": "MultiHeadAttention.DecodeSplitK", | |
| "shader": "attn-flash-decode-splitk.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query_query_t", "key_key_t", "value_value_t", "partial_out", "partial_stats", "params_kv_seq_scale"], | |
| "dispatch": { | |
| "x": "dim(shapes.queryT, 1) * noBiasSplitKCount", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| }, | |
| { | |
| "id": "merge", | |
| "name": "MultiHeadAttention.DecodeSplitKMerge", | |
| "shader": "attn-flash-decode-splitk-merge.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["partial_out_merge", "partial_stats_merge", "output_merge"], | |
| "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_flash_splitk", | |
| "priority": 24, | |
| "when": ["biasOk", "decodeSplitKBiasOk"], | |
| "requires": { "features": ["subgroups"] }, | |
| "derive": { | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputVec4": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "numSplits": "biasSplitKCount", | |
| "usesF16": "tensorDtypes.queryT == \"float16\"" | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "partialOut", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount * (dim(shapes.queryT, 2) / attrs.num_heads)]" | |
| }, | |
| { | |
| "id": "partialStats", | |
| "dtype": "float32", | |
| "shape": "[2 * dim(shapes.queryT, 0) * attrs.num_heads * biasSplitKCount]" | |
| } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "split_attention", | |
| "name": "MultiHeadAttention.DecodeSplitKBias", | |
| "shader": "attn-flash-decode-splitk.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query_query_t", "key_key_t", "value_value_t", "bias", "partial_out", "partial_stats", "params_kv_seq_scale"], | |
| "dispatch": { "x": "biasSplitKCount", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| }, | |
| { | |
| "id": "merge", | |
| "name": "MultiHeadAttention.DecodeSplitKMergeBias", | |
| "shader": "attn-flash-decode-splitk-merge.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["partial_out_merge", "partial_stats_merge", "bias", "output_merge"], | |
| "dispatch": { "x": 1, "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_materialized_gemm_f32", | |
| "priority": 42, | |
| "when": ["materializedF32Ok", "materializedWideSimdOk", "not present.biasT", "not materializedGemmFusedSoftmaxWorthIt"], | |
| "demoteWhen": ["narrowSubgroupRange"], | |
| "derive": { | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE", | |
| "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE", | |
| "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE", | |
| "materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM", | |
| "materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE", | |
| "materializedSoftmaxCols": "dim(shapes.keyT, 1)", | |
| "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4", | |
| "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"", | |
| "applyTileN": "materializedApplyTileN", | |
| "useSubgroups": "device.features.has(\"subgroups\")" | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "materializedScores", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1)]" | |
| } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "scores", | |
| "name": "MultiHeadAttention.MaterializedScores", | |
| "shader": "attn-materialized-score-f32.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query_f32", "key_f32", "scores", "params_scores"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| }, | |
| { | |
| "id": "softmax", | |
| "name": "MultiHeadAttention.MaterializedSoftmax", | |
| "shader": "attn-materialized-softmax-f32.wgsl.jinja", | |
| "derive": { "cacheVec4": "materializedCachedSoftmaxOk" }, | |
| "bindings": ["scores_softmax", "params_rows"], | |
| "dispatch": { | |
| "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)", | |
| "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "apply", | |
| "name": "MultiHeadAttention.MaterializedApply", | |
| "shader": "attn-materialized-apply-f32.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["scores_f32", "value_f32", "output_f32", "params_apply"], | |
| "dispatch": { | |
| "x": "ceilDiv(headDim, applyTileN)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_materialized_gemm_f32", | |
| "priority": 42, | |
| "when": ["materializedF32Ok", "materializedWideSimdOk", "biasOk", "not materializedGemmFusedSoftmaxWorthIt"], | |
| "derive": { | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE", | |
| "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE", | |
| "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE", | |
| "materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM", | |
| "materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE", | |
| "materializedSoftmaxCols": "dim(shapes.keyT, 1)", | |
| "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4", | |
| "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"", | |
| "applyTileN": "materializedApplyTileN", | |
| "useSubgroups": "device.features.has(\"subgroups\")" | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "materializedScores", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1)]" | |
| } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "scores", | |
| "name": "MultiHeadAttention.MaterializedScoresBias", | |
| "shader": "attn-materialized-score-f32.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query_f32", "key_f32", "bias_f32", "scores", "params_scores"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| }, | |
| { | |
| "id": "softmax", | |
| "name": "MultiHeadAttention.MaterializedSoftmaxBias", | |
| "shader": "attn-materialized-softmax-f32.wgsl.jinja", | |
| "derive": { "cacheVec4": "materializedCachedSoftmaxOk" }, | |
| "bindings": ["scores_softmax", "params_rows"], | |
| "dispatch": { | |
| "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)", | |
| "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "apply", | |
| "name": "MultiHeadAttention.MaterializedApplyBias", | |
| "shader": "attn-materialized-apply-f32.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "params_apply"], | |
| "dispatch": { | |
| "x": "ceilDiv(headDim, applyTileN)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_materialized_gemm_fused_f32", | |
| "priority": 43, | |
| "when": ["materializedF32Ok", "materializedWideSimdOk", "not present.biasT", "materializedGemmFusedSoftmaxWorthIt"], | |
| "demoteWhen": ["narrowSubgroupRange"], | |
| "derive": { | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE", | |
| "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE", | |
| "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE", | |
| "materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM", | |
| "applyTileN": "materializedApplyTileN", | |
| "statSlots": "materializedGemmStatSlots", | |
| "statQuerySeq": "dim(shapes.queryT, 1)" | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "materializedScores", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1)]" | |
| }, | |
| { "id": "materializedRowStats", "dtype": "float32", "shape": "[materializedRowStatsElements]" }, | |
| { "id": "materializedScorePartials", "dtype": "float32", "shape": "[materializedGemmScorePartialElements]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "scores", | |
| "name": "MultiHeadAttention.MaterializedScoresFused", | |
| "shader": "attn-materialized-score-f32.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"", "emitRowStats": true }, | |
| "bindings": ["query_f32", "key_f32", "scores", "scorePartials", "params_scores"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| }, | |
| { | |
| "id": "rowstats", | |
| "name": "MultiHeadAttention.MaterializedGemmRowStatsCombine", | |
| "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja", | |
| "derive": { "maxOnly": true }, | |
| "bindings": ["scorePartials_f32", "rowStats", "params_rows"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": 1, | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "apply", | |
| "name": "MultiHeadAttention.MaterializedApplyFused", | |
| "shader": "attn-materialized-apply-f32.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"", "fusedSoftmax": true }, | |
| "bindings": ["scores_f32", "value_f32", "output_f32", "rowStats_f32", "params_apply"], | |
| "dispatch": { | |
| "x": "ceilDiv(headDim, applyTileN)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_materialized_gemm_fused_f32", | |
| "priority": 43, | |
| "when": ["materializedF32Ok", "materializedWideSimdOk", "biasOk", "materializedGemmFusedSoftmaxWorthIt"], | |
| "derive": { | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "materializedInnerTile": "tunables.MATERIALIZED_INNER_TILE", | |
| "materializedQueryTile": "tunables.MATERIALIZED_QUERY_TILE", | |
| "materializedKeyTile": "tunables.MATERIALIZED_KEY_TILE", | |
| "materializedWorkgroupDim": "tunables.MATERIALIZED_WORKGROUP_DIM", | |
| "applyTileN": "materializedApplyTileN", | |
| "statSlots": "materializedGemmStatSlots", | |
| "statQuerySeq": "dim(shapes.queryT, 1)" | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "materializedScores", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1)]" | |
| }, | |
| { "id": "materializedRowStats", "dtype": "float32", "shape": "[materializedRowStatsElements]" }, | |
| { "id": "materializedScorePartials", "dtype": "float32", "shape": "[materializedGemmScorePartialElements]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "scores", | |
| "name": "MultiHeadAttention.MaterializedScoresBiasFused", | |
| "shader": "attn-materialized-score-f32.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"", "emitRowStats": true }, | |
| "bindings": ["query_f32", "key_f32", "bias_f32", "scores", "scorePartials", "params_scores"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.keyT, 1), tunables.MATERIALIZED_KEY_TILE)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| }, | |
| { | |
| "id": "rowstats", | |
| "name": "MultiHeadAttention.MaterializedGemmRowStatsCombineBias", | |
| "shader": "attn-materialized-rowstats-combine-f32.wgsl.jinja", | |
| "derive": { "maxOnly": true }, | |
| "bindings": ["scorePartials_f32", "rowStats", "params_rows"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1)), (materializedRowStatsWg)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": 1, | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "apply", | |
| "name": "MultiHeadAttention.MaterializedApplyBiasFused", | |
| "shader": "attn-materialized-apply-f32.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"", "fusedSoftmax": true }, | |
| "bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "rowStats_f32", "params_apply"], | |
| "dispatch": { | |
| "x": "ceilDiv(headDim, applyTileN)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), tunables.MATERIALIZED_QUERY_TILE)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_q32_broadcast", | |
| "priority": 30, | |
| "when": ["tensorDtypes.queryT == \"float16\"", "not present.biasT", "flashShapeOk", "(dim(shapes.queryT, 2) / attrs.num_heads) % 4 == 0", "(dim(shapes.queryT, 2) / attrs.num_heads) >= 64", "(dim(shapes.queryT, 2) / attrs.num_heads) <= 256", "dim(shapes.queryT, 1) >= 31", "wave32SubgroupsUsable"], | |
| "requires": { "features": ["subgroups", "shader-f16"] }, | |
| "derive": { | |
| "usesF16": true, | |
| "scalar": "\"f16\"", | |
| "inputElement": "\"vec4<f16>\"", | |
| "outputElement": "\"vec4<f16>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "kStep": 64, | |
| "qkGroups": 16 | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.FlashQ32Broadcast", | |
| "shader": "attn-flash-q32-broadcast.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "output", "params"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), 32)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_q32_shared", | |
| "priority": 29, | |
| "when": ["tensorDtypes.queryT == \"float16\"", "not present.biasT", "flashShapeOk", "(dim(shapes.queryT, 2) / attrs.num_heads) % 4 == 0", "(dim(shapes.queryT, 2) / attrs.num_heads) >= 64", "(dim(shapes.queryT, 2) / attrs.num_heads) <= 256", "dim(shapes.queryT, 1) >= 31", "ceilDiv(dim(shapes.queryT, 1), 64) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "((dim(shapes.queryT, 2) / attrs.num_heads) / 4) * 32 * 16 <= device.limits.maxComputeWorkgroupStorageSize"], | |
| "requires": { "features": ["shader-f16"] }, | |
| "derive": { | |
| "usesF16": true, | |
| "scalar": "\"f16\"", | |
| "inputElement": "\"vec4<f16>\"", | |
| "outputElement": "\"vec4<f16>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4", | |
| "kStep": 32, | |
| "qkGroups": 8, | |
| "qStep": 64 | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.FlashQ32Shared", | |
| "shader": "attn-flash-q32-broadcast.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"", "useSubgroups": "false" }, | |
| "bindings": ["query", "key", "value", "output", "params"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.queryT, 1), 64)", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_attn_bias", | |
| "priority": 0, | |
| "when": ["not present.biasT and fallbackMaskShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "hasKeyLimit": false, | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "kvHidden": "dim(shapes.queryT, 2)", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "workgroupSize": "portableWorkgroupSize" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention", | |
| "shader": "attn-online-scalar.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"], | |
| "dispatch": { | |
| "x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_attn_bias", | |
| "priority": 0, | |
| "when": ["biasOk and fallbackMaskShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "hasKeyLimit": false, | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "kvHidden": "dim(shapes.queryT, 2)", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "workgroupSize": "portableWorkgroupSize" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention", | |
| "shader": "attn-online-scalar.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"], | |
| "dispatch": { | |
| "x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias", | |
| "priority": 0, | |
| "when": ["not present.biasT and fallbackShapeOk and not flashShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "hasKeyLimit": false, | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "kvHidden": "dim(shapes.queryT, 2)", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "workgroupSize": "portableWorkgroupSize" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention", | |
| "shader": "attn-online-scalar.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "output", "params"], | |
| "dispatch": { | |
| "x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias", | |
| "priority": 0, | |
| "when": ["biasOk and fallbackShapeOk and not flashShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "hasKeyLimit": false, | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "outputElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "kvHidden": "dim(shapes.queryT, 2)", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "workgroupSize": "portableWorkgroupSize" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention", | |
| "shader": "attn-online-scalar.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "bias", "output", "params"], | |
| "dispatch": { | |
| "x": "min(dim(shapes.queryT, 1), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", | |
| "y": "attrs.num_heads", | |
| "z": "dim(shapes.queryT, 0)" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash", | |
| "priority": 20, | |
| "when": ["device.features.has(\"subgroups\")", "not present.biasT", "flashShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "combineSubgroups": true, | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-online.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "output", "params"], | |
| "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_flash", | |
| "priority": 20, | |
| "when": ["device.features.has(\"subgroups\")", "biasOk", "flashShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "combineSubgroups": true, | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-online.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "bias", "output", "params"], | |
| "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_nosg", | |
| "priority": 18, | |
| "when": ["true", "not present.biasT", "flashShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "combineSubgroups": false, | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.NoBiasOnlineFlashNoSg", | |
| "shader": "attn-flash-online.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "output", "params"], | |
| "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_flash_nosg", | |
| "priority": 17, | |
| "when": ["true", "biasOk", "flashShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "combineSubgroups": false, | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.BiasOnlineFlashNoSg", | |
| "shader": "attn-flash-online.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "bias", "output", "params"], | |
| "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_attn_bias", | |
| "priority": 20, | |
| "when": ["device.features.has(\"subgroups\")", "not present.biasT", "flashMaskShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "combineSubgroups": true, | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-online.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"], | |
| "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_flash_attn_bias", | |
| "priority": 20, | |
| "when": ["device.features.has(\"subgroups\")", "biasOk", "flashMaskShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "combineSubgroups": true, | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.Flash", | |
| "shader": "attn-flash-online.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"], | |
| "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_flash_attn_bias_nosg", | |
| "priority": 18, | |
| "when": ["true", "not present.biasT", "flashMaskShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "combineSubgroups": false, | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.NoBiasOnlineFlashNoSg", | |
| "shader": "attn-flash-online.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "attn_mask_main", "output", "params__uniform"], | |
| "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_flash_attn_bias_nosg", | |
| "priority": 17, | |
| "when": ["true", "biasOk", "flashMaskShapeOk"], | |
| "derive": { | |
| "hasCausal": true, | |
| "combineSubgroups": false, | |
| "maskElement": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "scalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "inputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "outputElement": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"", | |
| "inputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"", | |
| "qNumHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "qHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "kvHiddenV4": "dim(shapes.queryT, 2) / 4", | |
| "headDim": "dim(shapes.queryT, 2) / attrs.num_heads", | |
| "headDimV4": "(dim(shapes.queryT, 2) / attrs.num_heads) / 4" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "MultiHeadAttention.BiasOnlineFlashNoSg", | |
| "shader": "attn-flash-online.wgsl.jinja", | |
| "derive": { "layout": "\"bsh\"" }, | |
| "bindings": ["query", "key", "value", "attn_mask_main", "bias", "output", "params__uniform"], | |
| "dispatch": { "x": "dim(shapes.queryT, 1)", "y": "attrs.num_heads", "z": "dim(shapes.queryT, 0)" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_no_bias_materialized_sgmat_f32", | |
| "priority": 52, | |
| "when": ["materializedSgmatOk", "not present.biasT", "not materializedFusedSoftmaxWorthIt"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "qNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE", | |
| "materializedSoftmaxCols": "dim(shapes.keyT, 1)", | |
| "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4", | |
| "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"", | |
| "useSubgroups": true | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "materializedScores", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1)]" | |
| } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "scores", | |
| "name": "MultiHeadAttention.MaterializedScoresSgmat", | |
| "shader": "attn-materialized-sgmat-f32.wgsl.jinja", | |
| "derive": { "phase": "\"score\"" }, | |
| "bindings": ["query_f32", "key_f32", "scores", "params_scores"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| }, | |
| { | |
| "id": "softmax", | |
| "name": "MultiHeadAttention.MaterializedSoftmaxSgmat", | |
| "shader": "attn-materialized-softmax-f32.wgsl.jinja", | |
| "derive": { "cacheVec4": "materializedCachedSoftmaxOk" }, | |
| "bindings": ["scores_softmax", "params_rows"], | |
| "dispatch": { | |
| "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)", | |
| "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "apply", | |
| "name": "MultiHeadAttention.MaterializedApplySgmat", | |
| "shader": "attn-materialized-sgmat-f32.wgsl.jinja", | |
| "derive": { "phase": "\"apply\"" }, | |
| "bindings": ["scores_f32", "value_f32", "output_f32", "params_apply"], | |
| "dispatch": { | |
| "x": "ceilDiv(headDim, materializedSgmatKeyTile)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "qkv_bias_materialized_sgmat_f32", | |
| "priority": 52, | |
| "when": ["materializedSgmatOk", "biasOk", "not materializedFusedSoftmaxWorthIt"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "derive": { | |
| "qNumHeads": "attrs.num_heads", | |
| "qHidden": "dim(shapes.queryT, 2)", | |
| "materializedSoftmaxWg": "materializedCachedSoftmaxWg if materializedCachedSoftmaxOk else tunables.MATERIALIZED_SOFTMAX_WORKGROUP_SIZE", | |
| "materializedSoftmaxCols": "dim(shapes.keyT, 1)", | |
| "materializedSoftmaxCols4": "dim(shapes.keyT, 1) / 4", | |
| "softmaxElementType": "\"vec4<f32>\" if materializedCachedSoftmaxOk else \"f32\"", | |
| "useSubgroups": true | |
| }, | |
| "intermediates": [ | |
| { | |
| "id": "materializedScores", | |
| "dtype": "float32", | |
| "shape": "[dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1) * dim(shapes.keyT, 1)]" | |
| } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "scores", | |
| "name": "MultiHeadAttention.MaterializedScoresSgmatBias", | |
| "shader": "attn-materialized-sgmat-f32.wgsl.jinja", | |
| "derive": { "phase": "\"score\"" }, | |
| "bindings": ["query_f32", "key_f32", "bias_f32", "scores", "params_scores"], | |
| "dispatch": { | |
| "x": "ceilDiv(dim(shapes.keyT, 1), materializedSgmatKeyTile)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| }, | |
| { | |
| "id": "softmax", | |
| "name": "MultiHeadAttention.MaterializedSoftmaxSgmatBias", | |
| "shader": "attn-materialized-softmax-f32.wgsl.jinja", | |
| "derive": { "cacheVec4": "materializedCachedSoftmaxOk" }, | |
| "bindings": ["scores_softmax", "params_rows"], | |
| "dispatch": { | |
| "x": "min(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)", | |
| "y": "ceilDiv(dim(shapes.queryT, 0) * attrs.num_heads * dim(shapes.queryT, 1), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "apply", | |
| "name": "MultiHeadAttention.MaterializedApplySgmatBias", | |
| "shader": "attn-materialized-sgmat-f32.wgsl.jinja", | |
| "derive": { "phase": "\"apply\"" }, | |
| "bindings": ["scores_f32", "value_f32", "bias_f32", "output_f32", "params_apply"], | |
| "dispatch": { | |
| "x": "ceilDiv(headDim, materializedSgmatKeyTile)", | |
| "y": "ceilDiv(dim(shapes.queryT, 1), materializedSgmatQueryTile)", | |
| "z": "dim(shapes.queryT, 0) * attrs.num_heads" | |
| } | |
| } | |
| ] | |
| } | |
| ] | |
| } | |