diff --git a/backends/vulkan/op_registry.py b/backends/vulkan/op_registry.py index e2df399122c..248b5924df4 100644 --- a/backends/vulkan/op_registry.py +++ b/backends/vulkan/op_registry.py @@ -720,9 +720,8 @@ def is_reduce_node_supported_by_general_impl(node: torch.fx.Node) -> bool: if isinstance(dims_reduced, (list, tuple)) and not 1 <= len(dims_reduced) <= 2: return False - keepdim = get_keepdim_setting(node) - # keepdim = False is not supported yet for general implementation - if isinstance(keepdim, bool) and not keepdim: + # any.dim can repack the reduced texture after removing the reduction axis. + if not get_keepdim_setting(node) and node.target != exir_ops.edge.aten.any.dim: return False if utils.ndim_of(node.args[0]) == 4: @@ -812,6 +811,20 @@ def register_reduce_cpp_ops(): ) +@update_features(exir_ops.edge.aten.any.dim) +def register_any_dim(): + return OpFeatures( + inputs_storage=utils.ANY_TEXTURE, + inputs_dtypes=utils.BOOL_T, + supports_resize=True, + supports_highdim=True, + are_node_inputs_supported_fn=lambda node: ( + utils.ndim_of(node.args[0]) > 0 and is_reduce_node_supported(node) + ), + pick_io_storage_fn=pick_storage_for_reduce, + ) + + # ============================================================================= # ArgReduce.cpp # ============================================================================= diff --git a/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl b/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl index f37a0fc773a..4365256174e 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl +++ b/backends/vulkan/runtime/graph/ops/glsl/reduce.glsl @@ -10,6 +10,7 @@ #define PRECISION ${PRECISION} #define VEC4_T ${texel_load_type(DTYPE, STORAGE)} +#define T ${texel_load_component_type(DTYPE, STORAGE)} ${define_active_storage_type(STORAGE)} @@ -45,7 +46,7 @@ layout(constant_id = 6) const int NWORKERS = 4; #define MAX_NTHREADS 256 -shared vec4 shared_vecs[MAX_NTHREADS]; +shared VEC4_T shared_vecs[MAX_NTHREADS]; #include "indexing_utils.h" #include "indexing.glslh" @@ -123,7 +124,7 @@ void reduce_nonpacked_dim( // behaviour that hangs some GPUs. They still take a shared memory slot, but // it is one that no in-bounds group aggregates over, so what they leave in it // is never read. - vec4 accum = vec4(0); + VEC4_T accum = VEC4_T(0); if (in_bounds) { scan_pos[reduce_dim] = 0; accum = INIT_ACCUM(load_texel(tin, scan_pos)); @@ -195,10 +196,10 @@ void reduce_packed_dim( // behaviour that hangs some GPUs. They still take a shared memory slot, but // it is one that no in-bounds group aggregates over, so what they leave in it // is never read. - vec4 accum = vec4(0); + VEC4_T accum = VEC4_T(0); if (in_bounds) { scan_pos[reduce_dim] = 0; - accum = INIT_ACCUM(vec4(load_texel(tin, scan_pos).x)); + accum = INIT_ACCUM(VEC4_T(load_texel(tin, scan_pos).x)); // Partially accumulate over elements i, i + NWORKERS, i + 2*NWORKERS, ... // of the reduction row @@ -212,7 +213,7 @@ void reduce_packed_dim( // padding elements are ignored if (scan_pos[reduce_dim] == safe_idx(tin_limits, reduce_dim) - 1 && nspill > 0) { - const vec4 intex = load_texel(tin, scan_pos); + const VEC4_T intex = load_texel(tin, scan_pos); for (int i = 0; i < nspill; i++) { accum.x = UPDATE_ACCUM(accum.x, intex[i]); } @@ -233,13 +234,13 @@ void reduce_packed_dim( } // Each element of the texel is itself a partial maximum; iterate over the // texel to find the actual maximum - float accum_final = accum.x; + T accum_final = accum.x; [[unroll]] for (int i = 1; i < 4; i++) { accum_final = UPDATE_ACCUM(accum[i], accum_final); } scan_pos[reduce_dim] = tid.x; - write_texel(tout, scan_pos, POSTPROCESS(vec4(accum_final, 0, 0, 0))); + write_texel(tout, scan_pos, POSTPROCESS(VEC4_T(accum_final, 0, 0, 0))); } } @@ -251,7 +252,10 @@ void main() { gl_LocalInvocationID[reduce_dim], gl_LocalInvocationID[group_dim]); - const bool in_bounds = all(lessThan(scan_pos, tin_limits)); + // Reducing an empty dimension still produces one output element. + ivec3 out_limits = tin_limits; + out_limits[reduce_dim] = 1; + const bool in_bounds = all(lessThan(scan_pos, out_limits)); // reduce_dim and packed_dim are specialization constants, so this branch is // uniform across the work group and safe to take around a barrier. diff --git a/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml b/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml index f85d8405b2d..550aeea093a 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml +++ b/backends/vulkan/runtime/graph/ops/glsl/reduce.yaml @@ -17,6 +17,9 @@ reduce: - VALUE: float shader_variants: - NAME: sum + - NAME: any_uint8 + DTYPE: uint8 + UPDATE_ACCUM: max(accum, new_val) - NAME: mean POSTPROCESS: (accum / tin_sizes[reduce_dim]) - NAME: amax diff --git a/backends/vulkan/runtime/graph/ops/glsl/reduce_per_row_buffer.yaml b/backends/vulkan/runtime/graph/ops/glsl/reduce_per_row_buffer.yaml index e5a94165b96..e4850fb1a44 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/reduce_per_row_buffer.yaml +++ b/backends/vulkan/runtime/graph/ops/glsl/reduce_per_row_buffer.yaml @@ -20,6 +20,10 @@ reduce_per_row_buffer: - VALUE: int32 shader_variants: - NAME: sum_per_row_buffer + - NAME: any_per_row_buffer_uint8 + DTYPE: uint8 + UPDATE_ACCUM_FN: update_accum_amax + MERGE_ACCUM_FN: merge_accum_amax - NAME: mean_per_row_buffer POSTPROCESS_ACCUM_FN: postprocess_accum_mean - NAME: amax_per_row_buffer diff --git a/backends/vulkan/runtime/graph/ops/impl/Reduce.cpp b/backends/vulkan/runtime/graph/ops/impl/Reduce.cpp index 4c5b14976c0..2fade74cc44 100644 --- a/backends/vulkan/runtime/graph/ops/impl/Reduce.cpp +++ b/backends/vulkan/runtime/graph/ops/impl/Reduce.cpp @@ -10,6 +10,7 @@ #include #include +#include #include #include @@ -419,11 +420,50 @@ DEFINE_REDUCE_FN(mean, 4) DEFINE_REDUCE_FN(amax, 3) DEFINE_REDUCE_FN(amin, 3) +void resize_squeezed_reduce_node( + ComputeGraph* graph, + const std::vector& args, + const std::vector& resize_args) { + const ValueRef out = args.at(0).refs.at(0); + const ValueRef in = args.at(1).refs.at(0); + std::vector sizes = graph->sizes_of(in); + const int64_t dim = normalize( + graph->extract_scalar(resize_args.at(0)), sizes.size()); + sizes.erase(sizes.begin() + dim); + graph->virtual_resize(out, sizes); +} + +void any_dim(ComputeGraph& graph, const std::vector& args) { + if (graph.is_buffer_storage(args[0])) { + VK_CHECK_COND( + normalize( + graph.extract_scalar(args[1]), graph.dim_of(args[0])) == + graph.dim_of(args[0]) - 1); + return add_reduce_per_row_node(graph, args[0], args[2], args[3], "any"); + } + if (graph.extract_scalar(args[2])) { + return add_reduce_node(graph, args[0], args[1], args[3], "any"); + } + + std::vector sizes = graph.sizes_of(args[0]); + sizes.at(normalize(graph.extract_scalar(args[1]), sizes.size())) = 1; + TmpTensor reduced( + &graph, + sizes, + graph.dtype_of(args[0]), + graph.storage_type_of(args[0]), + graph.estimate_memory_layout_of(args[0])); + add_reduce_node(graph, args[0], args[1], reduced, "any"); + return add_view_copy_node( + graph, reduced, args[3], {args[1]}, resize_squeezed_reduce_node); +} + REGISTER_OPERATORS { VK_REGISTER_OP(aten.sum.dim_IntList, sum); VK_REGISTER_OP(aten.mean.dim, mean); VK_REGISTER_OP(aten.amax.default, amax); VK_REGISTER_OP(aten.amin.default, amin); + VK_REGISTER_OP(aten.any.dim, any_dim); } } // namespace vkcompute diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index 5b62563b944..ef1aa9a7856 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -149,6 +149,27 @@ def _run( ) ) + def test_partition_any_unsupported_inputs(self): + class AnyDim(torch.nn.Module): + def forward(self, x): + return torch.any(x, dim=0, keepdim=True) + + for x in ( + torch.tensor(True), + torch.tensor([0, 1], dtype=torch.int32), + torch.tensor([0, 1], dtype=torch.uint8), + torch.tensor([0.0, 1.0]), + ): + with self.subTest(shape=x.shape, dtype=x.dtype): + edge = to_edge_transform_and_lower( + export(AnyDim(), (x,)), + partitioner=[VulkanPartitioner({"require_dynamic_shapes": True})], + ) + self.assertNotIn( + torch.ops.higher_order.executorch_call_delegate, + [node.target for node in edge.exported_program().graph.nodes], + ) + def test_dynamic_gelu(self): for approximate in ("none", "tanh"): for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER): @@ -254,6 +275,127 @@ def forward(self, x): ) self._run(edge, model, inputs, atol=0, rtol=0) + def test_any_without_keepdim_uses_textures(self): + class AnyDim(torch.nn.Module): + def __init__(self, keepdim): + super().__init__() + self.keepdim = keepdim + + def forward(self, x): + mask = x > 0 + if self.keepdim is None: + reduced = torch.any(mask, dim=-1) + else: + reduced = torch.any(mask, dim=-1, keepdim=self.keepdim) + return torch.logical_not(reduced), torch.any(reduced, dim=-1) + + inputs = [ + ((torch.arange(2 * s * 5).reshape(2, s, 5) % 17).float() - 14,) + for s in (16, 3, 31, 2, 16) + ] + for keepdim in (None, False): + with self.subTest(keepdim=keepdim): + model = AnyDim(keepdim) + edge = self._lower( + model, + inputs[0], + ({1: Dim("s", min=2, max=32)},), + ) + for graph in _vulkan_graphs(edge): + for value in graph.values: + tensor = value.value + if isinstance(tensor, VkTensor) and tensor.constant_id < 0: + self.assertEqual( + tensor.storage_type, VkStorageType.TEXTURE_3D + ) + self._run(edge, model, inputs, atol=0, rtol=0) + + def test_any_texture_without_keepdim_shapes(self): + class AnyDim(torch.nn.Module): + def __init__(self, dim): + super().__init__() + self.dim = dim + + def forward(self, x): + return torch.any(x, dim=self.dim) + + for shape, dim, supported in ( + ((7,), 0, True), + ((0,), 0, True), + ((3, 5), 0, True), + ((3, 5), -1, True), + ((0, 3), 0, True), + ((2, 0, 3), 1, True), + ((2, 0, 3), -1, True), + ((1, 3, 5, 7), 1, True), + ((1, 3, 5, 7), -2, True), + ((2, 3, 5, 7), -1, True), + ((2, 3, 5, 7), 2, True), + ((1, 3, 1, 7), 1, True), + ((1, 3, 5, 7), 0, False), + ((2, 3, 5, 7), 1, False), + ): + with self.subTest(shape=shape, dim=dim): + x = torch.zeros(shape, dtype=torch.bool) + if x.numel() > 0: + x[tuple(size // 2 for size in shape)] = True + model = AnyDim(dim) + edge = self._lower(model, (x,), fully_delegated=supported) + if not supported: + self.assertEqual(_vulkan_graphs(edge), []) + for graph in _vulkan_graphs(edge): + for value_id in graph.input_ids + graph.output_ids: + tensor = graph.values[value_id].value + self.assertEqual(tensor.storage_type, VkStorageType.TEXTURE_3D) + self._run( + edge, + model, + [(x,), (torch.zeros_like(x),), (torch.ones_like(x),)], + atol=0, + rtol=0, + ) + + def test_dynamic_any_dim(self): + class AnyDim(torch.nn.Module): + def __init__(self, dim, keepdim): + super().__init__() + self.dim = dim + self.keepdim = keepdim + + def forward(self, x): + return torch.any(x, dim=self.dim, keepdim=self.keepdim) + + for dim, keepdim in ( + (-1, True), + (-1, False), + (1, True), + (1, False), + (0, False), + ): + for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER): + with self.subTest(dim=dim, keepdim=keepdim, storage=storage): + model = AnyDim(dim, keepdim) + inputs = [] + for s in (16, 3, 31, 2, 0, 16): + x = torch.zeros(2, s, 5, dtype=torch.bool) + if s not in (0, 3): + x[0, s // 2, 1] = True + x[1, -1, 3] = True + x[1, 0, 4] = True + inputs.append((x,)) + seq = Dim("s", min=0, max=32) + unsupported = dim != -1 and storage == VkStorageType.BUFFER + edge = self._lower( + model, + inputs[0], + ({1: seq},), + storage, + fully_delegated=not unsupported, + ) + if unsupported: + self.assertEqual(_vulkan_graphs(edge), []) + self._run(edge, model, inputs, atol=0, rtol=0) + def test_dynamic_logical_not(self): class LogicalNot(torch.nn.Module): def forward(self, x): @@ -706,7 +848,7 @@ def __init__(self, op, dim): def forward(self, x): return self.op(x, dim=self.dim, keepdim=True) - for op in (torch.sum, torch.mean, torch.amax): + for op in (torch.any, torch.sum, torch.mean, torch.amax): for batch, dim, supported in ( (1, 0, False), (2, 0, False), @@ -717,7 +859,11 @@ def forward(self, x): ): with self.subTest(op=op, batch=batch, dim=dim): values = torch.arange(batch * 3 * 4 * 5).reshape(batch, 3, 4, 5) - x = -((values * 37 + 11) % values.numel() + 1).float() / 7 + x = ( + values % 7 == 0 + if op == torch.any + else -((values * 37 + 11) % values.numel() + 1).float() / 7 + ) model = Reduce(op, dim) edge = self._lower(model, (x,), fully_delegated=supported) if not supported: