Xenova's picture
Xenova HF Staff
sync 6fdf6301e2bb
f7870ad verified
Raw History Blame
7.01 kB
{
"domain": "com.microsoft",
"name": "NGramHashMapping",
"sinceVersion": 1,
"inputs": {
"inputIdsT": { "onnx": "input_ids", "dtype": "M", "rank": 2 },
"multipliersT": { "onnx": "multipliers", "dtype": "M", "rank": 1 },
"vocabSizesT": { "onnx": "vocab_sizes", "dtype": "M", "rank": 1 },
"pastIdsT": { "onnx": "past_ids", "dtype": "M", "rank": 2, "optional": true }
},
"outputs": {
"hashIdsT": {
"onnx": "hash_ids",
"dtype": "M",
"rank": 3,
"shape": "[dim(shapes.inputIdsT, 0), dim(shapes.inputIdsT, 1), (attrs.max_ngram_size - 1) * attrs.n_head_per_ngram]"
},
"presentIdsT": {
"onnx": "present_ids",
"dtype": "M",
"rank": 2,
"shape": "[dim(shapes.inputIdsT, 0), attrs.max_ngram_size - 1]"
}
},
"attributes": { "max_ngram_size": {}, "n_head_per_ngram": {}, "pad_id": {} },
"attributeConstraints": {
"max_ngram_size": { "required": true },
"n_head_per_ngram": { "required": true },
"pad_id": { "required": true }
},
"typeConstraints": { "M": ["int32"] },
"tunables": { "WORKGROUP_SIZE": { "default": 256 }, "ITEMS_PER_LANE": { "default": 4 } },
"derive": {
"foldedDispatchCapacity": "min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
"maxNgramSize": "attrs.max_ngram_size",
"headsPerNgram": "attrs.n_head_per_ngram",
"stateLength": "maxNgramSize - 1",
"numHeads": "stateLength * headsPerNgram",
"stagedTableScalars": "maxNgramSize + numHeads",
"padId": "attrs.pad_id",
"idsRankOk": "ranks.inputIdsT == 2",
"batchSize": "dim(shapes.inputIdsT, 0) if idsRankOk else 0",
"seqLength": "dim(shapes.inputIdsT, 1) if idsRankOk else 0",
"spanLength": "seqLength + stateLength",
"sharedTileSlots": "max(1, min(tunables.WORKGROUP_SIZE, spanLength, floor(device.limits.maxComputeWorkgroupStorageSize / (max(1, stateLength) * 4))))",
"sharedSpanTiles": "ceilDiv(spanLength, sharedTileSlots)",
"sharedDispatchOk": "sharedSpanTiles <= foldedDispatchCapacity",
"sharedMemoryOk": "sharedTileSlots * stateLength * 4 <= device.limits.maxComputeWorkgroupStorageSize",
"tileSlots": "tunables.WORKGROUP_SIZE * tunables.ITEMS_PER_LANE",
"spanTiles": "ceilDiv(spanLength, max(1, tileSlots))",
"idScalar": "dtypes.M",
"attrsOk": "maxNgramSize >= 2 and maxNgramSize <= 16 and headsPerNgram >= 1 and numHeads <= 64 and padId >= -2147483648 and padId <= 2147483647",
"dtypeOk": "tensorDtypes.inputIdsT == \"int32\" and tensorDtypes.multipliersT == \"int32\" and tensorDtypes.vocabSizesT == \"int32\" and tensorDtypes.hashIdsT == \"int32\" and tensorDtypes.presentIdsT == \"int32\"",
"shapeOk": "idsRankOk and ranks.multipliersT == 1 and dim(shapes.multipliersT, 0) == maxNgramSize and ranks.vocabSizesT == 1 and dim(shapes.vocabSizesT, 0) == numHeads and ranks.hashIdsT == 3 and sameShape(shapes.hashIdsT, [batchSize, seqLength, numHeads]) and ranks.presentIdsT == 2 and sameShape(shapes.presentIdsT, [batchSize, stateLength])",
"dispatchOk": "batchSize <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and spanTiles <= foldedDispatchCapacity",
"workgroupOk": "tunables.WORKGROUP_SIZE > 0 and tunables.WORKGROUP_SIZE <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.WORKGROUP_SIZE <= device.limits.maxComputeWorkgroupSizeX and tunables.ITEMS_PER_LANE >= 1 and floor(tunables.ITEMS_PER_LANE) == tunables.ITEMS_PER_LANE",
"baseContract": "attrsOk and dtypeOk and shapeOk and dispatchOk and workgroupOk",
"noPastContract": "baseContract and not present.pastIdsT",
"pastContract": "baseContract and present.pastIdsT and tensorDtypes.pastIdsT == \"int32\" and ranks.pastIdsT == 2 and sameShape(shapes.pastIdsT, [batchSize, stateLength])"
},
"bindings": {
"input_ids": { "arg": "inputIdsT", "elementType": "$idScalar" },
"multipliers": { "arg": "multipliersT", "elementType": "$idScalar", "length": "$maxNgramSize" },
"vocab_sizes": { "arg": "vocabSizesT", "elementType": "$idScalar", "length": "$numHeads" },
"past_ids": { "arg": "pastIdsT", "elementType": "$idScalar" },
"hash_ids": { "arg": "hashIdsT", "elementType": "$idScalar" },
"present_ids": { "arg": "presentIdsT", "elementType": "$idScalar" },
"params": { "struct": [{ "name": "seqLength", "type": "u32", "value": "seqLength" }] }
},
"variants": [
{
"id": "past",
"when": ["pastContract"],
"derive": { "hasPast": true },
"passes": [
{
"id": "main",
"name": "NGramHashMapping.Past",
"shader": "ngram-hash-mapping.wgsl.jinja",
"bindings": ["input_ids", "multipliers", "vocab_sizes", "past_ids", "hash_ids", "present_ids", "params"],
"dispatch": {
"x": "min(spanTiles, DISPATCH_FOLD_WIDTH)",
"y": "batchSize",
"z": "ceilDiv(spanTiles, DISPATCH_FOLD_WIDTH)"
}
}
]
},
{
"id": "past_shared_prefix",
"when": ["pastContract", "sharedDispatchOk", "sharedMemoryOk"],
"derive": { "hasPast": true },
"passes": [
{
"id": "main",
"name": "NGramHashMapping.Past.SharedPrefix",
"shader": "ngram-hash-mapping-shared-prefix.wgsl.jinja",
"bindings": ["input_ids", "multipliers", "vocab_sizes", "past_ids", "hash_ids", "present_ids", "params"],
"dispatch": {
"x": "min(sharedSpanTiles, DISPATCH_FOLD_WIDTH)",
"y": "batchSize",
"z": "ceilDiv(sharedSpanTiles, DISPATCH_FOLD_WIDTH)"
}
}
],
"priority": 10,
"demoteWhen": ["stagedTableScalars < 59"]
},
{
"id": "fresh",
"when": ["noPastContract"],
"derive": { "hasPast": false },
"passes": [
{
"id": "main",
"name": "NGramHashMapping.Fresh",
"shader": "ngram-hash-mapping.wgsl.jinja",
"bindings": ["input_ids", "multipliers", "vocab_sizes", "hash_ids", "present_ids", "params"],
"dispatch": {
"x": "min(spanTiles, DISPATCH_FOLD_WIDTH)",
"y": "batchSize",
"z": "ceilDiv(spanTiles, DISPATCH_FOLD_WIDTH)"
}
}
]
},
{
"id": "fresh_shared_prefix",
"when": ["noPastContract", "sharedDispatchOk", "sharedMemoryOk"],
"derive": { "hasPast": false },
"passes": [
{
"id": "main",
"name": "NGramHashMapping.Fresh.SharedPrefix",
"shader": "ngram-hash-mapping-shared-prefix.wgsl.jinja",
"bindings": ["input_ids", "multipliers", "vocab_sizes", "hash_ids", "present_ids", "params"],
"dispatch": {
"x": "min(sharedSpanTiles, DISPATCH_FOLD_WIDTH)",
"y": "batchSize",
"z": "ceilDiv(sharedSpanTiles, DISPATCH_FOLD_WIDTH)"
}
}
],
"priority": 10,
"demoteWhen": ["stagedTableScalars < 59"]
}
]
}