From bcab9d05737d396af7166c882baf26985f78cc56 Mon Sep 17 00:00:00 2001 From: Mergen Nachin Date: Tue, 29 Sep 2026 11:02:04 -0400 Subject: [PATCH] [INITIAL] Recreate the Vulkan transformer and conformance stack with ghstack [ghstack-poisoned] --- .github/workflows/pull.yml | 1 + .../serialization/vulkan_graph_builder.py | 8 +----- backends/vulkan/test/targets.bzl | 11 ++++++++ .../vulkan/test/test_vulkan_graph_builder.py | 28 +++++++++++++++++++ 4 files changed, 41 insertions(+), 7 deletions(-) diff --git a/.github/workflows/pull.yml b/.github/workflows/pull.yml index 2f91799b9a4..a70a824f223 100644 --- a/.github/workflows/pull.yml +++ b/.github/workflows/pull.yml @@ -1682,6 +1682,7 @@ jobs: # route in the future. python -m unittest backends/vulkan/test/test_vulkan_delegate.py -k "*pt2e*" python -m unittest backends/vulkan/test/test_vulkan_delegate.py -k "*torchao*" + python -m unittest backends/vulkan/test/test_vulkan_graph_builder.py test-coreml-bc-macos: needs: [changed-files, run-decision] diff --git a/backends/vulkan/serialization/vulkan_graph_builder.py b/backends/vulkan/serialization/vulkan_graph_builder.py index 3de60966422..8703d391ce0 100644 --- a/backends/vulkan/serialization/vulkan_graph_builder.py +++ b/backends/vulkan/serialization/vulkan_graph_builder.py @@ -231,13 +231,7 @@ def create_null_value(self) -> int: return new_id def get_or_create_scalar_value(self, scalar: _ScalarType) -> int: - scalar_key = scalar - # Since Python considers 1 and True to be "equivalent" (as well as 0 and False) - # to distinguish entries in the dictionary, if scalar is bool then convert it - # to a string representation to use as a key for the dictionary - if isinstance(scalar, bool): - scalar_key = str(scalar) - + scalar_key = (type(scalar), repr(scalar)) if scalar_key in self.const_scalar_to_value_ids: return self.const_scalar_to_value_ids[scalar_key] diff --git a/backends/vulkan/test/targets.bzl b/backends/vulkan/test/targets.bzl index 77e74f9c8fc..5734e733195 100644 --- a/backends/vulkan/test/targets.bzl +++ b/backends/vulkan/test/targets.bzl @@ -29,6 +29,17 @@ def define_common_targets(is_fbcode = False): ], ) + python_unittest( + name = "test_vulkan_graph_builder", + srcs = ["test_vulkan_graph_builder.py"], + deps = [ + "//caffe2:torch", + "//executorch/backends/vulkan/serialization:lib", + "//executorch/backends/vulkan:vulkan_preprocess", + "//executorch/exir:lib", + ], + ) + python_unittest( name = "test_vulkan_passes", srcs = [ diff --git a/backends/vulkan/test/test_vulkan_graph_builder.py b/backends/vulkan/test/test_vulkan_graph_builder.py index 65afc3a2542..c180308e77a 100644 --- a/backends/vulkan/test/test_vulkan_graph_builder.py +++ b/backends/vulkan/test/test_vulkan_graph_builder.py @@ -8,12 +8,40 @@ import torch from executorch.backends.vulkan.serialization.vulkan_graph_builder import VkGraphBuilder +from executorch.backends.vulkan.serialization.vulkan_graph_schema import ( + Bool, + Double, + Int, +) from executorch.backends.vulkan.vulkan_preprocess import apply_passes from executorch.exir import to_edge from executorch.exir.backend.utils import DelegateMappingBuilder from executorch.exir.passes import SpecPropPass +class TestVkGraphBuilderScalarTensor(unittest.TestCase): + def test_scalar_cache_preserves_types_and_signed_zero(self): + program = torch.export.export(torch.nn.Identity(), (torch.ones(1),)) + builder = VkGraphBuilder( + program, DelegateMappingBuilder(generated_identifiers=True) + ) + scalars = (1.0, 1, True, 0.0, -0.0, 0, False) + expected = ( + Double(1.0), + Int(1), + Bool(True), + Double(0.0), + Double(-0.0), + Int(0), + Bool(False), + ) + ids = [builder.get_or_create_scalar_value(value) for value in scalars] + self.assertEqual(len(set(ids)), len(scalars)) + for value, value_id, serialized in zip(scalars, ids, expected): + self.assertEqual(builder.get_or_create_scalar_value(value), value_id) + self.assertEqual(repr(builder.values[value_id].value), repr(serialized)) + + class TestVkGraphBuilderInputIds(unittest.TestCase): """The serialized input list has to match the delegate call's arguments.