Xenova's picture
Xenova HF Staff
sync 2e7068faf55e
0af9165 verified
Raw
History Blame
20 kB
{
"domain": "ai.onnx",
"name": "InstanceNormalization",
"sinceVersion": 6,
"description": "Applies instance normalization to the input: `y = scale * (x - mean) / sqrt(variance + epsilon) + B`, where `mean` and `variance` are computed per instance per channel over the spatial dimensions. Equivalent to batch normalization with a batch size of one per channel.",
"inputs": [
{ "role": "input", "dtype": "T", "description": "Input tensor of shape `(N x C x D1 x ... x Dn)`; at least 3-D." },
{
"role": "scale",
"dtype": "T",
"rank": 1,
"description": "1-D scale tensor of size C, one scale factor per channel."
},
{ "role": "B", "dtype": "T", "rank": 1, "description": "1-D bias tensor of size C, one bias value per channel." }
],
"outputs": [
{
"role": "output",
"dtype": "T",
"rank": "ranks.input",
"description": "Normalized output tensor; same shape as the input.",
"shape": "shapes.input"
}
],
"attributes": { "epsilon": 0.00001 },
"attributeDescriptions": {
"epsilon": "Small constant added to the variance before taking the square root to avoid division by zero."
},
"typeConstraints": { "T": ["float32", "float16"] },
"args": {
"input": { "kind": "tensor", "semantic": "input", "role": "input" },
"scale": { "kind": "tensor", "semantic": "scale", "role": "input" },
"b": { "kind": "tensor", "semantic": "B", "role": "input" },
"output": { "kind": "tensor", "semantic": "output", "role": "output" }
},
"tunables": {
"WORKGROUP_SIZE": 256,
"MAX_STATS_SPLITS": 256,
"STATS_VALUES_PER_SPLIT": 2048,
"SPLIT_STATS_MIN_SPATIAL": 65536,
"SPLIT_STATS_MAX_PLANES": 256,
"COMBINE_WORKGROUP_SIZE": 64,
"BATCHED_MIN_PLANES_PER_WORKGROUP": 8
},
"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\"))",
"instanceContractOk": "f16Ok(dtypes.T) and ranks.input >= 3 and ranks.output == ranks.input and sameShape(shapes.output, shapes.input) and ranks.scale == 1 and ranks.B == 1 and dim(shapes.scale, 0) == dim(shapes.input, 1) and dim(shapes.B, 0) == dim(shapes.input, 1)",
"instancePlanes": "dim(shapes.input, 0) * dim(shapes.input, 1)",
"instanceSpatial": "inner(shapes.input, 1)",
"normDeviceWorkgroupCap": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)",
"normWorkgroupCap": "max(1, pow2ceil(normDeviceWorkgroupCap + 1) / 2)",
"normSubgroupMin": "device.adapterInfo.subgroupMinSize if has(device.adapterInfo, \"subgroupMinSize\") else 1",
"normSubgroupMax": "device.adapterInfo.subgroupMaxSize if has(device.adapterInfo, \"subgroupMaxSize\") else 32",
"hasSubgroupId": "device.features.has(\"subgroups\") and device.wgslLanguageFeatures.has(\"subgroup_id\")",
"instanceScalarWorkgroup": "min(normWorkgroupCap, pow2ceil(instanceSpatial))",
"instanceVec4Workgroup": "min(normWorkgroupCap, pow2ceil(instanceSpatial / 4))",
"instanceVec4SubgroupEfficient": "not hasSubgroupId or instanceVec4Workgroup >= normSubgroupMin",
"instanceBatchedVec4Lanes": "instanceVec4Workgroup",
"instanceBatchedVec4PlanesPerWorkgroup": "max(1, floor(normWorkgroupCap / instanceBatchedVec4Lanes))",
"instanceBatchedVec4Workgroups": "ceilDiv(instancePlanes, instanceBatchedVec4PlanesPerWorkgroup)",
"instanceBatchedVec4StorageBytes": "normWorkgroupCap * 2 * 4 + instanceBatchedVec4PlanesPerWorkgroup * 4",
"instanceRowWorkgroupBytes": "normWorkgroupCap * 2 * 4",
"instanceStatsBytes": "instancePlanes * 2 * 4",
"instanceStatsFits": "instanceStatsBytes <= device.limits.maxStorageBufferBindingSize and instanceStatsBytes <= device.limits.maxBufferSize",
"instanceRowCovered": "instanceContractOk and instanceRowWorkgroupBytes <= device.limits.maxComputeWorkgroupStorageSize",
"instanceSplitCount": "min(tunables.MAX_STATS_SPLITS, device.limits.maxComputeWorkgroupsPerDimension, pow2ceil(ceilDiv(instanceSpatial, tunables.STATS_VALUES_PER_SPLIT)))",
"instancePartialBytes": "instancePlanes * instanceSplitCount * 2 * 4",
"splitStatsCovered": "instanceRowCovered and instanceStatsFits and instancePlanes <= tunables.SPLIT_STATS_MAX_PLANES and instancePlanes <= device.limits.maxComputeWorkgroupsPerDimension and instanceSpatial >= tunables.SPLIT_STATS_MIN_SPATIAL and instancePartialBytes <= device.limits.maxStorageBufferBindingSize and instancePartialBytes <= device.limits.maxBufferSize",
"splitStatsPreferred": "splitStatsCovered and instancePlanes < normSubgroupMax"
},
"bindingSets": {
"planeIo": [
{
"name": "x",
"arg": "input",
"semantic": "input",
"buffer": { "type": "read-only-storage" },
"elementType": "$ioElement"
},
{
"name": "scale",
"arg": "scale",
"semantic": "scale",
"buffer": { "type": "read-only-storage" },
"elementType": "$T"
},
{ "name": "bias", "arg": "b", "semantic": "B", "buffer": { "type": "read-only-storage" }, "elementType": "$T" },
{
"name": "y",
"arg": "output",
"semantic": "output",
"buffer": { "type": "storage" },
"elementType": "$ioElement"
}
],
"plane": [
{
"name": "x",
"arg": "input",
"semantic": "input",
"buffer": { "type": "read-only-storage" },
"elementType": "$ioElement"
},
{
"name": "scale",
"arg": "scale",
"semantic": "scale",
"buffer": { "type": "read-only-storage" },
"elementType": "$T"
},
{ "name": "bias", "arg": "b", "semantic": "B", "buffer": { "type": "read-only-storage" }, "elementType": "$T" },
{
"name": "y",
"arg": "output",
"semantic": "output",
"buffer": { "type": "storage" },
"elementType": "$ioElement"
},
{
"name": "params",
"semantic": "kernel.params",
"buffer": { "type": "uniform" },
"struct": {
"name": "Params",
"fields": [
{ "name": "rows", "type": "u32", "value": "dim(shapes.input, 0) * dim(shapes.input, 1)" },
{
"name": "rowStride",
"type": "u32",
"value": "max(1, min(dim(shapes.input, 0) * dim(shapes.input, 1), device.limits.maxComputeWorkgroupsPerDimension))"
}
]
}
}
],
"planeBatched": [
{
"name": "x",
"arg": "input",
"semantic": "input",
"buffer": { "type": "read-only-storage" },
"elementType": "$ioElement"
},
{
"name": "scale",
"arg": "scale",
"semantic": "scale",
"buffer": { "type": "read-only-storage" },
"elementType": "$T"
},
{ "name": "bias", "arg": "b", "semantic": "B", "buffer": { "type": "read-only-storage" }, "elementType": "$T" },
{
"name": "y",
"arg": "output",
"semantic": "output",
"buffer": { "type": "storage" },
"elementType": "$ioElement"
},
{
"name": "params",
"semantic": "kernel.params",
"buffer": { "type": "uniform" },
"struct": {
"name": "Params",
"fields": [{ "name": "rows", "type": "u32", "value": "dim(shapes.input, 0) * dim(shapes.input, 1)" }]
}
}
],
"applyScalar": [
{
"name": "input",
"arg": "input",
"semantic": "input",
"buffer": { "type": "read-only-storage" },
"elementType": "$T"
},
{ "name": "stats", "semantic": "stats", "buffer": { "type": "read-only-storage" }, "elementType": "f32" },
{
"name": "scale",
"arg": "scale",
"semantic": "scale",
"buffer": { "type": "read-only-storage" },
"elementType": "$T"
},
{ "name": "bias", "arg": "b", "semantic": "B", "buffer": { "type": "read-only-storage" }, "elementType": "$T" },
{ "name": "output", "arg": "output", "semantic": "output", "buffer": { "type": "storage" }, "elementType": "$T" },
{
"name": "params",
"semantic": "kernel.params",
"buffer": { "type": "uniform" },
"struct": {
"name": "Params",
"fields": [
{ "name": "count", "type": "u32", "value": "numel(shapes.output)" },
{ "name": "channels", "type": "u32", "value": "dim(shapes.input, 1)" },
{ "name": "spatial", "type": "u32", "value": "instanceSpatial" }
]
}
}
],
"applyVec4": [
{
"name": "input",
"arg": "input",
"semantic": "input",
"buffer": { "type": "read-only-storage" },
"elementType": "$vectorScalar"
},
{ "name": "stats", "semantic": "stats", "buffer": { "type": "read-only-storage" }, "elementType": "f32" },
{
"name": "scale",
"arg": "scale",
"semantic": "scale",
"buffer": { "type": "read-only-storage" },
"elementType": "$T"
},
{ "name": "bias", "arg": "b", "semantic": "B", "buffer": { "type": "read-only-storage" }, "elementType": "$T" },
{
"name": "output",
"arg": "output",
"semantic": "output",
"buffer": { "type": "storage" },
"elementType": "$vectorScalar"
},
{
"name": "params",
"semantic": "kernel.params",
"buffer": { "type": "uniform" },
"struct": {
"name": "Params",
"fields": [
{ "name": "count", "type": "u32", "value": "numel(shapes.output) / 4" },
{ "name": "channels", "type": "u32", "value": "dim(shapes.input, 1)" },
{ "name": "spatial", "type": "u32", "value": "instanceSpatial" }
]
}
}
],
"splitPartials": [
{
"name": "input",
"arg": "input",
"semantic": "input",
"buffer": { "type": "read-only-storage" },
"elementType": "$splitInputElement"
},
{ "name": "partials", "semantic": "partials", "buffer": { "type": "storage" }, "elementType": "f32" },
{
"name": "params",
"semantic": "kernel.params",
"buffer": { "type": "uniform" },
"struct": {
"name": "Params",
"fields": [
{ "name": "planes", "type": "u32", "value": "instancePlanes" },
{ "name": "spatial", "type": "u32", "value": "instanceSpatial" }
]
}
}
],
"splitCombine": [
{ "name": "partials", "semantic": "partials", "buffer": { "type": "read-only-storage" }, "elementType": "f32" },
{ "name": "stats", "semantic": "stats", "buffer": { "type": "storage" }, "elementType": "f32" },
{
"name": "params",
"semantic": "kernel.params",
"buffer": { "type": "uniform" },
"struct": {
"name": "Params",
"fields": [
{ "name": "planes", "type": "u32", "value": "instancePlanes" },
{ "name": "spatial", "type": "u32", "value": "instanceSpatial" },
{ "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }
]
}
}
]
},
"variants": [
{
"id": "plane_batched_vec4",
"priority": 115,
"when": ["instanceRowCovered", "instanceSpatial % 4 == 0", "instanceSpatial >= 4", "instancePlanes >= normWorkgroupCap", "instanceBatchedVec4PlanesPerWorkgroup >= tunables.BATCHED_MIN_PLANES_PER_WORKGROUP", "instanceBatchedVec4StorageBytes <= device.limits.maxComputeWorkgroupStorageSize"],
"demoteWhen": ["reportedNonWave32Adapter and instancePlanes <= device.limits.maxComputeWorkgroupsPerDimension"],
"constants": {
"usesF16": "dtypes.T == \"f16\"",
"ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"",
"vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
"hidden": "instanceSpatial",
"hiddenVec": "instanceSpatial / 4",
"channels": "dim(shapes.input, 1) if dim(shapes.input, 1) > 0 else 1",
"epsilon": "attrs.epsilon",
"workgroupSize": "normWorkgroupCap",
"lanesPerPlane": "instanceBatchedVec4Lanes",
"planesPerWorkgroup": "instanceBatchedVec4PlanesPerWorkgroup"
},
"passes": [
{
"id": "main",
"name": "InstanceNormalization.PlaneBatchedVec4",
"shader": "instance-normalization-batched-planes-vec4.wgsl.jinja",
"bindings": "planeBatched",
"dispatch": { "workgroups": "instanceBatchedVec4Workgroups" }
}
]
},
{
"id": "plane_subgroup_vec4",
"priority": 110,
"requires": { "features": [] },
"when": ["instanceRowCovered", "inner(shapes.input, 1) % 4 == 0", "instanceVec4SubgroupEfficient"],
"constants": { "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"", "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" },
"passes": [
{
"id": "main",
"name": "InstanceNormalization.plane_subgroup_vec4",
"source": {
"shader": "norm-row-stats.wgsl.jinja",
"inputs": {
"mode": "\"instance\"",
"vec4": true,
"scalar": "dtypes.T",
"usesF16": "dtypes.T == \"f16\"",
"hidden": "instanceSpatial",
"wg": "instanceVec4Workgroup",
"epsilon": "attrs.epsilon",
"channels": "dim(shapes.input, 1) if dim(shapes.input, 1) > 0 else 1",
"hiddenVec": "instanceSpatial / 4",
"vecType": "\"vec4<\" ~ dtypes.T ~ \">\"",
"combineSubgroups": "hasSubgroupId"
}
},
"subgroupCollectivesWidth": "portable",
"bindings": "plane",
"dispatch": { "workgroups": "instancePlanes" }
}
]
},
{
"id": "plane_subgroup_vec4_scalar_io",
"priority": 111,
"requires": { "features": ["subgroups"] },
"when": ["dtypes.T == \"f32\"", "wave32Adapter", "device.wgslLanguageFeatures.has(\"subgroup_id\")", "instanceRowCovered", "instanceSpatial % 4 == 0", "instanceVec4SubgroupEfficient"],
"constants": { "ioElement": "dtypes.T" },
"passes": [
{
"id": "main",
"name": "InstanceNormalization.plane_subgroup_vec4_scalar_io",
"source": {
"shader": "norm-row-stats.wgsl.jinja",
"inputs": {
"mode": "\"instance\"",
"vec4": true,
"scalarIo": true,
"scalar": "dtypes.T",
"usesF16": false,
"hidden": "instanceSpatial",
"wg": "instanceVec4Workgroup",
"epsilon": "attrs.epsilon",
"channels": "dim(shapes.input, 1) if dim(shapes.input, 1) > 0 else 1",
"hiddenVec": "instanceSpatial / 4",
"vecType": "\"vec4<f32>\"",
"combineSubgroups": true
}
},
"subgroupCollectivesWidth": "portable",
"bindings": "plane",
"dispatch": { "workgroups": "instancePlanes" }
}
]
},
{
"id": "plane_subgroup",
"priority": 100,
"requires": { "features": [] },
"when": ["instanceRowCovered"],
"constants": { "ioElement": "dtypes.T" },
"passes": [
{
"id": "main",
"name": "InstanceNormalization.plane_subgroup",
"source": {
"shader": "norm-row-stats.wgsl.jinja",
"inputs": {
"mode": "\"instance\"",
"vec4": false,
"scalar": "dtypes.T",
"usesF16": "dtypes.T == \"f16\"",
"hidden": "instanceSpatial",
"wg": "instanceScalarWorkgroup",
"epsilon": "attrs.epsilon",
"channels": "dim(shapes.input, 1) if dim(shapes.input, 1) > 0 else 1",
"combineSubgroups": "hasSubgroupId"
}
},
"subgroupCollectivesWidth": "portable",
"bindings": "plane",
"dispatch": { "workgroups": "instancePlanes" }
}
]
},
{
"id": "plane_splitk_vec4",
"priority": 121,
"when": ["splitStatsPreferred", "instanceSpatial % 4 == 0"],
"constants": {
"vectorized": true,
"usesF16": "dtypes.T == \"f16\"",
"vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
"useSubgroups": "hasSubgroupId",
"splitInputElement": "\"vec4<\" ~ dtypes.T ~ \">\"",
"workgroupSize": "normWorkgroupCap",
"split": "instanceSplitCount",
"combineWorkgroupSize": "min(tunables.COMBINE_WORKGROUP_SIZE, normWorkgroupCap)",
"applyWorkgroupSize": "normWorkgroupCap"
},
"intermediates": [
{ "id": "partials", "dtype": "float32", "shape": "[instancePlanes * instanceSplitCount, 2]" },
{ "id": "stats", "dtype": "float32", "shape": "[instancePlanes, 2]" }
],
"passes": [
{
"id": "partials",
"name": "InstanceNormalization.SplitKPartialsVec4",
"shader": "instance-normalization-splitk-partials.wgsl.jinja",
"subgroupCollectivesWidth": "portable",
"bindings": "splitPartials",
"dispatch": { "workgroups": "instancePlanes", "z": "instanceSplitCount" }
},
{
"id": "combine",
"name": "InstanceNormalization.SplitKCombine",
"shader": "instance-normalization-splitk-combine.wgsl.jinja",
"bindings": "splitCombine",
"dispatch": { "threads": "instancePlanes", "workgroupSize": "constants.combineWorkgroupSize" }
},
{
"id": "apply",
"name": "InstanceNormalization.ApplyVec4",
"shader": "instance-normalization-apply.wgsl.jinja",
"bindings": "applyVec4",
"dispatch": { "threads": "numel(shapes.output) / 4", "workgroupSize": "constants.applyWorkgroupSize" }
}
]
},
{
"id": "plane_splitk",
"priority": 120,
"when": ["splitStatsPreferred"],
"constants": {
"scalar": "dtypes.T",
"usesF16": "dtypes.T == \"f16\"",
"splitInputElement": "dtypes.T",
"workgroupSize": "normWorkgroupCap",
"split": "instanceSplitCount",
"combineWorkgroupSize": "min(tunables.COMBINE_WORKGROUP_SIZE, normWorkgroupCap)",
"applyWorkgroupSize": "normWorkgroupCap"
},
"intermediates": [
{ "id": "partials", "dtype": "float32", "shape": "[instancePlanes * instanceSplitCount, 2]" },
{ "id": "stats", "dtype": "float32", "shape": "[instancePlanes, 2]" }
],
"passes": [
{
"id": "partials",
"name": "InstanceNormalization.SplitKPartials",
"shader": "instance-normalization-splitk-partials.wgsl.jinja",
"bindings": "splitPartials",
"dispatch": { "workgroups": "instancePlanes", "z": "instanceSplitCount" }
},
{
"id": "combine",
"name": "InstanceNormalization.SplitKCombine",
"shader": "instance-normalization-splitk-combine.wgsl.jinja",
"bindings": "splitCombine",
"dispatch": { "threads": "instancePlanes", "workgroupSize": "constants.combineWorkgroupSize" }
},
{
"id": "apply",
"name": "InstanceNormalization.Apply",
"shader": "instance-normalization-apply.wgsl.jinja",
"bindings": "applyScalar",
"dispatch": { "threads": "numel(shapes.output)", "workgroupSize": "constants.applyWorkgroupSize" }
}
]
}
]
}