llvm / llvm/torch-mlir

Add support for aten.miopen_batch_norm

Open
#4,476 1 comment 0 reactions 1 assignee View on GitHub

@rkayaith is already working on this.

Since Feb 23, 2026.

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

Description

On ROCm 7.1+ (MIOpen 3.5+), PyTorch decomposes `BatchNorm2d` with channels-last inputs to `aten.miopen_batch_norm` instead of `aten._native_batch_norm_legit_functional`. torch-mlir doesn't have a lowering for this op, so it gets imported as a generic `torch.operator` and fails legalization.

## Reproducer

```mlir
// repro.mlir
module @module {
func.func @main(%arg0: !torch.vtensor<[8,64,56,56],bf16>) -> !torch.vtensor<[8,64,56,56],bf16> {
%weight = torch.vtensor.literal(dense<1.0> : tensor<64xf32>) : !torch.vtensor<[64],f32>
%bias = torch.vtensor.literal(dense<0.0> : tensor<64xf32>) : !torch.vtensor<[64],f32>
%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.operator "torch.aten.miopen_batch_norm"(
%arg0, %weight, %bias, %running_mean, %running_var,
%false, %momentum, %eps
) : (!torch.vtensor<[8,64,56,56],bf16>,
!torch.vtensor<[64],f32>, !torch.vtensor<[64],f32>,
!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>
}
}
```

```
$ torch-mlir-opt repro.mlir --torch-function-to-torch-backend-pipeline
repro.mlir:12:12: error: failed to legalize operation 'torch.operator' that was explicitly marked illegal
%0:3 = torch.operator "torch.aten.miopen_batch_norm"(
^
```

### IREE Turbine reproducer

```python
import torch
import iree.turbine.aot as aot

model = torch.nn.BatchNorm2d(64).to(device="cuda", memory_format=torch.channels_last)
model.eval()
x = torch.randn(8, 64, 56, 56, device="cuda", dtype=torch.bfloat16).to(
memory_format=torch.channels_last
)

exported = aot.export(model, x)
exported.print_readable()
exported.compile(save_to=None)
```

## Suggested fix

Add a decomposition mapping `aten.miopen_batch_norm → aten.native_batch_norm`. The ops have the same signature and compatible output semantics. PyTorch Inductor has an equivalent decomposition: [torch/_inductor/decomposition.py#L867-L896](https://github.com/pytorch/pytorch/blob/449b1768410104d3ed79d3bcfe4ba1d65c7f22c0/torch/_inductor/decomposition.py#L867-L896)

## Environment

- PyTorch `2.10.0+rocm7.1`
- https://github.com/iree-org/iree-turbine/commit/391729bf9a123e9dcaf5faf449f480179eeb6107

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.