{ "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"] } ] }