Description
FoldBatchnormToConv2D rewrites a Conv2D followed by BatchNorm even when the
BatchNorm has scale=False or center=False. The replacement always uses both
gamma and beta, so the pass changes the result of a well-formed executable module.
Environment
- TVM
0.26.dev0, executed at e269315c90e3a061c9e1c77b370ce883b1b223f4
- Ubuntu 24.04 x86-64, LLVM 18.1.3 CPU target
- Current upstream
main checked at 043b57fbe9b84aba6c0525fc77b624ef8b7bbe2f
- The relevant Python source is byte-identical at both revisions (SHA-256
c591c953ce9de0f8b539149ab466692572ae608978bbb7ca4875fe919f3c588e)
- I did not have a binary built from current
main, so the runtime reproduction below is from the executed revision
Minimal reproduction
import numpy as np
import tvm
from tvm import relax
from tvm.script import ir as I
from tvm.script import relax as R
@I.ir_module
class Module:
@R.function
def main(
x: R.Tensor((1, 1, 2, 2), "float32"),
weight: R.Tensor((2, 1, 1, 1), "float32"),
gamma: R.Tensor((2,), "float32"),
beta: R.Tensor((2,), "float32"),
mean: R.Tensor((2,), "float32"),
variance: R.Tensor((2,), "float32"),
):
conv = R.nn.conv2d(x, weight, data_layout="NCHW", kernel_layout="OIHW", out_layout="NCHW", out_dtype="float32")
return R.nn.batch_norm(conv, gamma, beta, mean, variance, axis=1, scale=False, center=True, training=False)[0]
params = {
"weight": tvm.runtime.tensor(np.array([[[[1.5]]], [[[-2.0]]]], "float32")),
"gamma": tvm.runtime.tensor(np.array([2.0, 3.0], "float32")),
"beta": tvm.runtime.tensor(np.array([5.0, 7.0], "float32")),
"mean": tvm.runtime.tensor(np.array([1.0, -1.0], "float32")),
"variance": tvm.runtime.tensor(np.array([4.0, 9.0], "float32")),
}
source = relax.transform.BindParams("main", params)(Module)
target = relax.transform.FoldBatchnormToConv2D()(source)
assert relax.analysis.check_well_formed(source, check_ty=True)
assert relax.analysis.check_well_formed(target, check_ty=True)
def run(mod):
ex = relax.build(mod, target="llvm", relax_pipeline="default", exec_mode="bytecode")
vm = relax.VirtualMachine(ex, tvm.cpu())
x = tvm.runtime.tensor(np.array([[[[1, -2], [3, 4]]]], "float32"))
return vm["main"](x).numpy()
print(np.max(np.abs(run(source) - run(target))))
Expected behavior
The pass should preserve a disabled affine term or leave this BatchNorm unchanged. gamma must not affect a scale=False result, and beta must not affect a center=False result.
Actual behavior
The source and target are both well formed and executable, but their outputs differ. Independent runs covered channels 1 through 6 for both scale=False and center=False: all 12 variants differed, with maximum absolute differences from about 0.875 to 5.04. The adjacent scale=True, center=True controls matched within 1e-6.
Likely cause
fold_batch_norm_to_conv2d_for_inference.py matches any BatchNorm, reads only epsilon, and always computes with both gamma and beta. It never checks scale or center before replacing the operator.
Duplicate check
I searched open and closed TVM issues and PRs for this pass and the two attributes, and did not find the same cause. #18000 is a different multi-op pipeline report; #17654 is the original implementation PR.
I can send a fix PR and help with follow-up testing if this diagnosis looks right.
Triage
Description
FoldBatchnormToConv2Drewrites a Conv2D followed by BatchNorm even when theBatchNorm has
scale=Falseorcenter=False. The replacement always uses bothgammaandbeta, so the pass changes the result of a well-formed executable module.Environment
0.26.dev0, executed ate269315c90e3a061c9e1c77b370ce883b1b223f4mainchecked at043b57fbe9b84aba6c0525fc77b624ef8b7bbe2fc591c953ce9de0f8b539149ab466692572ae608978bbb7ca4875fe919f3c588e)main, so the runtime reproduction below is from the executed revisionMinimal reproduction
Expected behavior
The pass should preserve a disabled affine term or leave this BatchNorm unchanged.
gammamust not affect ascale=Falseresult, andbetamust not affect acenter=Falseresult.Actual behavior
The source and target are both well formed and executable, but their outputs differ. Independent runs covered channels 1 through 6 for both
scale=Falseandcenter=False: all 12 variants differed, with maximum absolute differences from about0.875to5.04. The adjacentscale=True, center=Truecontrols matched within1e-6.Likely cause
fold_batch_norm_to_conv2d_for_inference.pymatches any BatchNorm, reads onlyepsilon, and always computes with bothgammaandbeta. It never checksscaleorcenterbefore replacing the operator.Duplicate check
I searched open and closed TVM issues and PRs for this pass and the two attributes, and did not find the same cause. #18000 is a different multi-op pipeline report; #17654 is the original implementation PR.
I can send a fix PR and help with follow-up testing if this diagnosis looks right.
Triage