Skip to content
42 changes: 41 additions & 1 deletion backends/vulkan/partitioner/vulkan_partitioner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = [
Expand All @@ -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()
)
Expand All @@ -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
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
)
Expand Down
2 changes: 2 additions & 0 deletions backends/vulkan/test/targets.bzl
Original file line number Diff line number Diff line change
Expand Up @@ -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
],
)

Expand Down
42 changes: 42 additions & 0 deletions backends/vulkan/test/test_vulkan_delegate.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@

import ctypes
import functools
import operator
import unittest
from typing import Tuple

Expand All @@ -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,
Expand Down Expand Up @@ -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):
"""
Expand Down
91 changes: 91 additions & 0 deletions backends/vulkan/test/test_vulkan_dynamic.py
Original file line number Diff line number Diff line change
Expand Up @@ -236,6 +236,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):
Expand Down
Loading