diff --git a/backends/vulkan/op_registry.py b/backends/vulkan/op_registry.py index add4d01a78e..f1fbb786d1f 100644 --- a/backends/vulkan/op_registry.py +++ b/backends/vulkan/op_registry.py @@ -646,7 +646,12 @@ def register_q8ta_pixel_shuffle(): def get_dims_reduced(node: torch.fx.Node) -> Union[int, List[int]]: ndim = utils.ndim_of(node.args[0]) assert ndim is not None - dims_reduced = None + dims_reduced = ( + [] + if node.target + in (exir_ops.edge.aten.amax.default, exir_ops.edge.aten.amin.default) + else None + ) if len(node.args) >= 2: dims_reduced = node.args[1] @@ -690,7 +695,7 @@ def is_reduce_node_supported_by_per_row_impl(node: torch.fx.Node) -> bool: def is_reduce_node_supported_by_general_impl(node: torch.fx.Node) -> bool: dims_reduced = get_dims_reduced(node) # Only 1D and 2D reductions are supported at the moment. - if isinstance(dims_reduced, (list, tuple)) and len(dims_reduced) > 2: + if isinstance(dims_reduced, (list, tuple)) and not 1 <= len(dims_reduced) <= 2: return False keepdim = get_keepdim_setting(node) @@ -698,10 +703,24 @@ def is_reduce_node_supported_by_general_impl(node: torch.fx.Node) -> bool: if isinstance(keepdim, bool) and not keepdim: return False + if utils.ndim_of(node.args[0]) == 4: + dims = [dims_reduced] if isinstance(dims_reduced, int) else dims_reduced + # Textures fold batch into channels; neither axis can be reduced across batches. + if 0 in dims or ( + 1 in dims and utils.upper_bound_size(node.args[0].meta["val"].shape[0]) != 1 + ): + return False + return True def is_reduce_node_supported(node: torch.fx.Node) -> bool: + if ( + node.target in (exir_ops.edge.aten.sum.dim_IntList, exir_ops.edge.aten.mean.dim) + and (len(node.args) < 2 or node.args[1] is None) + and utils.ndim_of(node.args[0]) != 1 + ): + return False return is_reduce_node_supported_by_per_row_impl( node ) or is_reduce_node_supported_by_general_impl(node) @@ -776,6 +795,15 @@ def register_reduce_cpp_ops(): # ============================================================================= +def is_argreduce_node_supported(node: torch.fx.Node) -> bool: + ndim = utils.ndim_of(node.args[0]) + assert ndim is not None + dim = node.args[1] if len(node.args) > 1 else None + if dim is None: + return ndim == 1 + return ndim > 0 and utils.normalize_dims(dim, ndim) == ndim - 1 + + @update_features( [ exir_ops.edge.aten.argmax.default, @@ -784,13 +812,12 @@ def register_reduce_cpp_ops(): ) def register_argreduce_cpp_ops(): return OpFeatures( - inputs_storage=utils.ANY_STORAGE, + inputs_storage=utils.CONTIGUOUS_BUFFER, inputs_dtypes=utils.FP_T, outputs_dtypes=utils.INT_T, supports_resize=True, supports_highdim=True, - are_node_inputs_supported_fn=is_reduce_node_supported, - pick_io_storage_fn=pick_storage_for_reduce, + are_node_inputs_supported_fn=is_argreduce_node_supported, ) diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index f911d7ca22d..d97d659d916 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -166,6 +166,34 @@ def test_gelu_with_singleton_dimensions(self): edge = self._lower(model, (x,), storage=storage) self._run(edge, model, [(x,)], atol=5e-6, rtol=5e-6) + def test_4d_reductions(self): + class Reduce(torch.nn.Module): + def __init__(self, op, dim): + super().__init__() + self.op = op + self.dim = dim + + def forward(self, x): + return self.op(x, dim=self.dim, keepdim=True) + + for op in (torch.sum, torch.mean, torch.amax): + for batch, dim, supported in ( + (1, 0, False), + (2, 0, False), + (2, 1, False), + (1, 1, True), + (2, 2, True), + (2, -1, True), + ): + 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 + model = Reduce(op, dim) + edge = self._lower(model, (x,), fully_delegated=supported) + if not supported: + self.assertEqual(_vulkan_graphs(edge), []) + self._run(edge, model, [(x,)]) + def test_buffer_reduction_range(self): class Reduce(torch.nn.Module): def __init__(self, op): @@ -185,6 +213,73 @@ def forward(self, x): edge = self._lower(model, (x,), storage=VkStorageType.BUFFER) self._run(edge, model, [(x,)], atol=0, rtol=0) + def test_argreduce_dims(self): + class Reduce(torch.nn.Module): + def __init__(self, op, dim, keepdim): + super().__init__() + self.op = op + self.dim = dim + self.keepdim = keepdim + + def forward(self, x): + return self.op(x, dim=self.dim, keepdim=self.keepdim) + + for op in (torch.argmax, torch.argmin): + for keepdim in (True, False): + for shape, dim, supported in ( + ((1, 8), None, False), + ((3, 8), None, False), + ((3, 8), 0, False), + ((3, 8), -2, False), + ((8,), None, True), + ((3, 8), 1, True), + ((3, 8), -1, True), + ): + with self.subTest(op=op, keepdim=keepdim, shape=shape, dim=dim): + x = ((torch.arange(math.prod(shape)) * 5 + 3) % 17).float() + x = x.reshape(shape) + model = Reduce(op, dim, keepdim) + edge = self._lower(model, (x,), fully_delegated=supported) + graphs = _vulkan_graphs(edge) + if supported: + (graph,) = graphs + for value_id in graph.input_ids + graph.output_ids: + self.assertEqual( + graph.values[value_id].value.storage_type, + VkStorageType.BUFFER, + ) + else: + self.assertEqual(len(graphs), 0) + self._run(edge, model, [(x,)], atol=0, rtol=0) + + def test_unsupported_reduction_dims_fall_back(self): + class Reduce(torch.nn.Module): + def __init__(self, op, keepdim, dims): + super().__init__() + self.op = op + self.keepdim = keepdim + self.dims = dims + + def forward(self, x): + return self.op(x, dim=self.dims, keepdim=self.keepdim) + + for op in (torch.sum, torch.mean, torch.amax, torch.amin): + for keepdim in (False, True): + for dims, shape in ( + ([], (8,)), + ([], (2, 3, 5)), + (None, (8, 3)), + (None, (1, 8)), + ): + if dims is None and op not in (torch.sum, torch.mean): + continue + with self.subTest(op=op, keepdim=keepdim, dims=dims, shape=shape): + x = torch.linspace(-4, 3, math.prod(shape)).reshape(shape) + model = Reduce(op, keepdim, dims) + edge = self._lower(model, (x,), fully_delegated=False) + self.assertEqual(_vulkan_graphs(edge), []) + self._run(edge, model, [(x,)]) + def test_int32_buffer_reduction_shader_range(self): from executorch.extension.pybindings.portable_lib import ( _load_for_executorch_from_buffer,