From 2d0e6fc72c5f78e79a367d5c09035a4a5f196322 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mateusz=20S=C5=82uszniak?= Date: Sun, 30 Aug 2026 14:54:51 +0200 Subject: [PATCH] [ET-VK] Lower eligible conv1d as conv2d over a singleton height dim conv1d.glsl computes one output element per invocation with no tiling and no reuse between invocations, and its work grid packs texels along the batch dim, wasting three of every four lanes at batch 1. On an Adreno 840 it reaches about 30 GFLOP/s where the tiled linear reaches 900 GFLOP/s in the same graph. The conv2d im2col + GEMM path already solves this, so rewrite 1-D convolutions that are certain to reach it into a 2-D convolution over a singleton height dim. Only convs with groups 1, unit dilation, kernel > 1, batch 1 and out_channels >= kIm2colMinCOut are rewritten, so nothing moves onto a shader that has not been compared against conv1d. The Whisper-tiny encoder goes from 153.2 ms to 88.8 ms, from 0.84x XNNPACK to 1.44x. Fixes #22329 --- backends/vulkan/_passes/__init__.py | 2 + backends/vulkan/_passes/conv1d_as_conv2d.py | 157 +++++++++++++++++++ backends/vulkan/_passes/targets.bzl | 15 ++ backends/vulkan/test/test_vulkan_delegate.py | 27 ++++ backends/vulkan/vulkan_preprocess.py | 2 + 5 files changed, 203 insertions(+) create mode 100644 backends/vulkan/_passes/conv1d_as_conv2d.py diff --git a/backends/vulkan/_passes/__init__.py b/backends/vulkan/_passes/__init__.py index 1afaf48dde7..cdf44d87a63 100644 --- a/backends/vulkan/_passes/__init__.py +++ b/backends/vulkan/_passes/__init__.py @@ -6,6 +6,7 @@ # pyre-strict +from executorch.backends.vulkan._passes.conv1d_as_conv2d import Conv1dAsConv2dPass from executorch.backends.vulkan._passes.fold_qdq import FoldQDQPass from executorch.backends.vulkan._passes.fuse_patterns import FusePatternsPass from executorch.backends.vulkan._passes.fuse_quantized_ops import ( @@ -28,6 +29,7 @@ from executorch.backends.vulkan._passes.tag_memory_meta_pass import TagMemoryMetaPass __all__ = [ + "Conv1dAsConv2dPass", "FoldQDQPass", "FusePatternsPass", "FuseQuantizedOpsTransform", diff --git a/backends/vulkan/_passes/conv1d_as_conv2d.py b/backends/vulkan/_passes/conv1d_as_conv2d.py new file mode 100644 index 00000000000..a67a0248408 --- /dev/null +++ b/backends/vulkan/_passes/conv1d_as_conv2d.py @@ -0,0 +1,157 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +# pyre-strict + +from typing import List, Optional + +import torch + +from executorch.backends.transforms.utils import ( + get_param_tensor, + is_param_node, + set_param_tensor, +) +from executorch.exir.dialects._ops import ops as exir_ops +from executorch.exir.pass_base import ExportPass, PassResult + +from torch.export import ExportedProgram + +# Mirrors kIm2colMinCOut in Convolution.cpp. At or above this output channel +# count, should_use_conv2d_im2col() selects the im2col + GEMM path on every +# vendor, so a rewritten conv is guaranteed to land there rather than on the +# direct conv2d shader, which has not been compared against conv1d. +_IM2COL_MIN_C_OUT = 128 + + +def _reshape_fake(val: torch.Tensor, sizes: List[int]) -> torch.Tensor: + fake_mode = getattr(val, "fake_mode", None) + if fake_mode is None: + return val.reshape(sizes) + with fake_mode: + return val.reshape(sizes) + + +class Conv1dAsConv2dPass(ExportPass): + """ + Rewrite eligible 1-D convolutions into a 2-D convolution over a singleton + height dimension. + + conv1d.glsl computes one output element per invocation with no tiling and no + reuse between invocations, and its global work grid packs texels along the + batch dim, wasting three of every four lanes at batch 1. It measures around + 30x slower than the conv2d im2col + GEMM path for the same MAC count. + Expressing the convolution as 2-D lets the existing conv2d machinery + (im2col selection, memory layout tagging, weight prepacking) handle it with + no runtime changes. + + Only convolutions that are certain to reach the im2col path are rewritten, + so nothing silently moves onto a shader that has not been compared against + conv1d. The weight is reshaped in place rather than through a view node + because the conv2d implementation reads it as a constant TensorRef. + """ + + def __init__(self) -> None: + super().__init__() + self._exported_program: Optional[ExportedProgram] = None + + def _eligible(self, node: torch.fx.Node) -> bool: + assert self._exported_program is not None + in_node, weight_node = node.args[0], node.args[1] + if not isinstance(in_node, torch.fx.Node): + return False + if not isinstance(weight_node, torch.fx.Node): + return False + in_shape = in_node.meta["val"].shape + w_shape = weight_node.meta["val"].shape + if len(in_shape) != 3 or len(w_shape) != 3: + return False + # The weight has to be a constant so it can be reshaped in place. + if not is_param_node(self._exported_program, weight_node): + return False + transposed, groups = node.args[6], node.args[8] + if transposed or groups != 1: + return False + # conv2d does not support batched input at all, so a rewrite there + # would turn a working conv1d into a throw. + if in_shape[0] != 1: + return False + # Depthwise and pointwise 1-D convs keep their own shaders. + if w_shape[0] < _IM2COL_MIN_C_OUT or w_shape[2] == 1: + return False + # im2col requires unit dilation. + return list(node.args[5]) == [1] + + def _rewrite(self, graph: torch.fx.Graph, node: torch.fx.Node) -> None: + assert self._exported_program is not None + view = exir_ops.edge.aten.view_copy.default + in_node, weight_node = node.args[0], node.args[1] + + in_shape = list(in_node.meta["val"].shape) + w_shape = list(weight_node.meta["val"].shape) + out_shape = list(node.meta["val"].shape) + in_4d_shape = [in_shape[0], in_shape[1], 1, in_shape[2]] + w_4d_shape = [w_shape[0], w_shape[1], 1, w_shape[2]] + out_4d_shape = [out_shape[0], out_shape[1], 1, out_shape[2]] + + # The weight tensor is contiguous, so inserting a singleton dim is a + # pure metadata change. A weight shared by several convs is only + # reshaped once; the rank check in _eligible() skips it afterwards. + weight = get_param_tensor(self._exported_program, weight_node) + assert weight is not None + set_param_tensor( + self._exported_program, weight_node, weight.reshape(w_4d_shape) + ) + weight_node.meta["val"] = _reshape_fake( + weight_node.meta["val"], w_4d_shape + ) + + with graph.inserting_before(node): + in_4d = graph.call_function(view, args=(in_node, in_4d_shape)) + in_4d.meta = dict(in_node.meta) + in_4d.meta["val"] = _reshape_fake(in_node.meta["val"], in_4d_shape) + + def pad(arg: List[int], lead: int) -> List[int]: + return [lead, arg[0]] + + node.args = ( + in_4d, + weight_node, + node.args[2], + pad(node.args[3], 1), # stride + pad(node.args[4], 0), # padding + pad(node.args[5], 1), # dilation + node.args[6], + pad(node.args[7], 0), # output_padding + node.args[8], + ) + + with graph.inserting_after(node): + out_3d = graph.call_function(view, args=(node, out_shape)) + out_3d.meta = dict(node.meta) + node.replace_all_uses_with(out_3d) + out_3d.args = (node, out_shape) + node.meta["val"] = _reshape_fake(node.meta["val"], out_4d_shape) + + def call(self, graph_module: torch.fx.GraphModule) -> PassResult: + assert self._exported_program is not None + + modified = False + for node in list(graph_module.graph.nodes): + if node.op != "call_function": + continue + if node.target != exir_ops.edge.aten.convolution.default: + continue + if not self._eligible(node): + continue + self._rewrite(graph_module.graph, node) + modified = True + + if not modified: + return PassResult(graph_module, False) + + graph_module.recompile() + return PassResult(graph_module, True) diff --git a/backends/vulkan/_passes/targets.bzl b/backends/vulkan/_passes/targets.bzl index 89565a8f5c8..4d2f19bba17 100644 --- a/backends/vulkan/_passes/targets.bzl +++ b/backends/vulkan/_passes/targets.bzl @@ -93,6 +93,20 @@ def define_common_targets(is_fbcode = False): ], ) + runtime.python_library( + name = "conv1d_as_conv2d", + srcs = ["conv1d_as_conv2d.py"], + visibility = [ + "//executorch/backends/...", + ], + deps = [ + "//caffe2:torch", + "//executorch/backends/transforms:utils", + "//executorch/exir:pass_base", + "//executorch/exir/dialects:lib", + ], + ) + runtime.python_library( name = "fold_qdq", srcs = ["fold_qdq.py"], @@ -144,6 +158,7 @@ def define_common_targets(is_fbcode = False): "//executorch/examples/...", ], deps = [ + ":conv1d_as_conv2d", ":fold_qdq", ":fuse_patterns", ":fuse_quantized_ops", diff --git a/backends/vulkan/test/test_vulkan_delegate.py b/backends/vulkan/test/test_vulkan_delegate.py index c6915d37684..078f15904a7 100644 --- a/backends/vulkan/test/test_vulkan_delegate.py +++ b/backends/vulkan/test/test_vulkan_delegate.py @@ -1065,6 +1065,33 @@ def forward(self, x): sample_inputs, ) + def test_vulkan_backend_conv1d_as_conv2d(self): + # Eligible for Conv1dAsConv2dPass: groups 1, unit dilation, kernel > 1, + # batch 1 and out_channels >= the im2col threshold, so this lowers via + # the conv2d im2col + GEMM path rather than conv1d.glsl. + class Conv1dModule(torch.nn.Module): + def __init__(self): + super().__init__() + self.conv = torch.nn.Conv1d( + in_channels=32, + out_channels=128, + kernel_size=3, + stride=2, + padding=1, + bias=True, + ) + + def forward(self, x): + return self.conv(x) + + conv1d_module = Conv1dModule() + sample_inputs = (torch.randn(size=(1, 32, 64), dtype=torch.float32),) + + self.lower_module_and_test_output( + conv1d_module, + sample_inputs, + ) + @disable_test("layer norm compute shader not working with swiftshader") def test_vulkan_backend_native_layer_norm(self): class NativeLayerNormModule(torch.nn.Module): diff --git a/backends/vulkan/vulkan_preprocess.py b/backends/vulkan/vulkan_preprocess.py index f7d6955ce26..dcea5934b53 100644 --- a/backends/vulkan/vulkan_preprocess.py +++ b/backends/vulkan/vulkan_preprocess.py @@ -17,6 +17,7 @@ ViewCopyToSqueezeUnsqueezePass, ) from executorch.backends.vulkan._passes import ( + Conv1dAsConv2dPass, FoldQDQPass, FuseQuantizedOpsTransform, insert_prepack_nodes, @@ -190,6 +191,7 @@ def preprocess( # noqa: C901 AddmmToLinearTransform(), InsertDtypePromotionPass(), FusePatternsPass(), + Conv1dAsConv2dPass(), FuseClampPass(), RemoveRedundantOpsTransform(), FuseQuantizedOpsTransform(),