From 61b4851f6ea8cc30af5eab2bc38166782dbebeef Mon Sep 17 00:00:00 2001 From: Mergen Nachin Date: Mon, 28 Sep 2026 18:19:33 -0400 Subject: [PATCH] [Vulkan] Keep unrepresentable and symbolic scalars out of the delegate The binary-scalar and comparison kernels receive scalars as int32 or float, yet the partitioner accepted NaN, integers outside int32 and symbolic scalars, which lowered with wrong values or failed in the graph builder. is_scalar_value_supported rejects them for pow.Tensor_Scalar and the comparison scalar ops. The partitioner also rejects SymFloat and SymBool values and keeps SymInt producers on CPU when a consumer is unsupported; node support is memoized so that check stays linear. Authored with OpenAI Codex; split planned with Claude Code. --- 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):