Skip to content
Open
22 changes: 21 additions & 1 deletion backends/vulkan/op_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@

# pyre-unsafe

import math
import operator
from typing import Any, Callable, Dict, List, Optional, Tuple, Union

Expand Down Expand Up @@ -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
),
)


Expand All @@ -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
),
)


Expand Down Expand Up @@ -1890,6 +1907,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
),
)


Expand Down
24 changes: 21 additions & 3 deletions backends/vulkan/partitioner/vulkan_partitioner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")
Expand Down
121 changes: 121 additions & 0 deletions backends/vulkan/test/test_vulkan_dynamic.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,8 @@

import torch

import torch.nn.functional as F

from executorch.backends.vulkan.partitioner.vulkan_partitioner import (
parse_compile_options,
VulkanPartitioner,
Expand Down Expand Up @@ -208,6 +210,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):
Expand Down
Loading