Skip to content

[Bug][Relax] DecomposeOpsForTraining rejects BatchNorm with a valid negative axis #20496

Description

@Yuhx141

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

  • needs-triage
  • type: bug

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