{ "domain": "com.microsoft", "name": "PagedAttention", "sinceVersion": 1, "inputs": { "queryT": { "onnx": "query", "dtype": "T", "rank": 2 }, "keyT": { "onnx": "key", "dtype": "T", "rank": 2, "optional": true }, "valueT": { "onnx": "value", "dtype": "T", "rank": 2, "optional": true }, "keyCacheT": { "onnx": "key_cache", "dtype": "T", "rank": 4 }, "valueCacheT": { "onnx": "value_cache", "dtype": "T", "rank": 4 }, "cumulativeSequenceLengthT": { "onnx": "cumulative_sequence_length", "dtype": "S", "rank": 1, "storage": "int32" }, "pastSeqlensT": { "onnx": "past_seqlens", "dtype": "S", "rank": 1, "storage": "int32" }, "blockTableT": { "onnx": "block_table", "dtype": "S", "rank": 2, "storage": "int32" }, "slotMappingT": { "onnx": "slot_mapping", "dtype": "S", "rank": 1, "optional": true, "storage": "int32" } }, "outputs": { "outputT": { "onnx": "output", "dtype": "T", "rank": 2, "shape": "[dim(shapes.queryT, 0), attrs.num_heads * headSize]" }, "keyCacheT": { "onnx": "key_cache", "dtype": "T", "rank": 4, "optional": true, "shape": "shapes.keyCacheT" }, "valueCacheT": { "onnx": "value_cache", "dtype": "T", "rank": 4, "optional": true, "shape": "shapes.valueCacheT" } }, "attributes": { "is_causal": { "default": 1 }, "kv_num_heads": {}, "num_heads": {}, "scale": {} }, "attributeConstraints": { "is_causal": { "values": [1] }, "kv_num_heads": { "required": true }, "num_heads": { "required": true } }, "typeConstraints": { "T": ["float16"], "S": ["int32"] }, "tunables": { "WORKGROUP_SIZE": { "default": 64 }, "SCATTER_WORKGROUP_SIZE": { "default": 64 }, "MAX_SPLITS": { "default": 16 }, "SPLIT_MIN_KEYS": { "default": 128 }, "SPLIT_TARGET_WORKGROUPS": { "default": 1024 } }, "derive": { "tokenCount": "dim(shapes.queryT, 0)", "blockSize": "dim(shapes.keyCacheT, 1)", "headSize": "dim(shapes.keyCacheT, 3)", "maxBlocks": "dim(shapes.blockTableT, 1)", "batchSize": "dim(shapes.cumulativeSequenceLengthT, 0) - 1", "qHidden": "attrs.num_heads * headSize", "kvHidden": "attrs.kv_num_heads * headSize", "qPerKv": "attrs.num_heads / max(1, attrs.kv_num_heads)", "headVec": "headSize / 4", "cacheVec4Ok": "headSize % 4 == 0", "packedStride": "qHidden + 2 * kvHidden", "cacheShapeOk": "ranks.keyCacheT == 4 and ranks.valueCacheT == 4 and sameShape(shapes.valueCacheT, shapes.keyCacheT) and dim(shapes.keyCacheT, 2) == attrs.kv_num_heads and blockSize > 0 and headSize > 0 and tensorDtypes.keyCacheT == tensorDtypes.queryT and tensorDtypes.valueCacheT == tensorDtypes.queryT", "headLayoutOk": "attrs.num_heads > 0 and attrs.kv_num_heads > 0 and attrs.num_heads % attrs.kv_num_heads == 0", "scheduleShapeOk": "ranks.cumulativeSequenceLengthT == 1 and ranks.pastSeqlensT == 1 and ranks.blockTableT == 2 and batchSize >= 1 and dim(shapes.pastSeqlensT, 0) == batchSize and dim(shapes.blockTableT, 0) == batchSize and maxBlocks >= 1", "ioShapeOk": "ranks.queryT == 2 and ranks.outputT == 2 and dim(shapes.outputT, 0) == tokenCount and dim(shapes.outputT, 1) == qHidden and tensorDtypes.outputT == tensorDtypes.queryT and f16Ok(tensorDtypes.queryT)", "separateKv": "present.keyT and present.valueT and ranks.keyT == 2 and ranks.valueT == 2 and dim(shapes.keyT, 0) == tokenCount and dim(shapes.valueT, 0) == tokenCount and dim(shapes.keyT, 1) == kvHidden and dim(shapes.valueT, 1) == kvHidden and tensorDtypes.keyT == tensorDtypes.queryT and tensorDtypes.valueT == tensorDtypes.queryT and dim(shapes.queryT, 1) == qHidden", "packedKv": "not present.keyT and not present.valueT and dim(shapes.queryT, 1) == packedStride", "slotShapeOk": "dim(shapes.slotMappingT, 0) == tokenCount and ranks.slotMappingT == 1 if present.slotMappingT else true", "pagedContractOk": "cacheShapeOk and headLayoutOk and scheduleShapeOk and ioShapeOk and slotShapeOk", "dispatchFits": "tokenCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.num_heads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and tunables.WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.SCATTER_WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup", "pagedBaseWorkgroups": "tokenCount * attrs.kv_num_heads", "pagedAttnQueryFloats": "2 * headSize * qPerKv", "pagedAttnLaneFloats": "3 * qPerKv + 1", "pagedAttnWgBudget": "floor(device.limits.maxComputeWorkgroupStorageSize / 4 - pagedAttnQueryFloats) / pagedAttnLaneFloats", "pagedAttnWg": "max(32, min(tunables.WORKGROUP_SIZE, pow2ceil(max(1, floor(pagedAttnWgBudget)) + 1) / 2))", "pagedAttnStorageBytes": "(pagedAttnQueryFloats + pagedAttnLaneFloats * pagedAttnWg) * 4", "pagedAttnStorageFits": "pagedAttnStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and pagedAttnWg <= device.limits.maxComputeInvocationsPerWorkgroup and pagedAttnWg <= device.limits.maxComputeWorkgroupSizeX", "pagedKeyCeiling": "maxBlocks * blockSize", "pagedSplitWg": "min(256, max(32, pow2ceil(headVec)))", "pagedSplitStorageBytes": "(2 * headSize * qPerKv + (3 * qPerKv + 1) * pagedSplitWg) * 4", "pagedSplitStorageFits": "pagedSplitStorageBytes <= device.limits.maxComputeWorkgroupStorageSize", "pagedSplitCap": "max(1, min(tunables.MAX_SPLITS, pagedKeyCeiling / tunables.SPLIT_MIN_KEYS))", "numSplits": "max(1, min(pagedSplitCap, ceilDiv(tunables.SPLIT_TARGET_WORKGROUPS, max(1, pagedBaseWorkgroups))))", "pagedSplitScratchBytes": "tokenCount * attrs.num_heads * numSplits * headSize * 4", "pagedSplitStatsBytes": "2 * tokenCount * attrs.num_heads * numSplits * 4", "pagedSplitFits": "numSplits >= 2 and cacheVec4Ok and headVec >= 1 and pagedSplitScratchBytes <= device.limits.maxStorageBufferBindingSize and pagedSplitScratchBytes <= device.limits.maxBufferSize and pagedSplitStatsBytes <= device.limits.maxStorageBufferBindingSize and pagedSplitStatsBytes <= device.limits.maxBufferSize and numSplits <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and headVec <= device.limits.maxComputeInvocationsPerWorkgroup and headVec <= device.limits.maxComputeWorkgroupSizeX and pagedSplitWg <= device.limits.maxComputeInvocationsPerWorkgroup and pagedSplitWg <= device.limits.maxComputeWorkgroupSizeX", "aScalar": "dtypes.T", "scalar": "dtypes.T", "numHeads": "attrs.num_heads", "kvNumHeads": "attrs.kv_num_heads", "packedQkv": "not present.keyT", "hasSlotMapping": "present.slotMappingT", "attnWorkgroup": "pagedAttnWg", "cacheVec": "(\"vec4\" if dtypes.T == \"f16\" else \"vec4\") if cacheVec4Ok else dtypes.T", "headDimV4": "headVec", "headDim": "headSize", "qNumHeads": "attrs.num_heads", "qHiddenV4": "qHidden / 4", "hasBias": false, "outVec": "\"vec4\" if dtypes.T == \"f16\" else \"vec4\"" }, "when": ["pagedContractOk", "dispatchFits"], "bindings": { "key": { "arg": "keyT", "buffer": "read-only-storage", "elementType": "$aScalar" }, "value": { "arg": "valueT", "buffer": "read-only-storage", "elementType": "$aScalar" }, "key_cache": { "arg": "keyCacheT", "buffer": "storage", "elementType": "$aScalar" }, "value_cache": { "arg": "valueCacheT", "buffer": "storage", "elementType": "$aScalar" }, "cumulative_sequence_length": { "arg": "cumulativeSequenceLengthT", "buffer": "read-only-storage", "elementType": "i32" }, "past_seqlens": { "arg": "pastSeqlensT", "buffer": "read-only-storage", "elementType": "i32" }, "block_table": { "arg": "blockTableT", "buffer": "read-only-storage", "elementType": "i32" }, "params": { "buffer": "uniform", "struct": [ { "name": "batchSize", "type": "u32", "value": "batchSize" }, { "name": "scatterCount", "type": "u32", "value": "tokenCount * attrs.kv_num_heads * headSize" } ] }, "slot_mapping": { "arg": "slotMappingT", "buffer": "read-only-storage", "elementType": "i32" }, "params_2": { "name": "params", "buffer": "uniform", "struct": [{ "name": "scatterCount", "type": "u32", "value": "tokenCount * attrs.kv_num_heads * headSize" }] }, "query": { "arg": "queryT", "buffer": "read-only-storage", "elementType": "$aScalar" }, "key_cache_2": { "arg": "keyCacheT", "name": "key_cache", "buffer": "read-only-storage", "elementType": "$cacheVec" }, "value_cache_2": { "arg": "valueCacheT", "name": "value_cache", "buffer": "read-only-storage", "elementType": "$cacheVec" }, "params_3": { "name": "params", "buffer": "uniform", "struct": [ { "name": "batchSize", "type": "u32", "value": "batchSize" }, { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" } ] } }, "variants": [ { "id": "separate_derived_splitk", "priority": 10, "when": ["pagedSplitFits", "pagedSplitStorageFits", "separateKv", "not present.slotMappingT"], "requires": { "features": ["shader-f16"] }, "derive": { "attnWorkgroup": "pagedSplitWg" }, "intermediates": [ { "id": "partialOut", "dtype": "float32", "shape": "[tokenCount * attrs.num_heads * numSplits * headSize]" }, { "id": "partialStats", "dtype": "float32", "shape": "[2 * tokenCount * attrs.num_heads * numSplits]" } ], "passes": [ { "id": "scatter", "name": "PagedAttention.ScatterKV", "shader": "paged-scatter-kv.wgsl.jinja", "bindings": ["key", "value", "key_cache", "value_cache", "cumulative_sequence_length", "past_seqlens", "block_table", "params"], "dispatch": { "x": "min(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "y": "ceilDiv(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "z": 1 } }, { "id": "split_attention", "name": "PagedAttention.AttendSplitK", "shader": "paged-attention.wgsl.jinja", "derive": { "splitK": true }, "bindings": [ "query", "key_cache_2", "value_cache_2", "cumulative_sequence_length", "past_seqlens", "block_table", { "scratch": "partialOut", "name": "partial_out", "elementType": "vec4" }, { "scratch": "partialStats", "name": "partial_stats", "elementType": "vec2" }, "params_3" ], "dispatch": { "x": "tokenCount", "y": "attrs.kv_num_heads", "z": "numSplits" } }, { "id": "merge", "name": "PagedAttention.AttendSplitKMerge", "shader": "attn-flash-decode-splitk-merge.wgsl.jinja", "derive": { "layout": "\"bsh\"" }, "bindings": [ { "scratch": "partialOut", "name": "partial_out", "buffer": "read-only-storage", "elementType": "vec4" }, { "scratch": "partialStats", "name": "partial_stats", "buffer": "read-only-storage", "elementType": "vec2" }, { "arg": "outputT", "name": "output", "elementType": "$outVec" } ], "dispatch": { "x": 1, "y": "attrs.num_heads", "z": "tokenCount" } } ] }, { "id": "separate_derived", "when": ["pagedAttnStorageFits", "separateKv", "not present.slotMappingT"], "requires": { "features": ["shader-f16"] }, "passes": [ { "id": "scatter", "name": "PagedAttention.ScatterKV", "shader": "paged-scatter-kv.wgsl.jinja", "bindings": ["key", "value", "key_cache", "value_cache", "cumulative_sequence_length", "past_seqlens", "block_table", "params"], "dispatch": { "x": "min(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "y": "ceilDiv(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "z": 1 } }, { "id": "main", "name": "PagedAttention.Attend", "shader": "paged-attention.wgsl.jinja", "bindings": [ "query", "key_cache_2", "value_cache_2", "cumulative_sequence_length", "past_seqlens", "block_table", { "arg": "outputT", "name": "output", "elementType": "$aScalar" }, "params_3" ], "dispatch": { "x": "tokenCount", "y": "attrs.kv_num_heads" } } ] }, { "id": "separate_slot_splitk", "priority": 10, "when": ["pagedSplitFits", "pagedSplitStorageFits", "separateKv", "present.slotMappingT"], "requires": { "features": ["shader-f16"] }, "derive": { "attnWorkgroup": "pagedSplitWg" }, "intermediates": [ { "id": "partialOut", "dtype": "float32", "shape": "[tokenCount * attrs.num_heads * numSplits * headSize]" }, { "id": "partialStats", "dtype": "float32", "shape": "[2 * tokenCount * attrs.num_heads * numSplits]" } ], "passes": [ { "id": "scatter", "name": "PagedAttention.ScatterKV", "shader": "paged-scatter-kv.wgsl.jinja", "bindings": ["key", "value", "key_cache", "value_cache", "slot_mapping", "params_2"], "dispatch": { "x": "min(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "y": "ceilDiv(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "z": 1 } }, { "id": "split_attention", "name": "PagedAttention.AttendSplitK", "shader": "paged-attention.wgsl.jinja", "derive": { "splitK": true }, "bindings": [ "query", "key_cache_2", "value_cache_2", "cumulative_sequence_length", "past_seqlens", "block_table", { "scratch": "partialOut", "name": "partial_out", "elementType": "vec4" }, { "scratch": "partialStats", "name": "partial_stats", "elementType": "vec2" }, "params_3" ], "dispatch": { "x": "tokenCount", "y": "attrs.kv_num_heads", "z": "numSplits" } }, { "id": "merge", "name": "PagedAttention.AttendSplitKMerge", "shader": "attn-flash-decode-splitk-merge.wgsl.jinja", "derive": { "layout": "\"bsh\"" }, "bindings": [ { "scratch": "partialOut", "name": "partial_out", "buffer": "read-only-storage", "elementType": "vec4" }, { "scratch": "partialStats", "name": "partial_stats", "buffer": "read-only-storage", "elementType": "vec2" }, { "arg": "outputT", "name": "output", "elementType": "$outVec" } ], "dispatch": { "x": 1, "y": "attrs.num_heads", "z": "tokenCount" } } ] }, { "id": "separate_slot", "when": ["pagedAttnStorageFits", "separateKv", "present.slotMappingT"], "requires": { "features": ["shader-f16"] }, "passes": [ { "id": "scatter", "name": "PagedAttention.ScatterKV", "shader": "paged-scatter-kv.wgsl.jinja", "bindings": ["key", "value", "key_cache", "value_cache", "slot_mapping", "params_2"], "dispatch": { "x": "min(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "y": "ceilDiv(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "z": 1 } }, { "id": "main", "name": "PagedAttention.Attend", "shader": "paged-attention.wgsl.jinja", "bindings": [ "query", "key_cache_2", "value_cache_2", "cumulative_sequence_length", "past_seqlens", "block_table", { "arg": "outputT", "name": "output", "elementType": "$aScalar" }, "params_3" ], "dispatch": { "x": "tokenCount", "y": "attrs.kv_num_heads" } } ] }, { "id": "packed_derived_splitk", "priority": 10, "when": ["pagedSplitFits", "pagedSplitStorageFits", "packedKv", "not present.slotMappingT"], "requires": { "features": ["shader-f16"] }, "derive": { "attnWorkgroup": "pagedSplitWg" }, "intermediates": [ { "id": "partialOut", "dtype": "float32", "shape": "[tokenCount * attrs.num_heads * numSplits * headSize]" }, { "id": "partialStats", "dtype": "float32", "shape": "[2 * tokenCount * attrs.num_heads * numSplits]" } ], "passes": [ { "id": "scatter", "name": "PagedAttention.ScatterKV", "shader": "paged-scatter-kv.wgsl.jinja", "bindings": ["query", "key_cache", "value_cache", "cumulative_sequence_length", "past_seqlens", "block_table", "params"], "dispatch": { "x": "min(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "y": "ceilDiv(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "z": 1 } }, { "id": "split_attention", "name": "PagedAttention.AttendSplitK", "shader": "paged-attention.wgsl.jinja", "derive": { "splitK": true }, "bindings": [ "query", "key_cache_2", "value_cache_2", "cumulative_sequence_length", "past_seqlens", "block_table", { "scratch": "partialOut", "name": "partial_out", "elementType": "vec4" }, { "scratch": "partialStats", "name": "partial_stats", "elementType": "vec2" }, "params_3" ], "dispatch": { "x": "tokenCount", "y": "attrs.kv_num_heads", "z": "numSplits" } }, { "id": "merge", "name": "PagedAttention.AttendSplitKMerge", "shader": "attn-flash-decode-splitk-merge.wgsl.jinja", "derive": { "layout": "\"bsh\"" }, "bindings": [ { "scratch": "partialOut", "name": "partial_out", "buffer": "read-only-storage", "elementType": "vec4" }, { "scratch": "partialStats", "name": "partial_stats", "buffer": "read-only-storage", "elementType": "vec2" }, { "arg": "outputT", "name": "output", "elementType": "$outVec" } ], "dispatch": { "x": 1, "y": "attrs.num_heads", "z": "tokenCount" } } ] }, { "id": "packed_derived", "when": ["pagedAttnStorageFits", "packedKv", "not present.slotMappingT"], "requires": { "features": ["shader-f16"] }, "passes": [ { "id": "scatter", "name": "PagedAttention.ScatterKV", "shader": "paged-scatter-kv.wgsl.jinja", "bindings": ["query", "key_cache", "value_cache", "cumulative_sequence_length", "past_seqlens", "block_table", "params"], "dispatch": { "x": "min(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "y": "ceilDiv(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "z": 1 } }, { "id": "main", "name": "PagedAttention.Attend", "shader": "paged-attention.wgsl.jinja", "bindings": [ "query", "key_cache_2", "value_cache_2", "cumulative_sequence_length", "past_seqlens", "block_table", { "arg": "outputT", "name": "output", "elementType": "$aScalar" }, "params_3" ], "dispatch": { "x": "tokenCount", "y": "attrs.kv_num_heads" } } ] }, { "id": "packed_slot_splitk", "priority": 10, "when": ["pagedSplitFits", "pagedSplitStorageFits", "packedKv", "present.slotMappingT"], "requires": { "features": ["shader-f16"] }, "derive": { "attnWorkgroup": "pagedSplitWg" }, "intermediates": [ { "id": "partialOut", "dtype": "float32", "shape": "[tokenCount * attrs.num_heads * numSplits * headSize]" }, { "id": "partialStats", "dtype": "float32", "shape": "[2 * tokenCount * attrs.num_heads * numSplits]" } ], "passes": [ { "id": "scatter", "name": "PagedAttention.ScatterKV", "shader": "paged-scatter-kv.wgsl.jinja", "bindings": ["query", "key_cache", "value_cache", "slot_mapping", "params_2"], "dispatch": { "x": "min(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "y": "ceilDiv(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "z": 1 } }, { "id": "split_attention", "name": "PagedAttention.AttendSplitK", "shader": "paged-attention.wgsl.jinja", "derive": { "splitK": true }, "bindings": [ "query", "key_cache_2", "value_cache_2", "cumulative_sequence_length", "past_seqlens", "block_table", { "scratch": "partialOut", "name": "partial_out", "elementType": "vec4" }, { "scratch": "partialStats", "name": "partial_stats", "elementType": "vec2" }, "params_3" ], "dispatch": { "x": "tokenCount", "y": "attrs.kv_num_heads", "z": "numSplits" } }, { "id": "merge", "name": "PagedAttention.AttendSplitKMerge", "shader": "attn-flash-decode-splitk-merge.wgsl.jinja", "derive": { "layout": "\"bsh\"" }, "bindings": [ { "scratch": "partialOut", "name": "partial_out", "buffer": "read-only-storage", "elementType": "vec4" }, { "scratch": "partialStats", "name": "partial_stats", "buffer": "read-only-storage", "elementType": "vec2" }, { "arg": "outputT", "name": "output", "elementType": "$outVec" } ], "dispatch": { "x": 1, "y": "attrs.num_heads", "z": "tokenCount" } } ] }, { "id": "packed_slot", "when": ["pagedAttnStorageFits", "packedKv", "present.slotMappingT"], "requires": { "features": ["shader-f16"] }, "passes": [ { "id": "scatter", "name": "PagedAttention.ScatterKV", "shader": "paged-scatter-kv.wgsl.jinja", "bindings": ["query", "key_cache", "value_cache", "slot_mapping", "params_2"], "dispatch": { "x": "min(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "y": "ceilDiv(ceilDiv((tokenCount * attrs.kv_num_heads * headSize), (tunables.SCATTER_WORKGROUP_SIZE)), 65535)", "z": 1 } }, { "id": "main", "name": "PagedAttention.Attend", "shader": "paged-attention.wgsl.jinja", "bindings": [ "query", "key_cache_2", "value_cache_2", "cumulative_sequence_length", "past_seqlens", "block_table", { "arg": "outputT", "name": "output", "elementType": "$aScalar" }, "params_3" ], "dispatch": { "x": "tokenCount", "y": "attrs.kv_num_heads" } } ] } ] }