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):