From e0bb5f4f6c3c7eb812ab8515a12c1a91a8c36a90 Mon Sep 17 00:00:00 2001 From: Mergen Nachin Date: Mon, 28 Sep 2026 18:20:37 -0400 Subject: [PATCH] [Vulkan] Honor downcast_64_bit=False in the partitioner With downcast_64_bit disabled, the graph builder refuses int64 and float64 tensors, but the partitioner still claimed those nodes, so lowering failed late or the delegate aborted at load. The partitioner now rejects nodes that need 64-bit dtypes in that mode, and keeps a fusable pattern entirely on CPU when it touches a 64-bit non-constant tensor rather than splitting it. Constant parameters are exempt because quantization parameters are folded away during fusion. Authored with OpenAI Codex; split planned with Claude Code. --- .../vulkan/partitioner/vulkan_partitioner.py | 42 ++++++++- backends/vulkan/test/targets.bzl | 2 + backends/vulkan/test/test_vulkan_delegate.py | 42 +++++++++ backends/vulkan/test/test_vulkan_dynamic.py | 91 +++++++++++++++++++ 4 files changed, 176 insertions(+), 1 deletion(-) diff --git a/backends/vulkan/partitioner/vulkan_partitioner.py b/backends/vulkan/partitioner/vulkan_partitioner.py index 08c9f3bbd65..0920c8bd02b 100644 --- a/backends/vulkan/partitioner/vulkan_partitioner.py +++ b/backends/vulkan/partitioner/vulkan_partitioner.py @@ -43,6 +43,7 @@ from torch.fx.passes.infra.partitioner import CapabilityBasedPartitioner from torch.fx.passes.operator_support import OperatorSupportBase +from torch.utils._pytree import tree_leaves # pyre-ignore ops_not_to_decompose = [ @@ -67,12 +68,15 @@ def __init__( fusable_subgraphs: Optional[List[PatternMatch]] = None, nn_module_blocklist: Optional[Set[str]] = None, nn_module_allowlist: Optional[Set[str]] = None, + downcast_64_bit: bool = True, + constant_nodes: Optional[Set[torch.fx.Node]] = None, ) -> None: super().__init__() self.texture_limits: utils.ImageExtents = texture_limits self.buffer_limit = buffer_limit self.require_dynamic_shapes = require_dynamic_shape self.skip_bool_tensors = skip_bool_tensors + self.downcast_64_bit = downcast_64_bit self.operator_blocklist: Set[OpKey] = ( operator_blocklist if operator_blocklist is not None else set() ) @@ -82,8 +86,22 @@ def __init__( ) # Create a set of all nodes that are part of fusable subgraphs for quick lookup self.fusable_nodes: Set[torch.fx.Node] = set() + self.unsupported_fusable_nodes: Set[torch.fx.Node] = set() for match in self.fusable_subgraphs: - self.fusable_nodes.update(match.all_nodes) + nodes = { + node for node in match.all_nodes if isinstance(node, torch.fx.Node) + } + self.fusable_nodes.update(nodes) + inputs = {arg for node in nodes for arg in node.all_input_nodes} + if not downcast_64_bit and any( + isinstance(value, torch.Tensor) + and value.dtype in (torch.int64, torch.float64) + for node in (nodes | inputs) - (constant_nodes or set()) + for value in tree_leaves(node.meta.get("val")) + ): + # Keep the whole pattern outside Vulkan instead of splitting a fusion. + # Constant quantization parameters may disappear during fusion. + self.unsupported_fusable_nodes.update(nodes) self.nn_module_blocklist = nn_module_blocklist self.nn_module_allowlist = nn_module_allowlist @@ -207,6 +225,10 @@ def _is_node_supported(self, node: torch.fx.Node) -> bool: return self._node_support[node] def _check_node_support(self, node: torch.fx.Node) -> bool: # noqa: C901 + if node in self.unsupported_fusable_nodes: + self.log_skip(node, "fusable pattern requires 64-bit tensor downcasting") + return False + if any( isinstance(arg.meta.get("val"), (torch.SymFloat, torch.SymBool)) for arg in [node, *node.all_input_nodes] @@ -245,6 +267,17 @@ def _check_node_support(self, node: torch.fx.Node) -> bool: # noqa: C901 if node in self.fusable_nodes: return True + if not self.downcast_64_bit: + native_dtypes = utils.DtypeSetList( + utils.ALL_T - {torch.int64, torch.float64} + ) + dtype_valid, dtype_reason = utils.check_node_dtypes( + node, native_dtypes, native_dtypes + ) + if not dtype_valid: + self.log_skip(node, f"{dtype_reason} with downcast_64_bit disabled") + return False + target = node.target if ( node.target == torch.ops.higher_order.auto_functionalized @@ -433,6 +466,13 @@ def partition(self, exported_program: ExportedProgram) -> PartitionResult: fusable_subgraphs=fusable_subgraphs, nn_module_blocklist=self.nn_module_blocklist, nn_module_allowlist=self.nn_module_allowlist, + downcast_64_bit=self.options.get("downcast_64_bit", True), + constant_nodes={ + node + for node in exported_program.graph.nodes + if utils.is_param_node(exported_program, node) + and not utils.is_mutable_buffer_node(node, exported_program) + }, ), allows_single_node_partition=True, ) diff --git a/backends/vulkan/test/targets.bzl b/backends/vulkan/test/targets.bzl index d18e508341d..0051a74a530 100644 --- a/backends/vulkan/test/targets.bzl +++ b/backends/vulkan/test/targets.bzl @@ -22,10 +22,12 @@ def define_common_targets(is_fbcode = False): "//executorch/backends/transforms:convert_dtype_pass", "//executorch/backends/vulkan:vulkan_preprocess", "//executorch/backends/vulkan/partitioner:vulkan_partitioner", + "//executorch/backends/vulkan/quantizer:vulkan_quantizer", "//executorch/exir:lib", "//executorch/extension/pybindings:portable_lib", # @manual "//executorch/extension/pytree:pylib", "//executorch/kernels/portable:custom_ops_generated_lib", + "//pytorch/ao:torchao", # @manual ], ) diff --git a/backends/vulkan/test/test_vulkan_delegate.py b/backends/vulkan/test/test_vulkan_delegate.py index 674c3c7837f..ca6212a5982 100644 --- a/backends/vulkan/test/test_vulkan_delegate.py +++ b/backends/vulkan/test/test_vulkan_delegate.py @@ -8,6 +8,7 @@ import ctypes import functools +import operator import unittest from typing import Tuple @@ -16,6 +17,10 @@ import torch.nn.functional as F from executorch.backends.transforms.convert_dtype_pass import I64toI32 from executorch.backends.vulkan.partitioner.vulkan_partitioner import VulkanPartitioner +from executorch.backends.vulkan.quantizer.vulkan_quantizer import ( + get_symmetric_quantization_config as get_vulkan_quantization_config, + VulkanQuantizer, +) from executorch.backends.vulkan.vulkan_preprocess import VulkanBackend from executorch.backends.xnnpack.quantizer.xnnpack_quantizer import ( get_symmetric_quantization_config, @@ -2754,6 +2759,43 @@ def apply_quantization(self): quantized_linear_module_gemm, sample_inputs_gemm, atol=1e-2, rtol=1e-2 ) + def test_vulkan_backend_pt2e_quantized_linear_without_downcasting(self): + torch.manual_seed(0) + sample_inputs = (torch.randn(4, 64),) + quantizer = VulkanQuantizer().set_global(get_vulkan_quantization_config()) + model = prepare_pt2e( + export(torch.nn.Linear(64, 32).eval(), sample_inputs, strict=True).module(), + quantizer, + ) + model(*sample_inputs) + model = convert_pt2e(model) + self.assertTrue(any(buffer.dtype == torch.int64 for buffer in model.buffers())) + + for downcast in (False, True): + with self.subTest(downcast=downcast): + edge = lower_module( + model, + sample_inputs, + compile_options={"downcast_64_bit": downcast}, + ) + 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], + ) + program_buffer = edge.to_executorch().buffer + module = _load_for_executorch_from_buffer(program_buffer) + self.assert_outputs_equal( + module.run_method("forward", sample_inputs), + model(*sample_inputs), + atol=1e-5, + rtol=1e-5, + ) + @disable_test("Cannot run on swiftshader due to no integer dot product support") def test_vulkan_backend_xnnpack_pt2e_quantized_linear_sequence(self): """ diff --git a/backends/vulkan/test/test_vulkan_dynamic.py b/backends/vulkan/test/test_vulkan_dynamic.py index 822dccce989..51ebe2ca5f5 100644 --- a/backends/vulkan/test/test_vulkan_dynamic.py +++ b/backends/vulkan/test/test_vulkan_dynamic.py @@ -231,6 +231,97 @@ 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_64_bit_arithmetic_without_downcasting(self): + class Arithmetic(torch.nn.Module): + def forward(self, x): + return x + x, x.to(torch.float32) + 1 + + for dtype in (torch.int64, torch.float64): + with self.subTest(dtype=dtype): + model = Arithmetic() + inputs = [ + (torch.arange(3 * s, dtype=dtype).reshape(3, s),) for s in (7, 2) + ] + edge = self._lower( + model, + inputs[0], + ({1: Dim("s", min=2, max=16)},), + fully_delegated=False, + downcast_64_bit=False, + ) + graphs = _vulkan_graphs(edge) + self.assertTrue(graphs) + for graph in graphs: + for value in graph.values: + if isinstance(value.value, VkTensor): + self.assertNotIn( + value.value.datatype, + (VkDataType.INT64, VkDataType.FLOAT64), + ) + self._run(edge, model, inputs, atol=0, rtol=0) + + def test_64_bit_fusion_inputs_without_downcasting(self): + class SelectScalar(torch.nn.Module): + def __init__(self, narrow): + super().__init__() + self.narrow = narrow + + def forward(self, x, pos): + value = pos[0].item() + if self.narrow: + torch._check(value >= 0) + torch._check(value <= 6) + return x.narrow(1, value, 2) + 1 + return x * value + + inputs = [(torch.randn(2, 8), torch.tensor([pos])) for pos in (3, 5, 0)] + for narrow in (False, True): + for downcast in (False, True): + with self.subTest(narrow=narrow, downcast=downcast): + model = SelectScalar(narrow) + edge = self._lower( + model, + inputs[0], + fully_delegated=False, + downcast_64_bit=downcast, + ) + for graph in _vulkan_graphs(edge): + for value in graph.values: + if isinstance(value.value, VkTensor): + self.assertNotIn( + value.value.datatype, + (VkDataType.INT64, VkDataType.FLOAT64), + ) + self._run(edge, model, inputs, atol=0, rtol=0) + + def test_quantized_embedding_without_downcasting(self): + from torchao.quantization.granularity import PerGroup + from torchao.quantization.quant_api import IntxWeightOnlyConfig, quantize_ + from torchao.utils import unwrap_tensor_subclass + + torch.manual_seed(0) + model = torch.nn.Sequential(torch.nn.Embedding(64, 128)).eval() + quantize_( + model, + IntxWeightOnlyConfig(weight_dtype=torch.int4, granularity=PerGroup(32)), + filter_fn=lambda module, fqn: isinstance(module, torch.nn.Embedding), + ) + unwrap_tensor_subclass(model) + inputs = [(torch.tensor(indices),) for indices in ([0, 5, 63, 7], [3, 3, 1, 0])] + for downcast in (False, True): + with self.subTest(downcast=downcast): + if downcast and USING_SWIFTSHADER: + self.skipTest("Quantized embedding requires 8-bit storage buffers") + edge = self._lower( + model, + inputs[0], + fully_delegated=downcast, + downcast_64_bit=downcast, + ) + self.assertEqual(bool(_vulkan_graphs(edge)), downcast) + if downcast: + self._run(edge, model, inputs) + def test_dynamic_scalar_values_fall_back(self): class DynamicScalars(torch.nn.Module): def forward(self, x):