Skip to content

Honor the dtype argument of aten::mean when exporting without dim - #3009

Open
om singhal (Om-singhaI) wants to merge 1 commit into
microsoft:mainfrom
Om-singhaI:fix/aten-mean-dtype
Open

Honor the dtype argument of aten::mean when exporting without dim#3009
om singhal (Om-singhaI) wants to merge 1 commit into
microsoft:mainfrom
Om-singhaI:fix/aten-mean-dtype

Conversation

@Om-singhaI

Copy link
Copy Markdown

Honor the dtype argument of aten::mean when exporting without dim

Fixes #3008

Summary

The no dim overload of aten::mean declares dtype in its schema (mean(Tensor self, *, ScalarType? dtype=None)), but aten_mean in onnxscript/function_libs/torch_lib/ops/core.py did not accept the argument. Exporting torch.mean(x, dtype=torch.float64) for a float32 input succeeded silently, emitted only ReduceMean -> Squeeze, declared a FLOAT output, and accumulated in float32. For the input [[1e8, 1.0, -1e8]] PyTorch returns 0.3333333333333333 as float64 while the exported model returned 0.0 as float32.

Changes

  • aten_mean is now trace_only=True, takes dtype: int = -1, and when a dtype is given casts self before ReduceMean. The no dtype path is unchanged (ReduceMean -> Squeeze).
  • aten_mean_complex takes the same argument and raises NotImplementedError when it is supplied, matching aten_mean_dim_complex and aten_sum_complex.
  • New ops.aten.mean.dtype OpInfo in tests/function_libs/torch_lib/extra_opinfo.py (sample_inputs_mean_dtype) registered against core_ops.aten_mean in ops_test_data.py. It yields make_tensor samples of shapes (5, 5), (5,) and () plus the precision sensitive tensor [[1e8, 1.0, -1e8]], all with dtype=torch.float64.

Why the cast happens before the reduction

aten_mean_dim (#2885) and aten_sum cast the reduced result after the reduction. That is not sufficient here: PyTorch accumulates in the requested dtype, so for [1e8, 1.0, -1e8] the float32 mean is 0.0 (the 1.0 is lost when added to 1e8) while the float64 mean is 1/3. Casting after ReduceMean would produce a DOUBLE tensor holding 0.0, which still mismatches PyTorch. Casting the input first reproduces PyTorch semantics. The new test includes this sample specifically so that a cast after reduction implementation cannot pass by accident.

Verification

Reporter's script (torch.onnx.export(..., dynamo=True) of torch.mean(x, dtype=torch.float64) with float32 input) on pristine main at a39c0a5:

torch eager result: 0.3333333333333333 dtype: torch.float64
ONNX output elem_type: 1 (FLOAT), expected 11 (DOUBLE)
ORT result: 0.0 dtype: float32

With this change:

torch eager result: 0.3333333333333333 dtype: torch.float64
ONNX output elem_type: 11 (DOUBLE)
ORT result: 0.3333333333333333 dtype: float64

pytest tests/function_libs/torch_lib/ops_test.py -k mean:

State Result
Pristine main, without the new test 8 passed, 48 skipped, 8 xfailed, 60 subtests passed
New test with the source change stashed 4 failed, 10 passed, 48 skipped, 8 xfailed, 60 subtests passed. All four ops_aten_mean_dtype samples fail with TypeInferenceError: Inferred elem type differs from existing elem type: (1) vs (11)
New test with a cast placed after the reduction 1 failed, 8 passed, 50 skipped, 8 xfailed, 63 subtests passed. Only the [[1e8, 1.0, -1e8]] sample fails: Expected 0.3333333333333333 but got 0.0
New test with this change 8 passed, 50 skipped, 8 xfailed, 64 subtests passed

The two additional skips in the fixed state are the function proto validity checks, which skip for traced functions.

ruff check and ruff format --check (ruff 0.15.1, the lintrunner pinned version) pass on the three changed files.

Environment: Python 3.10, torch 2.13.0 (CPU), onnx 1.22.0, onnxruntime 1.23.2, macOS arm64.

The no dim overload of aten::mean declares a dtype keyword in its schema, but the torchlib implementation did not accept it. Exporting torch.mean(x, dtype=torch.float64) for a float32 input succeeded silently, produced a FLOAT output and accumulated in float32.

aten_mean is now a traced function that takes dtype and casts the input before ReduceMean, so the accumulation happens in the requested type exactly as PyTorch does. This differs on purpose from aten_mean_dim and aten_sum, which cast the reduced result. For the input 1e8, 1.0 and negative 1e8 the float32 mean is 0.0 while the float64 mean is one third, so casting after the reduction would still give the wrong value.

aten_mean_complex gains the same argument and raises NotImplementedError when it is supplied, matching aten_mean_dim_complex.

A new ops.aten.mean.dtype OpInfo exercises the fix, including a precision sensitive sample that fails when the cast is applied after the reduction.

Fixes microsoft#3008

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

While comparing this with the duplicate implementation I had opened, I checked out commit c48d820 locally and verified the surviving approach directly. Ruff check and format check pass on all three changed files. Exporting torch.mean([1e8, 1, -1e8], dtype=torch.float64) produces Cast -> ReduceMean -> Squeeze, returns float64, and matches PyTorch at 1/3. The OpInfo coverage is stronger than a single e2e regression, and accepting dtype on the complex overload while explicitly rejecting unsupported conversion keeps its schema behavior consistent. This looks correct to me.

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

Labels

None yet

Projects

Development

Successfully merging this pull request may close these issues.

aten::mean silently ignores dtype and exports the wrong output type

2 participants