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