| { |
| "domain": "ai.onnx", |
| "name": "SimplifiedLayerNormalization", |
| "conformance": "legacy-default-domain", |
| "sinceVersion": 1, |
| "inputs": { "x": { "onnx": "X", "dtype": "T" }, "scale": { "dtype": "V" } }, |
| "outputs": { |
| "y": { "onnx": "Y", "dtype": "V", "rank": "ranks.x", "shape": "shapes.x" }, |
| "invStdVar": { |
| "onnx": "inv_std_var", |
| "dtype": "U", |
| "rank": "ranks.x", |
| "optional": true, |
| "shape": "prefix(shapes.x, axisNorm) + fill(1, ranks.x - axisNorm)" |
| } |
| }, |
| "attributes": { |
| "axis": { "default": -1 }, |
| "epsilon": { "default": 0.00001 }, |
| "stash_type": { "default": 1 }, |
| "keep_dims": { "default": 1 } |
| }, |
| "attributeConstraints": { "stash_type": { "values": [1] }, "keep_dims": { "values": [1] } }, |
| "typeConstraints": { "T": ["float32", "float16"], "V": ["float32", "float16"], "U": ["float32"] }, |
| "tunables": { |
| "WORKGROUP_SIZE": { "default": 256 }, |
| "SPLIT_MAX_ROWS": { "default": 256 }, |
| "SPLIT_MIN_HIDDEN": { "default": 16384 }, |
| "SPLIT_TARGET_ELEMENTS": { "default": 4096 }, |
| "MAX_SPLITS": { "default": 64 } |
| }, |
| "derive": { |
| "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)", |
| "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32", |
| "reportedNonWave32Adapter": "not wave32Adapter and (has(device.adapterInfo, \"subgroupMinSize\") or has(device.adapterInfo, \"subgroupMaxSize\"))", |
| "normMaxWorkgroup": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)", |
| "hasSubgroupId": "device.features.has(\"subgroups\") and device.wgslLanguageFeatures.has(\"subgroup_id\")", |
| "axisNorm": "attrs.axis if attrs.axis >= 0 else attrs.axis + ranks.x", |
| "normRows": "outer(shapes.x, axisNorm)", |
| "normHidden": "dim(shapes.x, axisNorm) * inner(shapes.x, axisNorm)", |
| "normRowStride": "max(1, min(normRows, min(device.limits.maxComputeWorkgroupsPerDimension, 65535)))", |
| "rowWg": "min(normMaxWorkgroup, pow2ceil(max(1, normHidden)))", |
| "baseOk": "ranks.x >= 1 and sameShape(shapes.y, shapes.x) and ranks.scale >= 0 and ranks.scale <= ranks.x and broadcastable(shapes.scale, shapes.x) and attrs.axis + ranks.x >= 0 and attrs.axis < ranks.x and normHidden > 0 and attrs.stash_type == onnxDtypeCode(\"float32\") and f16Ok(dtypes.T) and f16Ok(dtypes.V)", |
| "lastAxisOk": "baseOk and (attrs.axis == -1 or attrs.axis == ranks.x - 1)", |
| "suffixAxisOk": "baseOk and ranks.x >= 2 and not (attrs.axis == -1 or attrs.axis == ranks.x - 1)", |
| "noStats": "not present.invStdVar", |
| "statsOk": "present.invStdVar and ranks.invStdVar == ranks.x and sameShape(prefix(shapes.invStdVar, axisNorm), prefix(shapes.x, axisNorm)) and numel(suffix(shapes.invStdVar, axisNorm)) == 1", |
| "sameDtype": "dtypes.T == dtypes.V", |
| "splitCount": "min(tunables.MAX_SPLITS, pow2ceil(ceilDiv(normHidden, tunables.SPLIT_TARGET_ELEMENTS)))", |
| "splitScratchBytes": "normRows * splitCount * 4", |
| "splitFits": "normRows <= tunables.SPLIT_MAX_ROWS and splitCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and splitScratchBytes <= device.limits.maxStorageBufferBindingSize and splitScratchBytes <= device.limits.maxBufferSize" |
| }, |
| "bindings": { |
| "x": { "buffer": "read-only-storage", "elementType": "$xElement" }, |
| "scale": { "buffer": "read-only-storage", "elementType": "$ioElement" }, |
| "y": { "buffer": "storage", "elementType": "$ioElement" }, |
| "params": { |
| "buffer": "uniform", |
| "struct": [ |
| { "name": "rows", "type": "u32", "value": "normRows" }, |
| { "name": "rowStride", "type": "u32", "value": "normRowStride" } |
| ] |
| }, |
| "inv_std_out": { "arg": "invStdVar", "buffer": "storage", "elementType": "f32" }, |
| "partials_2": { "name": "partials", "buffer": "read-only-storage", "elementType": "f32" } |
| }, |
| "variants": [ |
| { |
| "id": "last_axis", |
| "priority": 1, |
| "when": ["lastAxisOk", "noStats"], |
| "derive": { |
| "scalar": "dtypes.V", |
| "xElement": "dtypes.T", |
| "ioElement": "dtypes.V", |
| "hiddenSize": "normHidden", |
| "workgroupSize": "rowWg", |
| "epsilon": "attrs.epsilon" |
| }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "SimplifiedLayerNormalization.Row", |
| "shader": "rms-normalization.wgsl.jinja", |
| "derive": { |
| "xShape": "shapes.x", |
| "scaleShape": "shapes.scale", |
| "xRank": "ranks.x", |
| "scaleRank": "ranks.scale", |
| "writeStats": false, |
| "rmsScaleAfterCast": false |
| }, |
| "bindings": ["x", "scale", "y", "params"], |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 } |
| } |
| ] |
| }, |
| { |
| "id": "last_axis_stats", |
| "priority": 2, |
| "when": ["lastAxisOk", "statsOk"], |
| "derive": { |
| "scalar": "dtypes.V", |
| "xElement": "dtypes.T", |
| "ioElement": "dtypes.V", |
| "hiddenSize": "normHidden", |
| "workgroupSize": "rowWg", |
| "epsilon": "attrs.epsilon" |
| }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "SimplifiedLayerNormalization.Row", |
| "shader": "rms-normalization.wgsl.jinja", |
| "derive": { |
| "xShape": "shapes.x", |
| "scaleShape": "shapes.scale", |
| "xRank": "ranks.x", |
| "scaleRank": "ranks.scale", |
| "writeStats": true, |
| "rmsScaleAfterCast": false |
| }, |
| "bindings": ["x", "scale", "y", "inv_std_out", "params"], |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 } |
| } |
| ] |
| }, |
| { |
| "id": "suffix_axis", |
| "priority": 10, |
| "when": ["suffixAxisOk", "noStats"], |
| "derive": { |
| "scalar": "dtypes.V", |
| "xElement": "dtypes.T", |
| "ioElement": "dtypes.V", |
| "hiddenSize": "normHidden", |
| "workgroupSize": "rowWg", |
| "epsilon": "attrs.epsilon" |
| }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "SimplifiedLayerNormalization.Row", |
| "shader": "rms-normalization.wgsl.jinja", |
| "derive": { |
| "xShape": "shapes.x", |
| "scaleShape": "shapes.scale", |
| "xRank": "ranks.x", |
| "scaleRank": "ranks.scale", |
| "writeStats": false, |
| "rmsScaleAfterCast": false |
| }, |
| "bindings": ["x", "scale", "y", "params"], |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 } |
| } |
| ] |
| }, |
| { |
| "id": "suffix_axis_stats", |
| "priority": 11, |
| "when": ["suffixAxisOk", "statsOk"], |
| "derive": { |
| "scalar": "dtypes.V", |
| "xElement": "dtypes.T", |
| "ioElement": "dtypes.V", |
| "hiddenSize": "normHidden", |
| "workgroupSize": "rowWg", |
| "epsilon": "attrs.epsilon" |
| }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "SimplifiedLayerNormalization.Row", |
| "shader": "rms-normalization.wgsl.jinja", |
| "derive": { |
| "xShape": "shapes.x", |
| "scaleShape": "shapes.scale", |
| "xRank": "ranks.x", |
| "scaleRank": "ranks.scale", |
| "writeStats": true, |
| "rmsScaleAfterCast": false |
| }, |
| "bindings": ["x", "scale", "y", "inv_std_out", "params"], |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 } |
| } |
| ] |
| }, |
| { |
| "id": "suffix_axis_splitk", |
| "priority": 15, |
| "when": ["baseOk", "ranks.x >= 2", "noStats", "splitFits"], |
| "demoteWhen": ["reportedNonWave32Adapter", "normHidden < tunables.SPLIT_MIN_HIDDEN"], |
| "derive": { |
| "scalar": "dtypes.V", |
| "xElement": "dtypes.T", |
| "ioElement": "dtypes.V", |
| "hiddenSize": "normHidden", |
| "workgroupSize": "normMaxWorkgroup", |
| "split": "splitCount", |
| "epsilon": "attrs.epsilon", |
| "normalizeRows": "normRows" |
| }, |
| "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[normRows * splitCount]" }], |
| "passes": [ |
| { |
| "id": "partials", |
| "name": "SimplifiedLayerNormalization.SplitKPartials", |
| "shader": "rms-normalization-splitk-partials.wgsl.jinja", |
| "bindings": ["x", { "name": "partials", "buffer": "storage", "elementType": "f32" }, "params"], |
| "dispatch": { |
| "x": "min(normRows, DISPATCH_FOLD_WIDTH)", |
| "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)", |
| "z": "splitCount" |
| } |
| }, |
| { |
| "id": "normalize", |
| "name": "SimplifiedLayerNormalization.SplitKNormalize", |
| "shader": "rms-normalization-splitk-normalize.wgsl.jinja", |
| "derive": { |
| "xShape": "shapes.x", |
| "scaleShape": "shapes.scale", |
| "xRank": "ranks.x", |
| "scaleRank": "ranks.scale", |
| "writeStats": false, |
| "rmsScaleAfterCast": false, |
| "normalizeBlocks": "max(split, min(min(device.limits.maxComputeWorkgroupsPerDimension, 65535), ceilDiv(workgroupSize, max(1, normalizeRows)), ceilDiv(hiddenSize, workgroupSize * 4)))", |
| "normalizeChunk": "ceilDiv(hiddenSize, normalizeBlocks)" |
| }, |
| "bindings": ["x", "scale", "partials_2", "y", "params"], |
| "dispatch": { |
| "x": "min(normRows, DISPATCH_FOLD_WIDTH)", |
| "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)", |
| "z": "normalizeBlocks" |
| } |
| } |
| ] |
| }, |
| { |
| "id": "suffix_axis_splitk_stats", |
| "priority": 16, |
| "when": ["baseOk", "ranks.x >= 2", "statsOk", "splitFits"], |
| "demoteWhen": ["reportedNonWave32Adapter", "normHidden < tunables.SPLIT_MIN_HIDDEN"], |
| "derive": { |
| "scalar": "dtypes.V", |
| "xElement": "dtypes.T", |
| "ioElement": "dtypes.V", |
| "hiddenSize": "normHidden", |
| "workgroupSize": "normMaxWorkgroup", |
| "split": "splitCount", |
| "epsilon": "attrs.epsilon", |
| "normalizeRows": "normRows" |
| }, |
| "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[normRows * splitCount]" }], |
| "passes": [ |
| { |
| "id": "partials", |
| "name": "SimplifiedLayerNormalization.SplitKPartials", |
| "shader": "rms-normalization-splitk-partials.wgsl.jinja", |
| "bindings": ["x", { "name": "partials", "buffer": "storage", "elementType": "f32" }, "params"], |
| "dispatch": { |
| "x": "min(normRows, DISPATCH_FOLD_WIDTH)", |
| "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)", |
| "z": "splitCount" |
| } |
| }, |
| { |
| "id": "normalize", |
| "name": "SimplifiedLayerNormalization.SplitKNormalize", |
| "shader": "rms-normalization-splitk-normalize.wgsl.jinja", |
| "derive": { |
| "xShape": "shapes.x", |
| "scaleShape": "shapes.scale", |
| "xRank": "ranks.x", |
| "scaleRank": "ranks.scale", |
| "writeStats": true, |
| "rmsScaleAfterCast": false, |
| "normalizeBlocks": "max(split, min(min(device.limits.maxComputeWorkgroupsPerDimension, 65535), ceilDiv(workgroupSize, max(1, normalizeRows)), ceilDiv(hiddenSize, workgroupSize * 4)))", |
| "normalizeChunk": "ceilDiv(hiddenSize, normalizeBlocks)" |
| }, |
| "bindings": ["x", "scale", "partials_2", "y", "inv_std_out", "params"], |
| "dispatch": { |
| "x": "min(normRows, DISPATCH_FOLD_WIDTH)", |
| "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)", |
| "z": "normalizeBlocks" |
| } |
| } |
| ] |
| }, |
| { |
| "id": "last_axis_row_vec4", |
| "priority": 110, |
| "when": ["lastAxisOk", "sameDtype", "noStats", "ranks.scale >= 1", "numel(shapes.scale) == dim(shapes.x, -1)", "dim(shapes.scale, -1) == dim(shapes.x, -1)", "dim(shapes.x, -1) % 4 == 0"], |
| "derive": { |
| "scalar": "dtypes.T", |
| "xElement": "\"vec4<\" ~ dtypes.T ~ \">\"", |
| "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"" |
| }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "SimplifiedLayerNormalization.LastAxisRow", |
| "shader": "norm-row-stats.wgsl.jinja", |
| "derive": { |
| "modeSpec": "\"rms\"", |
| "vec4": true, |
| "writeStats": false, |
| "rmsScaleAfterCast": false, |
| "scalar": "dtypes.T", |
| "usesF16Spec": "dtypes.T == \"f16\"", |
| "hidden": "dim(shapes.x, -1)", |
| "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1) / 4)))", |
| "epsilon": "attrs.epsilon", |
| "hiddenVec": "dim(shapes.x, -1) / 4", |
| "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"", |
| "combineSubgroups": "hasSubgroupId" |
| }, |
| "bindings": ["x", "scale", "y", "params"], |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }, |
| "subgroupCollectivesWidth": "portable" |
| } |
| ] |
| }, |
| { |
| "id": "last_axis_row", |
| "priority": 100, |
| "when": ["lastAxisOk", "sameDtype", "noStats", "ranks.scale >= 1", "numel(shapes.scale) == dim(shapes.x, -1)", "dim(shapes.scale, -1) == dim(shapes.x, -1)"], |
| "derive": { "scalar": "dtypes.T", "xElement": "dtypes.T", "ioElement": "dtypes.T" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "SimplifiedLayerNormalization.LastAxisRow", |
| "shader": "norm-row-stats.wgsl.jinja", |
| "derive": { |
| "modeSpec": "\"rms\"", |
| "vec4": false, |
| "writeStats": false, |
| "rmsScaleAfterCast": false, |
| "scalar": "dtypes.T", |
| "usesF16Spec": "dtypes.T == \"f16\"", |
| "hidden": "dim(shapes.x, -1)", |
| "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1))))", |
| "epsilon": "attrs.epsilon", |
| "hiddenVec": 1, |
| "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"", |
| "combineSubgroups": "hasSubgroupId" |
| }, |
| "bindings": ["x", "scale", "y", "params"], |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }, |
| "subgroupCollectivesWidth": "portable" |
| } |
| ] |
| }, |
| { |
| "id": "last_axis_row_vec4_stats", |
| "priority": 112, |
| "when": ["lastAxisOk", "sameDtype", "statsOk", "ranks.scale >= 1", "numel(shapes.scale) == dim(shapes.x, -1)", "dim(shapes.scale, -1) == dim(shapes.x, -1)", "dim(shapes.x, -1) % 4 == 0"], |
| "derive": { |
| "scalar": "dtypes.T", |
| "xElement": "\"vec4<\" ~ dtypes.T ~ \">\"", |
| "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"" |
| }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "SimplifiedLayerNormalization.LastAxisRow", |
| "shader": "norm-row-stats.wgsl.jinja", |
| "derive": { |
| "modeSpec": "\"rms\"", |
| "vec4": true, |
| "writeStats": true, |
| "rmsScaleAfterCast": false, |
| "scalar": "dtypes.T", |
| "usesF16Spec": "dtypes.T == \"f16\"", |
| "hidden": "dim(shapes.x, -1)", |
| "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1) / 4)))", |
| "epsilon": "attrs.epsilon", |
| "hiddenVec": "dim(shapes.x, -1) / 4", |
| "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"", |
| "combineSubgroups": "hasSubgroupId" |
| }, |
| "bindings": ["x", "scale", "y", "inv_std_out", "params"], |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }, |
| "subgroupCollectivesWidth": "portable" |
| } |
| ] |
| }, |
| { |
| "id": "last_axis_row_stats", |
| "priority": 102, |
| "when": ["lastAxisOk", "sameDtype", "statsOk", "ranks.scale >= 1", "numel(shapes.scale) == dim(shapes.x, -1)", "dim(shapes.scale, -1) == dim(shapes.x, -1)"], |
| "derive": { "scalar": "dtypes.T", "xElement": "dtypes.T", "ioElement": "dtypes.T" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "SimplifiedLayerNormalization.LastAxisRow", |
| "shader": "norm-row-stats.wgsl.jinja", |
| "derive": { |
| "modeSpec": "\"rms\"", |
| "vec4": false, |
| "writeStats": true, |
| "rmsScaleAfterCast": false, |
| "scalar": "dtypes.T", |
| "usesF16Spec": "dtypes.T == \"f16\"", |
| "hidden": "dim(shapes.x, -1)", |
| "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1))))", |
| "epsilon": "attrs.epsilon", |
| "hiddenVec": 1, |
| "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"", |
| "combineSubgroups": "hasSubgroupId" |
| }, |
| "bindings": ["x", "scale", "y", "inv_std_out", "params"], |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 }, |
| "subgroupCollectivesWidth": "portable" |
| } |
| ] |
| } |
| ] |
| } |
|
|