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