From ced541a633403320c1e76edcace617e56863b77e Mon Sep 17 00:00:00 2001 From: Mergen Nachin Date: Mon, 28 Sep 2026 18:06:23 -0400 Subject: [PATCH] [Vulkan] Key serialized scalar constants by type and value VkGraphBuilder deduplicated scalar constants in a dict keyed on the Python value, so 1.0, 1 and True (and 0.0, -0.0 and 0) collapsed onto whichever entry was created first and were serialized with the wrong type or sign. Keying on (type, repr) keeps distinct scalars separate while still sharing identical ones. test_vulkan_graph_builder.py is also added to Buck and the Vulkan pull job, where it was not running. Authored with OpenAI Codex; split planned with Claude Code. --- .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.