Xenova's picture
Xenova HF Staff
sync 6fdf6301e2bb
ba08609 verified
Raw History Blame
11.6 kB
{
"domain": "ai.onnx",
"name": "RMSNormalization",
"sinceVersion": 23,
"inputs": { "x": { "onnx": "X", "dtype": "T" }, "scale": { "dtype": "V" } },
"outputs": { "y": { "onnx": "Y", "dtype": "V", "rank": "ranks.x", "shape": "shapes.x" } },
"attributes": { "axis": { "default": -1 }, "epsilon": { "default": 0.00001 }, "stash_type": { "default": 1 } },
"attributeConstraints": { "stash_type": { "values": [1, 10] } },
"typeConstraints": { "T": ["float32", "float16"], "V": ["float32", "float16"] },
"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\"))",
"variableSubgroup16To32": "device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 16 and device.adapterInfo.subgroupMaxSize == 32",
"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",
"normalizedRows": "outer(shapes.x, axisNorm)",
"normalizedHidden": "dim(shapes.x, axisNorm) * inner(shapes.x, axisNorm)",
"normalizedDispatchRows": "0 if normalizedHidden == 0 else normalizedRows",
"normalizedWorkgroupHidden": "max(1, normalizedHidden)",
"normalizationShapeOk": "ranks.x >= 1 and ranks.scale >= 0 and ranks.scale <= ranks.x and sameShape(shapes.y, shapes.x) and attrs.axis + ranks.x >= 0 and attrs.axis < ranks.x and broadcastable(shapes.scale, shapes.x) and f16Ok(dtypes.T) and f16Ok(dtypes.V)",
"baseOk": "normalizationShapeOk and attrs.stash_type == onnxDtypeCode(\"float32\")",
"stashF16Ok": "normalizationShapeOk and attrs.stash_type == onnxDtypeCode(\"float16\")",
"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)"
},
"bindings": {
"x": { "elementType": "$xElement" },
"scale": { "elementType": "$ioElement" },
"y": { "elementType": "$ioElement" },
"params": {
"struct": [
{ "name": "rows", "type": "u32", "value": "normalizedRows" },
{
"name": "rowStride",
"type": "u32",
"value": "max(1, min(normalizedRows, min(device.limits.maxComputeWorkgroupsPerDimension, 65535)))"
}
]
},
"params__uniform": {
"name": "params",
"struct": [
{ "name": "rows", "type": "u32", "value": "splitRows" },
{
"name": "rowStride",
"type": "u32",
"value": "max(1, min(splitRows, min(device.limits.maxComputeWorkgroupsPerDimension, 65535)))"
}
]
}
},
"variants": [
{
"id": "stash_f16_serial",
"priority": 1000,
"when": ["stashF16Ok"],
"derive": {
"scalar": "dtypes.V",
"xElement": "dtypes.T",
"ioElement": "dtypes.V",
"hiddenSize": "normalizedHidden",
"epsilon": "attrs.epsilon"
},
"passes": [
{
"id": "main",
"name": "RMSNormalization.StashF16Serial",
"shader": "rms-normalization-stash-f16-serial.wgsl.jinja",
"derive": {
"xShape": "shapes.x",
"scaleShape": "shapes.scale",
"xRank": "ranks.x",
"scaleRank": "ranks.scale"
},
"bindings": ["x", "scale", "y", "params"],
"dispatch": {
"x": "min(normalizedDispatchRows, 65535)",
"y": "ceilDiv(normalizedDispatchRows, 65535)",
"z": 1
}
}
]
},
{
"id": "suffix_axis_splitk",
"priority": 15,
"when": ["baseOk", "ranks.x >= 2", "normalizedRows <= tunables.SPLIT_MAX_ROWS", "normalizedHidden >= tunables.SPLIT_MIN_HIDDEN", "min(tunables.MAX_SPLITS, pow2ceil(ceilDiv(normalizedHidden, tunables.SPLIT_TARGET_ELEMENTS))) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "normalizedRows * min(tunables.MAX_SPLITS, pow2ceil(ceilDiv(normalizedHidden, tunables.SPLIT_TARGET_ELEMENTS))) * 4 <= device.limits.maxStorageBufferBindingSize", "normalizedRows * min(tunables.MAX_SPLITS, pow2ceil(ceilDiv(normalizedHidden, tunables.SPLIT_TARGET_ELEMENTS))) * 4 <= device.limits.maxBufferSize", "not (variableSubgroup16To32 and dtypes.T == \"f16\" and dtypes.V == \"f32\")"],
"demoteWhen": ["reportedNonWave32Adapter and not variableSubgroup16To32"],
"derive": {
"splitRows": "normalizedRows",
"splitHidden": "normalizedHidden",
"split": "min(tunables.MAX_SPLITS, pow2ceil(ceilDiv(splitHidden, tunables.SPLIT_TARGET_ELEMENTS)))",
"scalar": "dtypes.V",
"xElement": "dtypes.T",
"ioElement": "dtypes.V",
"hiddenSize": "splitHidden",
"workgroupSize": "normMaxWorkgroup",
"epsilon": "attrs.epsilon",
"normalizeRows": "splitRows"
},
"intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[splitRows * split]" }],
"passes": [
{
"id": "partials",
"name": "RMSNormalization.SplitKPartials",
"shader": "rms-normalization-splitk-partials.wgsl.jinja",
"bindings": ["x", { "name": "partials", "elementType": "f32" }, "params__uniform"],
"dispatch": {
"x": "min(splitRows, DISPATCH_FOLD_WIDTH)",
"y": "ceilDiv(splitRows, DISPATCH_FOLD_WIDTH)",
"z": "split"
}
},
{
"id": "normalize",
"name": "RMSNormalization.SplitKNormalize",
"shader": "rms-normalization-splitk-normalize.wgsl.jinja",
"derive": {
"xShape": "shapes.x",
"scaleShape": "shapes.scale",
"xRank": "ranks.x",
"scaleRank": "ranks.scale",
"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",
{ "name": "partials", "buffer": "read-only-storage", "elementType": "f32" },
"y",
"params__uniform"
],
"dispatch": {
"x": "min(splitRows, DISPATCH_FOLD_WIDTH)",
"y": "ceilDiv(splitRows, DISPATCH_FOLD_WIDTH)",
"z": "normalizeBlocks"
}
}
]
},
{
"id": "last_axis",
"priority": 0,
"when": ["lastAxisOk"],
"derive": {
"scalar": "dtypes.V",
"xElement": "dtypes.T",
"ioElement": "dtypes.V",
"hiddenSize": "normalizedHidden",
"workgroupSize": "min(normMaxWorkgroup, pow2ceil(normalizedWorkgroupHidden))",
"epsilon": "attrs.epsilon"
},
"passes": [
{
"id": "main",
"name": "RMSNormalization",
"shader": "rms-normalization.wgsl.jinja",
"derive": {
"compensateHalfStats": true,
"xShape": "shapes.x",
"scaleShape": "shapes.scale",
"xRank": "ranks.x",
"scaleRank": "ranks.scale"
},
"bindings": ["x", "scale", "y", "params"],
"dispatch": {
"x": "min(normalizedDispatchRows, 65535)",
"y": "ceilDiv(normalizedDispatchRows, 65535)",
"z": 1
}
}
]
},
{
"id": "suffix_axis",
"priority": 10,
"when": ["suffixAxisOk"],
"derive": {
"scalar": "dtypes.V",
"xElement": "dtypes.T",
"ioElement": "dtypes.V",
"hiddenSize": "normalizedHidden",
"workgroupSize": "min(normMaxWorkgroup, pow2ceil(normalizedWorkgroupHidden))",
"epsilon": "attrs.epsilon"
},
"passes": [
{
"id": "main",
"name": "RMSNormalization.SuffixAxis",
"shader": "rms-normalization.wgsl.jinja",
"derive": {
"compensateHalfStats": true,
"xShape": "shapes.x",
"scaleShape": "shapes.scale",
"xRank": "ranks.x",
"scaleRank": "ranks.scale"
},
"bindings": ["x", "scale", "y", "params"],
"dispatch": {
"x": "min(normalizedDispatchRows, 65535)",
"y": "ceilDiv(normalizedDispatchRows, 65535)",
"z": 1
}
}
]
},
{
"id": "last_axis_row_vec4",
"priority": 110,
"when": ["lastAxisOk", "dtypes.T == dtypes.V", "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": { "xElement": "\"vec4<\" ~ dtypes.T ~ \">\"", "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"" },
"passes": [
{
"id": "main",
"name": "RMSNormalization.LastAxisRow",
"shader": "norm-row-stats.wgsl.jinja",
"derive": {
"compensateHalfStats": true,
"vec4": true,
"scalar": "dtypes.T",
"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(normalizedDispatchRows, 65535)",
"y": "ceilDiv(normalizedDispatchRows, 65535)",
"z": 1
}
}
]
},
{
"id": "last_axis_row",
"priority": 100,
"when": ["lastAxisOk", "dtypes.T == dtypes.V", "ranks.scale >= 1", "numel(shapes.scale) == dim(shapes.x, -1)", "dim(shapes.scale, -1) == dim(shapes.x, -1)", "true"],
"derive": { "xElement": "dtypes.T", "ioElement": "dtypes.T" },
"passes": [
{
"id": "main",
"name": "RMSNormalization.LastAxisRow",
"shader": "norm-row-stats.wgsl.jinja",
"derive": {
"compensateHalfStats": true,
"vec4": false,
"scalar": "dtypes.T",
"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(normalizedDispatchRows, 65535)",
"y": "ceilDiv(normalizedDispatchRows, 65535)",
"z": 1
}
}
]
}
]
}