Skip to content

[Bug][Relax][ONNX] Training BatchNormalization uses the opposite momentum convention #20503

Description

@Yuhx141

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])

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions