diff --git a/backends/vulkan/op_registry.py b/backends/vulkan/op_registry.py index add4d01a78e..f694a3e66f3 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) diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index 089c14eae4b..cfb8de62f54 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -161,6 +161,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): @@ -180,6 +208,34 @@ def forward(self, x): edge = self._lower(model, (x,), storage=VkStorageType.BUFFER) 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_reduction_special_values(self): class Reduce(torch.nn.Module): def __init__(self, op):