llvm / llvm/torch-mlir

`DecomposeAtenNativeBatchNormOp` uses wrong dtype for running stats reshape

Open
#4,480 0 comments 0 reactions 1 assignee View on GitHub

@rkayaith is already working on this.

Since Mar 2, 2026.

Dominant language
C++
Stars
1.9k
Forks
736
Avg merge
5d 22h
Merged PRs (30d)
15

Description

Repo: llvm/torch-mlir
Title: `DecomposeAtenNativeBatchNormOp` uses wrong dtype for running stats reshape

---

`DecomposeAtenNativeBatchNormOp` reshapes `running_mean` and `running_var` from `[C]` to `[1,C,1,...]` to broadcast with the input. The result type of the reshape uses the *input* dtype instead of the running stats dtype, producing an invalid `aten.view` when the types differ (e.g., bf16 input with f32 running stats).

The bug is at [DecomposeComplexOps.cpp:L8497-L8499](https://github.com/llvm/torch-mlir/blob/17ae730d5f6e9e1a9eaa7ad12a28d14e6ba113ac/lib/Dialect/Torch/Transforms/DecomposeComplexOps.cpp#L8497-L8499):
```cpp
Type dtype = cast(input.getType()).getOptionalDtype();
Type reshapeType = ValueTensorType::get(
context, llvm::ArrayRef(runningStatsShapeInt), dtype);
```

`dtype` should come from `runningMean`, not `input`.

### Reproducer

```mlir
// bn_mixed_precision.mlir
func.func @main(%arg0: !torch.vtensor<[8,64,56,56],bf16>)
-> !torch.vtensor<[8,64,56,56],bf16> {
%none = torch.constant.none
%running_mean = torch.vtensor.literal(dense<0.0> : tensor<64xf32>) : !torch.vtensor<[64],f32>
%running_var = torch.vtensor.literal(dense<1.0> : tensor<64xf32>) : !torch.vtensor<[64],f32>
%false = torch.constant.bool false
%momentum = torch.constant.float 1.000000e-01
%eps = torch.constant.float 1.000000e-05
%0:3 = torch.aten.native_batch_norm %arg0, %none, %none,
%running_mean, %running_var, %false, %momentum, %eps :
!torch.vtensor<[8,64,56,56],bf16>, !torch.none, !torch.none,
!torch.vtensor<[64],f32>, !torch.vtensor<[64],f32>,
!torch.bool, !torch.float, !torch.float
-> !torch.vtensor<[8,64,56,56],bf16>,
!torch.vtensor<[0],bf16>, !torch.vtensor<[0],bf16>
return %0#0 : !torch.vtensor<[8,64,56,56],bf16>
}
```

Running just the decomposition pass shows the invalid view ops:
```
$ torch-mlir-opt bn_mixed_precision.mlir --torch-decompose-complex-ops | grep aten.view
%1 = torch.aten.view %running_mean, %0 : !torch.vtensor<[64],f32>, !torch.list -> !torch.vtensor<[1,64,1,1],bf16>
^^^ ^^^^
```

Running the full pipeline crashes during constant folding of the invalid view (once #4479 is fixed, this will be a verification error instead):
```
$ torch-mlir-opt bn_mixed_precision.mlir --torch-function-to-torch-backend-pipeline
torch-mlir-opt: .../mlir/lib/IR/BuiltinAttributes.cpp:973:
Assertion `floatAttr.getType() == eltType && "expected float attribute type to equal element type"' failed.
```

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.