From f588b1e072f99d72af76ad9fc0ae6520068e28dc Mon Sep 17 00:00:00 2001 From: Mergen Nachin Date: Tue, 29 Sep 2026 11:02:05 -0400 Subject: [PATCH] [INITIAL] Recreate the Vulkan transformer and conformance stack with ghstack [ghstack-poisoned] --- backends/vulkan/op_registry.py | 22 +++- .../vulkan/partitioner/vulkan_partitioner.py | 24 +++- backends/vulkan/test/test_vulkan_dynamic.py | 121 ++++++++++++++++++ 3 files changed, 163 insertions(+), 4 deletions(-) diff --git a/backends/vulkan/op_registry.py b/backends/vulkan/op_registry.py index f694a3e66f3..b6483eabd3b 100644 --- a/backends/vulkan/op_registry.py +++ b/backends/vulkan/op_registry.py @@ -6,6 +6,7 @@ # pyre-unsafe +import math import operator from typing import Any, Callable, Dict, List, Optional, Tuple, Union @@ -318,13 +319,26 @@ def register_bool_binary_ops(): # ============================================================================= +def is_scalar_value_supported(value: Any, dtype: torch.dtype) -> bool: + if type(value) not in (bool, int, float): + return False + if isinstance(value, float) and math.isnan(value): + return False + if dtype in utils.INT_T: + return -(2**31) <= value <= 2**31 - 1 + return True + + @update_features(exir_ops.edge.aten.pow.Tensor_Scalar) -def register_pow_tensor_scalar(): +def register_binary_scalar_ops(): return OpFeatures( inputs_storage=utils.ANY_STORAGE, inputs_dtypes=utils.FP_T, supports_resize=True, supports_highdim=True, + are_node_inputs_supported_fn=lambda node: is_scalar_value_supported( + node.args[1], node.meta["val"].dtype + ), ) @@ -336,6 +350,9 @@ def register_eq_scalar(): outputs_dtypes=utils.BOOL_T, supports_resize=True, supports_highdim=True, + are_node_inputs_supported_fn=lambda node: is_scalar_value_supported( + node.args[1], node.args[0].meta["val"].dtype + ), ) @@ -1882,6 +1899,9 @@ def register_compare_scalar_ops(): outputs_dtypes=utils.BOOL_T, supports_resize=True, supports_highdim=True, + are_node_inputs_supported_fn=lambda node: is_scalar_value_supported( + node.args[1], node.args[0].meta["val"].dtype + ), ) diff --git a/backends/vulkan/partitioner/vulkan_partitioner.py b/backends/vulkan/partitioner/vulkan_partitioner.py index 298581ebef7..08c9f3bbd65 100644 --- a/backends/vulkan/partitioner/vulkan_partitioner.py +++ b/backends/vulkan/partitioner/vulkan_partitioner.py @@ -87,6 +87,7 @@ def __init__( self.nn_module_blocklist = nn_module_blocklist self.nn_module_allowlist = nn_module_allowlist + self._node_support: Dict[torch.fx.Node, bool] = {} def op_node_is_compatible( # noqa: C901: Function is too complex self, node: torch.fx.Node, features: Optional[OpFeatures] = None @@ -198,10 +199,27 @@ def log_skip(self, node: torch.fx.Node, reason: str) -> None: def is_node_supported( self, submodules: Mapping[str, torch.nn.Module], node: torch.fx.Node ) -> bool: - r = self._is_node_supported(node) - return r + return self._is_node_supported(node) + + def _is_node_supported(self, node: torch.fx.Node) -> bool: + if node not in self._node_support: + self._node_support[node] = self._check_node_support(node) + return self._node_support[node] + + def _check_node_support(self, node: torch.fx.Node) -> bool: # noqa: C901 + if any( + isinstance(arg.meta.get("val"), (torch.SymFloat, torch.SymBool)) + for arg in [node, *node.all_input_nodes] + ): + self.log_skip(node, "symbolic float or bool values are not supported") + return False + + if utils.is_symint_node(node) and any( + not self._is_node_supported(user) for user in node.users + ): + self.log_skip(node, "symbolic scalar has an unsupported consumer") + return False - def _is_node_supported(self, node: torch.fx.Node) -> bool: # noqa: C901 # Check if tensor node dtype is supported by vulkan if utils.is_tensor_node(node) and not utils.io_dtypes_are_supported(node): self.log_skip(node, "dtype not supported") diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index 6feb41c65a5..822dccce989 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -15,6 +15,8 @@ import torch +import torch.nn.functional as F + from executorch.backends.vulkan.partitioner.vulkan_partitioner import VulkanPartitioner from executorch.backends.vulkan.serialization.vulkan_graph_schema import ( @@ -203,6 +205,125 @@ def test_constant_bool_mask(self): ) self._run(edge, model, inputs, atol=0, rtol=0) + def test_nan_scalars_fall_back(self): + class NanScalar(torch.nn.Module): + def __init__(self, kind): + super().__init__() + self.kind = kind + + def forward(self, x): + if self.kind == "where": + return torch.where(x > 0, x, torch.nan) + if self.kind == "masked_fill": + return x.masked_fill(x > 0, torch.nan) + if self.kind == "full": + return torch.full_like(x, torch.nan) + if self.kind == "scalar_tensor": + return torch.scalar_tensor(torch.nan) + if self.kind == "pow": + return x**torch.nan + return torch.ops.aten.mul.Scalar(x, torch.nan) + + inputs = [(torch.tensor([-1.0, 0.0, 1.0, 2.0]),)] + for kind in ("where", "masked_fill", "full", "scalar_tensor", "pow", "mul"): + with self.subTest(kind=kind): + model = NanScalar(kind) + edge = self._lower(model, inputs[0], fully_delegated=False) + self._run(edge, model, inputs, atol=0, rtol=0, equal_nan=True) + + def test_dynamic_scalar_values_fall_back(self): + class DynamicScalars(torch.nn.Module): + def forward(self, x): + n = x.shape[0] + value = n * 2 + return ( + x**value, + torch.ops.aten.mul.Scalar(x, value), + torch.full((n,), value), + torch.scalar_tensor(value, dtype=torch.int64), + torch.ops.aten.mul.Scalar(x, n * 0.5), + x + torch.full((n,), n * 0.5), + torch.full((n,), 0.5), + F.gelu(x), + torch.clamp(x, max=n * 0.5), + F.leaky_relu(x, negative_slope=n * 0.1), + ) + + model = DynamicScalars() + inputs = [(torch.linspace(-0.9, 4.1, n),) for n in (4, 2, 7, 3, 4)] + edge = self._lower( + model, inputs[0], ({0: Dim("n", min=2, max=8)},), fully_delegated=False + ) + self.assertTrue(_vulkan_graphs(edge)) + self._run(edge, model, inputs) + + def test_dynamic_compare_scalars_fall_back(self): + class Compare(torch.nn.Module): + def __init__(self, op): + super().__init__() + self.op = op + + def forward(self, x): + return self.op(x, x.shape[1]) + + inputs = [ + (torch.arange(2 * s, dtype=torch.float32).reshape(2, s),) + for s in (16, 3, 31, 2, 16) + ] + for op in (torch.eq, torch.ne, torch.lt, torch.le, torch.gt, torch.ge): + with self.subTest(op=op): + model = Compare(op) + edge = self._lower( + model, + inputs[0], + ({1: Dim("s", min=2, max=32)},), + fully_delegated=False, + ) + self.assertEqual(_vulkan_graphs(edge), []) + self._run(edge, model, inputs, atol=0, rtol=0) + + def test_compare_scalar_values_fall_back(self): + class Compare(torch.nn.Module): + def __init__(self, op, value): + super().__init__() + self.op = op + self.value = value + + def forward(self, x): + return self.op(x, self.value) + + for op in (torch.eq, torch.ne, torch.lt, torch.le, torch.gt, torch.ge): + for x, value in ( + (torch.tensor([-1.0, 0.0, 1.0, 2.0]), torch.nan), + (torch.tensor([-3, 0, 1, 7], dtype=torch.int32), 2**40), + (torch.tensor([-(2**40), 0, 2**40, 2**40 + 1]), 2**40), + ): + with self.subTest(op=op, dtype=x.dtype, value=value): + model = Compare(op, value) + inputs = [(x,)] + edge = self._lower(model, inputs[0], fully_delegated=False) + self.assertEqual(_vulkan_graphs(edge), []) + self._run(edge, model, inputs, atol=0, rtol=0) + + def test_compare_scalar_values(self): + class Compare(torch.nn.Module): + def __init__(self, op, value): + super().__init__() + self.op = op + self.value = value + + def forward(self, x): + return self.op(x, self.value) + + for op in (torch.eq, torch.ne, torch.lt, torch.le, torch.gt, torch.ge): + for dtype in (torch.int32, torch.float32): + for value in (2.0, 2.5, -1.5): + with self.subTest(op=op, dtype=dtype, value=value): + x = torch.arange(-7, 14, dtype=dtype).reshape(3, 7) + model = Compare(op, value) + edge = self._lower(model, (x,)) + self._run(edge, model, [(x,)], atol=0, rtol=0) + def test_4d_reductions(self): class Reduce(torch.nn.Module): def __init__(self, op, dim):