Xenova's picture
Xenova HF Staff
sync 6fdf6301e2bb
b635b3b verified
Raw History Blame
8.02 kB
{
"domain": "ai.onnx",
"name": "GroupNormalization",
"sinceVersion": 21,
"inputs": {
"x": { "onnx": "X", "dtype": "T" },
"scale": { "dtype": "T", "rank": 1 },
"bias": { "dtype": "T", "rank": 1 }
},
"outputs": { "y": { "onnx": "Y", "dtype": "T", "rank": "ranks.x", "shape": "shapes.x" } },
"attributes": { "epsilon": { "default": 0.00001 }, "stash_type": { "default": 1 }, "num_groups": {} },
"attributeConstraints": { "num_groups": { "required": true }, "stash_type": { "values": [1, 10] } },
"typeConstraints": { "T": ["float32", "float16"] },
"tunables": {
"WORKGROUP_SIZE": { "default": 256 },
"MAX_STATS_SPLITS": { "default": 64 },
"STATS_VALUES_PER_SPLIT": { "default": 4096 },
"SPLIT_STATS_MIN_HIDDEN": { "default": 65536 },
"SPLIT_STATS_MAX_ROWS": { "default": 256 }
},
"derive": {
"deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
"groupAttributesOk": "attrs.num_groups >= 1",
"groupShapeOk": "groupAttributesOk and f16Ok(dtypes.T) and ranks.x >= 3 and ranks.scale == 1 and ranks.bias == 1 and ranks.y == ranks.x and sameShape(shapes.y, shapes.x) and dim(shapes.scale, 0) == dim(shapes.x, 1) and dim(shapes.bias, 0) == dim(shapes.x, 1) and dim(shapes.x, 1) % attrs.num_groups == 0",
"groupContractOk": "groupShapeOk and attrs.stash_type == onnxDtypeCode(\"float32\")",
"groupStashF16Ok": "groupShapeOk and attrs.stash_type == onnxDtypeCode(\"float16\")",
"groupRows": "dim(shapes.x, 0) * attrs.num_groups if groupAttributesOk else 0",
"groupSpatial": "inner(shapes.x, 1)",
"groupChannelsPerGroup": "dim(shapes.x, 1) / attrs.num_groups if groupAttributesOk else 0",
"groupHidden": "groupChannelsPerGroup * groupSpatial",
"normDeviceWorkgroupCap": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)",
"normWorkgroupCap": "max(1, pow2ceil(normDeviceWorkgroupCap + 1) / 2)",
"groupScalarWorkgroup": "min(normWorkgroupCap, pow2ceil(groupHidden))",
"groupVec4Workgroup": "min(normWorkgroupCap, pow2ceil(groupHidden / 4))",
"hasSubgroupId": "device.features.has(\"subgroups\") and device.wgslLanguageFeatures.has(\"subgroup_id\")",
"groupRowWorkgroupBytes": "normWorkgroupCap * 2 * 4",
"groupRowCovered": "groupContractOk and groupRowWorkgroupBytes <= device.limits.maxComputeWorkgroupStorageSize",
"groupSplitCount": "min(tunables.MAX_STATS_SPLITS, min(device.limits.maxComputeWorkgroupsPerDimension, 65535), pow2ceil(ceilDiv(groupHidden, tunables.STATS_VALUES_PER_SPLIT)))",
"groupPartialBytes": "groupRows * groupSplitCount * 2 * 4",
"groupSplitCovered": "groupRowCovered and groupRows <= tunables.SPLIT_STATS_MAX_ROWS and groupRows <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and groupHidden >= tunables.SPLIT_STATS_MIN_HIDDEN and groupPartialBytes <= device.limits.maxStorageBufferBindingSize and groupPartialBytes <= device.limits.maxBufferSize"
},
"bindings": {
"x": { "elementType": "$ioElement" },
"scale": { "elementType": "$scalar" },
"bias": { "elementType": "$scalar" },
"y": { "elementType": "$ioElement" },
"params": {
"struct": [
{ "name": "rows", "type": "u32", "value": "groupRows" },
{
"name": "rowStride",
"type": "u32",
"value": "max(1, min(groupRows, min(device.limits.maxComputeWorkgroupsPerDimension, 65535)))"
}
]
},
"x_apply": { "name": "x", "elementType": "$scalar" },
"params_rows": { "name": "params", "struct": [{ "name": "rows", "type": "u32", "value": "groupRows" }] }
},
"variants": [
{
"id": "group_stash_f16_serial",
"priority": 1000,
"when": ["groupStashF16Ok"],
"derive": {
"scalar": "dtypes.T",
"ioElement": "dtypes.T",
"hiddenSize": "groupHidden",
"spatial": "groupSpatial",
"channelsPerGroup": "groupChannelsPerGroup",
"numGroups": "attrs.num_groups",
"epsilon": "attrs.epsilon"
},
"passes": [
{
"id": "main",
"name": "GroupNormalization.StashF16Serial",
"shader": "group-normalization-stash-f16-serial.wgsl.jinja",
"bindings": ["x", "scale", "bias", "y", "params"],
"dispatch": { "x": "min(groupRows, 65535)", "y": "ceilDiv(groupRows, 65535)", "z": 1 }
}
]
},
{
"id": "group_splitk",
"priority": 120,
"when": ["groupSplitCovered"],
"derive": {
"scalar": "dtypes.T",
"hiddenSize": "groupHidden",
"spatial": "groupSpatial",
"channelsPerGroup": "groupChannelsPerGroup",
"numGroups": "attrs.num_groups",
"workgroupSize": "normWorkgroupCap",
"split": "groupSplitCount",
"epsilon": "attrs.epsilon"
},
"intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[groupRows * groupSplitCount, 2]" }],
"passes": [
{
"id": "partials",
"name": "GroupNormalization.SplitKPartials",
"shader": "group-normalization-splitk-partials.wgsl.jinja",
"bindings": ["x_apply", { "name": "partials", "elementType": "vec2<f32>" }, "params_rows"],
"dispatch": {
"x": "min(groupRows, DISPATCH_FOLD_WIDTH)",
"y": "ceilDiv(groupRows, DISPATCH_FOLD_WIDTH)",
"z": "groupSplitCount"
}
},
{
"id": "apply",
"name": "GroupNormalization.SplitKApply",
"shader": "group-normalization-splitk-apply.wgsl.jinja",
"bindings": [
"x_apply",
"scale",
"bias",
{ "name": "partials", "buffer": "read-only-storage", "elementType": "vec2<f32>" },
{ "arg": "y", "elementType": "$scalar" },
"params_rows"
],
"dispatch": {
"x": "min(groupRows, DISPATCH_FOLD_WIDTH)",
"y": "ceilDiv(groupRows, DISPATCH_FOLD_WIDTH)",
"z": "groupSplitCount"
}
}
]
},
{
"id": "group_subgroup_vec4",
"priority": 110,
"when": ["groupRowCovered", "groupSpatial % 4 == 0"],
"derive": { "scalar": "dtypes.T", "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"" },
"passes": [
{
"id": "main",
"name": "GroupNormalization.GroupSubgroupVec4",
"shader": "norm-row-stats.wgsl.jinja",
"derive": {
"vec4": true,
"scalar": "dtypes.T",
"hidden": "groupHidden",
"wg": "groupVec4Workgroup",
"epsilon": "attrs.epsilon",
"numGroupsSpec": "attrs.num_groups",
"cpg": "groupChannelsPerGroup",
"hiddenVec": "groupHidden / 4",
"vecType": "\"vec4<\" ~ dtypes.T ~ \">\"",
"spatialVec": "groupSpatial / 4",
"combineSubgroups": "hasSubgroupId"
},
"bindings": ["x", "scale", "bias", "y", "params"],
"dispatch": { "x": "min(groupRows, 65535)", "y": "ceilDiv(groupRows, 65535)", "z": 1 }
}
]
},
{
"id": "group_subgroup",
"priority": 100,
"when": ["groupRowCovered"],
"derive": { "scalar": "dtypes.T", "ioElement": "dtypes.T" },
"passes": [
{
"id": "main",
"name": "GroupNormalization.GroupSubgroup",
"shader": "norm-row-stats.wgsl.jinja",
"derive": {
"vec4": false,
"scalar": "dtypes.T",
"hidden": "groupHidden",
"wg": "groupScalarWorkgroup",
"epsilon": "attrs.epsilon",
"numGroupsSpec": "attrs.num_groups",
"cpg": "groupChannelsPerGroup",
"spatial": "groupSpatial",
"combineSubgroups": "hasSubgroupId"
},
"bindings": ["x", "scale", "bias", "y", "params"],
"dispatch": { "x": "min(groupRows, 65535)", "y": "ceilDiv(groupRows, 65535)", "z": 1 }
}
]
}
]
}