ai.onnx.GridSample / build /webgpu /manifest.json
Xenova's picture
Xenova HF Staff
sync 91d990483a17
01356ce verified
Raw
History Blame
8.75 kB
{
"domain": "ai.onnx",
"name": "GridSample",
"sinceVersion": 20,
"inputs": { "x": { "onnx": "X", "dtype": "T" }, "grid": { "dtype": "T" } },
"outputs": {
"y": {
"onnx": "Y",
"dtype": "T",
"rank": "ranks.grid",
"shape": "prefix(shapes.x, 2) + prefix(suffix(shapes.grid, 1), ranks.grid - 2)"
}
},
"attributes": {
"mode": { "default": "linear" },
"padding_mode": { "default": "zeros" },
"align_corners": { "default": 0 }
},
"attributeConstraints": {
"mode": { "values": ["linear", "nearest", "cubic"] },
"padding_mode": { "values": ["zeros", "border", "reflection"] },
"align_corners": { "values": [0, 1] }
},
"typeConstraints": { "T": ["float32", "float16"] },
"tunables": { "WORKGROUP_SIZE": { "default": 256 } },
"derive": {
"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\"))",
"deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
"rank4Ok": "ranks.x == 4 and ranks.grid == 4 and ranks.y == 4 and dim(shapes.grid, 0) == dim(shapes.x, 0) and dim(shapes.grid, 3) == 2 and dim(shapes.y, 0) == dim(shapes.x, 0) and dim(shapes.y, 1) == dim(shapes.x, 1) and dim(shapes.y, 2) == dim(shapes.grid, 1) and dim(shapes.y, 3) == dim(shapes.grid, 2) and f16Ok(dtypes.T)",
"channelWidth": 4,
"scalar": "dtypes.T",
"volumeWidth": "min(channelWidth, pow(2, log2ceil(max(1, dim(shapes.x, 1)))))",
"volumeWorkgroup": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)"
},
"bindings": {
"params": {
"buffer": "uniform",
"struct": [
{ "name": "count", "type": "u32", "value": "numel(shapes.y)" },
{ "name": "C", "type": "u32", "value": "dim(shapes.x, 1)" },
{ "name": "inH", "type": "u32", "value": "dim(shapes.x, 2)" },
{ "name": "inW", "type": "u32", "value": "dim(shapes.x, 3)" },
{ "name": "outH", "type": "u32", "value": "dim(shapes.y, 2)" },
{ "name": "outW", "type": "u32", "value": "dim(shapes.y, 3)" }
]
}
},
"variants": [
{
"id": "nchw_rank4_channel_vector",
"priority": 5,
"when": ["rank4Ok", "not reportedNonWave32Adapter", "dim(shapes.x, 1) >= 2"],
"passes": [
{
"id": "main",
"name": "GridSample.ChannelX4",
"shader": "grid-sample.wgsl.jinja",
"derive": {
"modeSpec": "attrs.mode",
"paddingMode": "attrs.padding_mode",
"alignCorners": "attrs.align_corners != 0",
"channelWidthSpec": "channelWidth",
"channelTail": "dim(shapes.x, 1) % channelWidth != 0"
},
"bindings": ["x", "grid", "y", "params"],
"dispatch": {
"x": "min(ceilDiv((dim(shapes.y, 0) * ceilDiv(dim(shapes.y, 1), channelWidth) * dim(shapes.y, 2) * dim(shapes.y, 3)), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((dim(shapes.y, 0) * ceilDiv(dim(shapes.y, 1), channelWidth) * dim(shapes.y, 2) * dim(shapes.y, 3)), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
}
]
},
{
"id": "nchw_rank4",
"when": ["rank4Ok"],
"passes": [
{
"id": "main",
"name": "GridSample",
"shader": "grid-sample.wgsl.jinja",
"derive": {
"modeSpec": "attrs.mode",
"paddingMode": "attrs.padding_mode",
"alignCorners": "attrs.align_corners != 0"
},
"bindings": ["x", "grid", "y", "params"],
"dispatch": {
"x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
}
]
},
{
"id": "ncdhw_rank5_channel_vector",
"priority": 15,
"when": ["ranks.x == 5", "ranks.grid == 5", "ranks.y == 5", "dim(shapes.grid, 0) == dim(shapes.x, 0)", "dim(shapes.grid, 4) == 3", "dim(shapes.y, 0) == dim(shapes.x, 0)", "dim(shapes.y, 1) == dim(shapes.x, 1)", "dim(shapes.y, 2) == dim(shapes.grid, 1)", "dim(shapes.y, 3) == dim(shapes.grid, 2)", "dim(shapes.y, 4) == dim(shapes.grid, 3)", "(attrs.mode == \"linear\" or attrs.mode == \"nearest\" or attrs.mode == \"cubic\")", "(attrs.padding_mode == \"zeros\" or attrs.padding_mode == \"border\" or attrs.padding_mode == \"reflection\")", "f16Ok(dtypes.T)", "not reportedNonWave32Adapter", "dim(shapes.x, 1) >= 2", "attrs.mode != \"cubic\" or attrs.padding_mode != \"border\""],
"passes": [
{
"id": "main",
"name": "GridSample.VolumetricChannels",
"shader": "grid-sample3d.wgsl.jinja",
"derive": {
"modeSpec": "attrs.mode",
"paddingMode": "attrs.padding_mode",
"alignCorners": "attrs.align_corners != 0",
"channelWidthSpec": "volumeWidth",
"channelTail": "dim(shapes.x, 1) % volumeWidth != 0"
},
"bindings": [
"x",
"grid",
"y",
{
"name": "params",
"struct": [
{
"name": "count",
"type": "u32",
"value": "dim(shapes.y, 0) * ceilDiv(dim(shapes.y, 1), volumeWidth) * dim(shapes.y, 2) * dim(shapes.y, 3) * dim(shapes.y, 4)"
},
{ "name": "C", "type": "u32", "value": "dim(shapes.x, 1)" },
{ "name": "inD", "type": "u32", "value": "dim(shapes.x, 2)" },
{ "name": "inH", "type": "u32", "value": "dim(shapes.x, 3)" },
{ "name": "inW", "type": "u32", "value": "dim(shapes.x, 4)" },
{ "name": "outD", "type": "u32", "value": "dim(shapes.y, 2)" },
{ "name": "outH", "type": "u32", "value": "dim(shapes.y, 3)" },
{ "name": "outW", "type": "u32", "value": "dim(shapes.y, 4)" }
]
}
],
"dispatch": {
"x": "min(ceilDiv((dim(shapes.y, 0) * ceilDiv(dim(shapes.y, 1), volumeWidth) * dim(shapes.y, 2) * dim(shapes.y, 3) * dim(shapes.y, 4)), (volumeWorkgroup)), 65535)",
"y": "ceilDiv(ceilDiv((dim(shapes.y, 0) * ceilDiv(dim(shapes.y, 1), volumeWidth) * dim(shapes.y, 2) * dim(shapes.y, 3) * dim(shapes.y, 4)), (volumeWorkgroup)), 65535)",
"z": 1
}
}
]
},
{
"id": "ncdhw_rank5",
"priority": 10,
"when": ["ranks.x == 5", "ranks.grid == 5", "ranks.y == 5", "dim(shapes.grid, 0) == dim(shapes.x, 0)", "dim(shapes.grid, 4) == 3", "dim(shapes.y, 0) == dim(shapes.x, 0)", "dim(shapes.y, 1) == dim(shapes.x, 1)", "dim(shapes.y, 2) == dim(shapes.grid, 1)", "dim(shapes.y, 3) == dim(shapes.grid, 2)", "dim(shapes.y, 4) == dim(shapes.grid, 3)", "(attrs.mode == \"linear\" or attrs.mode == \"nearest\" or attrs.mode == \"cubic\")", "(attrs.padding_mode == \"zeros\" or attrs.padding_mode == \"border\" or attrs.padding_mode == \"reflection\")", "f16Ok(dtypes.T)"],
"passes": [
{
"id": "main",
"name": "GridSample.Volumetric",
"shader": "grid-sample3d.wgsl.jinja",
"derive": {
"modeSpec": "attrs.mode",
"paddingMode": "attrs.padding_mode",
"alignCorners": "attrs.align_corners != 0"
},
"bindings": [
"x",
"grid",
"y",
{
"name": "params",
"struct": [
{ "name": "count", "type": "u32", "value": "numel(shapes.y)" },
{ "name": "C", "type": "u32", "value": "dim(shapes.x, 1)" },
{ "name": "inD", "type": "u32", "value": "dim(shapes.x, 2)" },
{ "name": "inH", "type": "u32", "value": "dim(shapes.x, 3)" },
{ "name": "inW", "type": "u32", "value": "dim(shapes.x, 4)" },
{ "name": "outD", "type": "u32", "value": "dim(shapes.y, 2)" },
{ "name": "outH", "type": "u32", "value": "dim(shapes.y, 3)" },
{ "name": "outW", "type": "u32", "value": "dim(shapes.y, 4)" }
]
}
],
"dispatch": {
"x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
}
]
}
]
}