Xenova's picture
Xenova HF Staff
sync 6fdf6301e2bb
77f8a08 verified
Raw History Blame
7.07 kB
{
"domain": "com.microsoft",
"name": "RotaryEmbedding",
"sinceVersion": 1,
"inputs": {
"x": { "onnx": "input", "dtype": "T" },
"positionIds": { "onnx": "position_ids", "dtype": "M", "storage": "uint32", "narrowing": "checked" },
"cos": { "onnx": "cos_cache", "dtype": "T", "rank": 2 },
"sin": { "onnx": "sin_cache", "dtype": "T", "rank": 2 }
},
"outputs": { "y": { "onnx": "output", "dtype": "T", "rank": "ranks.x", "shape": "shapes.x" } },
"attributes": {
"interleaved": { "default": 0 },
"is_packed_batching": { "default": 0 },
"num_heads": { "default": 0 },
"rotary_embedding_dim": { "default": 0 },
"scale": { "default": 1 }
},
"attributeConstraints": {
"interleaved": { "values": [0, 1] },
"is_packed_batching": { "values": [0, 1] },
"scale": { "values": [1] }
},
"typeConstraints": { "T": ["float32", "float16"], "M": ["int64"] },
"tunables": { "WORKGROUP_SIZE": { "default": 256 }, "ITEMS_PER_LANE": { "default": 2 } },
"derive": {
"foldedDispatchCapacity": "min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
"rank3": "ranks.x == 3",
"rank4": "ranks.x == 4",
"rankOk": "rank3 or rank4",
"cacheShapeOk": "ranks.cos == 2 and ranks.sin == 2 and sameShape(shapes.cos, shapes.sin)",
"cacheWidth": "dim(shapes.cos, 1) if cacheShapeOk else 0",
"batchSize": "dim(shapes.x, 0) if rankOk else 0",
"seqLength": "(dim(shapes.x, 1) if rank3 else dim(shapes.x, 2)) if rankOk else 0",
"hiddenSize": "dim(shapes.x, 2) if rank3 else 0",
"rank3HeadSize": "(hiddenSize / attrs.num_heads if attrs.num_heads > 0 and hiddenSize % max(1, attrs.num_heads) == 0 else 2 * cacheWidth) if rank3 else 0",
"headSize": "rank3HeadSize if rank3 else (dim(shapes.x, 3) if rank4 else 0)",
"numHeads": "(hiddenSize / headSize if headSize > 0 and hiddenSize % max(1, headSize) == 0 else 0) if rank3 else (dim(shapes.x, 1) if rank4 else 0)",
"rotaryDim": "attrs.rotary_embedding_dim if attrs.rotary_embedding_dim > 0 else headSize",
"halfRotaryDim": "rotaryDim / 2 if rotaryDim % 2 == 0 else 0",
"geometryOk": "rankOk and cacheShapeOk and headSize > 0 and numHeads > 0 and batchSize > 0 and seqLength > 0 and rotaryDim > 0 and rotaryDim % 2 == 0 and rotaryDim <= headSize and cacheWidth == halfRotaryDim",
"attrsOk": "attrs.num_heads >= 0 and attrs.rotary_embedding_dim >= 0 and (attrs.rotary_embedding_dim == 0 or attrs.num_heads > 0)",
"contractOk": "f16Ok(dtypes.T) and geometryOk and attrsOk and sameShape(shapes.x, shapes.y)",
"itemsPerLane": "tunables.ITEMS_PER_LANE if (rank3 or numHeads % max(1, tunables.ITEMS_PER_LANE) == 0) else 1",
"headBlocks": "ceilDiv(numHeads, max(1, itemsPerLane)) if rank4 else 1",
"sliceRows": "(seqLength * numHeads if rank3 else seqLength) if rankOk else 0",
"sliceCount": "(batchSize if rank3 else batchSize * headBlocks) if rankOk else 0",
"sliceDispatchOk": "sliceCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
"workgroupSize": "max(1, min(tunables.WORKGROUP_SIZE, min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)))",
"laneUnits": "workgroupSize * itemsPerLane if rank3 else workgroupSize",
"quadHeadUnits": "headSize / 8 if headSize % 8 == 0 else 0",
"quadHeadStride": "headSize / 4 if headSize % 4 == 0 else 0",
"quadRotStart": "rotaryDim / 4 if rotaryDim % 4 == 0 else 0",
"quadRotUnits": "rotaryDim / 8 if rotaryDim % 8 == 0 else 0",
"quadHasTail": "quadHeadUnits > quadRotUnits",
"quadBlocks": "ceilDiv(sliceRows * quadHeadUnits, laneUnits)",
"quadOk": "headSize % 8 == 0 and rotaryDim % 8 == 0 and quadBlocks <= foldedDispatchCapacity and sliceDispatchOk",
"pairHeadUnits": "ceilDiv(headSize, 2)",
"pairHasTail": "pairHeadUnits > halfRotaryDim",
"pairTailGuard": "headSize % 2 == 1",
"pairBlocks": "ceilDiv(sliceRows * pairHeadUnits, laneUnits)",
"pairOk": "pairBlocks <= foldedDispatchCapacity and sliceDispatchOk",
"posOffsetOk": "ranks.positionIds <= 1 and numel(shapes.positionIds) == 1",
"posTableOk": "ranks.positionIds == 2 and dim(shapes.positionIds, 0) == batchSize and dim(shapes.positionIds, 1) == seqLength",
"interleaved": "attrs.interleaved != 0",
"scalarType": "dtypes.T",
"vectorType": "\"vec4<\" ~ dtypes.T ~ \">\""
},
"when": ["contractOk"],
"bindings": {
"x": { "elementType": "$vectorType" },
"position_ids": { "arg": "positionIds", "elementType": "$M" },
"cos_cache": { "arg": "cos", "elementType": "$vectorType" },
"sin_cache": { "arg": "sin", "elementType": "$vectorType" },
"y": { "elementType": "$vectorType" },
"x_scalar": { "arg": "x", "name": "x", "elementType": "$T" },
"cos_scalar": { "arg": "cos", "name": "cos_cache", "elementType": "$T" },
"sin_scalar": { "arg": "sin", "name": "sin_cache", "elementType": "$T" },
"y_scalar": { "arg": "y", "name": "y", "elementType": "$T" },
"params": { "struct": [{ "name": "sliceRows", "type": "u32", "value": "sliceRows" }] }
},
"variants": [
{
"id": "quad",
"priority": 10,
"when": ["posTableOk or posOffsetOk", "quadOk"],
"derive": {
"rank": "ranks.x",
"posOffset": "posOffsetOk",
"useVec4": "true",
"headUnits": "quadHeadUnits",
"headStride": "quadHeadStride",
"rotStart": "quadRotStart",
"rotUnits": "quadRotUnits",
"hasTail": "quadHasTail",
"tailGuard": "false",
"castIn": "\"vec4<f32>\"",
"castOut": "vectorType"
},
"passes": [
{
"id": "main",
"name": "RotaryEmbedding.Quad",
"shader": "rotary-embedding-slices.wgsl.jinja",
"bindings": ["x", "position_ids", "cos_cache", "sin_cache", "y", "params"],
"dispatch": {
"x": "min(quadBlocks, DISPATCH_FOLD_WIDTH)",
"y": "ceilDiv(quadBlocks, DISPATCH_FOLD_WIDTH)",
"z": "sliceCount"
}
}
]
},
{
"id": "pair",
"priority": 0,
"when": ["posTableOk or posOffsetOk", "pairOk"],
"derive": {
"rank": "ranks.x",
"posOffset": "posOffsetOk",
"useVec4": "false",
"headUnits": "pairHeadUnits",
"headStride": "headSize",
"rotStart": "rotaryDim",
"rotUnits": "halfRotaryDim",
"hasTail": "pairHasTail",
"tailGuard": "pairTailGuard",
"castIn": "\"f32\"",
"castOut": "scalarType"
},
"passes": [
{
"id": "main",
"name": "RotaryEmbedding.Pair",
"shader": "rotary-embedding-slices.wgsl.jinja",
"bindings": ["x_scalar", "position_ids", "cos_scalar", "sin_scalar", "y_scalar", "params"],
"dispatch": {
"x": "min(pairBlocks, DISPATCH_FOLD_WIDTH)",
"y": "ceilDiv(pairBlocks, DISPATCH_FOLD_WIDTH)",
"z": "sliceCount"
}
}
]
}
]
}