Enable nvprims.transpose fusions for nvFuser (#86967) This PR allows transposes to be fused with other operations. If a fusion group is formed only from operations that just manipulate metadata in PyTorch (transpose, view, etc.) then this group is not sent to nvFuser. On top of that if we have converted to `nvprims` but then decided to not form a fusion group we modify the graph use `prim.impl_aten` attribute instead of calling `prim(*args, **kwargs)` that has a higher overhead. cc @kevinstephano @jjsjann123 Pull Request resolved: https://github.com/pytorch/pytorch/pull/86967 Approved by: https://github.com/jjsjann123, https://github.com/SherlockNoMad
diff --git a/test/test_prims.py b/test/test_prims.py index f1b8f89..6223a34 100644 --- a/test/test_prims.py +++ b/test/test_prims.py
@@ -877,6 +877,32 @@ @onlyCUDA @skipCUDAIfRocm + @dtypes(torch.float16, torch.float32) + def test_nvprims_view_partitioner(self, device, dtype): + # This test verifies that views that are not fused with other ops are + # correctly overriden to call aten implementation. + from torch.fx.experimental.proxy_tensor import make_fx + from torch._prims.context import TorchRefsNvfuserCapabilityMode + from torch._prims.nvfuser_executor import maybe_partition_graph + + make_arg = partial(make_tensor, device=device, dtype=dtype) + a = make_arg((4, 5)) + b = make_arg((5, 4)) + + def func(a, b): + aa = a.view(b.shape) + aa = aa.view(a.shape) + return aa.digamma() + + with TorchRefsNvfuserCapabilityMode(): + gm = make_fx(func)(a, b) + gm, _ = maybe_partition_graph(gm, False, False) + + out = gm(a, b) + self.assertEqual(out, func(a, b)) + + @onlyCUDA + @skipCUDAIfRocm @dtypes(torch.float32, torch.float16) def test_cpu_tensor(self, device, dtype): from torch.fx.experimental.proxy_tensor import make_fx
diff --git a/torch/_prims/__init__.py b/torch/_prims/__init__.py index 3248009..b54019e 100644 --- a/torch/_prims/__init__.py +++ b/torch/_prims/__init__.py
@@ -306,6 +306,7 @@ p.schema = schema p.prim_impl = _prim_impl p.prim_meta_impl = meta + p.impl_aten = impl_aten return _prim
diff --git a/torch/_prims/context.py b/torch/_prims/context.py index 2bcee06..203d73f 100644 --- a/torch/_prims/context.py +++ b/torch/_prims/context.py
@@ -254,10 +254,6 @@ class TorchRefsNvfuserCapabilityMode(TorchRefsMode): def __init__(self, *, skip_ops=()): aten_ops_to_skip = ( - "aten.transpose.int", - "aten.t.default", - "aten.unsqueeze.default", - "aten.permute.default", "aten._log_softmax.default", "aten._log_softmax_backward_data.default", "aten.expand.default",
diff --git a/torch/_prims/nvfuser_executor.py b/torch/_prims/nvfuser_executor.py index 01e566d..227e184 100644 --- a/torch/_prims/nvfuser_executor.py +++ b/torch/_prims/nvfuser_executor.py
@@ -30,7 +30,7 @@ DEFAULT_NVFUSER_PYTHON_CONFIG = MappingProxyType( { "use_python_fusion_cache": True, - "allow_single_op_fusion": True, + "allow_single_op_fusion": False, } ) @@ -268,6 +268,23 @@ ) +# A set of operators that are supported by nvFuser +# but should not form a fusion group solely on their own +_non_compute_ops = [ + "torch.ops." + str(getattr(torch.ops.nvprims, prim).default) + for prim in dir(torch.ops.nvprims) + if isinstance(getattr(torch.ops.nvprims, prim), torch._ops.OpOverloadPacket) + and getattr(torch.ops.nvprims, prim).return_type + == torch._prims_common.RETURN_TYPE.VIEW +] + +_allowed_single_node_partition_ops = [ + "torch.ops.nvprims.native_batch_norm.default", + "torch.ops.nvprims.var_mean.default", + "torch.ops.nvprims.var_mean.main", +] + + def _remove_empty_like_fill(gm: GraphModule): # Remove empty_like + fill nodes that prevent lowering to nvprims # This is a workaround for nonoptimal traces of C++ code `(1 - tensor)` @@ -325,7 +342,11 @@ # CapabilityBasedPartitioner modifies the graph in-place so we need to make a copy of the graph gm = deepcopy(gm) partitioner = CapabilityBasedPartitioner( - gm, supported_ops, allows_single_node_partition=allow_single_op_fusion + gm, + supported_ops, + allows_single_node_partition=allow_single_op_fusion, + non_compute_ops=_non_compute_ops, + allowed_single_node_partition_ops=_allowed_single_node_partition_ops, ) partitions = partitioner.propose_partitions() if len(partitions) == 0: @@ -350,6 +371,16 @@ NvfuserGraphModule(nvfuser_submodule, use_python_fusion_cache), ) + # Go through the graph and replace all the nodes that were converted to + # nvprims but won't be sent to nvFuser with a call to PyTorch's eager + # mode. This is necessary because torch.ops.* have higher overhead than + # calling the eager mode directly. + for node in partitioned_graph.graph.nodes: + if node.op == "call_function" and str(node.target).startswith("nvprims."): + if getattr(node.target, "impl_aten", None) is not None: + node.target = node.target.impl_aten + partitioned_graph.graph.eliminate_dead_code() + partitioned_graph.recompile() return partitioned_graph, any_unsupported else: return gm, any_unsupported
diff --git a/torch/_prims/nvfuser_prims.py b/torch/_prims/nvfuser_prims.py index d4132b3..f37a214 100644 --- a/torch/_prims/nvfuser_prims.py +++ b/torch/_prims/nvfuser_prims.py
@@ -538,6 +538,10 @@ p.return_type = torch._prims_common.RETURN_TYPE.NEW # type: ignore[attr-defined] +def _nvprims_view_impl_aten(a, original_shape, new_shape): + return a.reshape(new_shape) + + def register_view(): """This function is used to register the view function in torch.ops.view module.""" # View is implemented as a decomposition into prims.split_dim, @@ -568,7 +572,8 @@ for p in (prim_packet, prim): p.__doc__ = "Creates a tensor with the specified shape containing a copy of the data in a." p.impl_nvfuser = _nvfuser_impls["view"] - p.return_type = torch._prims_common.RETURN_TYPE.NEW # type: ignore[attr-defined] + p.return_type = torch._prims_common.RETURN_TYPE.VIEW # type: ignore[attr-defined] + p.impl_aten = _nvprims_view_impl_aten def register_nvprims(): @@ -594,3 +599,4 @@ p.__doc__ = main_prim.__doc__ p.impl_nvfuser = _nvfuser_impls[name] p.return_type = main_prim.return_type # type: ignore[attr-defined] + p.impl_aten = main_prim.impl_aten
diff --git a/torch/fx/passes/infra/partitioner.py b/torch/fx/passes/infra/partitioner.py index d582f98..5f5a808 100644 --- a/torch/fx/passes/infra/partitioner.py +++ b/torch/fx/passes/infra/partitioner.py
@@ -1,4 +1,4 @@ -from typing import Dict, List, Set, Iterable, Optional +from typing import Dict, List, Set, Iterable, Sequence, Optional from torch.fx.passes.utils.fuser_utils import fuse_by_partitions @@ -35,11 +35,19 @@ def __init__(self, graph_module: GraphModule, operator_support: OperatorSupportBase, - allows_single_node_partition: bool = False + allows_single_node_partition: bool = False, + non_compute_ops: Optional[Sequence[str]] = None, + allowed_single_node_partition_ops: Optional[Sequence[str]] = None, ) -> None: self.graph_module = graph_module self.operator_support = operator_support self.allows_single_node_partition = allows_single_node_partition + self.non_compute_ops = non_compute_ops if non_compute_ops is not None else [] + self.allowed_single_node_partition_ops = ( + allowed_single_node_partition_ops + if allowed_single_node_partition_ops is not None + else [] + ) def __is_node_supported(self, node: Node) -> bool: return ( @@ -169,7 +177,8 @@ # filter out single node partitions if not self.allows_single_node_partition: logger.debug("Filtering out single node partitions...") - non_compute_ops = {"torch.ops.aten.view", "_operator.getitem"} + default_non_compute_ops = {"torch.ops.aten.view", "_operator.getitem"} + non_compute_ops = default_non_compute_ops.union(set(self.non_compute_ops)) partitions_to_remove: List[int] = [] for id, partition in partitions_by_id.items(): compute_node_count = 0 @@ -177,6 +186,9 @@ if node.op == "call_function" and \ _get_qualified_name(node.target) not in non_compute_ops: # type: ignore[arg-type] compute_node_count += 1 + if node.op == "call_function" and \ + _get_qualified_name(node.target) in self.allowed_single_node_partition_ops: + compute_node_count += 1 if compute_node_count <= 1: partitions_to_remove.append(id) for id in partitions_to_remove: