Xenova's picture
Xenova HF Staff
sync 6fdf6301e2bb
66f0d7a verified
Raw History Blame
13.6 kB
{
"domain": "com.microsoft",
"name": "BiasSoftmax",
"sinceVersion": 1,
"inputs": { "data": { "dtype": "T" }, "bias": { "dtype": "T" } },
"outputs": { "output": { "dtype": "T", "rank": "ranks.data", "shape": "shapes.data" } },
"attributes": { "axis": { "default": 1 }, "is_inner_broadcast": {} },
"attributeConstraints": { "is_inner_broadcast": { "required": true } },
"typeConstraints": { "T": ["float32", "float16"] },
"tunables": {
"WORKGROUP_SIZE": { "default": 256 },
"BLOCK_COLS": { "default": 2048 },
"ONLINE_WORKGROUP_SIZE": { "default": 64 }
},
"derive": {
"axisNorm": "attrs.axis if attrs.axis >= 0 else attrs.axis + ranks.data",
"batchCount": "outer(shapes.data, axisNorm)",
"blockSize": "dim(shapes.data, axisNorm) * inner(shapes.data, axisNorm)",
"biasBlockCount": "numel(shapes.bias) / max(1, blockSize)",
"biasContract": "(numel(shapes.data) == 0 and numel(shapes.bias) == 0) or (blockSize > 0 and biasBlockCount > 0 and numel(shapes.bias) % blockSize == 0 and biasBlockCount <= batchCount and batchCount % biasBlockCount == 0)",
"deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
"subgroupRowsWorkgroup": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)",
"subgroupRowsLaneVecs": "ceilDiv(blockSize / 4, device.adapterInfo.subgroupMinSize) if (device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and device.adapterInfo.subgroupMinSize > 0) else 0",
"biasBaseWorkgroup": "min(tunables.ONLINE_WORKGROUP_SIZE, deviceWorkgroupCap)",
"biasMaxWorkgroup": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)",
"biasWorkgroupSize": "min(biasMaxWorkgroup, pow2ceil(ceilDiv(blockSize, 4)), max(biasBaseWorkgroup, pow2ceil(ceilDiv(biasBaseWorkgroup * biasMaxWorkgroup, max(1, batchCount)))))"
},
"when": ["numel(shapes.data) == numel(shapes.output)", "ranks.data >= 1", "attrs.axis + ranks.data >= 0", "attrs.axis < ranks.data", "biasContract", "f16Ok(dtypes.T)"],
"bindings": {
"data": { "elementType": "$scalar" },
"bias": { "elementType": "$scalar" },
"output": { "elementType": "$scalar" }
},
"variants": [
{
"id": "online_subgroup_rows_vec4",
"priority": 20,
"when": ["blockSize % 4 == 0", "has(device.adapterInfo, \"subgroupMinSize\")", "device.adapterInfo.subgroupMinSize >= 16", "blockSize >= 64", "subgroupRowsLaneVecs >= 1", "subgroupRowsLaneVecs <= 8", "batchCount >= 64", "has(device.adapterInfo, \"subgroupMaxSize\")", "device.adapterInfo.subgroupMaxSize <= subgroupRowsWorkgroup", "pow2ceil(subgroupRowsWorkgroup) == subgroupRowsWorkgroup", "ceilDiv(ceilDiv(batchCount, subgroupRowsWorkgroup / device.adapterInfo.subgroupMaxSize), min(device.limits.maxComputeWorkgroupsPerDimension, 65535)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
"requires": { "features": ["subgroups"] },
"derive": {
"scalar": "dtypes.T",
"vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
"workgroupSize": "subgroupRowsWorkgroup",
"vecsPerLane": "subgroupRowsLaneVecs",
"isInnerBroadcast": "attrs.is_inner_broadcast != 0",
"biasBlockCountSpec": "max(1, biasBlockCount)",
"innerRepeat": "max(1, batchCount / max(1, biasBlockCount))"
},
"passes": [
{
"id": "main",
"name": "BiasSoftmax.SubgroupRowsVec4",
"shader": "softmax-subgroup-rows.wgsl.jinja",
"derive": { "op": "\"biassoftmax\"" },
"bindings": [
{ "arg": "data", "name": "x", "elementType": "$vectorScalar" },
{ "arg": "bias", "name": "bias", "elementType": "$vectorScalar" },
{ "arg": "output", "name": "y", "elementType": "$vectorScalar" },
{
"name": "params",
"struct": [
{ "name": "rows", "type": "u32", "value": "batchCount" },
{ "name": "vecCols", "type": "u32", "value": "blockSize / 4" }
]
}
],
"dispatch": {
"x": "min(ceilDiv(batchCount, subgroupRowsWorkgroup / device.adapterInfo.subgroupMaxSize), 65535)",
"y": "ceilDiv(ceilDiv(batchCount, subgroupRowsWorkgroup / device.adapterInfo.subgroupMaxSize), 65535)",
"z": 1
}
}
]
},
{
"id": "online_workgroup_vec4",
"priority": 15,
"when": ["blockSize % 4 == 0", "blockSize > 32", "blockSize < 65536", "pow2ceil(biasWorkgroupSize) == biasWorkgroupSize", "ceilDiv(batchCount, min(device.limits.maxComputeWorkgroupsPerDimension, 65535)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
"derive": {
"scalar": "dtypes.T",
"vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
"workgroupSize": "biasWorkgroupSize",
"combineSubgroups": false,
"isInnerBroadcast": "attrs.is_inner_broadcast != 0",
"biasBlockCountSpec": "max(1, biasBlockCount)",
"innerRepeat": "max(1, batchCount / max(1, biasBlockCount))",
"scalarVec4": false
},
"passes": [
{
"id": "main",
"name": "BiasSoftmax.WorkgroupVec4",
"shader": "softmax-online.wgsl.jinja",
"derive": { "op": "\"biassoftmax\"", "useVec4": true },
"bindings": [
{ "arg": "data", "name": "x", "elementType": "$vectorScalar" },
{ "arg": "bias", "name": "bias", "elementType": "$vectorScalar" },
{ "arg": "output", "name": "y", "elementType": "$vectorScalar" },
{
"name": "params",
"struct": [
{ "name": "rows", "type": "u32", "value": "batchCount" },
{ "name": "vecCols", "type": "u32", "value": "blockSize / 4" }
]
}
],
"dispatch": { "x": "min(batchCount, 65535)", "y": "ceilDiv(batchCount, 65535)", "z": 1 }
}
]
},
{
"id": "online_workgroup_scalar_vec4",
"priority": 15,
"when": ["blockSize % 4 != 0", "blockSize > 32", "blockSize < 65536", "pow2ceil(biasWorkgroupSize) == biasWorkgroupSize", "ceilDiv(batchCount, min(device.limits.maxComputeWorkgroupsPerDimension, 65535)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
"derive": {
"scalar": "dtypes.T",
"vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"",
"workgroupSize": "biasWorkgroupSize",
"combineSubgroups": false,
"isInnerBroadcast": "attrs.is_inner_broadcast != 0",
"biasBlockCountSpec": "max(1, biasBlockCount)",
"innerRepeat": "max(1, batchCount / max(1, biasBlockCount))",
"scalarVec4": true
},
"passes": [
{
"id": "main",
"name": "BiasSoftmax.WorkgroupScalarVec4",
"shader": "softmax-online.wgsl.jinja",
"derive": { "op": "\"biassoftmax\"", "useVec4": true },
"bindings": [
{ "arg": "data", "name": "x", "elementType": "$scalar" },
{ "arg": "bias", "name": "bias", "elementType": "$scalar" },
{ "arg": "output", "name": "y", "elementType": "$scalar" },
{
"name": "params",
"struct": [
{ "name": "rows", "type": "u32", "value": "batchCount" },
{ "name": "vecCols", "type": "u32", "value": "ceilDiv(blockSize, 4)" },
{ "name": "cols", "type": "u32", "value": "blockSize" }
]
}
],
"dispatch": { "x": "min(batchCount, 65535)", "y": "ceilDiv(batchCount, 65535)", "z": 1 }
}
]
},
{
"id": "longrow_split",
"priority": 40,
"when": ["blockSize >= 65536", "batchCount > 0", "batchCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "ceilDiv(blockSize, tunables.BLOCK_COLS) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)"],
"derive": { "scalar": "dtypes.T", "combineSubgroups": false },
"intermediates": [
{
"id": "blockMax",
"dtype": "float32",
"shape": "[outer(shapes.data, axisNorm) * ceilDiv(dim(shapes.data, axisNorm) * inner(shapes.data, axisNorm), tunables.BLOCK_COLS)]"
},
{
"id": "blockSum",
"dtype": "float32",
"shape": "[outer(shapes.data, axisNorm) * ceilDiv(dim(shapes.data, axisNorm) * inner(shapes.data, axisNorm), tunables.BLOCK_COLS)]"
},
{ "id": "rowMax", "dtype": "float32", "shape": "[outer(shapes.data, axisNorm)]" },
{ "id": "rowSum", "dtype": "float32", "shape": "[outer(shapes.data, axisNorm)]" }
],
"passes": [
{
"id": "block_stats",
"name": "BiasSoftmax.LongRowBlockStats",
"shader": "softmax-longrow-stats.wgsl.jinja",
"derive": {
"stage": "\"block\"",
"biasRow": true,
"isInnerBroadcast": "attrs.is_inner_broadcast != 0",
"biasBlockCountSpec": "max(1, biasBlockCount)",
"innerRepeat": "max(1, batchCount / max(1, biasBlockCount))"
},
"bindings": [
"data",
"bias",
{ "name": "blockMax", "elementType": "f32" },
{ "name": "blockSum", "elementType": "f32" },
{
"name": "params",
"struct": [
{ "name": "blockSize", "type": "u32", "value": "blockSize" },
{ "name": "blocks", "type": "u32", "value": "ceilDiv(blockSize, tunables.BLOCK_COLS)" }
]
}
],
"dispatch": { "x": "ceilDiv(blockSize, tunables.BLOCK_COLS)", "y": "batchCount" }
},
{
"id": "row_stats",
"name": "BiasSoftmax.LongRowStats",
"shader": "softmax-longrow-stats.wgsl.jinja",
"derive": { "stage": "\"row\"" },
"bindings": [
{ "name": "blockMax", "buffer": "read-only-storage", "elementType": "f32" },
{ "name": "blockSum", "buffer": "read-only-storage", "elementType": "f32" },
{ "name": "rowMax", "elementType": "f32" },
{ "name": "rowSum", "elementType": "f32" },
{
"name": "params",
"struct": [{ "name": "blocks", "type": "u32", "value": "ceilDiv(blockSize, tunables.BLOCK_COLS)" }]
}
],
"dispatch": { "x": "batchCount" }
},
{
"id": "normalize",
"name": "BiasSoftmax.LongRowNormalize",
"shader": "bias-softmax-longrow-normalize.wgsl.jinja",
"derive": {
"isInnerBroadcast": "attrs.is_inner_broadcast != 0",
"biasBlockCountSpec": "max(1, biasBlockCount)",
"innerRepeat": "max(1, batchCount / max(1, biasBlockCount))"
},
"bindings": [
"data",
"bias",
{ "name": "rowMax", "buffer": "read-only-storage", "elementType": "f32" },
{ "name": "rowSum", "buffer": "read-only-storage", "elementType": "f32" },
"output",
{ "name": "params", "struct": [{ "name": "blockSize", "type": "u32", "value": "blockSize" }] }
],
"dispatch": { "x": "ceilDiv(blockSize, tunables.BLOCK_COLS)", "y": "batchCount" }
}
]
},
{
"id": "packed_rows",
"priority": 30,
"when": ["blockSize > 0", "blockSize <= 8", "batchCount >= 64"],
"derive": { "scalar": "dtypes.T", "combineSubgroups": false, "packedRows": true },
"passes": [
{
"id": "main",
"name": "BiasSoftmax.PackedRows",
"shader": "bias-softmax.wgsl.jinja",
"derive": {
"isInnerBroadcast": "attrs.is_inner_broadcast != 0",
"biasBlockCountSpec": "max(1, biasBlockCount)",
"innerRepeat": "max(1, batchCount / max(1, biasBlockCount))"
},
"bindings": [
"data",
"bias",
"output",
{
"name": "params",
"struct": [
{ "name": "blockSize", "type": "u32", "value": "blockSize" },
{ "name": "batchCount", "type": "u32", "value": "batchCount" }
]
}
],
"dispatch": {
"x": "min(ceilDiv((batchCount), (tunables.WORKGROUP_SIZE)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))",
"y": 1,
"z": 1
}
}
]
},
{
"id": "adaptive_row",
"priority": 10,
"when": ["numel(shapes.data) >= 0"],
"derive": { "packedRows": false, "scalar": "dtypes.T", "combineSubgroups": "device.features.has(\"subgroups\")" },
"passes": [
{
"id": "main",
"name": "BiasSoftmax.AdaptiveRow",
"shader": "bias-softmax.wgsl.jinja",
"derive": {
"isInnerBroadcast": "attrs.is_inner_broadcast != 0",
"biasBlockCountSpec": "max(1, biasBlockCount)",
"innerRepeat": "max(1, batchCount / max(1, biasBlockCount))"
},
"bindings": [
"data",
"bias",
"output",
{
"name": "params",
"struct": [
{ "name": "blockSize", "type": "u32", "value": "blockSize" },
{ "name": "batchCount", "type": "u32", "value": "batchCount" }
]
}
],
"dispatch": { "x": "min(batchCount, 65535)", "y": "ceilDiv(batchCount, 65535)", "z": 1 }
}
]
}
]
}