Version
- Executed revision:
e269315c90e3a061c9e1c77b370ce883b1b223f4
- Current
main source rechecked: 0587e59a696cb435ad4d328726a07daa8d966f3c
- TVM:
0.26.dev0
- Host: Ubuntu 24.04 x86-64, LLVM CPU
Reproducer
Run the complete script below from a TVM source environment.
It creates a valid opset-15 ONNX BatchNormalization with training_mode=1 and momentum=0.9, imports it with
from_onnx, runs DecomposeOpsForTraining, and compares the returned running mean with ONNX's reference evaluator.
Observed output:
TVM running_mean: [16.750002 30.350002 43.95001 ]
ONNX running_mean: [ 90.75 181.15 271.55 ]
Max absolute difference: 227.59998
I also checked channel counts 1 through 6. All six momentum=0.9 cases showed the same reversal, while the
corresponding momentum=0.5 controls matched.
Expected behavior
ONNX defines the training update as:
running_mean = input_mean * momentum + current_mean * (1 - momentum)
running_var = input_var * momentum + current_var * (1 - momentum)
The imported and decomposed Relax function should return those values.
Actual behavior and likely cause
The ONNX frontend forwards the ONNX value unchanged:
momentum = attr.get("momentum", 0.9)
...
relax.op.nn.batch_norm(..., momentum=momentum, training=bool(training_mode))
DecomposeOpsForTraining and TOPI interpret that same value using the opposite convention:
new_moving_mean = (1 - momentum) * moving_mean + momentum * data_mean;
new_moving_var = (1 - momentum) * moving_var + momentum * data_var;
The Relax Python documentation itself states the ONNX convention, but the existing Relax structural test encodes the
TOPI convention. The smallest frontend fix appears to be converting the imported value to 1 - momentum, unless the
Relax operator contract is changed consistently instead.
Both the direct attribute forwarding and the reversed decomposition remain in current main.
Duplicate check
PR #18704 added direct propagation of momentum and training_mode, but did not account for the two momentum
conventions. GitHub issue and pull-request searches for ONNX BatchNormalization momentum, Relax BatchNorm momentum,
and DecomposeOpsForTraining found no report of this coefficient reversal as of 2026-09-30.
Complete reproducer
import numpy as np
import onnx
import tvm
from onnx import TensorProto, helper
from onnx.reference import ReferenceEvaluator
from tvm import relax
from tvm.relax.frontend.onnx import from_onnx
names = ("x", "scale", "bias", "mean", "var")
shapes = ([2, 3, 2, 2], [3], [3], [3], [3])
inputs = [helper.make_tensor_value_info(n, TensorProto.FLOAT, s) for n, s in zip(names, shapes)]
outputs = [
helper.make_tensor_value_info("y", TensorProto.FLOAT, [2, 3, 2, 2]),
helper.make_tensor_value_info("running_mean", TensorProto.FLOAT, [3]),
helper.make_tensor_value_info("running_var", TensorProto.FLOAT, [3]),
]
node = helper.make_node(
"BatchNormalization",
list(names),
[value.name for value in outputs],
momentum=0.9,
training_mode=1,
)
model = helper.make_model(
helper.make_graph([node], "batch_norm", inputs, outputs),
opset_imports=[helper.make_opsetid("", 15)],
)
onnx.checker.check_model(model)
module = relax.transform.DecomposeOpsForTraining()(from_onnx(model, keep_params_in_input=True))
executable = relax.build(module, target="llvm", relax_pipeline="default", exec_mode="bytecode")
vm = relax.VirtualMachine(executable, tvm.cpu())
feeds = {
"x": np.arange(24, dtype="float32").reshape(2, 3, 2, 2),
"scale": np.ones(3, dtype="float32"),
"bias": np.zeros(3, dtype="float32"),
"mean": np.array([100, 200, 300], dtype="float32"),
"var": np.array([4, 9, 16], dtype="float32"),
}
tvm_output = vm["main"](*(tvm.runtime.tensor(feeds[name]) for name in names))
onnx_output = ReferenceEvaluator(model).run(None, feeds)
print("TVM running_mean: ", tvm_output[1].numpy())
print("ONNX running_mean:", onnx_output[1])
print("Max absolute difference:", np.max(np.abs(tvm_output[1].numpy() - onnx_output[1])))
np.testing.assert_allclose(tvm_output[1].numpy(), onnx_output[1])
Version
e269315c90e3a061c9e1c77b370ce883b1b223f4mainsource rechecked:0587e59a696cb435ad4d328726a07daa8d966f3c0.26.dev0Reproducer
Run the complete script below from a TVM source environment.
It creates a valid opset-15 ONNX
BatchNormalizationwithtraining_mode=1andmomentum=0.9, imports it withfrom_onnx, runsDecomposeOpsForTraining, and compares the returned running mean with ONNX's reference evaluator.Observed output:
I also checked channel counts 1 through 6. All six
momentum=0.9cases showed the same reversal, while thecorresponding
momentum=0.5controls matched.Expected behavior
ONNX defines the training update as:
The imported and decomposed Relax function should return those values.
Actual behavior and likely cause
The ONNX frontend forwards the ONNX value unchanged:
DecomposeOpsForTrainingand TOPI interpret that same value using the opposite convention:The Relax Python documentation itself states the ONNX convention, but the existing Relax structural test encodes the
TOPI convention. The smallest frontend fix appears to be converting the imported value to
1 - momentum, unless theRelax operator contract is changed consistently instead.
Both the direct attribute forwarding and the reversed decomposition remain in current
main.Duplicate check
PR #18704 added direct propagation of
momentumandtraining_mode, but did not account for the two momentumconventions. GitHub issue and pull-request searches for ONNX BatchNormalization momentum, Relax BatchNorm momentum,
and
DecomposeOpsForTrainingfound no report of this coefficient reversal as of 2026-09-30.Complete reproducer