Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
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 @@ -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
),
)


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 VulkanPartitioner

from executorch.backends.vulkan.serialization.vulkan_graph_schema import (
Expand Down Expand Up @@ -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):
Expand Down
Loading