Description
DecomposeOpsForTraining builds the BatchNorm reduction axes by comparing each non-negative
dimension index directly with the raw axis attribute. A valid negative axis therefore never
matches, so the pass reduces the channel dimension as well and then raises while normalizing the
rewritten BatchNorm call.
Environment
- Ubuntu 24.04.4 LTS, x86-64
- TVM
0.26.dev0
- Runtime reproduction:
e269315c90e3a061c9e1c77b370ce883b1b223f4
- LLVM 18.1.3, CPU bytecode VM
Current upstream main was 2393db61bbd0100f048b1a6d3c2519fb70914564 when checked on
2026-09-30. The relevant files have changed elsewhere, but the current pass still compares loop
indices with raw attrs->axis, while BatchNorm type inference accepts the same negative axis after
normalization. I did not have a binary build for that exact main revision, so the isolated pass
run below is from the stated fixed revision.
Minimal reproducer
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((2, 3, 4), "float32"),
gamma: R.Tensor((4,), "float32"),
beta: R.Tensor((4,), "float32"),
moving_mean: R.Tensor((4,), "float32"),
moving_var: R.Tensor((4,), "float32"),
):
return R.nn.batch_norm(
x, gamma, beta, moving_mean, moving_var, axis=-1
)
assert relax.analysis.check_well_formed(Module, check_ty=True)
relax.transform.DecomposeOpsForTraining()(Module) # raises ValueError
Observed error:
ValueError: batch_norm requires the input moving_mean to have as many dimensions as the
length of input axes. However, the given one has ndim 0, which is other than the length
of axes 1
Expected behavior
The pass should normalize the BatchNorm axis before selecting reduction dimensions. axis=-1
should behave like the equivalent last non-negative axis.
Actual behavior
The source passes check_well_formed(..., check_ty=True). For a negative axis, the pass reduces
every data dimension, creates scalar batch statistics, and raises when those statistics no longer
match the original rank-one moving statistics.
Ranks 2 through 7 reproduce the failure. Replacing only axis=-1 with the equivalent
axis=rank-1 gives six passing controls: the pass changes the module, source and target remain well
formed, and their LLVM VM outputs match.
In this build, the negative-axis source also reaches a separate limitation in BatchNorm
legalization (list.remove(x): x not in list), so I do not use the risk source as a runtime value
oracle. The claim here is the isolated pass failure on a well-formed module.
Suspected cause
MutateBatchNormForTraining in src/relax/transform/decompose_ops.cc contains:
for (int i = 0; i < ty->ndim; ++i) {
if (i != attrs->axis) {
reduce_axes.push_back(i);
}
}
i is never negative. The operator type checker already normalizes axes when validating the input,
so the pass should do the same before this loop.
Duplicate check
I searched existing TVM issues and pull requests for DecomposeOpsForTraining, BatchNorm, and
negative axes. I found reports about CUDA code generation and inference/training semantics, but none
with the same cause: using an unnormalized negative axis to build the reduction set.
I can send a fix PR and help with follow-up testing if this diagnosis looks right.
Triage
Description
DecomposeOpsForTrainingbuilds the BatchNorm reduction axes by comparing each non-negativedimension index directly with the raw
axisattribute. A valid negative axis therefore nevermatches, so the pass reduces the channel dimension as well and then raises while normalizing the
rewritten BatchNorm call.
Environment
0.26.dev0e269315c90e3a061c9e1c77b370ce883b1b223f4Current upstream
mainwas2393db61bbd0100f048b1a6d3c2519fb70914564when checked on2026-09-30. The relevant files have changed elsewhere, but the current pass still compares loop
indices with raw
attrs->axis, while BatchNorm type inference accepts the same negative axis afternormalization. I did not have a binary build for that exact
mainrevision, so the isolated passrun below is from the stated fixed revision.
Minimal reproducer
Observed error:
Expected behavior
The pass should normalize the BatchNorm axis before selecting reduction dimensions.
axis=-1should behave like the equivalent last non-negative axis.
Actual behavior
The source passes
check_well_formed(..., check_ty=True). For a negative axis, the pass reducesevery data dimension, creates scalar batch statistics, and raises when those statistics no longer
match the original rank-one moving statistics.
Ranks 2 through 7 reproduce the failure. Replacing only
axis=-1with the equivalentaxis=rank-1gives six passing controls: the pass changes the module, source and target remain wellformed, and their LLVM VM outputs match.
In this build, the negative-axis source also reaches a separate limitation in BatchNorm
legalization (
list.remove(x): x not in list), so I do not use the risk source as a runtime valueoracle. The claim here is the isolated pass failure on a well-formed module.
Suspected cause
MutateBatchNormForTraininginsrc/relax/transform/decompose_ops.cccontains:iis never negative. The operator type checker already normalizes axes when validating the input,so the pass should do the same before this loop.
Duplicate check
I searched existing TVM issues and pull requests for
DecomposeOpsForTraining, BatchNorm, andnegative axes. I found reports about CUDA code generation and inference/training semantics, but none
with the same cause: using an unnormalized negative axis to build the reduction set.
I can send a fix PR and help with follow-up testing if this diagnosis looks right.
Triage