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