Download build/webgpu/manifest.json from webgpu-kernels/com.microsoft.SparseAttention: direct link, hf CLI and curl.
- Browser
- Download file 40.3 kB
-
https://huggingface.co/kernels/webgpu-kernels/com.microsoft.SparseAttention/resolve/v1/build/webgpu/manifest.json
- Command line
-
hf download hf://webgpu-kernels/com.microsoft.SparseAttention@v1/build/webgpu/manifest.json
-
curl -L -o manifest.json https://huggingface.co/kernels/webgpu-kernels/com.microsoft.SparseAttention/resolve/v1/build/webgpu/manifest.json
40.3 kB
| { | |
| "domain": "com.microsoft", | |
| "name": "SparseAttention", | |
| "sinceVersion": 1, | |
| "inputs": { | |
| "queryT": { "onnx": "query", "dtype": "T", "rank": 3 }, | |
| "keyT": { "onnx": "key", "dtype": "T", "rank": 3, "optional": true }, | |
| "valueT": { "onnx": "value", "dtype": "T", "rank": 3, "optional": true }, | |
| "pastKeyT": { "onnx": "past_key", "dtype": "T", "rank": 4 }, | |
| "pastValueT": { "onnx": "past_value", "dtype": "T", "rank": 4 }, | |
| "blockRowIndicesT": { "onnx": "block_row_indices", "dtype": "M", "rank": 2, "storage": "int32" }, | |
| "blockColIndicesT": { "onnx": "block_col_indices", "dtype": "M", "rank": 2, "storage": "int32" }, | |
| "totalSequenceLengthT": { "onnx": "total_sequence_length", "dtype": "M", "storage": "int32" }, | |
| "keyTotalSequenceLengthsT": { "onnx": "key_total_sequence_lengths", "dtype": "M", "rank": 1, "storage": "int32" }, | |
| "cosCacheT": { "onnx": "cos_cache", "dtype": "T", "rank": 2, "optional": true }, | |
| "sinCacheT": { "onnx": "sin_cache", "dtype": "T", "rank": 2, "optional": true } | |
| }, | |
| "outputs": { | |
| "outputT": { "onnx": "output", "dtype": "T", "rank": 3, "shape": "[batchSize, seqLen, numHeads * headSize]" }, | |
| "pastKeyT": { "onnx": "past_key", "dtype": "T", "rank": 4, "shape": "shapes.pastKeyT" }, | |
| "pastValueT": { "onnx": "past_value", "dtype": "T", "rank": 4, "shape": "shapes.pastValueT" } | |
| }, | |
| "attributes": { | |
| "do_rotary": { "default": 0 }, | |
| "rotary_interleaved": { "default": 0 }, | |
| "num_heads": {}, | |
| "kv_num_heads": {}, | |
| "sparse_block_size": {}, | |
| "scale": {} | |
| }, | |
| "attributeConstraints": { | |
| "num_heads": { "required": true }, | |
| "kv_num_heads": { "required": true }, | |
| "sparse_block_size": { "required": true } | |
| }, | |
| "typeConstraints": { "T": ["float32", "float16"], "M": ["int32"] }, | |
| "tunables": { | |
| "WORKGROUP_SIZE": { "default": 128 }, | |
| "APPEND_WORKGROUP_SIZE": { "default": 256 }, | |
| "NARROW_MIN_WORKGROUPS": { "default": 1024 }, | |
| "QUERY_TILE": { "default": 4 }, | |
| "V_STAGE_MAX_WORKGROUPS": { "default": 512 }, | |
| "MATRIX_MIN_WORKGROUPS": { "default": 16 } | |
| }, | |
| "derive": { | |
| "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32", | |
| "canPinSubgroupSize32": "device.features.has(\"subgroups\") and device.features.has(\"subgroup-size-control\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize <= 32 and device.adapterInfo.subgroupMaxSize >= 32", | |
| "pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter", | |
| "wave32Effective": "wave32Adapter or pinSubgroupSize32", | |
| "batchSize": "dim(shapes.queryT, 0)", | |
| "seqLen": "dim(shapes.queryT, 1)", | |
| "numHeads": "attrs.num_heads", | |
| "kvNumHeads": "attrs.kv_num_heads", | |
| "sparseBlockSize": "attrs.sparse_block_size", | |
| "headSize": "dim(shapes.pastKeyT, 3)", | |
| "headVec": "headSize / 4", | |
| "sparseRequestedQueryTile": "max(tunables.QUERY_TILE, 8) if device.features.has(\"subgroups\") and wave32Effective and not device.features.has(\"shader-f16\") and ceilDiv(seqLen, 8) * batchSize * numHeads >= tunables.NARROW_MIN_WORKGROUPS else tunables.QUERY_TILE", | |
| "sparseWidthBound": "min(256, max(32, pow2ceil(headVec))) if ceilDiv(seqLen, max(1, min(sparseBlockSize, sparseRequestedQueryTile))) * batchSize * numHeads >= tunables.NARROW_MIN_WORKGROUPS else max(256, tunables.WORKGROUP_SIZE)", | |
| "sparseQueryTileCap": "max(1, floor((device.limits.maxComputeWorkgroupStorageSize / 4 - sparseWidthBound) / (2 * headSize + 3 * sparseWidthBound)))", | |
| "sparseQueryTileWant": "min(sparseRequestedQueryTile, min(sparseBlockSize, sparseQueryTileCap))", | |
| "sparseQueryTile": "1 if seqLen <= 1 else (16 if sparseQueryTileWant >= 16 and seqLen >= 16 else (8 if sparseQueryTileWant >= 8 and seqLen >= 8 else (4 if sparseQueryTileWant >= 4 and seqLen >= 4 else (2 if sparseQueryTileWant >= 2 and seqLen >= 2 else 1))))", | |
| "sparseQueryTiles": "ceilDiv(seqLen, sparseQueryTile)", | |
| "sparseAttnWorkgroups": "sparseQueryTiles * batchSize * numHeads", | |
| "sparseAttnWorkgroup": "min(256, max(32, pow2ceil(headVec))) if sparseAttnWorkgroups >= tunables.NARROW_MIN_WORKGROUPS else tunables.WORKGROUP_SIZE", | |
| "maxCacheSeq": "dim(shapes.pastKeyT, 2)", | |
| "numLayout": "dim(shapes.blockRowIndicesT, 0)", | |
| "maxBlocks": "dim(shapes.blockRowIndicesT, 1) - 1", | |
| "maxNnz": "dim(shapes.blockColIndicesT, 1)", | |
| "packedQkv": "not present.keyT", | |
| "qHidden": "numHeads * headSize", | |
| "kvHidden": "kvNumHeads * headSize", | |
| "packedStride": "(numHeads + 2 * kvNumHeads) * headSize", | |
| "doRotary": "attrs.do_rotary == 1", | |
| "rotaryHalf": "dim(shapes.cosCacheT, 1) if doRotary and present.cosCacheT and ranks.cosCacheT == 2 else 0", | |
| "rotaryDim": "2 * rotaryHalf", | |
| "useRotary": "doRotary and rotaryDim > 0", | |
| "rotaryInterleaved": "attrs.rotary_interleaved == 1", | |
| "qRotaryElements": "batchSize * numHeads * seqLen * headSize", | |
| "cacheShapeOk": "ranks.pastKeyT == 4 and ranks.pastValueT == 4 and dim(shapes.pastKeyT, 0) == batchSize and dim(shapes.pastKeyT, 1) == kvNumHeads and sameShape(shapes.pastValueT, shapes.pastKeyT)", | |
| "queryShapeOk": "dim(shapes.queryT, 2) == (packedStride if packedQkv else qHidden)", | |
| "kvShapeOk": "packedQkv or (present.valueT and ranks.keyT == 3 and ranks.valueT == 3 and dim(shapes.keyT, 0) == batchSize and dim(shapes.keyT, 1) == seqLen and dim(shapes.keyT, 2) == kvHidden and sameShape(shapes.valueT, shapes.keyT) and tensorDtypes.keyT == tensorDtypes.queryT and tensorDtypes.valueT == tensorDtypes.queryT)", | |
| "kvPairOk": "present.keyT == present.valueT", | |
| "rotaryPairOk": "not doRotary or (present.cosCacheT and present.sinCacheT and ranks.cosCacheT == 2 and ranks.sinCacheT == 2 and rotaryHalf % 8 == 0 and rotaryDim <= headSize and sameShape(shapes.sinCacheT, shapes.cosCacheT) and tensorDtypes.cosCacheT == tensorDtypes.queryT and tensorDtypes.sinCacheT == tensorDtypes.queryT)", | |
| "blockIndexShapeOk": "ranks.blockRowIndicesT == 2 and ranks.blockColIndicesT == 2 and dim(shapes.blockColIndicesT, 0) == numLayout and maxBlocks >= 1 and maxNnz >= 0 and maxNnz <= maxBlocks * maxBlocks and tensorDtypes.blockRowIndicesT == \"int32\" and tensorDtypes.blockColIndicesT == \"int32\"", | |
| "scheduleShapeOk": "(ranks.totalSequenceLengthT == 0 or ranks.totalSequenceLengthT == 1) and numel(shapes.totalSequenceLengthT) == 1 and ranks.keyTotalSequenceLengthsT == 1 and dim(shapes.keyTotalSequenceLengthsT, 0) == batchSize and tensorDtypes.totalSequenceLengthT == \"int32\" and tensorDtypes.keyTotalSequenceLengthsT == \"int32\"", | |
| "sparseAttnBaseLdsBytes": "(2 * sparseQueryTile * headSize + (3 * sparseQueryTile + 1) * sparseAttnWorkgroup) * 4", | |
| "geometryOk": "tunables.WORKGROUP_SIZE >= 1 and floor(tunables.WORKGROUP_SIZE) == tunables.WORKGROUP_SIZE and pow2ceil(tunables.WORKGROUP_SIZE) == tunables.WORKGROUP_SIZE and tunables.WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX and tunables.APPEND_WORKGROUP_SIZE >= 1 and floor(tunables.APPEND_WORKGROUP_SIZE) == tunables.APPEND_WORKGROUP_SIZE and tunables.APPEND_WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.APPEND_WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX and sparseQueryTiles <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and batchSize * numHeads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(ceilDiv(qRotaryElements, tunables.APPEND_WORKGROUP_SIZE), min(device.limits.maxComputeWorkgroupsPerDimension, 65535)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and sparseAttnBaseLdsBytes <= device.limits.maxComputeWorkgroupStorageSize and sparseAttnWorkgroup <= device.limits.maxComputeInvocationsPerWorkgroup and sparseAttnWorkgroup <= device.limits.maxComputeWorkgroupSizeX", | |
| "contract": "ranks.queryT == 3 and ranks.outputT == 3 and (tensorDtypes.queryT == \"float32\" or tensorDtypes.queryT == \"float16\") and f16Ok(dtypes.T) and tensorDtypes.pastKeyT == tensorDtypes.queryT and tensorDtypes.pastValueT == tensorDtypes.queryT and tensorDtypes.outputT == tensorDtypes.queryT and numHeads >= 1 and kvNumHeads >= 1 and numHeads % kvNumHeads == 0 and headSize >= 8 and headSize % 8 == 0 and (not doRotary or headSize % 16 == 0) and numLayout >= 1 and numHeads % numLayout == 0 and (sparseBlockSize == 16 or sparseBlockSize == 32 or sparseBlockSize == 64 or sparseBlockSize == 128) and cacheShapeOk and queryShapeOk and kvShapeOk and kvPairOk and rotaryPairOk and blockIndexShapeOk and scheduleShapeOk and dim(shapes.outputT, 0) == batchSize and dim(shapes.outputT, 1) == seqLen and dim(shapes.outputT, 2) == qHidden", | |
| "packedContract": "contract and packedQkv and not useRotary", | |
| "packedRotaryContract": "contract and packedQkv and useRotary", | |
| "separateContract": "contract and not packedQkv and not useRotary", | |
| "separateRotaryContract": "contract and not packedQkv and useRotary", | |
| "sparseValueParts": "max(1, min(floor(sparseAttnWorkgroup / headVec), floor((device.limits.maxComputeWorkgroupStorageSize - sparseAttnBaseLdsBytes) / (sparseQueryTile * headSize * 4))))", | |
| "sparseVStageWorthIt": "sparseAttnWorkgroups <= tunables.V_STAGE_MAX_WORKGROUPS and headVec <= sparseAttnWorkgroup and sparseAttnBaseLdsBytes + 16 * headSize * 4 <= device.limits.maxComputeWorkgroupStorageSize", | |
| "sparseSgmatWorkgroup": "256", | |
| "sparseSgmatTileM": "64", | |
| "sgmatQueryTiles": "ceilDiv(seqLen, sparseSgmatTileM)", | |
| "sgmatDirectQuery": "seqLen % sparseSgmatTileM == 0", | |
| "sparseSgmatTileN": "64 if (64 * 32 + 64 * 64 + 64 * 3 + 128 * 2) * 4 <= device.limits.maxComputeWorkgroupStorageSize else 32", | |
| "sparseSgmatTileK": "sparseSgmatTileN / 2", | |
| "sparseSgmatLdsBytes": "(64 * sparseSgmatTileK + 64 * sparseSgmatTileN + 64 * 3 + 128 * 2) * 4", | |
| "sparseSgmatGeometryOk": "sparseSgmatWorkgroup <= device.limits.maxComputeInvocationsPerWorkgroup and sparseSgmatWorkgroup <= device.limits.maxComputeWorkgroupSizeX and sgmatQueryTiles <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and batchSize * numHeads <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and sparseSgmatLdsBytes <= device.limits.maxComputeWorkgroupStorageSize", | |
| "sparseSgmatOk": "tensorDtypes.queryT == \"float32\" and seqLen >= 64 and sparseBlockSize % 64 == 0 and headSize % 32 == 0 and headSize <= 128 and maxCacheSeq % 64 == 0 and device.features.has(\"subgroups\") and wave32Effective and device.features.has(\"chromium-experimental-subgroup-matrix\") and sparseSgmatGeometryOk", | |
| "scalar": "dtypes.T", | |
| "cacheVec": "\"vec4<f16>\" if dtypes.T == \"f16\" else \"vec4<f32>\"", | |
| "attnWorkgroup": "sparseAttnWorkgroup", | |
| "usesRotary": "useRotary", | |
| "appendWorkgroupSize": "tunables.APPEND_WORKGROUP_SIZE", | |
| "sparseTailRows": "seqLen % sparseSgmatTileM", | |
| "sparsePrefixRows": "seqLen - sparseTailRows", | |
| "sparseTailWorkgroup": "min(256, max(32, pow2ceil(headVec))) if batchSize * numHeads >= tunables.NARROW_MIN_WORKGROUPS else tunables.WORKGROUP_SIZE", | |
| "sparseTailBaseLdsBytes": "(2 * sparseTailRows * headSize + (3 * sparseTailRows + 1) * sparseTailWorkgroup) * 4", | |
| "sparseTailValueParts": "max(1, min(floor(sparseTailWorkgroup / headVec), floor((device.limits.maxComputeWorkgroupStorageSize - sparseTailBaseLdsBytes) / (max(1, sparseTailRows) * headSize * 4))))", | |
| "sparseTailVStageWorthIt": "batchSize * numHeads <= tunables.V_STAGE_MAX_WORKGROUPS and headVec <= sparseTailWorkgroup and sparseTailBaseLdsBytes + 16 * headSize * 4 <= device.limits.maxComputeWorkgroupStorageSize", | |
| "sparseTailOk": "sparseTailRows > 0 and sparseTailRows <= sparseQueryTile and sparsePrefixRows >= sparseSgmatTileM and sparseTailBaseLdsBytes <= device.limits.maxComputeWorkgroupStorageSize and sparseTailWorkgroup <= device.limits.maxComputeInvocationsPerWorkgroup and sparseTailWorkgroup <= device.limits.maxComputeWorkgroupSizeX" | |
| }, | |
| "when": ["geometryOk"], | |
| "bindings": { | |
| "new_key": { "arg": "keyT", "elementType": "$scalar" }, | |
| "new_value": { "arg": "valueT", "elementType": "$scalar" }, | |
| "present_key": { "arg": "pastKeyT", "elementType": "$scalar" }, | |
| "present_value": { "arg": "pastValueT", "elementType": "$scalar" }, | |
| "key_total_sequence_lengths": { "arg": "keyTotalSequenceLengthsT", "elementType": "i32" }, | |
| "total_sequence_length": { "arg": "totalSequenceLengthT", "elementType": "i32" }, | |
| "params": { | |
| "struct": [ | |
| { "name": "batchSize", "type": "u32", "value": "batchSize" }, | |
| { "name": "seqLen", "type": "u32", "value": "seqLen" } | |
| ] | |
| }, | |
| "cos_cache": { "arg": "cosCacheT", "elementType": "$scalar" }, | |
| "sin_cache": { "arg": "sinCacheT", "elementType": "$scalar" }, | |
| "packed_qkv": { "arg": "queryT", "elementType": "$scalar" }, | |
| "query": { "arg": "queryT", "elementType": "$scalar" }, | |
| "present_key_packed": { | |
| "arg": "pastKeyT", | |
| "name": "present_key", | |
| "buffer": "read-only-storage", | |
| "elementType": "$cacheVec" | |
| }, | |
| "present_value_packed": { | |
| "arg": "pastValueT", | |
| "name": "present_value", | |
| "buffer": "read-only-storage", | |
| "elementType": "$cacheVec" | |
| }, | |
| "block_row_indices": { "arg": "blockRowIndicesT", "elementType": "i32" }, | |
| "block_col_indices": { "arg": "blockColIndicesT", "elementType": "i32" }, | |
| "output": { "arg": "outputT", "elementType": "$scalar" }, | |
| "params_packed": { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "seqLen", "type": "u32", "value": "seqLen" }, | |
| { "name": "scale", "type": "f32", "value": "attrs.scale if has(attrs, \"scale\") else 0" } | |
| ] | |
| }, | |
| "q_rotary": { "scratch": "QRotary", "buffer": "read-only-storage", "elementType": "f32" }, | |
| "present_key_scalar": { | |
| "arg": "pastKeyT", | |
| "name": "present_key", | |
| "buffer": "read-only-storage", | |
| "elementType": "$scalar" | |
| }, | |
| "present_value_scalar": { | |
| "arg": "pastValueT", | |
| "name": "present_value", | |
| "buffer": "read-only-storage", | |
| "elementType": "$scalar" | |
| }, | |
| "q_rotary_f32": { "scratch": "QRotary", "name": "q_rotary", "elementType": "f32" } | |
| }, | |
| "variants": [ | |
| { | |
| "id": "separate", | |
| "priority": 0, | |
| "when": ["separateContract"], | |
| "derive": { "valueParts": "sparseValueParts", "vStageWorthIt": "sparseVStageWorthIt and valueParts == 1" }, | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"new_key\"", "vSource": "\"new_value\"" }, | |
| "bindings": ["new_key", "new_value", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.Attention", | |
| "shader": "sparse-attention.wgsl.jinja", | |
| "derive": { "qTile": "sparseQueryTile", "queryOffset": "0", "promptTail": "false" }, | |
| "bindings": ["query", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "separate_sgmat", | |
| "priority": 20, | |
| "when": ["separateContract", "sparseSgmatOk"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"new_key\"", "vSource": "\"new_value\"" }, | |
| "bindings": ["new_key", "new_value", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.AttentionSgmat", | |
| "shader": "sparse-attention-sgmat.wgsl.jinja", | |
| "bindings": ["query", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" }, | |
| "derive": { "splitQueryTail": "false" } | |
| } | |
| ], | |
| "demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"] | |
| }, | |
| { | |
| "id": "separate_sgmat_tail", | |
| "priority": 21, | |
| "when": ["separateContract", "sparseSgmatOk", "sparseTailOk"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"new_key\"", "vSource": "\"new_value\"" }, | |
| "bindings": ["new_key", "new_value", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.AttentionSgmat", | |
| "shader": "sparse-attention-sgmat.wgsl.jinja", | |
| "bindings": ["query", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" }, | |
| "derive": { "splitQueryTail": "true" } | |
| }, | |
| { | |
| "id": "tail", | |
| "name": "SparseAttention.Tail", | |
| "shader": "sparse-attention.wgsl.jinja", | |
| "derive": { | |
| "qTile": "sparseTailRows", | |
| "queryOffset": "sparsePrefixRows", | |
| "attnWorkgroup": "sparseTailWorkgroup", | |
| "valueParts": "sparseTailValueParts", | |
| "vStageWorthIt": "sparseTailVStageWorthIt and sparseTailValueParts == 1", | |
| "promptTail": "true" | |
| }, | |
| "bindings": ["query", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "1", "y": "batchSize * numHeads" } | |
| } | |
| ], | |
| "demoteWhen": ["sparsePrefixRows / sparseSgmatTileM * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64", "headSize < sparseSgmatTileM"] | |
| }, | |
| { | |
| "id": "separate_rotary", | |
| "priority": 10, | |
| "when": ["separateRotaryContract"], | |
| "derive": { "valueParts": "sparseValueParts", "vStageWorthIt": "sparseVStageWorthIt and valueParts == 1" }, | |
| "intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }], | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"new_key\"", "vSource": "\"new_value\"" }, | |
| "bindings": ["new_key", "new_value", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "cos_cache", "sin_cache", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "qrotary", | |
| "name": "SparseAttention.QueryRotary", | |
| "shader": "sparse-q-rotary.wgsl.jinja", | |
| "bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.Attention", | |
| "shader": "sparse-attention.wgsl.jinja", | |
| "derive": { "qTile": "sparseQueryTile", "queryOffset": "0", "promptTail": "false" }, | |
| "bindings": ["q_rotary", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "separate_rotary_sgmat", | |
| "priority": 30, | |
| "when": ["separateRotaryContract", "sparseSgmatOk"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }], | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"new_key\"", "vSource": "\"new_value\"" }, | |
| "bindings": ["new_key", "new_value", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "cos_cache", "sin_cache", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "qrotary", | |
| "name": "SparseAttention.QueryRotary", | |
| "shader": "sparse-q-rotary.wgsl.jinja", | |
| "bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.AttentionSgmat", | |
| "shader": "sparse-attention-sgmat.wgsl.jinja", | |
| "bindings": ["q_rotary", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" }, | |
| "derive": { "splitQueryTail": "false" } | |
| } | |
| ], | |
| "demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"] | |
| }, | |
| { | |
| "id": "separate_rotary_sgmat_tail", | |
| "priority": 31, | |
| "when": ["separateRotaryContract", "sparseSgmatOk", "sparseTailOk"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }], | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"new_key\"", "vSource": "\"new_value\"" }, | |
| "bindings": ["new_key", "new_value", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "cos_cache", "sin_cache", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "qrotary", | |
| "name": "SparseAttention.QueryRotary", | |
| "shader": "sparse-q-rotary.wgsl.jinja", | |
| "bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.AttentionSgmat", | |
| "shader": "sparse-attention-sgmat.wgsl.jinja", | |
| "bindings": ["q_rotary", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" }, | |
| "derive": { "splitQueryTail": "true" } | |
| }, | |
| { | |
| "id": "tail", | |
| "name": "SparseAttention.Tail", | |
| "shader": "sparse-attention.wgsl.jinja", | |
| "derive": { | |
| "qTile": "sparseTailRows", | |
| "queryOffset": "sparsePrefixRows", | |
| "attnWorkgroup": "sparseTailWorkgroup", | |
| "valueParts": "sparseTailValueParts", | |
| "vStageWorthIt": "sparseTailVStageWorthIt and sparseTailValueParts == 1", | |
| "promptTail": "true" | |
| }, | |
| "bindings": ["q_rotary", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "1", "y": "batchSize * numHeads" } | |
| } | |
| ], | |
| "demoteWhen": ["sparsePrefixRows / sparseSgmatTileM * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64", "headSize < sparseSgmatTileM"] | |
| }, | |
| { | |
| "id": "packed", | |
| "priority": 0, | |
| "when": ["packedContract"], | |
| "derive": { "valueParts": "sparseValueParts", "vStageWorthIt": "sparseVStageWorthIt and valueParts == 1" }, | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"packed_qkv\"", "vSource": "\"packed_qkv\"" }, | |
| "bindings": ["packed_qkv", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.Attention", | |
| "shader": "sparse-attention.wgsl.jinja", | |
| "derive": { "qTile": "sparseQueryTile", "queryOffset": "0", "promptTail": "false" }, | |
| "bindings": ["query", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "packed_sgmat", | |
| "priority": 20, | |
| "when": ["packedContract", "sparseSgmatOk"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"packed_qkv\"", "vSource": "\"packed_qkv\"" }, | |
| "bindings": ["packed_qkv", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.AttentionSgmat", | |
| "shader": "sparse-attention-sgmat.wgsl.jinja", | |
| "bindings": ["query", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" }, | |
| "derive": { "splitQueryTail": "false" } | |
| } | |
| ], | |
| "demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"] | |
| }, | |
| { | |
| "id": "packed_sgmat_tail", | |
| "priority": 21, | |
| "when": ["packedContract", "sparseSgmatOk", "sparseTailOk"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"packed_qkv\"", "vSource": "\"packed_qkv\"" }, | |
| "bindings": ["packed_qkv", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.AttentionSgmat", | |
| "shader": "sparse-attention-sgmat.wgsl.jinja", | |
| "bindings": ["query", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" }, | |
| "derive": { "splitQueryTail": "true" } | |
| }, | |
| { | |
| "id": "tail", | |
| "name": "SparseAttention.Tail", | |
| "shader": "sparse-attention.wgsl.jinja", | |
| "derive": { | |
| "qTile": "sparseTailRows", | |
| "queryOffset": "sparsePrefixRows", | |
| "attnWorkgroup": "sparseTailWorkgroup", | |
| "valueParts": "sparseTailValueParts", | |
| "vStageWorthIt": "sparseTailVStageWorthIt and sparseTailValueParts == 1", | |
| "promptTail": "true" | |
| }, | |
| "bindings": ["query", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "1", "y": "batchSize * numHeads" } | |
| } | |
| ], | |
| "demoteWhen": ["sparsePrefixRows / sparseSgmatTileM * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64", "headSize < sparseSgmatTileM"] | |
| }, | |
| { | |
| "id": "packed_rotary", | |
| "priority": 10, | |
| "when": ["packedRotaryContract"], | |
| "derive": { "valueParts": "sparseValueParts", "vStageWorthIt": "sparseVStageWorthIt and valueParts == 1" }, | |
| "intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }], | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"packed_qkv\"", "vSource": "\"packed_qkv\"" }, | |
| "bindings": ["packed_qkv", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "cos_cache", "sin_cache", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "qrotary", | |
| "name": "SparseAttention.QueryRotary", | |
| "shader": "sparse-q-rotary.wgsl.jinja", | |
| "bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.Attention", | |
| "shader": "sparse-attention.wgsl.jinja", | |
| "derive": { "qTile": "sparseQueryTile", "queryOffset": "0", "promptTail": "false" }, | |
| "bindings": ["q_rotary", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sparseQueryTiles", "y": "batchSize * numHeads" } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "packed_rotary_sgmat", | |
| "priority": 30, | |
| "when": ["packedRotaryContract", "sparseSgmatOk"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }], | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"packed_qkv\"", "vSource": "\"packed_qkv\"" }, | |
| "bindings": ["packed_qkv", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "cos_cache", "sin_cache", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "qrotary", | |
| "name": "SparseAttention.QueryRotary", | |
| "shader": "sparse-q-rotary.wgsl.jinja", | |
| "bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.AttentionSgmat", | |
| "shader": "sparse-attention-sgmat.wgsl.jinja", | |
| "bindings": ["q_rotary", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" }, | |
| "derive": { "splitQueryTail": "false" } | |
| } | |
| ], | |
| "demoteWhen": ["sgmatQueryTiles * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64"] | |
| }, | |
| { | |
| "id": "packed_rotary_sgmat_tail", | |
| "priority": 31, | |
| "when": ["packedRotaryContract", "sparseSgmatOk", "sparseTailOk"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "intermediates": [{ "id": "QRotary", "dtype": "float32", "shape": "[qRotaryElements]" }], | |
| "passes": [ | |
| { | |
| "id": "append", | |
| "name": "SparseAttention.Append", | |
| "shader": "sparse-kv-append.wgsl.jinja", | |
| "derive": { "kvSource": "\"packed_qkv\"", "vSource": "\"packed_qkv\"" }, | |
| "bindings": ["packed_qkv", "present_key", "present_value", "key_total_sequence_lengths", "total_sequence_length", "cos_cache", "sin_cache", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * kvNumHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "qrotary", | |
| "name": "SparseAttention.QueryRotary", | |
| "shader": "sparse-q-rotary.wgsl.jinja", | |
| "bindings": ["query", "cos_cache", "sin_cache", "q_rotary_f32", "key_total_sequence_lengths", "total_sequence_length", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((batchSize * numHeads * seqLen * headSize), (appendWorkgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "attention", | |
| "name": "SparseAttention.AttentionSgmat", | |
| "shader": "sparse-attention-sgmat.wgsl.jinja", | |
| "bindings": ["q_rotary", "present_key_scalar", "present_value_scalar", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "sgmatQueryTiles", "y": "batchSize * numHeads" }, | |
| "derive": { "splitQueryTail": "true" } | |
| }, | |
| { | |
| "id": "tail", | |
| "name": "SparseAttention.Tail", | |
| "shader": "sparse-attention.wgsl.jinja", | |
| "derive": { | |
| "qTile": "sparseTailRows", | |
| "queryOffset": "sparsePrefixRows", | |
| "attnWorkgroup": "sparseTailWorkgroup", | |
| "valueParts": "sparseTailValueParts", | |
| "vStageWorthIt": "sparseTailVStageWorthIt and sparseTailValueParts == 1", | |
| "promptTail": "true" | |
| }, | |
| "bindings": ["q_rotary", "present_key_packed", "present_value_packed", "block_row_indices", "block_col_indices", "key_total_sequence_lengths", "total_sequence_length", "output", "params_packed"], | |
| "dispatch": { "x": "1", "y": "batchSize * numHeads" } | |
| } | |
| ], | |
| "demoteWhen": ["sparsePrefixRows / sparseSgmatTileM * batchSize * numHeads * sparseSgmatTileN < tunables.MATRIX_MIN_WORKGROUPS * 64", "headSize < sparseSgmatTileM"] | |
| } | |
| ] | |
| } | |