diff --git a/backends/vulkan/op_registry.py b/backends/vulkan/op_registry.py index b5ba2774e8e..6d672a6ceb4 100644 --- a/backends/vulkan/op_registry.py +++ b/backends/vulkan/op_registry.py @@ -329,7 +329,12 @@ def is_scalar_value_supported(value: Any, dtype: torch.dtype) -> bool: return True -@update_features(exir_ops.edge.aten.pow.Tensor_Scalar) +@update_features( + [ + exir_ops.edge.aten.pow.Tensor_Scalar, + exir_ops.edge.aten.mul.Scalar, + ] +) def register_binary_scalar_ops(): return OpFeatures( inputs_storage=utils.ANY_STORAGE, diff --git a/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_buffer.yaml b/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_buffer.yaml index 89e28ba9acd..ff0275908f5 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_buffer.yaml +++ b/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_buffer.yaml @@ -21,6 +21,8 @@ binary_scalar_buffer: - parameter_values: [int32, float] shader_variants: - NAME: pow_scalar_buffer + - NAME: mul_scalar_buffer + OPERATOR: X * Y - NAME: eq_scalar_buffer OPERATOR: X == Y IS_COMPARISON_OP: true diff --git a/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_texture.yaml b/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_texture.yaml index 7af25fe13a7..d4971257356 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_texture.yaml +++ b/backends/vulkan/runtime/graph/ops/glsl/binary_scalar_texture.yaml @@ -21,6 +21,8 @@ binary_scalar_texture: - parameter_values: [int32, float] shader_variants: - NAME: pow_scalar_texture3d + - NAME: mul_scalar_texture3d + OPERATOR: X * Y - NAME: eq_scalar_texture3d OPERATOR: equal(X, Y) IS_COMPARISON_OP: true diff --git a/backends/vulkan/runtime/graph/ops/impl/BinaryScalarOp.cpp b/backends/vulkan/runtime/graph/ops/impl/BinaryScalarOp.cpp index 352c2987180..b0e29e8d372 100644 --- a/backends/vulkan/runtime/graph/ops/impl/BinaryScalarOp.cpp +++ b/backends/vulkan/runtime/graph/ops/impl/BinaryScalarOp.cpp @@ -123,6 +123,10 @@ void pow_tensor_scalar(ComputeGraph& graph, const std::vector& args) { return add_binary_scalar_op_node(graph, args[0], args[1], args[2], "pow"); } +void mul_tensor_scalar(ComputeGraph& graph, const std::vector& args) { + return add_binary_scalar_op_node(graph, args[0], args[1], args[2], "mul"); +} + void eq_tensor_scalar(ComputeGraph& graph, const std::vector& args) { return add_binary_scalar_op_node(graph, args[0], args[1], args[2], "eq"); } @@ -149,6 +153,7 @@ void ge_tensor_scalar(ComputeGraph& graph, const std::vector& args) { REGISTER_OPERATORS { VK_REGISTER_OP(aten.pow.Tensor_Scalar, pow_tensor_scalar); + VK_REGISTER_OP(aten.mul.Scalar, mul_tensor_scalar); VK_REGISTER_OP(aten.eq.Scalar, eq_tensor_scalar); VK_REGISTER_OP(aten.ne.Scalar, ne_tensor_scalar); VK_REGISTER_OP(aten.lt.Scalar, lt_tensor_scalar); diff --git a/backends/vulkan/test/op_tests/cases.py b/backends/vulkan/test/op_tests/cases.py index c99b7ad58b0..3699ca35874 100644 --- a/backends/vulkan/test/op_tests/cases.py +++ b/backends/vulkan/test/op_tests/cases.py @@ -2271,8 +2271,8 @@ def get_index_tensor_inputs(): return test_suite -@register_test_suite("aten.pow.Tensor_Scalar") -def get_pow_tensor_scalar_inputs(): +@register_test_suite(["aten.pow.Tensor_Scalar", "aten.mul.Scalar"]) +def get_binary_scalar_inputs(): test_suite = VkTestSuite( [ ((M1,), 2.0), diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index 3b049925c06..af75c099112 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -347,6 +347,39 @@ def forward(self, x): ) self._run(edge, model, inputs, atol=0, rtol=0) + def test_signed_zero_scalars(self): + class SignedZero(torch.nn.Module): + def forward(self, x): + return ( + torch.ops.aten.mul.Scalar(x, -0.0), + torch.full_like(x, -0.0), + torch.scalar_tensor(-0.0, dtype=x.dtype), + ) + + model = SignedZero() + for dtype in (torch.float32, torch.float16): + for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER): + with self.subTest(dtype=dtype, storage=storage): + inputs = [(torch.arange(-7, 14, dtype=dtype).reshape(3, 7),)] + edge = self._lower( + model, inputs[0], storage=storage, fully_delegated=False + ) + self.assertTrue(_vulkan_graphs(edge)) + self.assertTrue( + all( + node.target + in ( + operator.getitem, + torch.ops.higher_order.executorch_call_delegate, + ) + for node in edge.exported_program().graph.nodes + if node.op == "call_function" + ) + ) + self._run( + edge, model, inputs, atol=0, rtol=0, check_signed_zero=True + ) + def test_integer_scalar_range_fallback(self): class LargeScalar(torch.nn.Module): def __init__(self, kind, value):