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