From b90a036b6775452332978a3acfac3da954492e4d Mon Sep 17 00:00:00 2001 From: LeSingh1 Date: Thu, 27 Aug 2026 09:56:10 -0700 Subject: [PATCH 1/2] Only take the single-bool-mask path when the mask is the only index torch's index op took the boolean-mask shortcut whenever exactly one index was a bool tensor, ignoring whether other axes were indexed too. For x[mask, j] that drops j entirely and returns whole slices instead of the elements torch selects, silently and with a different shape: torch : [0.0, 5.0] coreml : [[0.0, 1.0, 2.0], [3.0, 4.0, 5.0]] Require the mask to be the only index before taking the shortcut, so anything else falls through to the general path that pairs the indices up. Also read the mask out of indices before computing its rank. The old code used the rank of whatever the enumerate loop happened to leave behind in 'index', which is the last index rather than the mask. --- .../converters/mil/frontend/torch/ops.py | 13 ++++++------ .../mil/frontend/torch/test/test_torch_ops.py | 21 +++++++++++++++++++ 2 files changed, 28 insertions(+), 6 deletions(-) diff --git a/coremltools/converters/mil/frontend/torch/ops.py b/coremltools/converters/mil/frontend/torch/ops.py index 1400e6fa8..b0c7e2342 100644 --- a/coremltools/converters/mil/frontend/torch/ops.py +++ b/coremltools/converters/mil/frontend/torch/ops.py @@ -5739,16 +5739,17 @@ def _try_single_bool_index(x: Var, indices: List[Var], name: str) -> bool: The true value indicates whether the element should be selected among the masked axes The output c is a tensor with shape (2, N), where N is the number of elements of b satisfying condition > 0.1 """ - boolean_indices_axis = [] - for i, index in enumerate(indices): - if index is not None and types.is_bool(index.dtype): - boolean_indices_axis.append(i) + non_none_indices_axis = [i for i, index in enumerate(indices) if index is not None] + boolean_indices_axis = [i for i in non_none_indices_axis if types.is_bool(indices[i].dtype)] - if len(boolean_indices_axis) == 1: + # This shortcut only holds when the mask is the sole index. With another index present, + # e.g. x[mask, j], torch pairs the mask's True positions up with j elementwise instead of + # selecting whole slices, which is what the general path below does. + if len(boolean_indices_axis) == 1 and len(non_none_indices_axis) == 1: # get the True element indices axis = boolean_indices_axis[0] - axes = list(range(axis, axis + index.rank)) index = indices[axis] + axes = list(range(axis, axis + index.rank)) index = mb.non_zero(x=index) # transpose the masked axes to the beginning diff --git a/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py b/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py index 62065c7de..c56ad958c 100644 --- a/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py +++ b/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py @@ -11372,6 +11372,27 @@ def forward(self, x, y): minimum_deployment_target=minimum_deployment_target, ) + @pytest.mark.parametrize( + "compute_unit, backend, frontend", + itertools.product(compute_units, backends, frontends), + ) + def test_index_bool_mask_with_another_index(self, compute_unit, backend, frontend): + """A bool mask paired with a second index selects elements, not whole slices.""" + + class IndexModel(torch.nn.Module): + def forward(self, x): + mask = torch.tensor([True, True]) + j = torch.tensor([0, 2]) + return x[mask, j] + + self.run_compare_torch( + [(2, 3)], + IndexModel(), + frontend=frontend, + backend=backend, + compute_unit=compute_unit, + ) + @pytest.mark.parametrize( "compute_unit, backend, frontend, input_dtype, shape, minimum_deployment_target", itertools.product( From 57f064c89b39ba0bc683035d7848c2a7a6063ddd Mon Sep 17 00:00:00 2001 From: LeSingh1 Date: Sat, 29 Aug 2026 22:28:28 -0700 Subject: [PATCH 2/2] xfail the new index test on the torch.export frontends torch.export cannot trace a bool mask combined with a second index: the mask makes the intermediate size data dependent and export raises PendingUnbackedSymbolNotFound before the converter runs. This matches the existing xfail on test_index_int_index_case_10, which cites the same cause. --- .../converters/mil/frontend/torch/test/test_torch_ops.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py b/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py index c56ad958c..733c35ebd 100644 --- a/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py +++ b/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py @@ -11378,6 +11378,12 @@ def forward(self, x, y): ) def test_index_bool_mask_with_another_index(self, compute_unit, backend, frontend): """A bool mask paired with a second index selects elements, not whole slices.""" + if frontend in TORCH_EXPORT_BASED_FRONTENDS: + pytest.xfail( + "torch.export cannot trace this model: the bool mask makes the intermediate " + "size data dependent, so export raises PendingUnbackedSymbolNotFound before " + "the converter is reached." + ) class IndexModel(torch.nn.Module): def forward(self, x):