Skip to content

Commit 8ea353d

Browse files
Arm backend: Fix ConvTranspose2d batch norm fusion (pytorch#22517)
Pass `transpose=True` to PyTorch's batch norm fusion helper for transposed convolutions. Fuse grouped transposed convolutions one group at a time and restore the grouped parameter layout before decomposition. Add tests for unequal input and output channel counts, grouped fusion with and without bias, and non-affine batch norm. Authored with Codex. Change-Id: I7fbd523fa1627ce5d198aea5adbd083c91c0635f Signed-off-by: Yufeng Shi <yufeng.shi@arm.com>
1 parent 2c1da32 commit 8ea353d

4 files changed

Lines changed: 220 additions & 13 deletions

File tree

‎backends/arm/_passes/fuse_batch_norm2d_pass.py‎

Lines changed: 107 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,9 @@
1212
create_node,
1313
get_first_fake_tensor,
1414
)
15+
from executorch.backends.arm._passes.decompose_grouped_conv_pass import (
16+
DecomposeGroupedConvPass,
17+
)
1518
from executorch.backends.arm.common.debug import get_node_debug_info
1619
from executorch.backends.transforms.utils import (
1720
create_constant_placeholder,
@@ -27,11 +30,13 @@
2730

2831

2932
class FuseBatchNorm2dPass(ArmPass):
30-
"""Fuses the pattern convolution -> batchnorm by updating the weights and
31-
bias of the convolution and removing the batchnorm.
33+
"""Fuse convolution followed by BatchNorm.
34+
35+
Update the convolution weights and bias and remove the BatchNorm operation.
36+
3237
"""
3338

34-
_passes_required_after: Set[Type[ExportPass]] = set()
39+
_passes_required_after: Set[Type[ExportPass]] = {DecomposeGroupedConvPass}
3540

3641
def __init__(self, exported_program: ExportedProgram, *args, **kwargs):
3742
super().__init__(*args, **kwargs)
@@ -45,6 +50,79 @@ def get_bias_name(self, weight_node: Node, bias_node: Node | None) -> str:
4550
else:
4651
return weight_node.name + "_bias_fused_bn"
4752

53+
@staticmethod
54+
def _fuse_grouped_transposed_conv_bn_weights(
55+
conv_weight: torch.Tensor,
56+
conv_bias: torch.Tensor | None,
57+
bn_mean: torch.Tensor,
58+
bn_var: torch.Tensor,
59+
bn_epsilon: float,
60+
bn_weight: torch.Tensor | None,
61+
bn_bias: torch.Tensor | None,
62+
groups: int,
63+
) -> tuple[torch.Tensor, torch.Tensor]:
64+
"""Fuse BatchNorm into grouped transposed-convolution parameters.
65+
66+
This helper runs before ``DecomposeGroupedConvPass`` and transforms::
67+
68+
grouped ConvTranspose -> BatchNorm
69+
70+
into a grouped ConvTranspose with fused weights and bias. A transposed
71+
convolution weight has layout ``[Cin, Cout/groups, ...]``. The weight
72+
is split on its input-channel dimension, while the bias and BatchNorm
73+
parameters are split on their output-channel dimension. Each group is
74+
fused independently before the original grouped layout is restored.
75+
76+
Args:
77+
conv_weight (torch.Tensor): Grouped transposed-convolution weight.
78+
conv_bias (torch.Tensor | None): Convolution bias.
79+
bn_mean (torch.Tensor): BatchNorm running mean.
80+
bn_var (torch.Tensor): BatchNorm running variance.
81+
bn_epsilon (float): BatchNorm numerical-stability constant.
82+
bn_weight (torch.Tensor | None): BatchNorm weight.
83+
bn_bias (torch.Tensor | None): BatchNorm bias.
84+
groups (int): Number of convolution groups.
85+
86+
Returns:
87+
tuple[torch.Tensor, torch.Tensor]: Fused weight and bias in the
88+
original grouped layout.
89+
90+
Raises:
91+
RuntimeError: If the grouped channel dimensions are inconsistent.
92+
93+
"""
94+
if conv_weight.size(0) % groups != 0 or bn_mean.numel() % groups != 0:
95+
raise RuntimeError("Grouped transposed convolution has invalid channels")
96+
97+
input_channels_per_group = conv_weight.size(0) // groups
98+
output_channels_per_group = bn_mean.numel() // groups
99+
if conv_weight.size(1) != output_channels_per_group:
100+
raise RuntimeError("BatchNorm channels do not match convolution output")
101+
102+
fused_weights: list[torch.Tensor] = []
103+
fused_biases: list[torch.Tensor] = []
104+
for group in range(groups):
105+
input_start = group * input_channels_per_group
106+
input_end = input_start + input_channels_per_group
107+
output_start = group * output_channels_per_group
108+
output_end = output_start + output_channels_per_group
109+
output_slice = slice(output_start, output_end)
110+
111+
fused_weight, fused_bias = fuse_conv_bn_weights(
112+
conv_weight[input_start:input_end],
113+
conv_bias[output_slice] if conv_bias is not None else None,
114+
bn_mean[output_slice],
115+
bn_var[output_slice],
116+
bn_epsilon,
117+
bn_weight[output_slice] if bn_weight is not None else None,
118+
bn_bias[output_slice] if bn_bias is not None else None,
119+
transpose=True,
120+
)
121+
fused_weights.append(fused_weight)
122+
fused_biases.append(fused_bias)
123+
124+
return torch.cat(fused_weights, dim=0), torch.cat(fused_biases, dim=0)
125+
48126
def call(self, graph_module: torch.fx.GraphModule) -> PassResult: # noqa: C901
49127
modified = False
50128
constant_placeholders_to_delete = set()
@@ -176,15 +254,32 @@ def call(self, graph_module: torch.fx.GraphModule) -> PassResult: # noqa: C901
176254
)
177255

178256
# Fuse bn weights/bias with input weights/bias
179-
fused_weight, fused_bias = fuse_conv_bn_weights(
180-
input_weight_tensor,
181-
input_bias_tensor,
182-
bn_mean_tensor,
183-
bn_var_tensor,
184-
epsilon,
185-
bn_weight_tensor,
186-
bn_bias_tensor,
187-
)
257+
transposed = bool(input_node.args[6])
258+
groups = int(input_node.args[8])
259+
if transposed and groups > 1:
260+
fused_weight, fused_bias = (
261+
self._fuse_grouped_transposed_conv_bn_weights(
262+
input_weight_tensor,
263+
input_bias_tensor,
264+
bn_mean_tensor,
265+
bn_var_tensor,
266+
epsilon,
267+
bn_weight_tensor,
268+
bn_bias_tensor,
269+
groups,
270+
)
271+
)
272+
else:
273+
fused_weight, fused_bias = fuse_conv_bn_weights(
274+
input_weight_tensor,
275+
input_bias_tensor,
276+
bn_mean_tensor,
277+
bn_var_tensor,
278+
epsilon,
279+
bn_weight_tensor,
280+
bn_bias_tensor,
281+
transpose=transposed,
282+
)
188283

189284
# Create fused weights and bias to conv and replace conv args
190285
with graph_module.graph.inserting_before(input_weight_node):

‎backends/arm/test/ops/test_batch_norm.py‎

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
Input = Tuple[torch.Tensor]
2323
ATEN_BATCH_NORM = "torch.ops.aten.batch_norm.default"
2424
ATEN_CONV2D = "torch.ops.aten.conv2d.default"
25+
ATEN_CONV_TRANSPOSE2D = "torch.ops.aten.conv_transpose2d.input"
2526

2627

2728
@dataclass(frozen=True)
@@ -110,6 +111,30 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
110111
return self.batch_norm(self.conv2d(x))
111112

112113

114+
class BatchNorm2dConvTranspose(torch.nn.Module):
115+
aten_ops = [ATEN_CONV_TRANSPOSE2D, ATEN_BATCH_NORM]
116+
117+
def __init__(self, groups: int) -> None:
118+
super().__init__()
119+
self.conv_transpose2d = torch.nn.ConvTranspose2d(
120+
in_channels=4,
121+
out_channels=6,
122+
kernel_size=3,
123+
padding=1,
124+
groups=groups,
125+
)
126+
self.batch_norm = _make_batch_norm(
127+
6,
128+
affine=True,
129+
weight=torch.rand(6),
130+
bias=torch.rand(6),
131+
track_running_stats=True,
132+
)
133+
134+
def forward(self, x: torch.Tensor) -> torch.Tensor:
135+
return self.batch_norm(self.conv_transpose2d(x))
136+
137+
113138
class BatchNorm2dNoStats(torch.nn.Module):
114139
def __init__(
115140
self,
@@ -205,6 +230,27 @@ def test_native_batch_norm_legit_no_training_tosa_FP_conv_fuses_before_decompose
205230
pipeline.run()
206231

207232

233+
@common.parametrize("groups", {"groups=1": 1, "groups=2": 2})
234+
def test_conv_transpose_batch_norm_fuses_before_decompose_tosa_FP(
235+
groups: int,
236+
) -> None:
237+
model = BatchNorm2dConvTranspose(groups)
238+
pipeline = TosaPipelineFP[Input](
239+
model,
240+
(torch.rand(1, 4, 5, 6),),
241+
aten_op=model.aten_ops,
242+
)
243+
pipeline.count_tosa_ops(
244+
{
245+
"TRANSPOSE_CONV2D": groups,
246+
"CONCAT": int(groups > 1),
247+
"RSQRT": 0,
248+
"SUB": 0,
249+
}
250+
)
251+
pipeline.run()
252+
253+
208254
@common.parametrize("case", batch_norm_cases)
209255
def test_native_batch_norm_legit_no_training_tosa_INT_conv(case: BatchNormCase) -> None:
210256
test_data, model_params = case.make_input_and_parameters()
@@ -254,6 +300,17 @@ def test_native_batch_norm_legit_no_training_vgf_no_quant_conv(
254300
).run()
255301

256302

303+
@common.SkipIfNoModelConverter
304+
def test_grouped_conv_transpose_batch_norm_vgf_no_quant() -> None:
305+
model = BatchNorm2dConvTranspose(groups=2)
306+
VgfPipeline[Input](
307+
model,
308+
(torch.rand(1, 4, 5, 6),),
309+
aten_op=model.aten_ops,
310+
quantize=False,
311+
).run()
312+
313+
257314
@common.parametrize("case", batch_norm_cases)
258315
@common.SkipIfNoModelConverter
259316
def test_native_batch_norm_legit_no_training_vgf_quant_conv(

‎backends/arm/test/passes/test_fuse_batchnorm_pass.py‎

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -104,6 +104,47 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
104104
return x
105105

106106

107+
class MergeConvTransposeBN(torch.nn.Module):
108+
ops_before_pass: ClassVar[Dict[str, int]] = {
109+
"executorch_exir_dialects_edge__ops_aten__native_batch_norm_legit_no_training_default": 1,
110+
"executorch_exir_dialects_edge__ops_aten_convolution_default": 1,
111+
}
112+
ops_after_pass: ClassVar[Dict[str, int]] = {
113+
"executorch_exir_dialects_edge__ops_aten__native_batch_norm_legit_no_training_default": 0,
114+
"executorch_exir_dialects_edge__ops_aten_convolution_default": 1,
115+
}
116+
117+
def __init__(
118+
self,
119+
groups: int = 1,
120+
bias: bool = False,
121+
affine: bool = True,
122+
in_channels: int = 4,
123+
out_channels: int = 6,
124+
) -> None:
125+
super().__init__()
126+
self.conv_transpose2d = torch.nn.ConvTranspose2d(
127+
in_channels=in_channels,
128+
out_channels=out_channels,
129+
kernel_size=2,
130+
stride=2,
131+
groups=groups,
132+
bias=bias,
133+
)
134+
self.batch_norm2d = torch.nn.BatchNorm2d(out_channels, affine=affine)
135+
self.batch_norm2d.running_mean = torch.rand(out_channels)
136+
self.batch_norm2d.running_var = torch.rand(out_channels)
137+
if affine:
138+
self.batch_norm2d.weight = torch.nn.Parameter(torch.rand(out_channels))
139+
self.batch_norm2d.bias = torch.nn.Parameter(torch.rand(out_channels))
140+
141+
def get_inputs(self) -> input_t:
142+
return (torch.randn(1, self.conv_transpose2d.in_channels, 8, 8),)
143+
144+
def forward(self, x: torch.Tensor) -> torch.Tensor:
145+
return self.batch_norm2d(self.conv_transpose2d(x))
146+
147+
107148
class MergeMultipleUsersBN(torch.nn.Module):
108149
ops_before_pass: ClassVar[Dict[str, int]] = {
109150
"executorch_exir_dialects_edge__ops_aten__native_batch_norm_legit_no_training_default": 2,
@@ -154,6 +195,20 @@ def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
154195
"merge_two_of_two_bn_affine": cast(
155196
ModuleWithBatchNormAttrs, MergeTwosOfTwoBN(True)
156197
),
198+
"merge_conv_transpose_bn": cast(ModuleWithBatchNormAttrs, MergeConvTransposeBN()),
199+
"merge_grouped_conv_transpose_bn": cast(
200+
ModuleWithBatchNormAttrs, MergeConvTransposeBN(groups=2)
201+
),
202+
"merge_grouped_conv_transpose_bn_bias": cast(
203+
ModuleWithBatchNormAttrs, MergeConvTransposeBN(groups=2, bias=True)
204+
),
205+
"merge_grouped_conv_transpose_bn_no_affine": cast(
206+
ModuleWithBatchNormAttrs, MergeConvTransposeBN(groups=2, affine=False)
207+
),
208+
"merge_grouped_conv_transpose_bn_equal_channels": cast(
209+
ModuleWithBatchNormAttrs,
210+
MergeConvTransposeBN(groups=2, in_channels=4, out_channels=4),
211+
),
157212
"merge_multiple_users_bn_affine": cast(
158213
ModuleWithBatchNormAttrs, MergeMultipleUsersBN(True)
159214
),

‎docs/source/backends/arm-vgf/VGF_op_support.md‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -43,7 +43,7 @@ Total supported PyTorch APIs: **154**.
4343
| `torch.conv1d` | FP, INT | `FP32`, `INT8`, `INT4` | 8x8, 8x4 |
4444
| `torch.conv2d` | FP, INT | `FP32`, `FP16`, `BF16`, `INT8`, `INT16`, `INT4` | 8x8, 8x4, 16x8 |
4545
| `torch.conv3d` | FP, INT | `FP32`, `FP16`, `BF16`, `INT8`, `INT16`, `INT4` | 8x8, 8x4, 16x8 |
46-
| `torch.conv_transpose2d` | FP, INT | `FP16`, `BF16`, `INT8`, `INT16`, `INT4` | 8x8, 8x4, 16x8 |
46+
| `torch.conv_transpose2d` | FP, INT | `FP32`, `FP16`, `BF16`, `INT8`, `INT16`, `INT4` | 8x8, 8x4, 16x8 |
4747
| `torch.cos` | FP, INT | `FP16`, `BF16`, `INT8` | 8x8 |
4848
| `torch.cosh` | FP, INT | `FP32`, `INT8` | 8x8 |
4949
| `torch.cumsum` | FP, INT | `FP32`, `INT8` | 8x8 |

0 commit comments

Comments
 (0)