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