{ "domain": "com.microsoft", "name": "MRotaryEmbedding", "sinceVersion": 1, "inputs": { "x": { "onnx": "input", "dtype": "T" }, "positionIds": { "onnx": "position_ids", "dtype": "M", "rank": 3, "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 }, "mrope_layout": { "default": 0 }, "num_heads": { "default": 0 }, "rotary_embedding_dim": { "default": 0 }, "scale": { "default": 1 }, "mrope_section": {} }, "attributeConstraints": { "interleaved": { "values": [0, 1] }, "is_packed_batching": { "values": [0] }, "mrope_layout": { "values": [0, 1] }, "mrope_section": { "required": true } }, "typeConstraints": { "T": ["float32", "float16"], "M": ["int64"] }, "tunables": { "WORKGROUP_SIZE": { "default": 256 } }, "derive": { "rank3HeadSize": "dim(shapes.x, 2) / attrs.num_heads if attrs.num_heads is defined and attrs.num_heads > 0 else 0", "headSize": "rank3HeadSize if ranks.x == 3 else dim(shapes.x, 3)", "effectiveRotaryDim": "attrs.rotary_embedding_dim if attrs.rotary_embedding_dim != 0 else headSize", "tasksPerHead": "ceilDiv(headSize, 2)", "pairCount": "(numel(shapes.x) / max(1, headSize)) * tasksPerHead", "pairDispatchOk": "ceilDiv(pairCount, tunables.WORKGROUP_SIZE) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "halfRotaryDim": "dim(shapes.cos, 1)", "sectionsDefined": "attrs.mrope_section is defined and (attrs.mrope_section | length) == 3", "sectionsValid": "sectionsDefined and attrs.mrope_section[0] >= 0 and attrs.mrope_section[1] >= 0 and attrs.mrope_section[2] >= 0 and attrs.mrope_section[0] + attrs.mrope_section[1] + attrs.mrope_section[2] == halfRotaryDim", "commonContract": "f16Ok(dtypes.T) and sameShape(shapes.x, shapes.y) and (attrs.interleaved == 0 or attrs.interleaved == 1) and (attrs.mrope_layout == 0 or attrs.mrope_layout == 1) and attrs.num_heads >= 0 and attrs.rotary_embedding_dim >= 0 and (attrs.rotary_embedding_dim == 0 or attrs.num_heads > 0) and sameShape(shapes.cos, shapes.sin) and ranks.cos == 2 and ranks.sin == 2 and sectionsValid and ranks.positionIds == 3 and dim(shapes.positionIds, 0) == 3 and dim(shapes.positionIds, 1) == dim(shapes.x, 0) and dim(shapes.positionIds, 2) == dim(shapes.x, 1 if ranks.x == 3 else 2)", "rank3Contract": "commonContract and ranks.x == 3 and ranks.y == 3 and attrs.num_heads is defined and attrs.num_heads >= 1 and dim(shapes.x, 2) % attrs.num_heads == 0 and rank3HeadSize > 0 and effectiveRotaryDim % 2 == 0 and (attrs.rotary_embedding_dim == 0 or attrs.rotary_embedding_dim <= rank3HeadSize) and halfRotaryDim * 2 == effectiveRotaryDim", "rank4Contract": "commonContract and ranks.x == 4 and ranks.y == 4 and dim(shapes.x, 3) > 0 and effectiveRotaryDim % 2 == 0 and (attrs.rotary_embedding_dim == 0 or attrs.rotary_embedding_dim <= dim(shapes.x, 3)) and halfRotaryDim * 2 == effectiveRotaryDim" }, "when": ["pairDispatchOk"], "bindings": { "x": { "buffer": "read-only-storage", "elementType": "$T" }, "position_ids": { "arg": "positionIds", "buffer": "read-only-storage", "elementType": "u32" }, "cos_cache": { "arg": "cos", "buffer": "read-only-storage", "elementType": "$T" }, "sin_cache": { "arg": "sin", "buffer": "read-only-storage", "elementType": "$T" }, "y": { "buffer": "storage", "elementType": "$T" }, "params": { "buffer": "uniform", "struct": [ { "name": "pairCount", "type": "u32", "value": "pairCount" }, { "name": "batchSize", "type": "u32", "value": "dim(shapes.x, 0)" }, { "name": "sequenceLength", "type": "u32", "value": "dim(shapes.x, 1)" }, { "name": "numHeads", "type": "u32", "value": "attrs.num_heads" }, { "name": "headSize", "type": "u32", "value": "rank3HeadSize" }, { "name": "rotaryDim", "type": "u32", "value": "halfRotaryDim * 2" }, { "name": "halfRotaryDim", "type": "u32", "value": "halfRotaryDim" }, { "name": "section0", "type": "u32", "value": "attrs.mrope_section[0]" }, { "name": "section1", "type": "u32", "value": "attrs.mrope_section[1]" }, { "name": "section2", "type": "u32", "value": "attrs.mrope_section[2]" }, { "name": "scale", "type": "f32", "value": "attrs.scale" } ] }, "params_2": { "name": "params", "buffer": "uniform", "struct": [ { "name": "pairCount", "type": "u32", "value": "pairCount" }, { "name": "batchSize", "type": "u32", "value": "dim(shapes.x, 0)" }, { "name": "sequenceLength", "type": "u32", "value": "dim(shapes.x, 2)" }, { "name": "numHeads", "type": "u32", "value": "dim(shapes.x, 1)" }, { "name": "headSize", "type": "u32", "value": "dim(shapes.x, 3)" }, { "name": "rotaryDim", "type": "u32", "value": "halfRotaryDim * 2" }, { "name": "halfRotaryDim", "type": "u32", "value": "halfRotaryDim" }, { "name": "section0", "type": "u32", "value": "attrs.mrope_section[0]" }, { "name": "section1", "type": "u32", "value": "attrs.mrope_section[1]" }, { "name": "section2", "type": "u32", "value": "attrs.mrope_section[2]" }, { "name": "scale", "type": "f32", "value": "attrs.scale" } ] } }, "variants": [ { "id": "rank3", "when": ["rank3Contract"], "derive": { "interleaved": "attrs.interleaved != 0", "mropeSectioned": "attrs.mrope_layout == 0", "usesF16": "dtypes.T == \"f16\"", "scalar": "dtypes.T" }, "passes": [ { "id": "main", "name": "mrotary_embedding3d", "shader": "mrotary-embedding.wgsl.jinja", "derive": { "rank": 3 }, "bindings": ["x", "position_ids", "cos_cache", "sin_cache", "y", "params"], "dispatch": { "x": "min(ceilDiv((pairCount), (tunables.WORKGROUP_SIZE)), 65535)", "y": "ceilDiv(ceilDiv((pairCount), (tunables.WORKGROUP_SIZE)), 65535)", "z": 1 } } ] }, { "id": "rank4", "when": ["rank4Contract"], "derive": { "interleaved": "attrs.interleaved != 0", "mropeSectioned": "attrs.mrope_layout == 0", "usesF16": "dtypes.T == \"f16\"", "scalar": "dtypes.T" }, "passes": [ { "id": "main", "name": "mrotary_embedding4d", "shader": "mrotary-embedding.wgsl.jinja", "derive": { "rank": 4 }, "bindings": ["x", "position_ids", "cos_cache", "sin_cache", "y", "params_2"], "dispatch": { "x": "min(ceilDiv((pairCount), (tunables.WORKGROUP_SIZE)), 65535)", "y": "ceilDiv(ceilDiv((pairCount), (tunables.WORKGROUP_SIZE)), 65535)", "z": 1 } } ] } ] }