Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
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 @@ -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):
Expand Down
Loading