diff --git a/backends/vulkan/op_registry.py b/backends/vulkan/op_registry.py index b6483eabd3b..da7d8f0cde0 100644 --- a/backends/vulkan/op_registry.py +++ b/backends/vulkan/op_registry.py @@ -1602,12 +1602,21 @@ def register_full_cpp_ops(): # ============================================================================= -@update_features(exir_ops.edge.aten.scalar_tensor.default) +@update_features( + [ + exir_ops.edge.aten.scalar_tensor.default, + # EXIR deliberately keeps scalar_tensor in the ATen dialect. + torch.ops.aten.scalar_tensor.default, + ] +) def register_scalar_tensor(): return OpFeatures( inputs_storage=utils.CHANNELS_PACKED_TEXTURE, inputs_dtypes=utils.FP_INT_T, supports_resize=True, + are_node_inputs_supported_fn=lambda node: is_scalar_value_supported( + node.args[0], node.meta["val"].dtype + ), ) diff --git a/backends/vulkan/runtime/graph/ops/glsl/scalar_tensor.glsl b/backends/vulkan/runtime/graph/ops/glsl/scalar_tensor.glsl index 372c233ea7c..67a71aca9c5 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/scalar_tensor.glsl +++ b/backends/vulkan/runtime/graph/ops/glsl/scalar_tensor.glsl @@ -19,6 +19,7 @@ ${define_explicit_type_extensions(SCALAR_VALUE_TYPE)} ${define_active_storage_type(STORAGE)} #include "indexing_utils.h" +#include "convert.glslh" layout(std430) buffer; @@ -52,6 +53,8 @@ void main() { } VEC4_T outtex = VEC4_T(scalar_value); + $if DTYPE == "half": + outtex = round_to_half_rte(outtex); write_texel(t_out, pos, outtex); } diff --git a/backends/vulkan/runtime/graph/ops/glsl/scalar_tensor.yaml b/backends/vulkan/runtime/graph/ops/glsl/scalar_tensor.yaml index cd45b80c4dc..fff3d513c41 100644 --- a/backends/vulkan/runtime/graph/ops/glsl/scalar_tensor.yaml +++ b/backends/vulkan/runtime/graph/ops/glsl/scalar_tensor.yaml @@ -12,16 +12,14 @@ scalar_tensor: PACKING: C_packed STORAGE: texture3d generate_variant_forall: - DTYPE: - - VALUE: half - - VALUE: float - - VALUE: int32 - STORAGE: - - VALUE: texture3d - - VALUE: buffer - SCALAR_VALUE_TYPE: - - VALUE: float - - VALUE: int32 - - VALUE: bool + combination: + parameter_names: [DTYPE, STORAGE, SCALAR_VALUE_TYPE] + combos: + - parameter_values: [half, texture3d, float] + - parameter_values: [half, buffer, float] + - parameter_values: [float, texture3d, float] + - parameter_values: [float, buffer, float] + - parameter_values: [int32, texture3d, int32] + - parameter_values: [int32, buffer, int32] shader_variants: - NAME: scalar_tensor diff --git a/backends/vulkan/runtime/graph/ops/impl/ScalarTensor.cpp b/backends/vulkan/runtime/graph/ops/impl/ScalarTensor.cpp index ca2fe79b7c4..23bda24d169 100644 --- a/backends/vulkan/runtime/graph/ops/impl/ScalarTensor.cpp +++ b/backends/vulkan/runtime/graph/ops/impl/ScalarTensor.cpp @@ -16,17 +16,21 @@ namespace vkcompute { void scalar_tensor(ComputeGraph& graph, const std::vector& args) { // Extract the scalar value from the first argument ValueRef scalar_in = args[0]; - float scalar_value = graph.extract_scalar(scalar_in); // Get the output tensor reference ValueRef out = args[args.size() - 1]; + const vkapi::ScalarType scalar_dtype = + graph.dtype_of(out) == vkapi::kInt ? vkapi::kInt : vkapi::kFloat; + const vkapi::BufferBindInfo scalar_buffer = scalar_dtype == vkapi::kInt + ? graph.create_params_buffer(graph.extract_scalar(scalar_in)) + : graph.create_params_buffer(graph.extract_scalar(scalar_in)); std::string kernel_name("scalar_tensor"); kernel_name.reserve(kShaderNameReserve); add_dtype_suffix(kernel_name, graph.dtype_of(out)); add_storage_type_suffix(kernel_name, graph.storage_type_of(out)); - add_dtype_suffix(kernel_name, graph.dtype_of(scalar_in)); + add_dtype_suffix(kernel_name, scalar_dtype); graph.execute_nodes().emplace_back(new DispatchNode( graph, @@ -36,7 +40,7 @@ void scalar_tensor(ComputeGraph& graph, const std::vector& args) { // Inputs and Outputs {{out, vkapi::kWrite}}, // Shader params buffers - {graph.create_params_buffer(scalar_value)}, + {scalar_buffer}, // Push Constants {}, // Specialization Constants diff --git a/backends/vulkan/serialization/vulkan_graph_builder.py b/backends/vulkan/serialization/vulkan_graph_builder.py index 8703d391ce0..0c49880f927 100644 --- a/backends/vulkan/serialization/vulkan_graph_builder.py +++ b/backends/vulkan/serialization/vulkan_graph_builder.py @@ -474,10 +474,13 @@ def process_call_function_node(self, node) -> None: if not self.delegate_mapping_builder else self.delegate_mapping_builder.insert_delegate_mapping_entry(node) ) + operator_name = node.target.__name__ + if node.target == torch.ops.aten.scalar_tensor.default: + operator_name = "aten.scalar_tensor.default" self.chain.append( vk_graph_schema.OperatorCall( node_id=operator_node_id, # pyre-ignore[6]: this is going to be an int - name=node.target.__name__, + name=operator_name, args=operator_call_args, ), ) diff --git a/backends/vulkan/test/op_tests/cases.py b/backends/vulkan/test/op_tests/cases.py index ea98b1389a1..c99b7ad58b0 100644 --- a/backends/vulkan/test/op_tests/cases.py +++ b/backends/vulkan/test/op_tests/cases.py @@ -892,10 +892,12 @@ def get_scalar_tensor_inputs(): test_suite = VkTestSuite( [ (42.0,), + (42,), (3.14,), (2.72,), (0.0,), (-1.0,), + (-7,), (100.0,), ] ) diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index 51ebe2ca5f5..a44f5f1a253 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -32,6 +32,8 @@ from executorch.exir import EdgeCompileConfig, to_edge_transform_and_lower +from executorch.exir.dialects._ops import ops as exir_ops + from executorch.exir.lowered_backend_module import LoweredBackendModule from torch.export import Dim, export @@ -160,6 +162,44 @@ def test_dynamic_gelu(self): tolerance = 5e-6 if dtype == torch.float32 else 1e-3 self._run(edge, model, inputs, atol=tolerance, rtol=tolerance) + def test_scalar_tensor_values(self): + class WhereScalars(torch.nn.Module): + def __init__(self, positive, negative): + super().__init__() + self.positive = positive + self.negative = negative + + def forward(self, x): + return torch.where(x, self.positive, self.negative) + + inputs = [(torch.tensor([True, False, True, False]),)] + for positive, negative in ( + (3, -7.0), + (3.0, -7), + (16777217, -7), + (2**31 - 1, -(2**31)), + ): + with self.subTest(positive=positive, negative=negative): + model = WhereScalars(positive, negative) + fully_delegated = isinstance(positive, float) or isinstance( + negative, float + ) + edge = self._lower(model, inputs[0], fully_delegated=fully_delegated) + if not fully_delegated: + self.assertEqual( + [ + node.target + for node in edge.exported_program().graph.nodes + if node.op == "call_function" + and node.target != operator.getitem + ], + [ + torch.ops.higher_order.executorch_call_delegate, + exir_ops.edge.aten.where.self, + ], + ) + self._run(edge, model, inputs, atol=0, rtol=0) + def test_gelu_with_singleton_dimensions(self): for approximate in ("none", "tanh"): for shape in ((6, 1, 3), (2, 1, 3, 5)): @@ -172,6 +212,22 @@ def test_gelu_with_singleton_dimensions(self): edge = self._lower(model, (x,), storage=storage) self._run(edge, model, [(x,)], atol=5e-6, rtol=5e-6) + def test_scalar_tensor_dtypes(self): + class ScalarTensor(torch.nn.Module): + def __init__(self, dtype): + super().__init__() + self.dtype = dtype + + def forward(self, x): + return torch.scalar_tensor(2.5, dtype=self.dtype) + + inputs = [(torch.ones(1),)] + for dtype in (torch.float16, torch.float32, torch.int32): + with self.subTest(dtype=dtype): + model = ScalarTensor(dtype) + edge = self._lower(model, inputs[0]) + self._run(edge, model, inputs, atol=0, rtol=0) + def test_dynamic_logical_not(self): class LogicalNot(torch.nn.Module): def forward(self, x): @@ -205,6 +261,22 @@ def test_constant_bool_mask(self): ) self._run(edge, model, inputs, atol=0, rtol=0) + def test_scalar_types_before_conv_and_view(self): + class ScalarTypes(torch.nn.Module): + def __init__(self): + super().__init__() + self.conv = torch.nn.Conv2d(1, 2, 3, padding=1) + + def forward(self, x): + y = self.conv(torch.where(x > 0, 1.0, 0.5)) + return y.view(1, 2, x.shape[2], -1) + + torch.manual_seed(0) + model = ScalarTypes().eval() + inputs = [(torch.randn(1, 1, s, 5),) for s in (7, 2, 15, 7)] + edge = self._lower(model, inputs[0], ({2: Dim("s", min=2, max=16)},)) + self._run(edge, model, inputs) + def test_nan_scalars_fall_back(self): class NanScalar(torch.nn.Module): def __init__(self, kind): @@ -231,6 +303,36 @@ def forward(self, x): edge = self._lower(model, inputs[0], fully_delegated=False) self._run(edge, model, inputs, atol=0, rtol=0, equal_nan=True) + def test_integer_scalar_range_fallback(self): + class LargeScalar(torch.nn.Module): + def __init__(self, kind, value): + super().__init__() + self.kind = kind + self.value = value + + def forward(self, x): + if self.kind == "scalar_tensor": + return torch.scalar_tensor(self.value, dtype=torch.int64) + if self.kind == "full": + return torch.full(x.shape, self.value, dtype=torch.int64) + return torch.full_like(x, self.value, dtype=torch.int64) + + inputs = [(torch.zeros(2, 3),)] + for kind in ("scalar_tensor", "full", "full_like"): + for value in ( + 2**31 - 0.5, + -(2**31) - 0.5, + 2**31, + 2**40, + 2**63 - 1, + -(2**63), + ): + with self.subTest(kind=kind, value=value): + model = LargeScalar(kind, value) + 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_64_bit_arithmetic_without_downcasting(self): class Arithmetic(torch.nn.Module): def forward(self, x): diff --git a/backends/vulkan/test/test_vulkan_graph_builder.py b/backends/vulkan/test/test_vulkan_graph_builder.py index c180308e77a..419759785bc 100644 --- a/backends/vulkan/test/test_vulkan_graph_builder.py +++ b/backends/vulkan/test/test_vulkan_graph_builder.py @@ -41,6 +41,28 @@ def test_scalar_cache_preserves_types_and_signed_zero(self): self.assertEqual(builder.get_or_create_scalar_value(value), value_id) self.assertEqual(repr(builder.values[value_id].value), repr(serialized)) + def test_aten_scalar_tensor_keeps_namespace(self): + class Mask(torch.nn.Module): + def forward(self, x): + return torch.where(x, 0.0, -torch.inf) + + program = torch.export.export(Mask(), (torch.tensor([True, False]),)) + edge = to_edge(program) + program = apply_passes(edge.exported_program(), [SpecPropPass()]) + self.assertEqual( + sum( + node.target == torch.ops.aten.scalar_tensor.default + for node in program.graph.nodes + ), + 2, + ) + graph = VkGraphBuilder( + program, DelegateMappingBuilder(generated_identifiers=True) + ).build_graph() + names = [op.name for op in graph.chain] + self.assertEqual(names.count("aten.scalar_tensor.default"), 2) + self.assertNotIn("scalar_tensor.default", names) + class TestVkGraphBuilderInputIds(unittest.TestCase): """The serialized input list has to match the delegate call's arguments.