| { |
| "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 |
| } |
| } |
| ] |
| } |
| ] |
| } |
|
|