From 9725f3ca4a144d96cdaf4404eea559a598fa38e3 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Wed, 15 Jul 2026 23:53:36 +0000 Subject: [PATCH 1/5] Initial plan From 65ce6009fa239a3b6715508c51dd608e4030b739 Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 16 Jul 2026 00:01:24 +0000 Subject: [PATCH 2/5] Fix bool index_put mask broadcasting --- .../function_libs/torch_lib/ops/core.py | 60 ++++++++----------- .../function_libs/torch_lib/e2e_ops_tests.py | 37 ++++++++++++ 2 files changed, 62 insertions(+), 35 deletions(-) diff --git a/onnxscript/function_libs/torch_lib/ops/core.py b/onnxscript/function_libs/torch_lib/ops/core.py index adf1bad4b6..33162aa69b 100644 --- a/onnxscript/function_libs/torch_lib/ops/core.py +++ b/onnxscript/function_libs/torch_lib/ops/core.py @@ -4970,8 +4970,21 @@ def aten_index_put( See implementation of `torch.onnx.symbolic_opset11.index_put `_. """ - if any(index is not None and index.dtype == BOOL.dtype for index in indices): + bool_index_positions = [ + i for i, index in enumerate(indices) if index is not None and index.dtype == BOOL.dtype + ] + if len(indices) == 1 and bool_index_positions == [0]: return _aten_index_put_bool(self, indices, values, accumulate) + if bool_index_positions: + neg_1 = op.Constant(value_ints=[-1]) + indices = list(indices) + for i in bool_index_positions: + index = indices[i] + if len(index.shape) != 1: + raise NotImplementedError( + "Boolean index_put with mixed or multi-indices supports only 1-D boolean masks." + ) + indices[i] = op.Reshape(op.Transpose(op.NonZero(index), perm=[1, 0]), neg_1) # Ensure the number of indices matches the tensor rank by appending trailing Nones. self_rank = len(self.shape) @@ -4986,6 +4999,11 @@ def is_advanced_index(index): # Note: In this function, the index is assumed to be either None or an int64 Tensor. return index is not None + def index_rank(index_position: int) -> int: + if index_position in bool_index_positions: + return 1 + return len(indices[index_position].shape) + advanced_indices: list[int] = [] none_indices: list[int] = [] num_advanced_indices = 0 @@ -5024,15 +5042,15 @@ def same_shape(other_shape: Optional[ir.Shape]) -> bool: all_same_shape = all(same_shape(indices[i].shape) for i in advanced_indices) if not all_same_shape: # Broadcast advanced indices to a common shape. - advanced_index_rank = max(len(indices[i].shape) for i in advanced_indices) + advanced_index_rank = max(index_rank(i) for i in advanced_indices) shapes = [] for i in advanced_indices: index = indices[i] - index_rank = len(index.shape) + current_index_rank = index_rank(i) index_shape = op.Shape(index) - if index_rank < advanced_index_rank: + if current_index_rank < advanced_index_rank: padding = op.Constant( - value_ints=[1 for _ in range(advanced_index_rank - index_rank)] + value_ints=[1 for _ in range(advanced_index_rank - current_index_rank)] ) index_shape = op.Concat(padding, index_shape, axis=0) shapes.append(index_shape) @@ -5043,10 +5061,10 @@ def same_shape(other_shape: Optional[ir.Shape]) -> bool: ] else: advanced_indices_shape = op.Shape(indices[advanced_indices[0]]) - advanced_index_rank = len(indices[advanced_indices[0]].shape) + advanced_index_rank = index_rank(advanced_indices[0]) else: advanced_indices_shape = op.Shape(indices[advanced_indices[0]]) - advanced_index_rank = len(indices[advanced_indices[0]].shape) + advanced_index_rank = index_rank(advanced_indices[0]) # ONNX ScatterND supports only the case where all advanced indices appear first, # followed by None indices. So, we need to transpose self and values so that the @@ -5120,34 +5138,6 @@ def _aten_index_put_bool( """index_put(Tensor self, Tensor?[] indices, Tensor values, bool accumulate=False) -> Tensor""" bool_mask = indices[0] - if len(indices) > 1: - if any(index is None for index in indices): - raise NotImplementedError( - "Boolean index_put with multiple indices does not support None indices." - ) - - advanced_indices = [] - selected_positions = [] - minus_one = op.Constant(value_ints=[-1]) - for index in indices: - if index.dtype != BOOL.dtype or len(index.shape) != 1: - raise NotImplementedError( - "Boolean index_put with multiple indices supports only 1-D boolean masks." - ) - positions = op.Reshape(op.Transpose(op.NonZero(index), perm=[1, 0]), minus_one) - selected_positions.append(positions) - advanced_indices.append(op.Unsqueeze(positions, minus_one)) - onnx_index = op.Concat(*advanced_indices, axis=-1) - target_shape = op.Concat( - op.Shape(selected_positions[0]), - op.Slice(op.Shape(self), starts=[len(indices)], ends=[len(self.shape)], axes=[0]), - axis=0, - ) - expanded_values = op.Expand(values, target_shape) - return op.ScatterND( - self, onnx_index, expanded_values, reduction="add" if accumulate else None - ) - if bool_mask is None or bool_mask.dtype != BOOL.dtype: raise NotImplementedError( "Boolean index_put expects a boolean mask as the first index." diff --git a/tests/function_libs/torch_lib/e2e_ops_tests.py b/tests/function_libs/torch_lib/e2e_ops_tests.py index e6b481f2c0..019164308f 100644 --- a/tests/function_libs/torch_lib/e2e_ops_tests.py +++ b/tests/function_libs/torch_lib/e2e_ops_tests.py @@ -1232,6 +1232,43 @@ def forward(self, x, mask0, mask1, update): ) _testing.assert_onnx_program(onnx_program) + def test_index_put_bool_multi_mask_broadcast(self): + class Model(torch.nn.Module): + def forward(self, x, mask0, mask1, update): + return torch.ops.aten.index_put(x, [mask0, mask1], update) + + x = torch.zeros((3, 4), dtype=torch.float32) + mask0 = torch.tensor([False, True, False], dtype=torch.bool) + mask1 = torch.tensor([True, False, True, True], dtype=torch.bool) + update = torch.tensor([10.0, 20.0, 30.0], dtype=torch.float32) + onnx_program = torch.onnx.export( + Model(), + (x, mask0, mask1, update), + input_names=["x", "mask0", "mask1", "update"], + output_names=["output"], + opset_version=18, + dynamo=True, + ) + _testing.assert_onnx_program(onnx_program) + + def test_index_put_bool_mask_with_none_index(self): + class Model(torch.nn.Module): + def forward(self, x, mask, update): + return torch.ops.aten.index_put(x, [None, mask], update) + + x = torch.arange(6, dtype=torch.float32).reshape((2, 3)) + mask = torch.tensor([True, False, True], dtype=torch.bool) + update = torch.tensor(9.0, dtype=torch.float32) + onnx_program = torch.onnx.export( + Model(), + (x, mask, update), + input_names=["x", "mask", "update"], + output_names=["output"], + opset_version=18, + dynamo=True, + ) + _testing.assert_onnx_program(onnx_program) + def test_std_mean(self): """Test torch.std_mean which will be decomposed into prims.sum.""" From 632426234e7311631a2880243fc4572cd373458e Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 23 Jul 2026 14:37:21 +0000 Subject: [PATCH 3/5] Replace deprecated Codecov test results action --- .github/workflows/main.yaml | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/.github/workflows/main.yaml b/.github/workflows/main.yaml index 9c2b574a9b..28412dc1e5 100644 --- a/.github/workflows/main.yaml +++ b/.github/workflows/main.yaml @@ -78,9 +78,10 @@ jobs: token: ${{ secrets.CODECOV_TOKEN }} - name: Upload test results to Codecov if: ${{ !cancelled() }} - uses: codecov/test-results-action@v1 + uses: codecov/codecov-action@v5 with: token: ${{ secrets.CODECOV_TOKEN }} + report_type: test_results - name: Upload torchlib error reports if: always() uses: actions/upload-artifact@v7 From 6c844d38cde18cf8d8c8e7911b290e67569ca70c Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 23 Jul 2026 14:45:21 +0000 Subject: [PATCH 4/5] Fix arange and logit behavior for torch-nightly CI --- .../function_libs/torch_lib/ops/core.py | 91 +++++++++++++++++-- 1 file changed, 83 insertions(+), 8 deletions(-) diff --git a/onnxscript/function_libs/torch_lib/ops/core.py b/onnxscript/function_libs/torch_lib/ops/core.py index 33162aa69b..d7cfc60cb7 100644 --- a/onnxscript/function_libs/torch_lib/ops/core.py +++ b/onnxscript/function_libs/torch_lib/ops/core.py @@ -56,6 +56,9 @@ _INT64_MAX = 9223372036854775807 _INT64_MIN = -9223372036854775808 _MATH_PI = math.pi +_INT64_ARANGE_NON_INTEGRAL_USES_FLOAT_LENGTH = ( + torch.arange(3.1, dtype=torch.int64).numel() == 4 +) @torch_op("aten::_local_scalar_dense", trace_only=True) @@ -521,6 +524,56 @@ def _range_supported(dtype: int) -> bool: } +def _is_integral_dtype(dtype: int) -> bool: + return dtype in { + INT8.dtype, + INT16.dtype, + INT32.dtype, + INT64.dtype, + } + + +def _is_integral_scalar(arg: TRealUnlessFloat16OrInt8) -> bool: + if isinstance(arg, int): + return True + if isinstance(arg, float): + return False + return arg.dtype in { + INT8.dtype, + INT16.dtype, + INT32.dtype, + INT64.dtype, + } + + +def _arange_integral_dtype_with_non_integral_args( + start: TRealUnlessFloat16OrInt8, + end: TRealUnlessFloat16OrInt8, + step: TRealUnlessFloat16OrInt8, + dtype: int, +) -> TensorType: + """Implements torch.arange for integral dtypes when not all inputs are integral.""" + start_float = op.Cast(start, to=FLOAT.dtype) + end_float = op.Cast(end, to=FLOAT.dtype) + step_float = op.Cast(step, to=FLOAT.dtype) + + length = op.Cast( + op.Ceil(op.Div(op.Sub(end_float, start_float), step_float)), to=INT64.dtype + ) + index = op.Range(op.Constant(value_int=0), length, op.Constant(value_int=1)) + + if _range_supported(dtype): + index = op.Cast(index, to=dtype) + start = op.Cast(start, to=dtype) + step = op.Cast(step, to=dtype) + return op.Add(start, op.Mul(step, index)) + + start = op.Cast(start, to=INT64.dtype) + step = op.Cast(step, to=INT64.dtype) + result = op.Add(start, op.Mul(step, index)) + return op.Cast(result, to=dtype) + + def _integral_to_be_adjusted(dtype: int) -> bool: """Returns true if the dtype is special integral handled by torch.""" return dtype in { @@ -544,6 +597,12 @@ def aten_arange( zero = op.CastLike(0.0, end) one = op.CastLike(1.0, end) result = op.Range(zero, end, one) + elif ( + _is_integral_dtype(dtype) + and not _is_integral_scalar(end) + and (dtype != INT64.dtype or _INT64_ARANGE_NON_INTEGRAL_USES_FLOAT_LENGTH) + ): + result = _arange_integral_dtype_with_non_integral_args(0, end, 1, dtype) elif _range_supported(dtype): end = op.Cast(end, to=dtype) zero = op.Cast(0, to=dtype) @@ -576,6 +635,12 @@ def aten_arange_start( if dtype == -1 or dtype is None: one = op.CastLike(1.0, end) result = op.Range(start, end, one) + elif ( + _is_integral_dtype(dtype) + and not (_is_integral_scalar(start) and _is_integral_scalar(end)) + and (dtype != INT64.dtype or _INT64_ARANGE_NON_INTEGRAL_USES_FLOAT_LENGTH) + ): + result = _arange_integral_dtype_with_non_integral_args(start, end, 1, dtype) elif _range_supported(dtype): end = op.Cast(end, to=dtype) start = op.Cast(start, to=dtype) @@ -659,17 +724,26 @@ def aten_arange_start_step( end = op.Cast(end, to=FLOAT.dtype) step = op.Cast(step, to=FLOAT.dtype) result = op.Range(start, end, step) - elif _integral_to_be_adjusted(dtype): - # PyTorch arange op handles these integral types differently from INT64, - # so we have to adjust these arguments accordingly. - # https://github.com/pytorch/pytorch/blob/121cfb60c0817816fcbe2190303b7f6d05c77cf3/torch/_refs/__init__.py#L4794 - start, end, step = _adjust_args_for_arange_int_dtype(start, end, step) - result = op.Cast(op.Range(start, end, step), to=dtype) + elif ( + _is_integral_dtype(dtype) + and not ( + _is_integral_scalar(start) + and _is_integral_scalar(end) + and _is_integral_scalar(step) + ) + and (dtype != INT64.dtype or _INT64_ARANGE_NON_INTEGRAL_USES_FLOAT_LENGTH) + ): + result = _arange_integral_dtype_with_non_integral_args(start, end, step, dtype) elif dtype == INT64.dtype: end = op.Cast(end, to=dtype) start = op.Cast(start, to=dtype) step = op.Cast(step, to=dtype) result = op.Range(start, end, step) + elif _integral_to_be_adjusted(dtype): + # PyTorch arange op handles these integral types differently from INT64 + # when all arguments are integral. + start, end, step = _adjust_args_for_arange_int_dtype(start, end, step) + result = op.Cast(op.Range(start, end, step), to=dtype) else: # Cast input to float if dtype is not supported by Range, # because the input dtype may be e.g. bfloat16, @@ -5794,8 +5868,9 @@ def aten_logit(self: TFloat, eps: Optional[float] = None) -> TFloat: one_minus_eps = ir.tensor(1 - eps, dtype=self.dtype) eps = ir.tensor(eps, dtype=self.dtype) - temporary_self = op.Where(self <= one_minus_eps, self, one_minus_eps) - z = op.Where(temporary_self < eps, eps, temporary_self) + # Match torch.clamp behavior for eps > 0.5 by applying max then min. + z = op.Where(self < eps, eps, self) + z = op.Where(z <= one_minus_eps, z, one_minus_eps) return op.Log(op.Div(z, op.Sub(one, z))) From 8ad6041a875cf35e95f3ea503855fc4ed4b2fe4a Mon Sep 17 00:00:00 2001 From: "copilot-swe-agent[bot]" <198982749+Copilot@users.noreply.github.com> Date: Thu, 23 Jul 2026 14:55:18 +0000 Subject: [PATCH 5/5] Remove unrelated torch changes from Codecov PR --- .../function_libs/torch_lib/ops/core.py | 151 +++++------------- .../function_libs/torch_lib/e2e_ops_tests.py | 37 ----- 2 files changed, 43 insertions(+), 145 deletions(-) diff --git a/onnxscript/function_libs/torch_lib/ops/core.py b/onnxscript/function_libs/torch_lib/ops/core.py index d7cfc60cb7..adf1bad4b6 100644 --- a/onnxscript/function_libs/torch_lib/ops/core.py +++ b/onnxscript/function_libs/torch_lib/ops/core.py @@ -56,9 +56,6 @@ _INT64_MAX = 9223372036854775807 _INT64_MIN = -9223372036854775808 _MATH_PI = math.pi -_INT64_ARANGE_NON_INTEGRAL_USES_FLOAT_LENGTH = ( - torch.arange(3.1, dtype=torch.int64).numel() == 4 -) @torch_op("aten::_local_scalar_dense", trace_only=True) @@ -524,56 +521,6 @@ def _range_supported(dtype: int) -> bool: } -def _is_integral_dtype(dtype: int) -> bool: - return dtype in { - INT8.dtype, - INT16.dtype, - INT32.dtype, - INT64.dtype, - } - - -def _is_integral_scalar(arg: TRealUnlessFloat16OrInt8) -> bool: - if isinstance(arg, int): - return True - if isinstance(arg, float): - return False - return arg.dtype in { - INT8.dtype, - INT16.dtype, - INT32.dtype, - INT64.dtype, - } - - -def _arange_integral_dtype_with_non_integral_args( - start: TRealUnlessFloat16OrInt8, - end: TRealUnlessFloat16OrInt8, - step: TRealUnlessFloat16OrInt8, - dtype: int, -) -> TensorType: - """Implements torch.arange for integral dtypes when not all inputs are integral.""" - start_float = op.Cast(start, to=FLOAT.dtype) - end_float = op.Cast(end, to=FLOAT.dtype) - step_float = op.Cast(step, to=FLOAT.dtype) - - length = op.Cast( - op.Ceil(op.Div(op.Sub(end_float, start_float), step_float)), to=INT64.dtype - ) - index = op.Range(op.Constant(value_int=0), length, op.Constant(value_int=1)) - - if _range_supported(dtype): - index = op.Cast(index, to=dtype) - start = op.Cast(start, to=dtype) - step = op.Cast(step, to=dtype) - return op.Add(start, op.Mul(step, index)) - - start = op.Cast(start, to=INT64.dtype) - step = op.Cast(step, to=INT64.dtype) - result = op.Add(start, op.Mul(step, index)) - return op.Cast(result, to=dtype) - - def _integral_to_be_adjusted(dtype: int) -> bool: """Returns true if the dtype is special integral handled by torch.""" return dtype in { @@ -597,12 +544,6 @@ def aten_arange( zero = op.CastLike(0.0, end) one = op.CastLike(1.0, end) result = op.Range(zero, end, one) - elif ( - _is_integral_dtype(dtype) - and not _is_integral_scalar(end) - and (dtype != INT64.dtype or _INT64_ARANGE_NON_INTEGRAL_USES_FLOAT_LENGTH) - ): - result = _arange_integral_dtype_with_non_integral_args(0, end, 1, dtype) elif _range_supported(dtype): end = op.Cast(end, to=dtype) zero = op.Cast(0, to=dtype) @@ -635,12 +576,6 @@ def aten_arange_start( if dtype == -1 or dtype is None: one = op.CastLike(1.0, end) result = op.Range(start, end, one) - elif ( - _is_integral_dtype(dtype) - and not (_is_integral_scalar(start) and _is_integral_scalar(end)) - and (dtype != INT64.dtype or _INT64_ARANGE_NON_INTEGRAL_USES_FLOAT_LENGTH) - ): - result = _arange_integral_dtype_with_non_integral_args(start, end, 1, dtype) elif _range_supported(dtype): end = op.Cast(end, to=dtype) start = op.Cast(start, to=dtype) @@ -724,26 +659,17 @@ def aten_arange_start_step( end = op.Cast(end, to=FLOAT.dtype) step = op.Cast(step, to=FLOAT.dtype) result = op.Range(start, end, step) - elif ( - _is_integral_dtype(dtype) - and not ( - _is_integral_scalar(start) - and _is_integral_scalar(end) - and _is_integral_scalar(step) - ) - and (dtype != INT64.dtype or _INT64_ARANGE_NON_INTEGRAL_USES_FLOAT_LENGTH) - ): - result = _arange_integral_dtype_with_non_integral_args(start, end, step, dtype) + elif _integral_to_be_adjusted(dtype): + # PyTorch arange op handles these integral types differently from INT64, + # so we have to adjust these arguments accordingly. + # https://github.com/pytorch/pytorch/blob/121cfb60c0817816fcbe2190303b7f6d05c77cf3/torch/_refs/__init__.py#L4794 + start, end, step = _adjust_args_for_arange_int_dtype(start, end, step) + result = op.Cast(op.Range(start, end, step), to=dtype) elif dtype == INT64.dtype: end = op.Cast(end, to=dtype) start = op.Cast(start, to=dtype) step = op.Cast(step, to=dtype) result = op.Range(start, end, step) - elif _integral_to_be_adjusted(dtype): - # PyTorch arange op handles these integral types differently from INT64 - # when all arguments are integral. - start, end, step = _adjust_args_for_arange_int_dtype(start, end, step) - result = op.Cast(op.Range(start, end, step), to=dtype) else: # Cast input to float if dtype is not supported by Range, # because the input dtype may be e.g. bfloat16, @@ -5044,21 +4970,8 @@ def aten_index_put( See implementation of `torch.onnx.symbolic_opset11.index_put `_. """ - bool_index_positions = [ - i for i, index in enumerate(indices) if index is not None and index.dtype == BOOL.dtype - ] - if len(indices) == 1 and bool_index_positions == [0]: + if any(index is not None and index.dtype == BOOL.dtype for index in indices): return _aten_index_put_bool(self, indices, values, accumulate) - if bool_index_positions: - neg_1 = op.Constant(value_ints=[-1]) - indices = list(indices) - for i in bool_index_positions: - index = indices[i] - if len(index.shape) != 1: - raise NotImplementedError( - "Boolean index_put with mixed or multi-indices supports only 1-D boolean masks." - ) - indices[i] = op.Reshape(op.Transpose(op.NonZero(index), perm=[1, 0]), neg_1) # Ensure the number of indices matches the tensor rank by appending trailing Nones. self_rank = len(self.shape) @@ -5073,11 +4986,6 @@ def is_advanced_index(index): # Note: In this function, the index is assumed to be either None or an int64 Tensor. return index is not None - def index_rank(index_position: int) -> int: - if index_position in bool_index_positions: - return 1 - return len(indices[index_position].shape) - advanced_indices: list[int] = [] none_indices: list[int] = [] num_advanced_indices = 0 @@ -5116,15 +5024,15 @@ def same_shape(other_shape: Optional[ir.Shape]) -> bool: all_same_shape = all(same_shape(indices[i].shape) for i in advanced_indices) if not all_same_shape: # Broadcast advanced indices to a common shape. - advanced_index_rank = max(index_rank(i) for i in advanced_indices) + advanced_index_rank = max(len(indices[i].shape) for i in advanced_indices) shapes = [] for i in advanced_indices: index = indices[i] - current_index_rank = index_rank(i) + index_rank = len(index.shape) index_shape = op.Shape(index) - if current_index_rank < advanced_index_rank: + if index_rank < advanced_index_rank: padding = op.Constant( - value_ints=[1 for _ in range(advanced_index_rank - current_index_rank)] + value_ints=[1 for _ in range(advanced_index_rank - index_rank)] ) index_shape = op.Concat(padding, index_shape, axis=0) shapes.append(index_shape) @@ -5135,10 +5043,10 @@ def same_shape(other_shape: Optional[ir.Shape]) -> bool: ] else: advanced_indices_shape = op.Shape(indices[advanced_indices[0]]) - advanced_index_rank = index_rank(advanced_indices[0]) + advanced_index_rank = len(indices[advanced_indices[0]].shape) else: advanced_indices_shape = op.Shape(indices[advanced_indices[0]]) - advanced_index_rank = index_rank(advanced_indices[0]) + advanced_index_rank = len(indices[advanced_indices[0]].shape) # ONNX ScatterND supports only the case where all advanced indices appear first, # followed by None indices. So, we need to transpose self and values so that the @@ -5212,6 +5120,34 @@ def _aten_index_put_bool( """index_put(Tensor self, Tensor?[] indices, Tensor values, bool accumulate=False) -> Tensor""" bool_mask = indices[0] + if len(indices) > 1: + if any(index is None for index in indices): + raise NotImplementedError( + "Boolean index_put with multiple indices does not support None indices." + ) + + advanced_indices = [] + selected_positions = [] + minus_one = op.Constant(value_ints=[-1]) + for index in indices: + if index.dtype != BOOL.dtype or len(index.shape) != 1: + raise NotImplementedError( + "Boolean index_put with multiple indices supports only 1-D boolean masks." + ) + positions = op.Reshape(op.Transpose(op.NonZero(index), perm=[1, 0]), minus_one) + selected_positions.append(positions) + advanced_indices.append(op.Unsqueeze(positions, minus_one)) + onnx_index = op.Concat(*advanced_indices, axis=-1) + target_shape = op.Concat( + op.Shape(selected_positions[0]), + op.Slice(op.Shape(self), starts=[len(indices)], ends=[len(self.shape)], axes=[0]), + axis=0, + ) + expanded_values = op.Expand(values, target_shape) + return op.ScatterND( + self, onnx_index, expanded_values, reduction="add" if accumulate else None + ) + if bool_mask is None or bool_mask.dtype != BOOL.dtype: raise NotImplementedError( "Boolean index_put expects a boolean mask as the first index." @@ -5868,9 +5804,8 @@ def aten_logit(self: TFloat, eps: Optional[float] = None) -> TFloat: one_minus_eps = ir.tensor(1 - eps, dtype=self.dtype) eps = ir.tensor(eps, dtype=self.dtype) - # Match torch.clamp behavior for eps > 0.5 by applying max then min. - z = op.Where(self < eps, eps, self) - z = op.Where(z <= one_minus_eps, z, one_minus_eps) + temporary_self = op.Where(self <= one_minus_eps, self, one_minus_eps) + z = op.Where(temporary_self < eps, eps, temporary_self) return op.Log(op.Div(z, op.Sub(one, z))) diff --git a/tests/function_libs/torch_lib/e2e_ops_tests.py b/tests/function_libs/torch_lib/e2e_ops_tests.py index 019164308f..e6b481f2c0 100644 --- a/tests/function_libs/torch_lib/e2e_ops_tests.py +++ b/tests/function_libs/torch_lib/e2e_ops_tests.py @@ -1232,43 +1232,6 @@ def forward(self, x, mask0, mask1, update): ) _testing.assert_onnx_program(onnx_program) - def test_index_put_bool_multi_mask_broadcast(self): - class Model(torch.nn.Module): - def forward(self, x, mask0, mask1, update): - return torch.ops.aten.index_put(x, [mask0, mask1], update) - - x = torch.zeros((3, 4), dtype=torch.float32) - mask0 = torch.tensor([False, True, False], dtype=torch.bool) - mask1 = torch.tensor([True, False, True, True], dtype=torch.bool) - update = torch.tensor([10.0, 20.0, 30.0], dtype=torch.float32) - onnx_program = torch.onnx.export( - Model(), - (x, mask0, mask1, update), - input_names=["x", "mask0", "mask1", "update"], - output_names=["output"], - opset_version=18, - dynamo=True, - ) - _testing.assert_onnx_program(onnx_program) - - def test_index_put_bool_mask_with_none_index(self): - class Model(torch.nn.Module): - def forward(self, x, mask, update): - return torch.ops.aten.index_put(x, [None, mask], update) - - x = torch.arange(6, dtype=torch.float32).reshape((2, 3)) - mask = torch.tensor([True, False, True], dtype=torch.bool) - update = torch.tensor(9.0, dtype=torch.float32) - onnx_program = torch.onnx.export( - Model(), - (x, mask, update), - input_names=["x", "mask", "update"], - output_names=["output"], - opset_version=18, - dynamo=True, - ) - _testing.assert_onnx_program(onnx_program) - def test_std_mean(self): """Test torch.std_mean which will be decomposed into prims.sum."""