From 16fcd4b12ac3e907260cd25bc6b06215b91b28d5 Mon Sep 17 00:00:00 2001 From: LeSingh1 Date: Sun, 9 Aug 2026 14:59:54 -0700 Subject: [PATCH] Tell apart torch overloads that TorchScript reports under one node kind TorchScript drops the overload suffix, so several aten ops reach the converter as one kind with two different argument layouts. Three of them assume the wrong layout. In each case the two overloads have different arity, which is what the fix keys on. sum / mean: aten::sum(self, dtype) has 2 inputs, aten::sum.dim_IntList( self, dim, keepdim, dtype) has 4. The converter always read position 1 as dim, so torch.sum(x, dtype=torch.float32) passed the dtype enum as an axis and failed with "IndexError: pop index out of range". randint: aten::randint(high, size, ...) has 6 inputs, aten::randint.low( low, high, size, ...) has 7. Under TorchScript the converter always took the .low layout, so torch.randint(10, (2, 3)) read the size list as high and the dtype enum as the shape, failing with "TypeError: 'int' object is not iterable". round: aten::round.decimals(self, decimals) has 2 inputs and was rejected outright by expected=1. Accept it and scale by 10**decimals, since dropping the argument would silently change the result. --- .../converters/mil/frontend/torch/ops.py | 46 +++++++++++-- .../mil/frontend/torch/test/test_torch_ops.py | 64 +++++++++++++++++++ 2 files changed, 103 insertions(+), 7 deletions(-) diff --git a/coremltools/converters/mil/frontend/torch/ops.py b/coremltools/converters/mil/frontend/torch/ops.py index 7f747baf4..65f2b432a 100644 --- a/coremltools/converters/mil/frontend/torch/ops.py +++ b/coremltools/converters/mil/frontend/torch/ops.py @@ -2138,8 +2138,17 @@ def _parse_positional_args(context, node) -> Tuple[Var]: if context.frontend == TorchFrontend.TORCHSCRIPT: x = inputs[0] - dim = inputs[1] if nargs > 1 else None - keepdim = inputs[2] if nargs > 2 else False + if nargs == 2: + # TorchScript reports one node kind for both + # aten::sum(self, *, dtype) -> 2 inputs + # aten::sum.dim_IntList(self, dim, keepdim, dtype) -> 4 inputs + # (and likewise for mean). Only the second has a dim, so at 2 inputs + # position 1 is the dtype, which reduces over the whole tensor. + dim = None + keepdim = False + else: + dim = inputs[1] if nargs > 1 else None + keepdim = inputs[2] if nargs > 2 else False else: if node.kind in ("mean", "sum", "all", "any"): x = inputs[0] @@ -5699,12 +5708,18 @@ def new_full(context, node): def randint(context, node): def _parse_positional_args(context, node) -> Tuple[Var]: inputs = _get_inputs(context, node, min_expected=2) - if context.frontend == TorchFrontend.TORCHSCRIPT or node.kind == "randint.low": + # TorchScript reports the node kind `randint` for both + # aten::randint(high, size, *, dtype, layout, device, pin_memory) -> 6 inputs + # aten::randint.low(low, high, size, *, dtype, ...) -> 7 inputs + # which differ only by the extra leading `low`, so tell them apart by arity. + has_low = node.kind == "randint.low" or ( + context.frontend == TorchFrontend.TORCHSCRIPT and len(inputs) >= 7 + ) + if has_low: low = mb.cast(x=inputs[0], dtype="fp32") high = mb.cast(x=inputs[1], dtype="fp32") shape = inputs[2].val else: - assert node.kind == "randint" low = 0.0 high = mb.cast(x=inputs[0], dtype="fp32") shape = inputs[1].val @@ -7490,10 +7505,27 @@ def log1p(context, node): context.add(mb.log(x=x, epsilon=1.0, name=node.name)) -@register_torch_op(torch_alias=["round"]) +@register_torch_op(torch_alias=["round", "round.decimals"]) def _round(context, node): - inputs = _get_inputs(context, node, expected=1) - context.add(mb.round(x=inputs[0], name=node.name)) + # TorchScript reports the node kind `round` for both + # aten::round(self) -> 1 input + # aten::round.decimals(self, decimals) -> 2 inputs + inputs = _get_inputs(context, node, min_expected=1) + x = inputs[0] + decimals = inputs[1] if len(inputs) > 1 else None + decimals = _get_kwinputs(context, node, "decimals", default=[decimals])[0] + if decimals is not None and isinstance(decimals, Var): + decimals = decimals.val + + if decimals is None or decimals == 0: + context.add(mb.round(x=x, name=node.name)) + return + + # torch rounds to `decimals` places, i.e. round(x * 10**d) / 10**d + scale = np.float32(10.0**decimals) + scaled = mb.mul(x=x, y=scale) + rounded = mb.round(x=scaled) + context.add(mb.real_div(x=rounded, y=scale, name=node.name)) @register_torch_op 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 58860798b..55b4c29ea 100644 --- a/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py +++ b/coremltools/converters/mil/frontend/torch/test/test_torch_ops.py @@ -4836,6 +4836,24 @@ def forward(self, x): frontend=frontend, ) + @pytest.mark.parametrize("frontend", frontends) + def test_randint_no_low(self, frontend): + """ + aten::randint(high, size, ...) and aten::randint.low(low, high, size, ...) + share the node kind `randint` under TorchScript and differ only by the extra + leading `low`, so the single argument form must not be read as `low`. + """ + + class TestModel(nn.Module): + def forward(self, x): + return torch.randint(10, (2, 3)) + + model = TestModel().eval() + x = torch.randn((1, 3)) + torch_model = export_torch_model_to_frontend(model, (x,), frontend) + inputs = [ct.TensorType(shape=x.shape)] if frontend == TorchFrontend.TORCHSCRIPT else None + ct.convert(torch_model, inputs=inputs) + @pytest.mark.parametrize("frontend", frontends) def test_tuple_input(self, frontend): if frontend == TorchFrontend.EXECUTORCH: @@ -7028,6 +7046,29 @@ def test_elementwise_no_params(self, compute_unit, backend, frontend, shape, op_ shape, model, compute_unit=compute_unit, backend=backend, frontend=frontend ) + @pytest.mark.parametrize( + "compute_unit, backend, frontend, decimals", + itertools.product(compute_units, backends, frontends, [0, 1, 2, -1]), + ) + def test_round_decimals(self, compute_unit, backend, frontend, decimals): + """ + aten::round(self) and aten::round.decimals(self, decimals) share the node kind + `round` under TorchScript. + """ + + class Model(nn.Module): + def forward(self, x): + return torch.round(x, decimals=decimals) + + self.run_compare_torch( + (1, 3, 5, 8), + Model(), + compute_unit=compute_unit, + backend=backend, + frontend=frontend, + rand_range=(-20.0, 20.0), + ) + @pytest.mark.parametrize( "compute_unit, backend, frontend, shape", itertools.product( @@ -12569,6 +12610,29 @@ def test_sum(self, compute_unit, backend, frontend, input_dtype): compute_unit=compute_unit, ) + @pytest.mark.parametrize( + "compute_unit, backend, frontend, op", + itertools.product(compute_units, backends, frontends, [torch.sum, torch.mean]), + ) + def test_sum_mean_dtype(self, compute_unit, backend, frontend, op): + """ + aten::sum(self, dtype) and aten::sum.dim_IntList(self, dim, keepdim, dtype) + share the node kind `sum` under TorchScript, so the dtype must not be read + as a dim. + """ + + class Model(nn.Module): + def forward(self, x): + return op(x, dtype=torch.float32) + + self.run_compare_torch( + (5, 4), + Model(), + frontend=frontend, + backend=backend, + compute_unit=compute_unit, + ) + @pytest.mark.parametrize( "compute_unit, backend, frontend, shape, dim", itertools.product(