Skip to content
Draft
7 changes: 6 additions & 1 deletion backends/vulkan/op_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
5 changes: 5 additions & 0 deletions backends/vulkan/runtime/graph/ops/impl/BinaryScalarOp.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,10 @@ void pow_tensor_scalar(ComputeGraph& graph, const std::vector<ValueRef>& args) {
return add_binary_scalar_op_node(graph, args[0], args[1], args[2], "pow");
}

void mul_tensor_scalar(ComputeGraph& graph, const std::vector<ValueRef>& args) {
return add_binary_scalar_op_node(graph, args[0], args[1], args[2], "mul");
}

void eq_tensor_scalar(ComputeGraph& graph, const std::vector<ValueRef>& args) {
return add_binary_scalar_op_node(graph, args[0], args[1], args[2], "eq");
}
Expand All @@ -149,6 +153,7 @@ void ge_tensor_scalar(ComputeGraph& graph, const std::vector<ValueRef>& 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);
Expand Down
4 changes: 2 additions & 2 deletions backends/vulkan/test/op_tests/cases.py
Original file line number Diff line number Diff line change
Expand Up @@ -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),
Expand Down
51 changes: 51 additions & 0 deletions backends/vulkan/test/test_vulkan_dynamic.py
Original file line number Diff line number Diff line change
Expand Up @@ -352,6 +352,57 @@ 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_fp16_chained_mul_scalar(self):
class ChainedMul(torch.nn.Module):
def forward(self, x):
return torch.ops.aten.mul.Scalar(
torch.ops.aten.mul.Scalar(x, 1.0006), 1000
)

model = ChainedMul()
x = torch.ones(3, 7, dtype=torch.float16)
for storage in (VkStorageType.TEXTURE_3D, VkStorageType.BUFFER):
with self.subTest(storage=storage):
edge = self._lower(model, (x,), storage=storage)
operators = [
op.name for graph in _vulkan_graphs(edge) for op in graph.chain
]
self.assertEqual(operators.count("aten.mul.Scalar"), 2)
self._run(edge, model, [(x,)], atol=0, rtol=0)

def test_integer_scalar_range_fallback(self):
class LargeScalar(torch.nn.Module):
def __init__(self, kind, value):
Expand Down
Loading